From b0cbd25f649d9843adabf6dcdedc6cb5a54af507 Mon Sep 17 00:00:00 2001 From: leefer Date: Wed, 29 Jul 2026 18:00:19 +0800 Subject: [PATCH] refactor: establish feature application service layout --- alert_service.py | 96 +------------ backend/bootstrap/container.py | 6 +- backend/features/__init__.py | 1 + backend/features/alerts/__init__.py | 3 + backend/features/alerts/service.py | 95 +++++++++++++ backend/features/review/__init__.py | 3 + backend/features/review/trade_journal.py | 100 ++++++++++++++ backend/features/screener/__init__.py | 3 + backend/features/screener/tracking.py | 134 ++++++++++++++++++ docs/governance/stage-10-feature-services.md | 28 ++++ strategy_tracking.py | 135 +------------------ tests/test_feature_boundaries.py | 56 ++++++++ trade_journal.py | 101 +------------- 13 files changed, 432 insertions(+), 329 deletions(-) create mode 100644 backend/features/__init__.py create mode 100644 backend/features/alerts/__init__.py create mode 100644 backend/features/alerts/service.py create mode 100644 backend/features/review/__init__.py create mode 100644 backend/features/review/trade_journal.py create mode 100644 backend/features/screener/__init__.py create mode 100644 backend/features/screener/tracking.py create mode 100644 docs/governance/stage-10-feature-services.md create mode 100644 tests/test_feature_boundaries.py diff --git a/alert_service.py b/alert_service.py index bd81af4..f4264ee 100644 --- a/alert_service.py +++ b/alert_service.py @@ -1,95 +1,3 @@ -from __future__ import annotations +from backend.features.alerts.service import AlertService -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") +__all__ = ["AlertService"] diff --git a/backend/bootstrap/container.py b/backend/bootstrap/container.py index 7b87880..b2ea3f7 100644 --- a/backend/bootstrap/container.py +++ b/backend/bootstrap/container.py @@ -4,17 +4,17 @@ from dataclasses import dataclass from pathlib import Path from collections.abc import Callable -from alert_service import AlertService from backend.data import DataGateway, build_data_gateway 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 database import ReviewDatabase from ifind_client import IfindHttpClient from mentor_agent import MentorSkillRegistry from realtime_aggregator import WebRealtimeAggregator from screener import ScreenerEngine -from strategy_tracking import StrategyTrackingService -from trade_journal import TradeJournalService @dataclass(frozen=True) diff --git a/backend/features/__init__.py b/backend/features/__init__.py new file mode 100644 index 0000000..86a7393 --- /dev/null +++ b/backend/features/__init__.py @@ -0,0 +1 @@ +"""Feature-owned application services.""" diff --git a/backend/features/alerts/__init__.py b/backend/features/alerts/__init__.py new file mode 100644 index 0000000..f2ba0fe --- /dev/null +++ b/backend/features/alerts/__init__.py @@ -0,0 +1,3 @@ +from .service import AlertService + +__all__ = ["AlertService"] diff --git a/backend/features/alerts/service.py b/backend/features/alerts/service.py new file mode 100644 index 0000000..bd81af4 --- /dev/null +++ b/backend/features/alerts/service.py @@ -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") diff --git a/backend/features/review/__init__.py b/backend/features/review/__init__.py new file mode 100644 index 0000000..aff8464 --- /dev/null +++ b/backend/features/review/__init__.py @@ -0,0 +1,3 @@ +from .trade_journal import EMOTIONS, TRADE_ACTIONS, TradeJournalService + +__all__ = ["EMOTIONS", "TRADE_ACTIONS", "TradeJournalService"] diff --git a/backend/features/review/trade_journal.py b/backend/features/review/trade_journal.py new file mode 100644 index 0000000..30fb54e --- /dev/null +++ b/backend/features/review/trade_journal.py @@ -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) diff --git a/backend/features/screener/__init__.py b/backend/features/screener/__init__.py new file mode 100644 index 0000000..2fe3373 --- /dev/null +++ b/backend/features/screener/__init__.py @@ -0,0 +1,3 @@ +from .tracking import StrategyTrackingService + +__all__ = ["StrategyTrackingService"] diff --git a/backend/features/screener/tracking.py b/backend/features/screener/tracking.py new file mode 100644 index 0000000..f20c646 --- /dev/null +++ b/backend/features/screener/tracking.py @@ -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, + } diff --git a/docs/governance/stage-10-feature-services.md b/docs/governance/stage-10-feature-services.md new file mode 100644 index 0000000..88067dd --- /dev/null +++ b/docs/governance/stage-10-feature-services.md @@ -0,0 +1,28 @@ +# Stage 10: Feature Application Service Layout + +Date: 2026-07-29 + +## Result + +- Created the governed `backend/features/` 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. diff --git a/strategy_tracking.py b/strategy_tracking.py index f20c646..ad0416f 100644 --- a/strategy_tracking.py +++ b/strategy_tracking.py @@ -1,134 +1,3 @@ -from __future__ import annotations +from backend.features.screener.tracking import StrategyTrackingService -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, - } +__all__ = ["StrategyTrackingService"] diff --git a/tests/test_feature_boundaries.py b/tests/test_feature_boundaries.py new file mode 100644 index 0000000..085cc62 --- /dev/null +++ b/tests/test_feature_boundaries.py @@ -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() diff --git a/trade_journal.py b/trade_journal.py index 30fb54e..d227544 100644 --- a/trade_journal.py +++ b/trade_journal.py @@ -1,100 +1,3 @@ -from __future__ import annotations +from backend.features.review.trade_journal import EMOTIONS, TRADE_ACTIONS, TradeJournalService -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) +__all__ = ["EMOTIONS", "TRADE_ACTIONS", "TradeJournalService"]