feat: add AdversarialFilterConfig for Phase B post-hoc filter layer

This commit is contained in:
2026-07-14 15:47:21 -04:00
parent ac115d96fb
commit 334fbbc94d
3 changed files with 104 additions and 0 deletions
+38
View File
@@ -0,0 +1,38 @@
"""AdversarialFilterConfig 默认值 + YAML 加载。"""
from app.question_gen.adversarial_config import (
AdversarialFilterConfig,
load_adversarial_config,
)
def test_defaults():
cfg = AdversarialFilterConfig()
assert cfg.filter_task_types == ("Action Recognition",)
assert cfg.adversarial_max_rounds == 5
assert cfg.adversarial_agent_max_steps == 40
assert cfg.difficulty_warn_threshold == 0.85
def test_load_from_yaml(tmp_path):
p = tmp_path / "c.yaml"
p.write_text(
"adversarial_filter:\n"
" filter_task_types: [Action Recognition, Object Recognition]\n"
" adversarial_max_rounds: 3\n"
" adversarial_agent_max_steps: 20\n"
" difficulty_warn_threshold: 0.7\n",
encoding="utf-8",
)
cfg = load_adversarial_config(p)
assert cfg.filter_task_types == ("Action Recognition", "Object Recognition")
assert cfg.adversarial_max_rounds == 3
assert cfg.adversarial_agent_max_steps == 20
assert cfg.difficulty_warn_threshold == 0.7
def test_load_missing_section_uses_defaults(tmp_path):
p = tmp_path / "c.yaml"
p.write_text("question_gen_v2:\n per_type: 3\n", encoding="utf-8")
cfg = load_adversarial_config(p)
assert cfg == AdversarialFilterConfig()