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-...
|
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` 参数可临时覆盖。
|
||||||
|
|
||||||
## 项目结构
|
## 项目结构
|
||||||
|
|
||||||
```
|
```
|
||||||
|
|||||||
@@ -23,3 +23,7 @@ file =
|
|||||||
# 级联复检范围:豆包初筛得到这些结论的图片会再过 DeepSeek
|
# 级联复检范围:豆包初筛得到这些结论的图片会再过 DeepSeek
|
||||||
[cascade]
|
[cascade]
|
||||||
recheck = 无违规,违规不明,检测异常
|
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("-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("输出组织/报表生成失败")
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|||||||
@@ -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",
|
||||||
|
|||||||
Reference in New Issue
Block a user