100 lines
4.0 KiB
Python

#!/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()