refactor: establish feature application service layout

This commit is contained in:
leefer
2026-07-29 18:00:19 +08:00
parent 367fe71fbf
commit b0cbd25f64
13 changed files with 432 additions and 329 deletions
+2 -94
View File
@@ -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"]
+3 -3
View File
@@ -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)
+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,
}
@@ -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
View File
@@ -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"]
+56
View File
@@ -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
View File
@@ -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"]