migration: establish exact preserved app baseline

This commit is contained in:
leefer
2026-07-30 23:51:48 +08:00
commit e4a9b2e647
389 changed files with 126625 additions and 0 deletions
+1
View File
@@ -0,0 +1 @@
"""Feature-owned application services."""
+3
View File
@@ -0,0 +1,3 @@
from .service import AlertService
__all__ = ["AlertService"]
+95
View File
@@ -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")
+3
View File
@@ -0,0 +1,3 @@
from .trade_journal import EMOTIONS, TRADE_ACTIONS, TradeJournalService
__all__ = ["EMOTIONS", "TRADE_ACTIONS", "TradeJournalService"]
+100
View File
@@ -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)
+3
View File
@@ -0,0 +1,3 @@
from .tracking import StrategyTrackingService
__all__ = ["StrategyTrackingService"]
+134
View File
@@ -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,
}