You can not select more than 25 topics
Topics must start with a letter or number, can include dashes ('-') and can be up to 35 characters long.
140 lines
5.9 KiB
140 lines
5.9 KiB
"""Agnes Image 2.1 Flash 文生图 / 图生图 / 多图合成的核心逻辑。"""
|
|
|
|
import base64
|
|
import json
|
|
import os
|
|
import time
|
|
import urllib.error
|
|
import urllib.request
|
|
from pathlib import Path
|
|
from uuid import uuid4
|
|
|
|
from fastapi import HTTPException
|
|
|
|
API_BASE_URL = os.getenv("AGNES_API_BASE_URL", "https://apihub.agnes-ai.com/v1")
|
|
OUTPUT_DIR = Path("output")
|
|
MAX_RETRIES = 3
|
|
RETRYABLE_STATUS_CODES = {408, 429, 500, 502, 503, 504}
|
|
DEFAULT_MODEL = "agnes-image-2.1-flash"
|
|
VALID_RATIOS = {"1:1", "3:4", "4:3", "16:9", "9:16", "2:3", "3:2", "21:9"}
|
|
RATIO_DIMENSIONS = {
|
|
"1:1": {"1K": "1024x1024", "2K": "2048x2048", "3K": "3072x3072", "4K": "4096x4096"},
|
|
"3:4": {"1K": "864x1152", "2K": "1728x2304", "3K": "2592x3456", "4K": "3456x4608"},
|
|
"4:3": {"1K": "1152x864", "2K": "2304x1728", "3K": "3456x2592", "4K": "4608x3456"},
|
|
"16:9": {"1K": "1312x736", "2K": "2624x1472", "3K": "3936x2208", "4K": "5248x2944"},
|
|
"9:16": {"1K": "736x1312", "2K": "1472x2624", "3K": "2208x3936", "4K": "2944x5248"},
|
|
"2:3": {"1K": "832x1248", "2K": "1664x2496", "3K": "2496x3744", "4K": "3328x4992"},
|
|
"3:2": {"1K": "1248x832", "2K": "2496x1664", "3K": "3744x2496", "4K": "4992x3328"},
|
|
"21:9": {"1K": "1568x672", "2K": "3136x1344", "3K": "4704x2016", "4K": "6272x2688"},
|
|
}
|
|
|
|
|
|
def _resolve_size(size: str, ratio: str | None) -> str:
|
|
"""把 size 档位 (1K/2K/3K/4K) 配合 ratio 解析成实际像素尺寸,方便保存文件。"""
|
|
if ratio and ratio not in VALID_RATIOS:
|
|
raise HTTPException(status_code=400, detail=f"不支持的宽高比: {ratio}")
|
|
if size in {"1K", "2K", "3K", "4K"}:
|
|
if not ratio:
|
|
ratio = "1:1"
|
|
return RATIO_DIMENSIONS[ratio][size]
|
|
return size
|
|
|
|
|
|
def _build_payload(*, prompt: str, size: str, ratio: str | None, images: list[str] | None) -> dict:
|
|
"""组装请求体。注意 response_format 必须放在 extra_body 内。"""
|
|
payload: dict = {
|
|
"model": DEFAULT_MODEL,
|
|
"prompt": prompt,
|
|
"size": size,
|
|
}
|
|
extra_body: dict = {}
|
|
if ratio:
|
|
payload["ratio"] = ratio
|
|
if images:
|
|
extra_body["image"] = images
|
|
if extra_body:
|
|
payload["extra_body"] = extra_body
|
|
return payload
|
|
|
|
|
|
def _post_with_retry(url: str, headers: dict, payload: dict) -> dict:
|
|
"""调用 Agnes Image 接口,并对可恢复的上游错误进行指数退避重试。"""
|
|
data = json.dumps(payload).encode("utf-8")
|
|
for attempt in range(1, MAX_RETRIES + 1):
|
|
request = urllib.request.Request(url, data=data, headers=headers, method="POST")
|
|
try:
|
|
with urllib.request.urlopen(request, timeout=360) as response:
|
|
return json.loads(response.read().decode("utf-8"))
|
|
except urllib.error.HTTPError as e:
|
|
if e.code not in RETRYABLE_STATUS_CODES or attempt == MAX_RETRIES:
|
|
body = e.read().decode("utf-8", errors="replace")
|
|
raise HTTPException(status_code=502, detail=f"Agnes 图片服务错误: HTTP {e.code} {body}") from e
|
|
except (urllib.error.URLError, TimeoutError) as e:
|
|
if attempt == MAX_RETRIES:
|
|
raise HTTPException(status_code=504, detail="Agnes 图片服务连接或请求超时") from e
|
|
time.sleep(2 ** (attempt - 1))
|
|
|
|
|
|
def _download_image(url: str) -> bytes:
|
|
"""从 URL 或 data URI 下载图片。"""
|
|
if url.startswith("data:image/"):
|
|
try:
|
|
_, encoded = url.split(",", 1)
|
|
return base64.b64decode(encoded, validate=True)
|
|
except (ValueError, TypeError) as e:
|
|
raise HTTPException(status_code=502, detail="图片 data URL 格式无效") from e
|
|
|
|
request = urllib.request.Request(url, headers={"User-Agent": "Mozilla/5.0"})
|
|
for attempt in range(1, MAX_RETRIES + 1):
|
|
try:
|
|
with urllib.request.urlopen(request, timeout=120) as response:
|
|
return response.read()
|
|
except urllib.error.HTTPError as e:
|
|
if e.code not in RETRYABLE_STATUS_CODES or attempt == MAX_RETRIES:
|
|
raise HTTPException(status_code=502, detail=f"图片 URL 下载失败: HTTP {e.code}") from e
|
|
except (urllib.error.URLError, TimeoutError) as e:
|
|
if attempt == MAX_RETRIES:
|
|
raise HTTPException(status_code=504, detail="图片 URL 下载超时或网络连接失败") from e
|
|
time.sleep(2 ** (attempt - 1))
|
|
|
|
|
|
def generate_images(
|
|
*,
|
|
prompt: str,
|
|
size: str,
|
|
ratio: str | None,
|
|
images: list[str] | None,
|
|
profile_name: str,
|
|
):
|
|
"""生成图片并保存到本地,返回结果字典。支持文生图 / 图生图 / 多图合成。"""
|
|
api_key = os.getenv("AGNES_API_KEY")
|
|
if not api_key:
|
|
raise HTTPException(status_code=500, detail="未配置环境变量 AGNES_API_KEY")
|
|
|
|
resolved_size = _resolve_size(size, ratio)
|
|
payload = _build_payload(prompt=prompt, size=size, ratio=ratio, images=images)
|
|
headers = {
|
|
"Authorization": f"Bearer {api_key}",
|
|
"Content-Type": "application/json",
|
|
}
|
|
response = _post_with_retry(f"{API_BASE_URL}/images/generations", headers, payload)
|
|
|
|
profile_dir = OUTPUT_DIR / profile_name
|
|
profile_dir.mkdir(parents=True, exist_ok=True)
|
|
saved = []
|
|
for item in response.get("data") or []:
|
|
try:
|
|
if item.get("b64_json"):
|
|
raw = base64.b64decode(item["b64_json"], validate=True)
|
|
elif item.get("url"):
|
|
raw = _download_image(item["url"])
|
|
else:
|
|
raise HTTPException(status_code=502, detail="Agnes 图片服务未返回 b64_json 或 url")
|
|
except ValueError as e:
|
|
raise HTTPException(status_code=502, detail="Agnes Base64 图片数据格式无效") from e
|
|
|
|
filename = f"{DEFAULT_MODEL}-{uuid4().hex}.png"
|
|
(profile_dir / filename).write_bytes(raw)
|
|
saved.append({"filename": filename, "url": f"/output/{profile_name}/{filename}", "bytes": len(raw)})
|
|
|
|
return {"model": DEFAULT_MODEL, "size": resolved_size, "ratio": ratio or "1:1", "images": saved}
|
|
|