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, }