fix(security): pin DNS resolutions for SSRF-safe fetches

Install connect-time DNS validation for Hermes-owned direct httpx clients so SSRF-sensitive fetch paths dial a vetted IP instead of re-resolving after preflight. This preserves Host/SNI semantics for direct HTTP(S) connections and keeps proxy routing as an explicit trusted egress boundary.

Wire the guarded clients into media cache downloads, vision downloads, Skills Hub direct/raw fetches, and platform attachment fetch paths that already perform SSRF preflight and redirect validation.

Fixes #8033

Co-authored-by: Tom Qiao <zqiao@microsoft.com>
This commit is contained in:
Eugeniusz Gilewski
2026-07-03 18:39:14 +02:00
committed by Teknium
parent ca988df8d0
commit 42626da1ce
13 changed files with 581 additions and 49 deletions
+4 -4
View File
@@ -749,14 +749,14 @@ async def cache_image_from_url(url: str, ext: str = ".jpg", retries: int = 2) ->
Raises:
ValueError: If the URL targets a private/internal network (SSRF protection).
"""
from tools.url_safety import is_safe_url
from tools.url_safety import create_ssrf_safe_async_client, is_safe_url
if not is_safe_url(url):
raise ValueError(f"Blocked unsafe URL (SSRF protection): {safe_url_for_log(url)}")
import httpx
_log = logging.getLogger(__name__)
async with httpx.AsyncClient(
async with create_ssrf_safe_async_client(
timeout=30.0,
follow_redirects=True,
event_hooks={"response": [_ssrf_redirect_guard]},
@@ -869,14 +869,14 @@ async def cache_audio_from_url(url: str, ext: str = ".ogg", retries: int = 2) ->
Raises:
ValueError: If the URL targets a private/internal network (SSRF protection).
"""
from tools.url_safety import is_safe_url
from tools.url_safety import create_ssrf_safe_async_client, is_safe_url
if not is_safe_url(url):
raise ValueError(f"Blocked unsafe URL (SSRF protection): {safe_url_for_log(url)}")
import httpx
_log = logging.getLogger(__name__)
async with httpx.AsyncClient(
async with create_ssrf_safe_async_client(
timeout=30.0,
follow_redirects=True,
event_hooks={"response": [_ssrf_redirect_guard]},
+2 -1
View File
@@ -312,7 +312,8 @@ class QQAdapter(BasePlatformAdapter):
# Tighter keepalive pool so idle CLOSE_WAIT sockets drain
# faster behind proxies like Cloudflare Warp (#18451).
from gateway.platforms._http_client_limits import platform_httpx_limits
self._http_client = httpx.AsyncClient(
from tools.url_safety import create_ssrf_safe_async_client
self._http_client = create_ssrf_safe_async_client(
timeout=30.0,
follow_redirects=True,
event_hooks={"response": [_ssrf_redirect_guard]},
+2 -2
View File
@@ -220,7 +220,7 @@ async def download_url(
# SSRF protection: yuanbao downloads model-supplied and inbound URLs
# server-side. Reject private/internal targets up front, and re-validate
# every redirect hop so a public URL can't 302 to http://169.254.169.254/.
from tools.url_safety import is_safe_url
from tools.url_safety import create_ssrf_safe_async_client, is_safe_url
if not is_safe_url(url):
raise ValueError(f"Blocked unsafe URL (SSRF protection): {url}")
@@ -234,7 +234,7 @@ async def download_url(
)
max_bytes = max_size_mb * 1024 * 1024
async with httpx.AsyncClient(
async with create_ssrf_safe_async_client(
timeout=30.0,
follow_redirects=True,
event_hooks={"response": [_redirect_guard]},
+2 -2
View File
@@ -1937,13 +1937,13 @@ class MatrixAdapter(BasePlatformAdapter):
return b"".join(parts), ct, fname
raise ValueError("too many redirects")
except ImportError:
import httpx
from tools.url_safety import create_ssrf_safe_async_client
_httpx_kw: dict = {}
if self._proxy_url:
_httpx_kw["proxy"] = self._proxy_url
_httpx_kw["event_hooks"] = {"response": [_ssrf_redirect_guard]}
async with httpx.AsyncClient(**_httpx_kw) as http:
async with create_ssrf_safe_async_client(**_httpx_kw) as http:
async with http.stream(
"GET",
url,
+2 -4
View File
@@ -3554,7 +3554,7 @@ class SlackAdapter(BasePlatformAdapter):
if not self._app:
return SendResult(success=False, error="Not connected")
from tools.url_safety import is_safe_url
from tools.url_safety import create_ssrf_safe_async_client, is_safe_url
if not is_safe_url(image_url):
logger.warning("[Slack] Blocked unsafe image URL (SSRF protection)")
@@ -3563,8 +3563,6 @@ class SlackAdapter(BasePlatformAdapter):
)
try:
import httpx
async def _ssrf_redirect_guard(response):
"""Re-check redirect targets so public URLs cannot bounce into private IPs."""
from tools.url_safety import redirect_target_from_response
@@ -3573,7 +3571,7 @@ class SlackAdapter(BasePlatformAdapter):
raise ValueError("Blocked redirect to private/internal address")
# Download the image first
async with httpx.AsyncClient(
async with create_ssrf_safe_async_client(
timeout=30.0,
follow_redirects=True,
event_hooks={"response": [_ssrf_redirect_guard]},
+2 -4
View File
@@ -812,15 +812,13 @@ class TeamsAdapter(BasePlatformAdapter):
SSRF guard and follows redirects through the shared redirect guard,
matching the cache_*_from_url helpers in gateway.platforms.base.
"""
from tools.url_safety import is_safe_url
from tools.url_safety import create_ssrf_safe_async_client, is_safe_url
from gateway.platforms.base import _ssrf_redirect_guard
if not is_safe_url(url):
raise ValueError("Blocked unsafe attachment URL (SSRF protection)")
import httpx
async with httpx.AsyncClient(
async with create_ssrf_safe_async_client(
timeout=timeout,
follow_redirects=True,
event_hooks={"response": [_ssrf_redirect_guard]},
@@ -12,6 +12,7 @@ in this environment.
"""
import asyncio
import socket
import sys
from unittest.mock import AsyncMock, MagicMock, patch
@@ -235,6 +236,54 @@ class TestCacheImageFromUrl:
mock_sleep.assert_not_called()
class TestCacheImageFromUrlConnectGuard:
def test_blocks_private_dns_answer_at_connect_time(self, tmp_path, monkeypatch):
"""A hostname that rebinds after preflight must not reach TCP connect."""
monkeypatch.setattr("gateway.platforms.base.IMAGE_CACHE_DIR", tmp_path / "img")
for proxy_var in (
"HTTP_PROXY",
"HTTPS_PROXY",
"ALL_PROXY",
"http_proxy",
"https_proxy",
"all_proxy",
):
monkeypatch.delenv(proxy_var, raising=False)
answers = [
[(socket.AF_INET, socket.SOCK_STREAM, 6, "", ("93.184.216.34", 80))],
[(socket.AF_INET, socket.SOCK_STREAM, 6, "", ("169.254.169.254", 80))],
]
def fake_getaddrinfo(host, port, *args, **kwargs):
assert host == "rebind.test"
return answers.pop(0)
from httpcore._backends.auto import AutoBackend
async def fail_connect_tcp(
self,
host,
port,
timeout=None,
local_address=None,
socket_options=None,
):
raise AssertionError(f"TCP connect attempted for {host}:{port}")
monkeypatch.setattr(socket, "getaddrinfo", fake_getaddrinfo)
monkeypatch.setattr(AutoBackend, "connect_tcp", fail_connect_tcp)
async def run():
from gateway.platforms.base import cache_image_from_url
await cache_image_from_url("http://rebind.test/image.jpg", ext=".jpg", retries=0)
with pytest.raises(ValueError, match="during connect"):
asyncio.run(run())
assert answers == []
# ---------------------------------------------------------------------------
# cache_audio_from_url (base.py)
# ---------------------------------------------------------------------------
+33 -20
View File
@@ -802,7 +802,7 @@ class TestWellKnownSkillSource:
@patch("tools.skills_hub._write_index_cache")
@patch("tools.skills_hub._read_index_cache", return_value=None)
@patch("tools.skills_hub.httpx.get")
@patch("tools.skills_hub._ssrf_safe_http_get")
def test_search_reads_index_from_well_known_url(self, mock_get, _mock_read_cache, _mock_write_cache):
mock_get.return_value = MagicMock(
status_code=200,
@@ -824,7 +824,7 @@ class TestWellKnownSkillSource:
@patch("tools.skills_hub._write_index_cache")
@patch("tools.skills_hub._read_index_cache", return_value=None)
@patch("tools.skills_hub.httpx.get")
@patch("tools.skills_hub._ssrf_safe_http_get")
def test_search_accepts_domain_root_and_resolves_index(self, mock_get, _mock_read_cache, _mock_write_cache):
mock_get.return_value = MagicMock(
status_code=200,
@@ -839,7 +839,7 @@ class TestWellKnownSkillSource:
@patch("tools.skills_hub._write_index_cache")
@patch("tools.skills_hub._read_index_cache", return_value=None)
@patch("tools.skills_hub.httpx.get")
@patch("tools.skills_hub._ssrf_safe_http_get")
def test_inspect_fetches_skill_md_from_well_known_endpoint(self, mock_get, _mock_read_cache, _mock_write_cache):
def fake_get(url, *args, **kwargs):
if url.endswith("/index.json"):
@@ -861,7 +861,7 @@ class TestWellKnownSkillSource:
@patch("tools.skills_hub._write_index_cache")
@patch("tools.skills_hub._read_index_cache", return_value=None)
@patch("tools.skills_hub.httpx.get")
@patch("tools.skills_hub._ssrf_safe_http_get")
def test_fetch_downloads_skill_files_from_well_known_endpoint(self, mock_get, _mock_read_cache, _mock_write_cache):
def fake_get(url, *args, **kwargs):
if url.endswith("/index.json"):
@@ -954,7 +954,7 @@ class TestUrlSource:
assert self._source().search("anything") == []
# ── inspect ─────────────────────────────────────────────────────────
@patch("tools.skills_hub.httpx.get")
@patch("tools.skills_hub._ssrf_safe_http_get")
def test_inspect_reads_frontmatter_from_url(self, mock_get):
mock_get.return_value = MagicMock(
status_code=200,
@@ -978,31 +978,31 @@ class TestUrlSource:
assert meta.tags == ["sharing", "chat"]
assert meta.extra["awaiting_name"] is False
@patch("tools.skills_hub.httpx.get")
@patch("tools.skills_hub._ssrf_safe_http_get")
def test_inspect_returns_none_when_url_not_md(self, mock_get):
# _matches filters first — no HTTP call.
meta = self._source().inspect("https://example.com/not-a-skill")
assert meta is None
mock_get.assert_not_called()
@patch("tools.skills_hub.httpx.get")
@patch("tools.skills_hub._ssrf_safe_http_get")
def test_inspect_returns_none_on_404(self, mock_get):
mock_get.return_value = MagicMock(status_code=404)
assert self._source().inspect("https://example.com/SKILL.md") is None
@patch("tools.skills_hub.httpx.get")
@patch("tools.skills_hub._ssrf_safe_http_get")
def test_inspect_returns_none_on_http_error(self, mock_get):
mock_get.side_effect = httpx.HTTPError("boom")
assert self._source().inspect("https://example.com/SKILL.md") is None
@patch("tools.skills_hub.httpx.get")
@patch("tools.skills_hub._ssrf_safe_http_get")
@patch("tools.skills_hub.check_website_access", return_value=None)
@patch("tools.skills_hub.is_safe_url", return_value=False)
def test_inspect_blocks_private_url(self, _mock_safe, _mock_policy, mock_get):
assert self._source().inspect("http://127.0.0.1/SKILL.md") is None
mock_get.assert_not_called()
@patch("tools.skills_hub.httpx.get")
@patch("tools.skills_hub._ssrf_safe_http_get")
def test_inspect_flags_awaiting_name_when_unresolvable(self, mock_get):
# No frontmatter name + a URL path that can't produce a valid slug
# (``SKILL`` isn't a valid skill name).
@@ -1016,7 +1016,7 @@ class TestUrlSource:
assert meta.extra["awaiting_name"] is True
# ── fetch ───────────────────────────────────────────────────────────
@patch("tools.skills_hub.httpx.get")
@patch("tools.skills_hub._ssrf_safe_http_get")
def test_fetch_builds_single_file_bundle(self, mock_get):
skill_md = (
"---\n"
@@ -1037,7 +1037,7 @@ class TestUrlSource:
assert bundle.metadata["url"] == "https://sharethis.chat/SKILL.md"
assert bundle.metadata["awaiting_name"] is False
@patch("tools.skills_hub.httpx.get")
@patch("tools.skills_hub._ssrf_safe_http_get")
def test_fetch_falls_back_to_url_directory_name(self, mock_get):
# Frontmatter has no ``name:`` — we slug from the URL directory.
mock_get.return_value = MagicMock(
@@ -1049,7 +1049,7 @@ class TestUrlSource:
assert bundle.name == "my-skill"
assert bundle.metadata["awaiting_name"] is False
@patch("tools.skills_hub.httpx.get")
@patch("tools.skills_hub._ssrf_safe_http_get")
def test_fetch_falls_back_to_filename_when_no_parent_dir(self, mock_get):
mock_get.return_value = MagicMock(
status_code=200,
@@ -1060,7 +1060,7 @@ class TestUrlSource:
assert bundle.name == "my-skill"
assert bundle.metadata["awaiting_name"] is False
@patch("tools.skills_hub.httpx.get")
@patch("tools.skills_hub._ssrf_safe_http_get")
def test_fetch_awaiting_name_when_unresolvable(self, mock_get):
# Bare ``SKILL.md`` at the domain root with no frontmatter name.
mock_get.return_value = MagicMock(
@@ -1074,7 +1074,7 @@ class TestUrlSource:
# File content still present — CLI will reuse it after picking a name.
assert bundle.files["SKILL.md"].startswith("---\n")
@patch("tools.skills_hub.httpx.get")
@patch("tools.skills_hub._ssrf_safe_http_get")
def test_fetch_awaiting_name_rejects_sentinel_slug(self, mock_get):
# Frontmatter has no name AND the URL filename slug is ``README`` —
# our valid-name check rejects it, so we flag awaiting_name.
@@ -1087,7 +1087,7 @@ class TestUrlSource:
assert bundle.name == ""
assert bundle.metadata["awaiting_name"] is True
@patch("tools.skills_hub.httpx.get")
@patch("tools.skills_hub._ssrf_safe_http_get")
def test_fetch_ignores_unsafe_frontmatter_name_and_falls_through_to_slug(self, mock_get):
# Traversal / unsafe names are rejected by ``_is_valid_skill_name``;
# resolver falls through to URL slug (``my-skill`` here) and succeeds.
@@ -1099,12 +1099,12 @@ class TestUrlSource:
assert bundle is not None
assert bundle.name == "my-skill"
@patch("tools.skills_hub.httpx.get")
@patch("tools.skills_hub._ssrf_safe_http_get")
def test_fetch_returns_none_on_404(self, mock_get):
mock_get.return_value = MagicMock(status_code=404)
assert self._source().fetch("https://example.com/SKILL.md") is None
@patch("tools.skills_hub.httpx.get")
@patch("tools.skills_hub._ssrf_safe_http_get")
@patch("tools.skills_hub.check_website_access", return_value=None)
@patch("tools.skills_hub.is_safe_url", side_effect=[True, False])
def test_fetch_blocks_redirect_to_private_url(self, _mock_safe, _mock_policy, mock_get):
@@ -1115,14 +1115,27 @@ class TestUrlSource:
assert self._source().fetch("https://example.com/SKILL.md") is None
assert mock_get.call_count == 1
@patch("tools.skills_hub.httpx.get")
@patch("tools.skills_hub._ssrf_safe_http_get")
@patch("tools.skills_hub.check_website_access", return_value=None)
@patch("tools.skills_hub.is_safe_url", return_value=False)
def test_fetch_blocks_private_url(self, _mock_safe, _mock_policy, mock_get):
assert self._source().fetch("http://127.0.0.1/SKILL.md") is None
mock_get.assert_not_called()
@patch("tools.skills_hub.httpx.get")
@patch("tools.skills_hub._ssrf_safe_http_get")
@patch("tools.skills_hub.check_website_access", return_value=None)
@patch("tools.skills_hub.is_safe_url", return_value=True)
def test_fetch_blocks_connect_time_dns_rebind(self, _mock_safe, _mock_policy, mock_get):
from tools.url_safety import SSRFConnectionBlocked
mock_get.side_effect = SSRFConnectionBlocked(
"Blocked request to private/internal address during connect"
)
assert self._source().fetch("https://example.com/SKILL.md") is None
mock_get.assert_called_once_with("https://example.com/SKILL.md", timeout=20)
@patch("tools.skills_hub._ssrf_safe_http_get")
def test_fetch_skips_non_matching_identifier(self, mock_get):
assert self._source().fetch("owner/repo/skill") is None
mock_get.assert_not_called()
+7 -4
View File
@@ -212,8 +212,9 @@ class TestClawHubSource(unittest.TestCase):
self.assertEqual(meta.identifier, "self-improving-agent")
self.assertEqual(meta.tags, ["automation"])
@patch("tools.skills_hub._ssrf_safe_http_get")
@patch("tools.skills_hub.httpx.get")
def test_fetch_resolves_latest_version_and_downloads_raw_files(self, mock_get):
def test_fetch_resolves_latest_version_and_downloads_raw_files(self, mock_get, mock_safe_get):
def side_effect(url, *args, **kwargs):
if url.endswith("/skills/caldav-calendar"):
return _MockResponse(
@@ -233,11 +234,10 @@ class TestClawHubSource(unittest.TestCase):
]
},
)
if url == "https://files.example/skill-md":
return _MockResponse(status_code=200, text="# Skill")
return _MockResponse(status_code=404, json_data={})
mock_get.side_effect = side_effect
mock_safe_get.return_value = _MockResponse(status_code=200, text="# Skill")
bundle = self.src.fetch("caldav-calendar")
@@ -246,6 +246,7 @@ class TestClawHubSource(unittest.TestCase):
self.assertIn("SKILL.md", bundle.files)
self.assertEqual(bundle.files["SKILL.md"], "# Skill")
self.assertEqual(bundle.files["README.md"], "hello")
mock_safe_get.assert_called_once_with("https://files.example/skill-md", timeout=20)
@patch("tools.skills_hub.httpx.get")
def test_fetch_falls_back_to_versions_list(self, mock_get):
@@ -267,7 +268,8 @@ class TestClawHubSource(unittest.TestCase):
@patch("tools.skills_hub.check_website_access", return_value=None)
@patch("tools.skills_hub.is_safe_url")
@patch("tools.skills_hub.httpx.get")
def test_fetch_blocks_private_raw_url(self, mock_get, mock_safe, _mock_policy):
@patch("tools.skills_hub._ssrf_safe_http_get")
def test_fetch_blocks_private_raw_url(self, mock_safe_get, mock_get, mock_safe, _mock_policy):
def side_effect(url, *args, **kwargs):
if url.endswith("/skills/caldav-calendar"):
return _MockResponse(
@@ -297,6 +299,7 @@ class TestClawHubSource(unittest.TestCase):
self.assertIsNone(bundle)
self.assertEqual(mock_get.call_count, 3)
mock_safe_get.assert_not_called()
@patch("tools.skills_hub._write_index_cache")
@patch("tools.skills_hub._read_index_cache", return_value=None)
+132
View File
@@ -3,12 +3,19 @@
import socket
from unittest.mock import patch
import httpx
from tools.url_safety import (
is_safe_url,
async_is_safe_url,
is_always_blocked_url,
normalize_url_for_request,
redirect_target_from_response,
create_ssrf_safe_async_client,
SSRFConnectionBlocked,
_SSRFGuardedAsyncNetworkBackend,
_MAX_SSRF_CONNECT_IPS,
_resolved_http_connect_ips,
_is_blocked_ip,
_global_allow_private_urls,
_reset_allow_private_cache,
@@ -291,6 +298,131 @@ class TestAsyncIsSafeUrl:
assert await async_is_safe_url("http://localhost:8080/") is False
class TestSSRFGuardedHttpxClient:
def test_connect_resolution_caps_safe_ip_candidates(self):
answers = [
(socket.AF_INET, socket.SOCK_STREAM, 6, "", (f"93.184.216.{idx}", 80))
for idx in range(1, _MAX_SSRF_CONNECT_IPS + 4)
]
with patch("socket.getaddrinfo", return_value=answers):
ips = _resolved_http_connect_ips("example.com", 80, "http")
assert len(ips) == _MAX_SSRF_CONNECT_IPS
assert ips[0] == "93.184.216.1"
assert ips[-1] == f"93.184.216.{_MAX_SSRF_CONNECT_IPS}"
def test_connect_resolution_checks_private_ip_beyond_candidate_cap(self):
answers = [
(socket.AF_INET, socket.SOCK_STREAM, 6, "", (f"93.184.216.{idx}", 80))
for idx in range(1, _MAX_SSRF_CONNECT_IPS + 1)
]
answers.append(
(socket.AF_INET, socket.SOCK_STREAM, 6, "", ("169.254.169.254", 80))
)
with patch("socket.getaddrinfo", return_value=answers):
with pytest.raises(SSRFConnectionBlocked, match="metadata"):
_resolved_http_connect_ips("example.com", 80, "http")
@pytest.mark.asyncio
async def test_async_client_dials_validated_ip_not_hostname(self, monkeypatch):
"""Direct httpx fetches should connect to the vetted IP, not re-resolve hostnames."""
import httpcore
from httpcore._backends.auto import AutoBackend
for proxy_var in (
"HTTP_PROXY",
"HTTPS_PROXY",
"ALL_PROXY",
"http_proxy",
"https_proxy",
"all_proxy",
):
monkeypatch.delenv(proxy_var, raising=False)
monkeypatch.setattr(
socket,
"getaddrinfo",
lambda host, port, *args, **kwargs: [
(socket.AF_INET, socket.SOCK_STREAM, 6, "", ("93.184.216.34", port)),
],
)
connect_attempts = []
async def fake_connect_tcp(
self,
host,
port,
timeout=None,
local_address=None,
socket_options=None,
):
connect_attempts.append((host, port))
raise httpcore.ConnectError("stop before network")
monkeypatch.setattr(AutoBackend, "connect_tcp", fake_connect_tcp)
async with create_ssrf_safe_async_client(timeout=0.01, trust_env=False) as client:
with pytest.raises(httpx.ConnectError):
await client.get("http://example.com/image.png")
assert connect_attempts == [("93.184.216.34", 80)]
@pytest.mark.asyncio
async def test_async_backend_blocks_unix_socket_connects(self):
import contextvars
backend = _SSRFGuardedAsyncNetworkBackend(contextvars.ContextVar("test_schemes"))
with pytest.raises(SSRFConnectionBlocked, match="Unix socket"):
await backend.connect_unix_socket("/tmp/hermes.sock")
def test_async_client_rejects_unpatchable_custom_transport(self):
class CustomTransport(httpx.AsyncBaseTransport):
async def handle_async_request(self, request):
return httpx.Response(200, request=request)
with pytest.raises(SSRFConnectionBlocked, match="Unsupported async httpx transport"):
create_ssrf_safe_async_client(transport=CustomTransport())
@pytest.mark.asyncio
async def test_async_client_preserves_env_proxy_mounts(self, monkeypatch):
"""Installing the guard must not disable or rewrite httpx env proxy setup."""
for proxy_var in (
"HTTP_PROXY",
"HTTPS_PROXY",
"ALL_PROXY",
"http_proxy",
"https_proxy",
"all_proxy",
"NO_PROXY",
"no_proxy",
):
monkeypatch.delenv(proxy_var, raising=False)
monkeypatch.setenv("HTTPS_PROXY", "http://proxy.example:8080")
client = create_ssrf_safe_async_client(timeout=0.01)
try:
proxy_transports = [
transport
for transport in client.__dict__.get("_mounts", {}).values()
if transport is not None
]
assert proxy_transports
assert type(client._transport._pool._network_backend).__name__ == (
"_SSRFGuardedAsyncNetworkBackend"
)
assert all(
type(transport._pool._network_backend).__name__
!= "_SSRFGuardedAsyncNetworkBackend"
for transport in proxy_transports
)
finally:
await client.aclose()
class TestIsBlockedIp:
"""Direct tests for the _is_blocked_ip helper."""
+12 -2
View File
@@ -291,8 +291,18 @@ def _resolve_lock_install_path(install_path: str, skill_name: str) -> Path:
return target
def _ssrf_safe_http_get(url: str, *, timeout: int = 20) -> httpx.Response:
"""Fetch one URL with connect-time SSRF validation and no automatic redirects."""
from tools.url_safety import create_ssrf_safe_client
with create_ssrf_safe_client(timeout=timeout, follow_redirects=False) as client:
return client.get(url)
def _guarded_http_get(url: str, *, timeout: int = 20) -> Optional[httpx.Response]:
"""Fetch a URL with SSRF and redirect-target validation."""
from tools.url_safety import SSRFConnectionBlocked
current_url = url
for _ in range(_MAX_SKILL_FETCH_REDIRECTS + 1):
@@ -310,8 +320,8 @@ def _guarded_http_get(url: str, *, timeout: int = 20) -> Optional[httpx.Response
return None
try:
resp = httpx.get(current_url, timeout=timeout, follow_redirects=False)
except httpx.HTTPError as exc:
resp = _ssrf_safe_http_get(current_url, timeout=timeout)
except (SSRFConnectionBlocked, httpx.HTTPError) as exc:
logger.debug("Skills Hub fetch failed for %s: %s", current_url, exc)
return None
+326 -3
View File
@@ -12,11 +12,13 @@ that use 198.18.0.0/15 or 100.64.0.0/10). Even when disabled, cloud
metadata hostnames (metadata.google.internal, 169.254.169.254) are
**always** blocked — those are never legitimate agent targets.
Limitations (documented, not fixable at pre-flight level):
Limitations:
- DNS rebinding (TOCTOU): an attacker-controlled DNS server with TTL=0
can return a public IP for the check, then a private IP for the actual
connection. Fixing this requires connection-level validation (e.g.
Python's Champion library or an egress proxy like Stripe's Smokescreen).
connection. Hermes-owned direct httpx request paths should use
``create_ssrf_safe_client()`` / ``create_ssrf_safe_async_client()`` so the
same policy is applied immediately before TCP connect and the client
connects to the validated IP while preserving Host/SNI semantics.
- Redirect-based bypass is mitigated by httpx event hooks that re-validate
each redirect target in vision_tools, gateway platform adapters, and
media cache helpers. Web tools use third-party SDKs (Firecrawl/Tavily)
@@ -183,6 +185,8 @@ _TRUSTED_PRIVATE_IP_HOSTS = frozenset({
"multimedia.nt.qq.com.cn",
})
_MAX_SSRF_CONNECT_IPS = 8
# 100.64.0.0/10 (CGNAT / Shared Address Space, RFC 6598) is NOT covered by
# ipaddress.is_private — it returns False for both is_private and is_global.
# Must be blocked explicitly. Used by carrier-grade NAT, Tailscale/WireGuard
@@ -476,6 +480,325 @@ async def async_is_safe_url(url: str) -> bool:
return await asyncio.to_thread(is_safe_url, url)
class SSRFConnectionBlocked(ValueError):
"""Raised when connect-time DNS resolution violates the URL safety policy."""
def _safe_connect_scheme(host: str, port: int, schemes_by_origin: dict[tuple[str, int], str]) -> str:
return schemes_by_origin.get((host, port)) or ("https" if port == 443 else "http")
def _resolved_http_connect_ips(host: str, port: int, scheme: str) -> list[str]:
"""Resolve and validate *host* for one HTTP connect attempt.
Unlike :func:`is_safe_url`, this is called from the HTTP transport at the
time the TCP socket is about to be opened. It returns concrete IP strings
that the transport can dial directly, closing the DNS-rebinding gap between
pre-flight validation and connection setup for direct httpx clients.
"""
hostname = (host or "").strip().lower().rstrip(".")
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_all_private = _global_allow_private_urls()
allow_private_ip = _allows_private_ip_resolution(hostname, scheme)
try:
addr_info = socket.getaddrinfo(
hostname, port, socket.AF_UNSPEC, socket.SOCK_STREAM
)
except socket.gaierror as exc:
raise SSRFConnectionBlocked(
f"Blocked request - DNS resolution failed for: {hostname}"
) from exc
safe_ips: list[str] = []
seen: set[str] = set()
for _family, _, _, _, sockaddr in addr_info:
ip_str = sockaddr[0]
if "%" in ip_str:
ip_str = ip_str.split("%")[0]
try:
ip = ipaddress.ip_address(ip_str)
except ValueError as exc:
raise SSRFConnectionBlocked(
f"Blocked request - unparseable IP address {sockaddr[0]!r} for hostname {hostname}"
) from exc
if ip in _ALWAYS_BLOCKED_IPS or any(ip in net for net in _ALWAYS_BLOCKED_NETWORKS):
raise SSRFConnectionBlocked(
f"Blocked request to cloud metadata address during connect: {hostname} -> {ip_str}"
)
if not allow_all_private and not allow_private_ip and _is_blocked_ip(ip):
raise SSRFConnectionBlocked(
f"Blocked request to private/internal address during connect: {hostname} -> {ip_str}"
)
if ip_str not in seen and len(safe_ips) < _MAX_SSRF_CONNECT_IPS:
safe_ips.append(ip_str)
seen.add(ip_str)
if not safe_ips:
raise SSRFConnectionBlocked(f"Blocked request - DNS returned no results for: {hostname}")
return safe_ips
class _SSRFGuardedAsyncNetworkBackend:
def __init__(self, schemes_by_origin_var: Any):
from httpcore._backends.auto import AutoBackend
self._backend = AutoBackend()
self._schemes_by_origin_var = schemes_by_origin_var
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
schemes_by_origin = self._schemes_by_origin_var.get({})
scheme = _safe_connect_scheme(host, port, schemes_by_origin)
ips = await asyncio.to_thread(_resolved_http_connect_ips, host, port, scheme)
last_exc: Exception | None = None
for ip in ips:
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:
last_exc = exc
continue
if last_exc is not None:
raise last_exc
raise SSRFConnectionBlocked(f"Blocked request - DNS returned no usable IPs for: {host}")
async def connect_unix_socket(
self,
path: str,
timeout: float | None = None,
socket_options: Any = None,
) -> Any:
raise SSRFConnectionBlocked("Blocked Unix socket connection in SSRF-safe transport")
async def sleep(self, seconds: float) -> None:
await self._backend.sleep(seconds)
class _SSRFGuardedNetworkBackend:
def __init__(self, schemes_by_origin_var: Any):
from httpcore._backends.sync import SyncBackend
self._backend = SyncBackend()
self._schemes_by_origin_var = schemes_by_origin_var
def connect_tcp(
self,
host: str,
port: int,
timeout: float | None = None,
local_address: str | None = None,
socket_options: Any = None,
) -> Any:
import httpcore
schemes_by_origin = self._schemes_by_origin_var.get({})
scheme = _safe_connect_scheme(host, port, schemes_by_origin)
ips = _resolved_http_connect_ips(host, port, scheme)
last_exc: Exception | None = None
for ip in ips:
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:
last_exc = exc
continue
if last_exc is not None:
raise last_exc
raise SSRFConnectionBlocked(f"Blocked request - DNS returned no usable IPs for: {host}")
def connect_unix_socket(
self,
path: str,
timeout: float | None = None,
socket_options: Any = None,
) -> Any:
raise SSRFConnectionBlocked("Blocked Unix socket connection in SSRF-safe transport")
def sleep(self, seconds: float) -> None:
self._backend.sleep(seconds)
def _origin_scheme_context(request: Any) -> dict[tuple[str, int], str]:
host = request.url.host
port = request.url.port
scheme = request.url.scheme
if not host or port is None or scheme not in {"http", "https"}:
return {}
return {(host, port): scheme}
def ssrf_safe_async_http_transport(**kwargs: Any) -> Any:
"""Return an httpx async transport that pins direct TCP connects to vetted IPs."""
import contextvars
import httpx
schemes_by_origin_var = contextvars.ContextVar("hermes_ssrf_async_origin_schemes")
class _Transport(httpx.AsyncHTTPTransport):
def __init__(self, **transport_kwargs: Any):
super().__init__(**transport_kwargs)
self._pool._network_backend = _SSRFGuardedAsyncNetworkBackend( # type: ignore[attr-defined]
schemes_by_origin_var
)
async def handle_async_request(self, request: Any) -> Any:
token = schemes_by_origin_var.set(_origin_scheme_context(request))
try:
return await super().handle_async_request(request)
finally:
schemes_by_origin_var.reset(token)
return _Transport(**kwargs)
def ssrf_safe_http_transport(**kwargs: Any) -> Any:
"""Return an httpx sync transport that pins direct TCP connects to vetted IPs."""
import contextvars
import httpx
schemes_by_origin_var = contextvars.ContextVar("hermes_ssrf_origin_schemes")
class _Transport(httpx.HTTPTransport):
def __init__(self, **transport_kwargs: Any):
super().__init__(**transport_kwargs)
self._pool._network_backend = _SSRFGuardedNetworkBackend( # type: ignore[attr-defined]
schemes_by_origin_var
)
def handle_request(self, request: Any) -> Any:
token = schemes_by_origin_var.set(_origin_scheme_context(request))
try:
return super().handle_request(request)
finally:
schemes_by_origin_var.reset(token)
return _Transport(**kwargs)
def _install_ssrf_guard_on_async_transport(transport: Any, schemes_by_origin_var: Any) -> None:
state = getattr(transport, "__dict__", {}) if transport is not None else {}
if transport is None or state.get("_hermes_ssrf_guarded", False):
return
pool = state.get("_pool")
if pool is None or not hasattr(pool, "_network_backend"):
raise SSRFConnectionBlocked("Unsupported async httpx transport cannot be made SSRF-safe")
pool._network_backend = _SSRFGuardedAsyncNetworkBackend(schemes_by_origin_var)
handle_async_request = getattr(transport, "handle_async_request", None)
if handle_async_request is None:
raise SSRFConnectionBlocked("Unsupported async httpx transport cannot be made SSRF-safe")
async def guarded_handle_async_request(request: Any) -> Any:
token = schemes_by_origin_var.set(_origin_scheme_context(request))
try:
return await handle_async_request(request)
finally:
schemes_by_origin_var.reset(token)
transport.handle_async_request = guarded_handle_async_request
transport._hermes_ssrf_guarded = True
def _install_ssrf_guard_on_transport(transport: Any, schemes_by_origin_var: Any) -> None:
state = getattr(transport, "__dict__", {}) if transport is not None else {}
if transport is None or state.get("_hermes_ssrf_guarded", False):
return
pool = state.get("_pool")
if pool is None or not hasattr(pool, "_network_backend"):
raise SSRFConnectionBlocked("Unsupported httpx transport cannot be made SSRF-safe")
pool._network_backend = _SSRFGuardedNetworkBackend(schemes_by_origin_var)
handle_request = getattr(transport, "handle_request", None)
if handle_request is None:
raise SSRFConnectionBlocked("Unsupported httpx transport cannot be made SSRF-safe")
def guarded_handle_request(request: Any) -> Any:
token = schemes_by_origin_var.set(_origin_scheme_context(request))
try:
return handle_request(request)
finally:
schemes_by_origin_var.reset(token)
transport.handle_request = guarded_handle_request
transport._hermes_ssrf_guarded = True
def _install_ssrf_guard_on_async_client(client: Any) -> None:
import contextvars
schemes_by_origin_var = contextvars.ContextVar("hermes_ssrf_async_origin_schemes")
state = getattr(client, "__dict__", {})
_install_ssrf_guard_on_async_transport(
state.get("_transport"), schemes_by_origin_var
)
def _install_ssrf_guard_on_client(client: Any) -> None:
import contextvars
schemes_by_origin_var = contextvars.ContextVar("hermes_ssrf_origin_schemes")
state = getattr(client, "__dict__", {})
_install_ssrf_guard_on_transport(
state.get("_transport"), schemes_by_origin_var
)
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 original request hostname is preserved for Host,
SNI, and certificate verification. If httpx routes through a proxy, final
target resolution is delegated to that configured proxy; treat the proxy as
a trusted egress boundary.
"""
import httpx
client = httpx.AsyncClient(**kwargs)
_install_ssrf_guard_on_async_client(client)
return client
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.
+8 -3
View File
@@ -424,10 +424,13 @@ async def _download_image(image_url: str, destination: Path, max_retries: int =
if blocked:
raise PermissionError(blocked["message"])
from tools.url_safety import create_ssrf_safe_async_client
# Download the image with appropriate headers using async httpx
# Enable follow_redirects to handle image CDNs that redirect (e.g., Imgur, Picsum)
# SSRF: event_hooks validates each redirect target against private IP ranges
async with httpx.AsyncClient(
# SSRF: the client validates DNS at TCP connect time; event_hooks
# validate each redirect target against private IP ranges.
async with create_ssrf_safe_async_client(
timeout=_VISION_DOWNLOAD_TIMEOUT,
follow_redirects=True,
event_hooks={"response": [_ssrf_redirect_guard]},
@@ -1575,7 +1578,9 @@ async def _download_video(video_url: str, destination: Path, max_retries: int =
if blocked:
raise PermissionError(blocked["message"])
async with httpx.AsyncClient(
from tools.url_safety import create_ssrf_safe_async_client
async with create_ssrf_safe_async_client(
timeout=60.0,
follow_redirects=True,
event_hooks={"response": [_ssrf_redirect_guard]},