refactor(tools): vision/video/url_safety — inline single-use helpers, dedupe backends, compact docstrings
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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
@@ -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")
|
||||
|
||||
Reference in New Issue
Block a user