2322 lines
98 KiB
Python
2322 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)
|
|
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)
|
|
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()
|