from __future__ import annotations import io import json import logging import threading import unittest from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer from urllib.parse import parse_qs, urlparse from backend.data.datahub.client import DatahubClient from backend.data.datahub.errors import DatahubError from backend.data.datahub.redact import redact_text from backend.data.datahub.settings import DatahubSettings TOKEN = "super-secret-datahub-token" class FakeHubState: def __init__(self) -> None: self.mode = "ok" self.hits = 0 self.paths: list[str] = [] STATE = FakeHubState() class FakeHubHandler(BaseHTTPRequestHandler): def log_message(self, format: str, *args: object) -> None: return def do_GET(self) -> None: # noqa: N802 STATE.hits += 1 parsed = urlparse(self.path) STATE.paths.append(parsed.path) token = self.headers.get("X-Datahub-Token", "") if STATE.mode == "timeout": raise TimeoutError("simulated timeout") if token != TOKEN: self._json(401, {"error": {"code": "UNAUTHORIZED", "message": "missing or invalid X-Datahub-Token"}}) return if STATE.mode == "unpublished": self._json(404, {"error": {"code": "DATASET_NOT_PUBLISHED", "message": "daily 19990101 尚未发布", "expected_at": "15:05+08:00"}}) return if STATE.mode == "empty": self._json(200, {"schema_version": 1, "data": [], "meta": {"tier": "official", "stale": False, "staleness_seconds": 0}}) return if STATE.mode == "stale": self._json(200, {"schema_version": 1, "data": [{"ts_code": "600000.SH", "trade_date": "20240902", "close": 10.2, "volume": 100000, "amount": 2000000}], "meta": {"tier": "official", "stale": True, "staleness_seconds": 999999}}) return if STATE.mode == "invalid": self.send_response(200) self.send_header("Content-Type", "application/json") self.end_headers() self.wfile.write(b"not-json") return if parsed.path == "/v1/health": self._json(200, {"schema_version": 1, "data": {"status": "ok"}, "meta": {"tier": "official", "source": "datahub", "stale": False, "staleness_seconds": 0}}) return if parsed.path == "/v1/calendar": self._json(200, {"schema_version": 1, "data": [{"cal_date": "20240902", "is_open": True, "pretrade_date": "20240830"}], "meta": {"tier": "official", "trade_date": "20240902", "stale": False, "staleness_seconds": 0}}) return if parsed.path == "/v1/bars/daily": query = {key: values[-1] for key, values in parse_qs(parsed.query).items()} self._json(200, { "schema_version": 1, "data": [{ "ts_code": "600000.SH", "trade_date": query.get("date") or "20240902", "open": 10.11, "high": 10.25, "low": 10.01, "close": 10.20, "pct_chg": 1.2345, "volume": 100000.0, "amount": 2000000.0, "adj_factor": 1.1, }], "meta": {"tier": "official", "trade_date": "20240902", "stale": False, "staleness_seconds": 0, "source": "tushare:daily"}, }) return if parsed.path == "/v1/datasets/status": self._json(200, {"schema_version": 1, "data": [{"dataset": "daily", "state": "published", "trade_date": "20240902"}], "meta": {"tier": "official", "stale": False, "staleness_seconds": 0}}) return self._json(400, {"error": {"code": "INVALID_ARGUMENT", "message": f"unknown endpoint: {parsed.path}"}}) def _json(self, status: int, payload: dict) -> None: body = json.dumps(payload).encode("utf-8") self.send_response(status) self.send_header("Content-Type", "application/json; charset=utf-8") self.send_header("Content-Length", str(len(body))) self.end_headers() self.wfile.write(body) class DatahubClientTests(unittest.TestCase): @classmethod def setUpClass(cls) -> None: cls.server = ThreadingHTTPServer(("127.0.0.1", 0), FakeHubHandler) cls.thread = threading.Thread(target=cls.server.serve_forever, daemon=True) cls.thread.start() cls.base = f"http://127.0.0.1:{cls.server.server_address[1]}" @classmethod def tearDownClass(cls) -> None: cls.server.shutdown() cls.server.server_close() def setUp(self) -> None: STATE.mode = "ok" STATE.hits = 0 STATE.paths = [] self.client = DatahubClient(DatahubSettings(base_url=self.base, token=TOKEN, retries=1, timeout_seconds=2)) def test_health_envelope(self) -> None: response = self.client.health() self.assertEqual(response.schema_version, 1) self.assertEqual(response.data["status"], "ok") self.assertIn("stale", response.meta) def test_missing_and_bad_token_401(self) -> None: missing = DatahubClient(DatahubSettings(base_url=self.base, token="")) with self.assertRaises(DatahubError) as raised: missing.health() self.assertEqual(raised.exception.code, "NOT_CONFIGURED") bad = DatahubClient(DatahubSettings(base_url=self.base, token="wrong")) with self.assertRaises(DatahubError) as raised: bad.health() self.assertEqual(raised.exception.code, "UNAUTHORIZED") self.assertNotIn(TOKEN, str(raised.exception)) def test_unpublished_and_empty_and_stale_codes(self) -> None: STATE.mode = "unpublished" with self.assertRaises(DatahubError) as raised: self.client.daily_bars(date="19990101") self.assertEqual(raised.exception.code, "DATASET_NOT_PUBLISHED") STATE.mode = "empty" response = self.client.daily_bars(date="20240902") self.assertEqual(response.data, []) STATE.mode = "stale" stale = self.client.daily_bars(date="20240902") self.assertTrue(stale.meta["stale"]) def test_invalid_json_maps_to_internal(self) -> None: STATE.mode = "invalid" with self.assertRaises(DatahubError) as raised: self.client.health() self.assertEqual(raised.exception.code, "INTERNAL") def test_timeout_maps_and_retries(self) -> None: hits = {"n": 0} def boom(_request, timeout=None): hits["n"] += 1 raise TimeoutError("late") client = DatahubClient( DatahubSettings(base_url=self.base, token=TOKEN, retries=1, timeout_seconds=1), urlopen=boom, ) with self.assertRaises(DatahubError) as raised: client.health() self.assertEqual(raised.exception.code, "TIMEOUT") self.assertEqual(hits["n"], 2) def test_token_never_appears_in_error_text_or_logs(self) -> None: stream = io.StringIO() logger = logging.getLogger("xiaobai.datahub") handler = logging.StreamHandler(stream) logger.addHandler(handler) logger.setLevel(logging.DEBUG) try: with self.assertRaises(DatahubError): DatahubClient(DatahubSettings(base_url=self.base, token="wrong")).health() blob = stream.getvalue() + redact_text("header " + TOKEN, (TOKEN,)) self.assertNotIn(TOKEN, blob) self.assertIn("***", redact_text(TOKEN, (TOKEN,))) finally: logger.removeHandler(handler) def test_calendar_and_status_contract(self) -> None: calendar = self.client.calendar("20240901", "20240902") self.assertEqual(calendar.data[0]["cal_date"], "20240902") status = self.client.dataset_status("20240902") self.assertEqual(status.data[0]["dataset"], "daily") if __name__ == "__main__": unittest.main()