总工复核 🔴:安全边界要求新增观测功能必须可关闭、关闭后现有功能完全照旧, 但此前实现没有任何运行时开关。修复: - settings.py: 新增 Settings.observability_enabled 字段,沿用既有 DATAHUB_SCHEDULER 的环境变量模式,读取 DATAHUB_OBSERVABILITY (0/false/off 关闭,默认开启)。 - hub.py: Hub.__init__ 把 settings.observability_enabled 挂到 self.db 上,让 pipeline/realtime_serve/steward/admin_api 已经 在传的 db 参数直接带上开关,零额外改造。 - observability.py: 新增 is_enabled(db),缺失该属性时默认按启用处理 (向后兼容裸 HubDB 用例/测试)。关闭时 observe() 变成纯 透传(不计时、不分类、不碰数据库),record_call() 直接 no-op。 - admin_api.py: 4 个新只读端点关闭时返回明确的 {"enabled": false, ...空结构} 而不是静默返回旧数据。 - 新增 10 个测试:开关默认值/环境变量解析、关闭后 observe() 的透传语义 (含异常原样重新抛出)、关闭后 record_call() 零写入、关闭后重新开启恢复 记录、4 个 admin 端点在关闭态的响应结构。 全量测试 183/183 通过(新增 10 个,含此前 173 个零回归)。 Co-authored-by: Cursor <cursoragent@cursor.com> Co-authored-by: multica-agent <github@multica.ai>
307 lines
12 KiB
Python
307 lines
12 KiB
Python
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)
|
|
|
|
|
|
class _ToggleDB(HubDB):
|
|
"""A real HubDB subclass so we can flip the HEL-543 kill switch the same
|
|
way Hub.__init__ does, without needing a full Hub/Settings wiring."""
|
|
|
|
|
|
class KillSwitchTests(unittest.TestCase):
|
|
"""HEL-543 total-review 🔴: the observability side channel must be
|
|
disable-able at runtime, and disabling it must leave existing behavior
|
|
completely unchanged (pure passthrough, zero db access)."""
|
|
|
|
def setUp(self) -> None:
|
|
self.tmp = tempfile.TemporaryDirectory()
|
|
self.db = _ToggleDB(Path(self.tmp.name) / "hub.db")
|
|
|
|
def tearDown(self) -> None:
|
|
self.tmp.cleanup()
|
|
|
|
def test_is_enabled_defaults_true_when_attribute_absent(self):
|
|
# A bare HubDB (as used throughout the rest of this test suite, and
|
|
# by any pre-HEL-543 call site) must default to enabled.
|
|
self.assertTrue(observability.is_enabled(self.db))
|
|
self.assertTrue(observability.is_enabled(None))
|
|
|
|
def test_disabled_record_call_writes_nothing(self):
|
|
self.db.observability_enabled = False
|
|
observability.record_call(self.db, "eastmoney", "indices", status="ok", latency_ms=1)
|
|
self.assertEqual(self.db.fetchall("SELECT * FROM provider_call_log"), [])
|
|
self.assertEqual(self.db.fetchall("SELECT * FROM provider_health"), [])
|
|
|
|
def test_disabled_observe_is_a_pure_passthrough_on_success(self):
|
|
self.db.observability_enabled = False
|
|
sentinel = {"ts_code": "600000.SH"}
|
|
calls = {"n": 0}
|
|
|
|
def fn():
|
|
calls["n"] += 1
|
|
return sentinel
|
|
|
|
result = observability.observe(self.db, "eastmoney", "indices", fn)
|
|
self.assertIs(result, sentinel)
|
|
self.assertEqual(calls["n"], 1)
|
|
self.assertEqual(self.db.fetchall("SELECT * FROM provider_call_log"), [])
|
|
|
|
def test_disabled_observe_still_reraises_the_exact_exception(self):
|
|
self.db.observability_enabled = False
|
|
boom = RuntimeError("upstream exploded")
|
|
|
|
def fn():
|
|
raise boom
|
|
|
|
with self.assertRaises(RuntimeError) as ctx:
|
|
observability.observe(self.db, "eastmoney", "indices", fn)
|
|
self.assertIs(ctx.exception, boom)
|
|
self.assertEqual(self.db.fetchall("SELECT * FROM provider_call_log"), [])
|
|
|
|
def test_re_enabling_resumes_recording(self):
|
|
self.db.observability_enabled = False
|
|
observability.observe(self.db, "eastmoney", "indices", lambda: {"ok": True})
|
|
self.assertEqual(self.db.fetchall("SELECT * FROM provider_call_log"), [])
|
|
self.db.observability_enabled = True
|
|
observability.observe(self.db, "eastmoney", "indices", lambda: {"ok": True})
|
|
self.assertEqual(len(self.db.fetchall("SELECT * FROM provider_call_log")), 1)
|
|
|
|
|
|
class SettingsToggleTests(unittest.TestCase):
|
|
"""The kill switch follows the same env-var pattern as DATAHUB_SCHEDULER."""
|
|
|
|
def test_defaults_to_enabled(self):
|
|
from datahub.settings import load_settings
|
|
|
|
settings = load_settings(env={})
|
|
self.assertTrue(settings.observability_enabled)
|
|
|
|
def test_datahub_observability_zero_disables(self):
|
|
from datahub.settings import load_settings
|
|
|
|
settings = load_settings(env={"DATAHUB_OBSERVABILITY": "0"})
|
|
self.assertFalse(settings.observability_enabled)
|
|
|
|
def test_datahub_observability_off_disables(self):
|
|
from datahub.settings import load_settings
|
|
|
|
settings = load_settings(env={"DATAHUB_OBSERVABILITY": "off"})
|
|
self.assertFalse(settings.observability_enabled)
|
|
|
|
def test_datahub_observability_one_keeps_enabled(self):
|
|
from datahub.settings import load_settings
|
|
|
|
settings = load_settings(env={"DATAHUB_OBSERVABILITY": "1"})
|
|
self.assertTrue(settings.observability_enabled)
|
|
|
|
|
|
if __name__ == "__main__":
|
|
unittest.main()
|