refactor: enforce account scoped repository boundaries

This commit is contained in:
leefer
2026-07-29 17:54:45 +08:00
parent 6ac9571ca0
commit 367fe71fbf
12 changed files with 311 additions and 33 deletions
+17 -8
View File
@@ -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()
+7 -3
View File
@@ -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,
+21
View File
@@ -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",
]
+52
View File
@@ -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]]]: ...
+108
View File
@@ -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),
)
+1 -1
View File
@@ -265,7 +265,7 @@
},
{
"path": "server.py",
"bytes": 267835,
"bytes": 267824,
"lines": 5947
},
{
+27
View File
@@ -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.
+4 -4
View File
@@ -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]]:
+9 -9
View File
@@ -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]] = {}
+3 -3
View File
@@ -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)
+54
View File
@@ -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
View File
@@ -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]