From 8a7d1f3698c609b024ad441ee8a8f0a8a9ce4ed9 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E6=80=BB=E5=B7=A5?= Date: Tue, 8 Sep 2026 23:37:53 +0800 Subject: [PATCH] =?UTF-8?q?fix(HEL-494):=20=E7=BD=91=E7=AB=99=E5=B8=82?= =?UTF-8?q?=E5=9C=BA=E5=AE=A2=E6=88=B7=E7=AB=AF=E6=94=B9=E4=B8=BA=E7=BA=AF?= =?UTF-8?q?=E4=B8=AD=E6=9E=A2=20Facade=EF=BC=8C=E5=B9=B6=E8=BF=81=E7=A7=BB?= =?UTF-8?q?=20iFinD=20=E5=87=AD=E6=8D=AE=E5=88=B0=E4=B8=AD=E6=9E=A2?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 生产 gateway 不再读取 Tushare token 或实例化 TushareProvider/TushareClient;问财凭据经带鉴权的中枢接口加密入库,避免发版后 iFinD 未配置。 Co-authored-by: Cursor Co-authored-by: multica-agent --- backend/application.py | 4 ++ backend/data/datahub/bridge.py | 67 +++++++++++++++---------- backend/data/datahub/client.py | 9 ++++ backend/data/datahub/ifind_proxy.py | 18 +++++++ backend/data/gateway.py | 16 ++---- compose.datahub.yaml | 2 + config/architecture-inventory.json | 10 ++-- tests/test_data_gateway.py | 16 +++--- tests/test_datahub_bridge.py | 19 ++++---- tests/test_hub_exclusive.py | 76 +++++++++++++++++++++++++++-- tests/test_realtime_dashboard.py | 28 +++-------- xiaobai-datahub/compose.yaml | 2 + xiaobai-datahub/datahub/httpapp.py | 5 ++ xiaobai-datahub/datahub/hub.py | 12 +++++ xiaobai-datahub/tests/test_api.py | 34 +++++++++++++ 15 files changed, 235 insertions(+), 83 deletions(-) diff --git a/backend/application.py b/backend/application.py index d9b6121..1babcab 100644 --- a/backend/application.py +++ b/backend/application.py @@ -125,6 +125,10 @@ class DashboardService( ) self.data_gateway = self.container.data_gateway self.ifind = self.container.ifind + refresh = str(self._system_credentials.get("ifind_refresh_token") or "") + access = str(self._system_credentials.get("ifind_access_token") or "") + if refresh or access: + self.ifind.set_credentials(refresh, access) self.screener = self.container.screener self.strategy_tracking = self.container.strategy_tracking self.alert_service = self.container.alert_service diff --git a/backend/data/datahub/bridge.py b/backend/data/datahub/bridge.py index 0cb8bae..fccfe3b 100644 --- a/backend/data/datahub/bridge.py +++ b/backend/data/datahub/bridge.py @@ -2,7 +2,8 @@ from __future__ import annotations import logging import sys -from typing import Any, Callable +from threading import Lock +from typing import Any, Callable, ClassVar from backend.data.datahub.client import DatahubClient, DatahubResponse from backend.data.datahub.compare import compare_rows @@ -18,7 +19,13 @@ from backend.data.datahub.native import ( 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_daily import DailyMarketMixin +from backend.data.providers.tushare_dashboard import DashboardMixin +from backend.data.providers.tushare_dragon_tiger import DragonTigerMixin +from backend.data.providers.tushare_indices import IndexMixin +from backend.data.providers.tushare_industries import ShenwanIndustryMixin +from backend.data.providers.tushare_sectors import SectorMixin +from backend.data.providers.tushare_stocks import StockMixin from backend.data.providers.tushare_transport import TushareError LOGGER = logging.getLogger("xiaobai.datahub") @@ -304,11 +311,9 @@ class DatahubBridge: def query( self, api_name: str, - params: dict[str, Any] | None, - fields: str, - legacy_query: Callable[..., list[dict[str, Any]]], + params: dict[str, Any] | None = None, + fields: str = "", ) -> 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) @@ -542,21 +547,36 @@ def _finite(value: Any) -> float: return 0.0 -class DatahubAwareTushareClient: - def __init__(self, legacy: TushareClient, bridge: DatahubBridge) -> None: - self._legacy = legacy - self._bridge = bridge - # Mixins run as methods on the inner instance (dashboard / indices / - # getattr). Bind hub hooks and query onto that instance so real - # assembly cannot skip 8766. - self._legacy_query = legacy.query - legacy.query = self.query - 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 +class DatahubAwareTushareClient( + DashboardMixin, + IndexMixin, + ShenwanIndustryMixin, + SectorMixin, + DragonTigerMixin, + StockMixin, + DailyMarketMixin, +): + """Website market facade. Mixins call query(); query talks only to the hub.""" + + _realtime_reference_cache: ClassVar[dict[str, dict[str, Any]]] = {} + _realtime_reference_lock: ClassVar[Lock] = Lock() + _capital_cache: ClassVar[dict[str, dict[str, Any]]] = {} + _latest_realtime_market: ClassVar[dict[str, dict[str, Any]]] = {} + _stock_activity_cache: ClassVar[dict[str, dict[str, Any]]] = {} + _stock_listing_cache: ClassVar[dict[str, Any]] = {} + _stock_listing_lock: ClassVar[Lock] = Lock() + _suspension_cache: ClassVar[dict[str, dict[str, str] | None]] = {} + _suspension_lock: ClassVar[Lock] = Lock() + _sw_member_cache: ClassVar[dict[str, Any]] = {} + _sw_member_lock: ClassVar[Lock] = Lock() + + def __init__(self, first: Any, second: Any | None = None) -> None: + # Production: DatahubAwareTushareClient(bridge) + # Older tests: DatahubAwareTushareClient(unused_legacy, bridge) + self._bridge = second if second is not None else first + self.token = "datahub" + self.timeout = 30 + self.realtime_aggregator = None def query( self, @@ -564,7 +584,7 @@ class DatahubAwareTushareClient: params: dict[str, Any] | None = None, fields: str = "", ) -> list[dict[str, Any]]: - return self._bridge.query(api_name, params, fields, self._legacy_query) + return self._bridge.query(api_name, params, fields) def try_market_quotes(self, trade_date: str = "") -> list[dict[str, Any]] | None: return self._bridge.try_market_quotes(trade_date) @@ -583,6 +603,3 @@ class DatahubAwareTushareClient: def record_datahub_legacy(self, dataset: str, source: str = "", error: str = "") -> None: self._bridge.record_legacy(dataset, source, error) - - def __getattr__(self, name: str) -> Any: - return getattr(self._legacy, name) diff --git a/backend/data/datahub/client.py b/backend/data/datahub/client.py index 5d04aa1..496597d 100644 --- a/backend/data/datahub/client.py +++ b/backend/data/datahub/client.py @@ -96,6 +96,15 @@ class DatahubClient: {"api_name": api_name, "params": params or {}, "fields": fields}, ) + def put_ifind_credentials(self, refresh_token: str, access_token: str = "") -> DatahubResponse: + return self.post( + "/v1/credentials/ifind", + { + "ifind_refresh_token": refresh_token, + "ifind_access_token": access_token, + }, + ) + def sector_quote(self, code: str, date: str = "") -> DatahubResponse: payload: dict[str, Any] = {"code": code} if date: diff --git a/backend/data/datahub/ifind_proxy.py b/backend/data/datahub/ifind_proxy.py index 2b1cbc1..0e6c061 100644 --- a/backend/data/datahub/ifind_proxy.py +++ b/backend/data/datahub/ifind_proxy.py @@ -1,5 +1,6 @@ from __future__ import annotations +import logging import time from typing import Any @@ -7,6 +8,8 @@ from backend.data.datahub.bridge import DatahubBridge from backend.data.datahub.errors import DatahubError from backend.data.providers.ifind_client import IfindError +LOGGER = logging.getLogger("xiaobai.datahub") + class HubIfindProxy: """Website-facing iFinD facade. Talks only to xiaobai-datahub.""" @@ -15,16 +18,20 @@ class HubIfindProxy: self._datahub = datahub self._status: dict[str, Any] | None = None self._status_at = 0.0 + self._pending: tuple[str, str] | None = None @property def configured(self) -> bool: return bool(self.status().get("configured")) def set_credentials(self, refresh_token: str, access_token: str = "") -> None: + self._pending = (str(refresh_token or ""), str(access_token or "")) self._status = None self._status_at = 0.0 + self._flush_credentials() def status(self) -> dict[str, Any]: + self._flush_credentials() now = time.monotonic() if self._status is not None and now - self._status_at < 30: return dict(self._status) @@ -132,7 +139,18 @@ class HubIfindProxy: "sample_time": str(payload[0].get("time") or "") if payload else "", } + def _flush_credentials(self) -> None: + pending = self._pending + if pending is None or not self._datahub.settings.token: + return + try: + self._datahub.client.put_ifind_credentials(pending[0], pending[1]) + self._pending = None + except DatahubError: + LOGGER.warning("datahub ifind credential push failed; will retry") + def _rows(self, api_name: str, params: dict[str, Any]) -> list[dict[str, Any]]: + self._flush_credentials() try: response = self._datahub.client.query_api(api_name, params) except DatahubError as exc: diff --git a/backend/data/gateway.py b/backend/data/gateway.py index 4c8d343..0d01b2f 100644 --- a/backend/data/gateway.py +++ b/backend/data/gateway.py @@ -10,9 +10,8 @@ from backend.data.datahub import DatahubAwareTushareClient, DatahubBridge, Datah from backend.data.datahub.ifind_proxy import HubIfindProxy from backend.data.datahub.realtime_proxy import HubRealtimeProxy from backend.data.policy import DataSourcePolicy -from backend.data.providers import IfindProvider, TushareProvider +from backend.data.providers import IfindProvider from backend.data.quality import DataQualityGate, QualityEvidence, QualityReport -from backend.data.providers.tushare_client import TushareClient from backend.features.market.charts import MarketChartClient @@ -20,7 +19,6 @@ from backend.features.market.charts import MarketChartClient class DataGateway: policy: DataSourcePolicy quality: DataQualityGate - tushare_provider: TushareProvider ifind_provider: IfindProvider chart_data: MarketChartClient realtime_observer: HubRealtimeProxy @@ -34,12 +32,10 @@ class DataGateway: self, dataset_id: str = "", usage: DataUsage = "calculation", - ) -> TushareClient: + ) -> DatahubAwareTushareClient: if dataset_id: self.policy.assert_allowed(dataset_id, "tushare", usage) - legacy = self.tushare_provider.client() - legacy.realtime_aggregator = None - return DatahubAwareTushareClient(legacy, self.datahub) + return DatahubAwareTushareClient(self.datahub) def dataset_status(self, trade_date: str) -> list[dict[str, Any]] | None: return self.datahub.dataset_status(trade_date) @@ -102,9 +98,8 @@ def build_data_gateway( tushare_token_supplier: Callable[[], str] | None = None, datahub_settings: DatahubSettings | None = None, ) -> DataGateway: - token_supplier = tushare_token_supplier or ( - lambda: str(credentials.get("tushare_token") or "") - ) + # Website Tushare tokens are not used to assemble market clients. + del tushare_token_supplier policy = DataSourcePolicy.load() settings = datahub_settings or DatahubSettings.load(credentials=credentials) datahub_client = DatahubClient(settings) @@ -113,7 +108,6 @@ def build_data_gateway( return DataGateway( policy=policy, quality=DataQualityGate.load(policy), - tushare_provider=TushareProvider(token_supplier), ifind_provider=IfindProvider(ifind), chart_data=MarketChartClient(datahub), realtime_observer=HubRealtimeProxy(datahub), diff --git a/compose.datahub.yaml b/compose.datahub.yaml index 786e78c..70d3a90 100644 --- a/compose.datahub.yaml +++ b/compose.datahub.yaml @@ -20,6 +20,8 @@ services: DATAHUB_TOKEN: "${DATAHUB_TOKEN:?DATAHUB_TOKEN must be set}" DATAHUB_ADMIN_PASSWORD: "${DATAHUB_ADMIN_PASSWORD:?DATAHUB_ADMIN_PASSWORD must be set}" TUSHARE_TOKEN: "${TUSHARE_TOKEN:-}" + IFIND_REFRESH_TOKEN: "${IFIND_REFRESH_TOKEN:-}" + IFIND_ACCESS_TOKEN: "${IFIND_ACCESS_TOKEN:-}" DATAHUB_DB_PATH: /app/data/datahub.db DATAHUB_BACKUP_DIR: /app/data/backups TZ: Asia/Shanghai diff --git a/config/architecture-inventory.json b/config/architecture-inventory.json index 5757de6..0e8c63e 100644 --- a/config/architecture-inventory.json +++ b/config/architecture-inventory.json @@ -636,16 +636,16 @@ "bytes": 8357, "lines": 116 }, + { + "path": "backend/application.py", + "bytes": 6997, + "lines": 182 + }, { "path": "backend/features/screener/formula.py", "bytes": 6983, "lines": 146 }, - { - "path": "backend/application.py", - "bytes": 6751, - "lines": 178 - }, { "path": "backend/features/market/insights_popularity.py", "bytes": 6739, diff --git a/tests/test_data_gateway.py b/tests/test_data_gateway.py index 908b2a9..8c82d84 100644 --- a/tests/test_data_gateway.py +++ b/tests/test_data_gateway.py @@ -35,21 +35,23 @@ class DataGatewayTests(unittest.TestCase): with self.assertRaises(DataPolicyError): policy.assert_allowed("market.level2", "unresolved", "display") - def test_gateway_uses_live_token_supplier_and_hub_proxies(self) -> None: - token = {"value": "first"} + def test_gateway_uses_hub_facade_and_proxies(self) -> None: gateway = build_data_gateway( {"ifind_refresh_token": "refresh", "ifind_access_token": "access"}, - lambda: token["value"], + lambda: "must-not-be-used", ) - self.assertEqual(gateway.tushare().token, "first") - token["value"] = "second" - self.assertEqual(gateway.tushare().token, "second") + client = gateway.tushare() + self.assertEqual(client.token, "datahub") + self.assertIsNone(client.realtime_aggregator) + self.assertFalse(hasattr(client, "_legacy")) self.assertIs(gateway.ifind, gateway.ifind_provider.client) self.assertIs(gateway.chart_data.datahub, gateway.datahub) self.assertIsNone(gateway.chart_data.ifind) + from backend.data.datahub.bridge import DatahubAwareTushareClient from backend.data.datahub.ifind_proxy import HubIfindProxy from backend.data.datahub.realtime_proxy import HubRealtimeProxy + self.assertIsInstance(client, DatahubAwareTushareClient) self.assertIsInstance(gateway.ifind, HubIfindProxy) self.assertIsInstance(gateway.realtime_observer, HubRealtimeProxy) @@ -70,7 +72,6 @@ class DataGatewayTests(unittest.TestCase): "IfindProvider": {"backend/data/gateway.py"}, "MarketChartClient": {"backend/data/gateway.py"}, "TushareClient": {"backend/features/market/service.py"}, - "TushareProvider": {"backend/data/gateway.py"}, "DatahubClient": {"backend/data/gateway.py"}, "DatahubAwareTushareClient": {"backend/data/gateway.py"}, "DatahubBridge": {"backend/data/gateway.py"}, @@ -82,6 +83,7 @@ class DataGatewayTests(unittest.TestCase): "IfindHttpClient": set(), "EastmoneyChartClient": set(), "WebRealtimeAggregator": set(), + "TushareProvider": set(), } found_forbidden = {name: set() for name in forbidden} for path in (root / "backend").rglob("*.py"): diff --git a/tests/test_datahub_bridge.py b/tests/test_datahub_bridge.py index f3917ea..a390153 100644 --- a/tests/test_datahub_bridge.py +++ b/tests/test_datahub_bridge.py @@ -546,7 +546,7 @@ class DatahubBridgeTests(unittest.TestCase): self.assertEqual(chart[-1]["trade_date"], "2024-09-02") self.assertEqual(chart[-1]["close"], 10.4) - def test_gateway_tushare_assembly_binds_hooks_on_inner_client(self) -> None: + def test_gateway_tushare_facade_has_no_legacy_client(self) -> None: quotes = [ { "ts_code": f"{index:06d}.SZ", @@ -574,21 +574,20 @@ class DatahubBridgeTests(unittest.TestCase): ) gateway.datahub.client = hub_client wrapped = gateway.tushare() - inner = wrapped._legacy - self.assertTrue(callable(getattr(inner, "try_market_quotes", None))) - self.assertTrue(callable(getattr(inner, "try_index_quotes", None))) - self.assertTrue(callable(getattr(inner, "record_datahub_legacy", None))) - self.assertIs(inner.query.__self__, wrapped) - self.assertEqual(inner.query.__func__, wrapped.query.__func__) - self.assertFalse(hasattr(type(inner), "try_market_quotes")) - rows = inner.try_market_quotes("20240902") + self.assertFalse(hasattr(wrapped, "_legacy")) + self.assertIsNone(getattr(type(wrapped), "__getattr__", None)) + self.assertTrue(callable(getattr(type(wrapped), "try_market_quotes", None))) + self.assertTrue(callable(getattr(type(wrapped), "try_index_quotes", None))) + self.assertTrue(callable(getattr(type(wrapped), "record_datahub_legacy", None))) + self.assertTrue(callable(getattr(type(wrapped), "dashboard", None))) + rows = wrapped.try_market_quotes("20240902") self.assertGreaterEqual(len(rows or []), 200) self.assertIn("/v1/quotes/latest", hub_client.paths) hub_client.response = DatahubResponse( data=[dict(HUB_DAILY)], meta={"stale": False, "staleness_seconds": 0, "source": "tushare:daily"}, ) - daily = inner.query("daily", {"trade_date": "20240902"}, "ts_code,amount") + daily = wrapped.query("daily", {"trade_date": "20240902"}, "ts_code,amount") self.assertEqual(daily[0]["amount"], 2000.0) self.assertIn("/v1/bars/daily", hub_client.paths) diff --git a/tests/test_hub_exclusive.py b/tests/test_hub_exclusive.py index 0adc23a..faf9b57 100644 --- a/tests/test_hub_exclusive.py +++ b/tests/test_hub_exclusive.py @@ -2,6 +2,7 @@ from __future__ import annotations import ast import json +import re import unittest from pathlib import Path from unittest.mock import patch @@ -143,6 +144,26 @@ def hub_payload(request) -> dict: ], "meta": {"stale": False, "staleness_seconds": 0, "source": "tencent"}, } + if path == "/v1/auction": + return { + "schema_version": 1, + "data": [ + { + "ts_code": "600000.SH", + "trade_date": "20240902", + "close": 10.2, + "vol": 1000.0, + "amount": 2000.0, + } + ], + "meta": {"stale": False, "staleness_seconds": 0, "source": "datahub"}, + } + if path == "/v1/credentials/ifind": + return { + "schema_version": 1, + "data": {"configured": True, "access_ready": True, "access_expires_at": ""}, + "meta": {"source": "ifind"}, + } if path == "/v1/intraday/points": return { "schema_version": 1, @@ -184,7 +205,7 @@ def hub_payload(request) -> dict: ], "meta": {"source": "ifind"}, } - if api_name in {"daily", "rt_k"}: + if api_name in {"daily", "rt_k", "stk_auction"}: return { "schema_version": 1, "data": [{"ts_code": "600000.SH", "trade_date": "20240902", "close": 10.2, "amount": 2000.0}], @@ -254,10 +275,22 @@ class HubExclusiveWebsiteTests(unittest.TestCase): self.assertNotIn("IfindHttpClient", source) self.assertNotIn("EastmoneyChartClient", source) self.assertNotIn("WebRealtimeAggregator", source) + self.assertNotIn("TushareProvider", source) + self.assertIsNone(re.search(r"(? None: violations = [] @@ -273,7 +306,8 @@ class HubExclusiveWebsiteTests(unittest.TestCase): 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.assertNotIn("TushareProvider", gateway_src) + self.assertIsNone(re.search(r"(? None: @@ -298,13 +332,47 @@ class HubExclusiveWebsiteTests(unittest.TestCase): settings = _enabled_settings() with patch("urllib.request.urlopen", blocked_urlopen): gateway = build_data_gateway({"tushare_token": "tok"}, datahub_settings=settings) - gateway.datahub.client = DatahubClient(settings, urlopen=blocked_urlopen) + hub_client = DatahubClient(settings, urlopen=blocked_urlopen) + gateway.datahub.client = hub_client rows = gateway.ifind.wencai("涨停") quotes = gateway.realtime_observer.tencent_indices() chart = gateway.chart_data.stock_daily("600000", "20240902") + market = gateway.tushare() + market_quotes = market.try_quotes(["600000.SH"]) + auction = market.query("stk_auction", {"trade_date": "20240902"}, "") + gateway.ifind.set_credentials("refresh-token", "access-token") self.assertEqual(rows[0]["涨停原因"], "重组") self.assertEqual(len(quotes), 3) self.assertEqual(chart[-1]["close"], 10.2) + self.assertEqual(market_quotes[0]["close"], 10.2) + self.assertEqual(auction[0]["close"], 10.2) + self.assertIsNone(market.realtime_aggregator) + self.assertEqual(market.token, "datahub") + + def test_set_credentials_posts_to_hub_not_ifind(self) -> None: + seen: list[str] = [] + + def urlopen(request, timeout=None): + url = str(getattr(request, "full_url", None) or request) + seen.append(url) + if any(host in url for host in BLOCKED_HOSTS): + raise AssertionError(f"website opened blocked host: {url}") + return _Resp(hub_payload(request)) + + settings = _enabled_settings() + hub_client = DatahubClient(settings, urlopen=urlopen) + proxy = HubIfindProxy(DatahubBridge(settings, hub_client)) + proxy.set_credentials("refresh-token", "access-token") + self.assertTrue(any("/v1/credentials/ifind" in url for url in seen)) + self.assertFalse(any("51ifind.com" in url for url in seen)) + self.assertFalse(any("quantapi" in url for url in seen)) + + def test_compose_passes_ifind_env_to_hub(self) -> None: + overlay = (ROOT / "compose.datahub.yaml").read_text(encoding="utf-8") + standalone = (ROOT / "xiaobai-datahub" / "compose.yaml").read_text(encoding="utf-8") + for text in (overlay, standalone): + self.assertIn('IFIND_REFRESH_TOKEN: "${IFIND_REFRESH_TOKEN:-}"', text) + self.assertIn('IFIND_ACCESS_TOKEN: "${IFIND_ACCESS_TOKEN:-}"', text) if __name__ == "__main__": diff --git a/tests/test_realtime_dashboard.py b/tests/test_realtime_dashboard.py index 5f83631..39de1a4 100644 --- a/tests/test_realtime_dashboard.py +++ b/tests/test_realtime_dashboard.py @@ -380,10 +380,9 @@ class RealtimeDashboardTests(unittest.TestCase): def test_gateway_dashboard_uses_bound_market_quotes(self) -> None: from backend.data import build_data_gateway + from backend.data.datahub.bridge import DatahubAwareTushareClient from backend.data.datahub.client import DatahubResponse from backend.data.datahub.settings import DATASETS, DatahubSettings, DatasetFlags - from backend.data.gateway import DataGateway - from backend.data.providers.tushare import TushareProvider quotes = [ { @@ -442,30 +441,17 @@ class RealtimeDashboardTests(unittest.TestCase): datasets = {name: DatasetFlags(name) for name in DATASETS} datasets["quotes"] = DatasetFlags("quotes", read=True, shadow=False) settings = DatahubSettings(base_url="http://127.0.0.1:9", token="tok", datasets=datasets) - base = build_data_gateway({"tushare_token": "tok"}, datahub_settings=settings) - gateway = DataGateway( - policy=base.policy, - quality=base.quality, - tushare_provider=TushareProvider( - lambda: "tok", - client_factory=lambda token: FakeRealtimeClient(token), - ), - ifind_provider=base.ifind_provider, - chart_data=base.chart_data, - realtime_observer=base.realtime_observer, - datahub=base.datahub, - ) + gateway = build_data_gateway({"tushare_token": "tok"}, datahub_settings=settings) gateway.datahub.client = QuoteHub() wrapped = gateway.tushare() - inner = wrapped._legacy - inner.clock = lambda: datetime(2026, 7, 20, 10, 30, tzinfo=timezone(timedelta(hours=8))) - inner.realtime_aggregator = FakeFreeAggregator(fail=True) - TushareClient._realtime_reference_cache.clear() + wrapped.clock = lambda: datetime(2026, 7, 20, 10, 30, tzinfo=timezone(timedelta(hours=8))) + wrapped.realtime_aggregator = FakeFreeAggregator(fail=True) + DatahubAwareTushareClient._realtime_reference_cache.clear() dashboard = wrapped.dashboard("20260720") self.assertEqual(dashboard["meta"]["quote_source"], "datahub") self.assertIn("/v1/quotes/latest", gateway.datahub.client.calls) - self.assertTrue(callable(getattr(inner, "try_market_quotes", None))) - self.assertFalse(hasattr(type(inner), "try_market_quotes")) + self.assertTrue(callable(getattr(type(wrapped), "try_market_quotes", None))) + self.assertFalse(hasattr(wrapped, "_legacy")) if __name__ == "__main__": diff --git a/xiaobai-datahub/compose.yaml b/xiaobai-datahub/compose.yaml index d2a8c08..08601b0 100644 --- a/xiaobai-datahub/compose.yaml +++ b/xiaobai-datahub/compose.yaml @@ -14,6 +14,8 @@ services: DATAHUB_TOKEN: "${DATAHUB_TOKEN:?DATAHUB_TOKEN must be set}" DATAHUB_ADMIN_PASSWORD: "${DATAHUB_ADMIN_PASSWORD:?DATAHUB_ADMIN_PASSWORD must be set}" TUSHARE_TOKEN: "${TUSHARE_TOKEN:-}" + IFIND_REFRESH_TOKEN: "${IFIND_REFRESH_TOKEN:-}" + IFIND_ACCESS_TOKEN: "${IFIND_ACCESS_TOKEN:-}" DATAHUB_DB_PATH: /app/data/datahub.db DATAHUB_BACKUP_DIR: /app/data/backups TZ: Asia/Shanghai diff --git a/xiaobai-datahub/datahub/httpapp.py b/xiaobai-datahub/datahub/httpapp.py index c1147b8..8540b22 100644 --- a/xiaobai-datahub/datahub/httpapp.py +++ b/xiaobai-datahub/datahub/httpapp.py @@ -76,6 +76,11 @@ class HubRequestHandler(BaseHTTPRequestHandler): payload = self.hub.api.query_api(body) self._json(payload, HTTPStatus.OK) return + if path == "/v1/credentials/ifind" and method == "POST": + body = self._read_json() + payload = self.hub.put_ifind_credentials(body) + self._json(payload, HTTPStatus.OK) + return payload = self.hub.api.handle(path, parse_query(query)) self._json(payload, HTTPStatus.OK) diff --git a/xiaobai-datahub/datahub/hub.py b/xiaobai-datahub/datahub/hub.py index e9b20cf..ba0d924 100644 --- a/xiaobai-datahub/datahub/hub.py +++ b/xiaobai-datahub/datahub/hub.py @@ -1,6 +1,7 @@ from __future__ import annotations from pathlib import Path +from typing import Any from datahub.adapters.ifind import IfindAdapter from datahub.adapters.tushare import TushareAdapter @@ -52,6 +53,17 @@ class Hub: self.admin = AdminAPI(self.db, self.pipeline, self.scheduler, self.auth, ifind=self.ifind) self.static_dir = Path(__file__).resolve().parents[1] / "admin" + def put_ifind_credentials(self, body: dict[str, Any] | None) -> dict[str, Any]: + from datahub.serving import envelope + + payload = dict(body or {}) + refresh = str(payload.get("ifind_refresh_token") or "").strip() + access = str(payload.get("ifind_access_token") or "").strip() + self.auth.store_credential("ifind_refresh_token", refresh) + self.auth.store_credential("ifind_access_token", access) + self.ifind.set_credentials(refresh, access) + return envelope(self.ifind.status(), {"source": "ifind"}) + def start(self) -> None: if self.settings.scheduler_enabled: self.scheduler.start() diff --git a/xiaobai-datahub/tests/test_api.py b/xiaobai-datahub/tests/test_api.py index 6e33247..bb55d1d 100644 --- a/xiaobai-datahub/tests/test_api.py +++ b/xiaobai-datahub/tests/test_api.py @@ -61,6 +61,18 @@ class ApiContractTests(unittest.TestCase): self.hub.stop() self.tmp.cleanup() + def _post(self, path: str, body: dict, token: str | None = None) -> tuple[int, dict]: + headers = {"Content-Type": "application/json"} + if token is not None: + headers["X-Datahub-Token"] = token + raw = json.dumps(body).encode("utf-8") + req = Request(self.base + path, data=raw, headers=headers, method="POST") + try: + with urlopen(req, timeout=5) as resp: + return resp.status, json.loads(resp.read().decode()) + except HTTPError as exc: + return exc.code, json.loads(exc.read().decode()) + def _get(self, path: str, token: str | None = None) -> tuple[int, dict]: headers = {} if token is not None: @@ -145,6 +157,28 @@ class ApiContractTests(unittest.TestCase): self.assertNotIn("tushare-secret-token-xyz", blob) self.assertNotIn(self.token, blob) + def test_ifind_credentials_require_token_and_update_adapter(self) -> None: + status, body = self._post( + "/v1/credentials/ifind", + {"ifind_refresh_token": "refresh-secret", "ifind_access_token": "access-secret"}, + token=None, + ) + self.assertEqual(status, 401) + self.assertEqual(body["error"]["code"], "UNAUTHORIZED") + status, body = self._post( + "/v1/credentials/ifind", + {"ifind_refresh_token": "refresh-secret", "ifind_access_token": "access-secret"}, + token=self.token, + ) + self.assertEqual(status, 200, body) + self.assertTrue(body["data"]["configured"]) + self.assertTrue(body["data"]["access_ready"]) + self.assertTrue(self.hub.ifind.configured) + self.assertEqual(self.hub.auth.load_credential("ifind_refresh_token"), "refresh-secret") + blob = json.dumps(body) + self.assertNotIn("refresh-secret", blob) + self.assertNotIn("access-secret", blob) + if __name__ == "__main__": unittest.main()