import base64 import io import json import os import socket import sys import threading import unittest from concurrent.futures import CancelledError from types import SimpleNamespace from unittest import mock sys.path.insert(0, os.path.dirname(__file__)) from _helpers import TempDirMixin from app import ai, appconfig, db class _Response: def __init__(self, payload): self.payload = payload def __enter__(self): return self def __exit__(self, exc_type, exc, tb): return False def read(self, size=-1): return json.dumps(self.payload).encode("utf-8") class _RequestsResponse: def __init__(self, payload=None, status_code=200, content=b"", headers=None): self.payload = payload if payload is not None else {} self.status_code = status_code self.content = content self.headers = headers or {} self.text = json.dumps(self.payload, ensure_ascii=False) self.closed = False def json(self): return self.payload def iter_content(self, chunk_size=65536): if self.content: yield self.content def close(self): self.closed = True class AITests(TempDirMixin, unittest.TestCase): def _write_models(self, path, text=None, image=None): text = text or { "name": "Text", "category": "text", "enabled": True, "url": "https://example.invalid/v1/chat/completions", "model": "text-model", "api_key": "sk-text-secret", "api_type": "chat", "connect_timeout_seconds": 1, "timeout_seconds": 1, "extra_body": {"temperature": 0}, } image = image or { "name": "Image", "category": "image", "enabled": True, "url": "https://example.invalid/v1", "model": "image-model", "api_key": "sk-image-secret", "api_type": "images_edits", "connect_timeout_seconds": 1, "timeout_seconds": 1, "extra_body": {}, } appconfig.save_ai_models_config({"models": [text, image]}, path=path) def _config(self): cfg = appconfig.default_config() cfg["ai"]["backend"] = "direct" cfg["ai"]["default_text_model"] = "Text" cfg["ai"]["default_image_model"] = "Image" cfg["ai"]["retry"] = 1 cfg["ai"]["resolution"] = "512" cfg["ai"]["jpg_quality"] = 80 return cfg def _cmhub_config(self, temp_dir): cfg = self._config() cfg["ai"]["backend"] = "cmhub" cfg["ai"]["resolution"] = "512" cfg["ai"]["cmhub"] = { "base_url": "https://cmhub.example.com", "title_alias": "title-standard", "image_alias": "image-hd", "vision_alias": "vision-standard", "connect_timeout": 3, "download_with_curl": "false", "check_balance_before_batch": False, } key_path = os.path.join(temp_dir, "cmhub.json") 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"]) batch_id = db.create_batch([os.path.join(temp_dir, "input.xlsx")], path=cfg["db_path"]) db.insert_tasks( batch_id, [ { "source_file_abs": os.path.join(temp_dir, "input.xlsx"), "source_sheet": "商品", "source_row": index + 2, "account_name": "Excel主店", "alias": "alias-a", "item_id": "5110063951%s" % index, } for index in range(len(titles)) ], path=cfg["db_path"], ) tasks = db.list_tasks(batch_id=batch_id, path=cfg["db_path"]) for task, title in zip(tasks, titles): db.set_collected( task.id, title, os.path.join(temp_dir, "%s_old.jpg" % task.item_id), path=cfg["db_path"], ) return batch_id, db.list_tasks(batch_id=batch_id, path=cfg["db_path"]) def test_cover_title_context_and_cover_only_generation_needs(self): task = SimpleNamespace( stage="collected", status="success", old_title="旧标题", new_title=None, new_cover_path=None, apply_attempts=0, ) self.assertEqual("旧标题", ai.cover_title_context(task)) self.assertEqual( {"title": False, "cover": True}, ai.generation_needs(task, generate_mode="cover"), ) task.new_title = "新标题" self.assertEqual("新标题", ai.cover_title_context(task)) task.new_title = None task.old_title = "" self.assertEqual("", ai.cover_title_context(task)) self.assertFalse(ai.is_generatable_task(task, generate_mode="cover")) self.assertEqual( {"title": True, "cover": True}, ai.generation_needs(task, generate_mode="title_cover"), ) def _write_old_cover_files(self, tasks): from PIL import Image for index, task in enumerate(tasks): os.makedirs(os.path.dirname(task.old_cover_path), exist_ok=True) Image.new( "RGB", (16, 16), ((index * 35) % 255, 80, 120), ).save(task.old_cover_path, "JPEG") def _png_bytes(self, color=(200, 120, 80)): from PIL import Image generated = io.BytesIO() Image.new("RGB", (8, 8), color).save(generated, "PNG") return generated.getvalue() def test_gen_title_uses_configured_model_and_retries(self): with self.make_temp_dir() as temp_dir: models_path = os.path.join(temp_dir, "ai_models.json") self._write_models(models_path) calls = [] steps = [] def fake_urlopen(request, timeout=None): calls.append((request, timeout)) if len(calls) == 1: raise ai.urllib.error.URLError("temporary") return _Response({"choices": [{"message": {"content": " 新标题 "}}]}) with mock.patch("app.ai.urllib.request.urlopen", side_effect=fake_urlopen): title = ai.gen_title( "优化标题", "旧标题", config=self._config(), models_path=models_path, on_step=steps.append, ) self.assertEqual("新标题", title) self.assertEqual(2, len(calls)) retry_events = [step for step in steps if isinstance(step, dict)] self.assertEqual(1, len(retry_events)) self.assertEqual("title_request", retry_events[0]["step"]) self.assertEqual("retry", retry_events[0]["result"]) self.assertEqual(1, retry_events[0]["attempt"]) self.assertEqual(2, retry_events[0]["attempts"]) self.assertNotIn("sk-text-secret", retry_events[0]["detail"]) body = json.loads(calls[-1][0].data.decode("utf-8")) self.assertEqual("text-model", body["model"]) self.assertEqual(0, body["temperature"]) self.assertEqual(1, len(body["messages"])) self.assertIn("优化标题", body["messages"][0]["content"]) self.assertIn("旧标题:\n旧标题", body["messages"][0]["content"]) self.assertIn("请只返回新标题,不要解释。", body["messages"][0]["content"]) self.assertNotIn("sk-text-secret", body["messages"][0]["content"]) self.assert_removed(temp_dir) def test_gen_title_renders_old_title_placeholder_without_duplicate_append(self): with self.make_temp_dir() as temp_dir: models_path = os.path.join(temp_dir, "ai_models.json") self._write_models(models_path) calls = [] def fake_urlopen(request, timeout=None): calls.append(request) return _Response({"choices": [{"message": {"content": "新标题"}}]}) with mock.patch("app.ai.urllib.request.urlopen", side_effect=fake_urlopen): title = ai.gen_title( "请基于{旧标题}重写标题,保留 {商品id}", "原始标题A", config=self._config(), models_path=models_path, ) self.assertEqual("新标题", title) body = json.loads(calls[-1].data.decode("utf-8")) content = body["messages"][0]["content"] self.assertIn("请基于原始标题A重写标题", content) self.assertEqual(1, content.count("原始标题A")) self.assertNotIn("旧标题:", content) self.assertIn("{商品id}", content) self.assertIn("请只返回新标题,不要解释。", content) self.assert_removed(temp_dir) def test_gen_title_accepts_openai_compatible_base_url(self): with self.make_temp_dir() as temp_dir: models_path = os.path.join(temp_dir, "ai_models.json") self._write_models( models_path, text={ "name": "Text", "category": "text", "enabled": True, "url": "https://example.invalid/v1", "model": "text-model", "api_key": "sk-text-secret", "api_type": "chat", "connect_timeout_seconds": 1, "timeout_seconds": 1, "extra_body": {}, }, ) calls = [] steps = [] def fake_urlopen(request, timeout=None): calls.append(request) return _Response({"choices": [{"message": {"content": "新标题"}}]}) with mock.patch("app.ai.urllib.request.urlopen", side_effect=fake_urlopen): title = ai.gen_title( "优化标题", "旧标题", config=self._config(), models_path=models_path, on_step=steps.append, ) self.assertEqual("新标题", title) self.assertEqual( "https://example.invalid/v1/chat/completions", calls[0].full_url, ) self.assert_removed(temp_dir) def test_missing_model_fields_raise_clear_error_without_secret(self): with self.make_temp_dir() as temp_dir: models_path = os.path.join(temp_dir, "ai_models.json") self._write_models( models_path, text={ "name": "Text", "category": "text", "enabled": True, "url": "", "model": "text-model", "api_key": "sk-text-secret", "api_type": "chat", "connect_timeout_seconds": 1, "timeout_seconds": 1, "extra_body": {}, }, ) with self.assertRaises(ai.AIError) as raised: ai.gen_title("prompt", "old", config=self._config(), models_path=models_path) message = str(raised.exception) self.assertIn("缺少字段: url", message) self.assertNotIn("sk-text-secret", message) self.assert_removed(temp_dir) def test_gen_cover_saves_jpeg_with_resolution_and_quality(self): try: from PIL import Image except ImportError: self.skipTest("Pillow not installed") with self.make_temp_dir() as temp_dir: models_path = os.path.join(temp_dir, "ai_models.json") self._write_models(models_path) old_cover = os.path.join(temp_dir, "old.jpg") output = os.path.join(temp_dir, "new.jpg") Image.new("RGB", (16, 16), (20, 30, 40)).save(old_cover, "JPEG") generated = io.BytesIO() Image.new("RGB", (8, 8), (200, 120, 80)).save(generated, "PNG") b64_image = base64.b64encode(generated.getvalue()).decode("ascii") def fake_urlopen(request, timeout=None): body = request.data self.assertIn(b'name="model"', body) self.assertIn(b"image-model", body) self.assertIn(b'name="image[]"', body) self.assertIn(b'name="n"', body) self.assertIn(b"\r\n1\r\n", body) self.assertIn("multipart/form-data", request.headers["Content-type"]) return _Response({"data": [{"b64_json": b64_image}]}) with mock.patch("app.ai.urllib.request.urlopen", side_effect=fake_urlopen): result = ai.gen_cover( "生成封面", old_cover, output, config=self._config(), models_path=models_path, ) self.assertEqual(os.path.abspath(output), result) with Image.open(output) as saved: self.assertEqual((512, 512), saved.size) self.assertEqual("JPEG", saved.format) self.assert_removed(temp_dir) def test_image_edit_body_keeps_repeated_image_fields_in_order(self): with self.make_temp_dir() as temp_dir: first = os.path.join(temp_dir, "primary.jpg") second = os.path.join(temp_dir, "reference.png") with open(first, "wb") as fh: fh.write(b"first-image") with open(second, "wb") as fh: fh.write(b"second-image") body, content_type = ai._image_edit_body( {"model": "image-model", "extra_body": {"quality": "high"}}, "生成商品图", [first, second], "1k", ) self.assertIn("multipart/form-data", content_type) self.assertEqual(2, body.count(b'name="image[]"')) self.assertLess(body.index(b"primary.jpg"), body.index(b"reference.png")) self.assertIn(b'name="n"', body) self.assertIn(b"\r\n1\r\n", body) self.assert_removed(temp_dir) def test_direct_image_response_only_accepts_openai_data_fields(self): encoded = base64.b64encode(b"image-bytes").decode("ascii") self.assertEqual(encoded, ai._find_image_ref({"data": [{"b64_json": encoded}]})) self.assertEqual( "https://images.example.com/new.png", ai._find_image_ref({"data": [{"url": "https://images.example.com/new.png"}]}), ) self.assertIsNone(ai._find_image_ref({"choices": [{"image_url": "x"}]})) with self.assertRaisesRegex(ai.AIError, "OpenAI 图片编辑接口"): ai._extract_image_bytes({"message": {"url": "https://bad.example/x"}}, {}, {}) with self.assertRaisesRegex(ai.AIError, "只允许 http/https"): ai._extract_image_bytes({"data": [{"url": "file:///tmp/image.png"}]}, {}, {}) def test_cmhub_gen_title_uses_alias_and_emits_metadata(self): with self.make_temp_dir() as temp_dir: cfg, key_path = self._cmhub_config(temp_dir) calls = [] events = [] def fake_request(method, url, **kwargs): calls.append((method, url, kwargs)) return _RequestsResponse( { "titles": [" 新标题 "], "alias": "title-standard", "model_used": "provider-title-model", "points_cost": 1, "points_balance": 99, "call_id": "call-title-1", } ) with mock.patch.object(ai._cmhub_session(), "request", side_effect=fake_request): title = ai.gen_title( "优化标题", "旧标题", config=cfg, cmhub_config_path=key_path, on_event=events.append, ) self.assertEqual("新标题", title) self.assertEqual(1, len(calls)) method, url, kwargs = calls[0] self.assertEqual("POST", method) self.assertEqual("https://cmhub.example.com/api/v1/generate/title", url) self.assertEqual((3, 600), kwargs["timeout"]) self.assertEqual("Bearer sk-cmhub-secret", kwargs["headers"]["Authorization"]) payload = kwargs["json"] self.assertEqual("title-standard", payload["model"]) self.assertEqual("512", payload["resolution"]) self.assertIn("优化标题", payload["prompt"]) self.assertIn("旧标题", payload["prompt"]) self.assertTrue(events) self.assertEqual("meta", events[0]["result"]) self.assertEqual(99, events[0]["metadata"]["points_balance"]) self.assert_removed(temp_dir) def test_cmhub_gen_title_renders_old_title_placeholder_once(self): with self.make_temp_dir() as temp_dir: cfg, key_path = self._cmhub_config(temp_dir) calls = [] def fake_request(method, url, **kwargs): calls.append((method, url, kwargs)) return _RequestsResponse({"titles": ["新标题"]}) with mock.patch.object(ai._cmhub_session(), "request", side_effect=fake_request): title = ai.gen_title( "请参考 {旧标题} 输出更短标题,店铺变量 {店铺}", "原始标题B", config=cfg, cmhub_config_path=key_path, ) self.assertEqual("新标题", title) prompt = calls[0][2]["json"]["prompt"] self.assertIn("请参考 原始标题B 输出更短标题", prompt) self.assertEqual(1, prompt.count("原始标题B")) self.assertNotIn("旧标题:", prompt) self.assertIn("{店铺}", prompt) self.assertIn("请只返回新标题,不要解释。", prompt) self.assert_removed(temp_dir) def test_cmhub_missing_config_raises_clear_error(self): with self.make_temp_dir() as temp_dir: cfg = self._config() cfg["ai"]["backend"] = "cmhub" cfg["ai"]["cmhub"] = { "base_url": "", "title_alias": "", "image_alias": "", "connect_timeout": 3, } with self.assertRaises(ai.CMHubError) as raised: ai.gen_title( "prompt", "old", config=cfg, cmhub_config_path=os.path.join(temp_dir, "missing.json"), ) self.assertEqual("cmhub_not_configured", raised.exception.code) self.assertIn("请去设置配置 cmhub", str(raised.exception)) self.assert_removed(temp_dir) def test_cmhub_analyze_product_images_uses_vision_alias_and_safe_metadata(self): with self.make_temp_dir() as temp_dir: cfg, key_path = self._cmhub_config(temp_dir) cfg["ai"]["backend"] = "direct" first = os.path.join(temp_dir, "first.png") second = os.path.join(temp_dir, "second.jpg") with open(first, "wb") as fh: fh.write(b"first-image") with open(second, "wb") as fh: fh.write(b"second-image") calls = [] def fake_request(method, url, **kwargs): calls.append((method, url, kwargs)) return _RequestsResponse( { "text": "浅绿色连帽上衣,突出宽松版型与日常穿搭场景。", "alias": "vision-standard", "model_used": "vision-provider", "points_cost": 1, "points_balance": 231, "call_id": "vision-call-1", } ) with mock.patch.object(ai._cmhub_session(), "request", side_effect=fake_request): result = ai.analyze_product_images( "补充卖点:避免夸大", "输出语言:繁体中文", [first, second], config=cfg, cmhub_config_path=key_path, ) self.assertEqual("浅绿色连帽上衣,突出宽松版型与日常穿搭场景。", result["text"]) self.assertEqual(2, result["image_count"]) self.assertEqual(1, result["metadata"]["points_cost"]) self.assertEqual(231, result["metadata"]["points_balance"]) self.assertNotIn("image_base64", result) self.assertEqual("POST", calls[0][0]) self.assertEqual(1, len(calls)) self.assertEqual( "https://cmhub.example.com/api/v1/analyze/images", calls[0][1], ) payload = calls[0][2]["json"] self.assertEqual("vision-standard", payload["model"]) self.assertEqual(2, len(payload["images"])) self.assertTrue(payload["images"][0]["image_base64"].startswith("data:image/png;base64,")) self.assertTrue(payload["images"][1]["image_base64"].startswith("data:image/jpeg;base64,")) self.assertEqual({"temperature": 0.2}, payload["parameters"]) prompt = payload["prompt"] self.assertIn("同一个商品项目", prompt) self.assertIn("一组证据", prompt) self.assertIn("禁止输出“图1/图2/第N张”", prompt) self.assertIn("商品概述、可确认卖点、适用人群与场景、套图画面要求", prompt) self.assertIn("待确认或避免编造的信息", prompt) self.assertIn("可见差异", prompt) self.assertEqual((3, ai.CMHUB_VISION_READ_TIMEOUT_SECONDS), calls[0][2]["timeout"]) self.assert_removed(temp_dir) def test_cmhub_analyze_product_images_requires_vision_alias_and_local_limits(self): with self.make_temp_dir() as temp_dir: cfg, key_path = self._cmhub_config(temp_dir) cfg["ai"]["backend"] = "direct" source = os.path.join(temp_dir, "source.jpg") with open(source, "wb") as fh: fh.write(b"source") cfg["ai"]["cmhub"]["vision_alias"] = "" with self.assertRaises(ai.CMHubError) as raised: ai.analyze_product_images( "提示", "输出语言:繁体中文", [source], config=cfg, cmhub_config_path=key_path, ) self.assertEqual("cmhub_not_configured", raised.exception.code) self.assertIn("图片理解别名", str(raised.exception)) cfg["ai"]["cmhub"]["vision_alias"] = "vision-standard" paths = [] for index in range(ai.CMHUB_VISION_MAX_IMAGES + 1): path = os.path.join(temp_dir, "source-%d.jpg" % index) with open(path, "wb") as fh: fh.write(b"image") paths.append(path) with mock.patch.object(ai._cmhub_session(), "request") as request: with self.assertRaises(ai.AIError) as too_many: ai.analyze_product_images( "提示", "输出语言:繁体中文", paths, config=cfg, cmhub_config_path=key_path, ) self.assertIn("最多支持", str(too_many.exception)) request.assert_not_called() oversized = os.path.join(temp_dir, "oversized.jpg") with open(oversized, "wb") as fh: fh.truncate(ai.CMHUB_VISION_MAX_IMAGE_BYTES + 1) with mock.patch.object(ai._cmhub_session(), "request") as request: with self.assertRaises(ai.AIError) as too_large: ai.analyze_product_images( "提示", "输出语言:繁体中文", [oversized], config=cfg, cmhub_config_path=key_path, ) self.assertIn("超过10MiB", str(too_large.exception)) request.assert_not_called() self.assert_removed(temp_dir) def test_cmhub_analyze_product_images_read_timeout_is_not_retried(self): with self.make_temp_dir() as temp_dir: cfg, key_path = self._cmhub_config(temp_dir) source = os.path.join(temp_dir, "source.jpg") with open(source, "wb") as fh: fh.write(b"source") calls = [] def fake_request(method, url, **kwargs): calls.append((method, url, kwargs)) raise ai.requests.exceptions.ReadTimeout("slow") with mock.patch.object(ai._cmhub_session(), "request", side_effect=fake_request): with self.assertRaises(ai.CMHubError) as raised: ai.analyze_product_images( "提示", "输出语言:繁体中文", [source], config=cfg, cmhub_config_path=key_path, ) self.assertEqual("read_timeout", raised.exception.code) self.assertIn("结果未确认", str(raised.exception)) self.assertEqual(1, len(calls)) self.assertEqual((3, ai.CMHUB_VISION_READ_TIMEOUT_SECONDS), calls[0][2]["timeout"]) self.assert_removed(temp_dir) def test_cmhub_gen_cover_downloads_image_url_safely(self): try: from PIL import Image except ImportError: self.skipTest("Pillow not installed") with self.make_temp_dir() as temp_dir: cfg, key_path = self._cmhub_config(temp_dir) cfg["ai"]["resolution"] = "1k" old_cover = os.path.join(temp_dir, "old.jpg") output = os.path.join(temp_dir, "new.jpg") Image.new("RGB", (16, 16), (20, 30, 40)).save(old_cover, "JPEG") generated = io.BytesIO() Image.new("RGB", (8, 8), (200, 120, 80)).save(generated, "PNG") calls = [] downloads = [] events = [] steps = [] def fake_request(method, url, **kwargs): calls.append((method, url, kwargs)) return _RequestsResponse( { "image_url": "https://cdn.example.com/generated.png", "alias": "image-hd", "model_used": "provider-image-model", "points_cost": 8, "points_balance": 91, "call_id": "call-image-1", } ) def fake_get(url, **kwargs): downloads.append((url, kwargs)) return _RequestsResponse(content=generated.getvalue()) with mock.patch.object(ai._cmhub_session(), "request", side_effect=fake_request), \ mock.patch.object(ai._cmhub_session(), "get", side_effect=fake_get), \ mock.patch( "app.ai.socket.getaddrinfo", return_value=[ ( socket.AF_INET, socket.SOCK_STREAM, 6, "", ("93.184.216.34", 443), ) ], ): result = ai.gen_cover( "生成封面", old_cover, output, resolution="1k", config=cfg, cmhub_config_path=key_path, on_step=steps.append, on_event=events.append, ) self.assertEqual(os.path.abspath(output), result) payload = calls[0][2]["json"] self.assertEqual("image-hd", payload["model"]) self.assertEqual("1K", payload["resolution"]) self.assertEqual("1:1", payload["aspect_ratio"]) self.assertTrue(payload["image_base64"].startswith("data:image/jpeg;base64,")) self.assertEqual((3, ai.CMHUB_IMAGE_READ_TIMEOUT_SECONDS), calls[0][2]["timeout"]) self.assertEqual("https://cdn.example.com/generated.png", downloads[0][0]) self.assertEqual((3, ai.CMHUB_IMAGE_READ_TIMEOUT_SECONDS), downloads[0][1]["timeout"]) self.assertEqual(91, events[0]["metadata"]["points_balance"]) timed_steps = [ event for event in steps if isinstance(event, dict) and event.get("result") == "success" ] timed_step_names = [event.get("step") for event in timed_steps] self.assertIn("cover_request", timed_step_names) self.assertIn("cover_download", timed_step_names) self.assertIn("cover_save", timed_step_names) self.assertTrue( any( event.get("step") == "cover_download" and "下载完成" in event.get("detail", "") and "耗时" in event.get("detail", "") for event in timed_steps ) ) with Image.open(output) as saved: self.assertEqual((1024, 1024), saved.size) self.assert_removed(temp_dir) def test_cmhub_gen_cover_emits_debug_image_url_when_enabled(self): try: from PIL import Image except ImportError: self.skipTest("Pillow not installed") with self.make_temp_dir() as temp_dir: cfg, key_path = self._cmhub_config(temp_dir) old_cover = os.path.join(temp_dir, "old.jpg") output = os.path.join(temp_dir, "new.jpg") Image.new("RGB", (16, 16), (20, 30, 40)).save(old_cover, "JPEG") generated_png = self._png_bytes() steps = [] def fake_request(method, url, **kwargs): return _RequestsResponse( {"image_url": "https://cdn.example.com/generated.png?token=secret"} ) def fake_get(url, **kwargs): return _RequestsResponse(content=generated_png) with mock.patch.dict( os.environ, {"CMSHOPEE_DEBUG_CMHUB_IMAGE_URL": "1"}, ), mock.patch.object(ai._cmhub_session(), "request", side_effect=fake_request), \ mock.patch.object(ai._cmhub_session(), "get", side_effect=fake_get), \ mock.patch( "app.ai.socket.getaddrinfo", return_value=[ ( socket.AF_INET, socket.SOCK_STREAM, 6, "", ("93.184.216.34", 443), ) ], ): ai.gen_cover( "生成封面", old_cover, output, config=cfg, cmhub_config_path=key_path, on_step=steps.append, ) debug_events = [ event for event in steps if isinstance(event, dict) and event.get("step") == "cover_image_url" ] self.assertEqual(1, len(debug_events)) self.assertEqual("debug", debug_events[0]["result"]) self.assertTrue(debug_events[0]["debug_only"]) self.assertIn("https://cdn.example.com/generated.png", debug_events[0]["detail"]) self.assert_removed(temp_dir) def test_cmhub_gen_cover_retries_image_download_without_new_generation(self): try: from PIL import Image except ImportError: self.skipTest("Pillow not installed") with self.make_temp_dir() as temp_dir: cfg, key_path = self._cmhub_config(temp_dir) old_cover = os.path.join(temp_dir, "old.jpg") output = os.path.join(temp_dir, "new.jpg") Image.new("RGB", (16, 16), (20, 30, 40)).save(old_cover, "JPEG") generated_png = self._png_bytes() request_calls = [] download_calls = [] steps = [] def fake_request(method, url, **kwargs): request_calls.append((method, url, kwargs)) return _RequestsResponse( {"image_url": "https://cdn.example.com/generated.png"} ) def fake_get(url, **kwargs): download_calls.append((url, kwargs)) if len(download_calls) == 1: raise ai.requests.exceptions.ConnectionError("temporary") return _RequestsResponse(content=generated_png) with mock.patch.object(ai._cmhub_session(), "request", side_effect=fake_request), \ mock.patch.object(ai._cmhub_session(), "get", side_effect=fake_get), \ mock.patch( "app.ai.socket.getaddrinfo", return_value=[ ( socket.AF_INET, socket.SOCK_STREAM, 6, "", ("93.184.216.34", 443), ) ], ), mock.patch("app.ai.time.sleep"): result = ai.gen_cover( "生成封面", old_cover, output, config=cfg, cmhub_config_path=key_path, on_step=steps.append, ) self.assertEqual(os.path.abspath(output), result) self.assertEqual(1, len(request_calls)) self.assertEqual(2, len(download_calls)) retry_events = [ event for event in steps if isinstance(event, dict) and event.get("step") == "cover_download" and event.get("result") == "retry" ] self.assertEqual(1, len(retry_events)) self.assertEqual(1, retry_events[0]["attempt"]) self.assertEqual(ai.CMHUB_IMAGE_DOWNLOAD_ATTEMPTS, retry_events[0]["attempts"]) self.assert_removed(temp_dir) def test_cmhub_download_slow_warning_event(self): try: from PIL import Image # noqa: F401 except ImportError: self.skipTest("Pillow not installed") generated_png = self._png_bytes() steps = [] def fake_get(url, **kwargs): return _RequestsResponse(content=generated_png) with mock.patch.object(ai._cmhub_session(), "get", side_effect=fake_get), \ mock.patch( "app.ai.socket.getaddrinfo", return_value=[ ( socket.AF_INET, socket.SOCK_STREAM, 6, "", ("93.184.216.34", 443), ) ], ): image_bytes, elapsed = ai._download_cmhub_image_with_retry( "https://cdn.example.com/generated.png", connect_timeout=3, read_timeout=ai.CMHUB_IMAGE_READ_TIMEOUT_SECONDS, on_step=steps.append, slow_threshold=0, ) self.assertEqual(generated_png, image_bytes) self.assertGreaterEqual(elapsed, 0) self.assertTrue( any( isinstance(event, dict) and event.get("step") == "cover_download" and event.get("result") == "warning" and "图片下载较慢" in event.get("detail", "") for event in steps ) ) def test_cmhub_image_url_rejects_private_and_private_dns(self): with self.assertRaises(ai.AIError): ai._download_cmhub_image("http://127.0.0.1/a.png", 1, 1) with mock.patch( "app.ai.socket.getaddrinfo", return_value=[ (socket.AF_INET, socket.SOCK_STREAM, 6, "", ("10.0.0.2", 443)) ], ): with self.assertRaises(ai.AIError): ai._download_cmhub_image("https://cdn.example.com/a.png", 1, 1) def test_cmhub_image_download_uses_curl_without_url_in_argv(self): generated_png = self._png_bytes() url = "https://cdn.example.com/generated.png?token=secret-token" public_dns = [ (socket.AF_INET, socket.SOCK_STREAM, 6, "", ("93.184.216.34", 443)) ] calls = [] class FakeProcess: returncode = 0 def poll(self): return self.returncode def communicate(self): return b"", b"" def fake_popen(args, **kwargs): calls.append((args, kwargs)) self.assertIn("-K", args) config_path = args[args.index("-K") + 1] with open(config_path, "r", encoding="utf-8") as fh: self.assertIn(url, fh.read()) self.assertNotIn(url, args) self.assertIn("--connect-timeout", args) self.assertEqual("3", args[args.index("--connect-timeout") + 1]) self.assertIn("--max-time", args) self.assertEqual( str(ai.CMHUB_IMAGE_READ_TIMEOUT_SECONDS), args[args.index("--max-time") + 1], ) self.assertIn("--max-filesize", args) self.assertEqual(str(ai.CMHUB_IMAGE_MAX_BYTES), args[args.index("--max-filesize") + 1]) self.assertIn("--noproxy", args) self.assertEqual("*", args[args.index("--noproxy") + 1]) self.assertFalse(kwargs["shell"]) self.assertEqual(0x08000000, kwargs["creationflags"]) output_path = args[args.index("--output") + 1] with open(output_path, "wb") as fh: fh.write(generated_png) return FakeProcess() with mock.patch("app.ai.os.name", "nt"), \ mock.patch("app.ai.subprocess.CREATE_NO_WINDOW", 0x08000000, create=True), \ mock.patch("app.ai._find_system_curl", return_value=r"C:\Windows\System32\curl.exe"), \ mock.patch("app.ai.subprocess.Popen", side_effect=fake_popen), \ mock.patch("app.ai.socket.getaddrinfo", return_value=public_dns): image_bytes = ai._download_cmhub_image( url, connect_timeout=3, read_timeout=ai.CMHUB_IMAGE_READ_TIMEOUT_SECONDS, use_system_proxy=False, download_with_curl="true", ) self.assertEqual(generated_png, image_bytes) self.assertEqual(1, len(calls)) def test_cmhub_curl_hidden_window_kwargs_are_windows_only(self): with mock.patch("app.ai.os.name", "nt"), \ mock.patch("app.ai.subprocess.CREATE_NO_WINDOW", 0x08000000, create=True): self.assertEqual( {"creationflags": 0x08000000}, ai._subprocess_hidden_window_kwargs(), ) with mock.patch("app.ai.os.name", "posix"): self.assertEqual({}, ai._subprocess_hidden_window_kwargs()) def test_cmhub_image_download_skips_curl_for_private_url(self): with mock.patch("app.ai._find_system_curl", return_value=r"C:\Windows\System32\curl.exe"), \ mock.patch("app.ai.subprocess.Popen") as popen: with self.assertRaises(ai.AIError): ai._download_cmhub_image( "http://127.0.0.1/a.png", connect_timeout=3, read_timeout=ai.CMHUB_IMAGE_READ_TIMEOUT_SECONDS, download_with_curl="true", ) popen.assert_not_called() def test_cmhub_image_download_falls_back_to_requests_when_curl_fails(self): generated_png = self._png_bytes() public_dns = [ (socket.AF_INET, socket.SOCK_STREAM, 6, "", ("93.184.216.34", 443)) ] def fake_get(url, **kwargs): return _RequestsResponse(content=generated_png) class FailedProcess: returncode = 28 def poll(self): return self.returncode def communicate(self): return b"", b"timeout" with mock.patch("app.ai._find_system_curl", return_value=r"C:\Windows\System32\curl.exe"), \ mock.patch( "app.ai.subprocess.Popen", return_value=FailedProcess(), ) as popen, \ mock.patch.object(ai._cmhub_session(), "get", side_effect=fake_get) as get, \ mock.patch("app.ai.socket.getaddrinfo", return_value=public_dns): image_bytes = ai._download_cmhub_image( "https://cdn.example.com/generated.png", connect_timeout=3, read_timeout=ai.CMHUB_IMAGE_READ_TIMEOUT_SECONDS, download_with_curl="true", ) self.assertEqual(generated_png, image_bytes) self.assertEqual(1, popen.call_count) self.assertEqual(1, get.call_count) def test_cmhub_image_download_auto_without_curl_uses_requests(self): generated_png = self._png_bytes() public_dns = [ (socket.AF_INET, socket.SOCK_STREAM, 6, "", ("93.184.216.34", 443)) ] def fake_get(url, **kwargs): return _RequestsResponse(content=generated_png) with mock.patch("app.ai.os.name", "nt"), \ mock.patch("app.ai._find_system_curl", return_value=""), \ mock.patch("app.ai.subprocess.Popen") as popen, \ mock.patch.object(ai._cmhub_session(), "get", side_effect=fake_get), \ mock.patch("app.ai.socket.getaddrinfo", return_value=public_dns): image_bytes = ai._download_cmhub_image( "https://cdn.example.com/generated.png", connect_timeout=3, read_timeout=ai.CMHUB_IMAGE_READ_TIMEOUT_SECONDS, download_with_curl="auto", ) self.assertEqual(generated_png, image_bytes) popen.assert_not_called() def test_cmhub_image_download_auto_on_non_windows_uses_requests(self): generated_png = self._png_bytes() public_dns = [ (socket.AF_INET, socket.SOCK_STREAM, 6, "", ("93.184.216.34", 443)) ] def fake_get(url, **kwargs): return _RequestsResponse(content=generated_png) with mock.patch("app.ai.os.name", "posix"), \ mock.patch("app.ai._find_system_curl", return_value="/usr/bin/curl"), \ mock.patch("app.ai.subprocess.Popen") as popen, \ mock.patch.object(ai._cmhub_session(), "get", side_effect=fake_get), \ mock.patch("app.ai.socket.getaddrinfo", return_value=public_dns): image_bytes = ai._download_cmhub_image( "https://cdn.example.com/generated.png", connect_timeout=3, read_timeout=ai.CMHUB_IMAGE_READ_TIMEOUT_SECONDS, download_with_curl="auto", ) self.assertEqual(generated_png, image_bytes) popen.assert_not_called() def test_cmhub_requests_download_cancel_closes_response(self): stopped = {"value": False} class StreamingResponse(_RequestsResponse): def iter_content(self, chunk_size=65536): yield b"first" stopped["value"] = True yield b"second" response = StreamingResponse() with mock.patch.object( ai._cmhub_session(), "get", return_value=response, ): with self.assertRaises(CancelledError): ai._download_cmhub_image_with_requests( "https://cdn.example.com/generated.png", connect_timeout=3, read_timeout=30, max_bytes=ai.CMHUB_IMAGE_MAX_BYTES, should_stop=lambda: stopped["value"], ) self.assertTrue(response.closed) def test_cmhub_curl_download_cancel_terminates_process_and_cleans_temp_files(self): stopped = {"value": False} captured_paths = [] class RunningProcess: def __init__(self): self.returncode = None self.terminated = False self.killed = False def poll(self): return self.returncode def terminate(self): self.terminated = True self.returncode = -15 def kill(self): self.killed = True self.returncode = -9 def wait(self, timeout=None): return self.returncode def communicate(self): return b"", b"" process = RunningProcess() def fake_popen(args, **kwargs): captured_paths.extend( [ args[args.index("-K") + 1], args[args.index("--output") + 1], ] ) stopped["value"] = True return process with mock.patch( "app.ai._find_system_curl", return_value=r"C:\Windows\System32\curl.exe", ), mock.patch( "app.ai.subprocess.Popen", side_effect=fake_popen, ): with self.assertRaises(CancelledError): ai._download_cmhub_image_with_curl( "https://cdn.example.com/generated.png", connect_timeout=3, read_timeout=30, max_bytes=ai.CMHUB_IMAGE_MAX_BYTES, should_stop=lambda: stopped["value"], ) self.assertTrue(process.terminated) self.assertFalse(process.killed) self.assertTrue(captured_paths) self.assertTrue(all(not os.path.exists(path) for path in captured_paths)) def test_cmhub_upstream_error_retries_and_keeps_metadata(self): with self.make_temp_dir() as temp_dir: cfg, key_path = self._cmhub_config(temp_dir) calls = [] def fake_request(method, url, **kwargs): calls.append((method, url, kwargs)) if len(calls) == 1: return _RequestsResponse( {"error": {"code": "upstream_error", "message": "bad gateway"}}, status_code=502, ) return _RequestsResponse({"titles": ["新标题"], "points_balance": 10}) with mock.patch.object(ai._cmhub_session(), "request", side_effect=fake_request), \ mock.patch("app.ai.time.sleep"): title = ai.gen_title( "prompt", "old", config=cfg, cmhub_config_path=key_path, ) self.assertEqual("新标题", title) self.assertEqual(2, len(calls)) self.assert_removed(temp_dir) def test_cmhub_image_read_timeout_does_not_retry(self): try: from PIL import Image except ImportError: self.skipTest("Pillow not installed") with self.make_temp_dir() as temp_dir: cfg, key_path = self._cmhub_config(temp_dir) old_cover = os.path.join(temp_dir, "old.jpg") output = os.path.join(temp_dir, "new.jpg") Image.new("RGB", (16, 16), (20, 30, 40)).save(old_cover, "JPEG") calls = [] def fake_request(method, url, **kwargs): calls.append((method, url, kwargs)) raise ai.requests.exceptions.ReadTimeout("slow") with mock.patch.object(ai._cmhub_session(), "request", side_effect=fake_request): with self.assertRaises(ai.CMHubError) as raised: ai.gen_cover( "prompt", old_cover, output, retry=3, config=cfg, cmhub_config_path=key_path, ) self.assertEqual("read_timeout", raised.exception.code) self.assertEqual(1, len(calls)) self.assertEqual((3, ai.CMHUB_IMAGE_READ_TIMEOUT_SECONDS), calls[0][2]["timeout"]) self.assert_removed(temp_dir) def test_fetch_cmhub_models_returns_aliases(self): calls = [] def fake_request(method, url, **kwargs): calls.append((method, url, kwargs)) return _RequestsResponse( { "models": [ { "alias": "title-standard", "operation_type": "title", "requires_image": False, "pricing_status": "priced", "prices": [{"resolution": "512", "points_cost": 1}], } ] } ) with mock.patch.object(ai._cmhub_session(), "request", side_effect=fake_request): models = ai.fetch_cmhub_models("https://cmhub.example.com", "sk-cmhub-secret") self.assertEqual("GET", calls[0][0]) self.assertEqual("https://cmhub.example.com/api/v1/models", calls[0][1]) self.assertEqual("title-standard", models[0]["alias"]) def test_fetch_cmhub_models_404_returns_clear_base_url_error(self): calls = [] def fake_request(method, url, **kwargs): calls.append((method, url, kwargs)) return _RequestsResponse({"detail": "notfound"}, status_code=404) with mock.patch.object(ai._cmhub_session(), "request", side_effect=fake_request): with self.assertRaises(ai.CMHubError) as raised: ai.fetch_cmhub_models( "https://cmhub.example.com/api/v1/", "sk-cmhub-secret", ) self.assertEqual("GET", calls[0][0]) self.assertEqual("https://cmhub.example.com/api/v1/models", calls[0][1]) self.assertEqual("not_found", raised.exception.code) self.assertEqual(404, raised.exception.status) message = str(raised.exception) self.assertIn("cmhub 接口不存在", message) self.assertIn("/api/v1/models", message) self.assertNotIn("notfound", message) def test_fetch_cmhub_balance_returns_points(self): calls = [] def fake_request(method, url, **kwargs): calls.append((method, url, kwargs)) return _RequestsResponse({"user": {"id": "u1"}, "points_balance": 42}) with mock.patch.object(ai._cmhub_session(), "request", side_effect=fake_request): balance = ai.fetch_cmhub_balance("https://cmhub.example.com", "sk-cmhub-secret") self.assertEqual("GET", calls[0][0]) self.assertEqual("https://cmhub.example.com/api/v1/balance", calls[0][1]) self.assertEqual(42, balance["points_balance"]) def test_cmhub_image_concurrency_plan_caps_at_five(self): plan = ai.cmhub_image_concurrency_plan({"image_concurrency": 10}) self.assertEqual(10, plan["configured_image_concurrency"]) self.assertEqual(5, plan["request_concurrency"]) self.assertEqual(5, plan["download_concurrency"]) self.assertEqual(5, plan["limit"]) def test_generate_batch_cmhub_image_downloads_do_not_block_later_requests(self): try: from PIL import Image # noqa: F401 except ImportError: self.skipTest("Pillow not installed") with self.make_temp_dir() as temp_dir: cfg, key_path = self._cmhub_config(temp_dir) cfg["db_path"] = os.path.join(temp_dir, "cmshopee.db") cfg["image_dir"] = os.path.join(temp_dir, "images") cfg["ai"]["generate_mode"] = "cover" cfg["ai"]["generate_mode"] = "cover" cfg["ai"]["generate_cover"] = True cfg["ai"]["image_concurrency"] = 10 titles = ["旧标题%s" % index for index in range(7)] batch_id, tasks = self._collected_tasks(temp_dir, cfg, titles) self._write_old_cover_files(tasks) for index, task in enumerate(tasks): db.set_generated(task.id, "已有标题%s" % index, None, path=cfg["db_path"]) tasks = db.list_tasks(batch_id=batch_id, path=cfg["db_path"]) generated_png = self._png_bytes() public_dns = [ (socket.AF_INET, socket.SOCK_STREAM, 6, "", ("93.184.216.34", 443)) ] lock = threading.Lock() all_requests_seen = threading.Event() counters = {"request_count": 0, "active": 0, "max_active": 0} downloads = [] def fake_request(method, url, **kwargs): if str(method).upper() == "GET": task_id = url.rsplit("/", 1)[-1] index = int(task_id.rsplit("-", 1)[-1]) return _RequestsResponse( { "task_id": task_id, "status": "succeeded", "result": { "image_url": "https://cdn.example.com/generated-%s.png" % index }, } ) with lock: counters["request_count"] += 1 request_index = counters["request_count"] counters["active"] += 1 counters["max_active"] = max( counters["max_active"], counters["active"], ) if counters["request_count"] >= len(tasks): all_requests_seen.set() try: return _RequestsResponse( { "task_id": "cmhub-task-%s" % request_index, "status": "queued", } ) finally: with lock: counters["active"] -= 1 def fake_get(url, **kwargs): self.assertTrue( all_requests_seen.wait(2), "慢下载不能阻塞后续 cmhub 生图请求提交", ) with lock: downloads.append((url, kwargs)) return _RequestsResponse(content=generated_png) with mock.patch.object(ai._cmhub_session(), "request", side_effect=fake_request), \ mock.patch.object(ai._cmhub_session(), "get", side_effect=fake_get), \ mock.patch("app.ai.socket.getaddrinfo", return_value=public_dns): summary = ai.generate_batch( tasks, {"title": "标题提示", "cover": "封面 {新标题}"}, ai_cfg={ "config": cfg, "db_path": cfg["db_path"], "cmhub_config_path": key_path, }, ) self.assertTrue(summary["ok"]) self.assertEqual(len(tasks), summary["cover_done"]) self.assertEqual(len(tasks), summary["generated_done"]) self.assertEqual(len(tasks), counters["request_count"]) self.assertLessEqual(counters["max_active"], 5) self.assertEqual(len(tasks), len(downloads)) updated = db.list_tasks(batch_id=batch_id, path=cfg["db_path"]) self.assertTrue(all(os.path.exists(task.new_cover_path) for task in updated)) self.assert_removed(temp_dir) def test_generate_batch_cmhub_download_failure_does_not_request_image_again(self): try: from PIL import Image # noqa: F401 except ImportError: self.skipTest("Pillow not installed") with self.make_temp_dir() as temp_dir: cfg, key_path = self._cmhub_config(temp_dir) cfg["db_path"] = os.path.join(temp_dir, "cmshopee.db") cfg["image_dir"] = os.path.join(temp_dir, "images") cfg["ai"]["generate_cover"] = True cfg["ai"]["image_concurrency"] = 10 batch_id, tasks = self._collected_tasks(temp_dir, cfg, ["旧标题"]) self._write_old_cover_files(tasks) db.set_generated(tasks[0].id, "已有标题", None, path=cfg["db_path"]) tasks = db.list_tasks(batch_id=batch_id, path=cfg["db_path"]) public_dns = [ (socket.AF_INET, socket.SOCK_STREAM, 6, "", ("93.184.216.34", 443)) ] requests_seen = [] downloads_seen = [] def fake_request(method, url, **kwargs): if str(method).upper() == "GET": return _RequestsResponse( { "task_id": "cmhub-task-1", "status": "succeeded", "result": {"image_url": "https://cdn.example.com/generated.png"}, } ) requests_seen.append((method, url, kwargs)) return _RequestsResponse({"task_id": "cmhub-task-1", "status": "queued"}) def fake_get(url, **kwargs): downloads_seen.append((url, kwargs)) raise ai.requests.exceptions.ConnectionError("download failed") with mock.patch.object(ai._cmhub_session(), "request", side_effect=fake_request), \ mock.patch.object(ai._cmhub_session(), "get", side_effect=fake_get), \ mock.patch("app.ai.socket.getaddrinfo", return_value=public_dns): summary = ai.generate_batch( tasks, {"title": "标题提示", "cover": "封面 {新标题}"}, ai_cfg={ "config": cfg, "db_path": cfg["db_path"], "cmhub_config_path": key_path, }, ) self.assertFalse(summary["ok"]) self.assertEqual(0, summary["cover_done"]) self.assertEqual(1, summary["failed"]) self.assertEqual(1, len(requests_seen)) self.assertEqual(ai.CMHUB_IMAGE_DOWNLOAD_ATTEMPTS, len(downloads_seen)) updated = db.list_tasks(batch_id=batch_id, path=cfg["db_path"])[0] self.assertEqual("failed", updated.status) self.assertIn("下载 cmhub 图片失败", updated.last_error) self.assert_removed(temp_dir) def test_generate_batch_cmhub_async_submit_persists_task_id_before_poll(self): try: from PIL import Image # noqa: F401 except ImportError: self.skipTest("Pillow not installed") with self.make_temp_dir() as temp_dir: cfg, key_path = self._cmhub_config(temp_dir) cfg["db_path"] = os.path.join(temp_dir, "cmshopee.db") cfg["image_dir"] = os.path.join(temp_dir, "images") cfg["ai"]["generate_cover"] = True batch_id, tasks = self._collected_tasks(temp_dir, cfg, ["旧标题"]) self._write_old_cover_files(tasks) db.set_generated(tasks[0].id, "已有标题", None, path=cfg["db_path"]) task = db.list_tasks(batch_id=batch_id, path=cfg["db_path"])[0] generated_png = self._png_bytes() public_dns = [ (socket.AF_INET, socket.SOCK_STREAM, 6, "", ("93.184.216.34", 443)) ] poll_db_values = [] headers_seen = [] submit_timeouts = [] downloads = [] def fake_request(method, url, **kwargs): if str(method).upper() == "POST": headers_seen.append(dict(kwargs["headers"])) submit_timeouts.append(kwargs.get("timeout")) return _RequestsResponse( { "task_id": "cmhub-task-1", "status": "queued", "points_cost": 2, "points_balance": 80, "call_id": "call-image-1", }, status_code=202, ) poll_db_values.append(db.get_task(task.id, path=cfg["db_path"]).image_task_id) return _RequestsResponse( { "task_id": "cmhub-task-1", "status": "succeeded", "result": {"image_url": "/generated/images/generated.png"}, } ) def fake_get(url, **kwargs): downloads.append((url, kwargs)) return _RequestsResponse(content=generated_png) with mock.patch.object(ai._cmhub_session(), "request", side_effect=fake_request), \ mock.patch.object(ai._cmhub_session(), "get", side_effect=fake_get), \ mock.patch("app.ai.socket.getaddrinfo", return_value=public_dns): summary = ai.generate_batch( [task], {"title": "标题提示", "cover": "封面 {新标题}"}, ai_cfg={ "config": cfg, "db_path": cfg["db_path"], "cmhub_config_path": key_path, }, ) self.assertTrue(summary["ok"]) self.assertEqual(["cmhub-task-1"], poll_db_values) self.assertTrue(headers_seen[0]["Idempotency-Key"].startswith("cmshopee-task-")) self.assertTrue(headers_seen[0]["X-Client-Version"]) self.assertEqual([(3, 36)], submit_timeouts) self.assertEqual( "https://cmhub.example.com/generated/images/generated.png", downloads[0][0], ) updated = db.get_task(task.id, path=cfg["db_path"]) self.assertEqual("cmhub-task-1", updated.image_task_id) self.assertTrue(updated.image_task_key) self.assertTrue(os.path.exists(updated.new_cover_path)) self.assert_removed(temp_dir) def test_cmhub_async_image_url_extractor_accepts_image_urls_list(self): data = { "task_id": "cmhub-task-1", "status": "succeeded", "result": { "image_urls": [ "https://cdn.example.com/generated-a.png", "https://cdn.example.com/generated-b.png", ] }, } self.assertEqual( "https://cdn.example.com/generated-a.png", ai._extract_cmhub_image_url(data, "https://cmhub.example.com"), ) def test_generate_batch_cmhub_resumes_existing_image_task_without_submit(self): try: from PIL import Image # noqa: F401 except ImportError: self.skipTest("Pillow not installed") with self.make_temp_dir() as temp_dir: cfg, key_path = self._cmhub_config(temp_dir) cfg["db_path"] = os.path.join(temp_dir, "cmshopee.db") cfg["image_dir"] = os.path.join(temp_dir, "images") cfg["ai"]["generate_cover"] = True batch_id, tasks = self._collected_tasks(temp_dir, cfg, ["旧标题"]) self._write_old_cover_files(tasks) db.set_generated(tasks[0].id, "已有标题", None, path=cfg["db_path"]) key = db.ensure_image_task_key(tasks[0].id, path=cfg["db_path"]) db.set_image_task_submitted(tasks[0].id, "cmhub-task-resume", key, path=cfg["db_path"]) task = db.list_tasks(batch_id=batch_id, path=cfg["db_path"])[0] generated_png = self._png_bytes() public_dns = [ (socket.AF_INET, socket.SOCK_STREAM, 6, "", ("93.184.216.34", 443)) ] posts = [] def fake_request(method, url, **kwargs): if str(method).upper() == "POST": posts.append((method, url, kwargs)) return _RequestsResponse({"task_id": "unexpected", "status": "queued"}) return _RequestsResponse( { "task_id": "cmhub-task-resume", "status": "succeeded", "result": {"image_url": "https://cdn.example.com/resume.png"}, } ) def fake_get(url, **kwargs): return _RequestsResponse(content=generated_png) with mock.patch.object(ai._cmhub_session(), "request", side_effect=fake_request), \ mock.patch.object(ai._cmhub_session(), "get", side_effect=fake_get), \ mock.patch("app.ai.socket.getaddrinfo", return_value=public_dns): summary = ai.generate_batch( [task], {"title": "标题提示", "cover": "封面 {新标题}"}, ai_cfg={ "config": cfg, "db_path": cfg["db_path"], "cmhub_config_path": key_path, }, ) self.assertTrue(summary["ok"]) self.assertEqual([], posts) updated = db.get_task(task.id, path=cfg["db_path"]) self.assertEqual("cmhub-task-resume", updated.image_task_id) self.assertEqual(key, updated.image_task_key) self.assertTrue(os.path.exists(updated.new_cover_path)) self.assert_removed(temp_dir) def test_generate_batch_cmhub_failed_task_clears_image_task_state(self): try: from PIL import Image # noqa: F401 except ImportError: self.skipTest("Pillow not installed") with self.make_temp_dir() as temp_dir: cfg, key_path = self._cmhub_config(temp_dir) cfg["db_path"] = os.path.join(temp_dir, "cmshopee.db") cfg["image_dir"] = os.path.join(temp_dir, "images") cfg["ai"]["generate_cover"] = True batch_id, tasks = self._collected_tasks(temp_dir, cfg, ["旧标题"]) self._write_old_cover_files(tasks) db.set_generated(tasks[0].id, "已有标题", None, path=cfg["db_path"]) key = db.ensure_image_task_key(tasks[0].id, path=cfg["db_path"]) db.set_image_task_submitted(tasks[0].id, "cmhub-task-failed", key, path=cfg["db_path"]) task = db.list_tasks(batch_id=batch_id, path=cfg["db_path"])[0] def fake_request(method, url, **kwargs): return _RequestsResponse( { "task_id": "cmhub-task-failed", "status": "failed", "error": {"code": "upstream_timeout", "message": "上游超时"}, } ) with mock.patch.object(ai._cmhub_session(), "request", side_effect=fake_request): summary = ai.generate_batch( [task], {"title": "标题提示", "cover": "封面 {新标题}"}, ai_cfg={ "config": cfg, "db_path": cfg["db_path"], "cmhub_config_path": key_path, }, ) self.assertFalse(summary["ok"]) self.assertEqual(1, summary["failed"]) updated = db.get_task(task.id, path=cfg["db_path"]) self.assertIsNone(updated.image_task_id) self.assertIsNone(updated.image_task_key) self.assertIn("上游生成超时", updated.last_error) self.assert_removed(temp_dir) def test_generate_batch_cmhub_cancelled_poll_keeps_image_task_state(self): try: from PIL import Image # noqa: F401 except ImportError: self.skipTest("Pillow not installed") with self.make_temp_dir() as temp_dir: cfg, key_path = self._cmhub_config(temp_dir) cfg["db_path"] = os.path.join(temp_dir, "cmshopee.db") cfg["image_dir"] = os.path.join(temp_dir, "images") cfg["ai"]["generate_cover"] = True batch_id, tasks = self._collected_tasks(temp_dir, cfg, ["旧标题"]) self._write_old_cover_files(tasks) db.set_generated(tasks[0].id, "已有标题", None, path=cfg["db_path"]) task = db.list_tasks(batch_id=batch_id, path=cfg["db_path"])[0] events = [] def fake_request(method, url, **kwargs): if str(method).upper() == "POST": return _RequestsResponse({"task_id": "cmhub-task-cancel", "status": "queued"}) return _RequestsResponse({"task_id": "cmhub-task-cancel", "status": "running"}) def should_stop(): current = db.get_task(task.id, path=cfg["db_path"]) return bool(current and current.image_task_id) with mock.patch.object(ai._cmhub_session(), "request", side_effect=fake_request): summary = ai.generate_batch( [task], {"title": "标题提示", "cover": "封面 {新标题}"}, ai_cfg={ "config": cfg, "db_path": cfg["db_path"], "cmhub_config_path": key_path, "on_event": events.append, }, should_stop=should_stop, ) self.assertFalse(summary["ok"]) self.assertTrue(summary["cancelled"]) updated = db.get_task(task.id, path=cfg["db_path"]) self.assertEqual("cmhub-task-cancel", updated.image_task_id) self.assertTrue(updated.image_task_key) self.assertIsNone(updated.new_cover_path) self.assertTrue( any("服务端任务可能仍在完成" in event.get("detail", "") for event in events) ) self.assert_removed(temp_dir) def test_generate_batch_forwards_cmhub_metadata_event(self): with self.make_temp_dir() as temp_dir: cfg, key_path = self._cmhub_config(temp_dir) cfg["db_path"] = os.path.join(temp_dir, "cmshopee.db") cfg["image_dir"] = os.path.join(temp_dir, "images") batch_id, tasks = self._collected_tasks(temp_dir, cfg, ["旧标题A"]) events = [] def fake_request(method, url, **kwargs): return _RequestsResponse( { "titles": ["新标题A"], "alias": "title-standard", "points_cost": 1, "points_balance": 88, "call_id": "call-batch-1", } ) with mock.patch.object(ai._cmhub_session(), "request", side_effect=fake_request): summary = ai.generate_batch( tasks, {"title": "标题提示", "cover": "封面"}, ai_cfg={ "config": cfg, "db_path": cfg["db_path"], "cmhub_config_path": key_path, "on_event": events.append, }, ) self.assertTrue(summary["ok"]) self.assertTrue( any( event.get("metadata", {}).get("points_balance") == 88 and event.get("metadata", {}).get("call_id") == "call-batch-1" for event in events ) ) updated = db.list_tasks(batch_id=batch_id, path=cfg["db_path"]) self.assertEqual("generated", updated[0].stage) self.assert_removed(temp_dir) def test_generate_batch_persists_titles_and_covers_per_task(self): with self.make_temp_dir() as temp_dir: cfg = self._config() cfg["db_path"] = os.path.join(temp_dir, "cmshopee.db") cfg["image_dir"] = os.path.join(temp_dir, "images") cfg["ai"]["title_concurrency"] = 2 cfg["ai"]["image_concurrency"] = 2 cfg["ai"]["generate_cover"] = True batch_id, tasks = self._collected_tasks(temp_dir, cfg) cover_prompts = [] progress = [] def fake_title(title_prompt, old_title, **kwargs): self.assertEqual("标题提示", title_prompt) return "新" + old_title def fake_cover(cover_prompt, old_cover_path, out_path, **kwargs): cover_prompts.append(cover_prompt) os.makedirs(os.path.dirname(out_path), exist_ok=True) with open(out_path, "wb") as fh: fh.write(b"jpeg") return out_path with mock.patch("app.ai.gen_title", side_effect=fake_title), \ mock.patch("app.ai.gen_cover", side_effect=fake_cover): summary = ai.generate_batch( tasks, { "title": "标题提示", "cover": "封面 {新标题} {店铺} {商品id}", }, ai_cfg={ "config": cfg, "db_path": cfg["db_path"], "account_by_alias": { "alias-a": SimpleNamespace(account_name="主店", slug="main") }, }, on_progress=progress.append, ) self.assertTrue(summary["ok"]) self.assertEqual(2, summary["title_done"]) self.assertEqual(2, summary["cover_done"]) self.assertEqual(0, summary["failed"]) self.assertTrue(progress) updated = db.list_tasks(batch_id=batch_id, path=cfg["db_path"]) self.assertTrue(all(task.stage == "generated" for task in updated)) self.assertEqual({"新旧标题A", "新旧标题B"}, {task.new_title for task in updated}) expected_cover_paths = { os.path.abspath( os.path.join( cfg["image_dir"], str(task.batch_id), "main", f"{task.id}_{task.item_id}_new.jpg", ) ) for task in updated } self.assertEqual(expected_cover_paths, {task.new_cover_path for task in updated}) self.assertTrue(all(os.path.exists(task.new_cover_path) for task in updated)) self.assertIn("主店", "\n".join(cover_prompts)) self.assertIn("新旧标题A", "\n".join(cover_prompts)) self.assert_removed(temp_dir) def test_generate_batch_can_skip_cover_generation(self): with self.make_temp_dir() as temp_dir: cfg = self._config() cfg["db_path"] = os.path.join(temp_dir, "cmshopee.db") cfg["image_dir"] = os.path.join(temp_dir, "images") batch_id, tasks = self._collected_tasks(temp_dir, cfg) events = [] progress = [] def fake_title(title_prompt, old_title, **kwargs): self.assertEqual("标题提示", title_prompt) return "新" + old_title with mock.patch("app.ai.gen_title", side_effect=fake_title), \ mock.patch("app.ai.gen_cover") as gen_cover: summary = ai.generate_batch( tasks, {"title": "标题提示", "cover": "封面 {新标题}"}, ai_cfg={ "config": cfg, "db_path": cfg["db_path"], "on_event": events.append, }, on_progress=progress.append, ) self.assertTrue(summary["ok"]) self.assertFalse(summary["generate_cover"]) self.assertEqual(2, summary["title_done"]) self.assertEqual(0, summary["cover_done"]) self.assertEqual(0, summary["cover_total"]) self.assertEqual(2, summary["generated_done"]) self.assertEqual(0, summary["failed"]) gen_cover.assert_not_called() updated = db.list_tasks(batch_id=batch_id, path=cfg["db_path"]) self.assertTrue(all(task.stage == "generated" for task in updated)) self.assertTrue(all(task.status == "success" for task in updated)) self.assertEqual({"新旧标题A", "新旧标题B"}, {task.new_title for task in updated}) self.assertTrue(all(task.new_cover_path is None for task in updated)) self.assertEqual(0, progress[-1]["cover_total"]) self.assertEqual(2, progress[-1]["generated_done"]) self.assertTrue( any( event.get("phase") == "title" and event.get("step") == "db_write" and event.get("result") == "success" and event.get("detail") == "仅生成标题" for event in events ) ) self.assert_removed(temp_dir) def test_generate_batch_marks_failed_task_without_blocking_others(self): with self.make_temp_dir() as temp_dir: cfg = self._config() cfg["db_path"] = os.path.join(temp_dir, "cmshopee.db") cfg["image_dir"] = os.path.join(temp_dir, "images") cfg["ai"]["generate_cover"] = True batch_id, tasks = self._collected_tasks(temp_dir, cfg, ["好标题", "坏标题"]) def fake_title(title_prompt, old_title, **kwargs): if old_title == "坏标题": raise ai.AIError("标题生成失败") return "新" + old_title def fake_cover(cover_prompt, old_cover_path, out_path, **kwargs): os.makedirs(os.path.dirname(out_path), exist_ok=True) with open(out_path, "wb") as fh: fh.write(b"jpeg") return out_path with mock.patch("app.ai.gen_title", side_effect=fake_title), \ mock.patch("app.ai.gen_cover", side_effect=fake_cover): summary = ai.generate_batch( tasks, {"title": "标题提示", "cover": "封面 {新标题}"}, ai_cfg={"config": cfg, "db_path": cfg["db_path"]}, ) self.assertFalse(summary["ok"]) self.assertEqual(1, summary["title_done"]) self.assertEqual(1, summary["cover_done"]) self.assertEqual(1, summary["failed"]) updated = db.list_tasks(batch_id=batch_id, path=cfg["db_path"]) by_title = {task.old_title: task for task in updated} self.assertEqual("generated", by_title["好标题"].stage) self.assertEqual("success", by_title["好标题"].status) self.assertEqual("collected", by_title["坏标题"].stage) self.assertEqual("failed", by_title["坏标题"].status) self.assertIn("标题生成失败", by_title["坏标题"].last_error) self.assert_removed(temp_dir) def test_generate_batch_fills_missing_cover_without_regenerating_title(self): with self.make_temp_dir() as temp_dir: cfg = self._config() cfg["db_path"] = os.path.join(temp_dir, "cmshopee.db") cfg["image_dir"] = os.path.join(temp_dir, "images") cfg["ai"]["generate_mode"] = "cover" cfg["ai"]["generate_cover"] = True batch_id, tasks = self._collected_tasks(temp_dir, cfg, ["旧标题"]) db.set_generated(tasks[0].id, "手动标题", None, path=cfg["db_path"]) cover_only_task = db.list_tasks(batch_id=batch_id, path=cfg["db_path"])[0] progress = [] def fake_cover(cover_prompt, old_cover_path, out_path, **kwargs): os.makedirs(os.path.dirname(out_path), exist_ok=True) with open(out_path, "wb") as fh: fh.write(b"jpeg") return out_path with mock.patch("app.ai.gen_title") as gen_title, \ mock.patch("app.ai.gen_cover", side_effect=fake_cover) as gen_cover: summary = ai.generate_batch( [cover_only_task], {"title": "标题提示", "cover": "封面 {新标题}"}, ai_cfg={"config": cfg, "db_path": cfg["db_path"]}, on_progress=progress.append, ) self.assertTrue(summary["ok"]) self.assertEqual("cover", summary["generate_mode"]) self.assertTrue(summary["generate_cover"]) self.assertEqual(1, summary["total"]) self.assertEqual(0, summary["title_total"]) self.assertEqual(0, summary["title_done"]) self.assertEqual(1, summary["cover_total"]) self.assertEqual(1, summary["cover_done"]) self.assertEqual(1, summary["generated_done"]) self.assertEqual(0, progress[-1]["title_total"]) self.assertEqual(1, progress[-1]["cover_total"]) gen_title.assert_not_called() gen_cover.assert_called_once() updated = db.list_tasks(batch_id=batch_id, path=cfg["db_path"])[0] self.assertEqual("generated", updated.stage) self.assertEqual("success", updated.status) self.assertEqual("手动标题", updated.new_title) self.assertTrue(os.path.exists(updated.new_cover_path)) self.assert_removed(temp_dir) def test_generate_batch_cover_only_uses_old_title_and_keeps_new_title_empty(self): with self.make_temp_dir() as temp_dir: cfg = self._config() cfg["db_path"] = os.path.join(temp_dir, "cmshopee.db") cfg["image_dir"] = os.path.join(temp_dir, "images") cfg["ai"]["generate_mode"] = "cover" cfg["ai"]["generate_cover"] = True batch_id, tasks = self._collected_tasks(temp_dir, cfg, ["旧标题"]) cover_prompts = [] events = [] def fake_cover(cover_prompt, old_cover_path, out_path, **kwargs): cover_prompts.append(cover_prompt) os.makedirs(os.path.dirname(out_path), exist_ok=True) with open(out_path, "wb") as fh: fh.write(b"jpeg") return out_path with mock.patch("app.ai.gen_title") as gen_title, \ mock.patch("app.ai.gen_cover", side_effect=fake_cover) as gen_cover: summary = ai.generate_batch( tasks, {"title": "标题提示", "cover": "封面 {新标题} / {旧标题}"}, ai_cfg={ "config": cfg, "db_path": cfg["db_path"], "on_event": events.append, }, ) self.assertTrue(summary["ok"]) self.assertEqual(0, summary["title_total"]) self.assertEqual(1, summary["cover_total"]) gen_title.assert_not_called() gen_cover.assert_called_once() self.assertEqual(["封面 旧标题 / 旧标题"], cover_prompts) updated = db.get_task(tasks[0].id, path=cfg["db_path"]) self.assertIsNone(updated.new_title) self.assertTrue(os.path.exists(updated.new_cover_path)) self.assertTrue( any("使用旧标题作为封面参考" in str(event.get("detail") or "") for event in events) ) cfg["ai"]["generate_mode"] = "title" cfg["ai"]["generate_cover"] = False with mock.patch("app.ai.gen_title", return_value="后补新标题") as gen_title, \ mock.patch("app.ai.gen_cover") as gen_cover: title_summary = ai.generate_batch( [updated], {"title": "标题提示", "cover": "封面"}, ai_cfg={"config": cfg, "db_path": cfg["db_path"]}, ) self.assertTrue(title_summary["ok"]) gen_title.assert_called_once() gen_cover.assert_not_called() completed = db.get_task(tasks[0].id, path=cfg["db_path"]) self.assertEqual("后补新标题", completed.new_title) self.assertEqual(updated.new_cover_path, completed.new_cover_path) self.assert_removed(temp_dir) def test_generate_batch_fills_reset_cover_after_committed_history(self): with self.make_temp_dir() as temp_dir: cfg = self._config() cfg["db_path"] = os.path.join(temp_dir, "cmshopee.db") cfg["image_dir"] = os.path.join(temp_dir, "images") cfg["ai"]["generate_cover"] = True batch_id, tasks = self._collected_tasks(temp_dir, cfg, ["旧标题"]) db.set_generated(tasks[0].id, "手动标题", "old-new.jpg", path=cfg["db_path"]) db.set_applied(tasks[0].id, True, path=cfg["db_path"]) db.reset_generated( tasks[0].id, reset_title=False, reset_cover=True, path=cfg["db_path"], ) cover_only_task = db.list_tasks(batch_id=batch_id, path=cfg["db_path"])[0] def fake_cover(cover_prompt, old_cover_path, out_path, **kwargs): os.makedirs(os.path.dirname(out_path), exist_ok=True) with open(out_path, "wb") as fh: fh.write(b"jpeg") return out_path with mock.patch("app.ai.gen_title") as gen_title, \ mock.patch("app.ai.gen_cover", side_effect=fake_cover) as gen_cover: summary = ai.generate_batch( [cover_only_task], {"title": "标题提示", "cover": "封面 {新标题}"}, ai_cfg={"config": cfg, "db_path": cfg["db_path"]}, ) self.assertTrue(summary["ok"]) self.assertEqual(1, summary["total"]) self.assertEqual(0, summary["title_total"]) self.assertEqual(1, summary["cover_total"]) gen_title.assert_not_called() gen_cover.assert_called_once() updated = db.list_tasks(batch_id=batch_id, path=cfg["db_path"])[0] self.assertEqual("手动标题", updated.new_title) self.assertTrue(os.path.exists(updated.new_cover_path)) self.assertEqual(1, updated.committed) self.assertEqual(1, updated.apply_attempts) self.assert_removed(temp_dir) def test_generate_batch_handles_mixed_title_and_cover_gaps(self): with self.make_temp_dir() as temp_dir: cfg = self._config() cfg["db_path"] = os.path.join(temp_dir, "cmshopee.db") cfg["image_dir"] = os.path.join(temp_dir, "images") cfg["ai"]["generate_cover"] = True batch_id, tasks = self._collected_tasks(temp_dir, cfg, ["旧标题A", "旧标题B"]) db.set_generated(tasks[1].id, "已有标题B", None, path=cfg["db_path"]) mixed_tasks = db.list_tasks(batch_id=batch_id, path=cfg["db_path"]) cover_prompts = [] def fake_title(title_prompt, old_title, **kwargs): return "新" + old_title def fake_cover(cover_prompt, old_cover_path, out_path, **kwargs): cover_prompts.append(cover_prompt) os.makedirs(os.path.dirname(out_path), exist_ok=True) with open(out_path, "wb") as fh: fh.write(b"jpeg") return out_path with mock.patch("app.ai.gen_title", side_effect=fake_title) as gen_title, \ mock.patch("app.ai.gen_cover", side_effect=fake_cover) as gen_cover: summary = ai.generate_batch( mixed_tasks, {"title": "标题提示", "cover": "封面 {新标题}"}, ai_cfg={"config": cfg, "db_path": cfg["db_path"]}, ) self.assertTrue(summary["ok"]) self.assertEqual(2, summary["total"]) self.assertEqual(1, summary["title_total"]) self.assertEqual(1, summary["title_done"]) self.assertEqual(2, summary["cover_total"]) self.assertEqual(2, summary["cover_done"]) self.assertEqual(2, summary["generated_done"]) self.assertEqual(1, gen_title.call_count) self.assertEqual(2, gen_cover.call_count) updated = db.list_tasks(batch_id=batch_id, path=cfg["db_path"]) by_old_title = {task.old_title: task for task in updated} self.assertEqual("新旧标题A", by_old_title["旧标题A"].new_title) self.assertEqual("已有标题B", by_old_title["旧标题B"].new_title) self.assertTrue(all(os.path.exists(task.new_cover_path) for task in updated)) self.assertIn("新旧标题A", "\n".join(cover_prompts)) self.assertIn("已有标题B", "\n".join(cover_prompts)) self.assert_removed(temp_dir) def test_generate_batch_ignores_title_only_task_when_cover_disabled(self): with self.make_temp_dir() as temp_dir: cfg = self._config() cfg["db_path"] = os.path.join(temp_dir, "cmshopee.db") cfg["image_dir"] = os.path.join(temp_dir, "images") cfg["ai"]["generate_cover"] = False batch_id, tasks = self._collected_tasks(temp_dir, cfg, ["旧标题"]) db.set_generated(tasks[0].id, "已有标题", None, path=cfg["db_path"]) title_only_task = db.list_tasks(batch_id=batch_id, path=cfg["db_path"])[0] with mock.patch("app.ai.gen_title") as gen_title, \ mock.patch("app.ai.gen_cover") as gen_cover: summary = ai.generate_batch( [title_only_task], {"title": "标题提示", "cover": "封面"}, ai_cfg={"config": cfg, "db_path": cfg["db_path"]}, ) self.assertTrue(summary["ok"]) self.assertEqual(0, summary["total"]) self.assertEqual(0, summary["title_total"]) self.assertEqual(0, summary["cover_total"]) gen_title.assert_not_called() gen_cover.assert_not_called() self.assert_removed(temp_dir) def test_generate_batch_ignores_complete_generated_task(self): with self.make_temp_dir() as temp_dir: cfg = self._config() cfg["db_path"] = os.path.join(temp_dir, "cmshopee.db") cfg["image_dir"] = os.path.join(temp_dir, "images") cfg["ai"]["generate_cover"] = True batch_id, tasks = self._collected_tasks(temp_dir, cfg, ["旧标题"]) db.set_generated(tasks[0].id, "已有标题", "new.jpg", path=cfg["db_path"]) complete_task = db.list_tasks(batch_id=batch_id, path=cfg["db_path"])[0] with mock.patch("app.ai.gen_title") as gen_title, \ mock.patch("app.ai.gen_cover") as gen_cover: summary = ai.generate_batch( [complete_task], {"title": "标题提示", "cover": "封面"}, ai_cfg={"config": cfg, "db_path": cfg["db_path"]}, ) self.assertTrue(summary["ok"]) self.assertEqual(0, summary["total"]) self.assertEqual(0, summary["title_total"]) self.assertEqual(0, summary["cover_total"]) gen_title.assert_not_called() gen_cover.assert_not_called() self.assert_removed(temp_dir) def test_generate_batch_does_not_retry_apply_failed_records(self): with self.make_temp_dir() as temp_dir: cfg = self._config() cfg["db_path"] = os.path.join(temp_dir, "cmshopee.db") cfg["image_dir"] = os.path.join(temp_dir, "images") batch_id, tasks = self._collected_tasks(temp_dir, cfg, ["旧标题"]) db.set_generated(tasks[0].id, "新标题", "new.jpg", path=cfg["db_path"]) db.mark_failed(tasks[0].id, "apply", "更新失败", path=cfg["db_path"]) update_failed_task = db.list_tasks(batch_id=batch_id, path=cfg["db_path"])[0] with mock.patch("app.ai.gen_title") as gen_title, \ mock.patch("app.ai.gen_cover") as gen_cover: summary = ai.generate_batch( [update_failed_task], {"title": "标题提示", "cover": "封面"}, ai_cfg={"config": cfg, "db_path": cfg["db_path"]}, ) self.assertTrue(summary["ok"]) self.assertEqual(0, summary["total"]) gen_title.assert_not_called() gen_cover.assert_not_called() unchanged = db.list_tasks(batch_id=batch_id, path=cfg["db_path"])[0] self.assertEqual("generated", unchanged.stage) self.assertEqual("failed", unchanged.status) self.assertEqual(1, unchanged.apply_attempts) self.assert_removed(temp_dir) def test_generate_batch_stop_before_scheduling_keeps_tasks_collected(self): with self.make_temp_dir() as temp_dir: cfg = self._config() cfg["db_path"] = os.path.join(temp_dir, "cmshopee.db") cfg["image_dir"] = os.path.join(temp_dir, "images") batch_id, tasks = self._collected_tasks(temp_dir, cfg) with mock.patch("app.ai.gen_title") as gen_title, \ mock.patch("app.ai.gen_cover") as gen_cover: summary = ai.generate_batch( tasks, {"title": "标题提示", "cover": "封面"}, ai_cfg={"config": cfg, "db_path": cfg["db_path"]}, should_stop=lambda: True, ) self.assertFalse(summary["ok"]) self.assertTrue(summary["cancelled"]) self.assertEqual(0, summary["title_done"]) self.assertEqual(0, summary["cover_done"]) self.assertEqual(0, summary["failed"]) gen_title.assert_not_called() gen_cover.assert_not_called() updated = db.list_tasks(batch_id=batch_id, path=cfg["db_path"]) self.assertTrue(all(task.stage == "collected" for task in updated)) self.assert_removed(temp_dir) def test_cmhub_session_is_shared_singleton_with_pool(self): session_a = ai._cmhub_session() session_b = ai._cmhub_session() self.assertIs(session_a, session_b) # 连接池要覆盖生图并发 + 下载并发(各上限 5) self.assertGreaterEqual( ai.CMHUB_HTTP_POOL_SIZE, ai.CMHUB_IMAGE_CONCURRENCY_LIMIT * 2 ) adapter = session_a.get_adapter("https://cmhub.example.com") self.assertEqual(ai.CMHUB_HTTP_POOL_SIZE, adapter._pool_maxsize) def test_apply_cmhub_proxy_toggles_trust_env(self): try: session = ai._apply_cmhub_proxy(False) self.assertIs(session, ai._cmhub_session()) self.assertFalse(session.trust_env) ai._apply_cmhub_proxy(True) self.assertTrue(ai._cmhub_session().trust_env) finally: # 复位为默认(绕过系统代理),避免影响其它测试 ai._apply_cmhub_proxy(False) def test_cmhub_runtime_applies_and_returns_use_system_proxy(self): with mock.patch( "app.ai.appconfig.cmhub_config", return_value={ "base_url": "https://cmhub.example.com", "title_alias": "title-standard", "image_alias": "image-hd", "connect_timeout": 10, "use_system_proxy": True, }, ), mock.patch( "app.ai.appconfig.get_cmhub_api_key", return_value="sk_cmhub_test" ): try: runtime = ai._cmhub_runtime({}, "image", cmhub_config_path=None) self.assertTrue(runtime["use_system_proxy"]) self.assertTrue(ai._cmhub_session().trust_env) finally: ai._apply_cmhub_proxy(False) if __name__ == "__main__": unittest.main()