feat(settings): support gateway source switching
This commit is contained in:
@@ -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
@@ -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:
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user