diff --git a/changelog.d/tsk-6uymvv-ssrf-dns-pinning.md b/changelog.d/tsk-6uymvv-ssrf-dns-pinning.md new file mode 100644 index 000000000..9b9bb7035 --- /dev/null +++ b/changelog.d/tsk-6uymvv-ssrf-dns-pinning.md @@ -0,0 +1,14 @@ +- Security: the SSRF guard now pins each outbound connection to the address it + validated. Fetches of user-supplied URLs (browser proxy, extract, download, + Library web ingest, Knowledge article ingest, peer handshake delivery, + UnifiedPush) go through a guarded client whose connections resolve and check + the hostname once and connect to that answer, so a low-TTL nameserver can no + longer answer public to the check and 127.0.0.1 to the connection. TLS + verification is unchanged and still validates the original hostname. +- Fix: Library web ingest now reuses a single guarded client across an entire + redirect chain instead of building and tearing down a fresh one (new + connection pool, SSL context, pinned backend) on every hop. +- Fix: Knowledge article ingest now rejects a caller-supplied `fetch_client` + that is an `httpx.AsyncClient` but was not built by `guarded_async_client`, + instead of silently accepting an unguarded client and bypassing the SSRF + pin. diff --git a/tests/routes/desktop_browser/test_ssrf.py b/tests/routes/desktop_browser/test_ssrf.py index f17582f7f..1e0e2aad8 100644 --- a/tests/routes/desktop_browser/test_ssrf.py +++ b/tests/routes/desktop_browser/test_ssrf.py @@ -191,3 +191,124 @@ def test_rejects_when_only_ipv6_resolves_to_private(self): ): with pytest.raises(SsrfBlockedError): validate_url_or_raise("http://dual-stack.test/") + + +class TestPinnedTransport: + """The guarded client must connect to the address the guard checked.""" + + def test_validate_returns_the_addresses_it_approved(self): + from tinyagentos.routes.desktop_browser.ssrf import validate_url_or_raise + + with patch( + "tinyagentos.routes.desktop_browser.ssrf.socket.getaddrinfo", + return_value=[ + (2, 1, 6, "", ("93.184.216.34", 0)), + (2, 1, 6, "", ("93.184.216.35", 0)), + ], + ): + addrs = validate_url_or_raise("http://example.com/") + + # Resolver order is preserved: the first answer is the one a pinned + # connection uses, so it must not come back through a set. + assert addrs == ["93.184.216.34", "93.184.216.35"] + + def test_client_pins_the_connection_to_the_checked_address(self): + """The pin sits on the pool's network backend, below the request URL.""" + from tinyagentos.routes.desktop_browser.ssrf import ( + _PinnedResolutionBackend, + guarded_async_client, + ) + + client = guarded_async_client() + backend = client._transport._pool._network_backend + assert isinstance(backend, _PinnedResolutionBackend) + + def test_pinned_backend_hands_the_socket_a_checked_literal(self): + """connect_tcp resolves once, validates, and connects to that answer.""" + import asyncio + from unittest.mock import AsyncMock + + from tinyagentos.routes.desktop_browser.ssrf import _PinnedResolutionBackend + + inner = AsyncMock() + backend = _PinnedResolutionBackend(inner) + + with patch( + "tinyagentos.routes.desktop_browser.ssrf.socket.getaddrinfo", + return_value=[(2, 1, 6, "", ("93.184.216.34", 0))], + ): + asyncio.run(backend.connect_tcp("example.com", 443, timeout=5.0)) + + args, kwargs = inner.connect_tcp.await_args + assert args[0] == "93.184.216.34" + assert args[1] == 443 + assert kwargs["timeout"] == 5.0 + + def test_pinned_backend_refuses_a_blocked_answer(self): + import asyncio + + from tinyagentos.routes.desktop_browser.ssrf import ( + SsrfBlockedError, + _PinnedResolutionBackend, + ) + from unittest.mock import AsyncMock + + inner = AsyncMock() + backend = _PinnedResolutionBackend(inner) + + with patch( + "tinyagentos.routes.desktop_browser.ssrf.socket.getaddrinfo", + return_value=[(2, 1, 6, "", ("127.0.0.1", 0))], + ): + with pytest.raises(SsrfBlockedError): + asyncio.run(backend.connect_tcp("rebind.test", 80)) + + inner.connect_tcp.assert_not_awaited() + + def test_allow_private_reaches_the_pin(self): + """A LAN-facing caller pins too — it just permits RFC1918.""" + import asyncio + from unittest.mock import AsyncMock + + from tinyagentos.routes.desktop_browser.ssrf import _PinnedResolutionBackend + + inner = AsyncMock() + backend = _PinnedResolutionBackend(inner, allow_private=True) + + with patch( + "tinyagentos.routes.desktop_browser.ssrf.socket.getaddrinfo", + return_value=[(2, 1, 6, "", ("192.168.1.50", 0))], + ): + asyncio.run(backend.connect_tcp("nas.example.com", 80)) + + assert inner.connect_tcp.await_args[0][0] == "192.168.1.50" + + def test_tls_verification_stays_on_and_the_url_keeps_the_hostname(self): + """Pinning below the URL is what keeps certificate checks meaningful.""" + import ssl + + from tinyagentos.routes.desktop_browser.ssrf import guarded_async_client + + client = guarded_async_client() + request = client.build_request("GET", "https://example.com/page") + # The URL is untouched, so httpcore derives SNI and the cert hostname + # from the real name rather than from a pinned literal. + assert request.url.host == "example.com" + assert request.headers["host"] == "example.com" + + ssl_context = client._transport._pool._ssl_context + assert ssl_context.verify_mode == ssl.CERT_REQUIRED + assert ssl_context.check_hostname is True + + def test_unix_sockets_are_refused(self): + import asyncio + from unittest.mock import AsyncMock + + from tinyagentos.routes.desktop_browser.ssrf import ( + SsrfBlockedError, + _PinnedResolutionBackend, + ) + + backend = _PinnedResolutionBackend(AsyncMock()) + with pytest.raises(SsrfBlockedError): + asyncio.run(backend.connect_unix_socket("/run/taos.sock")) diff --git a/tests/test_knowledge_ingest.py b/tests/test_knowledge_ingest.py index 426a5a073..c2e94896e 100644 --- a/tests/test_knowledge_ingest.py +++ b/tests/test_knowledge_ingest.py @@ -82,6 +82,7 @@ async def pipeline(store, mock_http): p = IngestPipeline( store=store, http_client=mock_http, + fetch_client=mock_http, notifications=notif, category_engine=cat_engine, qmd_base_url="", # QMD disabled for unit tests @@ -200,6 +201,7 @@ async def test_summarise_called_when_llm_url_set(store): pipeline = IngestPipeline( store=store, http_client=mock_http, + fetch_client=mock_http, notifications=notif, category_engine=cat_engine, qmd_base_url="", # disable embed for this test @@ -240,6 +242,7 @@ async def test_embed_called_when_qmd_url_set(store): pipeline = IngestPipeline( store=store, http_client=mock_http, + fetch_client=mock_http, notifications=notif, category_engine=cat_engine, qmd_base_url="http://localhost:7832", @@ -284,6 +287,7 @@ async def test_semaphore_custom_max_concurrent(store, mock_http): p = IngestPipeline( store=store, http_client=mock_http, + fetch_client=mock_http, notifications=notif, category_engine=cat_engine, max_concurrent=2, @@ -301,6 +305,7 @@ async def test_max_concurrent_zero_raises(store, mock_http): IngestPipeline( store=store, http_client=mock_http, + fetch_client=mock_http, notifications=notif, category_engine=cat_engine, max_concurrent=0, @@ -332,6 +337,7 @@ async def counting_run(self, item_id: str) -> None: p = IngestPipeline( store=store, http_client=mock_http, + fetch_client=mock_http, notifications=notif, category_engine=cat_engine, max_concurrent=2, @@ -373,6 +379,7 @@ async def test_categories_from_caller_are_preserved(store): pipeline = IngestPipeline( store=store, http_client=mock_http, + fetch_client=mock_http, notifications=notif, category_engine=cat_engine, qmd_base_url="", @@ -413,3 +420,36 @@ async def test_download_article_blocks_internal_url(pipeline, store): # The guard raises before any HTTP call to the internal address. for call in pipeline._http_client.get.await_args_list: assert "127.0.0.1" not in str(call) + + +@pytest.mark.asyncio +async def test_fetch_client_must_be_guarded(store, mock_http): + """A caller-supplied fetch_client that IS an httpx.AsyncClient must carry + the SSRF-pinned transport. Otherwise a caller handing in a plain, + unguarded client (e.g. a shared app-wide client) would silently bypass + the guard for every article fetch through this pipeline.""" + import httpx + + def _handler(request): + # Should never actually be reached — the type check must fire first. + return httpx.Response( + 200, + headers={"content-type": "text/html"}, + text="

