- §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>
413 lines
13 KiB
Python
413 lines
13 KiB
Python
import base64
|
||
import binascii
|
||
import logging
|
||
import mimetypes
|
||
from dataclasses import dataclass, field
|
||
from pathlib import Path
|
||
from typing import Any, Iterable, Optional, Tuple
|
||
from urllib.parse import urljoin, urlparse
|
||
|
||
import requests
|
||
|
||
logger = logging.getLogger(__name__)
|
||
|
||
API_AUTO = "auto"
|
||
API_CHAT = "chat"
|
||
API_GEMINI = "gemini"
|
||
API_IMAGES = "images"
|
||
API_IMAGES_EDITS = "images_edits"
|
||
|
||
SUPPORTED_API_TYPES = {
|
||
API_AUTO,
|
||
API_CHAT,
|
||
API_GEMINI,
|
||
API_IMAGES,
|
||
API_IMAGES_EDITS,
|
||
}
|
||
|
||
_BASE64_KEYS = {"image_base64", "base64", "b64_json", "data"}
|
||
_IMAGE_EXTENSIONS = {".png", ".jpg", ".jpeg", ".webp", ".gif"}
|
||
|
||
# Default read timeout (seconds) per output resolution (docs/11-ai-outfit.md §8).
|
||
# Higher resolutions take longer to generate; a fixed timeout_seconds in the
|
||
# model config overrides this mapping.
|
||
RESOLUTION_TIMEOUTS = {"512": 180, "1K": 240, "2K": 360, "4K": 600}
|
||
|
||
|
||
def resolution_timeout(resolution, default=240):
|
||
"""Return the default read timeout in seconds for *resolution*."""
|
||
return RESOLUTION_TIMEOUTS.get(str(resolution).strip().upper(), default)
|
||
|
||
|
||
class AiImageServiceError(RuntimeError):
|
||
"""Base error for AI image generation service failures."""
|
||
|
||
|
||
class AiConfigError(ValueError):
|
||
"""Raised when model configuration is missing or invalid."""
|
||
|
||
|
||
@dataclass
|
||
class AiModelConfig:
|
||
"""AI image model configuration loaded from ai_models.json."""
|
||
|
||
url: str
|
||
model: str
|
||
api_key: str
|
||
api_type: str = API_AUTO
|
||
timeout_seconds: int = 0 # 0 = 按分辨率自动(resolution_timeout)
|
||
connect_timeout_seconds: int = 30
|
||
extra_body: Any = field(default_factory=dict)
|
||
|
||
@classmethod
|
||
def from_dict(cls, data):
|
||
if not isinstance(data, dict):
|
||
raise AiConfigError("AI 模型配置必须是对象")
|
||
return cls(
|
||
url=str(data.get("url", "") or ""),
|
||
model=str(data.get("model", "") or ""),
|
||
api_key=str(data.get("api_key", "") or ""),
|
||
api_type=str(data.get("api_type", API_AUTO) or API_AUTO),
|
||
timeout_seconds=_parse_optional_timeout(data.get("timeout_seconds")),
|
||
connect_timeout_seconds=_to_int(data.get("connect_timeout_seconds", 30), 30),
|
||
extra_body=data.get("extra_body") or {},
|
||
)
|
||
|
||
|
||
def api_config_errors(config):
|
||
"""Return human-readable configuration errors without raising."""
|
||
try:
|
||
cfg = _coerce_config(config)
|
||
except AiConfigError as exc:
|
||
return [str(exc)]
|
||
errors = []
|
||
|
||
if not cfg.url.strip():
|
||
errors.append("缺少 url")
|
||
else:
|
||
scheme = urlparse(cfg.url).scheme.lower()
|
||
if scheme not in ("http", "https"):
|
||
errors.append("url 必须是 http(s) 地址")
|
||
if not cfg.model.strip():
|
||
errors.append("缺少 model")
|
||
if not cfg.api_key.strip():
|
||
errors.append("缺少 api_key")
|
||
if cfg.api_type not in SUPPORTED_API_TYPES:
|
||
errors.append("api_type 不支持: {}".format(cfg.api_type))
|
||
if cfg.timeout_seconds < 0: # 0 = 按分辨率自动;负数 = 解析失败/非法
|
||
errors.append("timeout_seconds 必须大于 0")
|
||
if cfg.connect_timeout_seconds <= 0:
|
||
errors.append("connect_timeout_seconds 必须大于 0")
|
||
if not isinstance(cfg.extra_body, dict):
|
||
errors.append("extra_body 必须是对象")
|
||
|
||
return errors
|
||
|
||
|
||
def validate_api_config(config):
|
||
errors = api_config_errors(config)
|
||
if errors:
|
||
raise AiConfigError("; ".join(errors))
|
||
|
||
|
||
def detect_api_type(url, api_type=API_AUTO):
|
||
"""Resolve auto/chat/gemini/images/images_edits for a model endpoint."""
|
||
if api_type and api_type != API_AUTO:
|
||
if api_type not in SUPPORTED_API_TYPES:
|
||
raise AiConfigError("api_type 不支持: {}".format(api_type))
|
||
return api_type
|
||
|
||
lower_url = str(url).lower()
|
||
if "generatecontent" in lower_url or "gemini" in lower_url:
|
||
return API_GEMINI
|
||
if "/images/edits" in lower_url:
|
||
return API_IMAGES_EDITS
|
||
if "/images" in lower_url:
|
||
return API_IMAGES
|
||
return API_CHAT
|
||
|
||
|
||
def normalize_api_url(url, api_type):
|
||
"""Return an endpoint URL for the selected API type.
|
||
|
||
Full endpoint URLs are kept. Base URLs are completed using OpenAI-compatible
|
||
defaults or Gemini generateContent where appropriate.
|
||
"""
|
||
raw = str(url).strip()
|
||
if not raw:
|
||
return raw
|
||
|
||
endpoint = {
|
||
API_CHAT: "chat/completions",
|
||
API_IMAGES: "images/generations",
|
||
API_IMAGES_EDITS: "images/edits",
|
||
}.get(api_type)
|
||
|
||
lower_path = urlparse(raw).path.lower().rstrip("/")
|
||
if api_type == API_GEMINI:
|
||
if "generatecontent" in lower_path:
|
||
return raw
|
||
return _join_url(raw, "v1beta/models/{model}:generateContent")
|
||
|
||
if endpoint is None:
|
||
return raw
|
||
if lower_path.endswith("/" + endpoint):
|
||
return raw
|
||
if lower_path.endswith("/v1"):
|
||
return _join_url(raw, endpoint)
|
||
return _join_url(raw, "v1/" + endpoint)
|
||
|
||
|
||
def image_to_data_url(image_path):
|
||
"""Read a local image file and return data:<mime>;base64,..."""
|
||
path = Path(image_path)
|
||
mime_type = mimetypes.guess_type(str(path))[0] or "application/octet-stream"
|
||
with open(str(path), "rb") as f:
|
||
encoded = base64.b64encode(f.read()).decode("ascii")
|
||
return "data:{};base64,{}".format(mime_type, encoded)
|
||
|
||
|
||
def decode_image_data_url(data_url):
|
||
marker = ";base64,"
|
||
if marker not in data_url:
|
||
raise AiImageServiceError("Unsupported data URL image format")
|
||
return base64.b64decode(data_url.split(marker, 1)[1])
|
||
|
||
|
||
def build_payload(config, prompt, image_data_url, resolution="1K", aspect_ratio="1:1"):
|
||
"""Build JSON body for chat/gemini/images request types."""
|
||
cfg = _coerce_config(config)
|
||
api_type = detect_api_type(cfg.url, cfg.api_type)
|
||
|
||
if api_type == API_CHAT:
|
||
payload = {
|
||
"model": cfg.model,
|
||
"messages": [
|
||
{
|
||
"role": "user",
|
||
"content": [
|
||
{"type": "text", "text": prompt},
|
||
{"type": "image_url", "image_url": {"url": image_data_url}},
|
||
],
|
||
}
|
||
],
|
||
"stream": False,
|
||
}
|
||
elif api_type == API_GEMINI:
|
||
mime_type, data = _split_data_url(image_data_url)
|
||
payload = {
|
||
"contents": [
|
||
{
|
||
"parts": [
|
||
{"text": prompt},
|
||
{"inlineData": {"mimeType": mime_type, "data": data}},
|
||
]
|
||
}
|
||
],
|
||
"generationConfig": {"responseModalities": ["TEXT", "IMAGE"]},
|
||
}
|
||
elif api_type == API_IMAGES:
|
||
payload = {
|
||
"model": cfg.model,
|
||
"prompt": prompt,
|
||
"image_urls": [image_data_url],
|
||
"aspect_ratio": aspect_ratio,
|
||
"resolution": resolution,
|
||
"n": 1,
|
||
}
|
||
else:
|
||
raise AiConfigError("build_payload 不支持 {}".format(api_type))
|
||
|
||
payload.update(cfg.extra_body)
|
||
return payload
|
||
|
||
|
||
def build_multipart_fields(config, prompt, image_path, resolution="1K"):
|
||
"""Build data/files for images_edits requests.
|
||
|
||
Caller owns and must close the returned file object.
|
||
"""
|
||
cfg = _coerce_config(config)
|
||
handle = open(str(image_path), "rb")
|
||
mime_type = mimetypes.guess_type(str(image_path))[0] or "application/octet-stream"
|
||
data = {
|
||
"model": cfg.model,
|
||
"prompt": prompt,
|
||
"n": "1",
|
||
"size": _resolution_to_size(resolution),
|
||
}
|
||
data.update(cfg.extra_body)
|
||
files = {"image": (Path(image_path).name, handle, mime_type)}
|
||
return data, files
|
||
|
||
|
||
def extract_image_from_response(data, session=None, timeout=60):
|
||
"""Return image bytes found in an AI JSON response, or None.
|
||
|
||
Searches data-url/base64 fields first, then image URLs.
|
||
"""
|
||
for key, value in _walk_json_items(data):
|
||
if not isinstance(value, str):
|
||
continue
|
||
text = value.strip()
|
||
if text.startswith("data:image/"):
|
||
return decode_image_data_url(text)
|
||
if key and key.lower() in _BASE64_KEYS and _looks_like_base64(text):
|
||
try:
|
||
return base64.b64decode(text)
|
||
except (TypeError, ValueError, binascii.Error):
|
||
pass
|
||
|
||
for _key, value in _walk_json_items(data):
|
||
if isinstance(value, str) and _is_image_url(value):
|
||
return download_image_url(value, session=session, timeout=timeout)
|
||
|
||
return None
|
||
|
||
|
||
def download_image_url(url, session=None, timeout=60):
|
||
client = session or requests.Session()
|
||
if hasattr(client, "trust_env"):
|
||
client.trust_env = False
|
||
response = client.get(url, timeout=timeout)
|
||
response.raise_for_status()
|
||
return response.content
|
||
|
||
|
||
class ImageApiClient:
|
||
"""Small HTTP client for AI image APIs.
|
||
|
||
GUI code should not use requests directly; it should call this service from
|
||
worker threads and send progress back through Qt signals.
|
||
"""
|
||
|
||
def __init__(self, config, session=None):
|
||
self.config = _coerce_config(config)
|
||
self.session = session or requests.Session()
|
||
if hasattr(self.session, "trust_env"):
|
||
self.session.trust_env = False
|
||
|
||
def generate(self, prompt, image_path, resolution="1K", aspect_ratio="1:1"):
|
||
validate_api_config(self.config)
|
||
api_type = detect_api_type(self.config.url, self.config.api_type)
|
||
url = normalize_api_url(self.config.url, api_type)
|
||
if api_type == API_GEMINI:
|
||
url = url.replace("{model}", self.config.model)
|
||
|
||
headers = {"Authorization": "Bearer {}".format(self.config.api_key)}
|
||
read_timeout = (
|
||
self.config.timeout_seconds
|
||
if self.config.timeout_seconds > 0
|
||
else resolution_timeout(resolution)
|
||
)
|
||
timeout = (self.config.connect_timeout_seconds, read_timeout)
|
||
|
||
if api_type == API_IMAGES_EDITS:
|
||
data, files = build_multipart_fields(self.config, prompt, image_path, resolution)
|
||
try:
|
||
response = self.session.post(
|
||
url,
|
||
headers=headers,
|
||
data=data,
|
||
files=files,
|
||
timeout=timeout,
|
||
)
|
||
finally:
|
||
files["image"][1].close()
|
||
else:
|
||
data_url = image_to_data_url(image_path)
|
||
payload = build_payload(self.config, prompt, data_url, resolution, aspect_ratio)
|
||
headers["Content-Type"] = "application/json"
|
||
response = self.session.post(url, headers=headers, json=payload, timeout=timeout)
|
||
|
||
response.raise_for_status()
|
||
payload = response.json()
|
||
image_bytes = extract_image_from_response(
|
||
payload,
|
||
session=self.session,
|
||
timeout=read_timeout,
|
||
)
|
||
if not image_bytes:
|
||
raise AiImageServiceError("AI 响应中未找到图片")
|
||
return image_bytes
|
||
|
||
|
||
def _coerce_config(config):
|
||
if isinstance(config, AiModelConfig):
|
||
return config
|
||
return AiModelConfig.from_dict(config)
|
||
|
||
|
||
def _to_int(value, default):
|
||
try:
|
||
return int(value or default)
|
||
except (TypeError, ValueError):
|
||
return default
|
||
|
||
|
||
def _parse_optional_timeout(value):
|
||
"""Parse an optional read timeout.
|
||
|
||
Absent/blank -> 0 (auto by resolution); an explicit positive value overrides
|
||
the per-resolution default; an unparseable value -> -1 so validation flags it.
|
||
"""
|
||
if value is None or value == "":
|
||
return 0
|
||
try:
|
||
return int(value)
|
||
except (TypeError, ValueError):
|
||
return -1
|
||
|
||
|
||
def _join_url(base_url, suffix):
|
||
base = str(base_url).rstrip("/") + "/"
|
||
return urljoin(base, suffix)
|
||
|
||
|
||
def _split_data_url(data_url):
|
||
prefix, encoded = data_url.split(",", 1)
|
||
mime_type = prefix[len("data:"):].split(";", 1)[0]
|
||
return mime_type, encoded
|
||
|
||
|
||
def _resolution_to_size(resolution):
|
||
mapping = {
|
||
"512": "512x512",
|
||
"512px": "512x512",
|
||
"1K": "1024x1024",
|
||
"2K": "2048x2048",
|
||
"4K": "4096x4096",
|
||
}
|
||
return mapping.get(str(resolution), str(resolution))
|
||
|
||
|
||
def _walk_json_items(value, key=None):
|
||
# type: (Any, Optional[str]) -> Iterable[Tuple[Optional[str], Any]]
|
||
yield key, value
|
||
if isinstance(value, dict):
|
||
for child_key, child_value in value.items():
|
||
for item in _walk_json_items(child_value, str(child_key)):
|
||
yield item
|
||
elif isinstance(value, list):
|
||
for child_value in value:
|
||
for item in _walk_json_items(child_value, key):
|
||
yield item
|
||
|
||
|
||
def _looks_like_base64(text):
|
||
if len(text) < 8:
|
||
return False
|
||
try:
|
||
base64.b64decode(text, validate=True)
|
||
return True
|
||
except (TypeError, ValueError, binascii.Error):
|
||
return False
|
||
|
||
|
||
def _is_image_url(text):
|
||
parsed = urlparse(text.strip())
|
||
if parsed.scheme.lower() not in ("http", "https"):
|
||
return False
|
||
suffix = Path(parsed.path).suffix.lower()
|
||
return suffix in _IMAGE_EXTENSIONS
|