From cc9dfb1cc92f58fb1281bac3c4f049c460365daf Mon Sep 17 00:00:00 2001 From: m4 Date: Thu, 23 Jul 2026 21:44:57 +0800 Subject: [PATCH] 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 --- EvoScientist/image_gen/adapters/base.py | 2 +- EvoScientist/image_gen/adapters/openai.py | 58 +++++++++--- tests/test_image_gen_openai_adapter.py | 109 ++++++++++++++++++++++ 3 files changed, 153 insertions(+), 16 deletions(-) diff --git a/EvoScientist/image_gen/adapters/base.py b/EvoScientist/image_gen/adapters/base.py index 6f2db03..838c42b 100644 --- a/EvoScientist/image_gen/adapters/base.py +++ b/EvoScientist/image_gen/adapters/base.py @@ -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 diff --git a/EvoScientist/image_gen/adapters/openai.py b/EvoScientist/image_gen/adapters/openai.py index 4a586a8..607622d 100644 --- a/EvoScientist/image_gen/adapters/openai.py +++ b/EvoScientist/image_gen/adapters/openai.py @@ -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()) diff --git a/tests/test_image_gen_openai_adapter.py b/tests/test_image_gen_openai_adapter.py index 5de5813..64e519b 100644 --- a/tests/test_image_gen_openai_adapter.py +++ b/tests/test_image_gen_openai_adapter.py @@ -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