Should never be read.

", + ) + + notif = AsyncMock() + cat_engine = AsyncMock() + unguarded_client = httpx.AsyncClient(transport=httpx.MockTransport(_handler)) + pipeline = IngestPipeline( + store=store, + http_client=mock_http, + fetch_client=unguarded_client, + notifications=notif, + category_engine=cat_engine, + ) + try: + with pytest.raises(TypeError, match="guarded_async_client"): + await pipeline._download_article("https://example.com/a", "", {}) + finally: + await unguarded_client.aclose() diff --git a/tests/test_library.py b/tests/test_library.py index 8c11028d8..f394b1ff4 100644 --- a/tests/test_library.py +++ b/tests/test_library.py @@ -709,6 +709,60 @@ async def test_web_redirect_hop_validated(self, lib_store, storage_dir): # Both hops should have been validated assert mock_validate.call_count == 2 + @pytest.mark.asyncio + async def test_web_redirect_reuses_single_guarded_client(self, lib_store, storage_dir): + """One `guarded_async_client()` call must serve every hop of a + multi-hop redirect chain — not a fresh client (pool, SSL context, + pinned backend) built and torn down on each hop.""" + from unittest.mock import patch, MagicMock, AsyncMock + import tinyagentos.routes.desktop_browser.ssrf as ssrf_mod + + html = "

