From 4a2776aa72ee344c0922b9dd29451d53a3e65859 Mon Sep 17 00:00:00 2001 From: Kowen Hao Date: Tue, 26 May 2026 00:31:48 +0800 Subject: [PATCH] feat(image): add OpenAI image edit support --- agent/image_gen_provider.py | 24 +++ plugins/image_gen/openai/__init__.py | 192 +++++++++++++++++- .../plugins/image_gen/test_openai_provider.py | 75 +++++++ 3 files changed, 288 insertions(+), 3 deletions(-) diff --git a/agent/image_gen_provider.py b/agent/image_gen_provider.py index a7f1b8c31ff95..92f6d01d4379d 100644 --- a/agent/image_gen_provider.py +++ b/agent/image_gen_provider.py @@ -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, diff --git a/plugins/image_gen/openai/__init__.py b/plugins/image_gen/openai/__init__.py index 448f5bc45af38..c8ce0ee4af11a 100644 --- a/plugins/image_gen/openai/__init__.py +++ b/plugins/image_gen/openai/__init__.py @@ -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, @@ -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") @@ -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 @@ -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", @@ -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 " @@ -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 @@ -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) @@ -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 diff --git a/tests/plugins/image_gen/test_openai_provider.py b/tests/plugins/image_gen/test_openai_provider.py index 6411996130e1f..4cdd45e33fac2 100644 --- a/tests/plugins/image_gen/test_openai_provider.py +++ b/tests/plugins/image_gen/test_openai_provider.py @@ -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() @@ -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 @@ -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"]