diff --git a/alert_service.py b/alert_service.py index fc36b47..bd81af4 100644 --- a/alert_service.py +++ b/alert_service.py @@ -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() diff --git a/backend/bootstrap/container.py b/backend/bootstrap/container.py index 0060ad6..7b87880 100644 --- a/backend/bootstrap/container.py +++ b/backend/bootstrap/container.py @@ -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, diff --git a/backend/database/repositories/__init__.py b/backend/database/repositories/__init__.py new file mode 100644 index 0000000..d48c58c --- /dev/null +++ b/backend/database/repositories/__init__.py @@ -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", +] diff --git a/backend/database/repositories/ports.py b/backend/database/repositories/ports.py new file mode 100644 index 0000000..bfadf18 --- /dev/null +++ b/backend/database/repositories/ports.py @@ -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]]]: ... diff --git a/backend/database/repositories/sqlite.py b/backend/database/repositories/sqlite.py new file mode 100644 index 0000000..22e206e --- /dev/null +++ b/backend/database/repositories/sqlite.py @@ -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), + ) diff --git a/docs/governance/architecture-inventory.json b/docs/governance/architecture-inventory.json index b8f397d..8595f96 100644 --- a/docs/governance/architecture-inventory.json +++ b/docs/governance/architecture-inventory.json @@ -265,7 +265,7 @@ }, { "path": "server.py", - "bytes": 267835, + "bytes": 267824, "lines": 5947 }, { diff --git a/docs/governance/stage-09-repositories.md b/docs/governance/stage-09-repositories.md new file mode 100644 index 0000000..a4a105c --- /dev/null +++ b/docs/governance/stage-09-repositories.md @@ -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. diff --git a/server.py b/server.py index 2304088..0ac19f6 100644 --- a/server.py +++ b/server.py @@ -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]]: diff --git a/strategy_tracking.py b/strategy_tracking.py index 926545c..f20c646 100644 --- a/strategy_tracking.py +++ b/strategy_tracking.py @@ -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]] = {} diff --git a/tests/test_bootstrap_container.py b/tests/test_bootstrap_container.py index ca547f7..032e7c6 100644 --- a/tests/test_bootstrap_container.py +++ b/tests/test_bootstrap_container.py @@ -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) diff --git a/tests/test_repository_boundaries.py b/tests/test_repository_boundaries.py new file mode 100644 index 0000000..7427f4f --- /dev/null +++ b/tests/test_repository_boundaries.py @@ -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() diff --git a/trade_journal.py b/trade_journal.py index b5de4bf..30fb54e 100644 --- a/trade_journal.py +++ b/trade_journal.py @@ -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]