from __future__ import annotations import json import urllib.error import urllib.request from collections.abc import Callable from dataclasses import replace from datetime import datetime, timedelta from typing import Any from zoneinfo import ZoneInfo from backend.data.contracts import ( DataSource, DataUsage, ObservationMetadata, ProviderResult, SnapshotState, ) from backend.data.providers.base import ProviderError SHANGHAI = ZoneInfo("Asia/Shanghai") class TushareProvider: source = DataSource.TUSHARE url = "http://api.tushare.pro" def __init__(self, token: str | None | Callable[[], str | None], timeout: int = 20) -> None: self._token_provider = token if callable(token) else lambda: token self._timeout = timeout @property def configured(self) -> bool: return bool(self._token()) def calendar(self, start_date: str, end_date: str) -> ProviderResult: result = self._query( "trade_cal", {"exchange": "SSE", "start_date": _compact(start_date), "end_date": _compact(end_date)}, "cal_date,is_open,pretrade_date", unit="calendar_day", ) days = ( datetime.fromisoformat(end_date).date() - datetime.fromisoformat(start_date).date() ).days + 1 return ProviderResult( result.rows, replace(result.metadata, coverage=min(len(result.rows) / max(days, 1), 1)), ) def entities(self) -> ProviderResult: rows: list[dict[str, Any]] = [] for status in ("L", "P", "D"): result = self._query( "stock_basic", {"exchange": "", "list_status": status}, "ts_code,symbol,name,industry,list_status,list_date,delist_date", unit="entity", ) rows.extend(result.rows) return ProviderResult( tuple(rows), _metadata(self.source, "entity", min(len(rows) / 5300, 1)) ) def daily(self, entity_type: str, identifier: str, end_date: str) -> ProviderResult: api_name = "index_daily" if entity_type == "index" else "daily" if entity_type in {"sector", "theme"}: api_name = "ths_daily" end = datetime.strptime(_compact(end_date), "%Y%m%d") start = (end - timedelta(days=380)).strftime("%Y%m%d") return self._query( api_name, {"ts_code": identifier, "start_date": start, "end_date": end.strftime("%Y%m%d")}, "ts_code,trade_date,open,high,low,close,vol,amount,pct_chg", unit="yuan/share", adjustment="unadjusted", ) def minute(self, entity_type: str, identifier: str, trade_date: str) -> ProviderResult: if entity_type != "stock": raise ProviderError("Tushare minute charts only support stocks") date = _display(trade_date) return self._query( "stk_mins", { "ts_code": identifier, "freq": "1min", "start_date": f"{date} 09:30:00", "end_date": f"{date} 15:00:00", }, "ts_code,trade_time,open,high,low,close,vol,amount", unit="yuan/share", ) def _query( self, api_name: str, params: dict[str, Any], fields: str, *, unit: str, adjustment: str = "not_applicable", ) -> ProviderResult: if not self.configured: raise ProviderError("行情服务尚未配置") token = self._token() if not token: raise ProviderError("行情服务尚未配置") body = json.dumps( {"api_name": api_name, "token": token, "params": params, "fields": fields}, ensure_ascii=False, ).encode("utf-8") request = urllib.request.Request( self.url, data=body, headers={"Content-Type": "application/json", "User-Agent": "XiaobaiReview/2"}, method="POST", ) try: with urllib.request.urlopen(request, timeout=self._timeout) as response: payload = json.loads(response.read().decode("utf-8")) except (urllib.error.URLError, TimeoutError, OSError, json.JSONDecodeError) as exc: raise ProviderError("行情服务请求失败") from exc if payload.get("code") not in (None, 0): raise ProviderError(str(payload.get("msg") or "行情服务拒绝请求")) data = payload.get("data") or {} columns = data.get("fields") or [] rows = tuple(dict(zip(columns, item, strict=False)) for item in data.get("items") or []) return ProviderResult(rows, _metadata(self.source, unit, 1 if rows else 0, adjustment)) def _token(self) -> str: return str(self._token_provider() or "").strip() def _metadata( source: DataSource, unit: str, coverage: float, adjustment: str = "not_applicable" ) -> ObservationMetadata: return ObservationMetadata( source=source, observed_at=datetime.now(SHANGHAI), unit=unit, adjustment=adjustment, freshness_seconds=0, coverage=coverage, state=SnapshotState.ARCHIVE, usage=DataUsage.CALCULATION, ) def _compact(value: str) -> str: normalized = value.replace("-", "") if len(normalized) != 8 or not normalized.isdigit(): raise ProviderError("日期格式无效") return normalized def _display(value: str) -> str: compact = _compact(value) return f"{compact[:4]}-{compact[4:6]}-{compact[6:]}"