fix(image-gen): address review findings in OpenAI image adapter

Scope download Authorization header to the provider origin, translate
httpx errors in _download/edit into safe ImageGenError messages, and
prevent entry params from clobbering core payload keys; also harden
strip_data_uri against malformed values and drop dead code.

Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
This commit is contained in:
m4
2026-07-23 21:44:57 +08:00
parent 12e4f34005
commit cc9dfb1cc9
3 changed files with 153 additions and 16 deletions
+1 -1
View File
@@ -38,7 +38,7 @@ class ImageGenAdapter(Protocol):
def strip_data_uri(value: str) -> str:
if value.startswith("data:image/"):
if value.startswith("data:image/") and "," in value:
return value.split(",", 1)[1]
return value
+43 -15
View File
@@ -6,6 +6,7 @@ import asyncio
import mimetypes
from pathlib import Path
from typing import Any
from urllib.parse import urlparse
import httpx
@@ -101,15 +102,33 @@ class OpenAIImageAdapter:
return resp.json()
assert last_response is not None
self._raise_for_status(last_response)
raise ImageGenError("Image provider request failed")
def _base_origin(self) -> str:
base = (self.entry.resolved_base_url() or "https://api.openai.com").rstrip("/")
parsed = urlparse(base)
return f"{parsed.scheme}://{parsed.netloc}"
async def _download(self, url: str) -> bytes:
headers = {"Authorization": f"Bearer {self._api_key()}"}
if self._client is not None:
resp = await self._client.get(url, headers=headers)
else:
async with httpx.AsyncClient(timeout=self._timeout) as client:
resp = await client.get(url, headers=headers)
# Only send the API key to the provider's own origin; provider-returned
# URLs on other hosts (e.g. pre-signed blob storage) must not leak it.
headers: dict[str, str] = {}
parsed = urlparse(url)
if f"{parsed.scheme}://{parsed.netloc}" == self._base_origin():
headers["Authorization"] = f"Bearer {self._api_key()}"
try:
if self._client is not None:
resp = await self._client.get(url, headers=headers)
else:
async with httpx.AsyncClient(timeout=self._timeout) as client:
resp = await client.get(url, headers=headers)
except httpx.TimeoutException:
raise ImageGenError(
f"Image provider request timed out after {self._timeout:g}s"
) from None
except httpx.RequestError as exc:
raise ImageGenError(
f"Image provider request failed: {exc.__class__.__name__}"
) from None
self._raise_for_status(resp)
content_type = resp.headers.get("content-type", "")
if content_type and not content_type.startswith("image/"):
@@ -131,13 +150,13 @@ class OpenAIImageAdapter:
) -> list[bytes]:
base = self._endpoint_base()
payload: dict[str, Any] = {
**self.entry.params,
"model": self.entry.id,
"prompt": prompt,
"size": size,
"quality": quality,
"n": n,
"response_format": "b64_json",
**self.entry.params,
}
if background and background != "auto":
payload["background"] = background
@@ -190,14 +209,23 @@ class OpenAIImageAdapter:
with mask.open("rb") as mask_fh:
files.append(("mask", (mask.name, mask_fh.read(), mask_mime)))
headers = {"Authorization": f"Bearer {self._api_key()}"}
if self._client is not None:
resp = await self._client.post(
f"{base}/images/edits", headers=headers, data=data, files=files
)
else:
async with httpx.AsyncClient(timeout=self._timeout) as client:
resp = await client.post(
try:
if self._client is not None:
resp = await self._client.post(
f"{base}/images/edits", headers=headers, data=data, files=files
)
else:
async with httpx.AsyncClient(timeout=self._timeout) as client:
resp = await client.post(
f"{base}/images/edits", headers=headers, data=data, files=files
)
except httpx.TimeoutException:
raise ImageGenError(
f"Image provider request timed out after {self._timeout:g}s"
) from None
except httpx.RequestError as exc:
raise ImageGenError(
f"Image provider request failed: {exc.__class__.__name__}"
) from None
self._raise_for_status(resp)
return await self._extract(resp.json())
+109
View File
@@ -142,3 +142,112 @@ async def test_edit_posts_multipart(tmp_path):
image=image, mask=None, prompt="make it blue", size="1024x1024", quality="auto"
)
assert images == [PNG_BYTES]
async def test_download_to_foreign_host_omits_auth_header():
def handler(request: httpx.Request) -> httpx.Response:
if request.url.path == "/v1/images/generations":
return httpx.Response(
200, json={"data": [{"url": "https://files.example/img.png"}]}
)
if request.url.host == "files.example":
assert "Authorization" not in request.headers
return httpx.Response(
200, content=PNG_BYTES, headers={"content-type": "image/png"}
)
return httpx.Response(404)
adapter = _adapter(handler)
images = await adapter.generate(
prompt="a cat", size="1024x1024", quality="auto", background="auto", n=1
)
assert images == [PNG_BYTES]
async def test_download_to_same_host_sends_auth_header():
def handler(request: httpx.Request) -> httpx.Response:
if request.url.path == "/v1/images/generations":
return httpx.Response(
200,
json={"data": [{"url": "https://api.example.com/files/img.png"}]},
)
if request.url.path == "/files/img.png":
assert request.headers["Authorization"] == "Bearer sk-test"
return httpx.Response(
200, content=PNG_BYTES, headers={"content-type": "image/png"}
)
return httpx.Response(404)
adapter = _adapter(handler, base_url="https://api.example.com")
images = await adapter.generate(
prompt="a cat", size="1024x1024", quality="auto", background="auto", n=1
)
assert images == [PNG_BYTES]
async def test_download_network_error_wrapped():
def handler(request: httpx.Request) -> httpx.Response:
if request.url.path == "/v1/images/generations":
return httpx.Response(
200, json={"data": [{"url": "https://files.example/img.png"}]}
)
raise httpx.ConnectError("boom", request=request)
adapter = _adapter(handler)
with pytest.raises(ImageGenError) as excinfo:
await adapter.generate(
prompt="a cat", size="1024x1024", quality="auto", background="auto", n=1
)
message = str(excinfo.value)
assert "https://" not in message
assert "files.example" not in message
async def test_edit_network_error_wrapped(tmp_path):
image = tmp_path / "in.png"
image.write_bytes(PNG_BYTES)
def handler(request: httpx.Request) -> httpx.Response:
raise httpx.ConnectError("boom", request=request)
adapter = _adapter(handler)
with pytest.raises(ImageGenError) as excinfo:
await adapter.edit(
image=image,
mask=None,
prompt="make it blue",
size="1024x1024",
quality="auto",
)
assert "https://" not in str(excinfo.value)
async def test_params_cannot_clobber_core_keys():
import json
def handler(request: httpx.Request) -> httpx.Response:
payload = json.loads(request.content)
assert payload["model"] == "gpt-image-2"
assert payload["response_format"] == "b64_json"
return httpx.Response(200, json={"data": [{"b64_json": B64}]})
adapter = _adapter(
handler, params={"model": "evil", "response_format": "url"}
)
images = await adapter.generate(
prompt="a cat", size="1024x1024", quality="auto", background="auto", n=1
)
assert images == [PNG_BYTES]
async def test_malformed_data_uri_ignored():
def handler(request: httpx.Request) -> httpx.Response:
return httpx.Response(
200, json={"data": [{"b64_json": "data:image/png"}, {"b64_json": B64}]}
)
adapter = _adapter(handler)
images = await adapter.generate(
prompt="a cat", size="1024x1024", quality="auto", background="auto", n=1
)
assert PNG_BYTES in images