refactor(tools): vision/video/url_safety — inline single-use helpers, dedupe backends, compact docstrings

This commit is contained in:
Teknium
2026-09-02 23:28:33 -07:00
parent 8540940455
commit c4a9f2bf49
4 changed files with 482 additions and 989 deletions

View File

@@ -1,12 +1,10 @@
"""URL safety checks — blocks requests to private/internal network addresses (SSRF).
``security.allow_private_urls: true`` disables private-IP blocking for environments
whose DNS resolves public names to private/benchmark ranges; cloud metadata
hostnames/IPs are **always** blocked. DNS rebinding (TOCTOU) is closed for Hermes-owned
httpx paths by ``create_ssrf_safe_client()`` / ``create_ssrf_safe_async_client()``, which
re-apply the policy at TCP connect and dial the validated IP while preserving Host/SNI.
Redirect bypass is mitigated by response hooks re-validating each target
(``redirect_target_from_response``); third-party SDK fetches redirect server-side.
``security.allow_private_urls: true`` disables private-IP blocking (DNS that resolves public
names to private ranges); cloud metadata hostnames/IPs are **always** blocked. DNS rebinding
(TOCTOU) is closed for Hermes-owned httpx paths by ``create_ssrf_safe_[async_]client()``, which
re-apply the policy at TCP connect and dial the validated IP while preserving Host/SNI. Redirect
bypass is mitigated by response hooks re-validating each target (``redirect_target_from_response``).
"""
import ipaddress
@@ -26,9 +24,7 @@ logger = logging.getLogger(__name__)
# Proxy env vars: when set, the runtime should delegate DNS to the proxy.
_PROXY_ENV_VARS = ("HTTPS_PROXY", "https_proxy", "HTTP_PROXY", "http_proxy", "ALL_PROXY", "all_proxy")
_HTTP_SCHEMES = frozenset({"http", "https"})
_IPAddress = ipaddress.IPv4Address | ipaddress.IPv6Address
@@ -37,15 +33,11 @@ def _proxy_is_configured() -> bool:
def normalize_url_for_request(url: str) -> str:
"""Return an ASCII-safe HTTP URL for Hermes-owned URL tools (IRI -> URI).
Users/models often supply IRIs (``https://wttr.in/Köln``). Preserves URL syntax
and existing percent escapes while IDNA-encoding the host and percent-encoding
non-ASCII path/query/fragment text. URL tool inputs only — never shell commands.
"""
"""ASCII-safe HTTP URL for Hermes-owned URL tools (IRI -> URI, e.g. ``https://wttr.in/Köln``).
Preserves URL syntax and existing percent escapes while IDNA-encoding the host and
percent-encoding non-ASCII path/query/fragment text. URL tool inputs only — never shell commands."""
if not isinstance(url, str):
return url
raw = url.strip()
if not raw:
return raw
@@ -53,17 +45,13 @@ def normalize_url_for_request(url: str) -> str:
# Repair model-emitted whitespace between scheme separator and authority
# (``https:// docs.example``); that position is never meaningful in HTTP URLs.
raw = re.sub(r"^([A-Za-z][A-Za-z0-9+.-]*://)\s+", r"\1", raw)
try:
parsed = urlsplit(raw)
except ValueError:
return raw
if parsed.scheme.lower() not in _HTTP_SCHEMES:
return raw
netloc = parsed.netloc
hostname = parsed.hostname
netloc, hostname = parsed.netloc, parsed.hostname
if hostname:
try:
ascii_host = hostname.encode("idna").decode("ascii")
@@ -71,12 +59,9 @@ def normalize_url_for_request(url: str) -> str:
ascii_host = hostname
if ascii_host != hostname:
netloc = netloc.replace(hostname, ascii_host, 1)
path = quote(parsed.path, safe="/%:@!$&'()*+,;=")
query = quote(parsed.query, safe="/%:@!$&'()*+,;=?")
fragment = quote(parsed.fragment, safe="/%:@!$&'()*+,;=?")
return urlunsplit((parsed.scheme, netloc, path, query, fragment))
safe = "/%:@!$&'()*+,;="
return urlunsplit((parsed.scheme, netloc, quote(parsed.path, safe=safe),
quote(parsed.query, safe=safe + "?"), quote(parsed.fragment, safe=safe + "?")))
# Unambiguously credential-bearing query param names. Deliberately narrow: bare
@@ -86,17 +71,13 @@ _SENSITIVE_QUERY_PARAM_NAMES = frozenset({
"access_token", "api_key", "apikey", "auth_token", "authorization", "awsaccesskeyid",
"client_secret", "credential", "credentials", "jwt", "password", "passwd", "secret",
"session_id", "signature", "token", "x_amz_security_token", "x_amz_signature",
"x-amz-security-token", "x-amz-signature",
})
"x-amz-security-token", "x-amz-signature"})
def sensitive_query_param_name(url: str) -> Optional[str]:
"""Return the first credential-named query parameter in ``url`` (with a value), if any.
Checked before handing URLs to third-party fetch/browser backends: catches opaque
magic links, OAuth codes, signed-URL signatures and custom ``?token=...`` values
that prefix-based token redaction misses.
"""
"""First credential-named query parameter in ``url`` (with a value), if any. Checked before
handing URLs to third-party fetch/browser backends: catches opaque magic links, OAuth codes,
signed-URL signatures and custom ``?token=...`` values that prefix-based redaction misses."""
if not isinstance(url, str) or "?" not in url:
return None
try:
@@ -105,18 +86,15 @@ def sensitive_query_param_name(url: str) -> Optional[str]:
return None
if parsed.scheme.lower() not in _HTTP_SCHEMES or not parsed.query:
return None
for key, value in parse_qsl(parsed.query, keep_blank_values=True):
if value and unquote(key).lower() in _SENSITIVE_QUERY_PARAM_NAMES:
return key
return None
return next((key for key, value in parse_qsl(parsed.query, keep_blank_values=True)
if value and unquote(key).lower() in _SENSITIVE_QUERY_PARAM_NAMES), None)
# Cloud metadata hostnames — always blocked regardless of DNS or config toggle.
_BLOCKED_HOSTNAMES = frozenset({"metadata.google.internal", "metadata.goog"})
# Cloud metadata / credential endpoints (the #1 SSRF target) — always blocked.
# IPv4-mapped IPv6 forms are listed because resolvers may return ``::ffff:x.x.x.x``
# and ipaddress treats those as distinct.
# Cloud metadata / credential endpoints (the #1 SSRF target) — always blocked, also in
# IPv4-mapped IPv6 form (resolvers may return ``::ffff:x.x.x.x``; ipaddress treats those as distinct).
_METADATA_V4 = (
"169.254.169.254", # AWS/GCP/Azure/DO/Oracle metadata
"169.254.170.2", # AWS ECS task metadata (task IAM creds)
@@ -128,15 +106,12 @@ _ALWAYS_BLOCKED_IPS = frozenset(
| {ipaddress.ip_address("::ffff:" + ip) for ip in _METADATA_V4}
| {ipaddress.ip_address("fd00:ec2::254")} # AWS metadata (IPv6)
)
_ALWAYS_BLOCKED_NETWORKS = (
ipaddress.ip_network("169.254.0.0/16"), # entire link-local range (no legit agent target)
ipaddress.ip_network("::ffff:169.254.0.0/112"), # IPv4-mapped link-local range
)
# Entire link-local range (no legit agent target), plus its IPv4-mapped form.
_ALWAYS_BLOCKED_NETWORKS = tuple(ipaddress.ip_network(n) for n in ("169.254.0.0/16", "::ffff:169.254.0.0/112"))
# Exact HTTPS hostnames allowed to resolve to private/benchmark-space IPs
# (QQ media legitimately resolves to 198.18.0.0/15 behind local proxy infra).
_TRUSTED_PRIVATE_IP_HOSTS = frozenset({"multimedia.nt.qq.com.cn"})
_MAX_SSRF_CONNECT_IPS = 8
# 100.64.0.0/10 (CGNAT, RFC 6598) is neither is_private nor is_global in
@@ -144,27 +119,19 @@ _MAX_SSRF_CONNECT_IPS = 8
_CGNAT_NETWORK = ipaddress.ip_network("100.64.0.0/10")
# Global toggle cache (process lifetime; see _global_allow_private_urls).
_allow_private_resolved = False
_cached_allow_private: bool = False
_allow_private_resolved, _cached_allow_private = False, False
def _global_allow_private_urls() -> bool:
"""Return True when the user has opted out of private-IP blocking.
Priority: ``HERMES_ALLOW_PRIVATE_URLS`` env, ``security.allow_private_urls``,
legacy ``browser.allow_private_urls``. A multiplex gateway serves several
independently configured profiles in one process, so profile-scoped turns
(``get_hermes_home_override()`` set) bypass the process-global cache —
otherwise the first profile's opt-out would disable blocking for every later one.
"""
"""True when the user has opted out of private-IP blocking. Priority: ``HERMES_ALLOW_PRIVATE_URLS``
env, ``security.allow_private_urls``, legacy ``browser.allow_private_urls``. Profile-scoped turns
(``get_hermes_home_override()`` set) bypass the process-global cache — a multiplex gateway serves
several profiles in one process; the first profile's opt-out must not disable blocking for later ones."""
global _allow_private_resolved, _cached_allow_private
if get_hermes_home_override() is not None:
return _resolve_allow_private_urls()
if not _allow_private_resolved:
_allow_private_resolved = True
_cached_allow_private = _resolve_allow_private_urls()
_allow_private_resolved, _cached_allow_private = True, _resolve_allow_private_urls()
return _cached_allow_private
@@ -175,7 +142,6 @@ def _resolve_allow_private_urls() -> bool:
return True
if env_val in {"false", "0", "no"}:
return False # explicit false does not fall through to config
try:
from hermes_cli.config import read_raw_config
cfg = read_raw_config()
@@ -185,15 +151,13 @@ def _resolve_allow_private_urls() -> bool:
return True
except Exception:
pass # config unavailable (tests, early import) — keep default
return False
def _reset_allow_private_cache() -> None:
"""Reset the cached toggle — only for tests."""
global _allow_private_resolved, _cached_allow_private
_allow_private_resolved = False
_cached_allow_private = False
_allow_private_resolved = _cached_allow_private = False
def _normalize_hostname(host: Optional[str]) -> str:
@@ -201,7 +165,7 @@ def _normalize_hostname(host: Optional[str]) -> str:
def _parse_ip(hostname: str) -> Optional[_IPAddress]:
"""Return the IP object for a literal-IP hostname, else None."""
"""IP object for a literal-IP hostname, else None."""
try:
return ipaddress.ip_address(hostname)
except ValueError:
@@ -209,14 +173,11 @@ def _parse_ip(hostname: str) -> Optional[_IPAddress]:
def _iter_resolved_ips(addr_info: Any):
"""Yield ``(raw, ip_str, ip)`` per getaddrinfo answer.
``ip_str`` has any IPv6 scope ID (``%eth0``) stripped; ``ip`` is None when the
answer is still unparseable — each caller decides skip / fail-closed / raise.
"""
"""Yield ``(raw, ip_str, ip)`` per getaddrinfo answer. ``ip_str`` has any IPv6 scope ID
(``%eth0``) stripped; ``ip`` is None when still unparseable — each caller decides skip/fail/raise."""
for _family, _, _, _, sockaddr in addr_info:
raw = sockaddr[0]
ip_str = raw.split("%")[0] if "%" in raw else raw
ip_str = raw.split("%")[0]
yield raw, ip_str, _parse_ip(ip_str)
@@ -238,36 +199,28 @@ def _is_blocked_ip(ip: _IPAddress) -> bool:
def is_always_blocked_url(url: str) -> bool:
"""Return True when the URL targets the always-blocked floor (cloud metadata).
Narrower than ``is_safe_url``: only the sentinel hostnames/IPs, regardless of
backend, routing, or ``allow_private_urls`` — for callers that deliberately bypass
the full check (e.g. hybrid cloud browser routing private URLs to a local sidecar)
but must still enforce the non-negotiable floor. Returns False for ordinary
private/loopback URLs, DNS failures on non-sentinel hosts, and parse errors
(the caller's ordinary fail-closed path handles those).
"""
"""True when the URL targets the always-blocked floor (cloud metadata) — only the sentinel
hostnames/IPs, regardless of backend, routing, or ``allow_private_urls``. For callers that
deliberately bypass the full check (e.g. hybrid cloud browser routing private URLs to a local
sidecar) but must still enforce the floor. False for ordinary private/loopback URLs, DNS
failures and parse errors (the caller's ordinary fail-closed path handles those)."""
try:
hostname = _normalize_hostname(urlparse(url).hostname)
if not hostname:
return False
if hostname in _BLOCKED_HOSTNAMES:
logger.warning("Blocked request to internal hostname (always-blocked floor): %s", hostname)
return True
ip = _parse_ip(hostname)
if ip is not None:
if _is_always_blocked_ip(ip):
logger.warning("Blocked request to cloud metadata address (always-blocked floor): %s", hostname)
return True
return False
try:
addr_info = _getaddrinfo(hostname)
except socket.gaierror:
return False # DNS failure is not part of the floor; caller's path handles it
for raw, ip_str, resolved in _iter_resolved_ips(addr_info):
if resolved is None:
logger.warning("Unparseable IP address %r for hostname %s — skipping address", raw, hostname)
@@ -277,9 +230,7 @@ def is_always_blocked_url(url: str) -> bool:
"Blocked request to cloud metadata address (always-blocked floor): %s -> %s", hostname, ip_str
)
return True
return False
except Exception as exc:
# Parse/unexpected errors are not "always blocked"; caller decides fail-open/closed.
logger.debug("is_always_blocked_url error for %s: %s", url, exc)
@@ -292,11 +243,8 @@ def _allows_private_ip_resolution(hostname: str, scheme: str) -> bool:
def _resolved_ip_block_reason(ip: _IPAddress, allow_private: bool) -> Optional[str]:
"""Why a resolved answer must be rejected, or None if it may be dialed.
The metadata floor is checked first and ignores ``allow_private``; ordinary
private/internal classes are only blocked when ``allow_private`` is False.
"""
"""Why a resolved answer must be rejected, or None if it may be dialed. The metadata floor
ignores ``allow_private``; ordinary private/internal classes are blocked only when it is False."""
if _is_always_blocked_ip(ip):
return "cloud metadata address"
if not allow_private and _is_blocked_ip(ip):
@@ -305,12 +253,9 @@ def _resolved_ip_block_reason(ip: _IPAddress, allow_private: bool) -> Optional[s
def is_safe_url(url: str) -> bool:
"""Return True if the URL target is not a private/internal address.
Resolves the hostname and checks every answer. Fails closed on DNS errors and
unexpected exceptions. ``allow_private_urls`` skips private-IP blocking, but
cloud metadata endpoints remain blocked regardless.
"""
"""True if the URL target is not a private/internal address. Resolves the hostname and checks
every answer; fails closed on DNS errors and unexpected exceptions. ``allow_private_urls``
skips private-IP blocking, but cloud metadata endpoints remain blocked regardless."""
try:
parsed = urlparse(url)
hostname = _normalize_hostname(parsed.hostname)
@@ -325,11 +270,9 @@ def is_safe_url(url: str) -> bool:
if hostname in _BLOCKED_HOSTNAMES:
logger.warning("Blocked request to internal hostname: %s", hostname)
return False
allow_all_private = _global_allow_private_urls()
allow_private_ip = _allows_private_ip_resolution(hostname, scheme)
allow_private = allow_all_private or allow_private_ip
try:
addr_info = _getaddrinfo(hostname)
except socket.gaierror:
@@ -340,29 +283,23 @@ def is_safe_url(url: str) -> bool:
if _parse_ip(hostname) is None and _proxy_is_configured():
logger.debug(
"DNS resolution failed for %s — proxy configured, allowing through for proxy-side resolution",
hostname,
)
hostname)
return True
logger.warning("Blocked request — DNS resolution failed for: %s", hostname)
return False
for raw, ip_str, ip in _iter_resolved_ips(addr_info):
if ip is None:
logger.warning("Blocked request — unparseable IP address %r for hostname %s", raw, hostname)
return False
reason = _resolved_ip_block_reason(ip, allow_private)
if reason is not None:
logger.warning("Blocked request to %s: %s -> %s", reason, hostname, ip_str)
return False
if allow_all_private:
logger.debug("Allowing private/internal resolution (security.allow_private_urls=true): %s", hostname)
elif allow_private_ip:
logger.debug("Allowing trusted hostname despite private/internal resolution: %s", hostname)
return True
except Exception as exc:
# Fail closed: parsing edge cases must not become SSRF bypass vectors.
logger.warning("Blocked request — URL safety check error for %s: %s", url, exc)
@@ -370,7 +307,7 @@ def is_safe_url(url: str) -> bool:
async def async_is_safe_url(url: str) -> bool:
"""Same rules as :func:`is_safe_url`, with the blocking DNS work off the event loop."""
""":func:`is_safe_url` with the blocking DNS work off the event loop."""
return await asyncio.to_thread(is_safe_url, url)
@@ -379,40 +316,30 @@ class SSRFConnectionBlocked(ValueError):
def _resolved_http_connect_ips(host: str, port: int, scheme: str) -> list[str]:
"""Resolve and validate *host* at TCP-connect time; return dialable IP strings.
Closes the DNS-rebinding gap between pre-flight validation and connection
setup for direct httpx clients. The result is capped at ``_MAX_SSRF_CONNECT_IPS``,
but EVERY answer is validated.
"""
"""Resolve and validate *host* at TCP-connect time; return dialable IP strings. Closes the
DNS-rebinding gap between pre-flight validation and connect for direct httpx clients. The
result is capped at ``_MAX_SSRF_CONNECT_IPS``, but EVERY answer is validated."""
hostname = _normalize_hostname(host)
if not hostname:
raise SSRFConnectionBlocked("Blocked request with empty hostname")
if hostname in _BLOCKED_HOSTNAMES:
raise SSRFConnectionBlocked(f"Blocked request to internal hostname: {hostname}")
allow_private = _global_allow_private_urls() or _allows_private_ip_resolution(hostname, scheme)
try:
addr_info = _getaddrinfo(hostname, port)
except socket.gaierror as exc:
raise SSRFConnectionBlocked(f"Blocked request - DNS resolution failed for: {hostname}") from exc
safe_ips: list[str] = []
for raw, ip_str, ip in _iter_resolved_ips(addr_info):
if ip is None:
raise SSRFConnectionBlocked(
f"Blocked request - unparseable IP address {raw!r} for hostname {hostname}"
) from ValueError(f"{ip_str!r} does not appear to be an IPv4 or IPv6 address")
reason = _resolved_ip_block_reason(ip, allow_private)
if reason is not None:
raise SSRFConnectionBlocked(f"Blocked request to {reason} during connect: {hostname} -> {ip_str}")
if ip_str not in safe_ips and len(safe_ips) < _MAX_SSRF_CONNECT_IPS:
safe_ips.append(ip_str)
if not safe_ips:
raise SSRFConnectionBlocked(f"Blocked request - DNS returned no results for: {hostname}")
return safe_ips
@@ -420,11 +347,8 @@ def _resolved_http_connect_ips(host: str, port: int, scheme: str) -> list[str]:
class _SSRFGuardedBackendBase:
"""httpcore backend that re-resolves + validates at connect time and dials a vetted IP.
Host/SNI stay on the original hostname (the transport still sees ``host``);
Unix sockets are refused outright. Candidate IPs are tried in order and the
last connect error is re-raised so callers see the real network failure.
"""
Host/SNI stay on the original hostname; Unix sockets are refused outright. Candidate IPs are
tried in order and the last connect error is re-raised so callers see the real failure."""
def __init__(self, backend: Any, schemes_by_origin_var: Any):
self._backend = backend
@@ -434,6 +358,11 @@ class _SSRFGuardedBackendBase:
scheme = self._schemes_by_origin_var.get({}).get((host, port)) or ("https" if port == 443 else "http")
return _resolved_http_connect_ips(host, port, scheme)
@staticmethod
def _connect_errors() -> tuple:
import httpcore
return (httpcore.ConnectError, httpcore.ConnectTimeout)
@staticmethod
def _no_usable_ips(host: str, last_exc: Exception | None) -> Exception:
if last_exc is not None:
@@ -451,16 +380,12 @@ class _SSRFGuardedAsyncNetworkBackend(_SSRFGuardedBackendBase):
async def connect_tcp(self, host: str, port: int, timeout: float | None = None,
local_address: str | None = None, socket_options: Any = None) -> Any:
import httpcore
ips = await asyncio.to_thread(self._connect_ips, host, port)
last_exc: Exception | None = None
for ip in ips:
for ip in await asyncio.to_thread(self._connect_ips, host, port):
try:
return await self._backend.connect_tcp(
ip, port, timeout=timeout, local_address=local_address, socket_options=socket_options,
)
except (httpcore.ConnectError, httpcore.ConnectTimeout) as exc:
ip, port, timeout=timeout, local_address=local_address, socket_options=socket_options)
except self._connect_errors() as exc:
last_exc = exc
raise self._no_usable_ips(host, last_exc)
@@ -478,15 +403,12 @@ class _SSRFGuardedNetworkBackend(_SSRFGuardedBackendBase):
def connect_tcp(self, host: str, port: int, timeout: float | None = None,
local_address: str | None = None, socket_options: Any = None) -> Any:
import httpcore
last_exc: Exception | None = None
for ip in self._connect_ips(host, port):
try:
return self._backend.connect_tcp(
ip, port, timeout=timeout, local_address=local_address, socket_options=socket_options,
)
except (httpcore.ConnectError, httpcore.ConnectTimeout) as exc:
ip, port, timeout=timeout, local_address=local_address, socket_options=socket_options)
except self._connect_errors() as exc:
last_exc = exc
raise self._no_usable_ips(host, last_exc)
@@ -496,9 +418,7 @@ class _SSRFGuardedNetworkBackend(_SSRFGuardedBackendBase):
def _origin_scheme_context(request: Any) -> dict[tuple[str, int], str]:
host, port, scheme = request.url.host, request.url.port, request.url.scheme
if not host or port is None or scheme not in _HTTP_SCHEMES:
return {}
return {(host, port): scheme}
return {(host, port): scheme} if host and port is not None and scheme in _HTTP_SCHEMES else {}
@contextmanager
@@ -512,22 +432,17 @@ def _origin_scope(schemes_by_origin_var: Any, request: Any):
def _install_ssrf_guard_on_transport(transport: Any, schemes_by_origin_var: Any, *, is_async: bool = False) -> None:
"""Swap the transport's pool network backend for the SSRF-guarded one (idempotent).
Only the client's direct transport is guarded; proxy mounts are left alone so
final-target resolution is delegated to the (trusted) proxy egress.
"""
"""Swap the transport's pool network backend for the SSRF-guarded one (idempotent). Only the
direct transport is guarded; proxy mounts delegate final-target resolution to the trusted proxy."""
state = getattr(transport, "__dict__", {}) if transport is not None else {}
if transport is None or state.get("_hermes_ssrf_guarded", False):
return
label = "async httpx transport" if is_async else "httpx transport"
pool = state.get("_pool")
if pool is None or not hasattr(pool, "_network_backend"):
raise SSRFConnectionBlocked(f"Unsupported {label} cannot be made SSRF-safe")
backend_cls = _SSRFGuardedAsyncNetworkBackend if is_async else _SSRFGuardedNetworkBackend
pool._network_backend = backend_cls(schemes_by_origin_var)
method_name = "handle_async_request" if is_async else "handle_request"
handle = getattr(transport, method_name, None)
if handle is None:
@@ -540,7 +455,6 @@ def _install_ssrf_guard_on_transport(transport: Any, schemes_by_origin_var: Any,
def guarded_sync(request: Any) -> Any:
with _origin_scope(schemes_by_origin_var, request):
return handle(request)
setattr(transport, method_name, guarded_async if is_async else guarded_sync)
transport._hermes_ssrf_guarded = True
@@ -548,22 +462,16 @@ def _install_ssrf_guard_on_transport(transport: Any, schemes_by_origin_var: Any,
def _install_ssrf_guard_on_client(client: Any, *, is_async: bool = False) -> None:
"""Guard ``client._transport`` only; ``_mounts`` (env/explicit proxies) stay untouched."""
import contextvars
var_name = "hermes_ssrf_async_origin_schemes" if is_async else "hermes_ssrf_origin_schemes"
_install_ssrf_guard_on_transport(
getattr(client, "__dict__", {}).get("_transport"), contextvars.ContextVar(var_name), is_async=is_async,
)
getattr(client, "__dict__", {}).get("_transport"), contextvars.ContextVar(var_name), is_async=is_async)
def create_ssrf_safe_async_client(**kwargs: Any) -> Any:
"""Create an ``httpx.AsyncClient`` with connect-time SSRF validation.
Direct HTTP(S) connections are resolved, validated, and dialed by IP at
TCP-connect time while the request hostname is preserved for Host, SNI, and
certificate verification. Proxied requests delegate resolution to the proxy.
"""
"""``httpx.AsyncClient`` with connect-time SSRF validation: direct HTTP(S) connections are
resolved, validated, and dialed by IP while the hostname is preserved for Host, SNI, and
certificate verification. Proxied requests delegate resolution to the proxy."""
import httpx
client = httpx.AsyncClient(**kwargs)
_install_ssrf_guard_on_client(client, is_async=True)
return client
@@ -572,30 +480,19 @@ def create_ssrf_safe_async_client(**kwargs: Any) -> Any:
def create_ssrf_safe_client(**kwargs: Any) -> Any:
"""Create an ``httpx.Client`` with connect-time SSRF validation."""
import httpx
client = httpx.Client(**kwargs)
_install_ssrf_guard_on_client(client)
return client
def redirect_target_from_response(response: Any) -> Optional[str]:
"""Return the redirect target visible from inside an httpx response hook.
``response.next_request`` is frequently ``None`` inside hooks (populated later
by the redirect follower), which would make an SSRF redirect guard silently
never fire. Resolve from the ``Location`` header first (relative via
``urljoin``), falling back to ``next_request``.
"""
"""Redirect target visible from inside an httpx response hook. ``response.next_request`` is
frequently ``None`` inside hooks (populated later by the follower), which would make an SSRF
redirect guard silently never fire — so resolve from ``Location`` first, then ``next_request``."""
if not getattr(response, "is_redirect", False):
return None
headers = getattr(response, "headers", {}) or {}
location = headers.get("location")
location = (getattr(response, "headers", {}) or {}).get("location")
if location:
return urljoin(str(getattr(response, "url", "")), str(location))
next_request = getattr(response, "next_request", None)
if next_request:
return str(next_request.url)
return None
return str(next_request.url) if next_request else None

View File

@@ -3,9 +3,8 @@
(``agent/video_gen_provider.py`` ABC, ``agent/video_gen_registry.py``, ``plugins/video_gen/<name>/``).
Ships **no in-tree provider**: enable a plugin and select it in ``hermes tools`` → Video
Generation. Covers text-, image- and reference-to-video; the tool layer only does lightweight
validation and each provider clamps/ignores unsupported params inside ``generate``, so the
surface stays stable as providers ship. Video edit/extend are deliberately not exposed here.
Generation. The tool layer only does lightweight validation; each provider clamps/ignores
unsupported params inside ``generate``. Video edit/extend are deliberately not exposed here.
"""
from __future__ import annotations
@@ -19,8 +18,7 @@ from agent.video_gen_provider import (
COMMON_RESOLUTIONS,
DEFAULT_ASPECT_RATIO,
DEFAULT_RESOLUTION,
error_response,
)
error_response)
from tools.registry import registry, tool_error
logger = logging.getLogger(__name__)
@@ -75,23 +73,15 @@ VIDEO_GENERATE_SCHEMA: Dict[str, Any] = {
}
# ---------------------------------------------------------------------------
# Config readers (mirror image_generation_tool.py)
# ---------------------------------------------------------------------------
def _read_video_gen_key(key: str) -> Optional[str]:
"""Return the stripped ``video_gen.<key>`` string from config.yaml, or None."""
try:
from hermes_cli.config import load_config
cfg = load_config()
section = cfg.get("video_gen") if isinstance(cfg, dict) else None
value = section.get(key) if isinstance(section, dict) else None
from hermes_cli.config import cfg_get, load_config
value = cfg_get(load_config(), "video_gen", key)
except Exception as exc:
logger.debug("Could not read video_gen config: %s", exc)
return None
if isinstance(value, str) and value.strip():
return value.strip()
return None
return value.strip() if isinstance(value, str) and value.strip() else None
def _read_configured_video_provider() -> Optional[str]:
@@ -106,32 +96,21 @@ def _discovered_registry():
"""Import the provider registry after (idempotent) plugin discovery so user-installed plugins are visible."""
from agent import video_gen_registry
from hermes_cli.plugins import _ensure_plugins_discovered
_ensure_plugins_discovered()
return video_gen_registry, _ensure_plugins_discovered
# ---------------------------------------------------------------------------
# Availability check + provider resolution
# ---------------------------------------------------------------------------
def check_video_generation_requirements() -> bool:
"""True when at least one registered provider reports available."""
try:
registry_mod, _ = _discovered_registry()
for provider in registry_mod.list_providers():
try:
if provider.is_available():
return True
except Exception:
continue
return any(_provider_call(p, "is_available", False) for p in registry_mod.list_providers())
except Exception:
pass
return False
return False
def _resolve_active_provider():
"""Active provider or None; forces a discovery refresh on a miss (long-lived sessions
that started before a plugin was installed)."""
"""Active provider or None; a miss forces a discovery refresh (sessions older than the plugin install)."""
try:
registry_mod, ensure_discovered = _discovered_registry()
provider = registry_mod.get_active_provider()
@@ -147,32 +126,23 @@ def _resolve_active_provider():
def _missing_provider_error(configured: Optional[str]) -> str:
if configured:
return json.dumps(error_response(
error=(
f"video_gen.provider='{configured}' is set but no plugin "
f"registered that name. Run `hermes plugins list` to see "
f"installed video gen backends, or `hermes tools` → Video "
f"Generation to pick one."
),
error_type="provider_not_registered",
provider=configured,
))
error=(f"video_gen.provider='{configured}' is set but no plugin registered that name. "
f"Run `hermes plugins list` to see installed video gen backends, or "
f"`hermes tools` → Video Generation to pick one."),
error_type="provider_not_registered", provider=configured))
return json.dumps(error_response(
error=(
"No video generation backend is configured. Run `hermes tools` → "
"Video Generation to enable one (xAI, FAL, or Google Veo)."
),
error_type="no_provider_configured",
))
error=("No video generation backend is configured. Run `hermes tools` → "
"Video Generation to enable one (xAI, FAL, or Google Veo)."),
error_type="no_provider_configured"))
_BOOL_WORDS = {"true": True, "1": True, "yes": True, "on": True,
"false": False, "0": False, "no": False, "off": False}
# ---------------------------------------------------------------------------
# Handler
# ---------------------------------------------------------------------------
def _coerce_int(value: Any) -> Optional[int]:
if value is None or value == "":
return None
try:
return int(value)
return None if value is None or value == "" else int(value)
except (TypeError, ValueError):
return None
@@ -180,21 +150,15 @@ def _coerce_int(value: Any) -> Optional[int]:
def _coerce_bool(value: Any) -> Optional[bool]:
if isinstance(value, bool):
return value
if isinstance(value, str):
return {"true": True, "1": True, "yes": True, "on": True,
"false": False, "0": False, "no": False, "off": False}.get(value.strip().lower())
return None
return _BOOL_WORDS.get(value.strip().lower()) if isinstance(value, str) else None
def _normalize_reference_images(value: Any) -> Optional[List[str]]:
if value is None:
return None
if isinstance(value, str):
value = [value]
if not isinstance(value, (list, tuple)):
return None
out = [item.strip() for item in value if isinstance(item, str) and item.strip()]
return out or None
return [item.strip() for item in value if isinstance(item, str) and item.strip()] or None
def _handle_video_generate(args: Dict[str, Any], **_kw: Any) -> str:
@@ -205,29 +169,28 @@ def _handle_video_generate(args: Dict[str, Any], **_kw: Any) -> str:
# Confinement chokepoint (mirrors image_generate): non-local backends hand providers data: URLs.
from tools.image_generation_tool import _confine_source_images
image_url, reference_image_urls, confine_error = _confine_source_images(
image_url, reference_image_urls, task_id)
if confine_error is not None:
return confine_error
duration = _coerce_int(args.get("duration"))
aspect_ratio = (args.get("aspect_ratio") or DEFAULT_ASPECT_RATIO).strip() or DEFAULT_ASPECT_RATIO
resolution = (args.get("resolution") or DEFAULT_RESOLUTION).strip() or DEFAULT_RESOLUTION
negative_prompt = (args.get("negative_prompt") or "").strip() or None
audio = _coerce_bool(args.get("audio"))
seed = _coerce_int(args.get("seed"))
upscale = _coerce_bool(args.get("upscale"))
# Coerced BEFORE validation (ordering parity: a bad value raises before a missing prompt).
optional = {
"duration": _coerce_int(args.get("duration")),
"aspect_ratio": (args.get("aspect_ratio") or DEFAULT_ASPECT_RATIO).strip() or DEFAULT_ASPECT_RATIO,
"resolution": (args.get("resolution") or DEFAULT_RESOLUTION).strip() or DEFAULT_RESOLUTION,
"negative_prompt": (args.get("negative_prompt") or "").strip() or None,
"audio": _coerce_bool(args.get("audio")),
"seed": _coerce_int(args.get("seed")),
"upscale": _coerce_bool(args.get("upscale"))}
model_override = (args.get("model") or "").strip() or None
# Soft validation — providers do their own; a backend may accept image-only, our surface never does.
# Soft validation — providers do their own; our surface never accepts image-only.
if not prompt:
return tool_error("prompt is required for video generation")
if "operation" in args or "video_url" in args:
return tool_error(
"video_generate only supports text-to-video, image-to-video, and "
"reference-to-video; use a provider-specific tool for video edit/extend"
)
"reference-to-video; use a provider-specific tool for video edit/extend")
configured = _read_configured_video_provider()
provider = _resolve_active_provider()
if provider is None:
@@ -235,20 +198,9 @@ def _handle_video_generate(args: Dict[str, Any], **_kw: Any) -> str:
# Explicit arg wins, then config, then provider default.
model = model_override or _read_configured_video_model() or provider.default_model()
kwargs: Dict[str, Any] = {
"model": model,
"_model_override_explicit": bool(model_override),
"image_url": image_url,
"reference_image_urls": reference_image_urls,
"duration": duration,
"aspect_ratio": aspect_ratio,
"resolution": resolution,
"negative_prompt": negative_prompt,
"audio": audio,
"seed": seed,
"upscale": upscale,
}
"model": model, "_model_override_explicit": bool(model_override),
"image_url": image_url, "reference_image_urls": reference_image_urls, **optional}
# Drop None entries so providers see clean defaults.
kwargs = {k: v for k, v in kwargs.items() if v is not None}
pname = getattr(provider, "name", "?")
@@ -256,42 +208,28 @@ def _handle_video_generate(args: Dict[str, Any], **_kw: Any) -> str:
def _err(error: str, error_type: str) -> str:
return json.dumps(error_response(
error=error, error_type=error_type,
provider=getattr(provider, "name", ""), model=model or "", prompt=prompt,
))
provider=getattr(provider, "name", ""), model=model or "", prompt=prompt))
try:
result = provider.generate(prompt=prompt, **kwargs)
except TypeError as exc:
# An un-widened provider signature is a plugin bug, not a caller error.
logger.warning(
"video_gen provider '%s' rejected kwargs (signature too narrow): %s",
pname, exc,
)
logger.warning("video_gen provider '%s' rejected kwargs (signature too narrow): %s", pname, exc)
return _err(
f"Provider '{pname}' signature is "
f"out of date with the video_generate schema. Report this "
f"to the plugin author.",
"provider_contract",
)
f"Provider '{pname}' signature is out of date with the video_generate schema. "
f"Report this to the plugin author.",
"provider_contract")
except Exception as exc:
logger.warning("video_gen provider '%s' raised: %s", pname, exc)
return _err(f"Provider '{pname}' error: {exc}", "provider_exception")
if not isinstance(result, dict):
return _err("Provider returned a non-dict result", "provider_contract")
return json.dumps(result)
# ---------------------------------------------------------------------------
# Dynamic schema — reflect the active backend's actual capabilities
# ---------------------------------------------------------------------------
# Surfacing the per-model surface (modalities, enums, durations, audio/negative-prompt)
# means the model usually gets the call right first try. model_tools.get_tool_definitions()
# keys its cache on config.yaml mtime, so the schema rebuilds on provider/model change.
# Optional params advertised only when the provider's capabilities() sets the flag
# (order = schema property order).
# Dynamic schema — reflects the active backend's actual capabilities so the model usually gets
# the call right first try. model_tools.get_tool_definitions() keys its cache on config.yaml
# mtime, so the schema rebuilds on provider/model change. Optional params below are advertised
# only when the provider's capabilities() sets the flag (order = schema property order).
_CAPABILITY_PARAMS = (
("supports_negative_prompt", "negative_prompt", {
"type": "string",
@@ -334,34 +272,31 @@ _GENERIC_DESCRIPTION = (
def _schema(description: str, properties: Dict[str, Any]) -> Dict[str, Any]:
return {
"description": description,
"parameters": {"type": "object", "properties": properties, "required": ["prompt"]},
}
"parameters": {"type": "object", "properties": properties, "required": ["prompt"]}}
def _provider_call(provider: Any, method: str, default: Any) -> Any:
"""``provider.<method>()`` or ``default`` when it raises or returns a falsy value."""
try:
return getattr(provider, method)() or default
except Exception:
return default
def _build_dynamic_video_schema() -> Dict[str, Any]:
"""Render description AND params from capabilities() + the model's catalog entry; enums and
duration bounds tighten to the active model. Unadvertised args are still accepted (replay compat)."""
"""Description AND params from capabilities() + the model's catalog entry; enums and duration
bounds tighten to the active model. Unadvertised args are still accepted (replay compat)."""
static_props = VIDEO_GENERATE_SCHEMA["parameters"]["properties"]
parts: List[str] = [_GENERIC_DESCRIPTION]
configured_model = _read_configured_video_model()
provider = _resolve_active_provider()
if provider is None:
parts.append(
"\nNo video backend is available. Calls will return an error "
"until the user picks one via `hermes tools` → Video Generation."
)
"until the user picks one via `hermes tools` → Video Generation.")
return _schema("\n".join(parts), {"prompt": static_props["prompt"]})
try:
caps = provider.capabilities() or {}
except Exception:
caps = {}
try:
models = provider.list_models() or []
except Exception:
models = []
caps = _provider_call(provider, "capabilities", {})
models = _provider_call(provider, "list_models", [])
active_model = configured_model or provider.default_model()
model_meta = next((m for m in models if isinstance(m, dict) and m.get("id") == active_model), {})
@@ -374,11 +309,9 @@ def _build_dynamic_video_schema() -> Dict[str, Any]:
if "image" in model_modalities and "text" not in model_modalities:
parts.append(
"- this model is image-to-video only — image_url is REQUIRED; "
"text-only calls will be rejected"
)
"text-only calls will be rejected")
elif "text" in model_modalities and "image" not in model_modalities:
parts.append("- this model is text-to-video only — image_url is not supported")
effective_modalities = model_modalities or set(caps.get("modalities") or [])
can_i2v = "image" in effective_modalities
t2v = "text" in effective_modalities
@@ -386,33 +319,26 @@ def _build_dynamic_video_schema() -> Dict[str, Any]:
parts.append("- image-to-video only: image_url is REQUIRED")
elif not can_i2v:
parts.append("- text-to-video only (no image input)")
if provider.name == "xai":
parts.append(
"- chaining: for edit/extend pass the public HTTPS MP4 in `video` "
"or `public_url` from the prior Imagine result (files-cdn). For "
"image-to-video / reference-to-video pass public image URLs the "
"same way"
)
"same way")
try:
from tools.xai_http import xai_storage_notice_text
notice = xai_storage_notice_text("video_gen")
except Exception:
notice = ""
if notice:
parts.append(f"- storage: {notice}")
properties: Dict[str, Any] = {"prompt": static_props["prompt"]}
if can_i2v:
properties["image_url"] = {
"type": "string",
"description": (
"Public HTTPS URL of a still image to animate "
"(image-to-video). Omit for text-to-video."
),
}
"(image-to-video). Omit for text-to-video.")}
max_refs = int(caps.get("max_reference_images") or 0)
if max_refs > 0:
properties["reference_image_urls"] = {
@@ -421,10 +347,7 @@ def _build_dynamic_video_schema() -> Dict[str, Any]:
"maxItems": max_refs,
"description": (
f"Up to {max_refs} public HTTPS reference image URLs "
"(style or character refs)."
),
}
"(style or character refs).")}
min_duration = model_meta.get("min_duration", caps.get("min_duration"))
max_duration = model_meta.get("max_duration", caps.get("max_duration"))
duration_param = dict(static_props["duration"])
@@ -433,8 +356,7 @@ def _build_dynamic_video_schema() -> Dict[str, Any]:
duration_param["maximum"] = int(max_duration)
duration_param["description"] = (
f"Video duration in seconds ({min_duration}-{max_duration}). "
"Omit for the provider default."
)
"Omit for the provider default.")
properties["duration"] = duration_param
# Tighten enums to the active backend's actual sets when declared.
@@ -443,7 +365,6 @@ def _build_dynamic_video_schema() -> Dict[str, Any]:
if caps.get(caps_key):
param["enum"] = list(caps[caps_key])
properties[key] = param
for flag, key, param in _CAPABILITY_PARAMS:
if caps.get(flag):
properties[key] = param
@@ -451,8 +372,7 @@ def _build_dynamic_video_schema() -> Dict[str, Any]:
parts.append(
"- audio: native stereo audio is generated with every video "
"(always on; no toggle) — describe the desired sound in the "
"prompt"
)
"prompt")
properties["model"] = static_props["model"]
return _schema("\n".join(parts), properties)
@@ -466,5 +386,4 @@ registry.register(
requires_env=[],
is_async=False,
emoji="🎬",
dynamic_schema_overrides=_build_dynamic_video_schema,
)
dynamic_schema_overrides=_build_dynamic_video_schema)

