From 0ed5b9d9ff6fdcda9d846479977bd2367708f068 Mon Sep 17 00:00:00 2001 From: yeuimu <2197651308@qq.com> Date: Wed, 2 Sep 2026 16:06:59 +0800 Subject: [PATCH] =?UTF-8?q?config:=20=E6=96=B0=E5=A2=9E=20[run]=20mode=20?= =?UTF-8?q?=E8=BF=90=E8=A1=8C=E6=A8=A1=E5=BC=8F=E9=85=8D=E7=BD=AE=EF=BC=88?= =?UTF-8?q?cascade/doubao/deepseek=EF=BC=89?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - config.ini 全局生效,GUI 跟随配置,CLI --mode 参数可临时覆盖 - 必填项检查与首次配置对话框按模式过滤(仅豆包不要求 DeepSeek 配置,反之亦然) - 模板/示例/README 同步更新,新增 2 个配置测试(共 35 个全部通过) --- README.md | 4 ++++ config.example.ini | 4 ++++ src/violation_detector/cli.py | 13 +++++------ src/violation_detector/config.py | 37 +++++++++++++++++++++++--------- src/violation_detector/gui.py | 35 +++++++++++++++++++----------- tests/test_config.py | 22 +++++++++++++++++++ 6 files changed, 86 insertions(+), 29 deletions(-) diff --git a/README.md b/README.md index 63b86a2..5ba8bac 100644 --- a/README.md +++ b/README.md @@ -76,8 +76,12 @@ model_id = ep-xxxx api_key = sk-... model = deepseek-v4-flash-vision-exp max_tokens = 5000 +[run] +mode = cascade # cascade / doubao(仅豆包)/ deepseek(仅 DeepSeek) ``` +运行模式:config.ini `[run] mode` 全局生效(GUI 也跟随);CLI 的 `--mode` 参数可临时覆盖。 + ## 项目结构 ``` diff --git a/config.example.ini b/config.example.ini index 34e84c4..d8b660e 100644 --- a/config.example.ini +++ b/config.example.ini @@ -23,3 +23,7 @@ file = # 级联复检范围:豆包初筛得到这些结论的图片会再过 DeepSeek [cascade] recheck = 无违规,违规不明,检测异常 + +# 运行模式:cascade=豆包初筛+DeepSeek复检(默认)/ doubao=仅豆包 / deepseek=仅DeepSeek +[run] +mode = cascade diff --git a/src/violation_detector/cli.py b/src/violation_detector/cli.py index 70ff1bd..bcd1b37 100644 --- a/src/violation_detector/cli.py +++ b/src/violation_detector/cli.py @@ -22,9 +22,9 @@ def build_parser() -> argparse.ArgumentParser: p.add_argument("-o", "--output", default=None, help="报表输出目录(默认为图片文件夹)") p.add_argument("-c", "--config", default=None, help="配置文件路径(默认 config.ini)") p.add_argument("--prompt", default=None, help="提示词文件路径(覆盖配置)") - p.add_argument("--mode", choices=["cascade", "doubao", "deepseek"], default="cascade", - help="检测模式:cascade=豆包初筛+DeepSeek复检(默认),doubao=仅豆包," - "deepseek=仅 DeepSeek 全量") + p.add_argument("--mode", choices=["cascade", "doubao", "deepseek"], default=None, + help="检测模式(默认取 config.ini [run] mode,未配置则 cascade):" + "cascade=豆包初筛+DeepSeek复检,doubao=仅豆包,deepseek=仅 DeepSeek 全量") p.add_argument("--workers-ark", type=int, default=None, help="豆包并发数") p.add_argument("--workers-ds", type=int, default=None, help="DeepSeek 并发数") p.add_argument("--max-tokens", type=int, default=None, @@ -47,8 +47,9 @@ def main(argv=None) -> int: cfg.deepseek_workers = args.workers_ds if args.max_tokens: cfg.max_tokens = args.max_tokens + mode = args.mode or cfg.mode - missing = missing_fields(cfg) + missing = missing_fields(cfg, mode) if missing: print(f"配置不完整,请编辑 {DEFAULT_CONFIG} 填写以下必填项:") for m in missing: @@ -57,7 +58,7 @@ def main(argv=None) -> int: try: rows, summary = asyncio.run(run_detection( - args.folder, cfg, mode=args.mode, verify=not args.no_verify)) + args.folder, cfg, mode=mode, verify=not args.no_verify)) except FileNotFoundError as e: logger.error("路径或文件不存在:%s", e) return 1 @@ -72,7 +73,7 @@ def main(argv=None) -> int: try: out_dir = args.output or args.folder ts_dir = organize_output(rows, args.folder, out_dir) - xlsx = build_report(rows, str(ts_dir), {"mode": args.mode, **{ + xlsx = build_report(rows, str(ts_dir), {"mode": mode, **{ k: summary.get(k) for k in ("ds_calls", "ds_peak_cost", "ds_idle_cost")}}) except Exception: # noqa: BLE001 logger.exception("输出组织/报表生成失败") diff --git a/src/violation_detector/config.py b/src/violation_detector/config.py index 3bac5d2..3ecbc03 100644 --- a/src/violation_detector/config.py +++ b/src/violation_detector/config.py @@ -55,6 +55,10 @@ file = # 级联复检范围:豆包初筛得到这些结论的图片会再过 DeepSeek [cascade] recheck = 无违规,违规不明,检测异常 + +# 运行模式:cascade=豆包初筛+DeepSeek复检(默认)/ doubao=仅豆包 / deepseek=仅DeepSeek +[run] +mode = cascade """ # 环境变量优先于配置文件 @@ -69,6 +73,14 @@ ENV_OVERRIDES = { PRICE = {"miss": (3.0, 1.5), "cached": (0.10, 0.05), "output": (9.0, 4.5)} +VALID_MODES = ("cascade", "doubao", "deepseek") + + +def normalize_mode(mode: str) -> str: + """非法模式回退为 cascade。""" + return mode if mode in VALID_MODES else "cascade" + + @dataclass class AppConfig: # 豆包(火山方舟) @@ -86,6 +98,7 @@ class AppConfig: prompt_file: str = "" recheck_categories: list = field(default_factory=lambda: ["无违规", "违规不明", "检测异常"]) # 运行 + mode: str = "cascade" retries: int = 3 @property @@ -113,17 +126,19 @@ def ensure_config(path: str | None = None) -> Path: return ini -def missing_fields(cfg: AppConfig) -> list: - """级联模式必填项检查,返回缺失项中文名列表。""" +def missing_fields(cfg: AppConfig, mode: str = "cascade") -> list: + """按运行模式检查必填项,返回缺失项中文名列表。""" missing = [] - if not cfg.ark_api_key: - missing.append("豆包 API Key([ark] api_key)") - if not cfg.ark_model: - missing.append("豆包模型接入点([ark] model_id)") - if not cfg.deepseek_api_key: - missing.append("DeepSeek API Key([deepseek] api_key)") - if not cfg.deepseek_model: - missing.append("DeepSeek 模型([deepseek] model)") + if mode in ("cascade", "doubao"): + if not cfg.ark_api_key: + missing.append("豆包 API Key([ark] api_key)") + if not cfg.ark_model: + missing.append("豆包模型接入点([ark] model_id)") + if mode in ("cascade", "deepseek"): + if not cfg.deepseek_api_key: + missing.append("DeepSeek API Key([deepseek] api_key)") + if not cfg.deepseek_model: + missing.append("DeepSeek 模型([deepseek] model)") return missing @@ -152,6 +167,7 @@ def load_config(path: str | None = None) -> AppConfig: raw = get("cascade", "recheck") if raw: cfg.recheck_categories = [s.strip() for s in raw.split(",") if s.strip()] + cfg.mode = normalize_mode(get("run", "mode") or cfg.mode) for attr, env in ENV_OVERRIDES.items(): val = os.environ.get(env) @@ -172,6 +188,7 @@ def save_config(cfg: AppConfig, path: str | None = None): "max_tokens": str(cfg.max_tokens)} parser["prompt"] = {"file": cfg.prompt_file or ""} parser["cascade"] = {"recheck": ",".join(cfg.recheck_categories)} + parser["run"] = {"mode": cfg.mode} with ini.open("w", encoding="utf-8") as f: parser.write(f) diff --git a/src/violation_detector/gui.py b/src/violation_detector/gui.py index 3ac0cfb..7707732 100644 --- a/src/violation_detector/gui.py +++ b/src/violation_detector/gui.py @@ -32,8 +32,13 @@ class FirstRunDialog: ("deepseek_model", "DeepSeek 模型", "deepseek-v4-flash-vision-exp"), ] - def __init__(self, root, cfg): + def __init__(self, root, cfg, mode="cascade"): self.ok = False + self.required = [] + if mode in ("cascade", "doubao"): + self.required += ["ark_api_key", "ark_model"] + if mode in ("cascade", "deepseek"): + self.required += ["deepseek_api_key", "deepseek_model"] self.win = tk.Toplevel(root) self.win.title("首次配置 - 商品图合规检测工具") self.win.resizable(False, False) @@ -56,9 +61,11 @@ class FirstRunDialog: def _save(self): values = {attr: v.get().strip() for attr, v in self.vars.items()} - empty = [label for attr, label, _ in self.FIELDS if not values[attr]] + empty = [label for attr, label, _ in self.FIELDS + if attr in self.required and not values[attr]] if empty: - messagebox.showwarning("提示", "以下项不能为空:\n" + "\n".join(empty), parent=self.win) + messagebox.showwarning("提示", "以下必填项不能为空:\n" + "\n".join(empty), + parent=self.win) return self.values = values self.ok = True @@ -215,18 +222,20 @@ class App: messagebox.showwarning("提示", "自定义提示词文件不存在") return out_dir = self.var_output.get().strip() or folder + mode = self.cfg.mode # 模式来自 config.ini [run] mode,界面不暴露 # 配置检查:缺失则弹首次配置对话框 cfg = load_config() - if missing_fields(cfg): - self._ui_log("检测到 config.ini 配置不完整,请补齐必填项。") - dlg = FirstRunDialog(self.root, cfg) + if missing_fields(cfg, mode): + logger.info("检测到 config.ini 配置不完整(模式 %s),请补齐必填项。", mode) + dlg = FirstRunDialog(self.root, cfg, mode) if not dlg.ok: return for attr, val in dlg.values.items(): - setattr(cfg, attr, val) + if val: + setattr(cfg, attr, val) save_config(cfg) - self._ui_log("配置已保存到 config.ini") + logger.info("配置已保存到 config.ini") cfg.prompt_file = custom_prompt # 空 = 内置提示词 self.cfg = cfg @@ -235,7 +244,7 @@ class App: self.btn_open.configure(state="disabled") self.btn_folder.configure(state="disabled") self.stop_event.clear() - logger.info("===== 开始检测:%s(模式 cascade,输出:%s)=====", folder, out_dir) + logger.info("===== 开始检测:%s(模式 %s,输出:%s)=====", folder, mode, out_dir) self.var_progress_label.set("检测中…") def progress(done, total, _label=None): @@ -248,13 +257,13 @@ class App: try: rows, summary = loop.run_until_complete(run_detection( folder, cfg, progress=progress, - stop=self.stop_event.is_set, mode="cascade", verify=True)) + stop=self.stop_event.is_set, mode=mode, verify=True)) logger.info("正在整理输出(分类归档 + 生成报表)…") ts_dir = organize_output(rows, folder, out_dir) xlsx = build_report(rows, str(ts_dir), - {"mode": "cascade", **{k: summary.get(k) for k in - ("ds_calls", "ds_peak_cost", - "ds_idle_cost")}}) + {"mode": mode, **{k: summary.get(k) for k in + ("ds_calls", "ds_peak_cost", + "ds_idle_cost")}}) outdir = str(ts_dir) logger.info("输出目录:%s", ts_dir) logger.info("报表:%s", xlsx) diff --git a/tests/test_config.py b/tests/test_config.py index 661ca1e..728f8cc 100644 --- a/tests/test_config.py +++ b/tests/test_config.py @@ -41,6 +41,28 @@ def test_missing_fields_none_when_filled(): 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",