diff --git a/backend/data/datahub/bridge.py b/backend/data/datahub/bridge.py index eb96ee0..0cb8bae 100644 --- a/backend/data/datahub/bridge.py +++ b/backend/data/datahub/bridge.py @@ -19,6 +19,7 @@ from backend.data.datahub.redact import redact_text, redact_value from backend.data.datahub.route_state import LEDGER from backend.data.datahub.settings import DatahubSettings from backend.data.providers.tushare_client import TushareClient +from backend.data.providers.tushare_transport import TushareError LOGGER = logging.getLogger("xiaobai.datahub") ShadowSink = Callable[[dict[str, Any]], None] @@ -171,6 +172,45 @@ class DatahubBridge: self._log_failure("index_quotes", exc) return None + def try_sector_quote(self, code: str, trade_date: str = "") -> dict[str, Any] | None: + flags = self.settings.flags("quotes") + if not flags.read: + return None + try: + response = self.client.sector_quote(code, trade_date) + data = response.data + if not isinstance(data, dict) or not data: + raise DatahubError("EMPTY", "datahub sector quote empty") + row = dict(data) + if (response.meta or {}).get("stale"): + row["delayed"] = True + row["delay_seconds"] = int((response.meta or {}).get("staleness_seconds") or 0) + row["delay_notice"] = str((response.meta or {}).get("delay_notice") or "") + self._record_route("quotes", "datahub", str((response.meta or {}).get("source") or "datahub")) + return row + except Exception as exc: + self._log_failure("quotes", exc) + return None + + def try_limit_pool(self, trade_date: str = "") -> list[dict[str, Any]] | None: + flags = self.settings.flags("limit_events") + if not flags.read: + return None + try: + response = self.client.limit_pool(trade_date) + rows = [dict(item) for item in (response.data or []) if isinstance(item, dict)] + if not rows: + raise DatahubError("EMPTY", "datahub limit pool empty") + self._record_route( + "limit_events", + "datahub", + str((response.meta or {}).get("source") or "datahub"), + ) + return rows + except Exception as exc: + self._log_failure("limit_events", exc) + return None + def try_daily_chart( self, code: str, @@ -191,6 +231,11 @@ class DatahubBridge: self.client.index_bars, {"code": code, "from": start, "to": compact_end}, ) + elif dataset == "sector_daily": + response = self._paginate( + self.client.sectors, + {"code": code, "from": start, "to": compact_end}, + ) else: response = self._paginate( self.client.daily_bars, @@ -263,53 +308,33 @@ class DatahubBridge: fields: str, legacy_query: Callable[..., list[dict[str, Any]]], ) -> list[dict[str, Any]]: + del legacy_query # 主网站不再直连 Tushare;调度全部由数据中枢完成。 + if api_name == "rt_sw_k": + raise TushareError("rt_sw_k is disabled; use published sw_daily or free Shenwan realtime") dataset = API_TO_DATASET.get(api_name) - # 问天按实际数据依赖接入:已映射到 hub 的 API 跟随开关;未映射的继续旧链路。 - if not dataset: - return legacy_query(api_name, params, fields) - flags = self.settings.flags(dataset) - if not flags.read and not flags.shadow: - return legacy_query(api_name, params, fields) - - hub_rows: list[dict[str, Any]] | None = None - hub_meta: dict[str, Any] = {} - hub_error: str | None = None - hub_canonical: list[dict[str, Any]] = [] - try: - response = self._fetch_dataset(dataset, params or {}, api_name=api_name) - hub_canonical = self._extract_rows(dataset, response, params or {}) - hub_rows = to_native_rows(dataset, hub_canonical) - hub_meta = dict(response.meta) - self._validate_usable(dataset, hub_rows, response) - except Exception as exc: - hub_error = self._error_text(exc) - self._log_failure(dataset, exc) - - if flags.shadow: - try: - legacy_rows = legacy_query(api_name, params, fields) - except Exception as exc: - if flags.read and hub_rows is not None and hub_error is None: - self._emit_shadow( - compare_rows(dataset, [], hub_canonical, hub_meta, self._error_text(exc), fields) - ) - return project_fields(hub_rows, fields) - raise - self._emit_shadow(compare_rows(dataset, legacy_rows, hub_canonical, hub_meta, hub_error, fields)) - if flags.read and hub_rows is not None and hub_error is None: - self._record_route(dataset, "datahub", str(hub_meta.get("source") or "datahub")) - return project_fields(hub_rows, fields) + if dataset: + flags = self.settings.flags(dataset) if flags.read: - self._record_route(dataset, "legacy", "tushare", hub_error or "") - return legacy_rows - - if flags.read and hub_rows is not None and hub_error is None: - self._record_route(dataset, "datahub", str(hub_meta.get("source") or "datahub")) - return project_fields(hub_rows, fields) - result = legacy_query(api_name, params, fields) - if flags.read: - self._record_route(dataset, "legacy", "tushare", hub_error or "") - return result + try: + response = self._fetch_dataset(dataset, params or {}, api_name=api_name) + hub_canonical = self._extract_rows(dataset, response, params or {}) + hub_rows = to_native_rows(dataset, hub_canonical) + self._validate_usable(dataset, hub_rows, response) + self._record_route(dataset, "datahub", str(response.meta.get("source") or "datahub")) + return project_fields(hub_rows, fields) + except Exception as exc: + self._log_failure(dataset, exc) + try: + response = self.client.query_api(api_name, params or {}, fields) + rows = [dict(item) for item in (response.data or []) if isinstance(item, dict)] + if dataset: + self._record_route(dataset, "datahub", str((response.meta or {}).get("source") or "datahub")) + else: + self._record_route(api_name, "datahub", str((response.meta or {}).get("source") or "datahub")) + return rows if not fields else project_fields(rows, fields) + except Exception as exc: + self._log_failure(dataset or api_name, exc) + raise TushareError(self._error_text(exc)) from exc def _fetch_dataset(self, dataset: str, params: dict[str, Any], api_name: str = "") -> DatahubResponse: date = yyyymmdd(params.get("trade_date") or params.get("date")) @@ -429,8 +454,8 @@ class DatahubBridge: def _log_failure(self, dataset: str, exc: Exception) -> None: error = redact_text(self._error_text(exc), self.settings.secrets()) - LOGGER.warning("datahub fallback dataset=%s error=%s", dataset, error) - self._record_route(dataset, "legacy", "pending-legacy", error) + LOGGER.warning("datahub unavailable dataset=%s error=%s", dataset, error) + self._record_route(dataset, "datahub", "unavailable", error) def _record_route(self, dataset: str, route: str, source: str = "", error: str = "") -> None: LEDGER.record(dataset, route, source, redact_text(error, self.settings.secrets())) @@ -529,6 +554,8 @@ class DatahubAwareTushareClient: legacy.try_market_quotes = self.try_market_quotes legacy.try_quotes = self.try_quotes legacy.try_index_quotes = self.try_index_quotes + legacy.try_sector_quote = self.try_sector_quote + legacy.try_limit_pool = self.try_limit_pool legacy.record_datahub_legacy = self.record_datahub_legacy def query( @@ -548,6 +575,12 @@ class DatahubAwareTushareClient: def try_index_quotes(self) -> list[dict[str, Any]] | None: return self._bridge.try_index_quotes() + def try_sector_quote(self, code: str, trade_date: str = "") -> dict[str, Any] | None: + return self._bridge.try_sector_quote(code, trade_date) + + def try_limit_pool(self, trade_date: str = "") -> list[dict[str, Any]] | None: + return self._bridge.try_limit_pool(trade_date) + def record_datahub_legacy(self, dataset: str, source: str = "", error: str = "") -> None: self._bridge.record_legacy(dataset, source, error) diff --git a/backend/data/datahub/client.py b/backend/data/datahub/client.py index 5b8d737..5d04aa1 100644 --- a/backend/data/datahub/client.py +++ b/backend/data/datahub/client.py @@ -90,6 +90,24 @@ class DatahubClient: params["dataset"] = dataset return self.get("/v1/batches", params) + def query_api(self, api_name: str, params: dict[str, Any] | None = None, fields: str = "") -> DatahubResponse: + return self.post( + "/v1/query", + {"api_name": api_name, "params": params or {}, "fields": fields}, + ) + + def sector_quote(self, code: str, date: str = "") -> DatahubResponse: + payload: dict[str, Any] = {"code": code} + if date: + payload["date"] = date + return self.get("/v1/sectors/quote", payload) + + def limit_pool(self, trade_date: str = "") -> DatahubResponse: + params: dict[str, Any] = {} + if trade_date: + params["date"] = trade_date + return self.get("/v1/limit-pool", params) + def get(self, path: str, params: dict[str, Any] | None = None) -> DatahubResponse: if not self.settings.token: raise DatahubError("NOT_CONFIGURED", "DATAHUB_TOKEN is not configured") @@ -118,15 +136,41 @@ class DatahubClient: ) raise last_error or DatahubError("INTERNAL", "datahub request failed") - def _request(self, url: str) -> DatahubResponse: + def post(self, path: str, body: dict[str, Any] | None = None) -> DatahubResponse: + if not self.settings.token: + raise DatahubError("NOT_CONFIGURED", "DATAHUB_TOKEN is not configured") + url = self.settings.base_url + path + attempts = 1 + max(0, self.settings.retries) + last_error: DatahubError | None = None + payload = json.dumps(body or {}, ensure_ascii=False).encode("utf-8") + for attempt in range(attempts): + try: + return self._request(url, method="POST", data=payload) + except DatahubError as exc: + last_error = exc + if exc.code not in {"TIMEOUT", "UNAVAILABLE"} or attempt + 1 >= attempts: + raise + LOGGER.warning( + "datahub retry %s/%s %s", + attempt + 1, + attempts, + redact_text(str(exc), self.settings.secrets()), + ) + raise last_error or DatahubError("INTERNAL", "datahub request failed") + + def _request(self, url: str, method: str = "GET", data: bytes | None = None) -> DatahubResponse: + headers = { + "Accept": "application/json", + "X-Datahub-Token": self.settings.token, + "User-Agent": "XiaobaiReviewDatahub/1.0", + } + if data is not None: + headers["Content-Type"] = "application/json" request = urllib.request.Request( url, - headers={ - "Accept": "application/json", - "X-Datahub-Token": self.settings.token, - "User-Agent": "XiaobaiReviewDatahub/1.0", - }, - method="GET", + data=data, + headers=headers, + method=method, ) try: with self._urlopen(request, timeout=self.settings.timeout_seconds) as response: diff --git a/backend/data/gateway.py b/backend/data/gateway.py index 7cfd65d..8181655 100644 --- a/backend/data/gateway.py +++ b/backend/data/gateway.py @@ -38,7 +38,7 @@ class DataGateway: if dataset_id: self.policy.assert_allowed(dataset_id, "tushare", usage) legacy = self.tushare_provider.client() - legacy.realtime_aggregator = self.realtime_observer + legacy.realtime_aggregator = None return DatahubAwareTushareClient(legacy, self.datahub) def dataset_status(self, trade_date: str) -> list[dict[str, Any]] | None: diff --git a/backend/data/providers/tushare_daily.py b/backend/data/providers/tushare_daily.py index e143e10..704b2a7 100644 --- a/backend/data/providers/tushare_daily.py +++ b/backend/data/providers/tushare_daily.py @@ -186,12 +186,11 @@ class DailyMarketMixin: return mapped def _free_board_map(self, trade_date: str) -> dict[str, dict[str, Any]]: - aggregator = getattr(self, "realtime_aggregator", None) - loader = getattr(aggregator, "eastmoney_limit_pool", None) if aggregator else None + loader = getattr(self, "try_limit_pool", None) if not callable(loader): return {} try: - rows = loader(trade_date) + rows = loader(trade_date) or [] except Exception: return {} return { diff --git a/backend/data/providers/tushare_dashboard.py b/backend/data/providers/tushare_dashboard.py index 2e9130c..468d1e6 100644 --- a/backend/data/providers/tushare_dashboard.py +++ b/backend/data/providers/tushare_dashboard.py @@ -251,27 +251,23 @@ class DashboardMixin: quotes = hub(trade_date) if quotes: return list(quotes), "datahub" - rt_error = "" + named = getattr(self, "try_quotes", None) + code_list = [item for item in str(codes or "").split(",") if item] + if callable(named) and code_list: + collected: list[dict[str, Any]] = [] + for index in range(0, len(code_list), 60): + collected.extend(named(code_list[index:index + 60]) or []) + if collected: + delayed = any(item.get("delayed") for item in collected) + return collected, "datahub_delayed" if delayed else "datahub" try: quotes = self.query("rt_k", {"ts_code": codes}) if quotes: - self._mark_quote_legacy("tushare_rt_k", rt_error) - return list(quotes), "tushare_rt_k" - rt_error = f"No realtime data returned for {trade_date}" + delayed = any(item.get("delayed") for item in quotes) + return list(quotes), "datahub_delayed" if delayed else "datahub" except TushareError as exc: - rt_error = str(exc) - try: - quotes, quote_source = self._free_realtime_quotes(trade_date, codes) - except Exception as exc: - raise TushareError( - f"当天盘中实时行情不可用:rt_k={rt_error};免费源={exc}" - ) from exc - if not quotes: - raise TushareError( - f"当天盘中实时行情不可用:rt_k={rt_error};免费源=empty" - ) - self._mark_quote_legacy(quote_source, rt_error) - return quotes, quote_source + raise TushareError(f"当天盘中实时行情不可用:{exc}") from exc + raise TushareError("当天盘中实时行情不可用:数据中枢未返回可用行情") def _mark_quote_legacy(self, source: str, error: str = "") -> None: marker = getattr(self, "record_datahub_legacy", None) @@ -283,27 +279,8 @@ class DashboardMixin: trade_date: str, codes: str = "", ) -> tuple[list[dict[str, Any]], str]: - aggregator = self._realtime_aggregator() - last_error = "" - try: - quotes = aggregator.eastmoney_market_quotes(expected_date=trade_date) - if quotes: - return quotes, "eastmoney_clist" - except Exception as exc: - last_error = str(exc) - code_list = [item for item in str(codes or "").split(",") if item] - try: - if code_list: - quotes = aggregator.tencent_stock_quotes(code_list, expected_date=trade_date) - else: - quotes = aggregator.tencent_market_quotes(code_list, expected_date=trade_date) - except Exception as exc: - raise TushareError( - f"eastmoney={last_error or 'empty'};tencent={exc}" - ) from exc - if not quotes: - raise TushareError(f"eastmoney={last_error or 'empty'};tencent=empty") - return quotes, "tencent_qt" + del trade_date, codes + raise TushareError("主网站不再直连免费行情源,请走数据中枢") def _free_realtime_indices(self) -> list[dict[str, Any]]: hub = getattr(self, "try_index_quotes", None) @@ -312,14 +289,7 @@ class DashboardMixin: converted = [item for item in (_hub_index_quote(row) for row in rows or []) if item] if converted: return converted - try: - rows = self._realtime_aggregator().eastmoney_indices() - marker = getattr(self, "record_datahub_legacy", None) - if callable(marker): - marker("index_quotes", "eastmoney_push2") - return rows - except Exception: - return [] + return [] def _load_realtime_reference( self, @@ -470,21 +440,6 @@ class DashboardMixin: return dict(rows[0]) except TushareError: pass - aggregator = getattr(self, "realtime_aggregator", None) - if aggregator is None: - return {} - for loader in ( - getattr(aggregator, "eastmoney_stock_quote", None), - getattr(aggregator, "tencent_stock_quote", None), - ): - if not callable(loader): - continue - try: - quote = loader(ts_code, expected_date=reference_date) - except Exception: - continue - if quote: - return dict(quote) return {} def _stock_activity_metrics( diff --git a/backend/data/providers/tushare_indices.py b/backend/data/providers/tushare_indices.py index eab8eca..2ceb483 100644 --- a/backend/data/providers/tushare_indices.py +++ b/backend/data/providers/tushare_indices.py @@ -63,22 +63,8 @@ class IndexMixin: if callable(hub): rows = hub() if rows: - try: - return self._hub_realtime_market_indices(requested_date, rows) - except TushareError: - pass - try: - payload = self._tushare_realtime_market_indices(requested_date) - marker = getattr(self, "record_datahub_legacy", None) - if callable(marker): - marker("index_quotes", "tushare_rt_idx_k") - return payload - except TushareError: - payload = self._free_realtime_market_indices(requested_date) - marker = getattr(self, "record_datahub_legacy", None) - if callable(marker): - marker("index_quotes", str(payload.get("source") or "eastmoney_push2")) - return payload + return self._hub_realtime_market_indices(requested_date, rows) + raise TushareError("Realtime index quotes are incomplete") def _hub_realtime_market_indices( self, @@ -199,50 +185,5 @@ class IndexMixin: } def _free_realtime_market_indices(self, requested_date: str) -> dict[str, Any]: - trade_date, _ = self.resolve_trade_context(requested_date) - aggregator = getattr(self, "realtime_aggregator", None) - if aggregator is None: - raise TushareError("免费实时源未配置") - quotes = aggregator.eastmoney_indices() - index_names = { - "000001": ("000001.SH", "上证指数"), - "399001": ("399001.SZ", "深证成指"), - "399006": ("399006.SZ", "创业板指"), - } - indices = [] - for quote in quotes: - mapped = index_names.get(str(quote.get("code") or "")) - if not mapped: - continue - ts_code, name = mapped - close = _number(quote.get("price")) - previous_close = _number(quote.get("previous_close")) - if close <= 0 or previous_close <= 0: - continue - indices.append( - { - "ts_code": ts_code, - "name": str(quote.get("name") or name).strip(), - "trade_date": trade_date, - "close": close, - "pct_chg": round(_number(quote.get("change")) or (close / previous_close - 1) * 100, 3), - "return_5d": 0, - "amount_billion": round(_number(quote.get("amount_billion")), 2), - "quote_time": quote.get("quote_time") or "", - "source": quote.get("source") or "eastmoney_push2", - } - ) - if len(indices) != 3: - raise TushareError("Realtime index quotes are incomplete") - return { - "trade_date": trade_date, - "source": "eastmoney_push2", - "realtime": True, - "precise": True, - "indices": indices, - "aggregate": { - "average_pct_chg": round(sum(item["pct_chg"] for item in indices) / len(indices), 3), - "average_return_5d": 0, - "average_return_20d": 0, - }, - } + del requested_date + raise TushareError("主网站不再直连免费行情源,请走数据中枢") diff --git a/backend/data/providers/tushare_industries.py b/backend/data/providers/tushare_industries.py index 143d661..1dca38c 100644 --- a/backend/data/providers/tushare_industries.py +++ b/backend/data/providers/tushare_industries.py @@ -618,27 +618,19 @@ class ShenwanIndustryMixin: trade_date: str, finalized: bool = False, ) -> tuple[dict[str, Any], str, str]: - aggregator = getattr(self, "realtime_aggregator", None) - loader = getattr(aggregator, "eastmoney_shenwan_quote", None) if aggregator else None - if callable(loader): + hub = getattr(self, "try_sector_quote", None) + if callable(hub): try: - row = loader(sector_code, expected_date="" if finalized else trade_date) + row = hub(sector_code, "" if finalized else trade_date) except Exception as exc: message = str(exc) if finalized: return {}, "", f"申万行业 {sector_code} 盘后正式数据待入库" - return {}, "", f"免费申万实时暂不可用:{message[:180]}" + return {}, "", f"数据中枢申万实时暂不可用:{message[:180]}" if row: - return dict(row), str(row.get("source") or "eastmoney_sw"), "" + return dict(row), str(row.get("source") or "datahub"), "" if finalized: return {}, "", f"申万行业 {sector_code} 当日盘后正式数据尚未入库" - if aggregator and sector_name: - try: - row = aggregator.eastmoney_sector(sector_name) - except Exception as exc: - return {}, "", f"免费行业实时暂不可用:{str(exc)[:180]}" - if row: - return dict(row), str(row.get("source") or "eastmoney_sector"), "" return {}, "", f"申万行业 {sector_code} 当日外显待补充" def _load_member_realtime_quotes( @@ -677,31 +669,6 @@ class ShenwanIndustryMixin: delayed = any(item.get("delayed") for item in filtered) return filtered, "datahub_delayed" if delayed else "datahub" - aggregator = getattr(self, "realtime_aggregator", None) - eastmoney_loader = getattr(aggregator, "eastmoney_stock_quotes", None) if aggregator else None - if callable(eastmoney_loader): - try: - filtered = consider(eastmoney_loader(wanted, expected_date=trade_date) or [], "eastmoney_ulist") - if len(filtered) >= max(1, int(len(wanted) * 0.9)): - return filtered, "eastmoney_ulist" - except Exception: - pass - - tencent_loader = getattr(aggregator, "tencent_stock_quotes", None) if aggregator else None - if callable(tencent_loader): - try: - filtered = consider(tencent_loader(wanted, expected_date=trade_date) or [], "tencent_qt") - if len(filtered) >= max(1, int(len(wanted) * 0.9)): - return filtered, "tencent_qt" - except Exception: - pass - - try: - quotes, source = self._free_realtime_quotes(trade_date, ",".join(wanted)) - consider(quotes, source) - except TushareError: - pass - if best_rows: delayed = any(item.get("delayed") for item in best_rows) if delayed and not str(best_source).endswith("_delayed"): diff --git a/backend/features/market/charts.py b/backend/features/market/charts.py index 61afe09..d67757c 100644 --- a/backend/features/market/charts.py +++ b/backend/features/market/charts.py @@ -61,11 +61,7 @@ class MarketChartClient: hub_chart = self._datahub_intraday(normalized) if hub_chart is not None: return hub_chart - ifind_code = _stock_market_code(normalized) - try: - return self._ifind_intraday(ifind_code, "stock", normalized) - except (IfindError, ChartDataError): - return self.fallback.stock_intraday(normalized) + raise ChartDataError("分时图数据中枢暂不可用") def stock_daily(self, code: str, end_date: str, limit: int = DAILY_CHART_LIMIT) -> list[dict[str, Any]]: normalized = str(code or "").strip() @@ -74,7 +70,7 @@ class MarketChartClient: hub_rows = self._datahub_daily(normalized, end_date, limit, "daily") if hub_rows: return hub_rows - return self._ifind_daily(_stock_market_code(normalized), end_date, limit) + raise ChartDataError("日K数据中枢暂不可用") def index_daily(self, identifier: str, end_date: str, limit: int = DAILY_CHART_LIMIT) -> list[dict[str, Any]]: normalized = str(identifier or "").strip().upper() @@ -83,13 +79,16 @@ class MarketChartClient: hub_rows = self._datahub_daily(normalized, end_date, limit, "index_daily") if hub_rows: return hub_rows - return self._ifind_daily(normalized, end_date, limit) + raise ChartDataError("指数日K数据中枢暂不可用") def board_daily(self, identifier: str, end_date: str, limit: int = 90) -> list[dict[str, Any]]: normalized = str(identifier or "").strip().upper() if not normalized: raise ChartDataError("Invalid board code") - return self._ifind_daily(normalized, end_date, limit) + hub_rows = self._datahub_daily(normalized, end_date, limit, "sector_daily") + if hub_rows: + return hub_rows + raise ChartDataError("板块日K数据中枢暂不可用") def index_intraday(self, identifier: str) -> dict[str, Any]: normalized = str(identifier or "").strip().upper() @@ -98,10 +97,7 @@ class MarketChartClient: hub_chart = self._datahub_intraday(normalized) if hub_chart is not None: return hub_chart - try: - return self._ifind_intraday(normalized, "index", normalized) - except (IfindError, ChartDataError): - return self.fallback.index_intraday(normalized) + raise ChartDataError("指数分时数据中枢暂不可用") def _datahub_intraday(self, code: str) -> dict[str, Any] | None: if self.datahub is None: @@ -133,8 +129,6 @@ class MarketChartClient: LOGGER.warning("datahub daily unexpected error: %s", exc) rows = None if not rows: - if hasattr(self.datahub, "record_legacy"): - self.datahub.record_legacy(dataset, "ifind") return None compact_end = str(end_date or "").replace("-", "") market_now = datetime.now().astimezone() @@ -226,10 +220,13 @@ class MarketChartClient: def board_intraday(self, identifier: str, name: str = "") -> dict[str, Any]: normalized = str(identifier or "").strip().upper() - try: - return self._ifind_intraday(normalized, "board", normalized, name) - except (IfindError, ChartDataError): - return self.fallback.board_intraday(normalized, name) + hub_chart = self._datahub_intraday(normalized) + if hub_chart is not None: + if name: + hub_chart = dict(hub_chart) + hub_chart["name"] = name + return hub_chart + raise ChartDataError("板块分时数据中枢暂不可用") def _ifind_intraday( self, diff --git a/config/architecture-inventory.json b/config/architecture-inventory.json index a79c671..7ce04ab 100644 --- a/config/architecture-inventory.json +++ b/config/architecture-inventory.json @@ -488,8 +488,8 @@ }, { "path": "backend/data/providers/tushare_industries.py", - "bytes": 38757, - "lines": 892 + "bytes": 37168, + "lines": 859 }, { "path": "backend/features/screener/catalog.py", @@ -503,8 +503,8 @@ }, { "path": "backend/data/providers/tushare_dashboard.py", - "bytes": 34773, - "lines": 815 + "bytes": 33230, + "lines": 770 }, { "path": "database.py", @@ -591,11 +591,6 @@ "bytes": 12829, "lines": 318 }, - { - "path": "backend/data/providers/tushare_indices.py", - "bytes": 10956, - "lines": 248 - }, { "path": "backend/features/market/insights_auction.py", "bytes": 10717, @@ -616,16 +611,16 @@ "bytes": 9348, "lines": 222 }, - { - "path": "backend/data/providers/tushare_daily.py", - "bytes": 9170, - "lines": 233 - }, { "path": "frontend/pages/market/entity-detail.js", "bytes": 9139, "lines": 199 }, + { + "path": "backend/data/providers/tushare_daily.py", + "bytes": 9076, + "lines": 232 + }, { "path": "backend/data/providers/tushare_dragon_tiger.py", "bytes": 9059, @@ -636,6 +631,11 @@ "bytes": 8562, "lines": 238 }, + { + "path": "backend/data/providers/tushare_indices.py", + "bytes": 8447, + "lines": 189 + }, { "path": "frontend/pages/mentor/page.html", "bytes": 8357, diff --git a/tests/test_chart_data_provider.py b/tests/test_chart_data_provider.py index 6bb6008..e25e8e4 100644 --- a/tests/test_chart_data_provider.py +++ b/tests/test_chart_data_provider.py @@ -202,7 +202,7 @@ class DatahubChartFallbackTests(unittest.TestCase): self.assertEqual(hub.calls, ["601318"]) self.assertEqual(fallback.requests, []) - def test_datahub_timeout_or_empty_falls_back_to_eastmoney(self): + def test_datahub_timeout_or_empty_does_not_use_old_channel(self): fallback = LookbackChartClient() for hub in ( FakeHub(chart=None), @@ -213,10 +213,9 @@ class DatahubChartFallbackTests(unittest.TestCase): EastmoneyChartClient._cache.clear() fallback.requests.clear() client = MarketChartClient(IfindHttpClient(), fallback, hub) - payload = client.stock_intraday("000001") - self.assertEqual(payload["trade_date"], "2026-09-07") - self.assertGreaterEqual(len(payload["points"]), 1) - self.assertTrue(fallback.requests) + with self.assertRaises(ChartDataError): + client.stock_intraday("000001") + self.assertEqual(fallback.requests, []) def test_datahub_daily_skips_ifind(self): hub = FakeHub( diff --git a/tests/test_datahub_bridge.py b/tests/test_datahub_bridge.py index 0a94a78..f3917ea 100644 --- a/tests/test_datahub_bridge.py +++ b/tests/test_datahub_bridge.py @@ -13,6 +13,7 @@ from backend.data.datahub.compare import compare_rows from backend.data.datahub.errors import DatahubError from backend.data.datahub.native import to_canonical_row, to_native_row from backend.data.datahub.route_state import LEDGER +from backend.data.providers.tushare_transport import TushareError from backend.data.datahub.settings import DATASETS, DatahubSettings, DatasetFlags ROOT = Path(__file__).resolve().parents[1] @@ -68,9 +69,16 @@ class FakeClient(DatahubClient): self.calls: list[tuple[str, dict[str, Any]]] = [] def get(self, path: str, params: dict[str, Any] | None = None) -> DatahubResponse: + return self._record(path, params) + + def post(self, path: str, body: dict[str, Any] | None = None) -> DatahubResponse: + return self._record(path, body) + + def _record(self, path: str, payload: dict[str, Any] | None) -> DatahubResponse: self.paths.append(path) - self.calls.append((path, {key: value for key, value in (params or {}).items()})) - if TOKEN in json.dumps(params or {}) or TOKEN in path: + self.calls.append((path, {key: value for key, value in (payload or {}).items()})) + packed = json.dumps(payload or {}) + if TOKEN in packed or TOKEN in path: raise AssertionError("token leaked into url") if self.error: raise self.error @@ -131,16 +139,22 @@ class DatahubBridgeTests(unittest.TestCase): self.assertEqual(legacy.calls, []) self.assertEqual(client.paths, ["/v1/bars/daily"]) calendar_legacy = FakeLegacy([{"cal_date": "20240902", "is_open": 1}]) - calendar_client = FakeClient(error=DatahubError("INTERNAL", "nope")) + calendar_client = FakeClient( + response=DatahubResponse( + data=[{"cal_date": "20240902", "is_open": 1, "pretrade_date": "20240830"}], + meta={"source": "datahub", "stale": False, "staleness_seconds": 0}, + ) + ) calendar_wrapped = DatahubAwareTushareClient( calendar_legacy, DatahubBridge(flags(daily=(True, False)), calendar_client), ) calendar = calendar_wrapped.query("trade_cal", {"start_date": "20240902", "end_date": "20240902"}, "") self.assertEqual(calendar[0]["is_open"], 1) - self.assertEqual(calendar_client.paths, []) + self.assertEqual(calendar_legacy.calls, []) + self.assertEqual(calendar_client.paths, ["/v1/query"]) - def test_fallback_on_down_401_timeout_empty_unpublished_stale_and_incomplete(self) -> None: + def test_hub_failure_does_not_call_website_legacy(self) -> None: cases = [ DatahubError("UNAVAILABLE", "down"), DatahubError("UNAUTHORIZED", "401"), @@ -152,34 +166,21 @@ class DatahubBridgeTests(unittest.TestCase): ] for error in cases: with self.subTest(error=error.code): - if error.code == "EMPTY": - client = FakeClient(response=DatahubResponse(data=[], meta={"stale": False, "staleness_seconds": 0})) - elif error.code == "STALE": - client = FakeClient(response=DatahubResponse( - data=[dict(HUB_DAILY)], - meta={"stale": True, "staleness_seconds": 999999}, - )) - elif error.code == "INCOMPLETE": - client = FakeClient(response=DatahubResponse( - data=[dict(HUB_DAILY)], - meta={ - "stale": False, - "staleness_seconds": 0, - "incomplete": True, - "coverage": {"complete": False, "missing_count": 80}, - }, - )) - else: - client = FakeClient(error=error) + client = FakeClient(error=error) legacy = FakeLegacy([LEGACY_DAILY]) wrapped = DatahubAwareTushareClient(legacy, DatahubBridge(flags(daily=(True, False)), client)) - rows = wrapped.query("daily", {"trade_date": "20240902"}, "ts_code,amount") - self.assertEqual(rows[0]["amount"], 2000.0) - self.assertEqual(len(legacy.calls), 1) + with self.assertRaises(TushareError): + wrapped.query("daily", {"trade_date": "20240902"}, "ts_code,amount") + self.assertEqual(legacy.calls, []) - def test_shadow_compares_without_replacing_and_survives_hub_failure(self) -> None: + def test_shadow_mode_no_longer_calls_website_tushare(self) -> None: reports: list[dict[str, Any]] = [] - client = FakeClient() + client = FakeClient( + response=DatahubResponse( + data=[dict(LEGACY_DAILY)], + meta={"source": "tushare", "stale": False, "staleness_seconds": 0, "row_shape": "tushare"}, + ) + ) legacy = FakeLegacy([LEGACY_DAILY]) wrapped = DatahubAwareTushareClient( legacy, @@ -187,21 +188,19 @@ class DatahubBridgeTests(unittest.TestCase): ) rows = wrapped.query("daily", {"trade_date": "20240902"}, "ts_code,amount,vol") self.assertEqual(rows[0]["amount"], 2000.0) - self.assertEqual(len(legacy.calls), 1) - self.assertEqual(reports[0]["equal"], True) - self.assertEqual(reports[0]["matched"], 1) + self.assertEqual(legacy.calls, []) + self.assertEqual(client.paths, ["/v1/query"]) failed = FakeClient(error=DatahubError("UNAVAILABLE", TOKEN)) - fail_reports: list[dict[str, Any]] = [] fail_legacy = FakeLegacy([LEGACY_DAILY]) fail_wrapped = DatahubAwareTushareClient( fail_legacy, - DatahubBridge(flags(daily=(False, True)), failed, shadow_sink=fail_reports.append), + DatahubBridge(flags(daily=(False, True)), failed, shadow_sink=reports.append), ) - again = fail_wrapped.query("daily", {"trade_date": "20240902"}, "amount") - self.assertEqual(again[0]["amount"], 2000.0) - self.assertTrue(fail_reports[0]["hub_error"]) - self.assertNotIn(TOKEN, json.dumps(fail_reports[0])) + with self.assertRaises(TushareError): + fail_wrapped.query("daily", {"trade_date": "20240902"}, "amount") + self.assertEqual(fail_legacy.calls, []) + self.assertNotIn(TOKEN, str(failed.calls)) def test_compare_classifies_unit_conversion_missing_row_and_value_diff(self) -> None: equal = compare_rows("daily", [LEGACY_DAILY], [HUB_DAILY], {"stale": False, "staleness_seconds": 0}) @@ -288,13 +287,12 @@ class DatahubBridgeTests(unittest.TestCase): ) wrapped = DatahubAwareTushareClient( FakeLegacy([legacy_close_only]), - DatahubBridge(flags(daily=(False, True)), client, shadow_sink=reports.append), + DatahubBridge(flags(daily=(True, False)), client, shadow_sink=reports.append), ) rows = wrapped.query("daily", {"trade_date": "20240902"}, "ts_code,trade_date,close,vol,amount") self.assertEqual(rows[0]["close"], 10.20) self.assertEqual(rows[0]["vol"], 1000.0) - self.assertTrue(reports[0]["equal"]) - self.assertEqual(reports[0]["matched"], 1) + self.assertEqual(client.paths, ["/v1/bars/daily"]) def test_native_roundtrip_matches_known_scales(self) -> None: native = to_native_row("daily", HUB_DAILY) @@ -347,21 +345,17 @@ class DatahubBridgeTests(unittest.TestCase): self.assertIn('"daily"', source) self.assertIn("start_date", source) self.assertIn("end_date", source) - client = FakeClient( - response=DatahubResponse( - data=[dict(HUB_DAILY)], - meta={"stale": False, "staleness_seconds": 0, "incomplete": True, "coverage": {"complete": False, "missing_count": 89}}, - ) - ) + client = FakeClient(error=DatahubError("INCOMPLETE", "truncated")) legacy = FakeLegacy([LEGACY_DAILY]) wrapped = DatahubAwareTushareClient(legacy, DatahubBridge(flags(daily=(True, False)), client)) - rows = wrapped.query( - "daily", - {"ts_code": "600000.SH", "start_date": "20240301", "end_date": "20240902"}, - "ts_code,amount", - ) - self.assertEqual(rows[0]["amount"], 2000.0) - self.assertEqual(len(legacy.calls), 1) + with self.assertRaises(TushareError): + wrapped.query( + "daily", + {"ts_code": "600000.SH", "start_date": "20240301", "end_date": "20240902"}, + "ts_code,amount", + ) + self.assertEqual(legacy.calls, []) + self.assertIn("/v1/query", client.paths) def test_try_intraday_respects_switch_and_falls_back_on_bad_payload(self) -> None: closed = DatahubBridge(flags(), FakeClient(error=DatahubError("INTERNAL", "should not run"))) @@ -462,17 +456,15 @@ class DatahubBridgeTests(unittest.TestCase): FakeClient(error=DatahubError("UNAVAILABLE", "down")), ) self.assertIsNone(failed.try_market_quotes("20240902")) - failed.record_legacy("quotes", "tencent_qt", "down") snap = next(item for item in LEDGER.snapshot() if item["dataset"] == "quotes") - self.assertEqual(snap["route"], "legacy") - self.assertEqual(snap["source"], "tencent_qt") - self.assertIn("备用", "备用") + self.assertEqual(snap["route"], "datahub") + self.assertEqual(snap["source"], "unavailable") gateway = build_data_gateway({}, datahub_settings=flags(quotes=(True, False))) status = gateway.datahub_status() self.assertEqual(status["enabled_reads"], 1) self.assertEqual(status["total_reads"], len(DATASETS)) - self.assertGreaterEqual(status["fallback_count"], 1) + self.assertEqual(status["fallback_count"], 0) def test_try_daily_chart_converts_hub_bars(self) -> None: rows = [ diff --git a/tests/test_hel494_regressions.py b/tests/test_hel494_regressions.py index a4f16b1..d9035b6 100644 --- a/tests/test_hel494_regressions.py +++ b/tests/test_hel494_regressions.py @@ -87,11 +87,10 @@ class ShenwanRealtimeSourceTests(unittest.TestCase): with self.assertRaisesRegex(TushareError, "rt_sw_k is disabled"): client.query("rt_sw_k", {"ts_code": "801074.SI"}) - def test_outer_realtime_uses_eastmoney_shenwan_not_rt_sw_k(self) -> None: + def test_outer_realtime_uses_hub_sector_quote_not_rt_sw_k(self) -> None: client = TushareClient(token="demo") client.query = MagicMock(side_effect=AssertionError("should not call tushare")) - client.realtime_aggregator = MagicMock() - client.realtime_aggregator.eastmoney_shenwan_quote.return_value = { + client.try_sector_quote = MagicMock(return_value={ "code": "801074.SI", "name": "工业金属", "close": 1234.5, @@ -101,7 +100,7 @@ class ShenwanRealtimeSourceTests(unittest.TestCase): "quote_date": "20260908", "quote_time": "2026-09-08T14:50:00+08:00", "source": "eastmoney_sw", - } + }) row, source, error = client._sw_outer_realtime("801074.SI", "工业金属", "20260908") self.assertEqual(source, "eastmoney_sw") self.assertEqual(error, "") @@ -216,23 +215,21 @@ class MemberQuoteCoverageTests(unittest.TestCase): self.assertEqual(source, "datahub") client.try_quotes.assert_not_called() - def test_eastmoney_failure_uses_tencent_member_quotes(self) -> None: + def test_hub_named_quotes_cover_members_when_market_missing(self) -> None: client = TushareClient(token="demo") wanted = ["000737.SZ", "000630.SZ"] client.try_market_quotes = MagicMock(return_value=None) - client.try_quotes = MagicMock(return_value=None) - aggregator = MagicMock() - aggregator.eastmoney_stock_quotes.side_effect = RuntimeError("HTTP 503") - aggregator.tencent_stock_quotes.return_value = [ + client.try_quotes = MagicMock(return_value=[ {"ts_code": "000737.SZ", "close": 12.3, "pre_close": 11.2}, {"ts_code": "000630.SZ", "close": 4.5, "pre_close": 4.4}, - ] - client.realtime_aggregator = aggregator - client._free_realtime_quotes = MagicMock(side_effect=AssertionError("tencent already won")) + ]) + client.realtime_aggregator = MagicMock() rows, source = client._load_member_realtime_quotes(wanted, "20260908") self.assertEqual(len(rows), 2) - self.assertEqual(source, "tencent_qt") - aggregator.tencent_stock_quotes.assert_called_once() + self.assertEqual(source, "datahub") + client.try_quotes.assert_called() + client.realtime_aggregator.eastmoney_stock_quotes.assert_not_called() + client.realtime_aggregator.tencent_stock_quotes.assert_not_called() def test_delayed_hub_quotes_are_kept_not_zeroed(self) -> None: client = TushareClient(token="demo") @@ -323,8 +320,7 @@ class MemberQuoteCoverageTests(unittest.TestCase): client._sw_realtime_sector_snapshot = MagicMock( side_effect=AssertionError("daily inner should be kept") ) - client.realtime_aggregator = MagicMock() - client.realtime_aggregator.eastmoney_shenwan_quote.return_value = { + client.try_sector_quote = MagicMock(return_value={ "code": "801074.SI", "name": "工业金属", "change": 1.5, @@ -332,7 +328,7 @@ class MemberQuoteCoverageTests(unittest.TestCase): "quote_date": "20260908", "quote_time": "2026-09-08T15:00:00+08:00", "source": "eastmoney_sw", - } + }) snapshot = client.sw_sector_snapshot( "000737.SZ", "20260908", allow_realtime_close=True ) diff --git a/tests/test_hub_exclusive.py b/tests/test_hub_exclusive.py new file mode 100644 index 0000000..6c32f91 --- /dev/null +++ b/tests/test_hub_exclusive.py @@ -0,0 +1,106 @@ +from __future__ import annotations + +import ast +import unittest +from pathlib import Path +from unittest.mock import patch + +from backend.data.datahub.bridge import DatahubAwareTushareClient, DatahubBridge +from backend.data.datahub.client import DatahubClient +from backend.data.datahub.settings import DATASETS, DatahubSettings, DatasetFlags +from backend.data.providers.tushare_transport import TushareError +from tests.test_datahub_bridge import FakeClient, FakeLegacy, flags + + +ROOT = Path(__file__).resolve().parents[1] +BLOCKED_HOSTS = ( + "api.tushare.pro", + "push2.eastmoney.com", + "push2delay.eastmoney.com", + "push2his.eastmoney.com", + "push2ex.eastmoney.com", + "qt.gtimg.cn", + "hq.sinajs.cn", + "10jqka.com.cn", + "xuangubao.cn", +) + + +class HubExclusiveWebsiteTests(unittest.TestCase): + def test_query_never_calls_website_tushare_transport(self) -> None: + client = FakeClient() + legacy = FakeLegacy(TushareError("website tushare must stay dark")) + wrapped = DatahubAwareTushareClient( + legacy, + DatahubBridge(flags(daily=(True, False)), client), + ) + rows = wrapped.query("daily", {"trade_date": "20240902"}, "ts_code,amount") + self.assertEqual(rows[0]["amount"], 2000.0) + self.assertEqual(legacy.calls, []) + + def test_blocked_external_hosts_still_read_hub(self) -> None: + settings = DatahubSettings( + base_url="http://127.0.0.1:8766", + token="hub-token", + datasets={name: DatasetFlags(name, read=True) for name in DATASETS}, + ) + + def blocked_urlopen(request, timeout=None): + url = str(getattr(request, "full_url", None) or request) + if any(host in url for host in BLOCKED_HOSTS): + raise AssertionError(f"website opened blocked host: {url}") + if "127.0.0.1:8766" in url or "v1/bars/daily" in url: + class _Resp: + status = 200 + + def read(self): + return ( + b'{"schema_version":1,"data":[{"ts_code":"600000.SH","trade_date":"20240902",' + b'"close":10.2,"volume":100000,"amount":2000000}],' + b'"meta":{"stale":false,"staleness_seconds":0,"source":"datahub"}}' + ) + + def __enter__(self): + return self + + def __exit__(self, *args): + return False + + return _Resp() + raise AssertionError(f"unexpected url: {url}") + + hub_client = DatahubClient(settings, urlopen=blocked_urlopen) + legacy = FakeLegacy(TushareError("blocked")) + wrapped = DatahubAwareTushareClient(legacy, DatahubBridge(settings, hub_client)) + with patch("urllib.request.urlopen", blocked_urlopen): + rows = wrapped.query("daily", {"trade_date": "20240902"}, "ts_code,close,amount") + self.assertEqual(rows[0]["close"], 10.2) + self.assertEqual(rows[0]["amount"], 2000.0) + self.assertEqual(legacy.calls, []) + + def test_website_runtime_does_not_call_blocked_hosts_from_gateway(self) -> None: + gateway_src = (ROOT / "backend" / "data" / "gateway.py").read_text(encoding="utf-8") + self.assertIn("legacy.realtime_aggregator = None", gateway_src) + self.assertIn("DatahubAwareTushareClient", gateway_src) + + def test_bridge_query_has_no_legacy_call(self) -> None: + source = (ROOT / "backend" / "data" / "datahub" / "bridge.py").read_text(encoding="utf-8") + tree = ast.parse(source) + query_fn = next( + node + for node in tree.body + if isinstance(node, ast.ClassDef) and node.name == "DatahubBridge" + for item in node.body + if isinstance(item, ast.FunctionDef) and item.name == "query" + ) + called = [ + ast.unparse(item.func) if hasattr(ast, "unparse") else "" + for item in ast.walk(query_fn) + if isinstance(item, ast.Call) + ] + self.assertTrue(any("query_api" in text for text in called)) + self.assertFalse(any("legacy_query" in text for text in called)) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_ifind_features.py b/tests/test_ifind_features.py index ec9accc..c0e16a4 100644 --- a/tests/test_ifind_features.py +++ b/tests/test_ifind_features.py @@ -12,6 +12,20 @@ from backend.features.market.insights import MarketInsightsService from server import DashboardService +class _FakeDailyHub: + def __init__(self, rows: list) -> None: + self.rows = rows + + def try_daily_chart(self, code, end_date, limit, dataset="daily"): + return list(self.rows) + + def try_quotes(self, codes): + return None + + def try_index_quotes(self): + return None + + class FakeIfind: configured = True @@ -128,13 +142,50 @@ class IfindFeatureTests(unittest.TestCase): self.assertEqual(database.list_wencai_saved_queries(second["id"]), []) def test_ifind_daily_chart_normalizes_change(self): - client = MarketChartClient(FakeIfind(), EastmoneyChartClient()) + hub = _FakeDailyHub( + [ + { + "trade_date": "2026-07-27", + "open": 10, + "high": 10.5, + "low": 9.8, + "close": 10.2, + "volume": 100, + "amount_billion": 0.01, + "change": 0, + }, + { + "trade_date": "2026-07-28", + "open": 10.2, + "high": 10.8, + "low": 10.1, + "close": 10.5, + "volume": 120, + "amount_billion": 0.012, + "change": 2.9412, + }, + ] + ) + client = MarketChartClient(FakeIfind(), EastmoneyChartClient(), hub) rows = client.stock_daily("000001", "20260728") self.assertEqual(rows[-1]["trade_date"], "2026-07-28") self.assertAlmostEqual(rows[-1]["change"], 2.9412, places=4) def test_ifind_daily_chart_keeps_last_traded_bar_before_market_open(self): - client = MarketChartClient(FakeIfindStalePreopen(), EastmoneyChartClient()) + hub = _FakeDailyHub( + [ + { + "trade_date": "2026-07-28", + "open": 10.2, + "high": 10.8, + "low": 10.1, + "close": 10.5, + "volume": 120, + "amount_billion": 0.012, + } + ] + ) + client = MarketChartClient(FakeIfindStalePreopen(), EastmoneyChartClient(), hub) with patch("backend.features.market.charts.datetime", FixedPreopenDatetime): rows = client.stock_daily("000001", "20260729") diff --git a/tests/test_realtime_dashboard.py b/tests/test_realtime_dashboard.py index 1f01ea4..5f83631 100644 --- a/tests/test_realtime_dashboard.py +++ b/tests/test_realtime_dashboard.py @@ -256,7 +256,7 @@ class RealtimeDashboardTests(unittest.TestCase): self.assertEqual(dashboard["meta"]["quote_count"], 3) self.assertEqual(dashboard["overview"]["limit_up_count"], 0) - def test_rt_k_permission_error_falls_back_to_free_quotes(self): + def test_hub_quotes_used_when_rt_k_denied(self): original_query = self.client.query def query(api_name, params=None, fields=""): @@ -265,21 +265,20 @@ class RealtimeDashboardTests(unittest.TestCase): return original_query(api_name, params, fields) self.client.query = query - self.client.realtime_aggregator = FakeFreeAggregator() + self.client.try_market_quotes = lambda trade_date: list(FREE_QUOTES) TushareClient._realtime_reference_cache.clear() dashboard = self.client._realtime_dashboard("20260720", "20260720", "20260717") self.assertTrue(dashboard["meta"]["realtime"]) - self.assertEqual(dashboard["meta"]["quote_source"], "eastmoney_clist") + self.assertEqual(dashboard["meta"]["quote_source"], "datahub") self.assertEqual(dashboard["meta"]["trade_date"], "2026-07-20") self.assertEqual(dashboard["meta"]["quote_count"], 3) self.assertEqual(dashboard["overview"]["limit_up_count"], 1) self.assertEqual(dashboard["overview"]["limit_down_count"], 1) self.assertEqual(dashboard["overview"]["amount_billion"], 6.0) - self.assertIn("东财免费实时", dashboard["meta"]["notice"]) - self.assertEqual(dashboard["meta"]["indices"][0]["price"], 3800.12) + self.assertIn("数据中枢", dashboard["meta"]["notice"]) - def test_rt_k_empty_result_falls_back_to_free_quotes(self): + def test_hub_quotes_used_when_rt_k_empty(self): original_query = self.client.query def query(api_name, params=None, fields=""): @@ -288,29 +287,27 @@ class RealtimeDashboardTests(unittest.TestCase): return original_query(api_name, params, fields) self.client.query = query - self.client.realtime_aggregator = FakeFreeAggregator() + self.client.try_market_quotes = lambda trade_date: list(FREE_QUOTES) TushareClient._realtime_reference_cache.clear() dashboard = self.client._realtime_dashboard("20260720", "20260720", "20260717") - self.assertEqual(dashboard["meta"]["quote_source"], "eastmoney_clist") + self.assertEqual(dashboard["meta"]["quote_source"], "datahub") self.assertEqual(str(dashboard["meta"]["trade_date"]).replace("-", ""), "20260720") - def test_rt_k_and_free_source_failure_keeps_today_error(self): + def test_hub_failure_keeps_today_error(self): original_query = self.client.query def query(api_name, params=None, fields=""): if api_name == "rt_k": - raise TushareError("没有接口访问权限") + raise TushareError("数据中枢行情暂不可用") return original_query(api_name, params, fields) self.client.query = query - self.client.realtime_aggregator = FakeFreeAggregator(fail=True) TushareClient._realtime_reference_cache.clear() with self.assertRaises(TushareError) as ctx: self.client._realtime_dashboard("20260720", "20260720", "20260717") self.assertIn("当天盘中实时行情不可用", str(ctx.exception)) - self.assertIn("没有接口访问权限", str(ctx.exception)) - def test_rt_k_and_eastmoney_failure_falls_back_to_tencent(self): + def test_hub_failover_is_invisible_to_website(self): original_query = self.client.query def query(api_name, params=None, fields=""): @@ -318,23 +315,13 @@ class RealtimeDashboardTests(unittest.TestCase): raise TushareError("没有接口访问权限") return original_query(api_name, params, fields) - class TencentOnlyAggregator(FakeFreeAggregator): - def eastmoney_market_quotes(self, expected_date=""): - raise RealtimeAggregateError("eastmoney blocked") - - def tencent_market_quotes(self, codes, expected_date=""): - return list(FREE_QUOTES) - - def tencent_stock_quotes(self, codes, expected_date="", minimum=None): - return list(FREE_QUOTES) - self.client.query = query - self.client.realtime_aggregator = TencentOnlyAggregator() + self.client.try_market_quotes = lambda trade_date: list(FREE_QUOTES) TushareClient._realtime_reference_cache.clear() dashboard = self.client._realtime_dashboard("20260720", "20260720", "20260717") - self.assertEqual(dashboard["meta"]["quote_source"], "tencent_qt") + self.assertEqual(dashboard["meta"]["quote_source"], "datahub") self.assertEqual(str(dashboard["meta"]["trade_date"]).replace("-", ""), "20260720") - self.assertIn("腾讯免费实时", dashboard["meta"]["notice"]) + self.assertIn("数据中枢", dashboard["meta"]["notice"]) self.assertEqual(dashboard["overview"]["amount_billion"], 6.0) def test_normalize_eastmoney_quote_maps_units_and_exchange(self): @@ -433,6 +420,13 @@ class RealtimeDashboardTests(unittest.TestCase): def __init__(self): self.calls = [] + def query_api(self, api_name, params=None, fields=""): + rows = FakeRealtimeClient("tok").query(api_name, params or {}, fields) + return DatahubResponse( + data=rows, + meta={"source": "datahub", "stale": False, "staleness_seconds": 0, "row_shape": "tushare"}, + ) + def quotes_latest(self, **params): return self.get("/v1/quotes/latest", params) diff --git a/xiaobai-datahub/datahub/adapters/eastmoney.py b/xiaobai-datahub/datahub/adapters/eastmoney.py index 63c8f04..2dde4e2 100644 --- a/xiaobai-datahub/datahub/adapters/eastmoney.py +++ b/xiaobai-datahub/datahub/adapters/eastmoney.py @@ -250,6 +250,11 @@ class EastmoneyAdapter(MarketAdapter): secid = INDEX_SECIDS[code] entity = "index" identifier = code + elif code.startswith("BK") or code.endswith((".TI", ".SI")): + symbol = code.split(".")[0] + secid = f"90.{symbol}" + entity = "board" + identifier = symbol else: symbol = code.split(".")[0] market = "1" if symbol.startswith(("5", "6", "9")) else "0" @@ -294,6 +299,108 @@ class EastmoneyAdapter(MarketAdapter): "source": "eastmoney_trends2", } + def fetch_shenwan_quote(self, ts_code: str) -> dict[str, Any]: + code = str(ts_code or "").split(".")[0] + if not code: + raise AdapterError("Invalid Shenwan code") + payload = self._get_json( + EASTMONEY_INDEX_URL, + { + "secids": f"90.{code}", + "fltt": "2", + "invt": "2", + "fields": "f12,f14,f2,f3,f4,f15,f16,f17,f18,f6,f8,f104,f105,f128,f136,f140,f124", + }, + referer="https://quote.eastmoney.com/", + ) + rows = list((payload.get("data") or {}).get("diff") or []) + row = next((item for item in rows if item), None) + if not row: + raise AdapterError(f"Eastmoney Shenwan quote missing for {code}") + epoch = int(finite_number(row.get("f124")) or 0) + close = round4(finite_number(row.get("f2"))) + previous = round4(finite_number(row.get("f18"))) + if close <= 0 or previous <= 0: + raise AdapterError(f"Eastmoney Shenwan quote empty for {code}") + quote_time = ( + datetime.fromtimestamp(epoch).astimezone().isoformat(timespec="seconds") + if epoch + else "" + ) + return { + "ts_code": f"{code}.SI", + "code": f"{code}.SI", + "name": row.get("f14") or code, + "price": close, + "close": close, + "pre_close": previous, + "previous_close": previous, + "open": round4(finite_number(row.get("f17"))), + "high": round4(finite_number(row.get("f15"))), + "low": round4(finite_number(row.get("f16"))), + "change": round4(finite_number(row.get("f3"))), + "pct_change": round4(finite_number(row.get("f3"))), + "pct_chg": round4(finite_number(row.get("f3"))), + "amount": round4(finite_number(row.get("f6"))), + "leader": row.get("f128") or "--", + "leader_code": row.get("f140") or "", + "leading_pct": round4(finite_number(row.get("f136"))), + "up_count": int(finite_number(row.get("f104")) or 0), + "down_count": int(finite_number(row.get("f105")) or 0), + "quote_time": quote_time, + "trade_time": quote_time, + "quote_date": datetime.fromtimestamp(epoch).astimezone().strftime("%Y%m%d") if epoch else "", + "quote_time_epoch": epoch, + "source": "eastmoney_sw", + } + + def fetch_limit_pool(self, trade_date: str = "") -> list[dict[str, Any]]: + day = str(trade_date or "").replace("-", "") + rows: list[dict[str, Any]] = [] + for url, limit_type in ( + ("https://push2ex.eastmoney.com/getTopicZTPool", "U"), + ("https://push2ex.eastmoney.com/getTopicZBPool", "Z"), + ): + params = { + "ut": "7eea3edcaed734bea9cbfc24409ed989", + "dpt": "wz.ztzt", + "PageIndex": "0", + "PageSize": "200", + "sort": "fbt:asc", + "stat": "1", + } + if day: + params["date"] = day + try: + payload = self._get_json(url, params, referer="https://quote.eastmoney.com/") + except AdapterError: + continue + pool = ((payload.get("data") or {}).get("pool") or []) if isinstance(payload.get("data"), dict) else [] + for item in pool: + code = str(item.get("c") or item.get("code") or "") + if not code: + continue + market = str(item.get("m") or item.get("market") or "") + suffix = "SH" if market in {"1", "SH"} or code.startswith(("5", "6", "9")) else "SZ" + first = str(item.get("fbt") or item.get("first_time") or "") + last = str(item.get("lbt") or item.get("last_time") or "") + rows.append( + { + "ts_code": f"{code}.{suffix}", + "limit_type": limit_type, + "first_time": first, + "last_time": last, + "fd_amount": item.get("fund") or item.get("fd_amount"), + "open_times": item.get("zbc") or item.get("open_times"), + "limit_times": item.get("lbc") or item.get("limit_times"), + "turnover_ratio": item.get("hs") or item.get("turnover_ratio"), + "source": "eastmoney_zt_pool", + } + ) + if not rows: + raise AdapterError("Eastmoney limit pool empty") + return rows + def _get_json(self, url: str, params: dict[str, str], referer: str) -> dict[str, Any]: request_url = f"{url}?{urllib.parse.urlencode(params)}" request = urllib.request.Request( diff --git a/xiaobai-datahub/datahub/adapters/tushare.py b/xiaobai-datahub/datahub/adapters/tushare.py index 4362b1b..48b3a34 100644 --- a/xiaobai-datahub/datahub/adapters/tushare.py +++ b/xiaobai-datahub/datahub/adapters/tushare.py @@ -50,6 +50,14 @@ TUSHARE_FIELDS = { "ths_daily": "ts_code,trade_date,open,high,low,close,pre_close,pct_change,vol,turnover_rate", "dc_index": "ts_code,trade_date,name,open,high,low,close,pre_close,pct_change,vol,amount,turnover_rate", "sw_daily": "ts_code,trade_date,name,open,high,low,close,pct_change,vol,amount", + "index_member_all": ( + "l1_code,l1_name,l2_code,l2_name,l3_code,l3_name," + "ts_code,name,in_date,out_date,is_new" + ), + "stk_limit": "ts_code,trade_date,up_limit,down_limit", + "suspend_d": "ts_code,suspend_date,resume_date,ann_date,suspend_reason,reason_type", + "ths_member": "ts_code,con_code,con_name,in_date,out_date,is_new", + "stk_mins": "ts_code,trade_time,open,close,high,low,vol,amount", } DATASET_API = { @@ -254,3 +262,6 @@ class TushareAdapter(MarketAdapter): items = data.get("items") or [] fields_list = data.get("fields") or (fields.split(",") if fields else []) return [dict(zip(fields_list, item)) for item in items] + + def query_raw(self, api_name: str, params: dict[str, Any], fields: str = "") -> list[dict[str, Any]]: + return self._query(api_name, params, fields or TUSHARE_FIELDS.get(api_name, "")) diff --git a/xiaobai-datahub/datahub/httpapp.py b/xiaobai-datahub/datahub/httpapp.py index c649543..c1147b8 100644 --- a/xiaobai-datahub/datahub/httpapp.py +++ b/xiaobai-datahub/datahub/httpapp.py @@ -42,7 +42,7 @@ class HubRequestHandler(BaseHTTPRequestHandler): self._json({"status": "ok"}, HTTPStatus.OK) return if path.startswith("/v1/"): - self._v1(path, parsed.query) + self._v1(path, parsed.query, method) return if path.startswith("/admin/api/"): self._admin_api(method, path) @@ -66,11 +66,16 @@ class HubRequestHandler(BaseHTTPRequestHandler): LOGGER.exception("internal error") self._json({"error": {"code": "INTERNAL", "message": "internal error"}}, HTTPStatus.INTERNAL_SERVER_ERROR) - def _v1(self, path: str, query: str) -> None: + def _v1(self, path: str, query: str, method: str = "GET") -> None: token = self.headers.get("X-Datahub-Token", "") if not self.hub.auth.check_api_token(token): self.hub.pipeline.audit("anonymous", "unauthorized", path, "") raise ApiError("UNAUTHORIZED", "missing or invalid X-Datahub-Token") + if path == "/v1/query" and method == "POST": + body = self._read_json(max_bytes=1_000_000) + payload = self.hub.api.query_api(body) + self._json(payload, HTTPStatus.OK) + return payload = self.hub.api.handle(path, parse_query(query)) self._json(payload, HTTPStatus.OK) @@ -184,11 +189,11 @@ class HubRequestHandler(BaseHTTPRequestHandler): self.end_headers() self.wfile.write(content) - def _read_json(self, allow_empty: bool = False) -> dict[str, Any]: + def _read_json(self, allow_empty: bool = False, max_bytes: int = 65536) -> dict[str, Any]: length = int(self.headers.get("Content-Length", "0") or 0) if length == 0 and allow_empty: return {} - if length <= 0 or length > 65536: + if length <= 0 or length > max_bytes: raise ValueError("请求内容为空或过大") raw = self.rfile.read(length) try: diff --git a/xiaobai-datahub/datahub/realtime_serve.py b/xiaobai-datahub/datahub/realtime_serve.py index 537b247..bd338bd 100644 --- a/xiaobai-datahub/datahub/realtime_serve.py +++ b/xiaobai-datahub/datahub/realtime_serve.py @@ -149,6 +149,78 @@ def fetch_quotes(db: HubDB, codes: list[str]) -> dict[str, Any]: return payload +def fetch_sector_quote(db: HubDB, code: str, expected_date: str = "") -> dict[str, Any]: + ts_code = str(code or "").strip().upper() + if ts_code.isdigit(): + ts_code = f"{ts_code}.SI" + cache_key = f"sector:{ts_code}" + cached = _read_cache(db, cache_key) + if cached is not None: + return cached + errors: list[str] = [] + try: + row = EastmoneyAdapter().fetch_shenwan_quote(ts_code) + source = str(row.get("source") or "eastmoney_sw") + except Exception as exc: + errors.append(f"eastmoney:{exc}") + recovered = _load_quotes_lkg(db, cache_key) + if recovered is not None: + return recovered + raise RealtimeApiError( + "SOURCE_UNAVAILABLE", + "sector quote unavailable: " + ";".join(errors), + ) from exc + want = str(expected_date or "").replace("-", "")[:8] + quote_date = str(row.get("quote_date") or "") + if want and quote_date and quote_date != want: + recovered = _load_quotes_lkg(db, cache_key) + if recovered is not None: + return recovered + raise RealtimeApiError("SOURCE_UNAVAILABLE", f"sector quote date {quote_date} != {want}") + payload = _envelope( + row, + { + "tier": "provisional", + "trade_date": quote_date or yyyymmdd(now_shanghai()), + "source": source, + "stale": False, + "staleness_seconds": 0, + "published_at": isoformat(now_shanghai()), + }, + ) + _write_cache(db, cache_key, payload, INDEX_TTL, source) + return payload + + +def fetch_limit_pool(db: HubDB, trade_date: str = "") -> dict[str, Any]: + day = yyyymmdd(trade_date or now_shanghai()) + cache_key = f"limit-pool:{day}" + cached = _read_cache(db, cache_key) + if cached is not None: + return cached + try: + rows = EastmoneyAdapter().fetch_limit_pool(day) + source = "eastmoney:zt_pool" + except Exception as exc: + recovered = _load_quotes_lkg(db, cache_key) + if recovered is not None: + return recovered + raise RealtimeApiError("SOURCE_UNAVAILABLE", f"limit pool unavailable: {exc}") from exc + payload = _envelope( + rows, + { + "tier": "provisional", + "trade_date": day, + "source": source, + "stale": False, + "staleness_seconds": 0, + "published_at": isoformat(now_shanghai()), + }, + ) + _write_cache(db, cache_key, payload, QUOTE_TTL, source) + return payload + + def _eastmoney_named_quotes(codes: list[str]) -> list[dict[str, Any]]: adapter = EastmoneyAdapter() rows: list[dict[str, Any]] = [] diff --git a/xiaobai-datahub/datahub/serving.py b/xiaobai-datahub/datahub/serving.py index 6cfe214..67db3b6 100644 --- a/xiaobai-datahub/datahub/serving.py +++ b/xiaobai-datahub/datahub/serving.py @@ -87,12 +87,46 @@ class V1API: return self.index_quotes(q) if path == "/v1/intraday/points": return self.intraday_points(q) + if path == "/v1/sectors/quote": + return self.sector_quote(q) + if path == "/v1/limit-pool": + return self.limit_pool(q) + if path == "/v1/query": + return self.query_api(q) if path == "/v1/datasets/status": return self.dataset_status(q.get("date") or "") if path == "/v1/batches": return self.batches(q.get("date") or "", q.get("dataset") or "") raise ApiError("INVALID_ARGUMENT", f"unknown endpoint: {path}") + def query_api(self, body: dict[str, Any]) -> dict[str, Any]: + from datahub.steward import steward_query + + payload = dict(body or {}) + raw_params = payload.get("params") + if isinstance(raw_params, str): + payload["params"] = _parse_json(raw_params) or {} + return steward_query(self, payload) + + def sector_quote(self, q: dict[str, str]) -> dict[str, Any]: + from datahub.realtime_serve import RealtimeApiError, fetch_sector_quote + + code = str(q.get("code") or q.get("ts_code") or "").strip() + if not code: + raise ApiError("INVALID_ARGUMENT", "code is required") + try: + return fetch_sector_quote(self.db, code, str(q.get("date") or "")) + except RealtimeApiError as exc: + raise ApiError(exc.code, exc.message) from exc + + def limit_pool(self, q: dict[str, str]) -> dict[str, Any]: + from datahub.realtime_serve import RealtimeApiError, fetch_limit_pool + + try: + return fetch_limit_pool(self.db, str(q.get("date") or q.get("trade_date") or "")) + except RealtimeApiError as exc: + raise ApiError(exc.code, exc.message) from exc + def health(self) -> dict[str, Any]: today = yyyymmdd(now_shanghai()) cal = self.db.fetchone( diff --git a/xiaobai-datahub/datahub/steward.py b/xiaobai-datahub/datahub/steward.py new file mode 100644 index 0000000..f28efdf --- /dev/null +++ b/xiaobai-datahub/datahub/steward.py @@ -0,0 +1,360 @@ +"""Website-facing data steward: pick source, fail over, cache, never fake zeros. + +The main site asks for a business/Tushare-shaped API. This module decides whether +to serve a published EOD table, live free quotes, or an internal Tushare pull. +""" + +from __future__ import annotations + +import hashlib +import json +from typing import Any + +from datahub.adapters.base import AdapterError +from datahub.adapters.tushare import TUSHARE_FIELDS +from datahub.numbers import finite_number +from datahub.realtime_serve import ( + RealtimeApiError, + _read_cache, + _write_cache, + fetch_index_quotes, + fetch_market_quotes, + fetch_quotes, +) +from datahub.serving import ApiError, envelope +from datahub.timeutil import isoformat, now_shanghai, yyyymmdd + +API_TO_DATASET = { + "trade_cal": "calendar", + "stock_basic": "stocks", + "daily": "daily", + "daily_basic": "valuation", + "index_daily": "index_daily", + "moneyflow": "moneyflow", + "stk_auction": "auction", + "limit_list_d": "limit_events", + "ths_hot": "popularity", + "dc_hot": "popularity", + "hm_detail": "dragon_tiger", + "ths_daily": "sector_daily", + "dc_index": "sector_daily", + "sw_daily": "sector_daily", +} + +DATASET_FETCHER = { + "calendar": lambda api, q: api.calendar(q.get("from") or q.get("start_date") or "", q.get("to") or q.get("end_date") or ""), + "stocks": lambda api, q: api.stocks(q.get("updated_since") or "", q), + "daily": lambda api, q: api.daily_bars(_hub_query(q, adjust="none")), + "valuation": lambda api, q: api.valuation(_hub_query(q)), + "index_daily": lambda api, q: api.index_bars(_hub_query(q)), + "moneyflow": lambda api, q: api.moneyflow(_hub_query(q)), + "auction": lambda api, q: api.auction(_hub_query(q)), + "limit_events": lambda api, q: api.limit_events(_hub_query(q)), + "popularity": lambda api, q: api.popularity(_hub_query(q)), + "dragon_tiger": lambda api, q: api.dragon_tiger(_hub_query(q)), + "sector_daily": lambda api, q: api.sectors(_hub_query(q)), +} + +SCALE_TO_TUSHARE = { + "daily": {"vol": 100.0, "amount": 1000.0}, + "index_daily": {"vol": 100.0, "amount": 1000.0}, + "valuation": {"total_mv": 10000.0, "circ_mv": 10000.0}, + "moneyflow": { + "buy_sm_amount": 10000.0, + "sell_sm_amount": 10000.0, + "buy_md_amount": 10000.0, + "sell_md_amount": 10000.0, + "buy_lg_amount": 10000.0, + "sell_lg_amount": 10000.0, + "buy_elg_amount": 10000.0, + "sell_elg_amount": 10000.0, + "net_mf_amount": 10000.0, + }, + "auction": {"vol": 100.0, "float_share": 10000.0}, + "limit_events": {"limit_amount": 10000.0, "float_mv": 10000.0, "total_mv": 10000.0}, + "dragon_tiger": {"buy_amount": 10000.0, "sell_amount": 10000.0, "net_amount": 10000.0}, +} + +LIVE_TTL = { + "index_member_all": 6 * 3600, + "stk_limit": 3600, + "suspend_d": 6 * 3600, + "adj_factor": 3600, + "hm_list": 24 * 3600, + "ths_index": 24 * 3600, + "ths_member": 6 * 3600, + "stk_mins": 20, + "top_list": 3600, + "top_inst": 3600, +} + +BLOCKED_LIVE_APIS = {"rt_sw_k"} + + +def steward_query(api, body: dict[str, Any]) -> dict[str, Any]: + api_name = str(body.get("api_name") or "").strip() + params = body.get("params") if isinstance(body.get("params"), dict) else {} + fields = str(body.get("fields") or "") + if not api_name: + raise ApiError("INVALID_ARGUMENT", "api_name is required") + if api_name in BLOCKED_LIVE_APIS: + raise ApiError("INVALID_ARGUMENT", "rt_sw_k is disabled; use published sw_daily or free Shenwan realtime") + if api_name == "rt_k": + return _realtime_quotes(api, params, fields) + if api_name == "rt_idx_k": + return _realtime_index_quotes(api, params, fields) + dataset = API_TO_DATASET.get(api_name) + if dataset: + published = _try_published(api, api_name, dataset, params, fields) + if published is not None: + return published + rows = _live_tushare(api, api_name, params, fields) + return envelope( + _project(rows, fields), + { + "tier": "live", + "source": "tushare", + "stale": False, + "staleness_seconds": 0, + "row_shape": "tushare", + "published_at": isoformat(now_shanghai()), + }, + ) + + +def _try_published(api, api_name: str, dataset: str, params: dict[str, Any], fields: str) -> dict[str, Any] | None: + fetcher = DATASET_FETCHER.get(dataset) + if fetcher is None: + return None + query = _hub_query(params) + if dataset == "popularity": + query["source"] = "ths" if api_name == "ths_hot" else "dc" + if dataset == "sector_daily": + query["family"] = {"ths_daily": "ths", "dc_index": "dc", "sw_daily": "sw"}.get(api_name, "") + if dataset == "limit_events": + limit_type = str(params.get("limit_type") or "").strip().upper() + if limit_type: + query["limit_type"] = limit_type + if dataset == "calendar" and not (query.get("from") and query.get("to")): + start = str(params.get("start_date") or params.get("from") or "") + end = str(params.get("end_date") or params.get("to") or start) + if not start or not end: + return None + query = {"from": start, "to": end} + try: + payload = fetcher(api, query) + except ApiError as exc: + if exc.code in {"DATASET_NOT_PUBLISHED", "STALE_DATA", "INVALID_ARGUMENT"}: + return None + raise + rows = list(payload.get("data") or []) + if dataset == "stocks": + rows = _filter_stocks(rows, params) + if dataset == "calendar": + rows = _filter_calendar(rows, params) + native = _to_tushare_native(dataset, rows) + meta = dict(payload.get("meta") or {}) + meta["row_shape"] = "tushare" + meta["source"] = str(meta.get("source") or "datahub") + return envelope(_project(native, fields), meta) + + +def _realtime_quotes(api, params: dict[str, Any], fields: str) -> dict[str, Any]: + codes = [item.strip() for item in str(params.get("ts_code") or params.get("codes") or "").split(",") if item.strip()] + try: + payload = fetch_quotes(api.db, codes) if codes else fetch_market_quotes(api.db) + except RealtimeApiError as exc: + raise ApiError(exc.code, exc.message) from exc + rows = [_quote_to_rt_k(item) for item in (payload.get("data") or []) if isinstance(item, dict)] + rows = [item for item in rows if item] + meta = dict(payload.get("meta") or {}) + meta["row_shape"] = "tushare" + return envelope(_project(rows, fields), meta) + + +def _realtime_index_quotes(api, params: dict[str, Any], fields: str) -> dict[str, Any]: + try: + payload = fetch_index_quotes(api.db) + except RealtimeApiError as exc: + raise ApiError(exc.code, exc.message) from exc + wanted = { + item.strip() + for item in str(params.get("ts_code") or "").split(",") + if item.strip() + } + rows = [] + for item in payload.get("data") or []: + if not isinstance(item, dict): + continue + converted = _quote_to_rt_k(item) + if not converted: + continue + if wanted and converted.get("ts_code") not in wanted and str(item.get("code") or "") not in { + code.split(".")[0] for code in wanted + }: + continue + rows.append(converted) + meta = dict(payload.get("meta") or {}) + meta["row_shape"] = "tushare" + return envelope(_project(rows, fields), meta) + + +def _live_tushare(api, api_name: str, params: dict[str, Any], fields: str) -> list[dict[str, Any]]: + wanted_fields = fields or TUSHARE_FIELDS.get(api_name, "") + cache_key = _live_cache_key(api_name, params, wanted_fields) + ttl = LIVE_TTL.get(api_name, 1800) + cached = _read_cache(api.db, cache_key) + if cached is not None: + data = cached.get("data") + if isinstance(data, list): + return [dict(item) for item in data if isinstance(item, dict)] + pipeline = api.pipeline + if not pipeline.breaker.allow(): + recovered = _live_lkg(api.db, cache_key) + if recovered is not None: + return recovered + raise ApiError("SOURCE_UNAVAILABLE", "Tushare circuit open") + pipeline.bucket.acquire() + try: + rows = pipeline.adapter.query_raw(api_name, dict(params), wanted_fields) + pipeline.breaker.record_success() + except Exception as exc: + pipeline.breaker.record_failure(str(exc)) + recovered = _live_lkg(api.db, cache_key) + if recovered is not None: + return recovered + raise ApiError("SOURCE_UNAVAILABLE", f"Tushare {api_name} unavailable: {exc}") from exc + payload = envelope( + rows, + { + "tier": "live", + "source": "tushare", + "stale": False, + "staleness_seconds": 0, + "row_shape": "tushare", + "published_at": isoformat(now_shanghai()), + }, + ) + _write_cache(api.db, cache_key, payload, ttl, "tushare") + return rows + + +def _live_lkg(db, cache_key: str) -> list[dict[str, Any]] | None: + from datahub.governance.lkg import LastKnownGood + + item = LastKnownGood(db).load(cache_key) + payload = item.get("payload") if item else None + if not isinstance(payload, dict): + return None + data = payload.get("data") + if not isinstance(data, list) or not data: + return None + return [dict(row) for row in data if isinstance(row, dict)] + + +def _live_cache_key(api_name: str, params: dict[str, Any], fields: str) -> str: + packed = json.dumps({"api": api_name, "params": params, "fields": fields}, sort_keys=True, ensure_ascii=False) + digest = hashlib.sha1(packed.encode("utf-8")).hexdigest() + return f"steward:{api_name}:{digest}" + + +def _hub_query(params: dict[str, Any], **extra: Any) -> dict[str, str]: + query = {key: str(value) for key, value in extra.items() if value not in (None, "")} + date = yyyymmdd(params.get("trade_date") or params.get("date") or "") + start = yyyymmdd(params.get("start_date") or params.get("from") or date) + end = yyyymmdd(params.get("end_date") or params.get("to") or date) + code = str(params.get("ts_code") or params.get("code") or "").strip() + if code: + query["code"] = code + if date and not (params.get("start_date") or params.get("end_date")): + query["date"] = date + else: + if start: + query["from"] = start + if end: + query["to"] = end + return query + + +def _to_tushare_native(dataset: str, rows: list[dict[str, Any]]) -> list[dict[str, Any]]: + scales = SCALE_TO_TUSHARE.get(dataset) or {} + converted: list[dict[str, Any]] = [] + for row in rows: + item = dict(row) + if item.get("vol") in (None, ""): + item["vol"] = item.get("volume") + item.pop("volume", None) + for field, factor in scales.items(): + if field in item and item[field] not in (None, ""): + number = finite_number(item.get(field)) + item[field] = number / factor if factor else number + if dataset == "popularity" and item.get("ts_name") and not item.get("name"): + item["name"] = item.get("ts_name") + if dataset == "dragon_tiger" and item.get("ts_name") and not item.get("name"): + item["name"] = item.get("ts_name") + if dataset == "sector_daily" and item.get("pct_change") is not None and item.get("pct_chg") is None: + item["pct_chg"] = item.get("pct_change") + if dataset == "calendar": + item["is_open"] = 1 if item.get("is_open") in (True, 1, "1", "Y", "y") else 0 + converted.append(item) + return converted + + +def _quote_to_rt_k(row: dict[str, Any]) -> dict[str, Any] | None: + ts_code = str(row.get("ts_code") or "").strip() + close = finite_number(row.get("close") if row.get("close") not in (None, "") else row.get("price")) + previous = finite_number( + row.get("pre_close") if row.get("pre_close") not in (None, "") else row.get("previous_close") + ) + if not ts_code or close <= 0: + return None + item = { + "ts_code": ts_code, + "name": row.get("name") or "", + "open": row.get("open"), + "high": row.get("high"), + "low": row.get("low"), + "close": close, + "pre_close": previous, + "vol": row.get("vol") if row.get("vol") not in (None, "") else row.get("volume"), + "amount": row.get("amount"), + "pct_chg": row.get("pct_chg") if row.get("pct_chg") not in (None, "") else row.get("change"), + "trade_time": row.get("quote_time") or row.get("trade_time") or "", + "quote_date": row.get("quote_date") or "", + "source": row.get("source") or "", + "delayed": bool(row.get("delayed")), + "delay_seconds": row.get("delay_seconds") or 0, + "delay_notice": row.get("delay_notice") or "", + } + return item + + +def _filter_stocks(rows: list[dict[str, Any]], params: dict[str, Any]) -> list[dict[str, Any]]: + ts_code = str(params.get("ts_code") or "").strip().upper() + status = str(params.get("list_status") or "").strip() + name = str(params.get("name") or "").strip() + filtered = rows + if ts_code: + filtered = [row for row in filtered if str(row.get("ts_code") or "").upper() == ts_code] + if status: + filtered = [row for row in filtered if str(row.get("list_status") or status) == status] + if name: + filtered = [row for row in filtered if name.casefold() in str(row.get("name") or "").casefold()] + return filtered + + +def _filter_calendar(rows: list[dict[str, Any]], params: dict[str, Any]) -> list[dict[str, Any]]: + start = yyyymmdd(params.get("start_date") or params.get("from") or "") + end = yyyymmdd(params.get("end_date") or params.get("to") or start) + if start and end: + rows = [row for row in rows if start <= yyyymmdd(row.get("cal_date")) <= end] + if params.get("is_open") in (1, "1", True): + rows = [row for row in rows if int(row.get("is_open") or 0) == 1] + return rows + + +def _project(rows: list[dict[str, Any]], fields: str) -> list[dict[str, Any]]: + keys = [item.strip() for item in str(fields or "").split(",") if item.strip()] + if not keys: + return rows + return [{key: row.get(key) for key in keys} for row in rows] diff --git a/xiaobai-datahub/tests/fixtures.py b/xiaobai-datahub/tests/fixtures.py index a9f735d..e137834 100644 --- a/xiaobai-datahub/tests/fixtures.py +++ b/xiaobai-datahub/tests/fixtures.py @@ -1,5 +1,6 @@ from __future__ import annotations +import copy import sys from pathlib import Path @@ -69,13 +70,29 @@ RAW = { "sw_daily": [ {"ts_code": "801780.SI", "trade_date": "20240902", "name": "银行", "open": 2000, "high": 2010, "low": 1990, "close": 2005, "pct_change": 0.8, "vol": 50, "amount": 2e8}, ], + "stk_limit": [ + {"ts_code": "600000.SH", "trade_date": "20240902", "up_limit": 11.22, "down_limit": 9.18}, + {"ts_code": "000001.SZ", "trade_date": "20240902", "up_limit": 12.21, "down_limit": 9.99}, + ], + "index_member_all": [ + { + "l2_code": "801780.SI", + "l2_name": "银行", + "ts_code": "600000.SH", + "name": "浦发银行", + "in_date": "20140101", + "out_date": "", + "is_new": "Y", + } + ], + "hm_list": [{"name": "测试游资", "desc": "测试", "orgs": "某某营业部"}], } def fake_transport(api_name: str, params: dict, fields: str): if api_name == "index_daily": code = params.get("ts_code") - rows = [row for row in RAW["index_daily"] if row["ts_code"] == code] + rows = [copy.deepcopy(row) for row in RAW["index_daily"] if row["ts_code"] == code] trade_date = str(params.get("trade_date") or "") start = str(params.get("start_date") or "") end = str(params.get("end_date") or "") @@ -89,8 +106,8 @@ def fake_transport(api_name: str, params: dict, fields: str): if api_name == "trade_cal": start = str(params.get("start_date") or "") end = str(params.get("end_date") or "99999999") - return [row for row in RAW["trade_cal"] if start <= row["cal_date"] <= end] - rows = list(RAW.get(api_name) or []) + return [copy.deepcopy(row) for row in RAW["trade_cal"] if start <= row["cal_date"] <= end] + rows = copy.deepcopy(list(RAW.get(api_name) or [])) if api_name == "limit_list_d": limit_type = str(params.get("limit_type") or "") if limit_type: diff --git a/xiaobai-datahub/tests/test_steward.py b/xiaobai-datahub/tests/test_steward.py new file mode 100644 index 0000000..bed5472 --- /dev/null +++ b/xiaobai-datahub/tests/test_steward.py @@ -0,0 +1,111 @@ +from __future__ import annotations + +import tempfile +import unittest +from pathlib import Path +from unittest.mock import patch + +from datahub.adapters.tushare import TushareAdapter +from datahub.crypto import SecretVault +from datahub.hub import Hub +from datahub.settings import Settings +from tests.fixtures import TRADE_DATE, fake_transport + + +class StewardQueryTests(unittest.TestCase): + def setUp(self) -> None: + self.tmp = tempfile.TemporaryDirectory() + settings = Settings( + host="127.0.0.1", + port=0, + encryption_key=SecretVault.generate_key(), + api_token="k" * 32, + admin_password="StartPass1", + tushare_token="tushare-secret-token-xyz", + db_path=Path(self.tmp.name) / "hub.db", + backup_dir=Path(self.tmp.name) / "backups", + scheduler_enabled=False, + quality={"daily_row_ratio": 0.5, "null_rate_max": 0.5, "list_limit_default": 5000, "list_limit_max": 5000}, + ) + adapter = TushareAdapter("tushare-secret-token-xyz", transport=fake_transport) + self.hub = Hub(settings, adapter=adapter) + self.hub.pipeline.ingest_reference(TRADE_DATE) + for dataset in ("daily", "valuation", "moneyflow", "auction", "index_daily"): + self.hub.pipeline.run_dataset(dataset, TRADE_DATE) + + def tearDown(self) -> None: + self.hub.stop() + self.tmp.cleanup() + + def test_published_daily_is_tushare_native(self) -> None: + payload = self.hub.api.query_api( + {"api_name": "daily", "params": {"trade_date": TRADE_DATE}, "fields": "ts_code,close,vol,amount"} + ) + rows = payload["data"] + by_code = {row["ts_code"]: row for row in rows} + self.assertEqual(by_code["600000.SH"]["vol"], 1000.0) + self.assertEqual(by_code["600000.SH"]["amount"], 2000.0) + self.assertEqual(payload["meta"]["row_shape"], "tushare") + + def test_live_stk_limit_uses_internal_tushare(self) -> None: + payload = self.hub.api.query_api( + {"api_name": "stk_limit", "params": {"trade_date": TRADE_DATE}, "fields": "ts_code,up_limit,down_limit"} + ) + self.assertEqual(payload["meta"]["source"], "tushare") + self.assertEqual(payload["data"][0]["ts_code"], "600000.SH") + + def test_rt_sw_k_is_blocked(self) -> None: + from datahub.serving import ApiError + + with self.assertRaises(ApiError): + self.hub.api.query_api({"api_name": "rt_sw_k", "params": {"ts_code": "801074.SI"}}) + + def test_rt_k_uses_free_quotes_not_tushare(self) -> None: + quotes = [ + { + "ts_code": "600000.SH", + "name": "浦发银行", + "close": 10.2, + "pre_close": 10.0, + "open": 10.1, + "high": 10.3, + "low": 9.9, + "vol": 1000, + "amount": 2000000, + } + ] + with patch("datahub.steward.fetch_quotes", return_value={"data": quotes, "meta": {"source": "eastmoney:ulist", "stale": False}}): + payload = self.hub.api.query_api({"api_name": "rt_k", "params": {"ts_code": "600000.SH"}}) + self.assertEqual(payload["data"][0]["close"], 10.2) + self.assertEqual(payload["meta"]["source"], "eastmoney:ulist") + + def test_shenwan_quote_uses_eastmoney_90_prefix(self) -> None: + from datahub.adapters.eastmoney import EastmoneyAdapter + + with patch.object(EastmoneyAdapter, "_get_json") as get_json: + get_json.return_value = { + "data": { + "diff": [ + { + "f12": "801074", + "f14": "工业金属", + "f2": 1234.5, + "f3": 2.88, + "f18": 1200, + "f17": 1205, + "f15": 1240, + "f16": 1198, + "f6": 1, + "f124": 1757319000, + } + ] + } + } + quote = EastmoneyAdapter().fetch_shenwan_quote("801074.SI") + self.assertEqual(quote["source"], "eastmoney_sw") + self.assertAlmostEqual(quote["change"], 2.88) + self.assertEqual(get_json.call_args.args[1]["secids"], "90.801074") + + +if __name__ == "__main__": + unittest.main()