100 lines
4.0 KiB
Python
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()
|