"""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()