feat(product-suite): enable direct generation flow
This commit is contained in:
@@ -115,6 +115,18 @@ class ProductSuiteGuiTests(TempDirMixin, unittest.TestCase):
|
||||
)
|
||||
return project, assets
|
||||
|
||||
def _configure_direct_image_gateway(self, config):
|
||||
models_path = appconfig.ai_models_config_path(config)
|
||||
ai_config = appconfig.ai_config(config)
|
||||
ai_config["backend"] = "direct"
|
||||
config["ai"] = ai_config
|
||||
config["ai_models_path"] = models_path
|
||||
models_config = appconfig.default_ai_models_config()
|
||||
for model in models_config["models"]:
|
||||
if model["category"] == "image":
|
||||
model["api_key"] = "test-image-key"
|
||||
appconfig.save_ai_models_config(models_config, path=models_path)
|
||||
|
||||
def _create_history_job(
|
||||
self,
|
||||
project,
|
||||
@@ -153,7 +165,7 @@ class ProductSuiteGuiTests(TempDirMixin, unittest.TestCase):
|
||||
path=db_path,
|
||||
)
|
||||
|
||||
def test_direct_gateway_blocks_new_suite_actions_but_keeps_default_job_resume_available(self):
|
||||
def test_direct_gateway_enables_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)
|
||||
@@ -180,11 +192,98 @@ class ProductSuiteGuiTests(TempDirMixin, unittest.TestCase):
|
||||
tab._load_state(state)
|
||||
tab.refresh_gateway_state()
|
||||
|
||||
self.assertFalse(tab.generate_button.isEnabled())
|
||||
self.assertTrue(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.assertIn("不计点数", tab.generate_button.toolTip())
|
||||
|
||||
self.assert_removed(temp_dir)
|
||||
|
||||
def test_direct_gateway_preflight_and_confirmation_do_not_call_cmhub_catalog(self):
|
||||
with self.make_temp_dir() as temp_dir:
|
||||
config = self._config(temp_dir)
|
||||
project, assets = self._create_project_with_assets(temp_dir, config, 2)
|
||||
self._configure_direct_image_gateway(config)
|
||||
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
|
||||
state.prompt = "轻便耐用,适合日常使用"
|
||||
state.settings["per_image_primary"] = False
|
||||
state.settings["ratio"] = "3:4"
|
||||
tab._load_state(state)
|
||||
confirmations = []
|
||||
tab._confirm = lambda title, message, **kwargs: confirmations.append(
|
||||
(title, message, kwargs)
|
||||
) and False
|
||||
|
||||
with mock.patch.object(
|
||||
tab,
|
||||
"_cmhub_catalog_params",
|
||||
side_effect=AssertionError("自定义网关不应读取 cmhub 目录"),
|
||||
):
|
||||
self.assertFalse(tab.start_generation(state))
|
||||
|
||||
self.assertEqual(1, len(confirmations))
|
||||
message = confirmations[0][1]
|
||||
self.assertIn("生成来源:自定义网关", message)
|
||||
self.assertIn("自定义网关不计点数", message)
|
||||
self.assertIn("主体一致性可能弱于默认网关", message)
|
||||
self.assertIn("输出尺寸:1024x1536(接近比例生成)", message)
|
||||
self.assertNotIn("cmhub 点数", message)
|
||||
self.assertEqual([], image_studio.list_jobs(project.id, path=config["db_path"]))
|
||||
|
||||
self.assert_removed(temp_dir)
|
||||
|
||||
def test_direct_gateway_invalid_image_model_blocks_before_confirmation(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)
|
||||
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
|
||||
state.prompt = "轻便耐用,适合日常使用"
|
||||
tab._load_state(state)
|
||||
messages = []
|
||||
tab._message = lambda title, message, **kwargs: messages.append((title, message))
|
||||
|
||||
self.assertFalse(tab.start_generation(state))
|
||||
self.assertEqual("自定义网关配置不完整", messages[0][0])
|
||||
self.assertIn("请到⑤设置", messages[0][1])
|
||||
self.assertEqual([], image_studio.list_jobs(project.id, path=config["db_path"]))
|
||||
|
||||
self.assert_removed(temp_dir)
|
||||
|
||||
def test_direct_gateway_stop_waits_for_current_image_result(self):
|
||||
with self.make_temp_dir() as temp_dir:
|
||||
config = self._config(temp_dir)
|
||||
config["ai"] = appconfig.ai_config(config)
|
||||
config["ai"]["backend"] = "direct"
|
||||
tab = ProductSuiteTab(config=config, db_path=config["db_path"])
|
||||
self.addCleanup(tab.close)
|
||||
state = tab._displayed_state
|
||||
state.worker = mock.Mock()
|
||||
state.generation_source = image_studio.GENERATION_SOURCE_DIRECT
|
||||
statuses = []
|
||||
tab._confirm = lambda *args, **kwargs: True
|
||||
tab._status = lambda message, level=None: statuses.append((message, level))
|
||||
|
||||
tab.toggle_generation()
|
||||
|
||||
state.worker.cancel.assert_called_once()
|
||||
self.assertTrue(state.generation_stop_requested)
|
||||
self.assertEqual("正在停止...", tab.generate_button.text())
|
||||
self.assertEqual(("正在停止,等待当前图片返回后结束", "warning"), statuses[-1])
|
||||
|
||||
self.assert_removed(temp_dir)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user