Files
comfyui_o1key/clients/gpt_image_client.py
T
Jony ba920f2b66 Publish current ComfyUI O1Key code baseline
Replace the prior release tree with the current plugin, frontend, tests, and documentation. Document retired node IDs and the public Gitea update source.
2026-09-24 19:56:48 +08:00

2018 lines
81 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""
GPT Image API 客户端
新版节点请求走异步任务接口:
- POST /async/v1/generateImage
- GET /async/v1/tasks/{task_id}
旧同步接口保留兼容代码,但 GPT Image 节点不再使用:
- POST /v1/images/generations/
- POST /v1/images/edits
设计原则:
- 对 ComfyUI 节点暴露同步入口,内部提交异步任务并轮询
- 图片和蒙版以 data:image/png;base64,... 放入 JSON 请求体
- 响应优先读取 data.images[].url
"""
import asyncio
import base64
import binascii
import json
import math
import os
import time
from concurrent.futures import ThreadPoolExecutor
from io import BytesIO
from typing import Any, Callable, Dict, List, Optional
import aiohttp
import numpy as np
import torch
from PIL import Image, UnidentifiedImageError
from ..utils.config import get_api_key_or_raise, get_api_base_url
from ..utils.image_utils import tensor_to_pil, encode_image_to_base64
from ..utils.http_error import HTTP_ERROR_MESSAGES, get_friendly_message
from ..utils.http2_client import (
format_response_body_diagnostics,
read_response_body_with_diagnostics,
response_task_id,
validate_response_task_id,
)
from ..utils.o1key_image_catalog import (
GPT_IMAGE_BACKGROUND_OPTIONS,
GPT_IMAGE_MODEL_OPTIONS,
GPT_IMAGE_OUTPUT_FORMAT_OPTIONS,
UNIFIED_IMAGE_ROUTE_OPTIONS,
)
try:
from comfy.model_management import processing_interrupted, InterruptProcessingException
_INTERRUPT_AVAILABLE = True
except ImportError:
_INTERRUPT_AVAILABLE = False
InterruptProcessingException = RuntimeError
processing_interrupted = lambda: False
# ── 接口端点 ──────────────────────────────────────────────────────────────────
_ENDPOINT_GENERATIONS = "/v1/images/generations/"
_ENDPOINT_EDITS = "/v1/images/edits"
_ENDPOINT_ASYNC_GENERATE = "/async/v1/generateImage"
_ENDPOINT_ASYNC_TASK = "/async/v1/tasks/{task_id}"
# ── 主模型与模型线路映射(界面值 → API 模型名)──────────────────────────────
GPT_IMAGE_ROUTE_OPTIONS = UNIFIED_IMAGE_ROUTE_OPTIONS
GPT_IMAGE_MODEL_MATRIX = {
("gpt-image-2", "畅速"): "gpt-image-2-c-sp",
("gpt-image-2", "直连"): "gpt-image-2-c-sd",
("gpt-image-2", "专线"): "gpt-image-2",
("gpt-image-2.5-sunburst", "畅速"): "gpt-image-2.5-sunburst-sp",
("gpt-image-2.5-sunburst", "直连"): "gpt-image-2.5-sunburst-sd",
("gpt-image-2.5-sunburst", "专线"): "gpt-image-2.5-sunburst",
("gpt-image-2.5-flare", "畅速"): "gpt-image-2.5-flare-sp",
("gpt-image-2.5-flare", "直连"): "gpt-image-2.5-flare-sd",
("gpt-image-2.5-flare", "专线"): "gpt-image-2.5-flare",
}
_MODEL_NAME_MAP = {
"畅速": "gpt-image-2-c-sp",
"直连": "gpt-image-2-c-sd",
"专线": "gpt-image-2",
# 兼容旧工作流和直接调用节点的旧参数值。
"gpt-image-2-特价": "gpt-image-2-c-sp",
"gpt-image-2-官方": "gpt-image-2",
"gpt-image-2-次卡": "gpt-image-2-c-sp",
"gpt-image-2-按量": "gpt-image-2",
}
def resolve_gpt_image_model(model_name: str, route: str) -> str:
"""由主模型和线路解析实际 API 模型名,并兼容旧的单参数调用。"""
# 旧节点曾把线路或旧显示模型名直接作为“模型”参数传入。
if model_name in GPT_IMAGE_ROUTE_OPTIONS or model_name in _MODEL_NAME_MAP:
return _MODEL_NAME_MAP.get(model_name, model_name)
model = GPT_IMAGE_MODEL_MATRIX.get((model_name, route))
# 兼容外部调用把旧线路显示值放在 route 参数中的情况。
if model is None and route in _MODEL_NAME_MAP:
model = _MODEL_NAME_MAP[route]
if model is None:
raise ValueError(f"模型 '{model_name}' 不支持模型线路 '{route}'")
return model
# ── 超时 ──────────────────────────────────────────────────────────────────────
_REQUEST_TIMEOUT = 900 # 秒
_ASYNC_POLL_SCHEDULE = [5.0, 20.0]
_ASYNC_POLL_INTERVAL = 3.0
def _attach_original_image_info(image: Image.Image, data: bytes) -> Image.Image:
"""Keep provider bytes alongside decoded pixels for format-preserving saves."""
image_format = str(image.format or "").upper()
if image_format == "JPG":
image_format = "JPEG"
if image_format:
image.format = image_format
setattr(image, "_o1key_original_format", image_format)
setattr(image, "_o1key_original_bytes", data)
return image
_ASYNC_MAX_WAIT = 600.0
# Only use these after a task has been submitted: status/result GET requests
# are idempotent, while retrying task submission can create duplicate charges.
_RESPONSE_READ_MAX_RETRIES = 3
_RESPONSE_READ_BASE_DELAY = 1.0
_RESPONSE_READ_MAX_DELAY = 8.0
_RETRYABLE_RESULT_HTTP_STATUSES = {408, 425, 429, 500, 502, 503, 504}
REQUEST_LOG_ENABLED = False
POLL_LOG_ENABLED = False
RESPONSE_LOG_ENABLED = os.environ.get("O1KEY_RESPONSE_LOG", "0").strip().lower() not in {
"0", "false", "no", "off"
}
class _IncompleteInlineImageError(RuntimeError):
"""The task succeeded, but an inline result image is incomplete."""
class _RetryableResultHTTPError(OSError):
"""An idempotent task query or result download returned a transient status."""
class GptImageClient:
"""
GPT Image API 客户端
接口说明:
async generateImageJSON 提交,返回 task_id
tasks/{task_id}:轮询任务状态,成功后读取 data.images[].url
旧 generations / edits 同步接口保留为兼容代码。
"""
def __init__(self):
self.api_key = get_api_key_or_raise("O1KEY_API_KEY")
self.base_url = get_api_base_url()
# Batch nodes can turn these off while keeping the same client API for
# single-image diagnostics.
self.response_log_enabled = RESPONSE_LOG_ENABLED
self.poll_log_enabled = POLL_LOG_ENABLED
@staticmethod
def _response_retry_delay(attempt: int) -> float:
return min(
_RESPONSE_READ_BASE_DELAY * (2 ** attempt),
_RESPONSE_READ_MAX_DELAY,
)
async def _get_json_with_response_retry(
self,
session: aiohttp.ClientSession,
url: str,
headers: dict,
label: str,
expected_task_id: Optional[str] = None,
) -> dict:
"""Read a JSON GET response, retrying interrupted response bodies."""
for attempt in range(_RESPONSE_READ_MAX_RETRIES + 1):
try:
async with session.get(url, headers=headers) as resp:
status = resp.status
body, diagnostics = await read_response_body_with_diagnostics(resp)
if self.response_log_enabled:
self._log_original_response_body(
f"{label} status={status}",
body.decode("utf-8", errors="replace"),
)
if status != 200:
print(
f"[o1key GPT Image] 任务查询传输追踪"
f" | requested_task_id={expected_task_id or '<unknown>'}"
f" | response_task_id=<not-parsed>"
f" | {format_response_body_diagnostics(diagnostics)}"
f" | json=not-parsed"
)
text = body.decode("utf-8", errors="replace")
message = self._extract_error_message(text, status)
if status in _RETRYABLE_RESULT_HTTP_STATUSES:
raise _RetryableResultHTTPError(message)
raise RuntimeError(message)
try:
payload = json.loads(body)
except json.JSONDecodeError:
print(
f"[o1key GPT Image] 任务查询传输追踪"
f" | requested_task_id={expected_task_id or '<unknown>'}"
f" | response_task_id=<unavailable>"
f" | {format_response_body_diagnostics(diagnostics)}"
f" | json=invalid"
)
raise
actual_task_id = response_task_id(payload)
if expected_task_id is not None:
if actual_task_id is not None and actual_task_id != expected_task_id:
print(
f"[o1key GPT Image] 任务查询传输追踪"
f" | requested_task_id={expected_task_id}"
f" | response_task_id={actual_task_id}"
f" | task_id_check=mismatch"
f" | {format_response_body_diagnostics(diagnostics)}"
f" | json=valid"
)
validate_response_task_id(payload, expected_task_id)
return payload
except (
aiohttp.ClientError,
asyncio.TimeoutError,
OSError,
json.JSONDecodeError,
) as exc:
if attempt >= _RESPONSE_READ_MAX_RETRIES:
raise RuntimeError(
f"{label}响应读取失败,已重试 "
f"{_RESPONSE_READ_MAX_RETRIES} 次:{exc}"
) from None
delay = self._response_retry_delay(attempt)
print(
f"[o1key GPT Image] {label}响应不完整 "
f"({type(exc).__name__}){delay:.1f}s 后重试 "
f"({attempt + 1}/{_RESPONSE_READ_MAX_RETRIES})"
f" | requested_task_id={expected_task_id or '<unknown>'} | {exc}"
)
await asyncio.sleep(delay)
raise RuntimeError(f"{label}响应读取失败") # pragma: no cover
async def _download_image_with_response_retry(
self,
session: aiohttp.ClientSession,
url: str,
label: str,
) -> tuple[Image.Image, int, float]:
"""Download and fully decode an image, retrying interrupted GET reads."""
download_started = time.perf_counter()
for attempt in range(_RESPONSE_READ_MAX_RETRIES + 1):
try:
async with session.get(url, allow_redirects=True) as resp:
image_bytes, diagnostics = await read_response_body_with_diagnostics(resp)
if resp.status != 200:
print(
f"[o1key GPT Image] 结果下载传输追踪"
f" | {format_response_body_diagnostics(diagnostics)}"
)
message = f"{label}下载失败 HTTP {resp.status}"
if resp.status in _RETRYABLE_RESULT_HTTP_STATUSES:
raise _RetryableResultHTTPError(message)
raise RuntimeError(message)
image = Image.open(BytesIO(image_bytes))
image.load()
_attach_original_image_info(image, image_bytes)
return image, len(image_bytes), time.perf_counter() - download_started
except (
aiohttp.ClientError,
asyncio.TimeoutError,
OSError,
UnidentifiedImageError,
) as exc:
if attempt >= _RESPONSE_READ_MAX_RETRIES:
raise RuntimeError(
f"{label}下载失败,已重试 "
f"{_RESPONSE_READ_MAX_RETRIES} 次:{exc}"
) from None
delay = self._response_retry_delay(attempt)
print(
f"[o1key GPT Image] {label}下载中断 "
f"({type(exc).__name__}){delay:.1f}s 后重试 "
f"({attempt + 1}/{_RESPONSE_READ_MAX_RETRIES})..."
)
await asyncio.sleep(delay)
raise RuntimeError(f"{label}下载失败") # pragma: no cover
@staticmethod
def _format_transfer_size(size: int) -> str:
if size < 1024 * 1024:
return f"{size / 1024:.1f} KiB"
return f"{size / 1024 / 1024:.2f} MiB"
@staticmethod
def _format_transfer_rate(size: int, elapsed: float) -> str:
if elapsed <= 0:
return "∞ MiB/s"
return f"{size / 1024 / 1024 / elapsed:.2f} MiB/s"
# ── 认证头 ────────────────────────────────────────────────────────────────
def _auth_headers(self) -> dict:
return {"Authorization": f"Bearer {self.api_key}"}
def _json_headers(self) -> dict:
return {
"Authorization": f"Bearer {self.api_key}",
"Content-Type": "application/json",
}
@staticmethod
def _new_multipart_form() -> aiohttp.FormData:
try:
return aiohttp.FormData(default_to_multipart=True)
except TypeError:
form = aiohttp.FormData()
form._is_multipart = True
return form
@staticmethod
def _add_form_fields(form: aiohttp.FormData, fields: dict) -> None:
for key, value in fields.items():
if value is None:
continue
if isinstance(value, bool):
value = "true" if value else "false"
elif isinstance(value, (dict, list)):
value = json.dumps(value, ensure_ascii=False)
form.add_field(key, str(value))
# ── 图像转换工具 ──────────────────────────────────────────────────────────
# ── 请求体大小限制 ────────────────────────────────────────────────────────
_MAX_BODY_BYTES = 20 * 1024 * 1024 # 20 MB
_UNIFIED_BODY_LIMIT_BYTES = 18 * 1024 * 1024
_SMART_RESIZE_MIN_LONG_EDGE = 256
_SMART_RESIZE_MAX_ATTEMPTS = 32
@staticmethod
def _shrink_png_to_limit(png_bytes: bytes, max_bytes: int, label: str = "") -> bytes:
"""
若 PNG bytes 超过 max_bytes,按等比缩放反复压缩直到满足限制。
每次将面积缩小至约 80%(线性尺寸缩小至约 89.4%)。
"""
if len(png_bytes) <= max_bytes:
return png_bytes
img = Image.open(BytesIO(png_bytes))
w, h = img.size
original_size = len(png_bytes)
step = 0
while len(png_bytes) > max_bytes:
scale = 0.894 # sqrt(0.8),面积缩小 20%
w = max(1, int(w * scale))
h = max(1, int(h * scale))
img = img.resize((w, h), Image.LANCZOS)
buf = BytesIO()
img.save(buf, format="PNG")
png_bytes = buf.getvalue()
step += 1
tag = f" ({label})" if label else ""
print(
f"[o1key GPT Image] 图像{tag}超出 {max_bytes // (1024*1024)}MB 限制,"
f"已等比缩放 {step} 次:{original_size // 1024}KB → {len(png_bytes) // 1024}KB "
f"{w}×{h}"
)
return png_bytes
@staticmethod
def _png_bytes_to_data_url(png_bytes: bytes) -> str:
b64 = base64.b64encode(png_bytes).decode("ascii")
return f"data:image/png;base64,{b64}"
@staticmethod
def _json_body_size(body: dict) -> int:
return len(json.dumps(body, ensure_ascii=False, separators=(",", ":")).encode("utf-8"))
@staticmethod
def _format_body_size(size: int) -> str:
if size < 1024 * 1024:
return f"{size / 1024:.2f}KB"
return f"{size / 1024 / 1024:.2f}MB"
@staticmethod
def _shorten_data_urls_for_log(obj, max_len: int = 200):
if isinstance(obj, dict):
return {
key: GptImageClient._shorten_data_urls_for_log(value, max_len)
for key, value in obj.items()
}
if isinstance(obj, list):
return [GptImageClient._shorten_data_urls_for_log(item, max_len) for item in obj]
if isinstance(obj, str) and len(obj) > max_len:
header, separator, data = obj.partition(",")
if separator and header.lower().startswith("data:") and ";base64" in header.lower():
return f"{header},<base64 data, {len(data)} chars>"
candidate = obj.strip()
if len(candidate) % 4 == 0:
try:
base64.b64decode(candidate, validate=True)
except Exception:
pass
else:
return f"<base64 data, {len(candidate)} chars>"
return obj
def _log_original_request_body(self, label: str, body: dict) -> None:
if not REQUEST_LOG_ENABLED:
return
body_size = self._json_body_size(body)
print(
f"\n{'=' * 60}\n"
f"[o1key GPT Image] 原始请求体日志 | {label} | "
f"请求体积: {self._format_body_size(body_size)} "
f"(data URL 已折叠显示 base64 长度)\n"
f"{json.dumps(self._shorten_data_urls_for_log(body), ensure_ascii=False, indent=2)}\n"
f"{'=' * 60}\n"
)
def _log_original_response_body(self, label: str, text: str) -> None:
if not self.response_log_enabled:
return
size = len(text.encode("utf-8"))
try:
safe_body = json.dumps(
self._shorten_data_urls_for_log(json.loads(text)),
ensure_ascii=False,
indent=2,
)
except Exception:
safe_body = f"<unparseable response body, {size} bytes; content omitted>"
print(
f"\n{'=' * 60}\n"
f"[o1key GPT Image] 原始返回响应体日志 | {label} | "
f"响应体积: {self._format_body_size(size)} (base64 已折叠)\n"
f"{safe_body}\n"
f"{'=' * 60}\n"
)
@staticmethod
def _extract_error_message(payload_or_text, status_code: int = 0) -> str:
payload = payload_or_text
if isinstance(payload_or_text, str):
try:
payload = json.loads(payload_or_text)
except Exception:
return get_friendly_message(status_code, payload_or_text)
if isinstance(payload, dict):
error = payload.get("error")
if isinstance(error, str) and error.strip():
return get_friendly_message(0, error.strip())
if isinstance(error, dict):
msg = error.get("message") or error.get("msg") or error.get("error")
if msg:
return get_friendly_message(0, str(msg))
return json.dumps(error, ensure_ascii=False)
msg = payload.get("message") or payload.get("msg")
if msg:
return get_friendly_message(0, str(msg))
return get_friendly_message(status_code, str(payload_or_text))
@staticmethod
def _extract_error_detail(payload: dict) -> dict:
if not isinstance(payload, dict):
return {}
detail = payload.get("error_detail")
return detail if isinstance(detail, dict) else {}
@staticmethod
def _coerce_progress_percent(value: Any) -> Optional[int]:
if value is None:
return None
if isinstance(value, str):
text = value.strip().rstrip("%")
if not text:
return None
try:
value = float(text)
except ValueError:
return None
elif isinstance(value, (int, float)):
value = float(value)
else:
return None
if 0 <= value <= 1:
value *= 100
return max(0, min(100, int(round(value))))
@staticmethod
def _emit_progress(progress_callback: Optional[Callable[[int], None]], pct: int) -> None:
if progress_callback is None:
return
try:
progress_callback(pct)
except Exception as error:
print(f"[o1key GPT Image] progress callback failed: {error}")
@staticmethod
def _resize_png_bytes(source_image: Image.Image, scale: float) -> bytes:
if scale < 0.999:
width, height = source_image.size
new_width = max(1, round(width * scale))
new_height = max(1, round(height * scale))
image = source_image.resize((new_width, new_height), Image.Resampling.LANCZOS)
else:
image = source_image
buf = BytesIO()
image.save(buf, format="PNG")
return buf.getvalue()
@staticmethod
def _pil_to_png_bytes(image: Image.Image) -> bytes:
buf = BytesIO()
image.save(buf, format="PNG")
return buf.getvalue()
def _fit_png_assets_to_body_limit(
self,
assets: List[Dict[str, Any]],
build_body,
*,
resize_mode: Optional[str] = None,
) -> Dict[str, bytes]:
"""Fit exact JSON body size while preserving unified-node resize semantics.
``resize_mode=None`` preserves the legacy GPT client behavior: automatically
fit to the provider's 20 MiB limit. The unified image node passes an
explicit mode and therefore uses the same 18 MiB safety ceiling as Nano
Banana. Smart mode resizes the largest encoded asset group first and
always resamples from the original pixels. A reference image and its mask
share a group so their dimensions remain aligned.
"""
if resize_mode not in {None, "不缩放", "智能缩放"}:
raise ValueError("缩放图片参数无效")
body_limit = (
self._MAX_BODY_BYTES
if resize_mode is None
else self._UNIFIED_BODY_LIMIT_BYTES
)
effective_mode = "智能缩放" if resize_mode is None else resize_mode
asset_bytes = {asset["key"]: asset["bytes"] for asset in assets}
initial_size = self._json_body_size(build_body(asset_bytes))
if initial_size <= body_limit:
return asset_bytes
limit_mib = body_limit / (1024 * 1024)
if effective_mode == "不缩放":
raise ValueError(
f"GPT Image: 请求体 {initial_size / (1024 * 1024):.2f} MiB 超过 "
f"{limit_mib:.0f} MiB 上限;当前设置为“不缩放”,"
"请在上游缩小图片或改用“智能缩放”"
)
if not assets:
raise ValueError(
f"GPT Image: 请求体 {initial_size / (1024 * 1024):.2f} MiB 超过 "
f"{limit_mib:.0f} MiB 上限,且没有可缩放图片"
)
originals: Dict[str, Image.Image] = {}
groups: Dict[str, List[str]] = {}
labels: Dict[str, str] = {}
try:
for asset in assets:
with Image.open(BytesIO(asset["bytes"])) as image:
image.load()
originals[asset["key"]] = image.copy()
group = str(asset.get("resize_group") or asset["key"])
groups.setdefault(group, []).append(asset["key"])
labels[asset["key"]] = str(asset.get("label") or asset["key"])
scales = {group: 1.0 for group in groups}
final_sizes = {key: image.size for key, image in originals.items()}
body_size = initial_size
for _attempt in range(self._SMART_RESIZE_MAX_ATTEMPTS):
if body_size <= body_limit:
changes = [
f"{labels[key]} {originals[key].width}×{originals[key].height}"
f"→{final_sizes[key][0]}×{final_sizes[key][1]}"
for key in originals
if final_sizes[key] != originals[key].size
]
print(
f"[o1key GPT Image] 智能缩放完成 | "
f"{initial_size / (1024 * 1024):.2f}→"
f"{body_size / (1024 * 1024):.2f} MiB | "
+ "".join(changes)
)
return asset_bytes
candidates = []
for group, keys in groups.items():
min_scale = max(
min(1.0, self._SMART_RESIZE_MIN_LONG_EDGE / max(originals[key].size))
for key in keys
)
current_scale = scales[group]
if current_scale <= min_scale + 1e-6:
continue
encoded_length = sum(
4 * ((len(asset_bytes[key]) + 2) // 3)
for key in keys
)
candidates.append((encoded_length, group, min_scale))
if not candidates:
break
current_data_length, group, min_scale = max(candidates)
overflow = body_size - body_limit
desired_data_length = max(1024, current_data_length - overflow - 32 * 1024)
estimated_area_ratio = min(0.90, desired_data_length / current_data_length)
proposed_scale = scales[group] * math.sqrt(max(0.01, estimated_area_ratio)) * 0.98
new_scale = max(min_scale, min(scales[group] * 0.92, proposed_scale))
candidate_bytes = dict(asset_bytes)
candidate_sizes = dict(final_sizes)
for key in groups[group]:
candidate_bytes[key] = self._resize_png_bytes(originals[key], new_scale)
with Image.open(BytesIO(candidate_bytes[key])) as candidate_image:
candidate_sizes[key] = candidate_image.size
if all(candidate_sizes[key] == final_sizes[key] for key in groups[group]):
new_scale = max(min_scale, scales[group] * 0.85)
for key in groups[group]:
candidate_bytes[key] = self._resize_png_bytes(originals[key], new_scale)
with Image.open(BytesIO(candidate_bytes[key])) as candidate_image:
candidate_sizes[key] = candidate_image.size
asset_bytes = candidate_bytes
final_sizes = candidate_sizes
scales[group] = new_scale
body_size = self._json_body_size(build_body(asset_bytes))
final_size = self._json_body_size(build_body(asset_bytes))
raise ValueError(
f"GPT Image: 智能缩放后请求体仍为 {final_size / (1024 * 1024):.2f} MiB"
f"无法安全降至 {limit_mib:.0f} MiB;请减少参考图或在上游进一步缩小"
)
finally:
for image in originals.values():
image.close()
@staticmethod
def _tensor_to_png_bytes(tensor: torch.Tensor) -> bytes:
"""
单张 ComfyUI IMAGE tensor [1, H, W, C] 或 [H, W, C] → PNG bytes
"""
if tensor.dim() == 4:
tensor = tensor.squeeze(0) # [H, W, C]
arr = (tensor.cpu().numpy() * 255).clip(0, 255).astype(np.uint8)
img = Image.fromarray(arr)
buf = BytesIO()
img.save(buf, format="PNG")
return buf.getvalue()
@staticmethod
def _mask_tensor_to_rgba_png_bytes(mask: torch.Tensor, image_size: tuple) -> bytes:
"""
ComfyUI MASK tensor [1, H, W] 或 [H, W] → RGBA PNG bytes
白色区域(mask=1)→ 透明(alpha=0),即 API 将在此处生成新内容。
"""
if mask.dim() == 3:
mask = mask.squeeze(0) # [H, W]
h, w = mask.shape
ih, iw = image_size
# 尺寸不一致时给出提示(API 侧也会报错)
if (h, w) != (ih, iw):
raise ValueError(
f"蒙版尺寸 ({h}×{w}) 与图像尺寸 ({ih}×{iw}) 不一致,请保持相同尺寸"
)
alpha = ((1.0 - mask.cpu().numpy()) * 255).clip(0, 255).astype(np.uint8)
rgba = np.zeros((h, w, 4), dtype=np.uint8)
rgba[:, :, 3] = alpha # 只设 alphaRGB 全 0
buf = BytesIO()
Image.fromarray(rgba, mode="RGBA").save(buf, format="PNG")
return buf.getvalue()
@staticmethod
def _pil_list_to_tensor(images: List[Image.Image]) -> torch.Tensor:
"""
PIL Image 列表 → ComfyUI IMAGE tensor [B, H, W, C],值域 [0, 1]
RGBA 自动转换为 RGBA(保留透明通道)
"""
if not images:
placeholder = Image.new("RGBA", (512, 512), (128, 128, 128, 255))
images = [placeholder]
tensors = []
source_metadata = []
for img in images:
source_metadata.append({
"format": getattr(img, "_o1key_original_format", None) or img.format,
"bytes": getattr(img, "_o1key_original_bytes", None),
"modified": False,
})
arr = np.array(img.convert("RGBA")).astype(np.float32) / 255.0
tensors.append(torch.from_numpy(arr))
# 批量模式下 API 可能返回不同尺寸,统一 resize 到最大尺寸
max_h = max(t.shape[0] for t in tensors)
max_w = max(t.shape[1] for t in tensors)
aligned = []
for index, t in enumerate(tensors):
if t.shape[0] != max_h or t.shape[1] != max_w:
source_metadata[index]["modified"] = True
t = t.permute(2, 0, 1).unsqueeze(0) # [1, C, H, W]
t = torch.nn.functional.interpolate(
t, size=(max_h, max_w), mode="bilinear", align_corners=False
)
t = t.squeeze(0).permute(1, 2, 0) # [H, W, C]
aligned.append(t)
result = torch.stack(aligned, dim=0) # [B, H, W, 4]
result._o1key_source_metadata = source_metadata
return result
# ── 响应解析(通用) ─────────────────────────────────────────────────────
async def _parse_response(
self,
resp_json: dict,
session: aiohttp.ClientSession,
) -> List[Image.Image]:
"""
解析 data 列表,优先取 b64_json,回退到 url 下载
"""
if "error" in resp_json:
err = resp_json["error"]
msg = (
err.get("message") or err.get("msg") or json.dumps(err, ensure_ascii=False)
if isinstance(err, dict)
else str(err)
)
raise RuntimeError(f"API 返回错误: {msg}")
data_list = resp_json.get("data")
if data_list is None:
data_list = resp_json.get("images")
if data_list is None and (resp_json.get("b64_json") or resp_json.get("url")):
data_list = [resp_json]
if not data_list:
raise RuntimeError(
f"API 响应中未找到 data 字段,完整响应:\n"
f"{json.dumps(resp_json, ensure_ascii=False, indent=2)}"
)
images: List[Image.Image] = []
for idx, item in enumerate(data_list):
b64 = item.get("b64_json", "")
url = item.get("url", "")
if b64:
# 优先 base64(无需二次下载)
try:
img_bytes = base64.b64decode(b64)
img = Image.open(BytesIO(img_bytes))
img.load()
_attach_original_image_info(img, img_bytes)
images.append(img)
print(f"[o1key GPT Image] 第 {idx + 1} 张 base64 解码完成 "
f"({img.size[0]}×{img.size[1]})")
except Exception as e:
raise RuntimeError(f"第 {idx + 1} 张 base64 解码失败: {e}")
elif url and url.startswith("http"):
# 回退:下载 URL
img, byte_count, elapsed = await self._download_image_with_response_retry(
session, url, f"第 {idx + 1} 张图片"
)
images.append(img)
print(
f"[o1key GPT Image] 下载图像 {idx + 1}/{len(data_list)} 完成 | "
f"{img.size[0]}×{img.size[1]} | {self._format_transfer_size(byte_count)} | "
f"{elapsed:.2f}s | {self._format_transfer_rate(byte_count, elapsed)}"
)
else:
print(f"[o1key GPT Image] 警告:第 {idx + 1} 条数据既无 b64_json 也无 url,已跳过")
return images
@staticmethod
def _decode_b64_image(b64: str, label: str) -> Image.Image:
try:
img_bytes = base64.b64decode(b64)
img = Image.open(BytesIO(img_bytes))
img.load()
_attach_original_image_info(img, img_bytes)
print(f"[o1key GPT Image] {label} base64 解码完成 ({img.size[0]}×{img.size[1]})")
return img
except Exception as e:
raise RuntimeError(f"{label} base64 解码失败: {e}") from None
@staticmethod
def _decode_inline_result_image(b64: str, label: str) -> tuple[Image.Image, int]:
"""Strictly decode and fully load a completed task's inline image."""
try:
if not isinstance(b64, str) or not b64:
raise ValueError("empty base64 payload")
image_bytes = base64.b64decode(b64, validate=True)
if not image_bytes:
raise ValueError("decoded image is empty")
image = Image.open(BytesIO(image_bytes))
image.load()
_attach_original_image_info(image, image_bytes)
return image, len(image_bytes)
except (binascii.Error, ValueError, OSError, UnidentifiedImageError) as exc:
raise _IncompleteInlineImageError(f"{label}不完整或无法解码: {exc}") from None
async def _append_images_from_payload(
self,
payload: dict,
session: aiohttp.ClientSession,
images: List[Image.Image],
event_name: str = "",
) -> bool:
if isinstance(payload, dict) and "error" in payload:
err = payload["error"]
msg = (
err.get("message") or err.get("msg") or json.dumps(err, ensure_ascii=False)
if isinstance(err, dict)
else str(err)
)
raise RuntimeError(get_friendly_message(500, msg)) from None
if not isinstance(payload, dict):
return False
event_type = payload.get("type") or event_name
if "partial_image" in event_type:
return False
for key in ("data", "images"):
data_list = payload.get(key)
if isinstance(data_list, list):
parsed = await self._parse_response({"data": data_list}, session)
images.extend(parsed)
return True
if payload.get("b64_json") or payload.get("url"):
parsed = await self._parse_response({"data": [payload]}, session)
images.extend(parsed)
return True
image_obj = payload.get("image")
if isinstance(image_obj, dict) and (image_obj.get("b64_json") or image_obj.get("url")):
parsed = await self._parse_response({"data": [image_obj]}, session)
images.extend(parsed)
return True
return False
async def _parse_edit_stream_response(
self,
resp: aiohttp.ClientResponse,
session: aiohttp.ClientSession,
) -> List[Image.Image]:
"""
Parse /v1/images/edits SSE events and return final completed images.
Partial images are intentionally ignored so the node output stays unchanged.
"""
images: List[Image.Image] = []
buffer = ""
event_name = ""
data_lines = []
partial_count = 0
debug_body_parts = []
async def _handle_event():
nonlocal event_name, data_lines, partial_count, images
if not data_lines:
event_name = ""
return
data_str = "\n".join(data_lines).strip()
event_name = event_name.strip()
data_lines = []
if not data_str or data_str == "[DONE]":
return
try:
payload = json.loads(data_str)
except Exception:
raise RuntimeError(get_friendly_message(500, data_str)) from None
if isinstance(payload, dict):
event_type = payload.get("type") or event_name
else:
event_type = event_name
if "partial_image" in event_type:
partial_count += 1
return
await self._append_images_from_payload(payload, session, images, event_name)
async for raw_chunk in resp.content.iter_any():
chunk_text = raw_chunk.decode("utf-8", errors="ignore")
debug_body_parts.append(chunk_text)
buffer += chunk_text
while "\n" in buffer:
line, buffer = buffer.split("\n", 1)
line = line.rstrip("\r")
if line == "":
await _handle_event()
event_name = ""
continue
if line.startswith(":"):
continue
if line.startswith("event:"):
event_name = line[len("event:"):].strip()
elif line.startswith("data:"):
data_lines.append(line[len("data:"):].lstrip())
if buffer.strip():
data_lines.append(buffer.strip())
await _handle_event()
if partial_count:
print(f"[o1key GPT Image] 流式中间图 {partial_count} 张(已忽略,仅输出最终图)")
if not images:
raise RuntimeError("流式响应结束,但未收到最终图片")
return images
async def _parse_stream_text_response(
self,
text: str,
session: aiohttp.ClientSession,
) -> List[Image.Image]:
images: List[Image.Image] = []
partial_count = 0
event_name = ""
data_lines = []
async def _handle_event():
nonlocal event_name, data_lines, partial_count, images
if not data_lines:
event_name = ""
return
data_str = "\n".join(data_lines).strip()
event_name = event_name.strip()
data_lines = []
if not data_str or data_str == "[DONE]":
return
payload = json.loads(data_str)
if isinstance(payload, dict):
event_type = payload.get("type") or event_name
else:
event_type = event_name
if "partial_image" in event_type:
partial_count += 1
return
await self._append_images_from_payload(payload, session, images, event_name)
for raw_line in text.splitlines():
line = raw_line.rstrip("\r")
if line == "":
await _handle_event()
event_name = ""
continue
if line.startswith(":"):
continue
if line.startswith("event:"):
event_name = line[len("event:"):].strip()
elif line.startswith("data:"):
data_lines.append(line[len("data:"):].lstrip())
await _handle_event()
if partial_count:
print(f"[o1key GPT Image] 流式中间图 {partial_count} 张(已忽略,仅输出最终图)")
if not images:
raise RuntimeError("流式响应结束,但未收到最终图片")
return images
async def _parse_success_response(
self,
resp: aiohttp.ClientResponse,
session: aiohttp.ClientSession,
label: str = "",
) -> List[Image.Image]:
content_type = resp.headers.get("Content-Type", "").lower()
if "event-stream" in content_type:
return await self._parse_edit_stream_response(resp, session)
text = await resp.text()
self._log_original_response_body(f"{label or 'sync'} status={resp.status}", text)
stripped = text.lstrip()
if stripped.startswith("data:") or stripped.startswith("event:"):
return await self._parse_stream_text_response(text, session)
try:
resp_json = json.loads(text)
except Exception:
raise RuntimeError(f"响应 JSON 解析失败,原始内容:{text[:500]}")
return await self._parse_response(resp_json, session)
# ── 中断轮询 ──────────────────────────────────────────────────────────────
@staticmethod
async def _poll_interrupt():
"""每 0.5s 轮询一次 ComfyUI 中断标志"""
while True:
await asyncio.sleep(0.5)
if _INTERRUPT_AVAILABLE and processing_interrupted():
return
@staticmethod
async def _run_with_interrupt(coro):
"""
将异步任务与中断轮询并发执行。
如果用户点击取消,cancel 掉 coro 并抛出 InterruptProcessingException。
"""
if not _INTERRUPT_AVAILABLE:
return await coro
request_task = asyncio.ensure_future(coro)
interrupt_task = asyncio.ensure_future(GptImageClient._poll_interrupt())
done, pending = await asyncio.wait(
[request_task, interrupt_task],
return_when=asyncio.FIRST_COMPLETED,
)
for t in pending:
t.cancel()
try:
await t
except (asyncio.CancelledError, Exception):
pass
if interrupt_task in done and request_task not in done:
raise InterruptProcessingException()
return request_task.result()
# ── 新版异步 GPT Image 接口 ──────────────────────────────────────────────
def _build_async_generate_body(
self,
prompt: str,
model: str,
quality: str,
size: Optional[str],
n: int,
image_list: Optional[List[torch.Tensor]] = None,
mask_tensor: Optional[torch.Tensor] = None,
output_format: str = "png",
background: str = "auto",
moderation: Optional[str] = None,
resize_mode: Optional[str] = None,
) -> dict:
if int(n) != 1:
raise ValueError("GPT Image 单次请求的 n 必须为 1;多图请使用外层并发")
output_format = str(output_format or "png").strip().lower()
background = str(background or "auto").strip().lower()
moderation = str(moderation).strip().lower() if moderation is not None else None
if output_format not in GPT_IMAGE_OUTPUT_FORMAT_OPTIONS:
raise ValueError("GPT Image 输出格式无效")
if background not in GPT_IMAGE_BACKGROUND_OPTIONS:
raise ValueError("GPT Image 背景参数无效")
if moderation not in {None, "low"}:
raise ValueError("GPT Image 内容审查强度无效")
if background == "transparent" and output_format == "jpeg":
raise ValueError("GPT Image 透明背景仅支持 PNG 或 WebP 输出格式")
api_model = _MODEL_NAME_MAP.get(model, model)
assets: List[Dict[str, Any]] = []
image_keys: List[str] = []
mask_key = None
if image_list:
for idx_img, img_tensor in enumerate(image_list):
pil_images = tensor_to_pil(img_tensor)
if not pil_images:
continue
key = f"image_{idx_img}"
image_keys.append(key)
assets.append({
"key": key,
"label": f"参考图{idx_img + 1}",
"bytes": self._pil_to_png_bytes(pil_images[0]),
"resize_group": key,
})
if mask_tensor is not None:
if not image_list:
raise ValueError("提供了蒙版但未提供图片,请同时提供图片和蒙版")
first_tensor = image_list[0]
if first_tensor.dim() == 3:
first_tensor = first_tensor.unsqueeze(0)
image_size = (first_tensor.shape[1], first_tensor.shape[2])
mask_key = "mask"
assets.append({
"key": mask_key,
"label": "蒙版",
"bytes": self._mask_tensor_to_rgba_png_bytes(mask_tensor, image_size),
"resize_group": image_keys[0],
})
def _make_body(asset_bytes: Dict[str, bytes]) -> dict:
body = {
"model": api_model,
"prompt": prompt,
"images": [
self._png_bytes_to_data_url(asset_bytes[key])
for key in image_keys
if key in asset_bytes
],
"quality": quality,
"n": int(n),
"output_format": output_format,
"background": background,
}
if size:
body["size"] = size
if mask_key and mask_key in asset_bytes:
body["mask"] = {
"image_url": self._png_bytes_to_data_url(asset_bytes[mask_key])
}
if moderation == "low":
body["moderation"] = "low"
return body
original_asset_bytes = {asset["key"]: asset["bytes"] for asset in assets}
original_body = _make_body(original_asset_bytes)
self._log_original_request_body("async generateImage", original_body)
asset_bytes = self._fit_png_assets_to_body_limit(
assets,
_make_body,
resize_mode=resize_mode,
)
body = _make_body(asset_bytes)
body_size = self._json_body_size(body)
body_limit = (
self._MAX_BODY_BYTES
if resize_mode is None
else self._UNIFIED_BODY_LIMIT_BYTES
)
if body_size > body_limit:
raise RuntimeError(
f"请求体超过 {body_limit / (1024 * 1024):.0f} MiB"
f"{body_size / (1024 * 1024):.2f} MiB"
)
return body
async def _submit_generate_image_task(
self,
session: aiohttp.ClientSession,
payload: dict,
task_submitted_callback: Optional[Callable[[str, str, float], None]] = None,
) -> str:
url = f"{self.base_url}{_ENDPOINT_ASYNC_GENERATE}"
t0 = time.time()
async with session.post(
url,
json=payload,
headers=self._json_headers(),
) as resp:
elapsed = time.time() - t0
text = await resp.text()
self._log_original_response_body(
f"submit generateImage status={resp.status}",
text,
)
if resp.status not in (200, 201, 202):
raise RuntimeError(self._extract_error_message(text, resp.status))
try:
data = json.loads(text)
except Exception:
raise RuntimeError(f"提交响应 JSON 解析失败,原始内容:{text[:500]}") from None
task_id = data.get("task_id")
if not task_id:
raise RuntimeError(f"提交响应中未找到 task_id: {data}")
status = data.get("status", "")
if task_submitted_callback:
task_submitted_callback(task_id, status, elapsed)
else:
print(f"[o1key GPT Image] 异步任务已提交 | task_id={task_id} | status={status} | 耗时 {elapsed:.1f}s")
return task_id
async def _poll_generate_image_task(
self,
session: aiohttp.ClientSession,
task_id: str,
progress_callback: Optional[Callable[[int], None]] = None,
initial_delay: bool = True,
) -> dict:
url = f"{self.base_url}{_ENDPOINT_ASYNC_TASK.format(task_id=task_id)}"
start_time = time.time()
poll_count = 0
last_poll_at = start_time
while True:
if poll_count == 0 and not initial_delay:
next_poll_at = start_time
elif poll_count < len(_ASYNC_POLL_SCHEDULE):
next_poll_at = start_time + _ASYNC_POLL_SCHEDULE[poll_count]
else:
next_poll_at = last_poll_at + _ASYNC_POLL_INTERVAL
sleep_time = next_poll_at - time.time()
if sleep_time > 0:
await asyncio.sleep(sleep_time)
last_poll_at = time.time()
elapsed = last_poll_at - start_time
if elapsed > _ASYNC_MAX_WAIT:
raise RuntimeError(f"任务 {task_id} 超时(>{int(_ASYNC_MAX_WAIT)}秒),请稍后用 task_id 查询结果")
poll_count += 1
task = await self._get_json_with_response_retry(
session,
url,
self._auth_headers(),
label=f"任务 {task_id} 查询",
expected_task_id=task_id,
)
status = task.get("status", "UNKNOWN")
progress = task.get("progress")
progress_pct = self._coerce_progress_percent(progress)
if self.poll_log_enabled:
progress_text = f" | progress={progress}" if progress is not None else ""
print(f"[o1key GPT Image] 查询任务 #{poll_count} | task_id={task_id} | status={status}{progress_text}")
if status == "SUCCESS":
self._emit_progress(progress_callback, 100)
return task
if progress_pct is not None and progress_pct < 100:
self._emit_progress(progress_callback, progress_pct)
if status == "FAILURE":
error_message = self._extract_error_message(task, 500) or "生成失败"
detail = self._extract_error_detail(task)
code = detail.get("code")
category = detail.get("category")
parts = [p for p in (
f"code={code}" if code else "",
f"category={category}" if category else "",
f"task_id={task_id}",
) if p]
raise RuntimeError(f"{error_message} ({', '.join(parts)})")
if status not in ("SUBMITTED", "IN_PROGRESS"):
raise RuntimeError(f"未知任务状态 {status}: {task}")
async def _parse_async_task_images(
self,
task: dict,
session: aiohttp.ClientSession,
result_url_callback: Optional[Callable[[str], None]] = None,
log_downloads: bool = True,
log_prefix: str = "[o1key GPT Image]",
) -> List[Image.Image]:
data = task.get("data", {})
image_items = data.get("images") if isinstance(data, dict) else None
if not isinstance(image_items, list) or not image_items:
raise RuntimeError(f"任务结果中未找到 data.images: {task}")
images: List[Image.Image] = []
result_urls: List[str] = []
try:
for idx, item in enumerate(image_items, 1):
if not isinstance(item, dict):
continue
url = item.get("url") or item.get("image_url")
b64 = item.get("b64_json", "")
if url and isinstance(url, str) and url.startswith("data:image"):
header, separator, b64_data = url.partition(",")
if not separator or ";base64" not in header.lower():
raise _IncompleteInlineImageError(
f"第 {idx} 张内联 data URL 缺少有效 Base64 头"
)
started = time.perf_counter()
img, byte_count = self._decode_inline_result_image(
b64_data,
f"第 {idx} 张内联 data URL 图片",
)
images.append(img)
if log_downloads:
print(
f"{log_prefix} 结果图像 {idx}/{len(image_items)} 内联 data URL,无网络下载 | "
f"{img.size[0]}×{img.size[1]} | {self._format_transfer_size(byte_count)} | "
f"解码={time.perf_counter() - started:.2f}s"
)
elif url and isinstance(url, str) and url.startswith("http"):
if log_downloads:
print(f"{log_prefix} 下载图像 {idx}/{len(image_items)} 开始")
img, byte_count, elapsed = await self._download_image_with_response_retry(
session, url, f"第 {idx} 张图片"
)
images.append(img)
result_urls.append(url)
if log_downloads:
print(
f"{log_prefix} 下载图像 {idx}/{len(image_items)} 完成 | "
f"{img.size[0]}×{img.size[1]} | {self._format_transfer_size(byte_count)} | "
f"{elapsed:.2f}s | {self._format_transfer_rate(byte_count, elapsed)}"
)
elif b64:
started = time.perf_counter()
img, byte_count = self._decode_inline_result_image(
b64,
f"第 {idx} 张内联 Base64 图片",
)
images.append(img)
if log_downloads:
print(
f"{log_prefix} 结果图像 {idx}/{len(image_items)} 内联 base64,无网络下载 | "
f"{img.size[0]}×{img.size[1]} | {self._format_transfer_size(byte_count)} | "
f"解码={time.perf_counter() - started:.2f}s"
)
else:
print(f"[o1key GPT Image] 警告:第 {idx} 条结果既无 url 也无 b64_json,已跳过")
except Exception:
for image in images:
image.close()
raise
if not images:
raise RuntimeError("任务成功但没有可用图片结果")
if result_url_callback:
for url in dict.fromkeys(result_urls):
result_url_callback(url)
return images
@staticmethod
def _inline_result_retry_delay(task_id: str, attempt: int) -> float:
base_delay = GptImageClient._response_retry_delay(attempt)
stable_jitter = (sum(task_id.encode("utf-8")) % 500) / 1000.0
return base_delay + stable_jitter
async def _parse_completed_task_images_with_retry(
self,
task: dict,
session: aiohttp.ClientSession,
task_id: str,
progress_callback: Optional[Callable[[int], None]] = None,
result_url_callback: Optional[Callable[[str], None]] = None,
log_downloads: bool = True,
log_prefix: str = "[o1key GPT Image]",
) -> tuple[dict, List[Image.Image]]:
"""Re-fetch a completed task when its inline image is incomplete."""
current_task = task
for attempt in range(_RESPONSE_READ_MAX_RETRIES + 1):
try:
images = await self._parse_async_task_images(
current_task,
session,
result_url_callback=result_url_callback,
log_downloads=log_downloads,
log_prefix=log_prefix,
)
return current_task, images
except _IncompleteInlineImageError as exc:
if attempt >= _RESPONSE_READ_MAX_RETRIES:
raise RuntimeError(
f"GPT Image: task_id={task_id} 的内联结果在 "
f"{attempt + 1} 次获取后仍不完整: {exc}"
) from None
delay = self._inline_result_retry_delay(task_id, attempt)
print(
f"[o1key GPT Image] 内联结果不完整,{delay:.1f}s 后重新查询同一任务 "
f"| task_id={task_id} | retry={attempt + 1}/{_RESPONSE_READ_MAX_RETRIES}"
)
await asyncio.sleep(delay)
current_task = await self._poll_generate_image_task(
session,
task_id,
progress_callback=progress_callback,
initial_delay=False,
)
raise RuntimeError(f"GPT Image: task_id={task_id} 的内联结果读取失败") # pragma: no cover
async def _generate_image_task_async(
self,
prompt: str,
model: str,
quality: str,
size: Optional[str],
n: int,
seed: int,
image_tensor: Optional[List[torch.Tensor]] = None,
mask_tensor: Optional[torch.Tensor] = None,
output_format: str = "png",
background: str = "auto",
moderation: Optional[str] = None,
progress_callback: Optional[Callable[[int], None]] = None,
result_url_callback: Optional[Callable[[str], None]] = None,
log_downloads: bool = True,
log_prefix: str = "[o1key GPT Image]",
task_submitted_callback: Optional[Callable[[str, str, float], None]] = None,
log_request_start: bool = True,
task_completed_callback: Optional[Callable[[str, int, float, List[str]], None]] = None,
resize_mode: Optional[str] = None,
) -> List[Image.Image]:
body = self._build_async_generate_body(
prompt=prompt,
model=model,
quality=quality,
size=size,
n=n,
image_list=image_tensor,
mask_tensor=mask_tensor,
output_format=output_format,
background=background,
moderation=moderation,
resize_mode=resize_mode,
)
mode = "图像编辑" if mask_tensor is not None else ("图生图" if image_tensor else "文生图")
body_size = self._json_body_size(body)
if log_request_start:
print(
f"[o1key GPT Image] {mode} | 新异步接口 | 模型={model} | "
f"quality={quality} | size={size} | n={n} | body={body_size // 1024}KB"
)
connector = aiohttp.TCPConnector(ssl=False, force_close=True)
timeout = aiohttp.ClientTimeout(total=_ASYNC_MAX_WAIT + 120)
task_started = time.time()
async with aiohttp.ClientSession(connector=connector, timeout=timeout) as session:
task_id = await self._submit_generate_image_task(
session, body, task_submitted_callback=task_submitted_callback,
)
task = await self._poll_generate_image_task(
session,
task_id,
progress_callback=progress_callback,
)
upstream_urls: List[str] = []
def _on_result_url(url: str) -> None:
upstream_urls.append(url)
if result_url_callback:
result_url_callback(url)
task, images = await self._parse_completed_task_images_with_retry(
task,
session,
task_id,
progress_callback=progress_callback,
result_url_callback=_on_result_url,
log_downloads=log_downloads,
log_prefix=log_prefix,
)
if task_completed_callback:
task_completed_callback(task_id, len(images), time.time() - task_started, upstream_urls)
return images
async def _generate_special_price_images_parallel(
self,
prompt: str,
model: str,
quality: str,
size: Optional[str],
n: int,
seed: int,
image_tensor: Optional[List[torch.Tensor]] = None,
mask_tensor: Optional[torch.Tensor] = None,
output_format: str = "png",
background: str = "auto",
moderation: Optional[str] = None,
progress_callback: Optional[Callable[[int], None]] = None,
result_url_callback: Optional[Callable[[str], None]] = None,
log_downloads: bool = True,
log_prefix: str = "[o1key GPT Image]",
task_submitted_callback: Optional[Callable[[str, str, float], None]] = None,
log_request_start: bool = True,
task_completed_callback: Optional[Callable[[str, int, float, List[str]], None]] = None,
resize_mode: Optional[str] = None,
) -> List[Image.Image]:
"""多图生成:并发提交 n 个单图任务,每个请求固定传 n=1。"""
request_count = max(1, int(n))
progress_values = [0] * request_count
def _task_progress_callback(task_index: int):
if progress_callback is None:
return None
def _update(value: int):
try:
progress_values[task_index] = max(0, min(100, int(value)))
except (TypeError, ValueError):
return
progress_callback(int(sum(progress_values) / request_count))
return _update
if log_request_start:
print(
f"[o1key GPT Image] 多图并发模式 | 并发请求={request_count} | 每个请求 n=1"
)
tasks = [
asyncio.create_task(
self._generate_image_task_async(
prompt=prompt,
model=model,
quality=quality,
size=size,
n=1,
seed=seed,
image_tensor=image_tensor,
mask_tensor=mask_tensor,
output_format=output_format,
background=background,
moderation=moderation,
progress_callback=_task_progress_callback(task_index),
result_url_callback=result_url_callback,
log_downloads=log_downloads,
log_prefix=log_prefix,
task_submitted_callback=task_submitted_callback,
log_request_start=log_request_start,
task_completed_callback=task_completed_callback,
resize_mode=resize_mode,
)
)
for task_index in range(request_count)
]
results = await asyncio.gather(*tasks, return_exceptions=True)
images: List[Image.Image] = []
errors = []
for task_index, result in enumerate(results, 1):
if isinstance(result, BaseException):
errors.append(f"第{task_index}个请求: {str(result).splitlines()[0]}")
else:
images.extend(result)
if progress_callback is not None:
progress_callback(100)
if errors:
print(
f"[o1key GPT Image] 多图并发完成 | 成功={request_count - len(errors)} "
f"| 失败={len(errors)} | {'; '.join(errors)}"
)
if not images:
detail = "; ".join(errors) or "未返回图片"
raise RuntimeError(f"多图并发请求全部失败:{detail}")
return images
# ── 文生图 / 图生图(generations 接口)───────────────────────────────────
async def _generate_async(
self,
prompt: str,
model: str,
quality: str,
size: str,
n: int,
seed: int,
image_list: Optional[List[torch.Tensor]] = None,
) -> List[Image.Image]:
"""
调用 /v1/images/generations/ 接口。
当传入 image_list 时,以 data URI 格式内联图像(图生图)。
"""
# 模型名映射:UI 显示名 → API 参数名
api_model = _MODEL_NAME_MAP.get(model, model)
body: dict = {
"model": api_model,
"prompt": prompt,
"quality": quality,
"n": n,
"moderation": "low",
"partial_images": 0,
}
body["size"] = size if size else "auto"
image_files = []
# 图生图:multipart 方式上传参考图
if image_list is not None:
for idx_img, img_tensor in enumerate(image_list):
pil_images = tensor_to_pil(img_tensor)
img = pil_images[0]
buf = BytesIO()
img.save(buf, format="PNG")
png_bytes = buf.getvalue()
# 单张图像预算:20MB 按图数平摊,至少保留 1MB 给其他字段
per_image_budget = max(
1024 * 1024,
(self._MAX_BODY_BYTES - 1024 * 1024) // len(image_list),
)
# base64 膨胀约 4/3,所以 PNG 目标上限 = budget * 3/4
png_budget = int(per_image_budget * 3 / 4)
label = f"第{idx_img + 1}张" if len(image_list) > 1 else ""
png_bytes = self._shrink_png_to_limit(png_bytes, png_budget, label)
image_files.append(png_bytes)
mode = f"图生图(参考图 {len(image_files)} 张)"
else:
mode = "文生图"
url = f"{self.base_url}{_ENDPOINT_GENERATIONS}"
print(f"[o1key GPT Image] {mode} | 模型={model} | quality={quality} | "
f"size={size} | n={n}")
connector = aiohttp.TCPConnector(ssl=False, force_close=True)
timeout = aiohttp.ClientTimeout(total=_REQUEST_TIMEOUT)
def _build_multipart_form() -> aiohttp.FormData:
form = self._new_multipart_form()
self._add_form_fields(form, body)
image_field = "image[]" if len(image_files) > 1 else "image"
for idx_img, png_bytes in enumerate(image_files):
form.add_field(
image_field,
png_bytes,
filename=f"image_{idx_img + 1}.png",
content_type="image/png",
)
return form
async def _do_request():
async with aiohttp.ClientSession(connector=connector, timeout=timeout) as session:
t0 = time.time()
async with session.post(
url,
data=_build_multipart_form(),
headers=self._auth_headers(),
) as resp:
elapsed = time.time() - t0
if resp.status != 200:
text = await resp.text()
if resp.status in HTTP_ERROR_MESSAGES:
raise RuntimeError(HTTP_ERROR_MESSAGES[resp.status])
try:
err_json = json.loads(text)
err_obj = err_json.get("error", {})
msg = (
err_obj.get("message") or err_obj.get("msg") or text
if isinstance(err_obj, dict)
else str(err_obj) or text
)
except Exception:
msg = text
raise RuntimeError(get_friendly_message(resp.status, msg))
print(f"[o1key GPT Image] API 响应耗时 {elapsed:.1f}s")
return await self._parse_success_response(resp, session, "GENERATIONS")
return await self._run_with_interrupt(_do_request())
# ── 图像编辑(edits 接口,multipart/form-data)──────────────────────
async def _edit_async(
self,
prompt: str,
model: str,
quality: str,
size: str,
n: int,
seed: int,
image_list: List[torch.Tensor],
mask_tensor: Optional[torch.Tensor] = None,
) -> List[Image.Image]:
"""
调用 /v1/images/edits 接口(multipart/form-data)。
"""
# 模型名映射:UI 显示名 → API 参数名
api_model = _MODEL_NAME_MAP.get(model, model)
# 统一 tensors 为 [1,H,W,C] 格式,支持不同尺寸
normalized_tensors = []
for t in image_list:
if t.dim() == 3:
t = t.unsqueeze(0) # [H,W,C] → [1,H,W,C]
normalized_tensors.append(t)
num_images = len(normalized_tensors)
image_files = []
# 多图:用 multipart image/image[] 字段逐张上传
# 预算:20MB 按图数平摊,蒙版预留 1MB
mask_reserve = 1024 * 1024 if mask_tensor is not None else 0
per_image_budget = max(
1024 * 1024,
(self._MAX_BODY_BYTES - mask_reserve) // num_images,
)
for i, frame in enumerate(normalized_tensors):
img_bytes = self._tensor_to_png_bytes(frame)
label = f"第{i + 1}张" if num_images > 1 else ""
img_bytes = self._shrink_png_to_limit(img_bytes, per_image_budget, label)
image_files.append(img_bytes)
# 蒙版尺寸校验以第一张图为基准
first_tensor = normalized_tensors[0]
ih, iw = first_tensor.shape[1], first_tensor.shape[2]
mask_png = None
if mask_tensor is not None:
mask_png = self._mask_tensor_to_rgba_png_bytes(mask_tensor, (ih, iw))
mode = "图像编辑(带蒙版)"
else:
mode = "图像编辑(无蒙版)"
form_fields = {
"model": api_model,
"prompt": prompt,
"partial_images": 0,
"n": n,
"quality": quality,
"size": size if size else "auto",
"output_format": "png",
"background": "opaque",
"moderation": "low",
}
def _build_multipart_form() -> aiohttp.FormData:
form = self._new_multipart_form()
self._add_form_fields(form, form_fields)
for idx_img, img_bytes in enumerate(image_files):
form.add_field(
"image[]",
img_bytes,
filename=f"image_{idx_img + 1}.png",
content_type="image/png",
)
if mask_png is not None:
form.add_field(
"mask",
mask_png,
filename="mask.png",
content_type="image/png",
)
return form
url = f"{self.base_url}{_ENDPOINT_EDITS}"
print(f"[o1key GPT Image] {mode} | 模型={model} | 参考图={num_images}张 | "
f"quality={quality} | size={size} | n={n}")
connector = aiohttp.TCPConnector(ssl=False, force_close=True)
timeout = aiohttp.ClientTimeout(total=_REQUEST_TIMEOUT)
async def _do_request():
async with aiohttp.ClientSession(connector=connector, timeout=timeout) as session:
t0 = time.time()
async with session.post(
url,
data=_build_multipart_form(),
headers=self._auth_headers(),
) as resp:
elapsed = time.time() - t0
if resp.status != 200:
text = await resp.text()
if resp.status in HTTP_ERROR_MESSAGES:
raise RuntimeError(HTTP_ERROR_MESSAGES[resp.status])
try:
err_json = json.loads(text)
err_obj = err_json.get("error", {})
msg = (
err_obj.get("message") or err_obj.get("msg") or text
if isinstance(err_obj, dict)
else str(err_obj) or text
)
except Exception:
msg = text
raise RuntimeError(get_friendly_message(resp.status, msg))
print(f"[o1key GPT Image] API 响应耗时 {elapsed:.1f}s")
return await self._parse_success_response(resp, session, "EDITS")
return await self._run_with_interrupt(_do_request())
# ── 同步统一入口(供节点调用)────────────────────────────────────────────
def run_sync(
self,
prompt: str,
model: str,
quality: str,
size: str,
n: int,
seed: int,
image_tensor: Optional[List[torch.Tensor]] = None,
mask_tensor: Optional[torch.Tensor] = None,
) -> List[Image.Image]:
"""
同步入口,在独立线程中运行事件循环,避免与 ComfyUI 主循环冲突。
路由逻辑:
- 无 image_tensor → generations 接口(文生图,multipart/form-data
- 有 image_tensor → edits 接口(图生图/编辑,multipart/form-data
"""
use_edits = (image_tensor is not None)
if use_edits:
coro = self._edit_async(
prompt=prompt, model=model, quality=quality,
size=size, n=n, seed=seed,
image_list=image_tensor, mask_tensor=mask_tensor,
)
else:
coro = self._generate_async(
prompt=prompt, model=model, quality=quality,
size=size, n=n, seed=seed,
image_list=image_tensor,
)
def _run():
loop = asyncio.new_event_loop()
asyncio.set_event_loop(loop)
try:
return loop.run_until_complete(coro)
finally:
loop.close()
with ThreadPoolExecutor(max_workers=1) as executor:
future = executor.submit(_run)
try:
return future.result(timeout=_REQUEST_TIMEOUT + 30)
except TimeoutError:
raise RuntimeError(
f"o1key GPT Image 请求超时(>{_REQUEST_TIMEOUT}s),请检查网络或稍后重试"
)
async def generate_image_async(
self,
prompt: str,
model: str,
quality: str,
size: Optional[str],
n: int,
seed: int,
image_tensor: Optional[List[torch.Tensor]] = None,
mask_tensor: Optional[torch.Tensor] = None,
output_format: str = "png",
background: str = "auto",
moderation: Optional[str] = None,
progress_callback: Optional[Callable[[int], None]] = None,
special_price_parallel: bool = False,
result_url_callback: Optional[Callable[[str], None]] = None,
log_downloads: bool = True,
log_prefix: str = "[o1key GPT Image]",
task_submitted_callback: Optional[Callable[[str, str, float], None]] = None,
log_request_start: bool = True,
task_completed_callback: Optional[Callable[[str, int, float, List[str]], None]] = None,
resize_mode: Optional[str] = None,
) -> List[Image.Image]:
"""异步公共入口;多图始终拆成并发的 n=1 单图请求。"""
requested_n = max(1, int(n))
if requested_n > 1:
return await self._generate_special_price_images_parallel(
prompt=prompt,
model=model,
quality=quality,
size=size,
n=requested_n,
seed=seed,
image_tensor=image_tensor,
mask_tensor=mask_tensor,
output_format=output_format,
background=background,
moderation=moderation,
progress_callback=progress_callback,
result_url_callback=result_url_callback,
log_downloads=log_downloads,
log_prefix=log_prefix,
task_submitted_callback=task_submitted_callback,
log_request_start=log_request_start,
task_completed_callback=task_completed_callback,
resize_mode=resize_mode,
)
return await self._generate_image_task_async(
prompt=prompt,
model=model,
quality=quality,
size=size,
n=requested_n,
seed=seed,
image_tensor=image_tensor,
mask_tensor=mask_tensor,
output_format=output_format,
background=background,
moderation=moderation,
progress_callback=progress_callback,
result_url_callback=result_url_callback,
log_downloads=log_downloads,
log_prefix=log_prefix,
task_submitted_callback=task_submitted_callback,
log_request_start=log_request_start,
task_completed_callback=task_completed_callback,
resize_mode=resize_mode,
)
def generate_image_async_sync(
self,
prompt: str,
model: str,
quality: str,
size: Optional[str],
n: int,
seed: int,
image_tensor: Optional[List[torch.Tensor]] = None,
mask_tensor: Optional[torch.Tensor] = None,
output_format: str = "png",
background: str = "auto",
moderation: Optional[str] = None,
progress_callback: Optional[Callable[[int], None]] = None,
special_price_parallel: bool = False,
result_url_callback: Optional[Callable[[str], None]] = None,
log_downloads: bool = True,
log_prefix: str = "[o1key GPT Image]",
task_submitted_callback: Optional[Callable[[str, str, float], None]] = None,
log_request_start: bool = True,
task_completed_callback: Optional[Callable[[str, int, float, List[str]], None]] = None,
resize_mode: Optional[str] = None,
) -> List[Image.Image]:
"""
新版异步任务入口,供节点调用。
旧 run_sync 保留兼容,但 GPT Image 节点不再使用旧同步接口。
"""
coro = self.generate_image_async(
prompt=prompt,
model=model,
quality=quality,
size=size,
n=n,
seed=seed,
image_tensor=image_tensor,
mask_tensor=mask_tensor,
output_format=output_format,
background=background,
moderation=moderation,
progress_callback=progress_callback,
special_price_parallel=special_price_parallel,
result_url_callback=result_url_callback,
log_downloads=log_downloads,
log_prefix=log_prefix,
task_submitted_callback=task_submitted_callback,
log_request_start=log_request_start,
task_completed_callback=task_completed_callback,
resize_mode=resize_mode,
)
def _run():
loop = asyncio.new_event_loop()
asyncio.set_event_loop(loop)
try:
return loop.run_until_complete(self._run_with_interrupt(coro))
finally:
loop.close()
with ThreadPoolExecutor(max_workers=1) as executor:
future = executor.submit(_run)
try:
return future.result(timeout=_ASYNC_MAX_WAIT + 150)
except TimeoutError:
raise RuntimeError(
f"o1key GPT Image 异步任务超时(>{int(_ASYNC_MAX_WAIT)}秒),请稍后用 task_id 查询结果"
)
# ── 余额查询 ──────────────────────────────────────────────────────────────
async def _query_balance_async(self) -> dict:
url = f"{self.base_url}/api/usage/token"
connector = aiohttp.TCPConnector(ssl=False, force_close=True)
timeout = aiohttp.ClientTimeout(total=10)
async with aiohttp.ClientSession(connector=connector, timeout=timeout) as session:
async with session.get(url, headers=self._auth_headers()) as resp:
if resp.status != 200:
raise RuntimeError(f"余额查询失败 HTTP {resp.status}")
return await resp.json()
def query_balance_sync(self) -> dict:
def _run():
loop = asyncio.new_event_loop()
asyncio.set_event_loop(loop)
try:
return loop.run_until_complete(self._query_balance_async())
finally:
loop.close()
with ThreadPoolExecutor(max_workers=1) as executor:
return executor.submit(_run).result(timeout=15)
@staticmethod
def format_balance_info(balance_data: dict) -> str:
data = balance_data.get("data", {})
api_name = data.get("name", "未知")
total_available = data.get("total_available", 0)
balance_in_dollars = total_available / 500000
return f"当前余额:{balance_in_dollars:.2f} | API{api_name}"