Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
60 changes: 56 additions & 4 deletions plugins/image_gen/openai-codex/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -21,6 +21,7 @@

import json
import logging
from pathlib import Path
from typing import Any, Dict, List, Optional, Tuple

from agent.image_gen_provider import (
Expand Down Expand Up @@ -143,16 +144,65 @@ def _read_codex_access_token() -> Optional[str]:
return None


def _build_responses_payload(*, prompt: str, size: str, quality: str) -> Dict[str, Any]:
def _normalize_reference_images(reference_images: Any) -> Tuple[List[Dict[str, str]], List[str]]:
"""Convert user-supplied image refs into Responses ``input_image`` parts.

Accepts remote URLs, data URLs, or local file paths. Local files are encoded
as data URLs so the Codex Responses backend can consume them directly.
Returns ``(parts, invalid_refs)``.
"""
if reference_images in (None, "", []):
return [], []

refs = reference_images if isinstance(reference_images, list) else [reference_images]
parts: List[Dict[str, str]] = []
invalid: List[str] = []

for raw_ref in refs:
if not isinstance(raw_ref, str):
invalid.append(str(raw_ref))
continue
ref = raw_ref.strip()
if not ref:
invalid.append(str(raw_ref))
continue
if ref.startswith(("http://", "https://", "data:image/")):
parts.append({"type": "input_image", "image_url": ref})
continue
path = Path(ref).expanduser()
if path.exists() and path.is_file():
try:
from agent.image_routing import _file_to_data_url

data_url = _file_to_data_url(path)
except Exception as exc:
logger.debug("Could not encode reference image %s: %s", path, exc)
data_url = None
if data_url:
parts.append({"type": "input_image", "image_url": data_url})
continue
invalid.append(ref)

return parts, invalid


def _build_responses_payload(*, prompt: str, size: str, quality: str, reference_images: Any = None) -> Dict[str, Any]:
"""Build the Codex Responses request body for an image_generation call."""
image_parts, invalid_refs = _normalize_reference_images(reference_images)
if invalid_refs:
raise ValueError(
"Invalid reference_images entries: " + ", ".join(invalid_refs)
)
content_parts: List[Dict[str, str]] = [{"type": "input_text", "text": prompt}]
content_parts.extend(image_parts)
return {
"model": _CODEX_CHAT_MODEL,
"store": False,
"instructions": _CODEX_INSTRUCTIONS,
"input": [{
"type": "message",
"role": "user",
"content": [{"type": "input_text", "text": prompt}],
"content": content_parts,
}],
"tools": [{
"type": "image_generation",
Expand Down Expand Up @@ -242,7 +292,7 @@ def flush():
yield payload


def _collect_image_b64(token: str, *, prompt: str, size: str, quality: str) -> Optional[str]:
def _collect_image_b64(token: str, *, prompt: str, size: str, quality: str, reference_images: Any = None) -> Optional[str]:
"""Stream a Codex Responses image_generation call and return the b64 image."""
import httpx
from agent.auxiliary_client import _codex_cloudflare_headers
Expand All @@ -253,7 +303,7 @@ def _collect_image_b64(token: str, *, prompt: str, size: str, quality: str) -> O
"Authorization": f"Bearer {token}",
"Content-Type": "application/json",
})
payload = _build_responses_payload(prompt=prompt, size=size, quality=quality)
payload = _build_responses_payload(prompt=prompt, size=size, quality=quality, reference_images=reference_images)
timeout = httpx.Timeout(300.0, connect=30.0, read=300.0, write=30.0, pool=30.0)

image_b64: Optional[str] = None
Expand Down Expand Up @@ -335,6 +385,7 @@ def generate(
) -> Dict[str, Any]:
prompt = (prompt or "").strip()
aspect = resolve_aspect_ratio(aspect_ratio)
reference_images = kwargs.get("reference_images")

if not prompt:
return error_response(
Expand Down Expand Up @@ -388,6 +439,7 @@ def generate(
prompt=prompt,
size=size,
quality=meta["quality"],
reference_images=reference_images,
)
except Exception as exc:
logger.debug("Codex image generation failed", exc_info=True)
Expand Down
26 changes: 25 additions & 1 deletion tests/plugins/image_gen/test_openai_codex_provider.py
Original file line number Diff line number Diff line change
Expand Up @@ -129,7 +129,7 @@ def test_codex_stream_request_shape(self, provider, monkeypatch):

captured = {}

def _collect(token, *, prompt, size, quality):
def _collect(token, *, prompt, size, quality, reference_images=None):
captured.update(codex_plugin._build_responses_payload(
prompt=prompt,
size=size,
Expand Down Expand Up @@ -160,6 +160,30 @@ def _collect(token, *, prompt, size, quality):
assert tool["background"] == "opaque"
assert tool["partial_images"] == 1

def test_payload_includes_reference_image_url(self):
payload = codex_plugin._build_responses_payload(
prompt="swap only the face",
size="1024x1024",
quality="medium",
reference_images=["https://example.com/source.png"],
)
content = payload["input"][0]["content"]
assert content[0] == {"type": "input_text", "text": "swap only the face"}
assert content[1] == {"type": "input_image", "image_url": "https://example.com/source.png"}

def test_payload_includes_local_reference_image_as_data_url(self, tmp_path):
img_path = tmp_path / "ref.png"
img_path.write_bytes(bytes.fromhex(_PNG_HEX))
payload = codex_plugin._build_responses_payload(
prompt="swap only the face",
size="1024x1024",
quality="medium",
reference_images=[str(img_path)],
)
content = payload["input"][0]["content"]
assert content[1]["type"] == "input_image"
assert content[1]["image_url"].startswith("data:image/png;base64,")

def test_partial_image_event_used_when_done_missing(self):
"""If output_item.done is missing, partial_image_b64 is accepted."""
payload = {
Expand Down
126 changes: 122 additions & 4 deletions tests/tools/test_image_generation.py
Original file line number Diff line number Diff line change
Expand Up @@ -119,6 +119,19 @@ def test_nano_banana_portrait_uses_aspect_ratio(self, image_tool):
p = image_tool._build_fal_payload("fal-ai/nano-banana-pro", "hello", "portrait")
assert p["aspect_ratio"] == "9:16"

def test_nano_banana_2_landscape_uses_aspect_ratio(self, image_tool):
p = image_tool._build_fal_payload("fal-ai/nano-banana-2", "hello", "landscape")
assert p["aspect_ratio"] == "16:9"
assert "image_size" not in p

def test_nano_banana_2_square_uses_aspect_ratio(self, image_tool):
p = image_tool._build_fal_payload("fal-ai/nano-banana-2", "hello", "square")
assert p["aspect_ratio"] == "1:1"

def test_nano_banana_2_portrait_uses_aspect_ratio(self, image_tool):
p = image_tool._build_fal_payload("fal-ai/nano-banana-2", "hello", "portrait")
assert p["aspect_ratio"] == "9:16"


class TestGptLiteralFamily:
"""GPT-Image 1.5 uses literal size strings."""
Expand Down Expand Up @@ -221,6 +234,11 @@ def test_nano_banana_never_gets_image_size(self, image_tool):
assert "image_size" not in p
assert p["aspect_ratio"] == "16:9"

def test_nano_banana_2_never_gets_image_size(self, image_tool):
p = image_tool._build_fal_payload("fal-ai/nano-banana-2", "hi", "landscape", seed=1)
assert "image_size" not in p
assert p["aspect_ratio"] == "16:9"


# ---------------------------------------------------------------------------
# Default merging
Expand Down Expand Up @@ -363,11 +381,12 @@ def test_empty_aspect_defaults_to_landscape(self, image_tool):

class TestRegistryIntegration:

def test_schema_exposes_only_prompt_and_aspect_ratio_to_agent(self, image_tool):
"""The agent-facing schema must stay tight — model selection is a
user-level config choice, not an agent-level arg."""
def test_schema_exposes_prompt_aspect_ratio_and_reference_images_to_agent(self, image_tool):
"""The agent-facing schema stays tight: the agent can provide the prompt,
output framing, and optional reference images, but not provider/model
selection knobs."""
props = image_tool.IMAGE_GENERATE_SCHEMA["parameters"]["properties"]
assert set(props.keys()) == {"prompt", "aspect_ratio"}
assert set(props.keys()) == {"prompt", "aspect_ratio", "reference_images"}

def test_aspect_ratio_enum_is_three_values(self, image_tool):
enum = image_tool.IMAGE_GENERATE_SCHEMA["parameters"]["properties"]["aspect_ratio"]["enum"]
Expand Down Expand Up @@ -414,6 +433,105 @@ class OddResponse:
assert image_tool._extract_http_status(exc) is None


class TestFalReferenceImages:
def test_normalize_fal_reference_images_accepts_urls_and_data_urls(self, image_tool):
image_urls, invalid = image_tool._normalize_fal_reference_images([
"https://example.com/a.png",
"data:image/png;base64,abc123",
])
assert image_urls == [
"https://example.com/a.png",
"data:image/png;base64,abc123",
]
assert invalid == []

def test_normalize_fal_reference_images_encodes_local_file(self, image_tool, tmp_path, monkeypatch):
img = tmp_path / "ref.png"
img.write_bytes(b"png-bytes")

monkeypatch.setattr(
"agent.image_routing._file_to_data_url",
lambda path: "data:image/png;base64,LOCALDATA",
)

image_urls, invalid = image_tool._normalize_fal_reference_images([str(img)])
assert image_urls == ["data:image/png;base64,LOCALDATA"]
assert invalid == []

def test_reference_images_use_edit_endpoint_for_nano_banana_pro(self, image_tool, monkeypatch):
monkeypatch.setattr(image_tool, "_resolve_managed_fal_gateway", lambda: object())
monkeypatch.setattr(
image_tool,
"_normalize_fal_reference_images",
lambda refs: (["https://example.com/base.png", "https://example.com/face.png"], []),
)

submitted = {}

class _Handler:
def get(self):
return {"images": [{"url": "https://example.com/out.png"}]}

def _fake_submit(model, arguments):
submitted["model"] = model
submitted["arguments"] = arguments
return _Handler()

monkeypatch.setattr(image_tool, "_submit_fal_request", _fake_submit)
monkeypatch.setattr(image_tool, "_resolve_fal_model", lambda: (
"fal-ai/nano-banana-pro",
image_tool.FAL_MODELS["fal-ai/nano-banana-pro"],
))

result = image_tool.image_generate_tool(
prompt="swap face",
aspect_ratio="landscape",
reference_images=["a", "b"],
)
assert "https://example.com/out.png" in result
assert submitted["model"] == "fal-ai/nano-banana-pro/edit"
assert submitted["arguments"]["image_urls"] == [
"https://example.com/base.png",
"https://example.com/face.png",
]
assert submitted["arguments"]["aspect_ratio"] == "16:9"

def test_reference_images_use_edit_endpoint_for_nano_banana_2(self, image_tool, monkeypatch):
monkeypatch.setattr(image_tool, "_resolve_managed_fal_gateway", lambda: object())
monkeypatch.setattr(
image_tool,
"_normalize_fal_reference_images",
lambda refs: (["https://example.com/base.png"], []),
)

submitted = {}

class _Handler:
def get(self):
return {"images": [{"url": "https://example.com/out2.png"}]}

def _fake_submit(model, arguments):
submitted["model"] = model
submitted["arguments"] = arguments
return _Handler()

monkeypatch.setattr(image_tool, "_submit_fal_request", _fake_submit)
monkeypatch.setattr(image_tool, "_resolve_fal_model", lambda: (
"fal-ai/nano-banana-2",
image_tool.FAL_MODELS["fal-ai/nano-banana-2"],
))

result = image_tool.image_generate_tool(
prompt="edit image",
aspect_ratio="square",
reference_images=["a"],
)
assert "https://example.com/out2.png" in result
assert submitted["model"] == "fal-ai/nano-banana-2/edit"
assert submitted["arguments"]["image_urls"] == ["https://example.com/base.png"]
assert submitted["arguments"]["aspect_ratio"] == "1:1"


class TestManagedGatewayErrorTranslation:
"""4xx from the Nous managed gateway should be translated to a user-actionable message."""

Expand Down
42 changes: 42 additions & 0 deletions tests/tools/test_image_generation_tool_reference_images.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,42 @@
from __future__ import annotations

import json
from types import SimpleNamespace


def test_handle_image_generate_passes_reference_images_to_plugin(monkeypatch):
from tools import image_generation_tool

monkeypatch.setattr(
image_generation_tool,
"_read_configured_image_provider",
lambda: "openai-codex",
)
monkeypatch.setattr(
image_generation_tool,
"_read_configured_image_model",
lambda: None,
)

class _Provider:
name = "openai-codex"

def generate(self, **kwargs):
return {
"success": True,
"image": "/tmp/fake.png",
"provider": "openai-codex",
"echo": kwargs,
}

monkeypatch.setitem(__import__("sys").modules, "agent.image_gen_registry", SimpleNamespace(get_provider=lambda name: _Provider()))
monkeypatch.setitem(__import__("sys").modules, "hermes_cli.plugins", SimpleNamespace(_ensure_plugins_discovered=lambda force=False: None))

result = image_generation_tool._handle_image_generate({
"prompt": "edit this image",
"aspect_ratio": "landscape",
"reference_images": ["https://example.com/ref.png"],
})
payload = json.loads(result)
assert payload["success"] is True
assert payload["echo"]["reference_images"] == ["https://example.com/ref.png"]
Loading
Loading