config: 新增 [run] mode 运行模式配置(cascade/doubao/deepseek)
- config.ini 全局生效,GUI 跟随配置,CLI --mode 参数可临时覆盖 - 必填项检查与首次配置对话框按模式过滤(仅豆包不要求 DeepSeek 配置,反之亦然) - 模板/示例/README 同步更新,新增 2 个配置测试(共 35 个全部通过)
This commit is contained in:
@@ -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` 参数可临时覆盖。
|
||||
|
||||
## 项目结构
|
||||
|
||||
```
|
||||
|
||||
@@ -23,3 +23,7 @@ file =
|
||||
# 级联复检范围:豆包初筛得到这些结论的图片会再过 DeepSeek
|
||||
[cascade]
|
||||
recheck = 无违规,违规不明,检测异常
|
||||
|
||||
# 运行模式:cascade=豆包初筛+DeepSeek复检(默认)/ doubao=仅豆包 / deepseek=仅DeepSeek
|
||||
[run]
|
||||
mode = cascade
|
||||
|
||||
@@ -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("输出组织/报表生成失败")
|
||||
|
||||
@@ -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,13 +126,15 @@ 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 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:
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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():
|
||||
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,11 +257,11 @@ 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
|
||||
{"mode": mode, **{k: summary.get(k) for k in
|
||||
("ds_calls", "ds_peak_cost",
|
||||
"ds_idle_cost")}})
|
||||
outdir = str(ts_dir)
|
||||
|
||||
@@ -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",
|
||||
|
||||
Reference in New Issue
Block a user