feat: cap cmhub image concurrency

This commit is contained in:
chengma
2026-07-07 21:00:43 +08:00
parent 2928624019
commit 3d902d4907
12 changed files with 601 additions and 95 deletions
+183
View File
@@ -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)
+38
View File
@@ -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"])