Files

2324 lines
98 KiB
Python

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()