323 lines
11 KiB
Python
323 lines
11 KiB
Python
"""Shared helpers for the bundled ``image_gen`` provider plugins.
|
|
|
|
Every provider under ``plugins/image_gen/`` is loaded by path (module name
|
|
``hermes_plugins.image_gen__<name>``) and resolves this module through the
|
|
repo root on ``sys.path`` — the same way they import ``agent.image_gen_provider``.
|
|
Nothing here is a plugin: the scanner only looks at directories.
|
|
"""
|
|
|
|
from __future__ import annotations
|
|
|
|
import logging
|
|
import os
|
|
from dataclasses import dataclass
|
|
from typing import Any, Callable, Dict, Iterable, List, Optional, Tuple
|
|
|
|
from agent.image_gen_provider import (
|
|
error_response,
|
|
normalize_reference_images,
|
|
save_b64_image,
|
|
save_url_image,
|
|
)
|
|
|
|
logger = logging.getLogger(__name__)
|
|
|
|
# OpenAI-style ``size`` strings for the three semantic aspect ratios. Shared by
|
|
# every OpenAI-compatible ``images.generations`` backend so aspect_ratio means
|
|
# the same thing across providers.
|
|
OPENAI_SIZES: Dict[str, str] = {
|
|
"landscape": "1536x1024",
|
|
"square": "1024x1024",
|
|
"portrait": "1024x1536",
|
|
}
|
|
|
|
# gpt-image-2 quality tiers, exposed as three virtual model ids so the picker
|
|
# and ``image_gen.model`` behave like any other multi-model backend. All three
|
|
# hit the same API model with a different ``quality`` knob.
|
|
GPT_IMAGE_2_API_MODEL = "gpt-image-2"
|
|
GPT_IMAGE_2_DEFAULT = "gpt-image-2-medium"
|
|
GPT_IMAGE_2_TIERS: Dict[str, Dict[str, Any]] = {
|
|
"gpt-image-2-low": {
|
|
"display": "GPT Image 2 (Low)",
|
|
"speed": "~15s",
|
|
"strengths": "Fast iteration, lowest cost",
|
|
"quality": "low",
|
|
},
|
|
"gpt-image-2-medium": {
|
|
"display": "GPT Image 2 (Medium)",
|
|
"speed": "~40s",
|
|
"strengths": "Balanced — default",
|
|
"quality": "medium",
|
|
},
|
|
"gpt-image-2-high": {
|
|
"display": "GPT Image 2 (High)",
|
|
"speed": "~2min",
|
|
"strengths": "Highest fidelity, strongest prompt adherence",
|
|
"quality": "high",
|
|
},
|
|
}
|
|
|
|
PROMPT_REQUIRED = "Prompt is required and must be a non-empty string"
|
|
OPENAI_MISSING = "openai Python package not installed (pip install openai)"
|
|
|
|
ErrorFn = Callable[..., Dict[str, Any]]
|
|
|
|
|
|
def size_for(aspect: str) -> str:
|
|
"""OpenAI ``size`` string for a semantic aspect (square when unknown)."""
|
|
return OPENAI_SIZES.get(aspect, OPENAI_SIZES["square"])
|
|
|
|
|
|
def load_image_gen_config(sub: Optional[str] = None) -> Dict[str, Any]:
|
|
"""Read ``image_gen`` (or ``image_gen.<sub>``) from config.yaml; ``{}`` on any failure."""
|
|
label = "image_gen" if sub is None else f"image_gen.{sub}"
|
|
try:
|
|
from hermes_cli.config import load_config
|
|
|
|
cfg = load_config()
|
|
section = cfg.get("image_gen") if isinstance(cfg, dict) else None
|
|
if sub is not None:
|
|
section = section.get(sub) if isinstance(section, dict) else None
|
|
return section if isinstance(section, dict) else {}
|
|
except Exception as exc: # noqa: BLE001 - config is best-effort
|
|
logger.debug("Could not load %s config: %s", label, exc)
|
|
return {}
|
|
|
|
|
|
def resolve_static_model(
|
|
models: Dict[str, Dict[str, Any]],
|
|
default: str,
|
|
*,
|
|
env_var: str,
|
|
config_key: str,
|
|
explicit: Optional[str] = None,
|
|
include_top_level: bool = True,
|
|
config: Optional[Dict[str, Any]] = None,
|
|
) -> Tuple[str, Dict[str, Any]]:
|
|
"""Pick ``(model_id, meta)`` from a fixed catalog.
|
|
|
|
Precedence (first *known* id wins; unknown ids fall through rather than
|
|
error): explicit caller override → ``env_var`` → ``image_gen.<config_key>.model``
|
|
→ top-level ``image_gen.model`` (when ``include_top_level``) → ``default``.
|
|
"""
|
|
if isinstance(explicit, str) and explicit.strip() in models:
|
|
return explicit.strip(), models[explicit.strip()]
|
|
|
|
env_override = os.environ.get(env_var)
|
|
if env_override and env_override in models:
|
|
return env_override, models[env_override]
|
|
|
|
cfg = load_image_gen_config() if config is None else config
|
|
scoped = cfg.get(config_key)
|
|
candidates = [scoped.get("model") if isinstance(scoped, dict) else None]
|
|
if include_top_level:
|
|
candidates.append(cfg.get("model"))
|
|
for candidate in candidates:
|
|
if isinstance(candidate, str) and candidate in models:
|
|
return candidate, models[candidate]
|
|
|
|
return default, models[default]
|
|
|
|
|
|
def collect_source_images(
|
|
image_url: Optional[str],
|
|
reference_image_urls: Optional[List[str]],
|
|
limit: Optional[int] = None,
|
|
) -> List[str]:
|
|
"""Primary ``image_url`` first, then normalized references, clamped to ``limit``."""
|
|
sources: List[str] = []
|
|
if isinstance(image_url, str) and image_url.strip():
|
|
sources.append(image_url.strip())
|
|
sources.extend(normalize_reference_images(reference_image_urls) or [])
|
|
return sources[:limit] if limit is not None else sources
|
|
|
|
|
|
def catalog_rows(
|
|
models: Dict[str, Dict[str, Any]],
|
|
fields: Iterable[str] = ("display", "speed", "strengths", "price"),
|
|
*,
|
|
price: Optional[str] = None,
|
|
) -> List[Dict[str, Any]]:
|
|
"""Picker rows for a catalog: ``id`` plus ``fields`` read from each meta.
|
|
|
|
A missing ``display`` falls back to the model id, other missing fields to
|
|
``""``; ``price`` (when given) overrides the per-model value.
|
|
"""
|
|
rows = []
|
|
for model_id, meta in models.items():
|
|
row: Dict[str, Any] = {"id": model_id}
|
|
for field in fields:
|
|
row[field] = meta.get(field, model_id if field == "display" else "")
|
|
if price is not None:
|
|
row["price"] = price
|
|
rows.append(row)
|
|
return rows
|
|
|
|
|
|
def api_key_setup_schema(
|
|
name: str, badge: str, tag: str, *, key: str, prompt: str, url: str
|
|
) -> Dict[str, Any]:
|
|
"""``get_setup_schema()`` dict for a provider authenticated by one env var."""
|
|
return {
|
|
"name": name,
|
|
"badge": badge,
|
|
"tag": tag,
|
|
"env_vars": [{"key": key, "prompt": prompt, "url": url}],
|
|
}
|
|
|
|
|
|
def error_factory(provider: str, aspect: str, *, model: str = "", prompt: str = "") -> ErrorFn:
|
|
"""Return ``fail(error, error_type, **override)`` pre-bound to this call's context."""
|
|
|
|
def fail(error: str, error_type: str, **override: Any) -> Dict[str, Any]:
|
|
kwargs = dict(provider=provider, model=model, prompt=prompt, aspect_ratio=aspect)
|
|
kwargs.update(override)
|
|
return error_response(error=error, error_type=error_type, **kwargs)
|
|
|
|
return fail
|
|
|
|
|
|
def prompt_required_error(provider: str, aspect: str) -> Dict[str, Any]:
|
|
return error_response(
|
|
error=PROMPT_REQUIRED,
|
|
error_type="invalid_argument",
|
|
provider=provider,
|
|
aspect_ratio=aspect,
|
|
)
|
|
|
|
|
|
def openai_importable() -> bool:
|
|
try:
|
|
import openai # noqa: F401
|
|
except ImportError:
|
|
return False
|
|
return True
|
|
|
|
|
|
def import_openai(provider: str, aspect: str) -> Tuple[Any, Optional[Dict[str, Any]]]:
|
|
"""Return ``(openai_module, None)`` or ``(None, missing_dependency error)``."""
|
|
try:
|
|
import openai
|
|
except ImportError:
|
|
return None, error_response(
|
|
error=OPENAI_MISSING,
|
|
error_type="missing_dependency",
|
|
provider=provider,
|
|
aspect_ratio=aspect,
|
|
)
|
|
return openai, None
|
|
|
|
|
|
def materialize_image(
|
|
b64: Optional[str],
|
|
url: Optional[str],
|
|
*,
|
|
prefix: str,
|
|
label: str,
|
|
provider: str,
|
|
model: str,
|
|
prompt: str,
|
|
aspect: str,
|
|
log: logging.Logger = logger,
|
|
) -> Tuple[Optional[str], Optional[Dict[str, Any]]]:
|
|
"""Turn a ``(b64_json, url)`` result pair into a local image reference.
|
|
|
|
Returns ``(image_ref, None)`` or ``(None, error_dict)``. Base64 is always
|
|
cached (a write failure is an ``io_error``); a URL is cached best-effort
|
|
because provider URLs are often ephemeral, falling back to the bare URL so
|
|
a cache hiccup never destroys a successful generation.
|
|
"""
|
|
fail = error_factory(provider, aspect, model=model, prompt=prompt)
|
|
if b64:
|
|
try:
|
|
return str(save_b64_image(b64, prefix=prefix)), None
|
|
except Exception as exc: # noqa: BLE001
|
|
return None, fail(f"Could not save image to cache: {exc}", "io_error")
|
|
if url:
|
|
return cache_url_best_effort(url, prefix=prefix, label=label, log=log), None
|
|
return None, fail(f"{label} response contained neither b64_json nor URL", "empty_response")
|
|
|
|
|
|
def cache_url_best_effort(url: str, *, prefix: str, label: str, log: logging.Logger = logger) -> str:
|
|
"""Cache ``url`` locally; on failure warn and return the bare URL."""
|
|
try:
|
|
return str(save_url_image(url, prefix=prefix))
|
|
except Exception as exc: # noqa: BLE001
|
|
log.warning(
|
|
"%s image URL %s could not be cached (%s); falling back to bare URL.", label, url, exc,
|
|
)
|
|
return url
|
|
|
|
|
|
def requests_error_message(response: Any, exc: Exception) -> str:
|
|
"""``error.message`` from an HTTP error body, else the first 300 chars of it."""
|
|
try:
|
|
return response.json().get("error", {}).get("message", response.text[:300])
|
|
except Exception: # noqa: BLE001 - non-JSON / non-dict body
|
|
return response.text[:300] if response is not None else str(exc)
|
|
|
|
|
|
@dataclass
|
|
class HttpFailure:
|
|
"""One failed ``post_json`` attempt, pre-shaped as ``(error, error_type)``.
|
|
|
|
``kind`` is ``http`` / ``timeout`` / ``connection`` / ``request`` /
|
|
``invalid_json``; ``status`` and ``response`` are set for ``http`` only.
|
|
"""
|
|
|
|
kind: str
|
|
error: str
|
|
error_type: str
|
|
status: int = 0
|
|
message: str = ""
|
|
response: Any = None
|
|
|
|
|
|
def post_json(
|
|
url: str,
|
|
*,
|
|
headers: Dict[str, str],
|
|
payload: Dict[str, Any],
|
|
timeout: Any,
|
|
label: str,
|
|
error_message: Callable[[Any, Exception], str] = requests_error_message,
|
|
catch_request_exception: bool = False,
|
|
) -> Tuple[Optional[Any], Optional[HttpFailure]]:
|
|
"""POST ``payload`` and parse the JSON body: ``(body, None)`` or ``(None, failure)``.
|
|
|
|
``timeout`` is passed to ``requests`` verbatim; the timeout message reports
|
|
its read component. ``error_message(response, exc)`` extracts the HTTP error
|
|
text so each backend can keep its own body shape.
|
|
"""
|
|
import requests
|
|
|
|
read_timeout = timeout[1] if isinstance(timeout, tuple) else timeout
|
|
try:
|
|
response = requests.post(url, headers=headers, json=payload, timeout=timeout)
|
|
response.raise_for_status()
|
|
except requests.HTTPError as exc:
|
|
resp = exc.response
|
|
status = resp.status_code if resp is not None else 0
|
|
message = error_message(resp, exc)
|
|
return None, HttpFailure(
|
|
"http", f"{label} image generation failed ({status}): {message}", "api_error",
|
|
status=status, message=message, response=resp,
|
|
)
|
|
except requests.Timeout:
|
|
return None, HttpFailure(
|
|
"timeout", f"{label} image generation timed out ({int(read_timeout)}s)", "timeout"
|
|
)
|
|
except requests.ConnectionError as exc:
|
|
return None, HttpFailure("connection", f"{label} connection error: {exc}", "connection_error")
|
|
except requests.RequestException as exc:
|
|
if not catch_request_exception:
|
|
raise
|
|
return None, HttpFailure("request", f"{label} request failed: {exc}", "api_error")
|
|
|
|
try:
|
|
return response.json(), None
|
|
except Exception as exc: # noqa: BLE001
|
|
return None, HttpFailure(
|
|
"invalid_json", f"{label} returned invalid JSON: {exc}", "invalid_response"
|
|
)
|