from __future__ import annotations import unittest from tushare_client import TushareClient class FakeRealtimeClient(TushareClient): def query(self, api_name, params=None, fields=""): params = params or {} if api_name == "trade_cal": return [ { "cal_date": "20260720", "is_open": 1, "pretrade_date": "20260717", } ] if api_name == "stock_basic": return [ {"ts_code": "000001.SZ", "name": "甲", "industry": "银行"}, {"ts_code": "000002.SZ", "name": "乙", "industry": "地产"}, {"ts_code": "000003.SZ", "name": "丙", "industry": "元器件"}, ] if api_name == "stk_limit": return [ {"ts_code": "000001.SZ", "up_limit": 11.0, "down_limit": 9.0}, {"ts_code": "000002.SZ", "up_limit": 22.0, "down_limit": 18.0}, {"ts_code": "000003.SZ", "up_limit": 33.0, "down_limit": 27.0}, ] if api_name == "daily_basic": requested_code = str(params.get("ts_code") or "") rows = [ {"ts_code": "000001.SZ", "trade_date": "20260717", "float_share": 1000}, {"ts_code": "000002.SZ", "trade_date": "20260717", "float_share": 2000}, {"ts_code": "000003.SZ", "trade_date": "20260717", "float_share": 3000}, ] return [row for row in rows if not requested_code or row["ts_code"] == requested_code] if api_name == "daily": return [ {"ts_code": params.get("ts_code"), "trade_date": "20260713", "vol": 1000, "amount": 1}, {"ts_code": params.get("ts_code"), "trade_date": "20260714", "vol": 1000, "amount": 1}, {"ts_code": params.get("ts_code"), "trade_date": "20260715", "vol": 1000, "amount": 1}, {"ts_code": params.get("ts_code"), "trade_date": "20260716", "vol": 1000, "amount": 1}, {"ts_code": params.get("ts_code"), "trade_date": "20260717", "vol": 1000, "amount": 1}, ] if api_name == "limit_list_d": return [ { "ts_code": "000001.SZ", "name": "甲", "industry": "银行", "close": 10.0, "pct_chg": 10.0, "amount": 100000000, "limit_times": 2, } ] if api_name == "rt_k": rows = [ { "ts_code": "000001.SZ", "name": "甲", "pre_close": 10.0, "open": 10.1, "high": 11.0, "low": 10.0, "close": 11.0, "vol": 1000, "amount": 100000000, "num": 10, }, { "ts_code": "000002.SZ", "name": "乙", "pre_close": 20.0, "open": 19.5, "high": 20.0, "low": 18.0, "close": 18.0, "vol": 2000, "amount": 200000000, "num": 20, }, { "ts_code": "000003.SZ", "name": "丙", "pre_close": 30.0, "open": 31.0, "high": 33.0, "low": 30.0, "close": 32.0, "vol": 3000, "amount": 300000000, "num": 30, }, ] requested = { code for code in str(params.get("ts_code") or "").split(",") if code } return [row for row in rows if row["ts_code"] in requested] raise AssertionError(f"Unexpected API call: {api_name} {params}") class RealtimeDashboardTests(unittest.TestCase): def setUp(self): TushareClient._realtime_reference_cache.clear() TushareClient._capital_cache.clear() TushareClient._latest_realtime_market.clear() TushareClient._stock_activity_cache.clear() self.client = FakeRealtimeClient("test-token") def test_realtime_dashboard_classifies_pools_and_units(self): dashboard = self.client._realtime_dashboard("20260720", "20260720", "20260717") self.assertTrue(dashboard["meta"]["realtime"]) 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"]["broken_count"], 1) self.assertEqual(dashboard["overview"]["amount_billion"], 6.0) self.assertEqual(dashboard["limits"][0]["streak"], 3) self.assertEqual(dashboard["limits"][0]["amount_billion"], 1.0) def test_realtime_stock_quote_uses_cached_industry(self): self.client._load_realtime_reference("20260720", "20260717") quote = self.client.realtime_stock_quote("000003.SZ") self.assertEqual(quote["name"], "丙") self.assertEqual(quote["sector"], "元器件") self.assertAlmostEqual(quote["change"], 6.6667) self.assertEqual(quote["amount_billion"], 3.0) self.assertAlmostEqual(quote["turnover_rate"], 0.01) if __name__ == "__main__": unittest.main()