135 lines
5.3 KiB
Python
135 lines
5.3 KiB
Python
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,
|
|
}
|