config: 新增 [run] mode 运行模式配置(cascade/doubao/deepseek)

- config.ini 全局生效,GUI 跟随配置,CLI --mode 参数可临时覆盖
- 必填项检查与首次配置对话框按模式过滤(仅豆包不要求 DeepSeek 配置,反之亦然)
- 模板/示例/README 同步更新,新增 2 个配置测试(共 35 个全部通过)
This commit is contained in:
yeuimu
2026-09-02 16:06:59 +08:00
parent 0033b7a4d2
commit 0ed5b9d9ff
6 changed files with 86 additions and 29 deletions
+4
View File
@@ -76,8 +76,12 @@ model_id = ep-xxxx
api_key = sk-... api_key = sk-...
model = deepseek-v4-flash-vision-exp model = deepseek-v4-flash-vision-exp
max_tokens = 5000 max_tokens = 5000
[run]
mode = cascade # cascade / doubao(仅豆包)/ deepseek(仅 DeepSeek
``` ```
运行模式:config.ini `[run] mode` 全局生效(GUI 也跟随);CLI 的 `--mode` 参数可临时覆盖。
## 项目结构 ## 项目结构
``` ```
+4
View File
@@ -23,3 +23,7 @@ file =
# 级联复检范围:豆包初筛得到这些结论的图片会再过 DeepSeek # 级联复检范围:豆包初筛得到这些结论的图片会再过 DeepSeek
[cascade] [cascade]
recheck = 无违规,违规不明,检测异常 recheck = 无违规,违规不明,检测异常
# 运行模式:cascade=豆包初筛+DeepSeek复检(默认)/ doubao=仅豆包 / deepseek=仅DeepSeek
[run]
mode = cascade
+7 -6
View File
@@ -22,9 +22,9 @@ def build_parser() -> argparse.ArgumentParser:
p.add_argument("-o", "--output", default=None, help="报表输出目录(默认为图片文件夹)") p.add_argument("-o", "--output", default=None, help="报表输出目录(默认为图片文件夹)")
p.add_argument("-c", "--config", default=None, help="配置文件路径(默认 config.ini)") p.add_argument("-c", "--config", default=None, help="配置文件路径(默认 config.ini)")
p.add_argument("--prompt", default=None, help="提示词文件路径(覆盖配置)") p.add_argument("--prompt", default=None, help="提示词文件路径(覆盖配置)")
p.add_argument("--mode", choices=["cascade", "doubao", "deepseek"], default="cascade", p.add_argument("--mode", choices=["cascade", "doubao", "deepseek"], default=None,
help="检测模式cascade=豆包初筛+DeepSeek复检(默认),doubao=仅豆包," help="检测模式(默认取 config.ini [run] mode,未配置则 cascade):"
"deepseek=仅 DeepSeek 全量") "cascade=豆包初筛+DeepSeek复检,doubao=仅豆包,deepseek=仅 DeepSeek 全量")
p.add_argument("--workers-ark", type=int, default=None, help="豆包并发数") p.add_argument("--workers-ark", type=int, default=None, help="豆包并发数")
p.add_argument("--workers-ds", type=int, default=None, help="DeepSeek 并发数") p.add_argument("--workers-ds", type=int, default=None, help="DeepSeek 并发数")
p.add_argument("--max-tokens", type=int, default=None, p.add_argument("--max-tokens", type=int, default=None,
@@ -47,8 +47,9 @@ def main(argv=None) -> int:
cfg.deepseek_workers = args.workers_ds cfg.deepseek_workers = args.workers_ds
if args.max_tokens: if args.max_tokens:
cfg.max_tokens = 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: if missing:
print(f"配置不完整,请编辑 {DEFAULT_CONFIG} 填写以下必填项:") print(f"配置不完整,请编辑 {DEFAULT_CONFIG} 填写以下必填项:")
for m in missing: for m in missing:
@@ -57,7 +58,7 @@ def main(argv=None) -> int:
try: try:
rows, summary = asyncio.run(run_detection( 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: except FileNotFoundError as e:
logger.error("路径或文件不存在:%s", e) logger.error("路径或文件不存在:%s", e)
return 1 return 1
@@ -72,7 +73,7 @@ def main(argv=None) -> int:
try: try:
out_dir = args.output or args.folder out_dir = args.output or args.folder
ts_dir = organize_output(rows, args.folder, out_dir) 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")}}) k: summary.get(k) for k in ("ds_calls", "ds_peak_cost", "ds_idle_cost")}})
except Exception: # noqa: BLE001 except Exception: # noqa: BLE001
logger.exception("输出组织/报表生成失败") logger.exception("输出组织/报表生成失败")
+19 -2
View File
@@ -55,6 +55,10 @@ file =
# 级联复检范围:豆包初筛得到这些结论的图片会再过 DeepSeek # 级联复检范围:豆包初筛得到这些结论的图片会再过 DeepSeek
[cascade] [cascade]
recheck = 无违规,违规不明,检测异常 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)} 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 @dataclass
class AppConfig: class AppConfig:
# 豆包(火山方舟) # 豆包(火山方舟)
@@ -86,6 +98,7 @@ class AppConfig:
prompt_file: str = "" prompt_file: str = ""
recheck_categories: list = field(default_factory=lambda: ["无违规", "违规不明", "检测异常"]) recheck_categories: list = field(default_factory=lambda: ["无违规", "违规不明", "检测异常"])
# 运行 # 运行
mode: str = "cascade"
retries: int = 3 retries: int = 3
@property @property
@@ -113,13 +126,15 @@ def ensure_config(path: str | None = None) -> Path:
return ini return ini
def missing_fields(cfg: AppConfig) -> list: def missing_fields(cfg: AppConfig, mode: str = "cascade") -> list:
"""级联模式必填项检查,返回缺失项中文名列表。""" """按运行模式检查必填项,返回缺失项中文名列表。"""
missing = [] missing = []
if mode in ("cascade", "doubao"):
if not cfg.ark_api_key: if not cfg.ark_api_key:
missing.append("豆包 API Key[ark] api_key") missing.append("豆包 API Key[ark] api_key")
if not cfg.ark_model: if not cfg.ark_model:
missing.append("豆包模型接入点([ark] model_id") missing.append("豆包模型接入点([ark] model_id")
if mode in ("cascade", "deepseek"):
if not cfg.deepseek_api_key: if not cfg.deepseek_api_key:
missing.append("DeepSeek API Key[deepseek] api_key") missing.append("DeepSeek API Key[deepseek] api_key")
if not cfg.deepseek_model: if not cfg.deepseek_model:
@@ -152,6 +167,7 @@ def load_config(path: str | None = None) -> AppConfig:
raw = get("cascade", "recheck") raw = get("cascade", "recheck")
if raw: if raw:
cfg.recheck_categories = [s.strip() for s in raw.split(",") if s.strip()] 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(): for attr, env in ENV_OVERRIDES.items():
val = os.environ.get(env) val = os.environ.get(env)
@@ -172,6 +188,7 @@ def save_config(cfg: AppConfig, path: str | None = None):
"max_tokens": str(cfg.max_tokens)} "max_tokens": str(cfg.max_tokens)}
parser["prompt"] = {"file": cfg.prompt_file or ""} parser["prompt"] = {"file": cfg.prompt_file or ""}
parser["cascade"] = {"recheck": ",".join(cfg.recheck_categories)} parser["cascade"] = {"recheck": ",".join(cfg.recheck_categories)}
parser["run"] = {"mode": cfg.mode}
with ini.open("w", encoding="utf-8") as f: with ini.open("w", encoding="utf-8") as f:
parser.write(f) parser.write(f)
+19 -10
View File
@@ -32,8 +32,13 @@ class FirstRunDialog:
("deepseek_model", "DeepSeek 模型", "deepseek-v4-flash-vision-exp"), ("deepseek_model", "DeepSeek 模型", "deepseek-v4-flash-vision-exp"),
] ]
def __init__(self, root, cfg): def __init__(self, root, cfg, mode="cascade"):
self.ok = False 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 = tk.Toplevel(root)
self.win.title("首次配置 - 商品图合规检测工具") self.win.title("首次配置 - 商品图合规检测工具")
self.win.resizable(False, False) self.win.resizable(False, False)
@@ -56,9 +61,11 @@ class FirstRunDialog:
def _save(self): def _save(self):
values = {attr: v.get().strip() for attr, v in self.vars.items()} 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: if empty:
messagebox.showwarning("提示", "以下项不能为空:\n" + "\n".join(empty), parent=self.win) messagebox.showwarning("提示", "以下必填项不能为空:\n" + "\n".join(empty),
parent=self.win)
return return
self.values = values self.values = values
self.ok = True self.ok = True
@@ -215,18 +222,20 @@ class App:
messagebox.showwarning("提示", "自定义提示词文件不存在") messagebox.showwarning("提示", "自定义提示词文件不存在")
return return
out_dir = self.var_output.get().strip() or folder out_dir = self.var_output.get().strip() or folder
mode = self.cfg.mode # 模式来自 config.ini [run] mode,界面不暴露
# 配置检查:缺失则弹首次配置对话框 # 配置检查:缺失则弹首次配置对话框
cfg = load_config() cfg = load_config()
if missing_fields(cfg): if missing_fields(cfg, mode):
self._ui_log("检测到 config.ini 配置不完整,请补齐必填项。") logger.info("检测到 config.ini 配置不完整(模式 %s,请补齐必填项。", mode)
dlg = FirstRunDialog(self.root, cfg) dlg = FirstRunDialog(self.root, cfg, mode)
if not dlg.ok: if not dlg.ok:
return return
for attr, val in dlg.values.items(): for attr, val in dlg.values.items():
if val:
setattr(cfg, attr, val) setattr(cfg, attr, val)
save_config(cfg) save_config(cfg)
self._ui_log("配置已保存到 config.ini") logger.info("配置已保存到 config.ini")
cfg.prompt_file = custom_prompt # 空 = 内置提示词 cfg.prompt_file = custom_prompt # 空 = 内置提示词
self.cfg = cfg self.cfg = cfg
@@ -235,7 +244,7 @@ class App:
self.btn_open.configure(state="disabled") self.btn_open.configure(state="disabled")
self.btn_folder.configure(state="disabled") self.btn_folder.configure(state="disabled")
self.stop_event.clear() 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("检测中…") self.var_progress_label.set("检测中…")
def progress(done, total, _label=None): def progress(done, total, _label=None):
@@ -248,11 +257,11 @@ class App:
try: try:
rows, summary = loop.run_until_complete(run_detection( rows, summary = loop.run_until_complete(run_detection(
folder, cfg, progress=progress, 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("正在整理输出(分类归档 + 生成报表)…") logger.info("正在整理输出(分类归档 + 生成报表)…")
ts_dir = organize_output(rows, folder, out_dir) ts_dir = organize_output(rows, folder, out_dir)
xlsx = build_report(rows, str(ts_dir), xlsx = build_report(rows, str(ts_dir),
{"mode": "cascade", **{k: summary.get(k) for k in {"mode": mode, **{k: summary.get(k) for k in
("ds_calls", "ds_peak_cost", ("ds_calls", "ds_peak_cost",
"ds_idle_cost")}}) "ds_idle_cost")}})
outdir = str(ts_dir) outdir = str(ts_dir)
+22
View File
@@ -41,6 +41,28 @@ def test_missing_fields_none_when_filled():
assert missing_fields(cfg) == [] 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): def test_save_and_load_roundtrip(tmp_path):
ini = tmp_path / "config.ini" ini = tmp_path / "config.ini"
cfg = AppConfig(ark_api_key="ark-key", ark_model="ep-xyz", cfg = AppConfig(ark_api_key="ark-key", ark_model="ep-xyz",