from __future__ import annotations from typing import Any from backend.data.gateway import DataGateway, MarketDataUnavailable from backend.data.providers.base import ProviderError from backend.data.quality import DataQualityError from backend.features.market.sync import MarketSnapshotService, SnapshotSyncError from backend.http.errors import AppError class MarketService: def __init__(self, gateway: DataGateway, snapshots: MarketSnapshotService) -> None: self._gateway = gateway self._snapshots = snapshots def context(self, requested_date: str | None = None) -> dict[str, Any]: context = self._call(self._gateway.trade_context, requested_date) return _context(context) def summary(self, requested_date: str | None = None) -> dict[str, Any]: result = self._call(self._gateway.summary, requested_date) return {"context": _context(result["context"]), "values": result["values"]} def search(self, query: str) -> dict[str, Any]: normalized = " ".join(query.split()) items = self._gateway.search(normalized) if normalized else () labels = {"stock": "股票", "sector": "板块", "theme": "题材", "index": "指数"} groups = [] for entity_type in ("stock", "sector", "theme", "index"): groups.append( { "entity_type": entity_type, "label": labels[entity_type], "items": [ { "entity_type": item.entity_type, "identifier": item.identifier, "code": item.code, "name": item.name, "sector": item.sector, } for item in items if item.entity_type == entity_type ], } ) return {"query": normalized, "groups": groups} def chart(self, entity_type: str, identifier: str, interval: str) -> dict[str, Any]: series = self._call(self._gateway.chart, entity_type, identifier, interval) return { "entity_type": series.entity.entity_type, "identifier": series.entity.identifier, "code": series.entity.code, "name": series.entity.name, "interval": series.interval, "trade_date": series.trade_date, "observed_at": series.metadata.observed_at, "previous_close": series.previous_close, "range_start": "09:30" if interval == "minute" else None, "range_end": "15:00" if interval == "minute" else None, "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 ], } def refresh_reference(self) -> dict[str, int | str]: return self._call(self._gateway.refresh_reference) def sync_snapshot(self, requested_date: str | None = None) -> dict[str, Any]: return self._call(self._snapshots.sync, requested_date) def workspace(self, key: str, requested_date: str | None = None) -> dict[str, Any]: return self._call(self._snapshots.workspace, key, requested_date) def rotation_members( self, sector_name: str, requested_date: str | None = None ) -> dict[str, Any]: trade_date, representative = self._call( self._snapshots.rotation_member_target, requested_date, sector_name ) return self._call( self._gateway.sector_members, trade_date, sector_name, representative ) @staticmethod def _call(function, *args): try: return function(*args) except (MarketDataUnavailable, ProviderError, DataQualityError, SnapshotSyncError) as exc: raise AppError("market_data_unavailable", str(exc), 503) from exc def _context(context) -> dict[str, Any]: return { "requested_date": context.requested_date, "actual_date": context.actual_date, "previous_date": context.previous_date, "observed_at": context.observed_at, "state": context.state.value if context.state else None, "carried_forward": context.carried_forward, "message": context.message, }