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
24 changes: 24 additions & 0 deletions agent/image_gen_provider.py
Original file line number Diff line number Diff line change
Expand Up @@ -127,6 +127,30 @@ def default_model(self) -> Optional[str]:
return models[0].get("id")
return None

def supports_image_edit(self) -> bool:
"""Whether this provider supports mother-image editing / image-to-image."""
return False

def edit_image(
self,
image_path: str,
instruction: str,
aspect_ratio: str = DEFAULT_ASPECT_RATIO,
**kwargs: Any,
) -> Dict[str, Any]:
"""Edit an existing image.

Providers that support image editing should override this. The default
implementation returns a uniform unsupported response instead of raising.
"""
return error_response(
error=f"Provider '{self.name}' does not support image editing",
error_type="unsupported_operation",
provider=self.name,
prompt=instruction,
aspect_ratio=resolve_aspect_ratio(aspect_ratio),
)

@abc.abstractmethod
def generate(
self,
Expand Down
192 changes: 189 additions & 3 deletions plugins/image_gen/openai/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -25,8 +25,11 @@

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

from hermes_cli.config import get_env_value

from agent.image_gen_provider import (
DEFAULT_ASPECT_RATIO,
ImageGenProvider,
Expand Down Expand Up @@ -93,6 +96,30 @@ def _load_openai_config() -> Dict[str, Any]:
return {}


def _resolve_timeout(cfg: Optional[Dict[str, Any]] = None) -> float:
"""Resolve OpenAI image timeout in seconds."""
env_value = get_env_value("OPENAI_IMAGE_TIMEOUT") or os.environ.get("OPENAI_IMAGE_TIMEOUT")
if env_value:
try:
return float(env_value)
except ValueError:
logger.warning("Invalid OPENAI_IMAGE_TIMEOUT=%r; ignoring", env_value)

cfg = cfg if isinstance(cfg, dict) else _load_openai_config()
openai_cfg = cfg.get("openai") if isinstance(cfg.get("openai"), dict) else {}
for section in (openai_cfg, cfg):
if isinstance(section, dict):
value = section.get("timeout")
if value is None:
continue
try:
return float(value)
except (TypeError, ValueError):
logger.warning("Invalid image_gen timeout value=%r; ignoring", value)

return 180.0


def _resolve_model() -> Tuple[str, Dict[str, Any]]:
"""Decide which tier to use and return ``(model_id, meta)``."""
env_override = os.environ.get("OPENAI_IMAGE_MODEL")
Expand Down Expand Up @@ -134,7 +161,7 @@ def display_name(self) -> str:
return "OpenAI"

def is_available(self) -> bool:
if not os.environ.get("OPENAI_API_KEY"):
if not get_env_value("OPENAI_API_KEY") and not os.environ.get("OPENAI_API_KEY"):
return False
try:
import openai # noqa: F401
Expand All @@ -157,6 +184,9 @@ def list_models(self) -> List[Dict[str, Any]]:
def default_model(self) -> Optional[str]:
return DEFAULT_MODEL

def supports_image_edit(self) -> bool:
return True

def get_setup_schema(self) -> Dict[str, Any]:
return {
"name": "OpenAI",
Expand Down Expand Up @@ -188,7 +218,7 @@ def generate(
aspect_ratio=aspect,
)

if not os.environ.get("OPENAI_API_KEY"):
if not get_env_value("OPENAI_API_KEY") and not os.environ.get("OPENAI_API_KEY"):
return error_response(
error=(
"OPENAI_API_KEY not set. Run `hermes tools` → Image "
Expand All @@ -211,6 +241,8 @@ def generate(
)

tier_id, meta = _resolve_model()
cfg = _load_openai_config()
timeout_seconds = _resolve_timeout(cfg)
size = _SIZES.get(aspect, _SIZES["square"])

# gpt-image-2 returns b64_json unconditionally and REJECTS
Expand All @@ -224,7 +256,15 @@ def generate(
}

try:
client = openai.OpenAI()
api_key = get_env_value("OPENAI_API_KEY") or os.environ.get("OPENAI_API_KEY")
base_url = get_env_value("OPENAI_BASE_URL") or os.environ.get("OPENAI_BASE_URL")
client_kwargs: Dict[str, Any] = {
"api_key": api_key,
"timeout": timeout_seconds,
}
if base_url:
client_kwargs["base_url"] = base_url
client = openai.OpenAI(**client_kwargs)
response = client.images.generate(**payload)
except Exception as exc:
logger.debug("OpenAI image generation failed", exc_info=True)
Expand Down Expand Up @@ -305,6 +345,152 @@ def generate(
extra=extra,
)

def edit_image(
self,
image_path: str,
instruction: str,
aspect_ratio: str = DEFAULT_ASPECT_RATIO,
**kwargs: Any,
) -> Dict[str, Any]:
instruction = (instruction or "").strip()
aspect = resolve_aspect_ratio(aspect_ratio)

if not instruction:
return error_response(
error="Instruction is required and must be a non-empty string",
error_type="invalid_argument",
provider="openai",
aspect_ratio=aspect,
)

source = Path(image_path).expanduser()
if not source.is_absolute() or not source.exists() or not source.is_file():
return error_response(
error=f"Source image does not exist or is not a file: {source}",
error_type="invalid_argument",
provider="openai",
prompt=instruction,
aspect_ratio=aspect,
)

if not get_env_value("OPENAI_API_KEY") and not os.environ.get("OPENAI_API_KEY"):
return error_response(
error=(
"OPENAI_API_KEY not set. Run `hermes tools` → Image "
"Generation → OpenAI to configure, or `hermes setup` "
"to add the key."
),
error_type="auth_required",
provider="openai",
prompt=instruction,
aspect_ratio=aspect,
)

try:
import openai
except ImportError:
return error_response(
error="openai Python package not installed (pip install openai)",
error_type="missing_dependency",
provider="openai",
prompt=instruction,
aspect_ratio=aspect,
)

tier_id, meta = _resolve_model()
cfg = _load_openai_config()
timeout_seconds = _resolve_timeout(cfg)
size = _SIZES.get(aspect, _SIZES["square"])

try:
api_key = get_env_value("OPENAI_API_KEY") or os.environ.get("OPENAI_API_KEY")
base_url = get_env_value("OPENAI_BASE_URL") or os.environ.get("OPENAI_BASE_URL")
client_kwargs: Dict[str, Any] = {
"api_key": api_key,
"timeout": timeout_seconds,
}
if base_url:
client_kwargs["base_url"] = base_url
client = openai.OpenAI(**client_kwargs)

with source.open("rb") as image_file:
response = client.images.edit(
model=API_MODEL,
image=image_file,
prompt=instruction,
size=size,
quality=meta["quality"],
)
except Exception as exc:
logger.debug("OpenAI image edit failed", exc_info=True)
return error_response(
error=f"OpenAI image edit failed: {exc}",
error_type="api_error",
provider="openai",
model=tier_id,
prompt=instruction,
aspect_ratio=aspect,
)

data = getattr(response, "data", None) or []
if not data:
return error_response(
error="OpenAI returned no image data",
error_type="empty_response",
provider="openai",
model=tier_id,
prompt=instruction,
aspect_ratio=aspect,
)

first = data[0]
b64 = getattr(first, "b64_json", None)
url = getattr(first, "url", None)
revised_prompt = getattr(first, "revised_prompt", None)

if b64:
try:
saved_path = save_b64_image(b64, prefix=f"openai_edit_{tier_id}")
except Exception as exc:
return error_response(
error=f"Could not save edited image to cache: {exc}",
error_type="io_error",
provider="openai",
model=tier_id,
prompt=instruction,
aspect_ratio=aspect,
)
image_ref = str(saved_path)
elif url:
image_ref = url
else:
return error_response(
error="OpenAI edit response contained neither b64_json nor URL",
error_type="empty_response",
provider="openai",
model=tier_id,
prompt=instruction,
aspect_ratio=aspect,
)

extra: Dict[str, Any] = {
"size": size,
"quality": meta["quality"],
"source_image": str(source),
"instruction": instruction,
}
if revised_prompt:
extra["revised_prompt"] = revised_prompt

return success_response(
image=image_ref,
model=tier_id,
prompt=instruction,
aspect_ratio=aspect,
provider="openai",
extra=extra,
)


# ---------------------------------------------------------------------------
# Plugin entry point
Expand Down
75 changes: 75 additions & 0 deletions tests/plugins/image_gen/test_openai_provider.py
Original file line number Diff line number Diff line change
Expand Up @@ -38,6 +38,7 @@ def _tmp_hermes_home(tmp_path, monkeypatch):
@pytest.fixture
def provider(monkeypatch):
monkeypatch.setenv("OPENAI_API_KEY", "test-key")
monkeypatch.setattr(openai_plugin, "get_env_value", lambda key: "test-key" if key == "OPENAI_API_KEY" else None)
return openai_plugin.OpenAIImageGenProvider()


Expand Down Expand Up @@ -74,10 +75,12 @@ def test_catalog_entries_have_display_speed_strengths(self, provider):
class TestAvailability:
def test_no_api_key_unavailable(self, monkeypatch):
monkeypatch.delenv("OPENAI_API_KEY", raising=False)
monkeypatch.setattr(openai_plugin, "get_env_value", lambda key: None)
assert openai_plugin.OpenAIImageGenProvider().is_available() is False

def test_api_key_set_available(self, monkeypatch):
monkeypatch.setenv("OPENAI_API_KEY", "test")
monkeypatch.setattr(openai_plugin, "get_env_value", lambda key: "test" if key == "OPENAI_API_KEY" else None)
assert openai_plugin.OpenAIImageGenProvider().is_available() is True


Expand Down Expand Up @@ -270,3 +273,75 @@ def test_url_response_falls_back_to_bare_url_when_download_fails(self, provider)

assert result["success"] is True
assert result["image"] == "https://example.com/img.png"


class TestImageEdit:
def test_supports_image_edit(self, provider):
assert provider.supports_image_edit() is True

def test_edit_image_uses_get_env_value_for_auth(self, monkeypatch, tmp_path):
monkeypatch.delenv("OPENAI_API_KEY", raising=False)
monkeypatch.setattr(openai_plugin, "get_env_value", lambda key: "test-key" if key == "OPENAI_API_KEY" else None)
provider = openai_plugin.OpenAIImageGenProvider()
image = tmp_path / "source.png"
image.write_bytes(bytes.fromhex(_PNG_HEX))

fake_client = MagicMock()
fake_client.images.edit.return_value = _fake_response(b64=_b64_png())

with _patched_openai(fake_client):
result = provider.edit_image(str(image), "Edit it")

assert result["success"] is True

def test_edit_image_rejects_missing_file(self, provider, tmp_path):
result = provider.edit_image(
str(tmp_path / "missing.png"),
"Turn this into a half-body work portrait.",
)
assert result["success"] is False
assert result["error_type"] == "invalid_argument"

def test_edit_image_calls_openai_images_edit(self, provider, tmp_path):
image = tmp_path / "source.png"
image.write_bytes(bytes.fromhex(_PNG_HEX))

fake_client = MagicMock()
fake_client.images.edit.return_value = _fake_response(b64=_b64_png())

with _patched_openai(fake_client):
result = provider.edit_image(
str(image),
"Turn this into a half-body work photo while preserving the same face.",
aspect_ratio="portrait",
)

assert result["success"] is True
assert result["model"] == "gpt-image-2-medium"
assert result["provider"] == "openai"
assert result["aspect_ratio"] == "portrait"
assert result["source_image"] == str(image)
assert result["instruction"].startswith("Turn this into")

saved = Path(result["image"])
assert saved.exists()

call_kwargs = fake_client.images.edit.call_args.kwargs
assert call_kwargs["model"] == "gpt-image-2"
assert call_kwargs["quality"] == "medium"
assert call_kwargs["size"] == "1024x1536"
assert call_kwargs["prompt"].startswith("Turn this into")

def test_edit_image_api_error_returns_error_response(self, provider, tmp_path):
image = tmp_path / "source.png"
image.write_bytes(bytes.fromhex(_PNG_HEX))

fake_client = MagicMock()
fake_client.images.edit.side_effect = RuntimeError("edit boom")

with _patched_openai(fake_client):
result = provider.edit_image(str(image), "Edit it")

assert result["success"] is False
assert result["error_type"] == "api_error"
assert "edit boom" in result["error"]
Loading