From a0adef4f6fa62bdb58631974f85205b2da15f44b Mon Sep 17 00:00:00 2001 From: Teknium <127238744+teknium1@users.noreply.github.com> Date: Wed, 2 Sep 2026 21:52:25 -0700 Subject: [PATCH] =?UTF-8?q?refactor(tools):=20image=20gen=20=E2=80=94=20un?= =?UTF-8?q?ify=20provider=20result=20envelope,=20compact=20headers/imports?= =?UTF-8?q?/exception=20classes?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- tools/fal_common.py | 13 +--- tools/image_generation_catalog.py | 47 ++++---------- tools/image_generation_tool.py | 100 +++++++++--------------------- tools/image_source.py | 23 ++----- 4 files changed, 52 insertions(+), 131 deletions(-) diff --git a/tools/fal_common.py b/tools/fal_common.py index 7436ac7374..45055feff9 100644 --- a/tools/fal_common.py +++ b/tools/fal_common.py @@ -80,16 +80,9 @@ class _ManagedFalSyncClient: _require(self._request_handle_class, "fal_client.client.SyncRequestHandle") def submit( - self, - application: str, - arguments: Dict[str, Any], - *, - path: str = "", - hint: Optional[str] = None, - webhook_url: Optional[str] = None, - priority: Any = None, - headers: Optional[Dict[str, str]] = None, - start_timeout: Optional[Union[int, float]] = None, + self, application: str, arguments: Dict[str, Any], *, path: str = "", + hint: Optional[str] = None, webhook_url: Optional[str] = None, priority: Any = None, + headers: Optional[Dict[str, str]] = None, start_timeout: Optional[Union[int, float]] = None, ): url = self._queue_url_format + application if path: diff --git a/tools/image_generation_catalog.py b/tools/image_generation_catalog.py index d8e3c3660e..fc9e3ddfbe 100644 --- a/tools/image_generation_catalog.py +++ b/tools/image_generation_catalog.py @@ -1,52 +1,31 @@ """FAL image model catalog + upscaler constants for ``tools.image_generation_tool``. -Each entry declares how to translate the unified inputs (prompt + aspect_ratio) -into the model's native payload. ``size_style`` picks the family: -``"image_size_preset"`` (FAL preset enum), ``"aspect_ratio"`` (ratio enum), -``"gpt_literal"`` (literal "WxH" strings). ``supports`` / ``edit_supports`` are -whitelists — keys outside them are stripped so models never receive rejected -parameters. ``upscale`` (Clarity Upscaler chained after generation) is False -everywhere: Clarity redraws content (creativity 0.35) and degraded text/CJK/ -faces when default-on, so upscaling is strictly per-call opt-in. -Pricing strings are as-of-commit and allowed to drift. +Each entry translates the unified inputs (prompt + aspect_ratio) into the model's native +payload. ``size_style``: ``"image_size_preset"`` (FAL preset enum), ``"aspect_ratio"`` (ratio +enum), ``"gpt_literal"`` (literal "WxH"). ``supports`` / ``edit_supports`` are whitelists — +other keys are stripped so models never receive rejected parameters. ``upscale`` is False +everywhere: Clarity redraws content (creativity 0.35) and degraded text/CJK/faces when +default-on, so upscaling is strictly per-call opt-in. Pricing strings may drift. """ from typing import Any, Dict, Optional -_PRESET_SIZES = { - "landscape": "landscape_16_9", - "square": "square_hd", - "portrait": "portrait_16_9", -} +_PRESET_SIZES = {"landscape": "landscape_16_9", "square": "square_hd", "portrait": "portrait_16_9"} _ASPECT_SIZES = {"landscape": "16:9", "square": "1:1", "portrait": "9:16"} _DEFAULT_SIZES = {"image_size_preset": _PRESET_SIZES, "aspect_ratio": _ASPECT_SIZES} def _model( - display: str, - speed: str, - strengths: str, - price: str, - *, - style: str = "image_size_preset", - sizes: Optional[Dict[str, Any]] = None, - defaults: Dict[str, Any], - supports: set, - edit_endpoint: Optional[str] = None, - edit_supports: Optional[set] = None, + display: str, speed: str, strengths: str, price: str, *, style: str = "image_size_preset", + sizes: Optional[Dict[str, Any]] = None, defaults: Dict[str, Any], supports: set, + edit_endpoint: Optional[str] = None, edit_supports: Optional[set] = None, max_reference_images: Optional[int] = None, ) -> Dict[str, Any]: """Build one catalog entry; edit keys are present only for edit-capable models.""" entry: Dict[str, Any] = { - "display": display, - "speed": speed, - "strengths": strengths, - "price": price, - "size_style": style, - "sizes": sizes if sizes is not None else _DEFAULT_SIZES[style], - "defaults": defaults, - "supports": supports, - "upscale": False, + "display": display, "speed": speed, "strengths": strengths, "price": price, + "size_style": style, "sizes": sizes if sizes is not None else _DEFAULT_SIZES[style], + "defaults": defaults, "supports": supports, "upscale": False, } if edit_endpoint: entry["edit_endpoint"] = edit_endpoint diff --git a/tools/image_generation_tool.py b/tools/image_generation_tool.py index 874bfa6ef6..70f9ac7fb6 100644 --- a/tools/image_generation_tool.py +++ b/tools/image_generation_tool.py @@ -36,28 +36,14 @@ from tools.fal_common import ( _normalize_fal_queue_url_format, # noqa: F401 — re-exported for tests ) from tools.image_generation_catalog import ( # noqa: F401 — re-exported (plugins/tests/tools_config) - DEFAULT_ASPECT_RATIO, - DEFAULT_MODEL, - FAL_MODELS, - UPSCALER_CREATIVITY, - UPSCALER_DEFAULT_PROMPT, - UPSCALER_FACTOR, - UPSCALER_GUIDANCE_SCALE, - UPSCALER_MODEL, - UPSCALER_NEGATIVE_PROMPT, - UPSCALER_NUM_INFERENCE_STEPS, - UPSCALER_RESEMBLANCE, - UPSCALER_SAFETY_CHECKER, - VALID_ASPECT_RATIOS, + DEFAULT_ASPECT_RATIO, DEFAULT_MODEL, FAL_MODELS, UPSCALER_CREATIVITY, UPSCALER_DEFAULT_PROMPT, + UPSCALER_FACTOR, UPSCALER_GUIDANCE_SCALE, UPSCALER_MODEL, UPSCALER_NEGATIVE_PROMPT, + UPSCALER_NUM_INFERENCE_STEPS, UPSCALER_RESEMBLANCE, UPSCALER_SAFETY_CHECKER, VALID_ASPECT_RATIOS, ) from tools.managed_tool_gateway import resolve_managed_tool_gateway from tools.tool_backend_helpers import ( - NOUS_MANAGED_PROVIDER, - fal_key_is_configured, - managed_nous_tools_enabled, - nous_tool_gateway_unavailable_message, - read_selection, - selection_error, + NOUS_MANAGED_PROVIDER, fal_key_is_configured, managed_nous_tools_enabled, + nous_tool_gateway_unavailable_message, read_selection, selection_error, ) logger = logging.getLogger(__name__) @@ -68,9 +54,7 @@ _managed_fal_client_config = None _managed_fal_client_lock = threading.Lock() -# --------------------------------------------------------------------------- -# Managed FAL gateway (Nous Subscription) -# --------------------------------------------------------------------------- +# --- Managed FAL gateway (Nous Subscription) --- def _resolve_managed_fal_gateway(): """Managed gateway config for the stored `hermes tools` selection, or ``None`` for direct FAL. @@ -179,9 +163,7 @@ def _submit_fal_request(model: str, arguments: Dict[str, Any]): raise -# --------------------------------------------------------------------------- -# Config readers, model resolution + payload construction -# --------------------------------------------------------------------------- +# --- Config readers, model resolution + payload construction --- def _read_image_gen_key(key: str) -> Optional[str]: """Return the stripped ``image_gen.`` string from config.yaml, or None.""" try: @@ -277,9 +259,7 @@ def _build_fal_edit_payload(model_id, prompt, image_urls, aspect_ratio=DEFAULT_A return _build_payload(model_id, prompt, aspect_ratio, seed, overrides, image_urls=image_urls) -# --------------------------------------------------------------------------- -# Upscaler -# --------------------------------------------------------------------------- +# --- Upscaler --- def _upscale_image(image_url: str, original_prompt: str) -> Optional[Dict[str, Any]]: """Upscale via FAL's Clarity Upscaler; None on failure (caller keeps the original).""" try: @@ -314,9 +294,7 @@ def _upscale_image(image_url: str, original_prompt: str) -> Optional[Dict[str, A return None -# --------------------------------------------------------------------------- -# Artifact path hinting for non-local terminal backends -# --------------------------------------------------------------------------- +# --- Artifact path hinting for non-local terminal backends --- _CONTAINER_HOME_ENVS = {"DockerEnvironment", "SingularityEnvironment", "ModalEnvironment"} # No environment yet: only backends with deterministic cache roots can be translated without # side effects. SSH uses a shell-visible tilde path; its first sync uploads the cache file. @@ -395,9 +373,7 @@ def _postprocess_image_generate_result(raw: str, task_id: str | None = None) -> return json.dumps(payload, ensure_ascii=False) -# --------------------------------------------------------------------------- -# Tool entry point -# --------------------------------------------------------------------------- +# --- Tool entry point --- def _format_images(images: list, should_upscale: bool, prompt: str) -> list: """Normalize FAL result images, optionally chaining the upscaler (falls back to the original on failure).""" formatted = [] @@ -458,16 +434,11 @@ def _prepare_fal_request(model_id, meta, prompt, aspect_ratio, seed, overrides, def image_generate_tool( - prompt: str, - aspect_ratio: str = DEFAULT_ASPECT_RATIO, - num_inference_steps: Optional[int] = None, - guidance_scale: Optional[float] = None, - num_images: Optional[int] = None, - output_format: Optional[str] = None, - seed: Optional[int] = None, - image_url: Optional[str] = None, - reference_image_urls: Optional[list] = None, - upscale: Optional[bool] = None, + prompt: str, aspect_ratio: str = DEFAULT_ASPECT_RATIO, + num_inference_steps: Optional[int] = None, guidance_scale: Optional[float] = None, + num_images: Optional[int] = None, output_format: Optional[str] = None, + seed: Optional[int] = None, image_url: Optional[str] = None, + reference_image_urls: Optional[list] = None, upscale: Optional[bool] = None, ) -> str: """Generate (or, with source images + an ``edit_endpoint`` model, edit) an image via FAL. @@ -620,9 +591,7 @@ def check_image_generation_requirements() -> bool: return False -# --------------------------------------------------------------------------- -# Registry -# --------------------------------------------------------------------------- +# --- Registry --- from tools.registry import registry, tool_error IMAGE_GENERATE_SCHEMA = { @@ -659,14 +628,19 @@ IMAGE_GENERATE_SCHEMA = { } -# --------------------------------------------------------------------------- -# Plugin provider dispatch + managed-mode Krea routing -# --------------------------------------------------------------------------- +# --- Plugin provider dispatch + managed-mode Krea routing --- def _provider_error(error: str, error_type: str) -> str: """JSON error envelope shared by every provider-dispatch failure path.""" return json.dumps({"success": False, "image": None, "error": error, "error_type": error_type}) +def _provider_result(result, contract_error: str) -> str: + """JSON-encode a provider's dict result; anything else is a contract violation.""" + if not isinstance(result, dict): + return _provider_error(contract_error, "provider_contract") + return json.dumps(result) + + def _add_provider_kwargs(kwargs, image_url, reference_image_urls, upscale, model=None) -> Dict[str, Any]: """Add the optional ``provider.generate(**kwargs)`` args in place (edit args only when supplied).""" if model: @@ -684,11 +658,8 @@ def _add_provider_kwargs(kwargs, image_url, reference_image_urls, upscale, model def _dispatch_to_plugin_provider( - prompt: str, - aspect_ratio: str, - image_url: Optional[str] = None, - reference_image_urls: Optional[list] = None, - upscale: Optional[bool] = None, + prompt: str, aspect_ratio: str, image_url: Optional[str] = None, + reference_image_urls: Optional[list] = None, upscale: Optional[bool] = None, ): """JSON result from the selected plugin provider, or ``None`` to fall through to in-tree FAL. @@ -743,9 +714,7 @@ def _dispatch_to_plugin_provider( except Exception as exc: logger.warning("Image gen provider '%s' raised: %s", pname, exc) return _provider_error(f"Provider '{pname}' error: {exc}", "provider_exception") - if not isinstance(result, dict): - return _provider_error("Provider returned a non-dict result", "provider_contract") - return json.dumps(result) + return _provider_result(result, "Provider returned a non-dict result") # Native ``krea-2-*`` ids are served by the Krea managed gateway (managed mode only — @@ -760,11 +729,8 @@ def _normalize_krea_model(model_id: Optional[str]) -> Optional[str]: def _maybe_route_managed_krea( - prompt: str, - aspect_ratio: str, - image_url: Optional[str] = None, - reference_image_urls: Optional[list] = None, - upscale: Optional[bool] = None, + prompt: str, aspect_ratio: str, image_url: Optional[str] = None, + reference_image_urls: Optional[list] = None, upscale: Optional[bool] = None, ) -> Optional[str]: """JSON result from the managed Krea gateway, or ``None`` to fall through. @@ -800,9 +766,7 @@ def _maybe_route_managed_krea( except Exception as exc: # noqa: BLE001 logger.warning("Managed Krea routing failed: %s", exc) return _provider_error(f"Managed Krea generation error: {exc}", "provider_exception") - if not isinstance(result, dict): - return _provider_error("Krea provider returned a non-dict result", "provider_contract") - return json.dumps(result) + return _provider_result(result, "Krea provider returned a non-dict result") def _confine_source_images(image_url, reference_image_urls, task_id, *, permitted: tuple = ("image",)): @@ -862,9 +826,7 @@ def _handle_image_generate(args, **kw): return _postprocess_image_generate_result(raw, task_id=task_id) -# --------------------------------------------------------------------------- -# Dynamic schema — reflect the active backend's image-to-image capability -# --------------------------------------------------------------------------- +# --- Dynamic schema — reflect the active backend's image-to-image capability --- # Telling the model up front whether it can edit saves a wasted turn. Memoized by # config.yaml mtime in model_tools.get_tool_definitions(), so it rebuilds on switch. _NO_CAPABILITIES = {"modalities": ["text"], "max_reference_images": 0, "supports_upscale": False} diff --git a/tools/image_source.py b/tools/image_source.py index 509e667031..5d73c0a100 100644 --- a/tools/image_source.py +++ b/tools/image_source.py @@ -30,24 +30,11 @@ class ImageResolutionError(Exception): self.src, self.origin = src, origin -class UnsupportedScheme(ImageResolutionError): - pass - - -class SourceUnsafe(ImageResolutionError): # SSRF / path-allowlist - pass - - -class SourceTooLarge(ImageResolutionError): - pass - - -class SourceNotFound(ImageResolutionError): - pass - - -class NotAnImage(ImageResolutionError): - pass +class UnsupportedScheme(ImageResolutionError): ... +class SourceUnsafe(ImageResolutionError): ... # SSRF / path-allowlist +class SourceTooLarge(ImageResolutionError): ... +class SourceNotFound(ImageResolutionError): ... +class NotAnImage(ImageResolutionError): ... @dataclass