Final hop content, long enough to pass the readability minimum threshold for extraction.

" + item_id = await lib_store.create_item( + kind="url:web", + source_url="https://safe.example.com/start", + ) + item = await lib_store.get_item(item_id) + proc = WebProcessor(lib_store, storage_dir) + + # Two redirects then a final 200: 3 hops total. + mock_resp1 = MagicMock() + mock_resp1.status_code = 302 + mock_resp1.headers = {"location": "https://safe.example.com/hop1"} + mock_resp1.is_redirect = True + + mock_resp2 = MagicMock() + mock_resp2.status_code = 302 + mock_resp2.headers = {"location": "https://safe.example.com/final"} + mock_resp2.is_redirect = True + + mock_resp3 = _mock_httpx_response(html, 200) + + mock_client = MagicMock() + mock_client.__aenter__ = AsyncMock(return_value=mock_client) + mock_client.__aexit__ = AsyncMock(return_value=None) + mock_client.stream = MagicMock( + side_effect=[ + _mock_stream_ctx(mock_resp1), + _mock_stream_ctx(mock_resp2), + _mock_stream_ctx(mock_resp3), + ] + ) + + counting = MagicMock(side_effect=lambda *a, **kw: mock_client) + + with ( + patch.object(ssrf_mod, "guarded_async_client", counting), + patch.object(ssrf_mod, "validate_url_or_raise"), + ): + await proc.process(item) + + assert counting.call_count == 1, ( + f"guarded_async_client entered {counting.call_count} times " + "for a 3-hop fetch — it must be entered exactly once and reused " + "across every redirect hop" + ) + @pytest.mark.asyncio async def test_web_size_cap(self, lib_store, storage_dir): """Responses exceeding the size cap raise ValueError.""" diff --git a/tests/test_ssrf_rebinding.py b/tests/test_ssrf_rebinding.py new file mode 100644 index 000000000..35fbcb402 --- /dev/null +++ b/tests/test_ssrf_rebinding.py @@ -0,0 +1,237 @@ +"""DNS-rebinding regression for the shared SSRF guard. + +The guard used to resolve a hostname, validate every returned address and +then hand the caller nothing but permission. The caller's HTTP client then +performed its *own* lookup, so an attacker running the authoritative +nameserver could answer public to the check and loopback to the connection. +Blocking `http://127.0.0.1/` proves nothing about that bug -- it is already +refused at validation time. The bug only shows up when the two lookups +disagree, so that is what these tests script. + +`_ScriptedResolver` stands in for `socket.getaddrinfo` and answers a +different address on each call for the same hostname. The connection layer +is intercepted at `httpcore`'s network backend -- the exact place a real +connection resolves the name and opens the socket -- so the second lookup is +made for real and the socket really is opened, against a local stand-in for +whichever service that lookup pointed at. +""" + +from __future__ import annotations + +import asyncio +import ipaddress +import socket +import tempfile +from pathlib import Path +from unittest.mock import patch + +import pytest +import pytest_asyncio +from httpcore._backends.anyio import AnyIOBackend + +from tinyagentos.library_pipeline import WebProcessor +from tinyagentos.library_store import LibraryStore +from tinyagentos.routes.desktop_browser.ssrf import SsrfBlockedError + +# example.com — a plain public address, accepted by the guard. +_PUBLIC = "93.184.216.34" +_INTERNAL = "127.0.0.1" + +_EXTERNAL_PAGE = ( + b"Public page
" + b"

This is the public page the user actually asked for. It carries " + b"enough prose that the readability extractor keeps it instead of " + b"falling back to the bare tag stripper for very short documents.

" + b"
" +) +_INTERNAL_PAGE = ( + b"Admin
" + b"

INTERNAL-SERVICE-SECRET: this response only exists on the loopback " + b"interface and must never be reachable through a user-supplied URL, " + b"however the attacker's nameserver answers the second lookup.

