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 ( ChartPoint, ChartSeries, DataSource, DataUsage, MarketEntity, ObservationMetadata, ProviderResult, SnapshotState, TradeContext, ) from backend.data.heaven import historical_payload, realtime_payload, should_use_realtime from backend.data.policy import DataSourcePolicy from backend.data.providers.base import MarketDataProvider, ProviderError from backend.data.quality import DataQualityError, require_quality from backend.data.repository import MarketRepository from backend.data.screener_gateway import assemble_screener_inputs from backend.database.connection import Database SHANGHAI = ZoneInfo("Asia/Shanghai") class MarketDataUnavailable(RuntimeError): pass class DataGateway: def __init__( self, database: Database, repository: MarketRepository, providers: tuple[MarketDataProvider, ...], policy: DataSourcePolicy, ) -> None: self._database = database self._repository = repository self._providers = {provider.source: provider for provider in providers} self._policy = policy def refresh_reference(self, now: datetime | None = None) -> dict[str, int | str]: clock = now or datetime.now(SHANGHAI) provider = self._provider(DataSource.TUSHARE) start = (clock.date() - timedelta(days=370)).isoformat() end = (clock.date() + timedelta(days=40)).isoformat() calendar = require_quality("calendar", provider.calendar(start, end)) entities = require_quality("entities", provider.entities()) observed_at = clock.isoformat(timespec="seconds") with self._database.transaction() as connection: self._repository.replace_calendar( connection, calendar.rows, calendar.metadata.source.value, observed_at ) self._repository.replace_stocks( connection, entities.rows, entities.metadata.source.value, observed_at ) return { "calendar_days": len(calendar.rows), "entities": len(entities.rows), "observed_at": observed_at, } def trade_context( self, requested_date: str | None = None, now: datetime | None = None ) -> TradeContext: clock = now or datetime.now(SHANGHAI) requested = _date(requested_date or clock.date().isoformat()) with self._database.read() as connection: summary = self._repository.latest_summary(connection, requested) if summary is None: return TradeContext( requested_date=requested, actual_date=None, previous_date=None, observed_at=None, state=None, carried_forward=False, message="等待管理员首次同步真实行情", ) actual = str(summary["trade_date"]) observed_at = datetime.fromisoformat(str(summary["observed_at"])) state = SnapshotState(str(summary["state"])) carried = actual != requested with self._database.read() as connection: dates = self._repository.open_dates(connection, actual, 2) return TradeContext( requested_date=requested, actual_date=actual, previous_date=dates[1] if len(dates) > 1 else None, observed_at=observed_at, state=state, carried_forward=carried, message="沿用最近真实收盘快照" if carried else "", ) def latest_available_date(self, now: datetime | None = None) -> str: context = self.trade_context(now=now) return context.actual_date or (now or datetime.now(SHANGHAI)).date().isoformat() def summary(self, requested_date: str | None = None) -> dict[str, Any]: context = self.trade_context(requested_date) if context.actual_date is None: return {"context": context, "values": None} with self._database.read() as connection: row = self._repository.latest_summary(connection, context.actual_date) return {"context": context, "values": json.loads(str(row["payload_json"])) if row else None} def snapshot_inputs(self, trade_date: str, previous_trade_date: str) -> dict[str, Any]: provider = self._provider(DataSource.TUSHARE) self._policy.assert_allowed(provider.source, DataUsage.CALCULATION) return provider.snapshot_inputs(trade_date, previous_trade_date) def trading_dates(self, through: str, limit: int = 2) -> tuple[str, ...]: requested = _date(through) with self._database.read() as connection: return self._repository.open_dates(connection, requested, limit) def stock_directory(self) -> dict[str, dict[str, Any]]: with self._database.read() as connection: rows = self._repository.stock_directory(connection) return {str(row["identifier"]): dict(row) for row in rows} def insight_inputs( self, kind: str, trade_date: str, previous_trade_date: str = "", identifier: str = "", ) -> dict[str, Any]: provider = self._provider(DataSource.TUSHARE) self._policy.assert_allowed(provider.source, DataUsage.CALCULATION) return provider.market_insight(kind, trade_date, previous_trade_date, identifier) def screener_inputs( self, trade_date: str, history_days: int = 260 ) -> tuple[dict[str, Any], dict[str, float], list[str]]: dates = self.trading_dates(trade_date, history_days) if len(dates) < 21: raise MarketDataUnavailable("历史交易日不足21日,无法生成选股因子") chronological = tuple(reversed(dates)) provider = self._provider(DataSource.TUSHARE) self._policy.assert_allowed(provider.source, DataUsage.CALCULATION) raw = provider.screener_inputs(chronological) inputs, coverage = assemble_screener_inputs( self._database, self._repository, trade_date, chronological, raw ) return inputs, coverage, [provider.source.value, DataSource.LOCAL.value] def dynamic_auction( self, identifiers: tuple[str, ...], start_time: str, end_time: str ) -> ProviderResult: if not identifiers: raise MarketDataUnavailable("动态竞价候选范围为空") provider = self._provider(DataSource.IFIND) self._policy.assert_allowed(provider.source, DataUsage.CALCULATION) result = provider.realtime_snapshots(identifiers, start_time, end_time) if not result.rows: raise MarketDataUnavailable("当前动态竞价快照暂不可用") return result def sector_members( self, trade_date: str, sector_name: str, representative: str ) -> dict[str, Any]: with self._database.read() as connection: cached = self._repository.sector_members(connection, trade_date, sector_name) if cached: return json.loads(str(cached["payload_json"])) provider = self._provider(DataSource.TUSHARE) self._policy.assert_allowed(provider.source, DataUsage.CALCULATION) result = provider.sector_members(representative, trade_date) if not result.rows: raise MarketDataUnavailable("该板块暂无可核验的申万成分股") rows = sorted( (dict(row) for row in result.rows), key=lambda row: ( bool(row.get("quoted")), _number(row.get("change")), _number(row.get("amount")), ), reverse=True, ) payload = { "trade_date": trade_date, "sector_name": str(rows[0].get("sector_name") or sector_name), "sector_code": str(rows[0].get("sector_code") or ""), "member_count": len(rows), "quoted_count": sum(bool(row.get("quoted")) for row in rows), "coverage": round(result.metadata.coverage, 4), "items": [ { "identifier": str(row.get("ts_code") or ""), "code": str(row.get("ts_code") or "").split(".")[0], "name": str(row.get("name") or ""), "change": row.get("change"), "open": row.get("open"), "close": row.get("close"), "amount": row.get("amount"), "quoted": bool(row.get("quoted")), } for row in rows ], } with self._database.transaction() as connection: self._repository.save_sector_members( connection, trade_date=trade_date, sector_name=sector_name, sector_code=payload["sector_code"], observed_at=result.metadata.observed_at.isoformat(timespec="seconds"), source=result.metadata.source.value, coverage=result.metadata.coverage, payload=payload, ) return payload def heaven_trend_inputs( self, query: str, requested_date: str, now: datetime | None = None ) -> dict[str, Any]: clock = now or datetime.now(SHANGHAI) context = self.trade_context(requested_date, clock) if context.actual_date is None: raise MarketDataUnavailable("等待管理员首次同步真实收盘行情") stock = self._resolve_stock_query(query) provider = self._provider(DataSource.TUSHARE) self._policy.assert_allowed(provider.source, DataUsage.CALCULATION) if should_use_realtime(requested_date, context, clock): dates = self.trading_dates(requested_date, 2) if len(dates) < 2 or dates[0] != requested_date: raise MarketDataUnavailable("目标日期不是有效交易日") raw = provider.heaven_realtime_inputs(stock.identifier, dates[0], dates[1]) return realtime_payload( self._database, self._repository, stock, dates[0], dates[1], raw, clock, ) return historical_payload( self._database, self._repository, stock, context.actual_date, provider ) def search(self, query: str) -> tuple[MarketEntity, ...]: with self._database.read() as connection: return self._repository.search(connection, query) def chart( self, entity_type: str, identifier: str, interval: str, now: datetime | None = None, ) -> ChartSeries: if entity_type not in {"stock", "sector", "theme", "index"}: raise MarketDataUnavailable("不支持的行情标的类型") if interval not in {"day", "minute"}: raise MarketDataUnavailable("不支持的行情周期") entity = self._resolve_entity(entity_type, identifier) clock = now or datetime.now(SHANGHAI) with self._database.read() as connection: cached = self._repository.chart(connection, entity_type, entity.identifier, interval) if cached and not _chart_stale(cached, interval, clock): return _stored_chart(entity, cached) try: series = self._load_chart(entity, interval, clock) except (ProviderError, DataQualityError) as exc: if cached: return _stored_chart(entity, cached) raise MarketDataUnavailable("当前没有可用的真实行情数据") from exc self._save_chart(series) return series def _load_chart(self, entity: MarketEntity, interval: str, clock: datetime) -> ChartSeries: dataset = "daily_chart" if interval == "day" else "minute_chart" target_dates = self._chart_dates(clock) errors: list[Exception] = [] for source in self._policy.candidates(dataset, DataUsage.DISPLAY): provider = self._providers.get(source) if provider is None or not provider.configured: continue self._policy.assert_allowed(source, DataUsage.DISPLAY) for target in target_dates: try: raw = ( provider.daily(entity.entity_type, entity.identifier, target) if interval == "day" else provider.minute(entity.entity_type, entity.identifier, target) ) require_quality(dataset, raw) normalized = _normalize_chart(entity, interval, raw.rows, raw.metadata, clock) if normalized.points: return normalized except (ProviderError, DataQualityError, ValueError) as exc: errors.append(exc) if interval == "day": break raise MarketDataUnavailable("当前没有可用的真实行情数据") from ( errors[-1] if errors else None ) def _chart_dates(self, clock: datetime) -> tuple[str, ...]: today = clock.date().isoformat() with self._database.read() as connection: dates = self._repository.open_dates(connection, today, 8) if dates: return dates return tuple((clock.date() - timedelta(days=offset)).isoformat() for offset in range(8)) def _resolve_entity(self, entity_type: str, identifier: str) -> MarketEntity: normalized = identifier.strip().upper() with self._database.read() as connection: entity = self._repository.entity(connection, entity_type, normalized) if entity is None and entity_type == "stock" and normalized.isdigit(): matches = self._repository.search(connection, normalized, 4) entity = next( ( item for item in matches if item.entity_type == "stock" and item.code == normalized ), None, ) if entity: return entity if entity_type == "stock" and len(normalized) == 6 and normalized.isdigit(): suffix = ( "BJ" if normalized.startswith(("4", "8", "9")) else "SH" if normalized.startswith("6") else "SZ" ) return MarketEntity("stock", f"{normalized}.{suffix}", normalized, normalized) raise MarketDataUnavailable("未找到该行情标的") def _resolve_stock_query(self, query: str) -> MarketEntity: normalized = query.strip() if not normalized: raise MarketDataUnavailable("请输入股票代码或股票名称") with self._database.read() as connection: matches = tuple( item for item in self._repository.search(connection, normalized, 16) if item.entity_type == "stock" ) exact = [ item for item in matches if item.code.casefold() == normalized.casefold() or item.identifier.casefold() == normalized.casefold() or item.name.casefold() == normalized.casefold() ] if len(exact) == 1: return exact[0] if len(exact) > 1: raise MarketDataUnavailable("股票名称存在重名,请输入六位代码") raise MarketDataUnavailable("未找到该股票,请检查代码或名称") def _save_chart(self, series: ChartSeries) -> None: payload = { "previous_close": series.previous_close, "points": [ { "time": point.time, "open": point.open, "high": point.high, "low": point.low, "close": point.close, "volume": point.volume, "amount": point.amount, "average": point.average, } for point in series.points ], } with self._database.transaction() as connection: self._repository.save_chart( connection, entity_type=series.entity.entity_type, identifier=series.entity.identifier, interval=series.interval, trade_date=series.trade_date, observed_at=series.metadata.observed_at.isoformat(timespec="seconds"), source=series.metadata.source.value, usage=series.metadata.usage.value, adjustment=series.metadata.adjustment, coverage=series.metadata.coverage, payload=payload, ) def _provider(self, source: DataSource) -> MarketDataProvider: provider = self._providers.get(source) if provider is None or not provider.configured: raise MarketDataUnavailable("所需行情服务尚未配置") return provider def _normalize_chart( entity: MarketEntity, interval: str, rows: tuple[dict[str, Any], ...], metadata: ObservationMetadata, clock: datetime, ) -> ChartSeries: parsed: list[tuple[str, ChartPoint, float | None]] = [] for row in rows: stamp = str(row.get("trade_date") or row.get("time") or row.get("trade_time") or "") trade_date = _row_date(stamp) point_time = trade_date if interval == "day" else _row_time(stamp) close = _number(row.get("close")) open_price = _number(row.get("open")) high = _number(row.get("high")) low = _number(row.get("low")) if not trade_date or close <= 0 or open_price <= 0 or high <= 0 or low <= 0: continue if interval == "minute" and not "09:30" <= point_time <= "15:00": continue volume = _number(row.get("volume", row.get("vol"))) amount = _number(row.get("amount")) if metadata.source is DataSource.TUSHARE: volume *= 100 amount *= 1000 parsed.append( ( trade_date, ChartPoint( time=point_time, open=open_price, high=high, low=low, close=close, volume=volume, amount=amount, average=_optional_number(row.get("avgPrice", row.get("average"))), ), _optional_number(row.get("preClose", row.get("pre_close"))), ) ) parsed.sort(key=lambda item: (item[0], item[1].time)) if not parsed: raise DataQualityError("chart contains no valid points") if interval == "minute": latest = parsed[-1][0] parsed = [item for item in parsed if item[0] == latest] else: today = clock.date().isoformat() if parsed[-1][0] == today and not _valid_today_bar(parsed[-1][1], clock): parsed.pop() if not parsed: raise DataQualityError("chart contains no completed bar") parsed = parsed[-90:] previous = parsed[0][2] if previous is None and interval == "day" and len(parsed) > 1: previous = parsed[-2][1].close return ChartSeries( entity=entity, interval=interval, trade_date=parsed[-1][0], previous_close=previous, points=tuple(item[1] for item in parsed), metadata=metadata, ) def _valid_today_bar(point: ChartPoint, clock: datetime) -> bool: if clock.time() < time(9, 30): return False return ( point.volume > 0 and point.amount > 0 and point.high >= max(point.open, point.close) and point.low <= min(point.open, point.close) ) def _chart_stale(row: Any, interval: str, clock: datetime) -> bool: observed = datetime.fromisoformat(str(row["observed_at"])) if interval == "day": return row["trade_date"] < clock.date().isoformat() and clock.time() >= time(15, 5) return (clock - observed.astimezone(SHANGHAI)).total_seconds() > 30 def _stored_chart(entity: MarketEntity, row: Any) -> ChartSeries: payload = json.loads(str(row["payload_json"])) points = tuple(ChartPoint(**point) for point in payload.get("points") or []) metadata = ObservationMetadata( source=DataSource(str(row["source"])), observed_at=datetime.fromisoformat(str(row["observed_at"])), unit="yuan/share", adjustment=str(row["adjustment"]), freshness_seconds=0, coverage=float(row["coverage"]), state=SnapshotState.ARCHIVE, usage=DataUsage(str(row["usage"])), ) return ChartSeries( entity, str(row["interval"]), str(row["trade_date"]), payload.get("previous_close"), points, metadata, ) def _date(value: str) -> str: try: return date.fromisoformat(value).isoformat() except ValueError as exc: raise MarketDataUnavailable("日期格式无效") from exc def _row_date(value: str) -> str: compact = value[:10].replace("-", "") if len(compact) != 8 or not compact.isdigit(): return "" return f"{compact[:4]}-{compact[4:6]}-{compact[6:]}" def _row_time(value: str) -> str: if " " in value: return value.split(" ", 1)[1][:5] return value[-8:-3] if len(value) >= 8 else value[:5] def _number(value: Any) -> float: try: return float(value or 0) except (TypeError, ValueError): return 0.0 def _optional_number(value: Any) -> float | None: if value is None or value == "": return None try: number = float(value) except (TypeError, ValueError): return None return number if number == number else None