feat: add ai image service

This commit is contained in:
2026-06-18 17:40:56 +08:00
parent 10238d3da3
commit fe1dc7778e
3 changed files with 639 additions and 1 deletions
+255
View File
@@ -0,0 +1,255 @@
"""Tests for AI image service helpers and client request construction."""
import base64
import shutil
import sys
import tempfile
import unittest
from pathlib import Path
sys.path.insert(0, str(Path(__file__).parent.parent / "src"))
class TestAiImageServiceHelpers(unittest.TestCase):
def setUp(self):
self.tmp = Path(tempfile.mkdtemp())
self.image_path = self.tmp / "sample.png"
self.image_bytes = b"\x89PNG\r\n\x1a\nfake-png"
self.image_path.write_bytes(self.image_bytes)
def tearDown(self):
shutil.rmtree(str(self.tmp), ignore_errors=True)
def test_api_config_errors_require_core_fields(self):
from services.ai_image_service import api_config_errors
errors = api_config_errors({"url": "", "model": "", "api_key": ""})
self.assertIn("缺少 url", errors)
self.assertIn("缺少 model", errors)
self.assertIn("缺少 api_key", errors)
def test_api_config_errors_reject_non_object_config(self):
from services.ai_image_service import api_config_errors
self.assertEqual(api_config_errors(["bad"]), ["AI 模型配置必须是对象"])
def test_api_config_errors_reject_invalid_extra_body_and_timeout(self):
from services.ai_image_service import api_config_errors
errors = api_config_errors({
"url": "https://api.example.test",
"model": "m",
"api_key": "k",
"timeout_seconds": "bad",
"extra_body": ["bad"],
})
self.assertIn("timeout_seconds 必须大于 0", errors)
self.assertIn("extra_body 必须是对象", errors)
def test_detect_api_type_from_url(self):
from services.ai_image_service import (
API_CHAT,
API_GEMINI,
API_IMAGES_EDITS,
detect_api_type,
)
self.assertEqual(detect_api_type("https://x/v1/chat/completions"), API_CHAT)
self.assertEqual(detect_api_type("https://x/v1/images/edits"), API_IMAGES_EDITS)
self.assertEqual(detect_api_type("https://x/v1beta/models/m:generateContent"), API_GEMINI)
def test_normalize_openai_urls(self):
from services.ai_image_service import API_CHAT, API_IMAGES_EDITS, normalize_api_url
self.assertEqual(
normalize_api_url("https://api.example.test", API_CHAT),
"https://api.example.test/v1/chat/completions",
)
self.assertEqual(
normalize_api_url("https://api.example.test/v1", API_IMAGES_EDITS),
"https://api.example.test/v1/images/edits",
)
def test_image_to_data_url(self):
from services.ai_image_service import image_to_data_url
data_url = image_to_data_url(self.image_path)
self.assertTrue(data_url.startswith("data:image/png;base64,"))
encoded = data_url.split(",", 1)[1]
self.assertEqual(base64.b64decode(encoded), self.image_bytes)
def test_build_chat_payload(self):
from services.ai_image_service import build_payload
payload = build_payload(
{
"url": "https://api.example.test/v1/chat/completions",
"model": "m1",
"api_key": "k",
"api_type": "chat",
"extra_body": {"temperature": 0},
},
"生成穿搭",
"data:image/png;base64,abc",
)
self.assertEqual(payload["model"], "m1")
self.assertEqual(payload["messages"][0]["content"][0]["text"], "生成穿搭")
self.assertEqual(
payload["messages"][0]["content"][1]["image_url"]["url"],
"data:image/png;base64,abc",
)
self.assertEqual(payload["temperature"], 0)
def test_build_gemini_payload(self):
from services.ai_image_service import build_payload
payload = build_payload(
{
"url": "https://generativelanguage.googleapis.com/v1beta/models/m:generateContent",
"model": "gemini",
"api_key": "k",
"api_type": "gemini",
},
"prompt",
"data:image/jpeg;base64,aW1n",
)
inline = payload["contents"][0]["parts"][1]["inlineData"]
self.assertEqual(inline["mimeType"], "image/jpeg")
self.assertEqual(inline["data"], "aW1n")
def test_build_images_payload(self):
from services.ai_image_service import build_payload
payload = build_payload(
{
"url": "https://api.example.test/v1/images/generations",
"model": "img",
"api_key": "k",
"api_type": "images",
},
"prompt",
"data:image/png;base64,abc",
resolution="1K",
)
self.assertEqual(payload["image_urls"], ["data:image/png;base64,abc"])
self.assertEqual(payload["aspect_ratio"], "1:1")
self.assertEqual(payload["resolution"], "1K")
def test_build_multipart_fields(self):
from services.ai_image_service import build_multipart_fields
data, files = build_multipart_fields(
{
"url": "https://api.example.test/v1/images/edits",
"model": "edit",
"api_key": "k",
"api_type": "images_edits",
},
"prompt",
self.image_path,
resolution="1K",
)
try:
self.assertEqual(data["model"], "edit")
self.assertEqual(data["size"], "1024x1024")
self.assertEqual(files["image"][0], "sample.png")
self.assertEqual(files["image"][2], "image/png")
finally:
files["image"][1].close()
def test_extract_data_url_from_response(self):
from services.ai_image_service import extract_image_from_response
encoded = base64.b64encode(b"image-bytes").decode("ascii")
payload = {"choices": [{"message": {"content": "x", "image": "data:image/png;base64," + encoded}}]}
self.assertEqual(extract_image_from_response(payload), b"image-bytes")
def test_extract_base64_from_response(self):
from services.ai_image_service import extract_image_from_response
encoded = base64.b64encode(b"image-bytes").decode("ascii")
payload = {"data": [{"b64_json": encoded}]}
self.assertEqual(extract_image_from_response(payload), b"image-bytes")
def test_extract_image_url_downloads(self):
from services.ai_image_service import extract_image_from_response
session = _FakeSession()
payload = {"data": [{"url": "https://cdn.example.test/out.jpg"}]}
self.assertEqual(extract_image_from_response(payload, session=session), b"downloaded")
self.assertEqual(session.last_url, "https://cdn.example.test/out.jpg")
class TestImageApiClient(unittest.TestCase):
def setUp(self):
self.tmp = Path(tempfile.mkdtemp())
self.image_path = self.tmp / "sample.png"
self.image_path.write_bytes(b"\x89PNG\r\n\x1a\nfake-png")
def tearDown(self):
shutil.rmtree(str(self.tmp), ignore_errors=True)
def test_generate_posts_chat_payload_and_extracts_image(self):
from services.ai_image_service import ImageApiClient
session = _FakeSession()
config = {
"url": "https://api.example.test",
"model": "model-x",
"api_key": "secret",
"api_type": "chat",
}
result = ImageApiClient(config, session=session).generate("prompt", self.image_path)
self.assertEqual(result, b"image-bytes")
self.assertEqual(session.last_post_url, "https://api.example.test/v1/chat/completions")
self.assertEqual(session.last_headers["Authorization"], "Bearer secret")
self.assertEqual(session.last_json["model"], "model-x")
class _FakeResponse:
def __init__(self, payload=None, content=b"downloaded"):
self._payload = payload or {
"data": [
{"b64_json": base64.b64encode(b"image-bytes").decode("ascii")}
]
}
self.content = content
def raise_for_status(self):
return None
def json(self):
return self._payload
class _FakeSession:
def __init__(self):
self.trust_env = True
self.last_url = None
self.last_post_url = None
self.last_headers = None
self.last_json = None
def get(self, url, timeout=60):
self.last_url = url
return _FakeResponse(content=b"downloaded")
def post(self, url, headers=None, json=None, data=None, files=None, timeout=None):
self.last_post_url = url
self.last_headers = headers
self.last_json = json
return _FakeResponse()
if __name__ == "__main__":
unittest.main()