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_stock_filter_missing_from_active_snapshot_falls_back_inside_hub(self) -> None: calls = [] def transport(api_name, params, fields): calls.append((api_name, dict(params))) if api_name == "stock_basic" and params.get("list_status") == "D": return [ { "ts_code": "000627.SZ", "symbol": "000627", "name": "退市天茂", "list_status": "D", "list_date": "19961112", } ] return fake_transport(api_name, params, fields) self.hub.pipeline.adapter._transport = transport payload = self.hub.api.query_api( { "api_name": "stock_basic", "params": {"list_status": "D"}, "fields": "ts_code,name,list_status,list_date", } ) self.assertEqual(payload["meta"]["source"], "tushare") self.assertEqual(payload["data"][0]["list_status"], "D") self.assertIn(("stock_basic", {"list_status": "D"}), calls) 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_missing_published_sw_family_falls_back_inside_hub(self) -> None: self.hub.pipeline.run_eod_batch_e(TRADE_DATE) batch_id = self.hub.pipeline.active_batch("sector_daily", TRADE_DATE) with self.hub.db.write() as connection: connection.execute( "DELETE FROM eod_sector_daily WHERE trade_date = ? AND batch_id = ? AND family = 'sw'", (TRADE_DATE, batch_id), ) payload = self.hub.api.query_api( { "api_name": "sw_daily", "params": {"ts_code": "801780.SI", "trade_date": TRADE_DATE}, "fields": "ts_code,trade_date,name,pct_change", } ) self.assertEqual(payload["meta"]["source"], "tushare") self.assertEqual(payload["data"][0]["ts_code"], "801780.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()