refactor: establish feature application service layout
This commit is contained in:
+2
-94
@@ -1,95 +1,3 @@
|
|||||||
from __future__ import annotations
|
from backend.features.alerts.service import AlertService
|
||||||
|
|
||||||
import secrets
|
__all__ = ["AlertService"]
|
||||||
from datetime import date, datetime
|
|
||||||
from typing import Any
|
|
||||||
|
|
||||||
from app_config import validate_text
|
|
||||||
from backend.database.repositories import AlertRepository
|
|
||||||
|
|
||||||
|
|
||||||
class AlertService:
|
|
||||||
def __init__(self, repository: AlertRepository) -> None:
|
|
||||||
self.repository = repository
|
|
||||||
|
|
||||||
def create_manual(self, user_id: int, payload: dict[str, Any]) -> int:
|
|
||||||
title = validate_text(payload.get("title"), "提醒标题", 80, required=True)
|
|
||||||
content = validate_text(payload.get("content"), "提醒内容", 500)
|
|
||||||
code = validate_text(payload.get("code"), "股票代码", 12)
|
|
||||||
available_date = self.calendar_date(
|
|
||||||
str(payload.get("remind_date") or date.today().isoformat())
|
|
||||||
)
|
|
||||||
return self.repository.save_alert(
|
|
||||||
user_id=user_id,
|
|
||||||
kind="manual",
|
|
||||||
title=title,
|
|
||||||
content=content,
|
|
||||||
available_date=available_date,
|
|
||||||
code=code,
|
|
||||||
dedupe_key=f"manual:{secrets.token_hex(12)}",
|
|
||||||
)
|
|
||||||
|
|
||||||
def sync_strategy_tracking(self, user_id: int, tracking: dict[str, Any]) -> int:
|
|
||||||
synced = 0
|
|
||||||
today = date.today().strftime("%Y%m%d")
|
|
||||||
for batch in tracking.get("batches") or []:
|
|
||||||
items = batch.get("items") or []
|
|
||||||
summary = batch.get("summary") or {}
|
|
||||||
if not items:
|
|
||||||
continue
|
|
||||||
run_id = int(batch.get("run_id") or 0)
|
|
||||||
strategy_name = str(batch.get("strategy_name") or "选股策略")
|
|
||||||
observed = int(summary.get("observed") or 0)
|
|
||||||
completed = int(summary.get("completed") or 0)
|
|
||||||
if observed:
|
|
||||||
win_rate = summary.get("t1_win_rate")
|
|
||||||
suffix = f",当前红盘率 {win_rate:.1f}%" if win_rate is not None else ""
|
|
||||||
self.repository.save_alert(
|
|
||||||
user_id, "strategy_t1", f"{strategy_name} 已有 T+1 反馈",
|
|
||||||
f"{observed}/{len(items)} 只标的已有首日表现{suffix}。",
|
|
||||||
today, "", f"strategy:{run_id}:t1",
|
|
||||||
)
|
|
||||||
synced += 1
|
|
||||||
if completed == len(items):
|
|
||||||
average = summary.get("average_t5")
|
|
||||||
suffix = f",平均收益 {average:+.2f}%" if average is not None else ""
|
|
||||||
self.repository.save_alert(
|
|
||||||
user_id, "strategy_t5", f"{strategy_name} 五日跟踪完成",
|
|
||||||
f"本批 {len(items)} 只标的已完成 T+5 跟踪{suffix}。",
|
|
||||||
today, "", f"strategy:{run_id}:t5",
|
|
||||||
)
|
|
||||||
synced += 1
|
|
||||||
return synced
|
|
||||||
|
|
||||||
def list_alerts(
|
|
||||||
self, user_id: int, status: str = "all", as_of: str = ""
|
|
||||||
) -> dict[str, Any]:
|
|
||||||
if status not in {"all", "unread"}:
|
|
||||||
raise ValueError("提醒筛选不支持。")
|
|
||||||
compact_date = self.calendar_date(as_of or date.today().isoformat())
|
|
||||||
items = self.repository.list_alerts(user_id, compact_date, status == "unread")
|
|
||||||
for item in items:
|
|
||||||
item["due"] = str(item.get("available_date") or "") <= compact_date
|
|
||||||
return {
|
|
||||||
"items": items,
|
|
||||||
"unread_count": self.repository.count_unread_alerts(user_id, compact_date),
|
|
||||||
"as_of": compact_date,
|
|
||||||
}
|
|
||||||
|
|
||||||
def mark_read(self, user_id: int, alert_id: int) -> bool:
|
|
||||||
return self.repository.mark_alert_read(user_id, alert_id)
|
|
||||||
|
|
||||||
def mark_all_read(self, user_id: int, as_of: str) -> int:
|
|
||||||
return self.repository.mark_all_alerts_read(user_id, as_of)
|
|
||||||
|
|
||||||
def delete(self, user_id: int, alert_id: int) -> bool:
|
|
||||||
return self.repository.delete_alert(user_id, alert_id)
|
|
||||||
|
|
||||||
@staticmethod
|
|
||||||
def calendar_date(value: str) -> str:
|
|
||||||
compact = value.replace("-", "").strip()
|
|
||||||
try:
|
|
||||||
parsed = datetime.strptime(compact, "%Y%m%d")
|
|
||||||
except ValueError as exc:
|
|
||||||
raise ValueError("提醒日期格式应为 YYYY-MM-DD。") from exc
|
|
||||||
return parsed.strftime("%Y%m%d")
|
|
||||||
|
|||||||
@@ -4,17 +4,17 @@ from dataclasses import dataclass
|
|||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from collections.abc import Callable
|
from collections.abc import Callable
|
||||||
|
|
||||||
from alert_service import AlertService
|
|
||||||
from backend.data import DataGateway, build_data_gateway
|
from backend.data import DataGateway, build_data_gateway
|
||||||
from backend.database.repositories import RepositoryBundle, build_repository_bundle
|
from backend.database.repositories import RepositoryBundle, build_repository_bundle
|
||||||
|
from backend.features.alerts import AlertService
|
||||||
|
from backend.features.review import TradeJournalService
|
||||||
|
from backend.features.screener import StrategyTrackingService
|
||||||
from chart_data_provider import MarketChartClient
|
from chart_data_provider import MarketChartClient
|
||||||
from database import ReviewDatabase
|
from database import ReviewDatabase
|
||||||
from ifind_client import IfindHttpClient
|
from ifind_client import IfindHttpClient
|
||||||
from mentor_agent import MentorSkillRegistry
|
from mentor_agent import MentorSkillRegistry
|
||||||
from realtime_aggregator import WebRealtimeAggregator
|
from realtime_aggregator import WebRealtimeAggregator
|
||||||
from screener import ScreenerEngine
|
from screener import ScreenerEngine
|
||||||
from strategy_tracking import StrategyTrackingService
|
|
||||||
from trade_journal import TradeJournalService
|
|
||||||
|
|
||||||
|
|
||||||
@dataclass(frozen=True)
|
@dataclass(frozen=True)
|
||||||
|
|||||||
@@ -0,0 +1 @@
|
|||||||
|
"""Feature-owned application services."""
|
||||||
@@ -0,0 +1,3 @@
|
|||||||
|
from .service import AlertService
|
||||||
|
|
||||||
|
__all__ = ["AlertService"]
|
||||||
@@ -0,0 +1,95 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import secrets
|
||||||
|
from datetime import date, datetime
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
from app_config import validate_text
|
||||||
|
from backend.database.repositories import AlertRepository
|
||||||
|
|
||||||
|
|
||||||
|
class AlertService:
|
||||||
|
def __init__(self, repository: AlertRepository) -> None:
|
||||||
|
self.repository = repository
|
||||||
|
|
||||||
|
def create_manual(self, user_id: int, payload: dict[str, Any]) -> int:
|
||||||
|
title = validate_text(payload.get("title"), "提醒标题", 80, required=True)
|
||||||
|
content = validate_text(payload.get("content"), "提醒内容", 500)
|
||||||
|
code = validate_text(payload.get("code"), "股票代码", 12)
|
||||||
|
available_date = self.calendar_date(
|
||||||
|
str(payload.get("remind_date") or date.today().isoformat())
|
||||||
|
)
|
||||||
|
return self.repository.save_alert(
|
||||||
|
user_id=user_id,
|
||||||
|
kind="manual",
|
||||||
|
title=title,
|
||||||
|
content=content,
|
||||||
|
available_date=available_date,
|
||||||
|
code=code,
|
||||||
|
dedupe_key=f"manual:{secrets.token_hex(12)}",
|
||||||
|
)
|
||||||
|
|
||||||
|
def sync_strategy_tracking(self, user_id: int, tracking: dict[str, Any]) -> int:
|
||||||
|
synced = 0
|
||||||
|
today = date.today().strftime("%Y%m%d")
|
||||||
|
for batch in tracking.get("batches") or []:
|
||||||
|
items = batch.get("items") or []
|
||||||
|
summary = batch.get("summary") or {}
|
||||||
|
if not items:
|
||||||
|
continue
|
||||||
|
run_id = int(batch.get("run_id") or 0)
|
||||||
|
strategy_name = str(batch.get("strategy_name") or "选股策略")
|
||||||
|
observed = int(summary.get("observed") or 0)
|
||||||
|
completed = int(summary.get("completed") or 0)
|
||||||
|
if observed:
|
||||||
|
win_rate = summary.get("t1_win_rate")
|
||||||
|
suffix = f",当前红盘率 {win_rate:.1f}%" if win_rate is not None else ""
|
||||||
|
self.repository.save_alert(
|
||||||
|
user_id, "strategy_t1", f"{strategy_name} 已有 T+1 反馈",
|
||||||
|
f"{observed}/{len(items)} 只标的已有首日表现{suffix}。",
|
||||||
|
today, "", f"strategy:{run_id}:t1",
|
||||||
|
)
|
||||||
|
synced += 1
|
||||||
|
if completed == len(items):
|
||||||
|
average = summary.get("average_t5")
|
||||||
|
suffix = f",平均收益 {average:+.2f}%" if average is not None else ""
|
||||||
|
self.repository.save_alert(
|
||||||
|
user_id, "strategy_t5", f"{strategy_name} 五日跟踪完成",
|
||||||
|
f"本批 {len(items)} 只标的已完成 T+5 跟踪{suffix}。",
|
||||||
|
today, "", f"strategy:{run_id}:t5",
|
||||||
|
)
|
||||||
|
synced += 1
|
||||||
|
return synced
|
||||||
|
|
||||||
|
def list_alerts(
|
||||||
|
self, user_id: int, status: str = "all", as_of: str = ""
|
||||||
|
) -> dict[str, Any]:
|
||||||
|
if status not in {"all", "unread"}:
|
||||||
|
raise ValueError("提醒筛选不支持。")
|
||||||
|
compact_date = self.calendar_date(as_of or date.today().isoformat())
|
||||||
|
items = self.repository.list_alerts(user_id, compact_date, status == "unread")
|
||||||
|
for item in items:
|
||||||
|
item["due"] = str(item.get("available_date") or "") <= compact_date
|
||||||
|
return {
|
||||||
|
"items": items,
|
||||||
|
"unread_count": self.repository.count_unread_alerts(user_id, compact_date),
|
||||||
|
"as_of": compact_date,
|
||||||
|
}
|
||||||
|
|
||||||
|
def mark_read(self, user_id: int, alert_id: int) -> bool:
|
||||||
|
return self.repository.mark_alert_read(user_id, alert_id)
|
||||||
|
|
||||||
|
def mark_all_read(self, user_id: int, as_of: str) -> int:
|
||||||
|
return self.repository.mark_all_alerts_read(user_id, as_of)
|
||||||
|
|
||||||
|
def delete(self, user_id: int, alert_id: int) -> bool:
|
||||||
|
return self.repository.delete_alert(user_id, alert_id)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def calendar_date(value: str) -> str:
|
||||||
|
compact = value.replace("-", "").strip()
|
||||||
|
try:
|
||||||
|
parsed = datetime.strptime(compact, "%Y%m%d")
|
||||||
|
except ValueError as exc:
|
||||||
|
raise ValueError("提醒日期格式应为 YYYY-MM-DD。") from exc
|
||||||
|
return parsed.strftime("%Y%m%d")
|
||||||
@@ -0,0 +1,3 @@
|
|||||||
|
from .trade_journal import EMOTIONS, TRADE_ACTIONS, TradeJournalService
|
||||||
|
|
||||||
|
__all__ = ["EMOTIONS", "TRADE_ACTIONS", "TradeJournalService"]
|
||||||
@@ -0,0 +1,100 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import json
|
||||||
|
from datetime import date
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
from app_config import normalize_date, validate_stock_code, validate_text
|
||||||
|
from backend.database.repositories import TradeJournalRepository
|
||||||
|
|
||||||
|
|
||||||
|
TRADE_ACTIONS = {"buy": "买入", "sell": "卖出", "trim": "减仓", "add": "加仓", "watch": "观察"}
|
||||||
|
EMOTIONS = {"calm": "平静", "confident": "笃定", "hesitant": "犹豫", "anxious": "焦虑", "impulsive": "冲动"}
|
||||||
|
|
||||||
|
|
||||||
|
class TradeJournalService:
|
||||||
|
def __init__(self, repository: TradeJournalRepository) -> None:
|
||||||
|
self.repository = repository
|
||||||
|
|
||||||
|
def save(self, user_id: int, payload: dict[str, Any]) -> int:
|
||||||
|
trade_id = int(payload.get("id") or 0)
|
||||||
|
trade_date = normalize_date(str(payload.get("trade_date") or date.today().isoformat()))
|
||||||
|
code = validate_stock_code(str(payload.get("code") or ""))
|
||||||
|
name = validate_text(payload.get("name"), "股票名称", 40, required=True)
|
||||||
|
action = str(payload.get("action") or "")
|
||||||
|
if action not in TRADE_ACTIONS:
|
||||||
|
raise ValueError("交易动作不支持。")
|
||||||
|
emotion = str(payload.get("emotion") or "calm")
|
||||||
|
if emotion not in EMOTIONS:
|
||||||
|
raise ValueError("交易情绪不支持。")
|
||||||
|
price = self._number(payload.get("price"), "成交价格", 0, 1000000, required=True)
|
||||||
|
quantity = int(self._number(payload.get("quantity"), "成交数量", 0, 100000000))
|
||||||
|
position_pct = self._number(payload.get("position_pct"), "仓位", 0, 100)
|
||||||
|
pnl_amount = self._optional_number(payload.get("pnl_amount"), "盈亏金额", -1e12, 1e12)
|
||||||
|
pnl_pct = self._optional_number(payload.get("pnl_pct"), "盈亏比例", -1000, 10000)
|
||||||
|
thesis = validate_text(payload.get("thesis"), "交易逻辑", 2000)
|
||||||
|
execution = validate_text(payload.get("execution"), "执行复核", 2000)
|
||||||
|
raw_tags = payload.get("tags") or []
|
||||||
|
if isinstance(raw_tags, str):
|
||||||
|
raw_tags = [item.strip() for item in raw_tags.replace(",", ",").split(",")]
|
||||||
|
if not isinstance(raw_tags, list):
|
||||||
|
raise ValueError("交易标签格式不正确。")
|
||||||
|
tags = [validate_text(item, "交易标签", 20) for item in raw_tags if str(item).strip()][:8]
|
||||||
|
return self.repository.save_trade_entry(
|
||||||
|
user_id, trade_date, code, name, action, price, quantity, position_pct,
|
||||||
|
pnl_amount, pnl_pct, thesis, execution, emotion, tags, trade_id or None,
|
||||||
|
)
|
||||||
|
|
||||||
|
def list_entries(
|
||||||
|
self, user_id: int, start_date: str = "", end_date: str = "", code: str = ""
|
||||||
|
) -> dict[str, Any]:
|
||||||
|
start = normalize_date(start_date) if start_date else ""
|
||||||
|
end = normalize_date(end_date) if end_date else date.today().strftime("%Y%m%d")
|
||||||
|
if start and start > end:
|
||||||
|
raise ValueError("开始日期不能晚于结束日期。")
|
||||||
|
code = validate_stock_code(code) if code else ""
|
||||||
|
items = self.repository.list_trade_entries(user_id, start, end, code)
|
||||||
|
for item in items:
|
||||||
|
item["tags"] = json.loads(item.get("tags") or "[]")
|
||||||
|
item["action_label"] = TRADE_ACTIONS.get(item["action"], item["action"])
|
||||||
|
item["emotion_label"] = EMOTIONS.get(item["emotion"], item["emotion"])
|
||||||
|
realized = [item for item in items if item.get("pnl_pct") is not None]
|
||||||
|
return {"items": items, "summary": self._summary(items, realized)}
|
||||||
|
|
||||||
|
def delete(self, user_id: int, trade_id: int) -> bool:
|
||||||
|
return self.repository.delete_trade_entry(user_id, trade_id)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _summary(items: list[dict[str, Any]], realized: list[dict[str, Any]]) -> dict[str, Any]:
|
||||||
|
pnl_amounts = [float(item["pnl_amount"]) for item in realized if item.get("pnl_amount") is not None]
|
||||||
|
positions = [float(item["position_pct"]) for item in items if float(item.get("position_pct") or 0) > 0]
|
||||||
|
wins = sum(float(item.get("pnl_pct") or 0) > 0 for item in realized)
|
||||||
|
return {
|
||||||
|
"total": len(items),
|
||||||
|
"realized": len(realized),
|
||||||
|
"win_rate": round(wins / len(realized) * 100, 1) if realized else None,
|
||||||
|
"pnl_amount": round(sum(pnl_amounts), 2) if pnl_amounts else None,
|
||||||
|
"average_position": round(sum(positions) / len(positions), 1) if positions else None,
|
||||||
|
}
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _number(value: Any, label: str, minimum: float, maximum: float, required: bool = False) -> float:
|
||||||
|
if value in (None, ""):
|
||||||
|
if required:
|
||||||
|
raise ValueError(f"{label}不能为空。")
|
||||||
|
return 0.0
|
||||||
|
try:
|
||||||
|
parsed = float(value)
|
||||||
|
except (TypeError, ValueError) as exc:
|
||||||
|
raise ValueError(f"{label}格式不正确。") from exc
|
||||||
|
if parsed < minimum or parsed > maximum:
|
||||||
|
raise ValueError(f"{label}超出允许范围。")
|
||||||
|
return parsed
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def _optional_number(
|
||||||
|
cls, value: Any, label: str, minimum: float, maximum: float
|
||||||
|
) -> float | None:
|
||||||
|
if value in (None, ""):
|
||||||
|
return None
|
||||||
|
return cls._number(value, label, minimum, maximum, required=True)
|
||||||
@@ -0,0 +1,3 @@
|
|||||||
|
from .tracking import StrategyTrackingService
|
||||||
|
|
||||||
|
__all__ = ["StrategyTrackingService"]
|
||||||
@@ -0,0 +1,134 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
from backend.database.repositories import StrategyTrackingRepository
|
||||||
|
|
||||||
|
|
||||||
|
class StrategyTrackingService:
|
||||||
|
def __init__(self, repository: StrategyTrackingRepository) -> None:
|
||||||
|
self.repository = repository
|
||||||
|
|
||||||
|
def record_run(
|
||||||
|
self,
|
||||||
|
user_id: int,
|
||||||
|
run_id: int,
|
||||||
|
selection_date: str,
|
||||||
|
strategy_name: str,
|
||||||
|
candidates: list[dict[str, Any]],
|
||||||
|
) -> int:
|
||||||
|
return self.repository.save_strategy_tracks(
|
||||||
|
user_id, run_id, selection_date, strategy_name, candidates
|
||||||
|
)
|
||||||
|
|
||||||
|
def add_candidate(self, user_id: int, run_id: int, code: str) -> dict[str, Any]:
|
||||||
|
run = self.repository.get_screener_run(user_id, run_id)
|
||||||
|
if not run:
|
||||||
|
run = self.repository.get_screener_run(0, run_id)
|
||||||
|
if not run:
|
||||||
|
raise ValueError("选股结果不存在或不属于当前账号。")
|
||||||
|
normalized_code = str(code or "").strip().split(".")[0]
|
||||||
|
candidate = next(
|
||||||
|
(
|
||||||
|
item for item in run.get("candidates", [])
|
||||||
|
if str(item.get("code") or item.get("ts_code") or "").split(".")[0]
|
||||||
|
== normalized_code
|
||||||
|
),
|
||||||
|
None,
|
||||||
|
)
|
||||||
|
if not candidate:
|
||||||
|
raise ValueError("该股票不在本次选股结果中。")
|
||||||
|
added = self.record_run(
|
||||||
|
user_id,
|
||||||
|
run_id,
|
||||||
|
str(run.get("meta", {}).get("trade_date") or ""),
|
||||||
|
str(run.get("strategy_name") or "未命名策略"),
|
||||||
|
[candidate],
|
||||||
|
)
|
||||||
|
return {"added": added, "tracking": self.list_tracking(user_id)}
|
||||||
|
|
||||||
|
def remove_candidate(self, user_id: int, track_id: int) -> dict[str, Any]:
|
||||||
|
deleted = self.repository.delete_strategy_track(user_id, track_id)
|
||||||
|
return {"deleted": deleted, "tracking": self.list_tracking(user_id)}
|
||||||
|
|
||||||
|
def list_tracking(self, user_id: int, limit_batches: int = 12) -> dict[str, Any]:
|
||||||
|
tracks = self.repository.list_strategy_tracks(user_id, limit_batches)
|
||||||
|
if not tracks:
|
||||||
|
return {"batches": [], "summary": self._summary([])}
|
||||||
|
|
||||||
|
bars = self.repository.load_tracking_bars(
|
||||||
|
[(item["ts_code"], item["selection_date"]) for item in tracks], 5
|
||||||
|
)
|
||||||
|
batches: dict[int, dict[str, Any]] = {}
|
||||||
|
all_items: list[dict[str, Any]] = []
|
||||||
|
for track in tracks:
|
||||||
|
key = (track["ts_code"], track["selection_date"])
|
||||||
|
metrics = self.calculate_metrics(float(track["entry_price"]), bars.get(key, []))
|
||||||
|
item = {
|
||||||
|
"id": track["id"],
|
||||||
|
"code": track["code"],
|
||||||
|
"name": track["name"],
|
||||||
|
"sector": track["sector"],
|
||||||
|
"entry_price": round(float(track["entry_price"]), 2),
|
||||||
|
**metrics,
|
||||||
|
}
|
||||||
|
all_items.append(item)
|
||||||
|
batch = batches.setdefault(
|
||||||
|
int(track["run_id"]),
|
||||||
|
{
|
||||||
|
"run_id": int(track["run_id"]),
|
||||||
|
"selection_date": track["selection_date"],
|
||||||
|
"strategy_name": track["strategy_name"],
|
||||||
|
"items": [],
|
||||||
|
},
|
||||||
|
)
|
||||||
|
batch["items"].append(item)
|
||||||
|
|
||||||
|
ordered = list(batches.values())
|
||||||
|
for batch in ordered:
|
||||||
|
batch["summary"] = self._summary(batch["items"])
|
||||||
|
return {"batches": ordered, "summary": self._summary(all_items)}
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def calculate_metrics(entry_price: float, bars: list[dict[str, Any]]) -> dict[str, Any]:
|
||||||
|
valid = [row for row in bars[:5] if float(row.get("close") or 0) > 0]
|
||||||
|
if entry_price <= 0 or not valid:
|
||||||
|
return {
|
||||||
|
"observed_days": 0,
|
||||||
|
"status": "等待 T+1",
|
||||||
|
"t1_open": None,
|
||||||
|
"t1_close": None,
|
||||||
|
"t3_close": None,
|
||||||
|
"t5_close": None,
|
||||||
|
"max_gain": None,
|
||||||
|
"max_drawdown": None,
|
||||||
|
}
|
||||||
|
|
||||||
|
def change(price: Any) -> float:
|
||||||
|
return round((float(price or 0) / entry_price - 1) * 100, 2)
|
||||||
|
|
||||||
|
observed = len(valid)
|
||||||
|
return {
|
||||||
|
"observed_days": observed,
|
||||||
|
"status": "已完成" if observed >= 5 else f"跟踪中 {observed}/5",
|
||||||
|
"t1_open": change(valid[0]["open"]),
|
||||||
|
"t1_close": change(valid[0]["close"]),
|
||||||
|
"t3_close": change(valid[2]["close"]) if observed >= 3 else None,
|
||||||
|
"t5_close": change(valid[4]["close"]) if observed >= 5 else None,
|
||||||
|
"max_gain": max(change(row["high"]) for row in valid),
|
||||||
|
"max_drawdown": min(change(row["low"]) for row in valid),
|
||||||
|
}
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _summary(items: list[dict[str, Any]]) -> dict[str, Any]:
|
||||||
|
completed = [item for item in items if item.get("t5_close") is not None]
|
||||||
|
t1 = [float(item["t1_close"]) for item in items if item.get("t1_close") is not None]
|
||||||
|
t5 = [float(item["t5_close"]) for item in completed]
|
||||||
|
return {
|
||||||
|
"total": len(items),
|
||||||
|
"observed": len(t1),
|
||||||
|
"completed": len(completed),
|
||||||
|
"t1_win_rate": round(sum(value > 0 for value in t1) / len(t1) * 100, 1) if t1 else None,
|
||||||
|
"t5_win_rate": round(sum(value > 0 for value in t5) / len(t5) * 100, 1) if t5 else None,
|
||||||
|
"average_t5": round(sum(t5) / len(t5), 2) if t5 else None,
|
||||||
|
}
|
||||||
@@ -0,0 +1,28 @@
|
|||||||
|
# Stage 10: Feature Application Service Layout
|
||||||
|
|
||||||
|
Date: 2026-07-29
|
||||||
|
|
||||||
|
## Result
|
||||||
|
|
||||||
|
- Created the governed `backend/features/<feature>` source boundary.
|
||||||
|
- Moved alert behavior into the alerts feature.
|
||||||
|
- Moved trade-journal behavior into the review feature.
|
||||||
|
- Moved strategy-tracking behavior into the screener feature.
|
||||||
|
- Changed the application container to import feature-owned services directly.
|
||||||
|
- Reduced the former top-level service modules to compatibility exports, preserving existing
|
||||||
|
imports without maintaining duplicate implementations.
|
||||||
|
- Added dependency tests that prevent feature services from importing HTTP delivery code or
|
||||||
|
concrete market-provider adapters.
|
||||||
|
|
||||||
|
## Template
|
||||||
|
|
||||||
|
Each migrated feature owns its application behavior and depends on repository or gateway
|
||||||
|
ports. HTTP delivery and background jobs may invoke the service, but neither may become the
|
||||||
|
owner of its rules. The same layout is now available for the remaining market, account,
|
||||||
|
Mentor, Wentian, and screener services.
|
||||||
|
|
||||||
|
## Compatibility
|
||||||
|
|
||||||
|
No route, response shape, persisted data, service method, or top-level import path changes in
|
||||||
|
this stage. Compatibility exports are removed only after all internal and external callers
|
||||||
|
use the feature packages.
|
||||||
+2
-133
@@ -1,134 +1,3 @@
|
|||||||
from __future__ import annotations
|
from backend.features.screener.tracking import StrategyTrackingService
|
||||||
|
|
||||||
from typing import Any
|
__all__ = ["StrategyTrackingService"]
|
||||||
|
|
||||||
from backend.database.repositories import StrategyTrackingRepository
|
|
||||||
|
|
||||||
|
|
||||||
class StrategyTrackingService:
|
|
||||||
def __init__(self, repository: StrategyTrackingRepository) -> None:
|
|
||||||
self.repository = repository
|
|
||||||
|
|
||||||
def record_run(
|
|
||||||
self,
|
|
||||||
user_id: int,
|
|
||||||
run_id: int,
|
|
||||||
selection_date: str,
|
|
||||||
strategy_name: str,
|
|
||||||
candidates: list[dict[str, Any]],
|
|
||||||
) -> int:
|
|
||||||
return self.repository.save_strategy_tracks(
|
|
||||||
user_id, run_id, selection_date, strategy_name, candidates
|
|
||||||
)
|
|
||||||
|
|
||||||
def add_candidate(self, user_id: int, run_id: int, code: str) -> dict[str, Any]:
|
|
||||||
run = self.repository.get_screener_run(user_id, run_id)
|
|
||||||
if not run:
|
|
||||||
run = self.repository.get_screener_run(0, run_id)
|
|
||||||
if not run:
|
|
||||||
raise ValueError("选股结果不存在或不属于当前账号。")
|
|
||||||
normalized_code = str(code or "").strip().split(".")[0]
|
|
||||||
candidate = next(
|
|
||||||
(
|
|
||||||
item for item in run.get("candidates", [])
|
|
||||||
if str(item.get("code") or item.get("ts_code") or "").split(".")[0]
|
|
||||||
== normalized_code
|
|
||||||
),
|
|
||||||
None,
|
|
||||||
)
|
|
||||||
if not candidate:
|
|
||||||
raise ValueError("该股票不在本次选股结果中。")
|
|
||||||
added = self.record_run(
|
|
||||||
user_id,
|
|
||||||
run_id,
|
|
||||||
str(run.get("meta", {}).get("trade_date") or ""),
|
|
||||||
str(run.get("strategy_name") or "未命名策略"),
|
|
||||||
[candidate],
|
|
||||||
)
|
|
||||||
return {"added": added, "tracking": self.list_tracking(user_id)}
|
|
||||||
|
|
||||||
def remove_candidate(self, user_id: int, track_id: int) -> dict[str, Any]:
|
|
||||||
deleted = self.repository.delete_strategy_track(user_id, track_id)
|
|
||||||
return {"deleted": deleted, "tracking": self.list_tracking(user_id)}
|
|
||||||
|
|
||||||
def list_tracking(self, user_id: int, limit_batches: int = 12) -> dict[str, Any]:
|
|
||||||
tracks = self.repository.list_strategy_tracks(user_id, limit_batches)
|
|
||||||
if not tracks:
|
|
||||||
return {"batches": [], "summary": self._summary([])}
|
|
||||||
|
|
||||||
bars = self.repository.load_tracking_bars(
|
|
||||||
[(item["ts_code"], item["selection_date"]) for item in tracks], 5
|
|
||||||
)
|
|
||||||
batches: dict[int, dict[str, Any]] = {}
|
|
||||||
all_items: list[dict[str, Any]] = []
|
|
||||||
for track in tracks:
|
|
||||||
key = (track["ts_code"], track["selection_date"])
|
|
||||||
metrics = self.calculate_metrics(float(track["entry_price"]), bars.get(key, []))
|
|
||||||
item = {
|
|
||||||
"id": track["id"],
|
|
||||||
"code": track["code"],
|
|
||||||
"name": track["name"],
|
|
||||||
"sector": track["sector"],
|
|
||||||
"entry_price": round(float(track["entry_price"]), 2),
|
|
||||||
**metrics,
|
|
||||||
}
|
|
||||||
all_items.append(item)
|
|
||||||
batch = batches.setdefault(
|
|
||||||
int(track["run_id"]),
|
|
||||||
{
|
|
||||||
"run_id": int(track["run_id"]),
|
|
||||||
"selection_date": track["selection_date"],
|
|
||||||
"strategy_name": track["strategy_name"],
|
|
||||||
"items": [],
|
|
||||||
},
|
|
||||||
)
|
|
||||||
batch["items"].append(item)
|
|
||||||
|
|
||||||
ordered = list(batches.values())
|
|
||||||
for batch in ordered:
|
|
||||||
batch["summary"] = self._summary(batch["items"])
|
|
||||||
return {"batches": ordered, "summary": self._summary(all_items)}
|
|
||||||
|
|
||||||
@staticmethod
|
|
||||||
def calculate_metrics(entry_price: float, bars: list[dict[str, Any]]) -> dict[str, Any]:
|
|
||||||
valid = [row for row in bars[:5] if float(row.get("close") or 0) > 0]
|
|
||||||
if entry_price <= 0 or not valid:
|
|
||||||
return {
|
|
||||||
"observed_days": 0,
|
|
||||||
"status": "等待 T+1",
|
|
||||||
"t1_open": None,
|
|
||||||
"t1_close": None,
|
|
||||||
"t3_close": None,
|
|
||||||
"t5_close": None,
|
|
||||||
"max_gain": None,
|
|
||||||
"max_drawdown": None,
|
|
||||||
}
|
|
||||||
|
|
||||||
def change(price: Any) -> float:
|
|
||||||
return round((float(price or 0) / entry_price - 1) * 100, 2)
|
|
||||||
|
|
||||||
observed = len(valid)
|
|
||||||
return {
|
|
||||||
"observed_days": observed,
|
|
||||||
"status": "已完成" if observed >= 5 else f"跟踪中 {observed}/5",
|
|
||||||
"t1_open": change(valid[0]["open"]),
|
|
||||||
"t1_close": change(valid[0]["close"]),
|
|
||||||
"t3_close": change(valid[2]["close"]) if observed >= 3 else None,
|
|
||||||
"t5_close": change(valid[4]["close"]) if observed >= 5 else None,
|
|
||||||
"max_gain": max(change(row["high"]) for row in valid),
|
|
||||||
"max_drawdown": min(change(row["low"]) for row in valid),
|
|
||||||
}
|
|
||||||
|
|
||||||
@staticmethod
|
|
||||||
def _summary(items: list[dict[str, Any]]) -> dict[str, Any]:
|
|
||||||
completed = [item for item in items if item.get("t5_close") is not None]
|
|
||||||
t1 = [float(item["t1_close"]) for item in items if item.get("t1_close") is not None]
|
|
||||||
t5 = [float(item["t5_close"]) for item in completed]
|
|
||||||
return {
|
|
||||||
"total": len(items),
|
|
||||||
"observed": len(t1),
|
|
||||||
"completed": len(completed),
|
|
||||||
"t1_win_rate": round(sum(value > 0 for value in t1) / len(t1) * 100, 1) if t1 else None,
|
|
||||||
"t5_win_rate": round(sum(value > 0 for value in t5) / len(t5) * 100, 1) if t5 else None,
|
|
||||||
"average_t5": round(sum(t5) / len(t5), 2) if t5 else None,
|
|
||||||
}
|
|
||||||
|
|||||||
@@ -0,0 +1,56 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import ast
|
||||||
|
import unittest
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
|
||||||
|
ROOT = Path(__file__).resolve().parents[1]
|
||||||
|
FEATURES = ROOT / "backend" / "features"
|
||||||
|
|
||||||
|
|
||||||
|
class FeatureBoundaryTests(unittest.TestCase):
|
||||||
|
def test_feature_services_do_not_import_http_or_provider_adapters(self) -> None:
|
||||||
|
forbidden = {
|
||||||
|
"server",
|
||||||
|
"tushare_client",
|
||||||
|
"ifind_client",
|
||||||
|
"chart_data_provider",
|
||||||
|
"realtime_aggregator",
|
||||||
|
}
|
||||||
|
violations = []
|
||||||
|
for path in FEATURES.rglob("*.py"):
|
||||||
|
tree = ast.parse(path.read_text(encoding="utf-8"), filename=str(path))
|
||||||
|
for node in ast.walk(tree):
|
||||||
|
names = []
|
||||||
|
if isinstance(node, ast.Import):
|
||||||
|
names = [alias.name for alias in node.names]
|
||||||
|
elif isinstance(node, ast.ImportFrom) and node.module:
|
||||||
|
names = [node.module]
|
||||||
|
for name in names:
|
||||||
|
if name.split(".")[0] in forbidden:
|
||||||
|
violations.append(f"{path.relative_to(ROOT)} -> {name}")
|
||||||
|
self.assertEqual(violations, [])
|
||||||
|
|
||||||
|
def test_legacy_service_modules_are_compatibility_exports_only(self) -> None:
|
||||||
|
for filename in ("alert_service.py", "trade_journal.py", "strategy_tracking.py"):
|
||||||
|
tree = ast.parse((ROOT / filename).read_text(encoding="utf-8"))
|
||||||
|
definitions = [
|
||||||
|
node for node in tree.body
|
||||||
|
if isinstance(node, (ast.ClassDef, ast.FunctionDef, ast.AsyncFunctionDef))
|
||||||
|
]
|
||||||
|
self.assertEqual(definitions, [], filename)
|
||||||
|
|
||||||
|
def test_each_migrated_feature_owns_one_application_service(self) -> None:
|
||||||
|
expected = {
|
||||||
|
"alerts/service.py": "AlertService",
|
||||||
|
"review/trade_journal.py": "TradeJournalService",
|
||||||
|
"screener/tracking.py": "StrategyTrackingService",
|
||||||
|
}
|
||||||
|
for relative, class_name in expected.items():
|
||||||
|
tree = ast.parse((FEATURES / relative).read_text(encoding="utf-8"))
|
||||||
|
self.assertIn(class_name, {node.name for node in tree.body if isinstance(node, ast.ClassDef)})
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main()
|
||||||
+2
-99
@@ -1,100 +1,3 @@
|
|||||||
from __future__ import annotations
|
from backend.features.review.trade_journal import EMOTIONS, TRADE_ACTIONS, TradeJournalService
|
||||||
|
|
||||||
import json
|
__all__ = ["EMOTIONS", "TRADE_ACTIONS", "TradeJournalService"]
|
||||||
from datetime import date
|
|
||||||
from typing import Any
|
|
||||||
|
|
||||||
from app_config import normalize_date, validate_stock_code, validate_text
|
|
||||||
from backend.database.repositories import TradeJournalRepository
|
|
||||||
|
|
||||||
|
|
||||||
TRADE_ACTIONS = {"buy": "买入", "sell": "卖出", "trim": "减仓", "add": "加仓", "watch": "观察"}
|
|
||||||
EMOTIONS = {"calm": "平静", "confident": "笃定", "hesitant": "犹豫", "anxious": "焦虑", "impulsive": "冲动"}
|
|
||||||
|
|
||||||
|
|
||||||
class TradeJournalService:
|
|
||||||
def __init__(self, repository: TradeJournalRepository) -> None:
|
|
||||||
self.repository = repository
|
|
||||||
|
|
||||||
def save(self, user_id: int, payload: dict[str, Any]) -> int:
|
|
||||||
trade_id = int(payload.get("id") or 0)
|
|
||||||
trade_date = normalize_date(str(payload.get("trade_date") or date.today().isoformat()))
|
|
||||||
code = validate_stock_code(str(payload.get("code") or ""))
|
|
||||||
name = validate_text(payload.get("name"), "股票名称", 40, required=True)
|
|
||||||
action = str(payload.get("action") or "")
|
|
||||||
if action not in TRADE_ACTIONS:
|
|
||||||
raise ValueError("交易动作不支持。")
|
|
||||||
emotion = str(payload.get("emotion") or "calm")
|
|
||||||
if emotion not in EMOTIONS:
|
|
||||||
raise ValueError("交易情绪不支持。")
|
|
||||||
price = self._number(payload.get("price"), "成交价格", 0, 1000000, required=True)
|
|
||||||
quantity = int(self._number(payload.get("quantity"), "成交数量", 0, 100000000))
|
|
||||||
position_pct = self._number(payload.get("position_pct"), "仓位", 0, 100)
|
|
||||||
pnl_amount = self._optional_number(payload.get("pnl_amount"), "盈亏金额", -1e12, 1e12)
|
|
||||||
pnl_pct = self._optional_number(payload.get("pnl_pct"), "盈亏比例", -1000, 10000)
|
|
||||||
thesis = validate_text(payload.get("thesis"), "交易逻辑", 2000)
|
|
||||||
execution = validate_text(payload.get("execution"), "执行复核", 2000)
|
|
||||||
raw_tags = payload.get("tags") or []
|
|
||||||
if isinstance(raw_tags, str):
|
|
||||||
raw_tags = [item.strip() for item in raw_tags.replace(",", ",").split(",")]
|
|
||||||
if not isinstance(raw_tags, list):
|
|
||||||
raise ValueError("交易标签格式不正确。")
|
|
||||||
tags = [validate_text(item, "交易标签", 20) for item in raw_tags if str(item).strip()][:8]
|
|
||||||
return self.repository.save_trade_entry(
|
|
||||||
user_id, trade_date, code, name, action, price, quantity, position_pct,
|
|
||||||
pnl_amount, pnl_pct, thesis, execution, emotion, tags, trade_id or None,
|
|
||||||
)
|
|
||||||
|
|
||||||
def list_entries(
|
|
||||||
self, user_id: int, start_date: str = "", end_date: str = "", code: str = ""
|
|
||||||
) -> dict[str, Any]:
|
|
||||||
start = normalize_date(start_date) if start_date else ""
|
|
||||||
end = normalize_date(end_date) if end_date else date.today().strftime("%Y%m%d")
|
|
||||||
if start and start > end:
|
|
||||||
raise ValueError("开始日期不能晚于结束日期。")
|
|
||||||
code = validate_stock_code(code) if code else ""
|
|
||||||
items = self.repository.list_trade_entries(user_id, start, end, code)
|
|
||||||
for item in items:
|
|
||||||
item["tags"] = json.loads(item.get("tags") or "[]")
|
|
||||||
item["action_label"] = TRADE_ACTIONS.get(item["action"], item["action"])
|
|
||||||
item["emotion_label"] = EMOTIONS.get(item["emotion"], item["emotion"])
|
|
||||||
realized = [item for item in items if item.get("pnl_pct") is not None]
|
|
||||||
return {"items": items, "summary": self._summary(items, realized)}
|
|
||||||
|
|
||||||
def delete(self, user_id: int, trade_id: int) -> bool:
|
|
||||||
return self.repository.delete_trade_entry(user_id, trade_id)
|
|
||||||
|
|
||||||
@staticmethod
|
|
||||||
def _summary(items: list[dict[str, Any]], realized: list[dict[str, Any]]) -> dict[str, Any]:
|
|
||||||
pnl_amounts = [float(item["pnl_amount"]) for item in realized if item.get("pnl_amount") is not None]
|
|
||||||
positions = [float(item["position_pct"]) for item in items if float(item.get("position_pct") or 0) > 0]
|
|
||||||
wins = sum(float(item.get("pnl_pct") or 0) > 0 for item in realized)
|
|
||||||
return {
|
|
||||||
"total": len(items),
|
|
||||||
"realized": len(realized),
|
|
||||||
"win_rate": round(wins / len(realized) * 100, 1) if realized else None,
|
|
||||||
"pnl_amount": round(sum(pnl_amounts), 2) if pnl_amounts else None,
|
|
||||||
"average_position": round(sum(positions) / len(positions), 1) if positions else None,
|
|
||||||
}
|
|
||||||
|
|
||||||
@staticmethod
|
|
||||||
def _number(value: Any, label: str, minimum: float, maximum: float, required: bool = False) -> float:
|
|
||||||
if value in (None, ""):
|
|
||||||
if required:
|
|
||||||
raise ValueError(f"{label}不能为空。")
|
|
||||||
return 0.0
|
|
||||||
try:
|
|
||||||
parsed = float(value)
|
|
||||||
except (TypeError, ValueError) as exc:
|
|
||||||
raise ValueError(f"{label}格式不正确。") from exc
|
|
||||||
if parsed < minimum or parsed > maximum:
|
|
||||||
raise ValueError(f"{label}超出允许范围。")
|
|
||||||
return parsed
|
|
||||||
|
|
||||||
@classmethod
|
|
||||||
def _optional_number(
|
|
||||||
cls, value: Any, label: str, minimum: float, maximum: float
|
|
||||||
) -> float | None:
|
|
||||||
if value in (None, ""):
|
|
||||||
return None
|
|
||||||
return cls._number(value, label, minimum, maximum, required=True)
|
|
||||||
|
|||||||
Reference in New Issue
Block a user