refactor: enforce account scoped repository boundaries
This commit is contained in:
+17
-8
@@ -5,12 +5,12 @@ from datetime import date, datetime
|
||||
from typing import Any
|
||||
|
||||
from app_config import validate_text
|
||||
from database import ReviewDatabase
|
||||
from backend.database.repositories import AlertRepository
|
||||
|
||||
|
||||
class AlertService:
|
||||
def __init__(self, database: ReviewDatabase) -> None:
|
||||
self.database = database
|
||||
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)
|
||||
@@ -19,7 +19,7 @@ class AlertService:
|
||||
available_date = self.calendar_date(
|
||||
str(payload.get("remind_date") or date.today().isoformat())
|
||||
)
|
||||
return self.database.save_alert(
|
||||
return self.repository.save_alert(
|
||||
user_id=user_id,
|
||||
kind="manual",
|
||||
title=title,
|
||||
@@ -44,7 +44,7 @@ class AlertService:
|
||||
if observed:
|
||||
win_rate = summary.get("t1_win_rate")
|
||||
suffix = f",当前红盘率 {win_rate:.1f}%" if win_rate is not None else ""
|
||||
self.database.save_alert(
|
||||
self.repository.save_alert(
|
||||
user_id, "strategy_t1", f"{strategy_name} 已有 T+1 反馈",
|
||||
f"{observed}/{len(items)} 只标的已有首日表现{suffix}。",
|
||||
today, "", f"strategy:{run_id}:t1",
|
||||
@@ -53,7 +53,7 @@ class AlertService:
|
||||
if completed == len(items):
|
||||
average = summary.get("average_t5")
|
||||
suffix = f",平均收益 {average:+.2f}%" if average is not None else ""
|
||||
self.database.save_alert(
|
||||
self.repository.save_alert(
|
||||
user_id, "strategy_t5", f"{strategy_name} 五日跟踪完成",
|
||||
f"本批 {len(items)} 只标的已完成 T+5 跟踪{suffix}。",
|
||||
today, "", f"strategy:{run_id}:t5",
|
||||
@@ -67,15 +67,24 @@ class AlertService:
|
||||
if status not in {"all", "unread"}:
|
||||
raise ValueError("提醒筛选不支持。")
|
||||
compact_date = self.calendar_date(as_of or date.today().isoformat())
|
||||
items = self.database.list_alerts(user_id, compact_date, status == "unread")
|
||||
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.database.count_unread_alerts(user_id, compact_date),
|
||||
"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()
|
||||
|
||||
@@ -6,6 +6,7 @@ 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 chart_data_provider import MarketChartClient
|
||||
from database import ReviewDatabase
|
||||
from ifind_client import IfindHttpClient
|
||||
@@ -19,6 +20,7 @@ from trade_journal import TradeJournalService
|
||||
@dataclass(frozen=True)
|
||||
class ApplicationContainer:
|
||||
database: ReviewDatabase
|
||||
repositories: RepositoryBundle
|
||||
data_gateway: DataGateway
|
||||
ifind: IfindHttpClient
|
||||
screener: ScreenerEngine
|
||||
@@ -38,14 +40,16 @@ def build_application_container(
|
||||
tushare_token_supplier: Callable[[], str] | None = None,
|
||||
) -> ApplicationContainer:
|
||||
data_gateway = build_data_gateway(credentials, tushare_token_supplier)
|
||||
repositories = build_repository_bundle(database)
|
||||
return ApplicationContainer(
|
||||
database=database,
|
||||
repositories=repositories,
|
||||
data_gateway=data_gateway,
|
||||
ifind=data_gateway.ifind,
|
||||
screener=ScreenerEngine(database),
|
||||
strategy_tracking=StrategyTrackingService(database),
|
||||
alert_service=AlertService(database),
|
||||
trade_journal=TradeJournalService(database),
|
||||
strategy_tracking=StrategyTrackingService(repositories.strategy_tracking),
|
||||
alert_service=AlertService(repositories.alerts),
|
||||
trade_journal=TradeJournalService(repositories.trades),
|
||||
mentor_skills=MentorSkillRegistry(mentor_skills_dir, private_mentor_skills_dir),
|
||||
realtime_aggregator=data_gateway.realtime_observer,
|
||||
chart_data=data_gateway.chart_data,
|
||||
|
||||
@@ -0,0 +1,21 @@
|
||||
from .ports import AlertRepository, StrategyTrackingRepository, TradeJournalRepository
|
||||
from .sqlite import (
|
||||
RepositoryBundle,
|
||||
SQLiteAlertRepository,
|
||||
SQLiteStrategyTrackingRepository,
|
||||
SQLiteTradeJournalRepository,
|
||||
build_repository_bundle,
|
||||
require_user_id,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"AlertRepository",
|
||||
"RepositoryBundle",
|
||||
"SQLiteAlertRepository",
|
||||
"SQLiteStrategyTrackingRepository",
|
||||
"SQLiteTradeJournalRepository",
|
||||
"StrategyTrackingRepository",
|
||||
"TradeJournalRepository",
|
||||
"build_repository_bundle",
|
||||
"require_user_id",
|
||||
]
|
||||
@@ -0,0 +1,52 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any, Protocol
|
||||
|
||||
|
||||
class AlertRepository(Protocol):
|
||||
def save_alert(
|
||||
self, user_id: int, kind: str, title: str, content: str,
|
||||
available_date: str, code: str, dedupe_key: str,
|
||||
) -> int: ...
|
||||
|
||||
def list_alerts(
|
||||
self, user_id: int, as_of: str, unread_only: bool = False, limit: int = 100,
|
||||
) -> list[dict[str, Any]]: ...
|
||||
|
||||
def count_unread_alerts(self, user_id: int, as_of: str) -> int: ...
|
||||
|
||||
def mark_alert_read(self, user_id: int, alert_id: int) -> bool: ...
|
||||
|
||||
def mark_all_alerts_read(self, user_id: int, as_of: str) -> int: ...
|
||||
|
||||
def delete_alert(self, user_id: int, alert_id: int) -> bool: ...
|
||||
|
||||
|
||||
class TradeJournalRepository(Protocol):
|
||||
def save_trade_entry(self, *args: Any, **kwargs: Any) -> int: ...
|
||||
|
||||
def list_trade_entries(
|
||||
self, user_id: int, start_date: str = "", end_date: str = "",
|
||||
code: str = "", limit: int = 300,
|
||||
) -> list[dict[str, Any]]: ...
|
||||
|
||||
def delete_trade_entry(self, user_id: int, trade_id: int) -> bool: ...
|
||||
|
||||
|
||||
class StrategyTrackingRepository(Protocol):
|
||||
def save_strategy_tracks(
|
||||
self, user_id: int, run_id: int, selection_date: str,
|
||||
strategy_name: str, candidates: list[dict[str, Any]],
|
||||
) -> int: ...
|
||||
|
||||
def get_screener_run(self, user_id: int, run_id: int) -> dict[str, Any] | None: ...
|
||||
|
||||
def delete_strategy_track(self, user_id: int, track_id: int) -> bool: ...
|
||||
|
||||
def list_strategy_tracks(
|
||||
self, user_id: int, limit_batches: int = 12,
|
||||
) -> list[dict[str, Any]]: ...
|
||||
|
||||
def load_tracking_bars(
|
||||
self, targets: list[tuple[str, str]], limit: int = 5,
|
||||
) -> dict[tuple[str, str], list[dict[str, Any]]]: ...
|
||||
@@ -0,0 +1,108 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
from database import ReviewDatabase
|
||||
|
||||
|
||||
def require_user_id(value: int) -> int:
|
||||
user_id = int(value)
|
||||
if user_id <= 0:
|
||||
raise ValueError("A positive account owner is required")
|
||||
return user_id
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class SQLiteAlertRepository:
|
||||
database: ReviewDatabase
|
||||
|
||||
def save_alert(self, user_id: int, *args: Any, **kwargs: Any) -> int:
|
||||
return self.database.save_alert(require_user_id(user_id), *args, **kwargs)
|
||||
|
||||
def list_alerts(
|
||||
self, user_id: int, as_of: str, unread_only: bool = False, limit: int = 100,
|
||||
) -> list[dict[str, Any]]:
|
||||
return self.database.list_alerts(
|
||||
require_user_id(user_id), as_of, unread_only, limit
|
||||
)
|
||||
|
||||
def count_unread_alerts(self, user_id: int, as_of: str) -> int:
|
||||
return self.database.count_unread_alerts(require_user_id(user_id), as_of)
|
||||
|
||||
def mark_alert_read(self, user_id: int, alert_id: int) -> bool:
|
||||
return self.database.mark_alert_read(require_user_id(user_id), alert_id)
|
||||
|
||||
def mark_all_alerts_read(self, user_id: int, as_of: str) -> int:
|
||||
return self.database.mark_all_alerts_read(require_user_id(user_id), as_of)
|
||||
|
||||
def delete_alert(self, user_id: int, alert_id: int) -> bool:
|
||||
return self.database.delete_alert(require_user_id(user_id), alert_id)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class SQLiteTradeJournalRepository:
|
||||
database: ReviewDatabase
|
||||
|
||||
def save_trade_entry(self, user_id: int, *args: Any, **kwargs: Any) -> int:
|
||||
return self.database.save_trade_entry(require_user_id(user_id), *args, **kwargs)
|
||||
|
||||
def list_trade_entries(
|
||||
self, user_id: int, start_date: str = "", end_date: str = "",
|
||||
code: str = "", limit: int = 300,
|
||||
) -> list[dict[str, Any]]:
|
||||
return self.database.list_trade_entries(
|
||||
require_user_id(user_id), start_date, end_date, code, limit
|
||||
)
|
||||
|
||||
def delete_trade_entry(self, user_id: int, trade_id: int) -> bool:
|
||||
return self.database.delete_trade_entry(require_user_id(user_id), trade_id)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class SQLiteStrategyTrackingRepository:
|
||||
database: ReviewDatabase
|
||||
|
||||
def save_strategy_tracks(
|
||||
self, user_id: int, run_id: int, selection_date: str,
|
||||
strategy_name: str, candidates: list[dict[str, Any]],
|
||||
) -> int:
|
||||
return self.database.save_strategy_tracks(
|
||||
require_user_id(user_id), run_id, selection_date, strategy_name, candidates
|
||||
)
|
||||
|
||||
def get_screener_run(self, user_id: int, run_id: int) -> dict[str, Any] | None:
|
||||
owner_id = int(user_id)
|
||||
if owner_id < 0:
|
||||
raise ValueError("Account owner cannot be negative")
|
||||
return self.database.get_screener_run(owner_id, run_id)
|
||||
|
||||
def delete_strategy_track(self, user_id: int, track_id: int) -> bool:
|
||||
return self.database.delete_strategy_track(require_user_id(user_id), track_id)
|
||||
|
||||
def list_strategy_tracks(
|
||||
self, user_id: int, limit_batches: int = 12,
|
||||
) -> list[dict[str, Any]]:
|
||||
return self.database.list_strategy_tracks(
|
||||
require_user_id(user_id), limit_batches
|
||||
)
|
||||
|
||||
def load_tracking_bars(
|
||||
self, targets: list[tuple[str, str]], limit: int = 5,
|
||||
) -> dict[tuple[str, str], list[dict[str, Any]]]:
|
||||
return self.database.load_tracking_bars(targets, limit)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class RepositoryBundle:
|
||||
alerts: SQLiteAlertRepository
|
||||
trades: SQLiteTradeJournalRepository
|
||||
strategy_tracking: SQLiteStrategyTrackingRepository
|
||||
|
||||
|
||||
def build_repository_bundle(database: ReviewDatabase) -> RepositoryBundle:
|
||||
return RepositoryBundle(
|
||||
alerts=SQLiteAlertRepository(database),
|
||||
trades=SQLiteTradeJournalRepository(database),
|
||||
strategy_tracking=SQLiteStrategyTrackingRepository(database),
|
||||
)
|
||||
@@ -265,7 +265,7 @@
|
||||
},
|
||||
{
|
||||
"path": "server.py",
|
||||
"bytes": 267835,
|
||||
"bytes": 267824,
|
||||
"lines": 5947
|
||||
},
|
||||
{
|
||||
|
||||
@@ -0,0 +1,27 @@
|
||||
# Stage 09: Account-Scoped Repository Boundaries
|
||||
|
||||
Date: 2026-07-29
|
||||
|
||||
## Result
|
||||
|
||||
- Added narrow repository ports for alerts, the trade journal, and strategy tracking.
|
||||
- Added SQLite adapters that expose only the persistence operations each service requires.
|
||||
- Required a positive account owner on every private read, write, update, and delete path.
|
||||
- Kept the shared automatic screener-run lookup explicit as the sole zero-owner read in the
|
||||
strategy tracking adapter.
|
||||
- Routed the application container through a repository bundle.
|
||||
- Removed direct alert and trade deletion/update calls from the HTTP-facing dashboard service.
|
||||
- Preserved structural compatibility for isolated tests and gradual extraction from the legacy
|
||||
database facade.
|
||||
|
||||
## Boundary
|
||||
|
||||
Application services depend on repository protocols. The SQLite implementation may later be
|
||||
replaced without changing those services. Repository adapters do not call market providers,
|
||||
and provider adapters do not access user tables.
|
||||
|
||||
## Residual Migration
|
||||
|
||||
The legacy `ReviewDatabase` still contains the SQL behind these adapters and remains the
|
||||
compatibility facade for features not yet extracted. Subsequent feature stages can move SQL
|
||||
behind the same ports one feature at a time after account-isolation tests pass.
|
||||
@@ -1494,16 +1494,16 @@ class DashboardService:
|
||||
return {"id": alert_id, **self.alert_center()}
|
||||
|
||||
def mark_alert_read(self, alert_id: int) -> dict[str, Any]:
|
||||
self.database.mark_alert_read(self.current_user_id, alert_id)
|
||||
self.alert_service.mark_read(self.current_user_id, alert_id)
|
||||
return self.alert_center()
|
||||
|
||||
def mark_all_alerts_read(self, as_of: str = "") -> dict[str, Any]:
|
||||
compact_date = self.alert_service.calendar_date(as_of or date.today().isoformat())
|
||||
self.database.mark_all_alerts_read(self.current_user_id, compact_date)
|
||||
self.alert_service.mark_all_read(self.current_user_id, compact_date)
|
||||
return self.alert_center(as_of=compact_date)
|
||||
|
||||
def delete_alert(self, alert_id: int) -> dict[str, Any]:
|
||||
deleted = self.database.delete_alert(self.current_user_id, alert_id)
|
||||
deleted = self.alert_service.delete(self.current_user_id, alert_id)
|
||||
return {"deleted": deleted, **self.alert_center()}
|
||||
|
||||
def trade_entries(
|
||||
@@ -1598,7 +1598,7 @@ class DashboardService:
|
||||
return {"id": trade_id, **self.trade_entries()}
|
||||
|
||||
def delete_trade_entry(self, trade_id: int) -> dict[str, Any]:
|
||||
deleted = self.database.delete_trade_entry(self.current_user_id, trade_id)
|
||||
deleted = self.trade_journal.delete(self.current_user_id, trade_id)
|
||||
return {"deleted": deleted, **self.trade_entries()}
|
||||
|
||||
def assistant_messages(self) -> list[dict[str, Any]]:
|
||||
|
||||
@@ -2,12 +2,12 @@ from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
from database import ReviewDatabase
|
||||
from backend.database.repositories import StrategyTrackingRepository
|
||||
|
||||
|
||||
class StrategyTrackingService:
|
||||
def __init__(self, database: ReviewDatabase) -> None:
|
||||
self.database = database
|
||||
def __init__(self, repository: StrategyTrackingRepository) -> None:
|
||||
self.repository = repository
|
||||
|
||||
def record_run(
|
||||
self,
|
||||
@@ -17,14 +17,14 @@ class StrategyTrackingService:
|
||||
strategy_name: str,
|
||||
candidates: list[dict[str, Any]],
|
||||
) -> int:
|
||||
return self.database.save_strategy_tracks(
|
||||
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.database.get_screener_run(user_id, run_id)
|
||||
run = self.repository.get_screener_run(user_id, run_id)
|
||||
if not run:
|
||||
run = self.database.get_screener_run(0, run_id)
|
||||
run = self.repository.get_screener_run(0, run_id)
|
||||
if not run:
|
||||
raise ValueError("选股结果不存在或不属于当前账号。")
|
||||
normalized_code = str(code or "").strip().split(".")[0]
|
||||
@@ -48,15 +48,15 @@ class StrategyTrackingService:
|
||||
return {"added": added, "tracking": self.list_tracking(user_id)}
|
||||
|
||||
def remove_candidate(self, user_id: int, track_id: int) -> dict[str, Any]:
|
||||
deleted = self.database.delete_strategy_track(user_id, track_id)
|
||||
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.database.list_strategy_tracks(user_id, limit_batches)
|
||||
tracks = self.repository.list_strategy_tracks(user_id, limit_batches)
|
||||
if not tracks:
|
||||
return {"batches": [], "summary": self._summary([])}
|
||||
|
||||
bars = self.database.load_tracking_bars(
|
||||
bars = self.repository.load_tracking_bars(
|
||||
[(item["ts_code"], item["selection_date"]) for item in tracks], 5
|
||||
)
|
||||
batches: dict[int, dict[str, Any]] = {}
|
||||
|
||||
@@ -42,9 +42,9 @@ class BootstrapContainerTests(unittest.TestCase):
|
||||
)
|
||||
self.assertIs(container.database, database)
|
||||
self.assertIs(container.screener.database, database)
|
||||
self.assertIs(container.strategy_tracking.database, database)
|
||||
self.assertIs(container.alert_service.database, database)
|
||||
self.assertIs(container.trade_journal.database, database)
|
||||
self.assertIs(container.strategy_tracking.repository.database, database)
|
||||
self.assertIs(container.alert_service.repository.database, database)
|
||||
self.assertIs(container.trade_journal.repository.database, database)
|
||||
self.assertIs(container.chart_data.ifind, container.ifind)
|
||||
self.assertTrue(container.ifind.configured)
|
||||
|
||||
|
||||
@@ -0,0 +1,54 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import tempfile
|
||||
import unittest
|
||||
from pathlib import Path
|
||||
|
||||
from alert_service import AlertService
|
||||
from backend.bootstrap.container import build_application_container
|
||||
from backend.database.repositories import (
|
||||
SQLiteAlertRepository,
|
||||
SQLiteStrategyTrackingRepository,
|
||||
SQLiteTradeJournalRepository,
|
||||
)
|
||||
from database import ReviewDatabase
|
||||
|
||||
|
||||
class RepositoryBoundaryTests(unittest.TestCase):
|
||||
def setUp(self) -> None:
|
||||
self.temporary = tempfile.TemporaryDirectory()
|
||||
self.addCleanup(self.temporary.cleanup)
|
||||
self.database = ReviewDatabase(Path(self.temporary.name) / "review.db")
|
||||
|
||||
def test_container_injects_narrow_user_repositories(self) -> None:
|
||||
skills = Path(self.temporary.name) / "skills"
|
||||
private = Path(self.temporary.name) / "private"
|
||||
skills.mkdir()
|
||||
private.mkdir()
|
||||
container = build_application_container(self.database, {}, skills, private)
|
||||
self.assertIsInstance(container.repositories.alerts, SQLiteAlertRepository)
|
||||
self.assertIsInstance(container.repositories.trades, SQLiteTradeJournalRepository)
|
||||
self.assertIsInstance(
|
||||
container.repositories.strategy_tracking, SQLiteStrategyTrackingRepository
|
||||
)
|
||||
self.assertIs(container.alert_service.repository, container.repositories.alerts)
|
||||
self.assertIs(container.trade_journal.repository, container.repositories.trades)
|
||||
|
||||
def test_private_repositories_reject_missing_account_owner(self) -> None:
|
||||
alerts = SQLiteAlertRepository(self.database)
|
||||
trades = SQLiteTradeJournalRepository(self.database)
|
||||
tracking = SQLiteStrategyTrackingRepository(self.database)
|
||||
with self.assertRaises(ValueError):
|
||||
alerts.list_alerts(0, "20260729")
|
||||
with self.assertRaises(ValueError):
|
||||
trades.list_trade_entries(0)
|
||||
with self.assertRaises(ValueError):
|
||||
tracking.list_strategy_tracks(0)
|
||||
|
||||
def test_alert_service_keeps_database_compatible_structural_port(self) -> None:
|
||||
service = AlertService(self.database)
|
||||
self.assertIs(service.repository, self.database)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
+8
-5
@@ -5,7 +5,7 @@ from datetime import date
|
||||
from typing import Any
|
||||
|
||||
from app_config import normalize_date, validate_stock_code, validate_text
|
||||
from database import ReviewDatabase
|
||||
from backend.database.repositories import TradeJournalRepository
|
||||
|
||||
|
||||
TRADE_ACTIONS = {"buy": "买入", "sell": "卖出", "trim": "减仓", "add": "加仓", "watch": "观察"}
|
||||
@@ -13,8 +13,8 @@ EMOTIONS = {"calm": "平静", "confident": "笃定", "hesitant": "犹豫", "anxi
|
||||
|
||||
|
||||
class TradeJournalService:
|
||||
def __init__(self, database: ReviewDatabase) -> None:
|
||||
self.database = database
|
||||
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)
|
||||
@@ -40,7 +40,7 @@ class TradeJournalService:
|
||||
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.database.save_trade_entry(
|
||||
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,
|
||||
)
|
||||
@@ -53,7 +53,7 @@ class TradeJournalService:
|
||||
if start and start > end:
|
||||
raise ValueError("开始日期不能晚于结束日期。")
|
||||
code = validate_stock_code(code) if code else ""
|
||||
items = self.database.list_trade_entries(user_id, start, end, code)
|
||||
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"])
|
||||
@@ -61,6 +61,9 @@ class TradeJournalService:
|
||||
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]
|
||||
|
||||
Reference in New Issue
Block a user