feat: add cmhub AI backend

This commit is contained in:
chengma
2026-07-04 15:13:15 +08:00
parent c8a5e9ada8
commit 1efb095767
11 changed files with 1061 additions and 43 deletions
+319
View File
@@ -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()
+30
View File
@@ -1,3 +1,4 @@
import json
import os
import sys
import unittest
@@ -30,6 +31,35 @@ class AppConfigTests(TempDirMixin, unittest.TestCase):
self.assert_removed(temp_dir)
def test_cmhub_defaults_old_config_and_key_helper(self):
with self.make_temp_dir() as temp_dir:
config_path = os.path.join(temp_dir, "config.json")
cmhub_path = os.path.join(temp_dir, "config", "cmhub.json")
config = appconfig.load_config(config_path)
ai = appconfig.ai_config(config)
self.assertEqual("direct", ai["backend"])
self.assertEqual("direct", appconfig.ai_backend(config))
self.assertEqual("", appconfig.cmhub_config(config)["base_url"])
self.assertFalse(os.path.exists(cmhub_path))
self.assertEqual({"api_key": ""}, appconfig.load_cmhub_config(cmhub_path))
saved = appconfig.save_cmhub_config(
{"api_key": "sk-cmhub-123456"},
path=cmhub_path,
)
self.assertEqual("sk-cmhub-123456", saved["api_key"])
self.assertEqual("sk-cmhub-123456", appconfig.get_cmhub_api_key(cmhub_path))
self.assertEqual("sk-c***3456", appconfig.get_cmhub_api_key(cmhub_path, masked=True))
with open(config_path, "w", encoding="utf-8") as fh:
json.dump({"ai": {"resolution": "512"}}, fh)
migrated = appconfig.load_config(config_path)
self.assertEqual("direct", appconfig.ai_config(migrated)["backend"])
self.assertEqual(180, appconfig.response_timeout(migrated))
self.assert_removed(temp_dir)
def test_config_rejects_sensitive_fields(self):
with self.make_temp_dir() as temp_dir:
config_path = os.path.join(temp_dir, "config.json")