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
|
||||
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"]
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
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"]
|
||||
|
||||
@@ -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
|
||||
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"]
|
||||
|
||||
Reference in New Issue
Block a user