" + b"
" +) + + +class _ScriptedResolver: + """A `socket.getaddrinfo` stand-in that answers a low-TTL nameserver. + + `answers` maps a hostname to the addresses it hands out, one per call; + the last entry repeats once the script runs out. IP literals resolve to + themselves, as the real resolver does. + """ + + def __init__(self, answers: dict[str, list[str]]) -> None: + self._answers = answers + self.calls: dict[str, int] = {} + + def __call__(self, host, port=0, *args, **kwargs): + name = host.decode() if isinstance(host, bytes) else str(host) + try: + ipaddress.ip_address(name) + except ValueError: + script = self._answers.get(name.lower()) + if script is None: + raise socket.gaierror( + socket.EAI_NONAME, f"scripted resolver has no answer for {name!r}" + ) + index = min(self.calls.get(name.lower(), 0), len(script) - 1) + self.calls[name.lower()] = index + 1 + addr = script[index] + else: + addr = name + + sock_port = port if isinstance(port, int) else 0 + return [ + ( + socket.AF_INET, + socket.SOCK_STREAM, + socket.IPPROTO_TCP, + "", + (addr, sock_port), + ) + ] + + +async def _start_page_server(body: bytes) -> tuple[asyncio.AbstractServer, int]: + """Serve `body` once per connection over minimal HTTP/1.1 on loopback.""" + + async def handle(reader: asyncio.StreamReader, writer: asyncio.StreamWriter) -> None: + try: + await reader.readuntil(b"\r\n\r\n") + writer.write( + b"HTTP/1.1 200 OK\r\n" + b"Content-Type: text/html; charset=utf-8\r\n" + b"Content-Length: %d\r\n" + b"Connection: close\r\n\r\n" % len(body) + + body + ) + await writer.drain() + except (asyncio.IncompleteReadError, ConnectionError): + pass + finally: + writer.close() + + server = await asyncio.start_server(handle, "127.0.0.1", 0) + return server, server.sockets[0].getsockname()[1] + + +class _ConnectionRecorder: + """Intercepts the connection so the address it lands on is observable. + + Stands where `httpcore` opens the socket: it performs the connection's + own hostname lookup (the second lookup) exactly as the real backend + does, records the address that lookup produced, and then opens a real + socket to the local stand-in for that address -- the internal service + for a loopback answer, the public site otherwise. + """ + + def __init__(self, *, internal_port: int, external_port: int) -> None: + self.internal_port = internal_port + self.external_port = external_port + self.connected: list[str] = [] + self._real_connect_tcp = AnyIOBackend.connect_tcp + + async def connect_tcp( + self, + host: str, + port: int, + timeout: float | None = None, + local_address: str | None = None, + socket_options=None, + ): + addr = socket.getaddrinfo(host, port, 0, socket.SOCK_STREAM)[0][4][0] + self.connected.append(addr) + stand_in_port = ( + self.internal_port + if ipaddress.ip_address(addr).is_loopback + else self.external_port + ) + return await self._real_connect_tcp( + AnyIOBackend(), "127.0.0.1", stand_in_port, timeout=timeout, + ) + + def patched(self): + return patch.object(AnyIOBackend, "connect_tcp", self.connect_tcp) + + +@pytest_asyncio.fixture +async def lib_store(): + with tempfile.NamedTemporaryFile(suffix=".db", delete=False) as f: + db_path = Path(f.name) + store = LibraryStore(db_path) + await store.init() + yield store + await store.close() + db_path.unlink(missing_ok=True) + + +@pytest_asyncio.fixture +async def page_servers(): + """Loopback stand-ins for the internal service and the public site.""" + internal, internal_port = await _start_page_server(_INTERNAL_PAGE) + external, external_port = await _start_page_server(_EXTERNAL_PAGE) + recorder = _ConnectionRecorder( + internal_port=internal_port, external_port=external_port, + ) + yield recorder + for server in (internal, external): + server.close() + await server.wait_closed() + + +@pytest.fixture +def storage_dir(): + with tempfile.TemporaryDirectory() as d: + yield Path(d) + + +async def _fetch_through_web_processor( + store: LibraryStore, storage_dir: Path, url: str, +) -> list[dict]: + """Drive a real `validate_url_or_raise` caller end to end.""" + item_id = await store.create_item(kind="url:web", source_url=url, title="") + item = await store.get_item(item_id) + return await WebProcessor(store, storage_dir).process(item) + + +@pytest.mark.asyncio +async def test_second_lookup_to_loopback_is_refused(lib_store, storage_dir, page_servers): + """A nameserver that answers public, then loopback, must not be followed.""" + resolver = _ScriptedResolver({"rebind.test": [_PUBLIC, _INTERNAL]}) + + blocked: SsrfBlockedError | None = None + with patch("socket.getaddrinfo", resolver), page_servers.patched(): + try: + await _fetch_through_web_processor( + lib_store, storage_dir, "http://rebind.test/page", + ) + except SsrfBlockedError as e: + blocked = e + + reached = page_servers.connected[-1] if page_servers.connected else "(no connection)" + assert blocked is not None, ( + f"expected SsrfBlockedError, but the fetch reached {reached}" + ) + assert _INTERNAL not in page_servers.connected, ( + f"the connection was opened against {page_servers.connected}" + ) + + +@pytest.mark.asyncio +async def test_agreeing_lookups_still_fetch(lib_store, storage_dir, page_servers): + """Control: two lookups agreeing on a public address must still fetch.""" + resolver = _ScriptedResolver({"stable.test": [_PUBLIC, _PUBLIC]}) + + with patch("socket.getaddrinfo", resolver), page_servers.patched(): + artifacts = await _fetch_through_web_processor( + lib_store, storage_dir, "http://stable.test/page", + ) + + assert page_servers.connected == [_PUBLIC], page_servers.connected + text_artifacts = [a for a in artifacts if a["kind"] == "text"] + assert len(text_artifacts) == 1 + body = Path(text_artifacts[0]["path"]).read_text(encoding="utf-8") + assert "public page" in body + assert "INTERNAL-SERVICE-SECRET" not in body + diff --git a/tinyagentos/knowledge_ingest.py b/tinyagentos/knowledge_ingest.py index 107b2c972..f4fd0540a 100644 --- a/tinyagentos/knowledge_ingest.py +++ b/tinyagentos/knowledge_ingest.py @@ -87,9 +87,20 @@ def __init__( qmd_base_url: str = "", llm_base_url: str = "", max_concurrent: int = _INGEST_SEMAPHORE_SLOTS, + fetch_client: "httpx.AsyncClient | None" = None, ) -> None: + """`http_client` talks to our own services (the LLM backend, qmd), so + it must stay unguarded — they live on loopback. Article URLs are + user-supplied and are fetched with `fetch_client` instead, which + defaults to a fresh SSRF-pinned client per download; pass one only if + it is guarded too. If the object passed in is an `httpx.AsyncClient`, + this is enforced at fetch time — it must carry the + `SsrfGuardedAsyncTransport`, or a `TypeError` is raised. Test doubles + that are not `httpx.AsyncClient` (mocks, fakes) are exempt. + """ self._store = store self._http_client = http_client + self._fetch_client = fetch_client self._notifications = notifications self._category_engine = category_engine self._qmd_base_url = qmd_base_url @@ -286,27 +297,51 @@ async def _download_article( against the loopback / link-local / private-range blocklist before the request is issued, so an attacker-supplied URL (or a public URL that 302-redirects inward) cannot make the host fetch internal services. + The fetch itself goes through an SSRF-pinned client, so the address + that passed the blocklist is the address the socket is opened to. """ + from contextlib import nullcontext from urllib.parse import urljoin + import httpx + from tinyagentos.routes.desktop_browser.ssrf import ( SsrfBlockedError, + SsrfGuardedAsyncTransport, + guarded_async_client, validate_url_or_raise, ) + # A caller-supplied fetch_client that IS an httpx.AsyncClient must + # carry the SSRF-pinned transport, or the guard is silently bypassed + # for every article fetch through this pipeline (e.g. a shared + # app-wide client handed in by mistake). Test doubles that are not + # httpx.AsyncClient at all (mocks, fakes) pass through unchanged. + if self._fetch_client is not None and isinstance(self._fetch_client, httpx.AsyncClient): + if not isinstance( + getattr(self._fetch_client, "_transport", None), SsrfGuardedAsyncTransport + ): + raise TypeError("fetch_client must be built by guarded_async_client") + + client_cm = ( + nullcontext(self._fetch_client) + if self._fetch_client is not None + else guarded_async_client() + ) current_url = url resp = None - for _hop in range(_MAX_ARTICLE_REDIRECTS + 1): - validate_url_or_raise(current_url) # raises SsrfBlockedError - resp = await self._http_client.get( - current_url, timeout=30, follow_redirects=False - ) - if resp.is_redirect and resp.headers.get("location"): - current_url = urljoin(current_url, resp.headers["location"]) - continue - break - else: - raise SsrfBlockedError(f"too many redirects fetching {url!r}") + async with client_cm as http: + for _hop in range(_MAX_ARTICLE_REDIRECTS + 1): + validate_url_or_raise(current_url) # raises SsrfBlockedError + resp = await http.get( + current_url, timeout=30, follow_redirects=False + ) + if resp.is_redirect and resp.headers.get("location"): + current_url = urljoin(current_url, resp.headers["location"]) + continue + break + else: + raise SsrfBlockedError(f"too many redirects fetching {url!r}") resp.raise_for_status() html = resp.text diff --git a/tinyagentos/library_pipeline.py b/tinyagentos/library_pipeline.py index b5b187948..4f99f99f5 100644 --- a/tinyagentos/library_pipeline.py +++ b/tinyagentos/library_pipeline.py @@ -480,6 +480,7 @@ async def process(self, item: dict) -> list[dict]: from tinyagentos.routes.desktop_browser.ssrf import ( SsrfBlockedError, + guarded_async_client, validate_url_or_raise, ) @@ -488,13 +489,17 @@ async def process(self, item: dict) -> list[dict]: # — a hostile server streaming a multi-GB text/html body is OOM-safe. async def _fetch() -> tuple[str, str, bytes]: current_url = source_url - for _hop in range(self._MAX_WEB_REDIRECTS + 1): - validate_url_or_raise(current_url) + # One client (one pool, one SSL context, one pinned backend) serves + # every hop of the redirect chain — the inner backend re-resolves + # and re-validates per connection anyway (see ssrf.py), so reuse + # across hops is exactly what it was designed for. + async with guarded_async_client( + timeout=httpx.Timeout(30), + follow_redirects=False, + ) as client: + for _hop in range(self._MAX_WEB_REDIRECTS + 1): + validate_url_or_raise(current_url) - async with httpx.AsyncClient( - timeout=httpx.Timeout(30), - follow_redirects=False, - ) as client: async with client.stream("GET", current_url) as resp: status_code = resp.status_code content_type = resp.headers.get("content-type", "") @@ -527,10 +532,10 @@ async def _fetch() -> tuple[str, str, bytes]: body_chunks.append(chunk) encoding = resp.encoding or "utf-8" return content_type, encoding, b"".join(body_chunks) - else: - raise SsrfBlockedError( - f"too many redirects fetching {source_url!r}" - ) + else: + raise SsrfBlockedError( + f"too many redirects fetching {source_url!r}" + ) try: content_type, encoding, body = await asyncio.wait_for( diff --git a/tinyagentos/peer.py b/tinyagentos/peer.py index 17fc6b8b3..f62b2d4fd 100644 --- a/tinyagentos/peer.py +++ b/tinyagentos/peer.py @@ -260,12 +260,15 @@ async def deliver_handshake( from tinyagentos.routes.desktop_browser.ssrf import ( SsrfBlockedError, + guarded_async_client, validate_url_or_raise, ) own_client = http_client is None if own_client: - http_client = httpx.AsyncClient(timeout=15.0) + # Guarded: the endpoint URL is peer-supplied, so the address that + # passed the blocklist must be the address the POST connects to. + http_client = guarded_async_client(timeout=15.0) try: for ep in peer_endpoints: diff --git a/tinyagentos/push/unifiedpush.py b/tinyagentos/push/unifiedpush.py index 1f22cb16e..11430d5e6 100644 --- a/tinyagentos/push/unifiedpush.py +++ b/tinyagentos/push/unifiedpush.py @@ -7,7 +7,11 @@ import httpx -from tinyagentos.routes.desktop_browser.ssrf import SsrfBlockedError, validate_url_or_raise +from tinyagentos.routes.desktop_browser.ssrf import ( + SsrfBlockedError, + guarded_async_client, + validate_url_or_raise, +) logger = logging.getLogger(__name__) @@ -58,7 +62,10 @@ def _actions_for_row(row: dict) -> list[dict] | None: class HttpUnifiedPushSender: def __init__(self, *, client: httpx.AsyncClient | None = None): - self._client = client or httpx.AsyncClient() + # Push tokens are user-supplied URLs, so the default client pins each + # connection to the address the guard checked. allow_private mirrors + # the send() validation: a LAN distributor is fine, loopback is not. + self._client = client or guarded_async_client(allow_private=True) self._owns_client = client is None async def send(self, push_token: str, payload: dict) -> bool: diff --git a/tinyagentos/routes/desktop_browser/download.py b/tinyagentos/routes/desktop_browser/download.py index 8f1963bff..9af66c539 100644 --- a/tinyagentos/routes/desktop_browser/download.py +++ b/tinyagentos/routes/desktop_browser/download.py @@ -38,6 +38,7 @@ ) from tinyagentos.routes.desktop_browser.ssrf import ( SsrfBlockedError, + guarded_async_client, validate_url_or_raise, ) @@ -152,7 +153,7 @@ async def download_endpoint( ) # Manage the AsyncClient lifetime — don't close until the streamer finishes. - http = httpx.AsyncClient( + http = guarded_async_client( follow_redirects=False, timeout=_FETCH_TIMEOUT, cookies=cookies, ) diff --git a/tinyagentos/routes/desktop_browser/extract.py b/tinyagentos/routes/desktop_browser/extract.py index 1ffe0f8c3..4867335f0 100644 --- a/tinyagentos/routes/desktop_browser/extract.py +++ b/tinyagentos/routes/desktop_browser/extract.py @@ -32,6 +32,7 @@ from tinyagentos.routes.desktop_browser import router from tinyagentos.routes.desktop_browser.ssrf import ( SsrfBlockedError, + guarded_async_client, validate_url_or_raise, ) @@ -110,7 +111,9 @@ async def extract_endpoint( # following a redirect to an internal address after the initial SSRF gate passes. _MAX_HOPS = 5 response: httpx.Response | None = None - async with httpx.AsyncClient(follow_redirects=False, timeout=_FETCH_TIMEOUT) as http: + async with guarded_async_client( + follow_redirects=False, timeout=_FETCH_TIMEOUT, + ) as http: fetch_url = url for _hop in range(_MAX_HOPS): try: diff --git a/tinyagentos/routes/desktop_browser/proxy.py b/tinyagentos/routes/desktop_browser/proxy.py index c65e54f4a..a6cb5bbac 100644 --- a/tinyagentos/routes/desktop_browser/proxy.py +++ b/tinyagentos/routes/desktop_browser/proxy.py @@ -54,6 +54,7 @@ from tinyagentos.routes.desktop_browser.rewriter import rewrite_html from tinyagentos.routes.desktop_browser.ssrf import ( SsrfBlockedError, + guarded_async_client, validate_url_or_raise, ) @@ -307,7 +308,7 @@ async def _fetch_with_redirects() -> httpx.Response | None: hop_method = method hop_body = req_body _resp: httpx.Response | None = None - async with httpx.AsyncClient( + async with guarded_async_client( follow_redirects=False, timeout=_HOP_TIMEOUT, ) as http: for hop in range(_MAX_REDIRECTS + 1): diff --git a/tinyagentos/routes/desktop_browser/ssrf.py b/tinyagentos/routes/desktop_browser/ssrf.py index 0e33d9868..622b0dc7f 100644 --- a/tinyagentos/routes/desktop_browser/ssrf.py +++ b/tinyagentos/routes/desktop_browser/ssrf.py @@ -7,18 +7,38 @@ parses every target URL, resolves its hostname, and refuses to proceed if any resolved address is in the blocklist. +Validating is not enough on its own: a check that resolves the hostname +and then lets the HTTP client resolve it a second time can be defeated by +an attacker who runs the authoritative nameserver for that hostname and +answers public to the check and 127.0.0.1 to the connection. So the +address that was checked has to be the address that is connected to. That +is what `guarded_async_client` is for — it hands out an `httpx` client +whose connections resolve the hostname exactly once, validate that answer, +and open the socket to it. TLS is untouched: the request URL still carries +the hostname, so SNI and certificate verification still run against the +original name. + Usage: from tinyagentos.routes.desktop_browser.ssrf import ( SsrfBlockedError, + guarded_async_client, validate_url_or_raise, ) try: - validate_url_or_raise(target_url) + validate_url_or_raise(target_url) # fail fast, with a reason except SsrfBlockedError as e: return JSONResponse({"error": str(e)}, status_code=403) + async with guarded_async_client(timeout=30) as http: # enforced here + resp = await http.get(target_url) + +Any client that fetches a user-supplied URL must come from +`guarded_async_client`; a bare `httpx.AsyncClient` re-resolves the name and +reopens the hole. Clients that only talk to trusted local services (the LLM +backend, qmd) do not need it. + For redirect handling, callers must invoke validate_url_or_raise on EVERY redirect target (not just the initial URL). The `httpx` follow_redirects=True default does not give us a callback per redirect, @@ -29,8 +49,12 @@ import ipaddress import socket +import typing from urllib.parse import urlparse +import httpcore +import httpx + class SsrfBlockedError(Exception): """Raised when a URL fails SSRF validation.""" @@ -55,13 +79,20 @@ class SsrfBlockedError(Exception): ) -def validate_url_or_raise(url: str, *, allow_private: bool = False) -> None: +def validate_url_or_raise(url: str, *, allow_private: bool = False) -> list[str]: """Validate that `url` is safe to fetch. Parses the URL, checks scheme + hostname suffix, resolves DNS, and verifies every resolved address against the blocklist. Raises `SsrfBlockedError` on any failure. + Returns the resolved, checked addresses in resolver order. Permission + on its own is not enough: whoever fetches the URL has to connect to the + address that was checked, or an attacker-run nameserver can answer this + lookup public and the connection's lookup 127.0.0.1. Fetch with a + `guarded_async_client`, which resolves and validates inside the + connection itself. + Pass ``allow_private=True`` to permit RFC1918 addresses and their IPv6 unique-local equivalent (e.g. self-hosted LAN services) while still refusing loopback, link-local, multicast, reserved, and unspecified @@ -76,7 +107,18 @@ def validate_url_or_raise(url: str, *, allow_private: bool = False) -> None: if not parsed.hostname: raise SsrfBlockedError("URL has no hostname") - host = parsed.hostname.strip().lower() + return resolve_and_validate(parsed.hostname, allow_private=allow_private) + + +def resolve_and_validate(hostname: str, *, allow_private: bool = False) -> list[str]: + """Resolve `hostname` once and validate every address it answers with. + + Returns the checked addresses in resolver order — the first is the one + a connection should be opened to. Raises `SsrfBlockedError` if the + hostname carries a blocked suffix, does not resolve, or resolves to any + blocked address. See `validate_resolved_addr` for ``allow_private``. + """ + host = hostname.strip().lower() # Hostname-based blocklist (catches .local / .onion / .internal # before we even resolve DNS, since these may not resolve at all @@ -112,7 +154,10 @@ def validate_url_or_raise(url: str, *, allow_private: bool = False) -> None: results = socket.getaddrinfo(host, None) # results is a list of (family, type, proto, canonname, sockaddr). # sockaddr[0] is the address string for both AF_INET and AF_INET6. - addrs = list({r[4][0] for r in results}) + # Dedupe but keep first-seen order: `dict.fromkeys` preserves the + # resolver's own order, and the first answer is the one the pin + # connects to. + addrs = list(dict.fromkeys(r[4][0] for r in results)) except socket.gaierror as e: raise SsrfBlockedError(f"could not resolve hostname: {e}") from e @@ -122,6 +167,8 @@ def validate_url_or_raise(url: str, *, allow_private: bool = False) -> None: for addr in addrs: validate_resolved_addr(addr, allow_private=allow_private) + return addrs + def validate_resolved_addr(addr: str, *, allow_private: bool = False) -> None: """Validate that a resolved IP address is safe to connect to. @@ -166,6 +213,109 @@ def validate_resolved_addr(addr: str, *, allow_private: bool = False) -> None: ) +class _PinnedResolutionBackend(httpcore.AsyncNetworkBackend): + """Network backend that connects to the address it just validated. + + `httpcore` calls `connect_tcp` with the hostname from the request URL, + which is where the second, unchecked DNS lookup used to happen. This + backend does that lookup itself, runs the blocklist over the answer, + and hands the socket layer the literal address instead of the name — + so there is only ever one lookup, and it is the checked one. + + The request URL is left alone, so `httpcore` still derives SNI and the + certificate-verification hostname from the original name. + """ + + def __init__( + self, inner: httpcore.AsyncNetworkBackend, *, allow_private: bool = False, + ) -> None: + self._inner = inner + self._allow_private = allow_private + + async def connect_tcp( + self, + host: str, + port: int, + timeout: float | None = None, + local_address: str | None = None, + socket_options: typing.Iterable[typing.Any] | None = None, + ) -> httpcore.AsyncNetworkStream: + addrs = resolve_and_validate(host, allow_private=self._allow_private) + return await self._inner.connect_tcp( + addrs[0], + port, + timeout=timeout, + local_address=local_address, + socket_options=socket_options, + ) + + async def connect_unix_socket( + self, + path: str, + timeout: float | None = None, + socket_options: typing.Iterable[typing.Any] | None = None, + ) -> httpcore.AsyncNetworkStream: + # Nothing routes a user-supplied URL to a unix socket, and one would + # bypass the address blocklist entirely, so refuse rather than pass through. + raise SsrfBlockedError("unix-socket connections are not allowed") + + async def sleep(self, seconds: float) -> None: + await self._inner.sleep(seconds) + + +class SsrfGuardedAsyncTransport(httpx.AsyncHTTPTransport): + """`httpx` transport whose connections are pinned to a checked address. + + Everything above the socket is stock `httpx`: request URLs, redirects, + cookies, and TLS verification behave exactly as they do on the default + transport. Only the pool's network backend is swapped, for one that + resolves and validates the hostname as part of opening the connection. + """ + + def __init__(self, *, allow_private: bool = False, **kwargs: typing.Any) -> None: + super().__init__(**kwargs) + pool = getattr(self, "_pool", None) + if not hasattr(pool, "_network_backend"): + # An httpx/httpcore upgrade moved the seam. Fail loudly: handing + # back a transport that silently does not pin is the bug itself. + raise RuntimeError( + "SSRF pinning could not be installed — httpx's connection pool " + "no longer exposes _network_backend. Refusing to hand out an " + "unpinned client." + ) + pool._network_backend = _PinnedResolutionBackend( + pool._network_backend, allow_private=allow_private, + ) + + +def guarded_async_client( + *, + allow_private: bool = False, + verify: typing.Any = True, + http2: bool = False, + **kwargs: typing.Any, +) -> httpx.AsyncClient: + """An `httpx.AsyncClient` that only ever connects to checked addresses. + + Use this — not a bare `httpx.AsyncClient` — for every fetch of a URL + the user or a remote peer supplied. Remaining keyword arguments go to + `httpx.AsyncClient` (timeout, follow_redirects, cookies, headers, ...). + + TLS verification stays on (`verify` defaults to True) and still checks + the certificate against the original hostname; pinning happens below + the URL, so there is no reason to weaken it. Note that supplying a + transport means `httpx` no longer picks up HTTP_PROXY/HTTPS_PROXY from + the environment — proxying an untrusted fetch would defeat the pin + anyway, since the proxy would do the resolving. + """ + return httpx.AsyncClient( + transport=SsrfGuardedAsyncTransport( + allow_private=allow_private, verify=verify, http2=http2, + ), + **kwargs, + ) + + def _try_parse_encoded_ipv4(host: str) -> str | None: """Attempt to interpret `host` as an integer-encoded IPv4 address.