Files
cmbot/tests/test_ai_image_service.py
adminandClaude Opus 4.8 5ca615f7ae fix(ai-outfit): per-resolution timeout + _to_int default fallback
- §8 dynamic timeout: add resolution_timeout() (512/1K/2K/4K ->
  180/240/360/600s); generate() uses it for the POST and image download.
  timeout_seconds now defaults to 0 = auto-by-resolution; an explicit
  positive value overrides. Validation allows 0, rejects negatives.
- _to_int now returns the supplied default on parse failure (was 0), so a
  bad connect_timeout_seconds falls back to 30 instead of failing as 0.
- Tests: +5 (resolution map, auto vs override timeout, connect fallback).
- tasks.md §19.1/§19.2: tick unit-test item, note runner is pure-logic.

Co-Authored-By: Claude Opus 4.8 <noreply@anthropic.com>
2026-06-18 18:03:52 +08:00

311 lines
11 KiB
Python

"""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")
def test_generate_uses_resolution_timeout_when_unset(self):
from services.ai_image_service import ImageApiClient
session = _FakeSession()
config = {"url": "https://api.example.test", "model": "m", "api_key": "k", "api_type": "chat"}
ImageApiClient(config, session=session).generate("p", self.image_path, resolution="4K")
# timeout_seconds absent -> (connect 30, read 600 for 4K)
self.assertEqual(session.last_timeout, (30, 600))
def test_generate_explicit_timeout_overrides_resolution(self):
from services.ai_image_service import ImageApiClient
session = _FakeSession()
config = {
"url": "https://api.example.test",
"model": "m",
"api_key": "k",
"api_type": "chat",
"timeout_seconds": 120,
}
ImageApiClient(config, session=session).generate("p", self.image_path, resolution="4K")
self.assertEqual(session.last_timeout, (30, 120))
class TestTimeoutHelpers(unittest.TestCase):
def test_resolution_timeout_mapping(self):
from services.ai_image_service import resolution_timeout
self.assertEqual(resolution_timeout("512"), 180)
self.assertEqual(resolution_timeout("1K"), 240)
self.assertEqual(resolution_timeout("2k"), 360) # case-insensitive
self.assertEqual(resolution_timeout("4K"), 600)
self.assertEqual(resolution_timeout("weird"), 240) # unknown -> default
def test_blank_timeout_is_auto_and_valid(self):
from services.ai_image_service import AiModelConfig, api_config_errors
cfg = AiModelConfig.from_dict({"url": "https://x", "model": "m", "api_key": "k"})
self.assertEqual(cfg.timeout_seconds, 0) # 0 = auto by resolution
self.assertEqual(api_config_errors(cfg), [])
def test_invalid_connect_timeout_falls_back_to_default(self):
from services.ai_image_service import AiModelConfig
cfg = AiModelConfig.from_dict(
{"url": "https://x", "model": "m", "api_key": "k", "connect_timeout_seconds": "bad"}
)
self.assertEqual(cfg.connect_timeout_seconds, 30)
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
self.last_timeout = 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
self.last_timeout = timeout
return _FakeResponse()
if __name__ == "__main__":
unittest.main()