import io import os import sys import threading import time import unittest import base64 from concurrent.futures import CancelledError from unittest import mock sys.path.insert(0, os.path.dirname(__file__)) from _helpers import TempDirMixin from app import ai, appconfig, db, image_studio, image_studio_generation class ImageStudioGenerationTests(TempDirMixin, unittest.TestCase): def _png_bytes(self): from PIL import Image output = io.BytesIO() Image.new("RGB", (32, 32), (120, 80, 160)).save(output, format="PNG") return output.getvalue() def _config(self, temp_dir): return { "data_dir": temp_dir, "db_path": os.path.join(temp_dir, "cmshopee.db"), "image_dir": os.path.join(temp_dir, "images"), "ai": { "backend": "cmhub", "image_concurrency": 1, "retry": 0, "resolution": "1k", "jpg_quality": 90, "cmhub": { "base_url": "https://cmhub.example.com", "image_alias": "image-hd", "connect_timeout": 3, "download_with_curl": "false", }, }, } def _project_source(self, temp_dir): cfg = self._config(temp_dir) db.init_db(cfg["db_path"]) project = image_studio.create_or_get_project( account_alias="alias", account_slug="alias_slug", item_id="51100639510", path=cfg["db_path"], ) source_path = os.path.join(temp_dir, "source.png") with open(source_path, "wb") as fh: fh.write(self._png_bytes()) source = image_studio.add_asset( project.id, image_studio.ASSET_KIND_ORIGINAL, remote_url="https://cdn.example.com/source.png", local_path=source_path, source_order=1, path=cfg["db_path"], ) return cfg, project, source def _runtime(self): return { "base_url": "https://cmhub.example.com", "api_key": "sk-test", "alias": "image-hd", "connect_timeout": 3, "use_system_proxy": False, "download_with_curl": "false", } def _direct_config(self, temp_dir): cfg = self._config(temp_dir) models_path = os.path.join(temp_dir, "ai_models.json") cfg["ai"]["backend"] = "direct" cfg["ai"]["default_text_model"] = "Text" cfg["ai"]["default_image_model"] = "Direct Image" cfg["ai_models_path"] = models_path appconfig.save_ai_models_config( { "models": [ { "name": "Text", "category": "text", "enabled": True, "url": "https://text.example.com/v1", "model": "text-model", "api_key": "sk-text", "api_type": "chat", "connect_timeout_seconds": 3, "timeout_seconds": 10, "extra_body": {}, }, { "name": "Direct Image", "category": "image", "enabled": True, "url": "https://image.example.com/v1", "model": "image-model", "api_key": "sk-image", "api_type": "images_edits", "connect_timeout_seconds": 7, "timeout_seconds": 12, "extra_body": {}, }, ] }, path=models_path, ) return ai.freeze_runtime_config( cfg, models_path=models_path, include_cmhub=False, include_direct_models=True, ) def test_generate_image_jobs_creates_independent_jobs_and_assets(self): with self.make_temp_dir() as temp_dir: cfg, project, source = self._project_source(temp_dir) submitted = [] poll_seen_persisted = [] def fake_submit(method, url, api_key, **kwargs): self.assertEqual("POST", method) task_id = "cmhub-task-%d" % (len(submitted) + 1) submitted.append((task_id, kwargs["headers_extra"]["Idempotency-Key"])) return { "task_id": task_id, "status": "queued", "call_id": "call-%d" % len(submitted), "points_cost": 2, "points_balance": 100 - len(submitted) * 2, } def fake_poll(method, url, api_key, **kwargs): self.assertEqual("GET", method) task_id = url.rsplit("/", 1)[-1] jobs = image_studio.list_resumable_jobs(path=cfg["db_path"], project_id=project.id) poll_seen_persisted.append(any(job.task_id == task_id for job in jobs)) return { "task_id": task_id, "status": "succeeded", "result": {"image_url": "https://cdn.example.com/%s.png" % task_id}, } with mock.patch("app.image_studio_generation._runtime", return_value=self._runtime()), \ mock.patch( "app.image_studio_generation.ai._cmhub_call_with_retry", side_effect=fake_submit, ), \ mock.patch( "app.image_studio_generation.ai._cmhub_call_once", side_effect=fake_poll, ), \ mock.patch( "app.image_studio_generation.ai._download_cmhub_image_with_retry", return_value=(self._png_bytes(), 0.1), ), \ mock.patch("app.image_studio_generation.ai._sleep_cmhub_poll"): summary = image_studio_generation.generate_image_jobs( project.id, source.id, "完整提示词", 3, config=cfg, path=cfg["db_path"], ) self.assertEqual(3, summary["success"]) self.assertEqual(0, summary["failed"]) self.assertEqual(3, len(submitted)) self.assertEqual(3, len({key for _, key in submitted})) self.assertEqual([True, True, True], poll_seen_persisted) jobs = image_studio.list_resumable_jobs(path=cfg["db_path"], project_id=project.id) self.assertEqual([], jobs) all_jobs = [ image_studio.get_job(result["job"].id, path=cfg["db_path"]) for result in summary["jobs"] ] self.assertEqual(["succeeded", "succeeded", "succeeded"], [job.status for job in all_jobs]) self.assertTrue(all(job.task_id for job in all_jobs)) self.assertTrue(all(job.task_key for job in all_jobs)) self.assertEqual([2, 2, 2], [job.points_cost for job in all_jobs]) assets = image_studio.list_assets(project.id, kind="generated_main", path=cfg["db_path"]) self.assertEqual(3, len(assets)) self.assertTrue(all(os.path.isfile(asset.local_path) for asset in assets)) self.assertTrue(all(asset.parent_asset_id == source.id for asset in assets)) self.assertTrue(all(asset.prompt == "完整提示词" for asset in assets)) self.assert_removed(temp_dir) def test_direct_gateway_creates_direct_suite_job_before_execution(self): with self.make_temp_dir() as temp_dir: cfg, project, source = self._project_source(temp_dir) cfg["ai"]["backend"] = "direct" jobs = image_studio_generation.create_generation_jobs( project.id, source.id, "提交前来源快照", 1, config=cfg, path=cfg["db_path"], ) self.assertEqual(1, len(jobs)) self.assertEqual("direct", jobs[0].generation_source) self.assertEqual("openai_images_edits", jobs[0].provider) self.assertIsNone(jobs[0].task_id) self.assert_removed(temp_dir) def test_direct_job_with_invalid_task_id_is_not_sent_to_any_gateway(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="openai_images_edits", path=cfg["db_path"], ) conn = db.connect(cfg["db_path"]) try: conn.execute( "UPDATE image_studio_jobs SET task_id = 'custom-task-1' WHERE id = ?", (job.id,), ) conn.commit() finally: conn.close() job = image_studio.get_job(job.id, path=cfg["db_path"]) with mock.patch("app.image_studio_generation._runtime") as runtime, \ mock.patch("app.image_studio_generation._direct_runtime") as direct_runtime: summary = image_studio_generation.run_jobs( [job], config=cfg, path=cfg["db_path"], ) runtime.assert_not_called() direct_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_job_uses_one_ordered_edit_request_and_records_output_dimensions(self): with self.make_temp_dir() as temp_dir: _cfg, project, source = self._project_source(temp_dir) cfg = self._direct_config(temp_dir) reference_path = os.path.join(temp_dir, "reference.png") with open(reference_path, "wb") as fh: fh.write(self._png_bytes()) reference = image_studio.add_asset( project.id, image_studio.ASSET_KIND_ORIGINAL, local_path=reference_path, source_order=2, path=cfg["db_path"], ) job = image_studio.create_job( project.id, source_asset_id=source.id, reference_asset_ids=[reference.id], job_type="场景图", prompt="已冻结的套图提示词", generation_source="direct", provider="openai_images_edits", path=cfg["db_path"], ) seen = [] encoded = base64.b64encode(self._png_bytes()).decode("ascii") def fake_direct_call(model, body, config, attempts, **kwargs): seen.append((model, body, config, attempts, kwargs)) return {"data": [{"b64_json": encoded}]} with mock.patch( "app.image_studio_generation.ai._call_with_retry", side_effect=fake_direct_call, ), mock.patch("app.image_studio_generation._runtime") as cmhub_runtime, \ mock.patch("app.image_studio_generation.ai._cmhub_call_with_retry") as cmhub_submit, \ mock.patch("app.image_studio_generation.ai._cmhub_call_once") as cmhub_poll: summary = image_studio_generation.run_jobs( [job], aspect_ratio="3:4", config=cfg, path=cfg["db_path"], run_session_id="direct-test-session", ) self.assertEqual(1, summary["success"]) self.assertEqual("1024x1536", summary["output"]["requested_output_size"]) self.assertTrue(summary["output"]["approximate_ratio"]) self.assertEqual(1, len(seen)) _model, body, _config, attempts, kwargs = seen[0] self.assertEqual(1, attempts) self.assertEqual("multipart", kwargs["request_kind"]) self.assertEqual(2, body.count(b'name="image[]"')) self.assertLess(body.index(b"source.png"), body.index(b"reference.png")) self.assertIn(b'name="n"', body) self.assertIn(b"\r\n1\r\n", body) cmhub_runtime.assert_not_called() cmhub_submit.assert_not_called() cmhub_poll.assert_not_called() stored = image_studio.get_job(job.id, path=cfg["db_path"]) self.assertEqual("succeeded", stored.status) self.assertEqual("direct", stored.generation_source) self.assertEqual("openai_images_edits", stored.provider) self.assertIsNone(stored.task_id) self.assertEqual("direct-test-session", stored.run_session_id) asset = image_studio.get_asset(stored.output_asset_id, path=cfg["db_path"]) self.assertIsNone(asset.remote_url) self.assertEqual("3:4", asset.aspect_ratio) self.assertEqual("1024x1536", asset.requested_output_size) self.assertEqual((1024, 1536), (asset.rendered_width, asset.rendered_height)) self.assert_removed(temp_dir) def test_direct_job_saves_returned_image_after_stop_and_never_retries(self): with self.make_temp_dir() as temp_dir: _cfg, project, source = self._project_source(temp_dir) cfg = self._direct_config(temp_dir) job = image_studio.create_job( project.id, source_asset_id=source.id, prompt="停止后保存", generation_source="direct", provider="openai_images_edits", path=cfg["db_path"], ) stopped = {"value": False} encoded = base64.b64encode(self._png_bytes()).decode("ascii") def fake_direct_call(*args, **kwargs): self.assertEqual(1, args[3]) stopped["value"] = True return {"data": [{"b64_json": encoded}]} with mock.patch( "app.image_studio_generation.ai._call_with_retry", side_effect=fake_direct_call, ) as direct_call: summary = image_studio_generation.run_jobs( [job], config=cfg, path=cfg["db_path"], should_stop=lambda: stopped["value"], ) self.assertEqual(1, summary["success"]) self.assertEqual(1, direct_call.call_count) self.assertEqual("succeeded", image_studio.get_job(job.id, path=cfg["db_path"]).status) self.assert_removed(temp_dir) def test_startup_recovery_marks_only_stale_direct_running_jobs_failed(self): with self.make_temp_dir() as temp_dir: cfg, project, source = self._project_source(temp_dir) stale = image_studio.create_job( project.id, source_asset_id=source.id, prompt="已中断", generation_source="direct", provider="openai_images_edits", path=cfg["db_path"], ) active = image_studio.create_job( project.id, source_asset_id=source.id, prompt="仍在执行", generation_source="direct", provider="openai_images_edits", path=cfg["db_path"], ) image_studio.update_job_status( stale.id, "running", run_session_id="previous-session", path=cfg["db_path"], ) image_studio.update_job_status( active.id, "running", run_session_id="active-session", path=cfg["db_path"], ) recovered = image_studio.fail_stale_direct_jobs( active_run_session_ids=["active-session"], path=cfg["db_path"], ) self.assertEqual([stale.id], [job.id for job in recovered]) stale = image_studio.get_job(stale.id, path=cfg["db_path"]) active = image_studio.get_job(active.id, path=cfg["db_path"]) self.assertEqual("failed", stale.status) self.assertIn("程序中断", stale.error) self.assertEqual(image_studio.JOB_RECOVERY_REGENERATE, stale.recovery_action) self.assertEqual("running", active.status) 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) submitted_payloads = [] def fake_submit(method, url, api_key, **kwargs): submitted_payloads.append(dict(kwargs["payload"])) return {"task_id": "cmhub-ratio", "status": "queued"} with mock.patch("app.image_studio_generation._runtime", return_value=self._runtime()), \ mock.patch( "app.image_studio_generation.ai._cmhub_call_with_retry", side_effect=fake_submit, ), \ mock.patch( "app.image_studio_generation.ai._cmhub_call_once", return_value={ "task_id": "cmhub-ratio", "status": "succeeded", "result": {"image_url": "https://cdn.example.com/ratio.png"}, }, ), \ mock.patch( "app.image_studio_generation.ai._download_cmhub_image_with_retry", return_value=(self._png_bytes(), 0.1), ): summary = image_studio_generation.generate_image_jobs( project.id, source.id, "比例测试", 1, aspect_ratio="3:4", config=cfg, path=cfg["db_path"], ) self.assertEqual(1, summary["success"]) self.assertEqual("3:4", submitted_payloads[0]["aspect_ratio"]) self.assertEqual(1, len(submitted_payloads[0]["images"])) self.assertIn("image_base64", submitted_payloads[0]["images"][0]) self.assertNotIn("image_base64", submitted_payloads[0]) self.assert_removed(temp_dir) def test_build_cmhub_images_keeps_order_and_limits_to_eight_inputs(self): with self.make_temp_dir() as temp_dir: cfg, project, source = self._project_source(temp_dir) references = [] for index in range(2, 10): path = os.path.join(temp_dir, "source-%d.png" % index) with open(path, "wb") as fh: fh.write(self._png_bytes()) references.append( image_studio.add_asset( project.id, image_studio.ASSET_KIND_ORIGINAL, local_path=path, source_order=index, path=cfg["db_path"], ) ) with mock.patch( "app.image_studio_generation.ai._image_data_url", side_effect=lambda path: "data:image/png;base64,%s" % os.path.basename(path), ): images, omitted_count = image_studio_generation._build_cmhub_images( source, references, ) self.assertEqual(8, len(images)) self.assertEqual(1, omitted_count) self.assertTrue(images[0]["image_base64"].endswith("source.png")) self.assertTrue(images[-1]["image_base64"].endswith("source-8.png")) self.assert_removed(temp_dir) def test_build_cmhub_images_rejects_oversized_total_payload(self): with self.make_temp_dir() as temp_dir: _, _, source = self._project_source(temp_dir) with mock.patch( "app.image_studio_generation.ai._image_data_url", return_value="x" * 32, ), mock.patch.object( image_studio_generation, "CMHUB_IMAGE_STUDIO_MAX_TOTAL_INPUT_BYTES", 16, ): with self.assertRaisesRegex( image_studio_generation.ImageStudioGenerationError, "总大小超过32MiB", ): image_studio_generation._build_cmhub_images(source) self.assert_removed(temp_dir) def test_job_reference_snapshot_submits_ordered_images(self): with self.make_temp_dir() as temp_dir: cfg, project, source = self._project_source(temp_dir) reference_path = os.path.join(temp_dir, "reference.png") with open(reference_path, "wb") as fh: fh.write(self._png_bytes()) reference = image_studio.add_asset( project.id, image_studio.ASSET_KIND_ORIGINAL, local_path=reference_path, source_order=2, path=cfg["db_path"], ) job = image_studio.create_job( project.id, source_asset_id=source.id, reference_asset_ids=[reference.id], prompt="多图提示词", path=cfg["db_path"], ) submitted = [] def fake_submit(method, url, api_key, **kwargs): submitted.append(kwargs["payload"]) return {"task_id": "multi-image-task", "status": "queued"} with mock.patch("app.image_studio_generation._runtime", return_value=self._runtime()), \ mock.patch("app.image_studio_generation.ai._cmhub_call_with_retry", side_effect=fake_submit), \ mock.patch( "app.image_studio_generation.ai._cmhub_call_once", return_value={ "task_id": "multi-image-task", "status": "succeeded", "result": {"image_url": "https://cdn.example.com/multi.png"}, }, ), \ 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"], ) self.assertEqual(1, summary["success"]) self.assertEqual(2, len(submitted[0]["images"])) self.assertNotIn("image_base64", submitted[0]) self.assert_removed(temp_dir) def test_missing_reference_snapshot_fails_without_submitting(self): with self.make_temp_dir() as temp_dir: cfg, project, source = self._project_source(temp_dir) reference = image_studio.add_asset( project.id, image_studio.ASSET_KIND_ORIGINAL, local_path=os.path.join(temp_dir, "missing-reference.png"), source_order=2, path=cfg["db_path"], ) job = image_studio.create_job( project.id, source_asset_id=source.id, reference_asset_ids=[reference.id], prompt="多图提示词", 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_with_retry") as submit: summary = image_studio_generation.run_jobs( [job], config=cfg, path=cfg["db_path"], ) self.assertEqual(1, summary["failed"]) self.assertIn("参考图尚未下载", summary["jobs"][0]["error"]) submit.assert_not_called() self.assert_removed(temp_dir) def test_stop_after_download_discards_temporary_result(self): with self.make_temp_dir() as temp_dir: cfg, project, source = self._project_source(temp_dir) stopped = {"value": False} saved_paths = [] def fake_download( request_result, out_path, config, on_event, job_id, should_stop=None, ): os.makedirs(os.path.dirname(out_path), exist_ok=True) with open(out_path, "wb") as fh: fh.write(self._png_bytes()) saved_paths.append(out_path) stopped["value"] = True return out_path with mock.patch("app.image_studio_generation._runtime", return_value=self._runtime()), \ mock.patch( "app.image_studio_generation.ai._cmhub_call_with_retry", return_value={"task_id": "cmhub-stop", "status": "queued"}, ), \ mock.patch( "app.image_studio_generation.ai._cmhub_call_once", return_value={ "task_id": "cmhub-stop", "status": "succeeded", "result": {"image_url": "https://cdn.example.com/stop.png"}, }, ), \ mock.patch( "app.image_studio_generation._download_and_save_job_image", side_effect=fake_download, ): summary = image_studio_generation.generate_image_jobs( project.id, source.id, "停止测试", 1, config=cfg, path=cfg["db_path"], should_stop=lambda: stopped["value"], ) self.assertEqual(1, summary["cancelled"]) self.assertEqual([], image_studio.list_assets(project.id, kind="generated_main", path=cfg["db_path"])) self.assertEqual(1, len(saved_paths)) self.assertFalse(os.path.exists(saved_paths[0])) self.assert_removed(temp_dir) def test_download_cancel_maps_to_cancelled_and_keeps_resume_action(self): with self.make_temp_dir() as temp_dir: cfg, project, source = self._project_source(temp_dir) stopped = {"value": False} def fake_download(*args, **kwargs): self.assertIsNotNone(kwargs.get("should_stop")) stopped["value"] = True raise CancelledError() with mock.patch( "app.image_studio_generation._runtime", return_value=self._runtime(), ), mock.patch( "app.image_studio_generation.ai._cmhub_call_with_retry", return_value={"task_id": "cmhub-cancel-download", "status": "queued"}, ), mock.patch( "app.image_studio_generation.ai._cmhub_call_once", return_value={ "task_id": "cmhub-cancel-download", "status": "succeeded", "result": { "image_url": "https://cdn.example.com/cancel-download.png" }, }, ), mock.patch( "app.image_studio_generation.ai._download_cmhub_image_with_retry", side_effect=fake_download, ): summary = image_studio_generation.generate_image_jobs( project.id, source.id, "下载停止测试", 1, config=cfg, path=cfg["db_path"], should_stop=lambda: stopped["value"], ) self.assertEqual(1, summary["cancelled"]) job = summary["jobs"][0]["job"] stored = image_studio.get_job(job.id, path=cfg["db_path"]) self.assertEqual("cancelled", stored.status) self.assertEqual(image_studio.JOB_RECOVERY_RESUME, stored.recovery_action) self.assertEqual( [], image_studio.list_assets( project.id, kind="generated_main", path=cfg["db_path"], ), ) self.assert_removed(temp_dir) def test_run_jobs_cancels_queued_future_before_running_slot_is_released(self): with self.make_temp_dir() as temp_dir: cfg, project, source = self._project_source(temp_dir) cfg["ai"]["image_concurrency"] = 1 jobs = [ image_studio.create_job( project.id, source_asset_id=source.id, job_type="main", prompt="排队停止测试", path=cfg["db_path"], ) for _ in range(2) ] started_ids = [] first_started = threading.Event() release_first = threading.Event() stop_requested = threading.Event() result_holder = {} def fake_run_one(*args): job_id = int(args[0]) started_ids.append(job_id) if len(started_ids) == 1: first_started.set() release_first.wait(timeout=3) updated = image_studio.update_job_status( job_id, "cancelled", error="用户停止", recovery_action=image_studio.JOB_RECOVERY_REGENERATE, path=cfg["db_path"], ) return { "job": updated, "status": "cancelled", "error": "用户停止", } def run(): result_holder["summary"] = image_studio_generation.run_jobs( jobs, config=cfg, path=cfg["db_path"], should_stop=stop_requested.is_set, ) with mock.patch( "app.image_studio_generation._runtime", return_value=self._runtime(), ), mock.patch( "app.image_studio_generation._run_one_job_with_global_slot", side_effect=fake_run_one, ): thread = threading.Thread(target=run) thread.start() self.assertTrue(first_started.wait(timeout=2)) stop_requested.set() time.sleep(0.35) release_first.set() thread.join(timeout=3) self.assertFalse(thread.is_alive()) self.assertEqual([jobs[0].id], started_ids) summary = result_holder["summary"] self.assertEqual(2, summary["cancelled"]) self.assertEqual(2, len(summary["jobs"])) self.assertEqual( ["cancelled", "cancelled"], [ image_studio.get_job(job.id, path=cfg["db_path"]).status for job in jobs ], ) self.assert_removed(temp_dir) def test_resume_existing_job_polls_without_new_submit(self): with self.make_temp_dir() as temp_dir: cfg, project, source = self._project_source(temp_dir) job = image_studio.create_job( project.id, source_asset_id=source.id, job_type="main", prompt="续查提示词", path=cfg["db_path"], ) image_studio.set_job_submitted(job.id, "cmhub-task-resume", path=cfg["db_path"]) def fake_poll(method, url, api_key, **kwargs): return { "task_id": "cmhub-task-resume", "status": "succeeded", "result": {"image_url": "https://cdn.example.com/resume.png"}, } with mock.patch("app.image_studio_generation._runtime", return_value=self._runtime()), \ mock.patch("app.image_studio_generation.ai._cmhub_call_with_retry") as submit, \ mock.patch("app.image_studio_generation.ai._cmhub_call_once", side_effect=fake_poll), \ mock.patch( "app.image_studio_generation.ai._download_cmhub_image_with_retry", return_value=(self._png_bytes(), 0.1), ): summary = image_studio_generation.resume_image_jobs( project_id=project.id, config=cfg, path=cfg["db_path"], ) self.assertEqual(1, summary["success"]) submit.assert_not_called() updated = image_studio.get_job(job.id, path=cfg["db_path"]) self.assertEqual("succeeded", updated.status) self.assertEqual("cmhub-task-resume", updated.task_id) self.assert_removed(temp_dir) def test_download_failure_does_not_submit_again_or_create_asset(self): with self.make_temp_dir() as temp_dir: cfg, project, source = self._project_source(temp_dir) submit_count = 0 def fake_submit(method, url, api_key, **kwargs): nonlocal submit_count submit_count += 1 return {"task_id": "cmhub-task-1", "status": "queued"} def fake_poll(method, url, api_key, **kwargs): return { "task_id": "cmhub-task-1", "status": "succeeded", "result": {"image_url": "https://cdn.example.com/fail.png"}, } with mock.patch("app.image_studio_generation._runtime", return_value=self._runtime()), \ mock.patch( "app.image_studio_generation.ai._cmhub_call_with_retry", side_effect=fake_submit, ), \ mock.patch("app.image_studio_generation.ai._cmhub_call_once", side_effect=fake_poll), \ mock.patch( "app.image_studio_generation.ai._download_cmhub_image_with_retry", side_effect=RuntimeError("download failed"), ): summary = image_studio_generation.generate_image_jobs( project.id, source.id, "完整提示词", 1, config=cfg, path=cfg["db_path"], ) self.assertEqual(0, summary["success"]) self.assertEqual(1, summary["failed"]) self.assertEqual(1, submit_count) jobs = [result["job"] for result in summary["jobs"]] self.assertEqual("failed", image_studio.get_job(jobs[0].id, path=cfg["db_path"]).status) self.assertEqual([], image_studio.list_assets(project.id, kind="generated_main", path=cfg["db_path"])) self.assert_removed(temp_dir) def test_resume_failed_download_job_polls_without_new_submit(self): with self.make_temp_dir() as temp_dir: cfg, project, source = self._project_source(temp_dir) job = image_studio.create_job( project.id, source_asset_id=source.id, job_type="main", prompt="恢复下载", path=cfg["db_path"], ) image_studio.set_job_submitted( job.id, "cmhub-task-download", call_id="call-download", points_cost=2, points_balance=88, path=cfg["db_path"], ) image_studio.update_job_status( job.id, "failed", error="下载新封面失败", path=cfg["db_path"], ) failed_job = image_studio.get_job(job.id, path=cfg["db_path"]) self.assertEqual(image_studio.JOB_RECOVERY_RESUME, failed_job.recovery_action) resumable = image_studio.list_resumable_jobs( path=cfg["db_path"], project_id=project.id, include_failed_downloads=True, ) self.assertEqual([job.id], [item.id for item in resumable]) def fake_poll(method, url, api_key, **kwargs): self.assertEqual("GET", method) return { "task_id": "cmhub-task-download", "status": "succeeded", "result": {"image_url": "https://cdn.example.com/recovered.png"}, "points_balance": 88, } with mock.patch("app.image_studio_generation._runtime", return_value=self._runtime()), \ mock.patch("app.image_studio_generation.ai._cmhub_call_with_retry") as submit, \ mock.patch("app.image_studio_generation.ai._cmhub_call_once", side_effect=fake_poll), \ mock.patch( "app.image_studio_generation.ai._download_cmhub_image_with_retry", return_value=(self._png_bytes(), 0.1), ): summary = image_studio_generation.resume_image_jobs( project_id=project.id, config=cfg, path=cfg["db_path"], ) self.assertEqual(1, summary["success"]) submit.assert_not_called() updated = image_studio.get_job(job.id, path=cfg["db_path"]) self.assertEqual("succeeded", updated.status) self.assertEqual("cmhub-task-download", updated.task_id) self.assertEqual("call-download", updated.call_id) self.assertEqual(image_studio.JOB_RECOVERY_NONE, updated.recovery_action) assets = image_studio.list_assets(project.id, kind="generated_main", path=cfg["db_path"]) self.assertEqual(1, len(assets)) self.assert_removed(temp_dir) def test_failed_cmhub_task_marks_only_that_job_failed(self): with self.make_temp_dir() as temp_dir: cfg, project, source = self._project_source(temp_dir) submitted = [] def fake_submit(method, url, api_key, **kwargs): task_id = "cmhub-task-%d" % (len(submitted) + 1) submitted.append(task_id) return {"task_id": task_id, "status": "queued"} def fake_poll(method, url, api_key, **kwargs): task_id = url.rsplit("/", 1)[-1] if task_id.endswith("-2"): return {"task_id": task_id, "status": "failed", "error": {"message": "上游失败"}} return { "task_id": task_id, "status": "succeeded", "result": {"image_url": "https://cdn.example.com/%s.png" % task_id}, } with mock.patch("app.image_studio_generation._runtime", return_value=self._runtime()), \ mock.patch( "app.image_studio_generation.ai._cmhub_call_with_retry", side_effect=fake_submit, ), \ mock.patch("app.image_studio_generation.ai._cmhub_call_once", side_effect=fake_poll), \ mock.patch( "app.image_studio_generation.ai._download_cmhub_image_with_retry", return_value=(self._png_bytes(), 0.1), ): summary = image_studio_generation.generate_image_jobs( project.id, source.id, "完整提示词", 2, config=cfg, path=cfg["db_path"], ) self.assertEqual(1, summary["success"]) self.assertEqual(1, summary["failed"]) assets = image_studio.list_assets(project.id, kind="generated_main", path=cfg["db_path"]) self.assertEqual(1, len(assets)) failed_jobs = [ job for job in image_studio.list_resumable_jobs( path=cfg["db_path"], project_id=project.id, include_failed_downloads=True, ) if job.status == "failed" ] self.assertEqual([], failed_jobs) self.assert_removed(temp_dir) if __name__ == "__main__": unittest.main()