refactor(tools): image gen — unify provider result envelope, compact headers/imports/exception classes

This commit is contained in:
Teknium
2026-09-02 21:52:25 -07:00
parent b54e31e280
commit a0adef4f6f
4 changed files with 52 additions and 131 deletions

View File

@@ -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:

View File

@@ -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

View File

@@ -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.<key>`` 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}

View File

@@ -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