#!/usr/bin/env python3 """数据库规则/样例装载器离线测试:与文件装载等价、失败关闭、激活门同源。""" import pathlib import sys import unittest PROJECT_ROOT = pathlib.Path(__file__).resolve().parents[2] sys.path.insert(0, str(PROJECT_ROOT / "humanization" / "src")) from deai import load, load_db # noqa: E402 class _Result: def __init__(self, rows): self.rows = list(rows) def fetchall(self): return list(self.rows) class _FakeConn: """按表名返回 payload 行;可注入执行异常验证失败关闭。""" def __init__(self, rule_payloads=None, sample_payloads=None, *, error=None): self.rule_payloads = list(rule_payloads or []) self.sample_payloads = list(sample_payloads or []) self.error = error def execute(self, sql, params=None): if self.error is not None: raise self.error if "example_ai_flavor_sample" in sql: return _Result([(payload,) for payload in self.sample_payloads]) if "example_ai_flavor_rule" in sql: return _Result([(payload,) for payload in self.rule_payloads]) raise AssertionError(f"未预期的 SQL: {sql}") def _file_corpus(): samples = load.load_samples() rules = load.load_rules(samples=samples) return samples, rules class LoadDbContractTest(unittest.TestCase): def test_db_corpus_equals_file_corpus_and_same_fingerprint(self): samples, rules = _file_corpus() conn = _FakeConn( rule_payloads=[load_db.rule_row(rule)["payload"] for rule in sorted(rules.values(), key=lambda r: r["id"])], sample_payloads=[load_db.sample_row(s)["payload"] for s in sorted(samples.values(), key=lambda s: s["id"])], ) db_samples = load_db.load_samples_from_db(conn) db_rules = load_db.load_rules_from_db(conn, samples=db_samples) self.assertEqual(db_samples, samples) self.assertEqual(db_rules, rules) self.assertEqual(load_db.canonical_sha(rules), load_db.canonical_sha(db_rules)) self.assertEqual(load.rule_library_version(db_rules), load.rule_library_version(rules)) def test_db_failure_is_fail_closed_without_fallback(self): conn = _FakeConn(error=RuntimeError("connection refused")) with self.assertRaises(load.LoadError) as ctx: load_db.load_rules_from_db(conn) self.assertIn("失败关闭", str(ctx.exception)) self.assertIn("不回退", str(ctx.exception)) def test_active_rule_missing_samples_is_rejected_on_db_path(self): samples, rules = _file_corpus() broken = load_db.rule_row(rules["l001"])["payload"] broken = dict(broken) broken["samples"] = {"sf": [], "snf": ["snf-l001-01"], "boundary": ["b-l001-01"], "regression": ["reg-l001-01"]} conn = _FakeConn( rule_payloads=[broken], sample_payloads=[load_db.sample_row(s)["payload"] for s in samples.values()], ) with self.assertRaises(load.LoadError) as ctx: load_db.load_rules_from_db(conn, samples=load_db.load_samples_from_db(conn)) self.assertIn("l001", str(ctx.exception)) def test_duplicate_rule_id_is_rejected(self): samples, rules = _file_corpus() payload = load_db.rule_row(rules["l001"])["payload"] conn = _FakeConn(rule_payloads=[payload, payload], sample_payloads=[]) with self.assertRaisesRegex(load.LoadError, "重复"): load_db.load_rules_from_db(conn) def test_row_projection_sha_matches_canonical_payload(self): samples, rules = _file_corpus() for rule in rules.values(): row = load_db.rule_row(rule) self.assertEqual(row["content_sha256"], load_db.canonical_sha(rule)) self.assertEqual(row["payload"], rule) for sample in samples.values(): row = load_db.sample_row(sample) self.assertEqual(row["content_sha256"], load_db.canonical_sha(sample)) self.assertEqual(row["payload"], sample) if __name__ == "__main__": unittest.main()