Files
xiaobaifupan/backend/database/repositories/sqlite.py
T

109 lines
3.8 KiB
Python

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),
)