fix openai image reference edits

This commit is contained in:
sbyinin
2026-06-19 17:16:09 +08:00
committed by Xubin Ren
parent dd4d410ce2
commit adb737b614
2 changed files with 215 additions and 25 deletions
+104 -25
View File
@@ -955,6 +955,56 @@ class OpenAIImageGenerationClient(ImageGenerationProvider):
return model.split("/", 1)[1]
return model
async def _parse_images_response(self, payload: dict[str, Any]) -> list[str]:
client = self._client
owns_client = client is None
if owns_client:
client = httpx.AsyncClient(timeout=self.timeout)
try:
return await _openai_images_from_payload(client, payload)
finally:
if owns_client:
await client.aclose()
async def _post_image_edit(
self,
*,
headers: dict[str, str],
body: dict[str, Any],
reference_images: list[str],
) -> httpx.Response:
files: list[tuple[str, tuple[str, Any, str]]] = []
handles: list[Any] = []
try:
for path in reference_images:
p = Path(path)
raw = p.read_bytes()
mime = detect_image_mime(raw)
if mime is None:
raise ImageGenerationError(f"unsupported reference image: {p}")
handle = p.open("rb")
handles.append(handle)
files.append(("image[]", (p.name, handle, mime)))
client = self._client
if client is not None:
return await client.post(
f"{self.api_base}/images/edits",
headers=headers,
data=body,
files=files,
)
async with httpx.AsyncClient(timeout=self.timeout) as c:
return await c.post(
f"{self.api_base}/images/edits",
headers=headers,
data=body,
files=files,
)
finally:
for handle in handles:
handle.close()
async def generate(
self,
*,
@@ -967,21 +1017,18 @@ class OpenAIImageGenerationClient(ImageGenerationProvider):
if not self.api_key:
raise ImageGenerationError(self.missing_key_message)
if reference_images:
logger.warning(
"DALL-E models do not support reference images; "
"ignoring {} reference image(s) for {}",
len(reference_images),
model,
)
clean_model = self._strip_model_prefix(model)
headers = {
generation_headers = {
"Authorization": f"Bearer {self.api_key}",
"Content-Type": "application/json",
**self.extra_headers,
}
edit_headers = {
"Authorization": f"Bearer {self.api_key}",
**self.extra_headers,
}
clean_model = self._strip_model_prefix(model)
body: dict[str, Any] = {
"model": clean_model,
"prompt": prompt,
@@ -999,13 +1046,37 @@ class OpenAIImageGenerationClient(ImageGenerationProvider):
# Drop null-valued params so extraBody can opt out of defaults like response_format.
body = {key: value for key, value in body.items() if value is not None}
logger.info("OpenAI Images API request: POST {}/images/generations body={}", self.api_base, body)
refs = list(reference_images or [])
if refs:
if not _openai_is_gpt_image_model(clean_model):
raise ImageGenerationError(
f"OpenAI model '{clean_model}' does not support reference images; "
"use a GPT Image model"
)
edit_body = _openai_multipart_form_body(body)
logger.info(
"OpenAI Images API request: POST {}/images/edits body={} reference_images={}",
self.api_base,
edit_body,
len(refs),
)
response = await self._post_image_edit(
headers=edit_headers,
body=edit_body,
reference_images=refs,
)
else:
logger.info(
"OpenAI Images API request: POST {}/images/generations body={}",
self.api_base,
body,
)
response = await self._http_post(
f"{self.api_base}/images/generations",
headers=headers,
body=body,
)
response = await self._http_post(
f"{self.api_base}/images/generations",
headers=generation_headers,
body=body,
)
try:
response.raise_for_status()
@@ -1020,16 +1091,7 @@ class OpenAIImageGenerationClient(ImageGenerationProvider):
logger.info("OpenAI Images API response ({}): {}", response.status_code,
{k: v for k, v in payload.items() if k != "data"})
client = self._client
owns_client = client is None
if owns_client:
client = httpx.AsyncClient(timeout=self.timeout)
try:
images = await _openai_images_from_payload(client, payload)
finally:
if owns_client:
await client.aclose()
images = await self._parse_images_response(payload)
self._require_images(images, payload)
return GeneratedImageResponse(images=images, content="", raw=payload)
@@ -1260,6 +1322,23 @@ def _openai_size(
return "1024x1024"
def _openai_multipart_form_body(body: dict[str, Any]) -> dict[str, str]:
form: dict[str, str] = {}
for key, value in body.items():
if value is None:
continue
if isinstance(value, bool):
form[key] = "true" if value else "false"
elif isinstance(value, str | int | float):
form[key] = str(value)
else:
logger.warning(
"OpenAI image edit parameter '{}' is not a scalar form field; ignoring it",
key,
)
return form
def _openai_is_gpt_image_model(model: str) -> bool:
normalized = model.lower()
return normalized.startswith(("gpt-image", "chatgpt-image"))