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:
@@ -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
|
||||
|
||||
|
||||
@@ -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())
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user