File diff suppressed because it is too large Load Diff

View File

@@ -1,9 +1,8 @@
"""Image format detection, normalization and region cropping for vision tools.
Everything here runs BEFORE an image is base64-embedded. A vision tool result
is baked into immutable conversation history and re-sent every turn, so an
unsupported media type or corrupt bytes would wedge the session with a
non-retryable 400 on every resume — normalization must happen up front.
Everything here runs BEFORE an image is base64-embedded: a vision tool result is
baked into immutable history and re-sent every turn, so an unsupported media type
or corrupt bytes would wedge the session with a non-retryable 400 on every resume.
"""
from __future__ import annotations
@@ -30,8 +29,11 @@ _EXTENSION_MIME_TYPES = {
# Media types the major vision providers (Anthropic in particular) accept
# inline. SVG/BMP/TIFF are rejected with a non-retryable 400.
_ANTHROPIC_SUPPORTED_MEDIA_TYPES = frozenset(
{"image/jpeg", "image/png", "image/gif", "image/webp"}
_ANTHROPIC_SUPPORTED_MEDIA_TYPES = frozenset({"image/jpeg", "image/png", "image/gif", "image/webp"})
_MAGIC_MIME_TYPES = (
(b"\xff\xd8\xff", "image/jpeg"), ((b"GIF87a", b"GIF89a"), "image/gif"), (b"BM", "image/bmp"),
)
@@ -41,17 +43,12 @@ def _determine_mime_type(image_path: Path) -> str:
def _detect_image_mime_type_from_bytes(data: bytes) -> Optional[str]:
"""Magic-byte MIME sniff (authoritative; no extension trust).
Returns ``None`` for anything without a recognized header — including SVG,
which has no magic bytes (the resolver sniffs ``<svg`` and passes it
through for rasterization).
"""
"""Magic-byte MIME sniff (authoritative; no extension trust). ``None`` for anything without a
recognized header — including SVG, which has none (the resolver sniffs ``<svg`` itself)."""
header = data[:64]
if header.startswith(b"\x89PNG\r\n\x1a\n"):
# Magic bytes alone are insufficient: reject corrupt PNGs before they
# can be embedded. Pillow is optional — without it fall back to
# header-only sniffing; only an actual failed verify() rejects.
# Reject corrupt PNGs before they can be embedded. Pillow is optional —
# without it fall back to header-only sniffing; only a failed verify() rejects.
try:
from PIL import Image
except ImportError:
@@ -59,52 +56,41 @@ def _detect_image_mime_type_from_bytes(data: bytes) -> Optional[str]:
try:
with Image.open(BytesIO(data)) as image:
image.verify()
return "image/png"
except Exception:
return None
return "image/png"
if header.startswith(b"\xff\xd8\xff"):
return "image/jpeg"
if header.startswith((b"GIF87a", b"GIF89a")):
return "image/gif"
if header.startswith(b"BM"):
return "image/bmp"
for magic, mime in _MAGIC_MIME_TYPES:
if header.startswith(magic):
return mime
if len(header) >= 12 and header[:4] == b"RIFF" and header[8:12] == b"WEBP":
return "image/webp"
return None
def _supported_media_types() -> frozenset:
"""Formats the ACTIVE main model's server can decode.
The managed llama-server decodes with stb_image — no WebP — and an
undecodable image part fails SILENTLY (the model confabulates), so the set
is narrowed there and normalization converts those formats to PNG.
"""
"""Formats the ACTIVE main model's server can decode. The managed llama-server decodes with
stb_image — no WebP — and an undecodable image part fails SILENTLY (the model confabulates),
so the set is narrowed there and normalization converts those formats to PNG."""
try:
from agent.auxiliary_client import _runtime_main_value
from hermes_cli.local_runtime.capabilities import (
ACCEPTED_IMAGE_MIMES,
is_managed_provider,
)
if is_managed_provider(
str(_runtime_main_value("provider") or ""),
str(_runtime_main_value("base_url") or "")):
from agent.auxiliary_client import _runtime_main_value as _v
from hermes_cli.local_runtime.capabilities import ACCEPTED_IMAGE_MIMES, is_managed_provider
if is_managed_provider(str(_v("provider") or ""), str(_v("base_url") or "")):
return ACCEPTED_IMAGE_MIMES
except Exception: # noqa: BLE001 — best-effort narrowing only
except Exception: # best-effort narrowing only
pass
return _ANTHROPIC_SUPPORTED_MEDIA_TYPES
def _nonempty_file(path: Path) -> bool:
return path.exists() and path.stat().st_size > 0
def _rasterize_svg_to_png(svg_path: Path, out_path: Path) -> bool:
"""Best-effort SVG → PNG via cairosvg, svglib+reportlab, rsvg-convert, inkscape (all soft deps)."""
def _ok() -> bool:
return out_path.exists() and out_path.stat().st_size > 0
try:
import cairosvg # type: ignore
cairosvg.svg2png(url=str(svg_path), write_to=str(out_path))
return _ok()
return _nonempty_file(out_path)
except Exception:
pass
try:
@@ -113,23 +99,18 @@ def _rasterize_svg_to_png(svg_path: Path, out_path: Path) -> bool:
drawing = svg2rlg(str(svg_path))
if drawing is not None:
renderPM.drawToFile(drawing, str(out_path), fmt="PNG")
return _ok()
return _nonempty_file(out_path)
except Exception:
pass
import shutil
import subprocess
for cmd in (
["rsvg-convert", "-o", str(out_path), str(svg_path)],
["inkscape", str(svg_path), "--export-type=png",
f"--export-filename={out_path}"],
):
["inkscape", str(svg_path), "--export-type=png", f"--export-filename={out_path}"]):
if shutil.which(cmd[0]):
try:
subprocess.run(
cmd, check=True, capture_output=True, timeout=30,
stdin=subprocess.DEVNULL,
)
if _ok():
subprocess.run(cmd, check=True, capture_output=True, timeout=30, stdin=subprocess.DEVNULL)
if _nonempty_file(out_path):
return True
except Exception:
continue
@@ -137,58 +118,43 @@ def _rasterize_svg_to_png(svg_path: Path, out_path: Path) -> bool:
def _normalize_to_supported_image(
image_path: Path, detected_mime: str
) -> tuple[Optional[Path], Optional[str], Optional[str]]:
"""Ensure an image is in a provider-supported format.
Returns ``(path, mime, error)``: the input unchanged when already
supported; ``(new_png_path, "image/png", None)`` after conversion — a temp
file the CALLER must clean up; ``(None, None, message)`` when impossible.
SVG is rasterized; other Pillow-readable rasters (BMP, TIFF) re-encode to PNG.
"""
image_path: Path, detected_mime: str) -> tuple[Optional[Path], Optional[str], Optional[str]]:
"""Ensure an image is in a provider-supported format. Returns ``(path, mime, error)``: the input
unchanged when supported; ``(new_png_path, "image/png", None)`` after conversion — a temp file
the CALLER must clean up; ``(None, None, message)`` when impossible. SVG is rasterized; other
Pillow-readable rasters (BMP, TIFF) re-encode to PNG."""
if detected_mime in _supported_media_types():
return image_path, detected_mime, None
out_dir = get_hermes_dir("cache/vision", "temp_vision_images")
out_dir.mkdir(parents=True, exist_ok=True)
out_path = out_dir / f"converted_{uuid.uuid4()}.png"
if detected_mime == "image/svg+xml":
if _rasterize_svg_to_png(image_path, out_path):
return out_path, "image/png", None
return (
None,
None,
return None, None, (
"This is an SVG, which vision models cannot read directly, and no "
"SVG rasterizer is installed (tried cairosvg, svglib, rsvg-convert, "
"inkscape). Convert the SVG to PNG first — e.g. open it in a browser "
"and screenshot it, or install a rasterizer "
"(`pip install cairosvg`) — then re-run vision_analyze on the PNG.",
)
"(`pip install cairosvg`) — then re-run vision_analyze on the PNG.")
try:
from PIL import Image as _PILImage
with _PILImage.open(image_path) as _img:
if _img.mode not in ("RGB", "RGBA", "L"):
_img = _img.convert("RGBA")
_img.save(out_path, format="PNG")
if out_path.exists() and out_path.stat().st_size > 0:
if _nonempty_file(out_path):
return out_path, "image/png", None
except Exception as _exc:
logger.warning("Failed to normalize %s image to PNG: %s",
detected_mime, _exc)
return (
None,
None,
logger.warning("Failed to normalize %s image to PNG: %s", detected_mime, _exc)
return None, None, (
f"Image format {detected_mime!r} is not supported by the vision API "
f"and could not be converted to PNG (install Pillow for raster "
f"conversion). Convert it to PNG or JPEG and try again.",
)
f"conversion). Convert it to PNG or JPEG and try again.")
# Full raster validation runs on untrusted images in a shared CPU executor:
# bound animated-image work by frame count AND total decoded area so a compact
# file cannot monopolize a worker with unbounded frames.
# Full raster validation runs on untrusted images in a shared CPU executor: bound animated
# work by frame count AND total decoded area so a compact file cannot monopolize a worker.
_VISION_MAX_VALIDATED_FRAME_COUNT = 100
_VISION_MAX_VALIDATED_AGGREGATE_PIXELS = 100_000_000
@@ -196,17 +162,12 @@ _VISION_MAX_VALIDATED_AGGREGATE_PIXELS = 100_000_000
def _validate_raster_image_decodable(
image_path: Path,
max_frames: int = _VISION_MAX_VALIDATED_FRAME_COUNT,
max_pixels: int = _VISION_MAX_VALIDATED_AGGREGATE_PIXELS,
) -> Optional[str]:
"""Return an error unless Pillow can fully decode every frame.
Header sniffing and ``Image.open`` only inspect containers: a timed-out
download can look like a valid PNG with a truncated pixel stream. Without
Pillow the image passes unvalidated rather than rejecting everything.
"""
max_pixels: int = _VISION_MAX_VALIDATED_AGGREGATE_PIXELS) -> Optional[str]:
"""Return an error unless Pillow can fully decode every frame. Header sniffing and ``Image.open``
only inspect containers: a timed-out download can look like a valid PNG with a truncated pixel
stream. Without Pillow the image passes unvalidated rather than rejecting everything."""
try:
from PIL import Image as _PILImage
from PIL import ImageSequence as _PILImageSequence
from PIL import Image as _PILImage, ImageSequence as _PILImageSequence
except ImportError:
return None
try:
@@ -214,23 +175,19 @@ def _validate_raster_image_decodable(
image.verify()
with _PILImage.open(image_path) as image:
validated_pixels = 0
for frame_number, frame in enumerate(
_PILImageSequence.Iterator(image), start=1
):
for frame_number, frame in enumerate(_PILImageSequence.Iterator(image), start=1):
if frame_number > max_frames:
return (
"Image validation rejected animation: "
f"frame {frame_number} exceeds the maximum "
f"{max_frames} validated frames."
)
f"{max_frames} validated frames.")
next_validated_pixels = validated_pixels + frame.width * frame.height
if next_validated_pixels > max_pixels:
return (
"Image validation rejected animation: aggregate decoded "
f"pixel count would reach {next_validated_pixels} at frame "
f"{frame_number}, exceeding the maximum "
f"{max_pixels}."
)
f"{max_pixels}.")
frame.load()
validated_pixels = next_validated_pixels
except Exception as exc:
@@ -239,12 +196,9 @@ def _validate_raster_image_decodable(
def _image_exceeds_dimension(image_path: Path, max_dimension: int) -> bool:
"""True if the longest side exceeds ``max_dimension`` px.
Anthropic enforces an 8000px per-side cap independently of the byte cap.
Returns False (no forced resize) without Pillow or on unreadable files —
a missing soft dependency must never break the embed path.
"""
"""True if the longest side exceeds ``max_dimension`` px (Anthropic's 8000px per-side cap is
independent of bytes). False without Pillow or on unreadable files — a missing soft
dependency must never break the embed path."""
try:
from PIL import Image as _PILImage
with _PILImage.open(image_path) as _img:
@@ -254,57 +208,39 @@ def _image_exceeds_dimension(image_path: Path, max_dimension: int) -> bool:
def _crop_image_region(
image_path: Path,
region: Any,
offset_out: Optional[dict] = None,
image_path: Path, region: Any, offset_out: Optional[dict] = None
) -> tuple[Optional[Path], Optional[str], Optional[str]]:
"""Crop to ``region`` = [x1, y1, x2, y2] (original-image pixels).
Applied BEFORE downscaling so the crop gets the full resolution budget.
Coordinates clamp to the image bounds; a zero-area/inverted region is
rejected with an error naming the real dimensions. Returns
``(cropped_temp_path, mime, None)`` — caller owns cleanup — or
``(None, None, error)``. Ported from QwenLM/qwen-code zoom-image.ts (Apache-2.0).
"""
"""Crop to ``region`` = [x1, y1, x2, y2] (original-image pixels), BEFORE downscaling so the crop
gets the full resolution budget. Coordinates clamp to the image bounds; a zero-area/inverted
region is rejected with an error naming the real dimensions. Returns ``(cropped_temp_path,
mime, None)`` — caller owns cleanup — or ``(None, None, error)``.
Ported from QwenLM/qwen-code zoom-image.ts (Apache-2.0)."""
try:
from PIL import Image
except ImportError:
return None, None, (
"region cropping requires Pillow (`pip install Pillow`); "
"retry without the region parameter."
)
if (
not isinstance(region, (list, tuple))
or len(region) != 4
or not all(isinstance(v, (int, float)) and not isinstance(v, bool) for v in region)
):
"retry without the region parameter.")
if not (isinstance(region, (list, tuple)) and len(region) == 4
and all(isinstance(v, (int, float)) and not isinstance(v, bool) for v in region)):
return None, None, (
"Invalid region: expected [x1, y1, x2, y2] as four numbers "
"(pixel coordinates in the original image)."
)
"(pixel coordinates in the original image).")
try:
with Image.open(image_path) as img:
width, height = img.size
x1, y1, x2, y2 = (int(v) for v in region)
cx1 = max(0, min(x1, width))
cy1 = max(0, min(y1, height))
cx2 = max(0, min(x2, width))
cy2 = max(0, min(y2, height))
cx1, cy1, cx2, cy2 = (max(0, min(v, b)) for v, b in zip((x1, y1, x2, y2), (width, height) * 2))
if cx2 <= cx1 or cy2 <= cy1:
return None, None, (
f"Invalid region [{x1}, {y1}, {x2}, {y2}]: crops to zero "
f"area after clamping to the image bounds. The image is "
f"{width}x{height} px — pick x1<x2 and y1<y2 inside "
f"[0, 0, {width}, {height}]."
)
f"[0, 0, {width}, {height}].")
cropped = img.crop((cx1, cy1, cx2, cy2))
if offset_out is not None:
offset_out.update(x=cx1, y=cy1, width=cx2 - cx1, height=cy2 - cy1)
out_path = image_path.with_name(
f"{image_path.stem}_region_{uuid.uuid4().hex[:8]}.png"
)
out_path = image_path.with_name(f"{image_path.stem}_region_{uuid.uuid4().hex[:8]}.png")
if cropped.mode not in ("RGB", "RGBA", "L", "LA", "P"):
cropped = cropped.convert("RGB")
cropped.save(out_path, format="PNG")