#!/usr/bin/env python3 """humanization 种子同步工具离线测试:幂等、事件留痕、db-only 失败关闭、dry-run 不连库。""" import contextlib import io import json import pathlib import sys import unittest PROJECT_ROOT = pathlib.Path(__file__).resolve().parents[2] sys.path.insert(0, str(PROJECT_ROOT / "humanization" / "src")) sys.path.insert(0, str(PROJECT_ROOT / "humanization" / "tools")) from deai import load, load_db # noqa: E402 import seed_rules_db as seed_tool # noqa: E402 class _Txn: def __enter__(self): return self def __exit__(self, *exc): return False class _SeedConn: """记录写语句的假连接;现有行以 id->sha 注入。""" def __init__(self, existing_rules=None, existing_samples=None): self.existing_rules = dict(existing_rules or {}) self.existing_samples = dict(existing_samples or {}) self.writes = [] self.txn_count = 0 def transaction(self): self.txn_count += 1 return _Txn() def execute(self, sql, params=None): normalized = " ".join(sql.split()) if normalized.startswith("SELECT rule_id, content_sha256 FROM example_ai_flavor_rule"): return _Rows([(rid, sha) for rid, sha in sorted(self.existing_rules.items())]) if normalized.startswith("SELECT sample_id, content_sha256 FROM example_ai_flavor_sample"): return _Rows([(sid, sha) for sid, sha in sorted(self.existing_samples.items())]) kind = ( "rule_insert" if normalized.startswith("INSERT INTO example_ai_flavor_rule (") else "rule_update" if normalized.startswith("UPDATE example_ai_flavor_rule") else "sample_insert" if normalized.startswith("INSERT INTO example_ai_flavor_sample") else "sample_update" if normalized.startswith("UPDATE example_ai_flavor_sample") else "rule_event" if normalized.startswith("INSERT INTO example_ai_flavor_rule_event") else None ) if kind is None: raise AssertionError(f"未预期的 SQL: {normalized}") self.writes.append((kind, params)) return _Rows([]) class _Rows: def __init__(self, rows): self.rows = list(rows) def fetchall(self): return list(self.rows) def _corpus(): samples = load.load_samples() rules = load.load_rules(samples=samples) return samples, rules class SeedRulesDbTest(unittest.TestCase): def test_first_seed_inserts_all_and_records_events(self): samples, rules = _corpus() conn = _SeedConn() summary = seed_tool.seed(conn, rules=rules, samples=samples) self.assertEqual(summary["rules"]["inserted"], len(rules)) self.assertEqual(summary["samples"]["inserted"], len(samples)) self.assertEqual(summary["rules"]["updated"], 0) self.assertEqual(summary["events"], len(rules)) kinds = [kind for kind, _ in conn.writes] self.assertEqual(kinds.count("rule_insert"), len(rules)) self.assertEqual(kinds.count("sample_insert"), len(samples)) self.assertEqual(kinds.count("rule_event"), len(rules)) def test_second_seed_with_same_content_is_idempotent(self): samples, rules = _corpus() existing_rules = {rid: load_db.canonical_sha(rule) for rid, rule in rules.items()} existing_samples = {sid: load_db.canonical_sha(s) for sid, s in samples.items()} conn = _SeedConn(existing_rules=existing_rules, existing_samples=existing_samples) summary = seed_tool.seed(conn, rules=rules, samples=samples) self.assertEqual(summary["rules"]["unchanged"], len(rules)) self.assertEqual(summary["samples"]["unchanged"], len(samples)) self.assertEqual(conn.writes, []) def test_changed_rule_is_updated_with_event(self): samples, rules = _corpus() changed = dict(rules["l001"], fix_hint="更新后的修复提示") rules = dict(rules, l001=changed) existing_rules = {rid: "0" * 64 for rid in rules} existing_samples = {sid: load_db.canonical_sha(s) for sid, s in samples.items()} conn = _SeedConn(existing_rules=existing_rules, existing_samples=existing_samples) summary = seed_tool.seed(conn, rules=rules, samples=samples) self.assertEqual(summary["rules"]["updated"], len(rules)) self.assertEqual(summary["samples"]["unchanged"], len(samples)) rule_writes = [kind for kind, _ in conn.writes if kind in ("rule_update", "rule_event")] self.assertEqual(rule_writes.count("rule_update"), len(rules)) self.assertEqual(rule_writes.count("rule_event"), len(rules)) def test_db_only_rows_are_reported_and_strict_fails_closed(self): samples, rules = _corpus() conn = _SeedConn(existing_rules={"z999": "0" * 64}) summary = seed_tool.seed(conn, rules=rules, samples=samples) self.assertEqual(summary["db_only_rules"], ["z999"]) conn_strict = _SeedConn(existing_rules={"z999": "0" * 64}) with self.assertRaisesRegex(seed_tool.SeedError, "失败关闭"): seed_tool.seed(conn_strict, rules=rules, samples=samples, strict=True) def test_dry_run_does_not_touch_database(self): buffer = io.StringIO() with contextlib.redirect_stdout(buffer): code = seed_tool.main(["--dry-run"]) self.assertEqual(code, 0) plan = json.loads(buffer.getvalue()) self.assertEqual(plan["status"], "dry_run") samples, rules = _corpus() self.assertEqual(plan["rules"], len(rules)) self.assertEqual(plan["samples"], len(samples)) self.assertEqual(plan["library_version"], load.rule_library_version(rules)) if __name__ == "__main__": unittest.main()