Files
violation-detector/tests/test_config.py
T
yeuimu 0ed5b9d9ff config: 新增 [run] mode 运行模式配置(cascade/doubao/deepseek)
- config.ini 全局生效,GUI 跟随配置,CLI --mode 参数可临时覆盖
- 必填项检查与首次配置对话框按模式过滤(仅豆包不要求 DeepSeek 配置,反之亦然)
- 模板/示例/README 同步更新,新增 2 个配置测试(共 35 个全部通过)
2026-09-02 16:06:59 +08:00

93 lines
3.4 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
# -*- coding: utf-8 -*-
"""config 模块测试:模板生成、缺项检测、读写回环、提示词路径。"""
from violation_detector.config import (
AppConfig, ensure_config, load_config, missing_fields, save_config,
)
def test_ensure_config_creates_template(tmp_path):
ini = tmp_path / "config.ini"
path = ensure_config(str(ini))
assert path == ini
text = ini.read_text(encoding="utf-8")
assert "[ark]" in text and "[deepseek]" in text
def test_ensure_config_keeps_existing(tmp_path):
ini = tmp_path / "config.ini"
ini.write_text("[ark]\napi_key = abc\n", encoding="utf-8")
ensure_config(str(ini))
assert "api_key = abc" in ini.read_text(encoding="utf-8")
def test_missing_fields_all_empty():
# deepseek_model 有默认值;只有用户在 ini 里清空才视为缺失
missing = missing_fields(AppConfig(deepseek_model=""))
assert len(missing) == 4
assert any("豆包 API Key" in m for m in missing)
assert any("DeepSeek API Key" in m for m in missing)
def test_missing_fields_default_model_ok():
# 默认仅缺:豆包 Key、豆包接入点、DeepSeek KeyDeepSeek 模型有默认值)
missing = missing_fields(AppConfig())
assert len(missing) == 3
assert not any("DeepSeek 模型" in m for m in missing)
def test_missing_fields_none_when_filled():
cfg = AppConfig(ark_api_key="k1", ark_model="ep-1",
deepseek_api_key="sk-1", deepseek_model="m1")
assert missing_fields(cfg) == []
def test_missing_fields_mode_aware():
cfg = AppConfig(ark_api_key="k", ark_model="ep", deepseek_api_key="", deepseek_model="m")
# cascade 需要 DeepSeek Key
assert len(missing_fields(cfg, "cascade")) == 1
# doubao 模式不需要 DeepSeek 配置
assert missing_fields(cfg, "doubao") == []
# deepseek 模式不需要豆包配置
assert missing_fields(AppConfig(deepseek_api_key="sk"), "deepseek") == []
def test_mode_from_ini(tmp_path):
ini = tmp_path / "config.ini"
ini.write_text("[run]\nmode = deepseek\n", encoding="utf-8")
assert load_config(str(ini)).mode == "deepseek"
# 非法值回退 cascade
ini.write_text("[run]\nmode = turbo\n", encoding="utf-8")
assert load_config(str(ini)).mode == "cascade"
# 未配置默认 cascade
ini.write_text("[ark]\n", encoding="utf-8")
assert load_config(str(ini)).mode == "cascade"
def test_save_and_load_roundtrip(tmp_path):
ini = tmp_path / "config.ini"
cfg = AppConfig(ark_api_key="ark-key", ark_model="ep-xyz",
deepseek_api_key="sk-key", deepseek_model="ds-model",
ark_workers=8, max_tokens=4000, prompt_file="my.txt")
save_config(cfg, str(ini))
loaded = load_config(str(ini))
assert loaded.ark_api_key == "ark-key"
assert loaded.ark_model == "ep-xyz"
assert loaded.deepseek_api_key == "sk-key"
assert loaded.deepseek_model == "ds-model"
assert loaded.ark_workers == 8
assert loaded.max_tokens == 4000
assert loaded.prompt_file == "my.txt"
assert loaded.recheck_categories == ["无违规", "违规不明", "检测异常"]
def test_prompt_path_defaults_to_builtin():
cfg = AppConfig()
assert cfg.prompt_path.name == "prompts.txt"
def test_prompt_path_custom(tmp_path):
p = tmp_path / "p.txt"
p.write_text("x", encoding="utf-8")
cfg = AppConfig(prompt_file=str(p))
assert cfg.prompt_path == p