feat(settings): support gateway source switching

This commit is contained in:
chengma
2026-07-20 16:29:28 +08:00
parent ed1d8ee764
commit 6a0cd1c763
18 changed files with 746 additions and 77 deletions
+41
View File
@@ -106,6 +106,47 @@ class AITests(TempDirMixin, unittest.TestCase):
appconfig.save_cmhub_config({"api_key": "sk-cmhub-secret"}, path=key_path)
return cfg, key_path
def test_freeze_runtime_config_keeps_worker_credentials_and_models_in_memory(self):
with self.make_temp_dir() as temp_dir:
config = self._config()
config["ai"]["cmhub"] = {
"base_url": "https://cmhub.example.com",
"image_alias": "image-hd",
"connect_timeout": 3,
}
cmhub_path = os.path.join(temp_dir, "cmhub.json")
models_path = os.path.join(temp_dir, "ai_models.json")
self._write_models(models_path)
appconfig.save_cmhub_config({"api_key": "sk-before-save"}, path=cmhub_path)
snapshot = ai.freeze_runtime_config(
config,
cmhub_config_path=cmhub_path,
models_path=models_path,
include_cmhub=True,
)
appconfig.save_cmhub_config({"api_key": "sk-after-save"}, path=cmhub_path)
replacement_models = appconfig.list_ai_models(
path=models_path,
reveal_api_key=True,
)
for model in replacement_models:
model["name"] = "已保存后替换的%s模型" % model["category"]
appconfig.save_ai_models_config({"models": replacement_models}, path=models_path)
self.assertEqual(
"sk-before-save",
ai._cmhub_runtime(snapshot, "image", cmhub_path)["api_key"],
)
ai.validate_direct_generation_config(
snapshot,
"title",
models_path=models_path,
)
self.assertIn("direct_models", snapshot["_cmshopee_ai_runtime"])
self.assert_removed(temp_dir)
def _collected_tasks(self, temp_dir, cfg, titles=None):
titles = titles or ["旧标题A", "旧标题B"]
db.init_db(cfg["db_path"])
+42 -9
View File
@@ -2970,6 +2970,39 @@ class GuiTests(TempDirMixin, unittest.TestCase):
self.assert_removed(temp_dir)
def test_settings_gateway_selector_switches_panels_and_persists_only_on_save(self):
with self.make_temp_dir() as temp_dir:
cfg = self.make_config(temp_dir)
cfg["ai"] = appconfig.default_config()["ai"]
cfg["ai"]["backend"] = "cmhub"
tab = SettingsTab(
config=cfg,
config_path=cfg["config_path"],
ai_models_path=cfg["ai_models_path"],
)
self.addCleanup(tab.close)
self.assertTrue(tab.gateway_default_button.isChecked())
self.assertFalse(tab.gateway_custom_button.isChecked())
self.assertFalse(tab.cmhub_panel.isHidden())
self.assertTrue(tab.model_picker_panel.isHidden())
tab.gateway_custom_button.setChecked(True)
self.assertTrue(tab.gateway_custom_button.isChecked())
self.assertFalse(tab.gateway_default_button.isChecked())
self.assertTrue(tab.cmhub_panel.isHidden())
self.assertFalse(tab.model_picker_panel.isHidden())
self.assertFalse(tab.direct_role_panel.isHidden())
self.assertTrue(tab.is_dirty())
self.assertEqual("cmhub", appconfig.load_config(cfg["config_path"])["ai"]["backend"])
with mock.patch("app.gui.QMessageBox.information"):
self.assertTrue(tab.save_app_settings())
self.assertEqual("direct", appconfig.load_config(cfg["config_path"])["ai"]["backend"])
self.assert_removed(temp_dir)
def test_settings_tab_cmhub_alias_refresh_filters_unpriced_and_keeps_saved(self):
with self.make_temp_dir() as temp_dir:
cfg = self.make_config(temp_dir)
@@ -3103,8 +3136,8 @@ class GuiTests(TempDirMixin, unittest.TestCase):
"points_balance": 66,
}
)
self.assertIn("cmhub 账号「主账号」连接成功", tab.cmhub_result_label.text())
self.assertNotIn("cmhub 账号「cmhub_user」", tab.cmhub_result_label.text())
self.assertIn("账号「主账号」连接默认网关成功", tab.cmhub_result_label.text())
self.assertNotIn("账号「cmhub_user」连接默认网关成功", tab.cmhub_result_label.text())
self.assertIn("余额 66", tab.cmhub_result_label.text())
self.assertEqual(tab.cmhub_result_label.text(), statuses[-1])
@@ -3116,7 +3149,7 @@ class GuiTests(TempDirMixin, unittest.TestCase):
"points_balance": 67,
}
)
self.assertIn("cmhub 账号「备用账号」连接成功", tab.cmhub_result_label.text())
self.assertIn("账号「备用账号」连接默认网关成功", tab.cmhub_result_label.text())
tab._on_cmhub_finished(
{
@@ -3126,12 +3159,12 @@ class GuiTests(TempDirMixin, unittest.TestCase):
"points_balance": 77,
}
)
self.assertIn("cmhub 账号「o***r@example.com」连接成功", tab.cmhub_result_label.text())
self.assertIn("账号「o***r@example.com」连接默认网关成功", tab.cmhub_result_label.text())
self.assertNotIn("owner@example.com", tab.cmhub_result_label.text())
tab._on_cmhub_finished({"ok": True, "models": models, "points_balance": 88})
self.assertIn("cmhub 连接成功", tab.cmhub_result_label.text())
self.assertNotIn("cmhub 账号", tab.cmhub_result_label.text())
self.assertIn("默认网关连接成功", tab.cmhub_result_label.text())
self.assertNotIn("账号「", tab.cmhub_result_label.text())
self.assert_removed(temp_dir)
def test_settings_tab_tracks_dirty_state_and_programmatic_cmhub_refresh_is_clean(self):
@@ -3381,7 +3414,7 @@ class GuiTests(TempDirMixin, unittest.TestCase):
self.assertEqual("生图用时 0 秒", tab.cover_elapsed_label.text())
self.assertEqual(tab.title_elapsed_label.width(), tab.cover_elapsed_label.width())
self.assertEqual("generateCmhubBalanceLabel", tab.cmhub_balance_label.objectName())
self.assertEqual("cmhub余额:未获取", tab.cmhub_balance_label.text())
self.assertEqual("默认网关余额:未获取", tab.cmhub_balance_label.text())
self.assertTrue(tab.cmhub_balance_label.isHidden())
self.assertEqual(0, tab.title_progress_bar.value())
self.assertEqual(0, tab.cover_progress_bar.value())
@@ -3993,7 +4026,7 @@ class GuiTests(TempDirMixin, unittest.TestCase):
self.addCleanup(tab.close)
self.assertTrue(tab.cmhub_balance_label.isHidden())
self.assertEqual("cmhub余额:未获取", tab.cmhub_balance_label.text())
self.assertEqual("默认网关余额:未获取", tab.cmhub_balance_label.text())
tab._on_generate_progress(
{
@@ -4006,7 +4039,7 @@ class GuiTests(TempDirMixin, unittest.TestCase):
"points_balance": 88,
}
)
self.assertEqual("cmhub余额:88", tab.cmhub_balance_label.text())
self.assertEqual("默认网关余额:88", tab.cmhub_balance_label.text())
self.assertTrue(tab.cmhub_balance_label.isHidden())
with mock.patch("app.gui.QMessageBox.warning") as warning:
+98
View File
@@ -149,6 +149,104 @@ class ImageStudioGenerationTests(TempDirMixin, unittest.TestCase):
self.assert_removed(temp_dir)
def test_direct_gateway_rejects_new_suite_job_before_creation(self):
with self.make_temp_dir() as temp_dir:
cfg, project, source = self._project_source(temp_dir)
cfg["ai"]["backend"] = "direct"
with mock.patch("app.image_studio_generation.image_studio.create_job") as create_job:
with self.assertRaises(image_studio_generation.ImageStudioGenerationError):
image_studio_generation.create_generation_jobs(
project.id,
source.id,
"不应提交",
1,
config=cfg,
path=cfg["db_path"],
)
create_job.assert_not_called()
self.assert_removed(temp_dir)
def test_resume_rejects_non_default_gateway_task_without_reading_gateway_config(self):
with self.make_temp_dir() as temp_dir:
cfg, project, source = self._project_source(temp_dir)
cfg["ai"]["backend"] = "direct"
job = image_studio.create_job(
project.id,
source_asset_id=source.id,
job_type="白底图",
prompt="旧任务",
generation_source="direct",
provider="direct",
path=cfg["db_path"],
)
job = image_studio.set_job_submitted(
job.id,
"custom-task-1",
path=cfg["db_path"],
)
with mock.patch("app.image_studio_generation._runtime") as runtime:
summary = image_studio_generation.run_jobs(
[job],
config=cfg,
path=cfg["db_path"],
)
runtime.assert_not_called()
self.assertEqual(1, summary["total"])
self.assertEqual(1, summary["failed"])
self.assertIn("不属于默认网关", summary["jobs"][0]["error"])
self.assert_removed(temp_dir)
def test_direct_selection_still_resumes_submitted_default_gateway_task(self):
with self.make_temp_dir() as temp_dir:
cfg, project, source = self._project_source(temp_dir)
cfg["ai"]["backend"] = "direct"
job = image_studio.create_job(
project.id,
source_asset_id=source.id,
job_type="白底图",
prompt="已扣点图片",
generation_source="cmhub",
provider="cmhub",
path=cfg["db_path"],
)
job = image_studio.set_job_submitted(
job.id,
"cmhub-task-1",
path=cfg["db_path"],
)
with mock.patch(
"app.image_studio_generation._runtime",
return_value=self._runtime(),
), mock.patch(
"app.image_studio_generation.ai._cmhub_call_once",
return_value={
"task_id": "cmhub-task-1",
"status": "succeeded",
"result": {"image_url": "https://cdn.example.com/result.png"},
},
) as poll, mock.patch(
"app.image_studio_generation.ai._download_cmhub_image_with_retry",
return_value=(self._png_bytes(), 0.1),
):
summary = image_studio_generation.run_jobs(
[job],
config=cfg,
path=cfg["db_path"],
)
poll.assert_called_once()
self.assertEqual(1, summary["success"])
self.assertEqual("succeeded", image_studio.get_job(job.id, path=cfg["db_path"]).status)
self.assert_removed(temp_dir)
def test_generate_image_jobs_sends_selected_aspect_ratio(self):
with self.make_temp_dir() as temp_dir:
cfg, project, source = self._project_source(temp_dir)
+35
View File
@@ -153,6 +153,41 @@ class ProductSuiteGuiTests(TempDirMixin, unittest.TestCase):
path=db_path,
)
def test_direct_gateway_blocks_new_suite_actions_but_keeps_default_job_resume_available(self):
with self.make_temp_dir() as temp_dir:
config = self._config(temp_dir)
config["ai"] = appconfig.ai_config(config)
config["ai"]["backend"] = "direct"
project, assets = self._create_project_with_assets(temp_dir, config, 1)
job = image_studio.create_job(
project.id,
source_asset_id=assets[0].id,
job_type="白底图",
prompt="已提交图片",
generation_source="cmhub",
provider="cmhub",
path=config["db_path"],
)
image_studio.set_job_submitted(job.id, "cmhub-task-1", path=config["db_path"])
tab = ProductSuiteTab(config=config, db_path=config["db_path"])
self.addCleanup(tab.close)
state = tab._displayed_state
state.account_alias = "alias-a"
state.item_id = project.item_id
state.project_id = project.id
state.project_binding_state = project.binding_state
tab._load_state(state)
tab.refresh_gateway_state()
self.assertFalse(tab.generate_button.isEnabled())
self.assertFalse(tab.ai_write_button.isEnabled())
self.assertFalse(tab.resume_submitted_button.isHidden())
self.assertTrue(tab.resume_submitted_button.isEnabled())
self.assertIn("仅支持默认网关", tab.generate_button.toolTip())
self.assert_removed(temp_dir)
def test_tab_builds_suite_controls_without_old_detail_workspace(self):
with self.make_temp_dir() as temp_dir:
config = self._config(temp_dir)
+9 -6
View File
@@ -423,13 +423,16 @@ class WorkerTests(unittest.TestCase):
) as analyze, mock.patch("app.gui.workers.ai.gen_title") as gen_title:
result = worker.execute()
analyze.assert_called_once_with(
"补充要求",
"输出语言:繁体中文",
["first.jpg", "second.jpg"],
config={"ai": {"backend": "cmhub"}},
cmhub_config_path="cmhub.json",
analyze.assert_called_once()
args, kwargs = analyze.call_args
self.assertEqual(
("补充要求", "输出语言:繁体中文", ["first.jpg", "second.jpg"]),
args,
)
self.assertEqual(worker.config, kwargs["config"])
self.assertEqual("cmhub.json", kwargs["cmhub_config_path"])
self.assertIn("cmhub_api_key", worker.config["_cmshopee_ai_runtime"])
self.assertNotIn("direct_models", worker.config["_cmshopee_ai_runtime"])
gen_title.assert_not_called()
self.assertEqual(expected, result)