feat: cap cmhub image concurrency
This commit is contained in:
@@ -4,6 +4,7 @@ import json
|
||||
import os
|
||||
import socket
|
||||
import sys
|
||||
import threading
|
||||
import unittest
|
||||
from types import SimpleNamespace
|
||||
from unittest import mock
|
||||
@@ -127,6 +128,24 @@ class AITests(TempDirMixin, unittest.TestCase):
|
||||
)
|
||||
return batch_id, db.list_tasks(batch_id=batch_id, path=cfg["db_path"])
|
||||
|
||||
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")
|
||||
@@ -360,6 +379,7 @@ class AITests(TempDirMixin, unittest.TestCase):
|
||||
calls = []
|
||||
downloads = []
|
||||
events = []
|
||||
steps = []
|
||||
|
||||
def fake_request(method, url, **kwargs):
|
||||
calls.append((method, url, kwargs))
|
||||
@@ -399,6 +419,7 @@ class AITests(TempDirMixin, unittest.TestCase):
|
||||
resolution="1k",
|
||||
config=cfg,
|
||||
cmhub_config_path=key_path,
|
||||
on_step=steps.append,
|
||||
on_event=events.append,
|
||||
)
|
||||
|
||||
@@ -412,6 +433,22 @@ class AITests(TempDirMixin, unittest.TestCase):
|
||||
self.assertEqual("https://cdn.example.com/generated.png", downloads[0][0])
|
||||
self.assertEqual((3, 650), 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)
|
||||
|
||||
@@ -554,6 +591,152 @@ class AITests(TempDirMixin, unittest.TestCase):
|
||||
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_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):
|
||||
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(
|
||||
{
|
||||
"image_url": "https://cdn.example.com/generated-%s.png"
|
||||
% request_index
|
||||
}
|
||||
)
|
||||
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("app.ai.requests.request", side_effect=fake_request), \
|
||||
mock.patch("app.ai.requests.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):
|
||||
requests_seen.append((method, url, kwargs))
|
||||
return _RequestsResponse(
|
||||
{"image_url": "https://cdn.example.com/generated.png"}
|
||||
)
|
||||
|
||||
def fake_get(url, **kwargs):
|
||||
downloads_seen.append((url, kwargs))
|
||||
raise ai.requests.exceptions.ConnectionError("download failed")
|
||||
|
||||
with mock.patch("app.ai.requests.request", side_effect=fake_request), \
|
||||
mock.patch("app.ai.requests.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(1, 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_forwards_cmhub_metadata_event(self):
|
||||
with self.make_temp_dir() as temp_dir:
|
||||
cfg, key_path = self._cmhub_config(temp_dir)
|
||||
|
||||
@@ -1936,6 +1936,7 @@ class GuiTests(TempDirMixin, unittest.TestCase):
|
||||
with self.make_temp_dir() as temp_dir:
|
||||
cfg = self.make_config(temp_dir)
|
||||
cfg["ai"] = appconfig.ai_config(cfg)
|
||||
cfg["ai"]["backend"] = "direct"
|
||||
cfg["ai"]["generate_cover"] = True
|
||||
account = accounts.create_account("主店", "alias-a", debug_port=9222, config=cfg)
|
||||
batch_id = db.create_batch(["input.xlsx"], path=cfg["db_path"])
|
||||
@@ -1989,6 +1990,33 @@ class GuiTests(TempDirMixin, unittest.TestCase):
|
||||
"level": "warning",
|
||||
}
|
||||
)
|
||||
ai_cfg["on_event"](
|
||||
{
|
||||
"task": tasks[0],
|
||||
"phase": "cover",
|
||||
"step": "cover_request",
|
||||
"result": "success",
|
||||
"detail": "cmhub 已返回 image_url,耗时 91.2秒",
|
||||
}
|
||||
)
|
||||
ai_cfg["on_event"](
|
||||
{
|
||||
"task": tasks[0],
|
||||
"phase": "cover",
|
||||
"step": "cover_download",
|
||||
"result": "success",
|
||||
"detail": "下载完成,1.3MB,耗时 12.4秒",
|
||||
}
|
||||
)
|
||||
ai_cfg["on_event"](
|
||||
{
|
||||
"task": tasks[0],
|
||||
"phase": "cover",
|
||||
"step": "cover_save",
|
||||
"result": "success",
|
||||
"detail": "JPEG 已保存,耗时 1.1秒,文件 220.0KB",
|
||||
}
|
||||
)
|
||||
ai_cfg["on_task_update"](tasks[0].id, {"stage": "generated"})
|
||||
return {"ok": True, "total": 1, "title_done": 1, "cover_done": 1, "failed": 0}
|
||||
|
||||
@@ -2016,6 +2044,9 @@ class GuiTests(TempDirMixin, unittest.TestCase):
|
||||
self.assertIn("[开始] 本轮生成 1 条", joined_logs)
|
||||
self.assertIn("[图片] 1/1 商品 51100639510", joined_logs)
|
||||
self.assertIn("准备重试 1/2", joined_logs)
|
||||
self.assertIn("cmhub 已返回 image_url,耗时 91.2秒", joined_logs)
|
||||
self.assertIn("下载完成,1.3MB,耗时 12.4秒", joined_logs)
|
||||
self.assertIn("本地保存完成,JPEG 已保存,耗时 1.1秒,文件 220.0KB", joined_logs)
|
||||
self.assertIn("token=***", joined_logs)
|
||||
self.assertNotIn("SECRET-TOKEN", joined_logs)
|
||||
self.assertIn("[完成] AI 生成完成:标题1/1,图片1/1,失败0", joined_logs)
|
||||
@@ -2023,6 +2054,7 @@ class GuiTests(TempDirMixin, unittest.TestCase):
|
||||
event_messages = "\n".join(event.message for event in events)
|
||||
self.assertIn("[图片] 1/1 商品 51100639510", event_messages)
|
||||
self.assertIn("准备重试 1/2", event_messages)
|
||||
self.assertIn("下载完成,1.3MB,耗时 12.4秒", event_messages)
|
||||
|
||||
self.assert_removed(temp_dir)
|
||||
|
||||
@@ -2032,6 +2064,7 @@ class GuiTests(TempDirMixin, unittest.TestCase):
|
||||
cfg["ai"] = appconfig.ai_config(cfg)
|
||||
cfg["ai"]["backend"] = "cmhub"
|
||||
cfg["ai"]["generate_cover"] = True
|
||||
cfg["ai"]["image_concurrency"] = 10
|
||||
accounts.create_account("主店", "alias-a", debug_port=9222, config=cfg)
|
||||
batch_id = db.create_batch(["input.xlsx"], path=cfg["db_path"])
|
||||
db.insert_tasks(
|
||||
@@ -2096,6 +2129,9 @@ class GuiTests(TempDirMixin, unittest.TestCase):
|
||||
self.assertEqual(88, summary["points_balance"])
|
||||
self.assertEqual(88, progress[-1]["points_balance"])
|
||||
joined_logs = "\n".join(logs)
|
||||
self.assertIn("图片并发10", joined_logs)
|
||||
self.assertIn("cmhub实际生图并发5", joined_logs)
|
||||
self.assertIn("下载并发5", joined_logs)
|
||||
self.assertIn("[计费] 商品 51100639510", joined_logs)
|
||||
self.assertIn("别名 title-standard", joined_logs)
|
||||
self.assertIn("扣点 1", joined_logs)
|
||||
@@ -2193,6 +2229,7 @@ class GuiTests(TempDirMixin, unittest.TestCase):
|
||||
with self.make_temp_dir() as temp_dir:
|
||||
cfg = self.make_config(temp_dir)
|
||||
cfg["ai"] = appconfig.ai_config(cfg)
|
||||
cfg["ai"]["backend"] = "direct"
|
||||
cfg["ai"]["generate_cover"] = True
|
||||
accounts.create_account("主店", "alias-a", debug_port=9222, config=cfg)
|
||||
batch_id = db.create_batch(["input.xlsx"], path=cfg["db_path"])
|
||||
@@ -2337,6 +2374,7 @@ class GuiTests(TempDirMixin, unittest.TestCase):
|
||||
with self.make_temp_dir() as temp_dir:
|
||||
cfg = self.make_config(temp_dir)
|
||||
cfg["ai"] = appconfig.ai_config(cfg)
|
||||
cfg["ai"]["backend"] = "direct"
|
||||
cfg["ai"]["generate_cover"] = True
|
||||
accounts.create_account("主店", "alias-a", debug_port=9222, config=cfg)
|
||||
batch_id = db.create_batch(["input.xlsx"], path=cfg["db_path"])
|
||||
|
||||
Reference in New Issue
Block a user