feat: add cmhub AI backend
This commit is contained in:
@@ -2,6 +2,7 @@ import base64
|
||||
import io
|
||||
import json
|
||||
import os
|
||||
import socket
|
||||
import sys
|
||||
import unittest
|
||||
from types import SimpleNamespace
|
||||
@@ -28,6 +29,22 @@ class _Response:
|
||||
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)
|
||||
|
||||
def json(self):
|
||||
return self.payload
|
||||
|
||||
def iter_content(self, chunk_size=65536):
|
||||
if self.content:
|
||||
yield self.content
|
||||
|
||||
class AITests(TempDirMixin, unittest.TestCase):
|
||||
def _write_models(self, path, text=None, image=None):
|
||||
text = text or {
|
||||
@@ -65,6 +82,21 @@ class AITests(TempDirMixin, unittest.TestCase):
|
||||
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",
|
||||
"connect_timeout": 3,
|
||||
"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 _collected_tasks(self, temp_dir, cfg, titles=None):
|
||||
titles = titles or ["旧标题A", "旧标题B"]
|
||||
db.init_db(cfg["db_path"])
|
||||
@@ -240,6 +272,293 @@ class AITests(TempDirMixin, unittest.TestCase):
|
||||
|
||||
self.assert_removed(temp_dir)
|
||||
|
||||
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("app.ai.requests.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, 180), 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_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_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 = []
|
||||
|
||||
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("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=[
|
||||
(
|
||||
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_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("https://cdn.example.com/generated.png", downloads[0][0])
|
||||
self.assertEqual((3, 240), downloads[0][1]["timeout"])
|
||||
self.assertEqual(91, events[0]["metadata"]["points_balance"])
|
||||
with Image.open(output) as saved:
|
||||
self.assertEqual((1024, 1024), saved.size)
|
||||
|
||||
self.assert_removed(temp_dir)
|
||||
|
||||
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_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("app.ai.requests.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("app.ai.requests.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.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("app.ai.requests.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_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("app.ai.requests.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()
|
||||
|
||||
Reference in New Issue
Block a user