from __future__ import annotations import json from datetime import date, datetime, time, timedelta from typing import Any from zoneinfo import ZoneInfo from backend.data.contracts import ProviderResult, SnapshotState from backend.data.providers.base import ProviderError from backend.data.providers.tushare import TushareProvider from backend.data.repository import MarketRepository from backend.database.connection import Database from backend.features.market.sentiment import calculate_sentiment from backend.features.market.snapshot import build_snapshot SHANGHAI = ZoneInfo("Asia/Shanghai") class SnapshotSyncError(RuntimeError): pass class MarketSnapshotService: def __init__( self, database: Database, repository: MarketRepository, provider: TushareProvider ) -> None: self._database = database self._repository = repository self._provider = provider def sync( self, requested_date: str | None = None, now: datetime | None = None ) -> dict[str, Any]: clock = now or datetime.now(SHANGHAI) requested = _date(requested_date or clock.date().isoformat()) through = requested if requested == clock.date().isoformat() and clock.time() < time(9, 15): through = (clock.date() - timedelta(days=1)).isoformat() elif requested == clock.date().isoformat() and clock.time() < time(15, 10): raise SnapshotSyncError("盘中快照任务尚未开放,请保留最近真实收盘数据") with self._database.read() as connection: dates = self._repository.open_dates(connection, through, 2) active_count = self._repository.active_stock_count(connection) if len(dates) < 2: raise SnapshotSyncError("请先同步完整交易日历") if active_count <= 0: raise SnapshotSyncError("请先同步股票目录") trade_date, previous_date = dates[0], dates[1] try: inputs = self._provider.snapshot_inputs(trade_date, previous_date) except ProviderError as exc: raise SnapshotSyncError("收盘行情读取失败,已保留原有快照") from exc daily = inputs.get("daily") if not isinstance(daily, ProviderResult): raise SnapshotSyncError("收盘日线缺失,已保留原有快照") coverage = len(daily.rows) / active_count if coverage < 0.98: raise SnapshotSyncError(f"收盘日线覆盖率仅{coverage * 100:.1f}%,未写入不完整快照") for key in ("limit_up", "limit_down", "broken", "previous_limit_up", "price_limits"): value = inputs.get(key) if not isinstance(value, ProviderResult) or value.metadata.coverage < 1: raise SnapshotSyncError("涨跌停事件数据不完整,已保留原有快照") snapshot = build_snapshot(trade_date, previous_date, inputs) with self._database.read() as connection: rows = self._repository.summaries(connection, previous_date, 250) history = [json.loads(str(row["payload_json"])) for row in rows] sentiment = calculate_sentiment(snapshot, history) snapshot["sentiment"] = sentiment snapshot.update(snapshot["overview"]) snapshot["temperature"] = sentiment["score"] observed_at = datetime.combine( date.fromisoformat(trade_date), time(15), tzinfo=SHANGHAI ).isoformat(timespec="seconds") state = ( SnapshotState.FINAL if trade_date == clock.date().isoformat() else SnapshotState.ARCHIVE ) with self._database.transaction() as connection: self._repository.save_summary( connection, trade_date=trade_date, observed_at=observed_at, state=state.value, source="tushare", coverage=min(coverage, 1), payload=snapshot, ) return { "trade_date": trade_date, "observed_at": observed_at, "coverage": round(min(coverage, 1), 4), "stocks": len(daily.rows), "limit_up": len(snapshot["limits"]), "limit_down": len(snapshot["down_limits"]), "broken": len(snapshot["broken"]), "temperature": sentiment["score"], } def workspace(self, key: str, requested_date: str | None = None) -> dict[str, Any]: requested = _date(requested_date or datetime.now(SHANGHAI).date().isoformat()) with self._database.read() as connection: row = self._repository.latest_summary(connection, requested) history_rows = self._repository.summaries(connection, requested, 60) if row is None: return {"trade_date": None, "message": "等待管理员首次同步真实收盘行情"} payload = json.loads(str(row["payload_json"])) response: dict[str, Any] = { "trade_date": str(row["trade_date"]), "observed_at": str(row["observed_at"]), "carried_forward": str(row["trade_date"]) != requested, "message": "沿用最近真实收盘快照" if str(row["trade_date"]) != requested else "", "overview": payload.get("overview") or {}, } if key == "emotion": response["sentiment"] = payload.get("sentiment") or {} response["history"] = [ _history_item(json.loads(str(item["payload_json"]))) for item in history_rows ] elif key == "pool": response["items"] = payload.get("limits") or [] elif key == "broken": response["items"] = payload.get("broken") or [] elif key == "limit-down": response["items"] = payload.get("down_limits") or [] elif key == "yesterday": response["items"] = payload.get("yesterday_limits") or [] elif key == "performance": response["items"] = payload.get("limit_performance") or [] else: raise SnapshotSyncError("不支持的市场工作区") return response def _history_item(payload: dict[str, Any]) -> dict[str, Any]: sentiment = payload.get("sentiment") or {} stats = sentiment.get("stats") or {} overview = payload.get("overview") or {} return { "trade_date": payload.get("trade_date"), "temperature": sentiment.get("score"), "phase": sentiment.get("phase"), "direction": sentiment.get("direction"), "positive_rate": stats.get("positive_rate"), "seal_rate": overview.get("seal_rate"), "limit_up": overview.get("limit_up"), "broken": overview.get("broken"), "limit_down": overview.get("limit_down"), "max_height": stats.get("max_height"), "amount": overview.get("amount"), } def _date(value: str) -> str: try: return date.fromisoformat(value).isoformat() except ValueError as exc: raise SnapshotSyncError("日期格式无效") from exc