from __future__ import annotations import tempfile import unittest from contextlib import contextmanager from pathlib import Path from datahub import observability from datahub.db import HubDB class _BrokenDB: """A db double whose write() always raises, to prove fail-open.""" @contextmanager def write(self): raise RuntimeError("disk is full") yield None # pragma: no cover - unreachable, keeps this a generator def fetchall(self, sql, params=()): raise RuntimeError("disk is full") def fetchone(self, sql, params=()): raise RuntimeError("disk is full") class ClassifyRowsTests(unittest.TestCase): def test_empty_list_is_flagged_empty(self): status, reason, age = observability.classify_rows([]) self.assertEqual(status, "empty") self.assertEqual(reason, "no_rows_returned") self.assertIsNone(age) def test_empty_dict_result_is_flagged_empty(self): status, reason, _ = observability.classify_rows({}) self.assertEqual(status, "empty") self.assertEqual(reason, "no_rows_returned") def test_missing_required_field_is_flagged(self): rows = [{"ts_code": "600000.SH", "close": 10.2}, {"ts_code": "000001.SZ"}] status, reason, _ = observability.classify_rows(rows, required_fields=("close",), freshness_field=None) self.assertEqual(status, "missing_fields") self.assertIn("close", reason) def test_fresh_rows_are_ok(self): import time rows = [{"ts_code": "600000.SH", "close": 10.2, "quote_time_epoch": int(time.time())}] status, reason, age = observability.classify_rows(rows) self.assertEqual(status, "ok") self.assertEqual(reason, "") self.assertIsNotNone(age) self.assertLess(age, 5) def test_stale_rows_are_flagged(self): import time rows = [{"ts_code": "600000.SH", "close": 10.2, "quote_time_epoch": int(time.time()) - 3600}] status, reason, age = observability.classify_rows(rows, max_age_seconds=300) self.assertEqual(status, "stale") self.assertEqual(reason, "data_age_exceeds_threshold") self.assertGreaterEqual(age, 3600 - 5) def test_classifier_never_raises_on_garbage_input(self): status, reason, age = observability.classify_rows(object()) self.assertEqual(status, "empty") self.assertIsNone(age) # Malformed rows inside a list must not raise either. status, _, _ = observability.classify_rows(["not-a-dict", 123, None]) self.assertEqual(status, "empty") class ClassifyErrorTests(unittest.TestCase): def test_blocked_page_markers_are_detected(self): status, reason = observability.classify_error("eastmoney request failed: Expecting value: line 1 column 1") self.assertEqual(status, "blocked") self.assertEqual(reason, "response_looks_like_intercept_page") def test_timeout_is_detected(self): status, _ = observability.classify_error("tencent request failed: timed out") self.assertEqual(status, "timeout") def test_generic_error_falls_back(self): status, reason = observability.classify_error("connection reset by peer") self.assertEqual(status, "error") self.assertEqual(reason, "") def test_never_raises_on_none(self): status, reason = observability.classify_error(None) self.assertEqual(status, "error") class RecordCallTests(unittest.TestCase): def setUp(self) -> None: self.tmp = tempfile.TemporaryDirectory() self.db = HubDB(Path(self.tmp.name) / "hub.db") def tearDown(self) -> None: self.tmp.cleanup() def test_record_call_writes_log_and_health(self): observability.record_call(self.db, "eastmoney", "indices", status="ok", latency_ms=42) log_rows = self.db.fetchall("SELECT * FROM provider_call_log") self.assertEqual(len(log_rows), 1) self.assertEqual(log_rows[0]["provider"], "eastmoney") self.assertEqual(log_rows[0]["interface"], "indices") self.assertEqual(log_rows[0]["status"], "ok") health = self.db.fetchone( "SELECT * FROM provider_health WHERE provider = ? AND interface = ?", ("eastmoney", "indices"), ) self.assertIsNotNone(health) self.assertEqual(health["state"], "ok") self.assertEqual(health["consec_failures"], 0) def test_consecutive_failures_increment_and_reset(self): observability.record_call(self.db, "tencent", "named_quotes", status="error", error="boom") observability.record_call(self.db, "tencent", "named_quotes", status="error", error="boom again") health = self.db.fetchone( "SELECT * FROM provider_health WHERE provider = ? AND interface = ?", ("tencent", "named_quotes"), ) self.assertEqual(health["consec_failures"], 2) self.assertEqual(health["state"], "error") observability.record_call(self.db, "tencent", "named_quotes", status="ok") health = self.db.fetchone( "SELECT * FROM provider_health WHERE provider = ? AND interface = ?", ("tencent", "named_quotes"), ) self.assertEqual(health["consec_failures"], 0) self.assertEqual(health["state"], "ok") def test_none_db_is_a_silent_noop(self): # Must not raise even though there is nowhere to write. observability.record_call(None, "ifind", "wencai", status="ok") def test_broken_db_write_does_not_raise(self): observability.record_call(_BrokenDB(), "eastmoney", "indices", status="ok") class ObserveTests(unittest.TestCase): def setUp(self) -> None: self.tmp = tempfile.TemporaryDirectory() self.db = HubDB(Path(self.tmp.name) / "hub.db") def tearDown(self) -> None: self.tmp.cleanup() def test_returns_exact_success_value_unmodified(self): sentinel = {"ts_code": "600000.SH", "close": 10.2} result = observability.observe(self.db, "eastmoney", "indices", lambda: sentinel) self.assertIs(result, sentinel) rows = self.db.fetchall("SELECT * FROM provider_call_log") self.assertEqual(len(rows), 1) self.assertEqual(rows[0]["status"], "ok") def test_reraises_exact_exception_on_failure(self): boom = ValueError("upstream exploded") def fn(): raise boom with self.assertRaises(ValueError) as ctx: observability.observe(self.db, "eastmoney", "indices", fn) self.assertIs(ctx.exception, boom) rows = self.db.fetchall("SELECT * FROM provider_call_log") self.assertEqual(len(rows), 1) self.assertEqual(rows[0]["status"], "error") self.assertIn("upstream exploded", rows[0]["error"]) def test_classify_downgrades_success_to_stale_without_changing_return_value(self): sentinel = [{"ts_code": "600000.SH", "quote_time_epoch": 1}] result = observability.observe( self.db, "eastmoney", "indices", lambda: sentinel, classify=lambda rows: observability.classify_rows(rows), ) self.assertIs(result, sentinel) row = self.db.fetchone("SELECT * FROM provider_call_log") self.assertEqual(row["status"], "stale") def test_broken_db_never_breaks_a_successful_call(self): sentinel = {"ok": True} result = observability.observe(_BrokenDB(), "eastmoney", "indices", lambda: sentinel) self.assertIs(result, sentinel) def test_broken_db_never_masks_a_real_failure(self): def fn(): raise RuntimeError("real upstream failure") with self.assertRaises(RuntimeError) as ctx: observability.observe(_BrokenDB(), "eastmoney", "indices", fn) self.assertEqual(str(ctx.exception), "real upstream failure") def test_classifier_exception_does_not_break_the_call(self): sentinel = {"ok": True} def bad_classify(_result): raise KeyError("classifier bug") result = observability.observe(self.db, "eastmoney", "indices", lambda: sentinel, classify=bad_classify) self.assertIs(result, sentinel) row = self.db.fetchone("SELECT * FROM provider_call_log") # A classifier bug must degrade to "ok", never silently drop the row # nor claim the call failed when it did not. self.assertEqual(row["status"], "ok") def test_none_db_is_transparent_passthrough(self): sentinel = object() result = observability.observe(None, "eastmoney", "indices", lambda: sentinel) self.assertIs(result, sentinel) if __name__ == "__main__": unittest.main()