migration: establish exact preserved app baseline
This commit is contained in:
@@ -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)
|
||||
Reference in New Issue
Block a user