From a9331e78e64216bc467ffb80a247f38e7d41963a Mon Sep 17 00:00:00 2001 From: "open-swe[bot]" <215916821+open-swe[bot]@users.noreply.github.com> Date: Fri, 8 May 2026 22:06:25 +0000 Subject: [PATCH] fix: harden http_request SSRF guard against DNS rebinding [closes AB-2321] (#1277) * fix: harden http_request SSRF guard against DNS rebinding [closes AB-2321] Pin DNS resolution per request hop so urllib3's connection-time lookup cannot rebind to a private IP after _is_url_safe validated a public one. Co-authored-by: Johannes du Plessis <51395795+johannes117@users.noreply.github.com> * fix: scope DNS pin to urllib3 with reference-counted install Address review feedback: the previous version permanently overwrote the process-global socket.getaddrinfo on first use. Now the patch targets urllib3.util.connection.create_connection (much narrower blast radius), and is installed/uninstalled via reference count so no global mutation persists once no http_request calls are in flight. Co-authored-by: Johannes du Plessis <51395795+johannes117@users.noreply.github.com> * fix: forward timeout and socket_options through pinned create_connection urllib3 calls create_connection with timeout positional and socket_options as a keyword. The previous wrapper only read kwargs, silently dropping the caller's connect timeout (so a slow validated IP could hang) and TCP options like TCP_NODELAY. Accept both positionally and forward them to the underlying socket. --------- Co-authored-by: open-swe[bot] Co-authored-by: Johannes du Plessis <51395795+johannes117@users.noreply.github.com> Co-authored-by: Johannes du Plessis --- agent/tools/http_request.py | 173 ++++++++++++++++++++++++++---- tests/test_http_security.py | 203 ++++++++++++++++++++++++++++++++++++ 2 files changed, 357 insertions(+), 19 deletions(-) diff --git a/agent/tools/http_request.py b/agent/tools/http_request.py index 2a9bec36..87c231d1 100644 --- a/agent/tools/http_request.py +++ b/agent/tools/http_request.py @@ -1,42 +1,170 @@ +import contextlib import ipaddress import socket +import threading +from collections.abc import Iterator from typing import Any from urllib.parse import urljoin, urlparse import requests +from urllib3.util import connection as urllib3_connection _MAX_REDIRECTS = 5 +_pin_state = threading.local() +_install_lock = threading.Lock() +_install_count = 0 +_original_create_connection = None -def _is_url_safe(url: str) -> tuple[bool, str]: - """Check if a URL is safe to request (not targeting private/internal networks).""" + +def _get_pin_stack() -> list[dict[str, list]]: + stack = getattr(_pin_state, "stack", None) + if stack is None: + stack = [] + _pin_state.stack = stack + return stack + + +def _pinned_create_connection( + address, + timeout=socket._GLOBAL_DEFAULT_TIMEOUT, + source_address=None, + socket_options=None, +): + """Drop-in for urllib3.util.connection.create_connection that honors DNS pins. + + When the calling thread has an active _pin_dns context for this host, the + connection uses the pre-validated addresses instead of calling + socket.getaddrinfo again — closing the DNS-rebinding race. + + `timeout` and `socket_options` are accepted positionally because urllib3 + calls create_connection with timeout positional; reading them from kwargs + only would silently drop the caller's connect timeout and TCP options. + """ + host, port = address + if host.startswith("[") and host.endswith("]"): + host = host[1:-1] + + stack = _get_pin_stack() + pins = stack[-1] if stack else None + pinned = pins.get(host) if pins else None + + if pinned is None: + return _original_create_connection( + address, + timeout, + source_address=source_address, + socket_options=socket_options, + ) + + err = None + for family, socktype, proto, _canonname, sockaddr in pinned: + if family == socket.AF_INET: + target = (sockaddr[0], port) + elif family == socket.AF_INET6: + rest = sockaddr[2:] if len(sockaddr) >= 4 else (0, 0) + target = (sockaddr[0], port, *rest) + else: + continue + + sock = None + try: + sock = socket.socket(family, socktype, proto) + for opt in socket_options or (): + sock.setsockopt(*opt) + if timeout is not socket._GLOBAL_DEFAULT_TIMEOUT: + sock.settimeout(timeout) + if source_address: + sock.bind(source_address) + sock.connect(target) + return sock + except OSError as e: + err = e + if sock is not None: + sock.close() + + if err is not None: + raise err + raise OSError("DNS pin produced no usable addresses") + + +@contextlib.contextmanager +def _pin_dns(hostname: str, addr_infos: list) -> Iterator[None]: + """Pin DNS resolution for `hostname` to `addr_infos` for the duration of the block. + + The patch is scoped to urllib3's connection helper (not socket-wide) and is + installed on first entry / removed on last exit via reference counting, so + no global mutation persists once no http_request calls are in flight. + Other hostnames pass through to the original resolver. Per-thread scope + (`threading.local`) keeps concurrent requests on other threads unaffected. + """ + global _install_count, _original_create_connection + + with _install_lock: + if _install_count == 0: + _original_create_connection = urllib3_connection.create_connection + urllib3_connection.create_connection = _pinned_create_connection + _install_count += 1 + + stack = _get_pin_stack() + pins: dict[str, list] = dict(stack[-1]) if stack else {} + pins[hostname] = addr_infos + stack.append(pins) + + try: + yield + finally: + stack.pop() + with _install_lock: + _install_count -= 1 + if _install_count == 0 and _original_create_connection is not None: + urllib3_connection.create_connection = _original_create_connection + _original_create_connection = None + + +def _resolve_and_validate(url: str) -> tuple[bool, str, str | None, list | None]: + """Resolve a URL's hostname and check every address is safe to contact. + + Returns (is_safe, reason, hostname, addr_infos). When safe, the caller must + use _pin_dns(hostname, addr_infos) so the subsequent connection cannot pick + up a different (e.g. DNS-rebound) address. + """ try: parsed = urlparse(url) if parsed.scheme not in {"http", "https"}: - return False, f"Unsupported URL scheme: {parsed.scheme or ''}" + return False, f"Unsupported URL scheme: {parsed.scheme or ''}", None, None hostname = parsed.hostname if not hostname: - return False, "Could not parse hostname from URL" + return False, "Could not parse hostname from URL", None, None try: addr_infos = socket.getaddrinfo(hostname, None) except socket.gaierror: - return False, f"Could not resolve hostname: {hostname}" + return False, f"Could not resolve hostname: {hostname}", hostname, None + + if not addr_infos: + return False, f"Could not resolve hostname: {hostname}", hostname, None for addr_info in addr_infos: ip_str = addr_info[4][0] try: ip = ipaddress.ip_address(ip_str) except ValueError: - continue + return False, f"Could not parse resolved address: {ip_str}", hostname, None if ip.is_private or ip.is_loopback or ip.is_link_local or ip.is_reserved: - return False, f"URL resolves to blocked address: {ip_str}" + return False, f"URL resolves to blocked address: {ip_str}", hostname, None - return True, "" + return True, "", hostname, addr_infos except Exception as e: # noqa: BLE001 - return False, f"URL validation error: {e}" + return False, f"URL validation error: {e}", None, None + + +def _is_url_safe(url: str) -> tuple[bool, str]: + """Check if a URL is safe to request (not targeting private/internal networks).""" + is_safe, reason, _, _ = _resolve_and_validate(url) + return is_safe, reason def _blocked_response(url: str, reason: str) -> dict[str, Any]: @@ -56,23 +184,30 @@ def _request_with_safe_redirects( timeout: int, **kwargs: Any, ) -> tuple[requests.Response | None, dict[str, Any] | None]: - """Issue a request while validating every redirect target before following it.""" + """Issue a request while validating every redirect target before following it. + + The hostname is resolved once per hop and the connection is forced to use + the validated addresses, closing the DNS-rebinding race where a controlled + resolver returns a public IP at validation time and a private IP at connect + time. + """ current_method = method.upper() current_url = url request_kwargs = dict(kwargs) for redirect_count in range(_MAX_REDIRECTS + 1): - is_safe, reason = _is_url_safe(current_url) - if not is_safe: + is_safe, reason, hostname, addr_infos = _resolve_and_validate(current_url) + if not is_safe or hostname is None or addr_infos is None: return None, _blocked_response(current_url, reason) - response = requests.request( - current_method, - current_url, - timeout=timeout, - allow_redirects=False, - **request_kwargs, - ) + with _pin_dns(hostname, addr_infos): + response = requests.request( + current_method, + current_url, + timeout=timeout, + allow_redirects=False, + **request_kwargs, + ) if not response.is_redirect and not response.is_permanent_redirect: return response, None diff --git a/tests/test_http_security.py b/tests/test_http_security.py index 02d36527..15b69d83 100644 --- a/tests/test_http_security.py +++ b/tests/test_http_security.py @@ -1,6 +1,7 @@ from __future__ import annotations import importlib +import socket as real_socket import sys import types @@ -20,6 +21,16 @@ _PERMANENT_REDIRECT_CODES = {301, 308} _NO_JSON = object() +def _addr_info(ip: str, port: int | None = None) -> tuple: + return ( + real_socket.AF_INET, + real_socket.SOCK_STREAM, + 6, + "", + (ip, port or 0), + ) + + class FakeResponse: def __init__( self, @@ -72,6 +83,12 @@ def test_fetch_url_blocks_private_ip_without_issuing_a_request(monkeypatch) -> N def test_fetch_url_blocks_redirects_to_private_ips(monkeypatch) -> None: calls: list[tuple[str, str, bool]] = [] + def fake_getaddrinfo(host, port, *args, **kwargs): # type: ignore[no-untyped-def] + ip = "93.184.216.34" if host == "example.com" else host + return [_addr_info(ip, port)] + + monkeypatch.setattr(http_request_tool.socket, "getaddrinfo", fake_getaddrinfo) + def fake_request( method: str, url: str, *, timeout: int, allow_redirects: bool, **kwargs ) -> FakeResponse: # type: ignore[no-untyped-def] @@ -90,3 +107,189 @@ def test_fetch_url_blocks_redirects_to_private_ips(monkeypatch) -> None: assert result["status_code"] == 0 assert result["url"] == "http://169.254.169.254/latest/meta-data" assert "Request blocked" in result["error"] + + +class _FakeSocket: + """Records connect() targets without performing real network I/O.""" + + instances: list = [] + + def __init__(self, family, socktype, proto): + self.family = family + self.socktype = socktype + self.proto = proto + self.connected_to = None + self.timeout = None + self.sockopts: list = [] + self.closed = False + _FakeSocket.instances.append(self) + + def settimeout(self, t): + self.timeout = t + + def setsockopt(self, *opt): + self.sockopts.append(opt) + + def bind(self, _addr): + pass + + def connect(self, address): + self.connected_to = address + + def close(self): + self.closed = True + + +def test_pinned_dns_blocks_rebinding_to_private_ip(monkeypatch) -> None: + """A resolver that flips public -> private must not be able to rebind. + + Validation sees a public IP; a later resolution would return 127.0.0.1. + The connection layer (urllib3's create_connection) must observe the pinned + public IP, not the private IP. + """ + hostname = "rebind.example.com" + public_addr = "93.184.216.34" + private_addr = "127.0.0.1" + + call_count = {"n": 0} + + def fake_getaddrinfo(host, port, *args, **kwargs): # type: ignore[no-untyped-def] + call_count["n"] += 1 + ip = public_addr if call_count["n"] == 1 else private_addr + return [_addr_info(ip, port)] + + monkeypatch.setattr(http_request_tool.socket, "getaddrinfo", fake_getaddrinfo) + + _FakeSocket.instances = [] + monkeypatch.setattr(http_request_tool.socket, "socket", _FakeSocket) + + def fake_request(method, url, *, timeout, allow_redirects, **kwargs): # type: ignore[no-untyped-def] + # Drive urllib3's connection helper the way urllib3 itself would. + http_request_tool.urllib3_connection.create_connection((hostname, 80)) + return FakeResponse(status_code=200, url=url, text="ok") + + monkeypatch.setattr(http_request_tool.requests, "request", fake_request) + + result = http_request_tool.http_request(f"http://{hostname}/probe") + + assert len(_FakeSocket.instances) == 1 + sock = _FakeSocket.instances[0] + assert sock.connected_to == (public_addr, 80), ( + f"Connection step must target pinned public IP, got {sock.connected_to}" + ) + assert result["status_code"] == 200 + + +def test_rebinding_to_only_private_ips_is_blocked(monkeypatch) -> None: + """If the very first resolution returns a private IP, validation must reject.""" + hostname = "evil.example.com" + private_addr = "169.254.169.254" + + def fake_getaddrinfo(host, port, *args, **kwargs): # type: ignore[no-untyped-def] + return [_addr_info(private_addr, port)] + + monkeypatch.setattr(http_request_tool.socket, "getaddrinfo", fake_getaddrinfo) + + def fail_request(*args, **kwargs): # type: ignore[no-untyped-def] + raise AssertionError("request should not be issued for blocked URLs") + + monkeypatch.setattr(http_request_tool.requests, "request", fail_request) + + result = http_request_tool.http_request(f"http://{hostname}/") + + assert result["status_code"] == 0 + assert "Request blocked" in result["content"] + + +def test_pin_does_not_affect_other_hostnames(monkeypatch) -> None: + """The pinned create_connection must only override the validated hostname.""" + hostname = "pinned.example.com" + public_addr = "93.184.216.34" + other_hostname = "other.example.com" + + addr_infos = [_addr_info(public_addr)] + + fallthrough_calls: list = [] + + def fake_original_create_connection(address, *args, **kwargs): # type: ignore[no-untyped-def] + fallthrough_calls.append(address) + return ("fallthrough", address) + + monkeypatch.setattr( + http_request_tool.urllib3_connection, + "create_connection", + fake_original_create_connection, + ) + + with http_request_tool._pin_dns(hostname, addr_infos): + # The pinned wrapper is now installed; calling it for the pinned host + # must NOT delegate to the real create_connection. + try: + pinned_sock = http_request_tool._pinned_create_connection((hostname, 80)) + if isinstance(pinned_sock, real_socket.socket): + assert pinned_sock.getpeername()[0] == public_addr or True + pinned_sock.close() + except OSError: + # Expected — no actual server at the pinned IP. The point is that + # the fallthrough was NOT used. + pass + + # Other host MUST fall through to the (mocked) real resolver. + other_result = http_request_tool._pinned_create_connection((other_hostname, 443)) + + assert fallthrough_calls == [(other_hostname, 443)], ( + f"Pin must only override the pinned hostname, got fallthrough calls: {fallthrough_calls}" + ) + assert other_result == ("fallthrough", (other_hostname, 443)) + + +def test_pin_install_count_unwinds() -> None: + """After all _pin_dns blocks exit, urllib3's create_connection is restored.""" + sentinel_original = http_request_tool.urllib3_connection.create_connection + addr_infos = [_addr_info("93.184.216.34")] + + with http_request_tool._pin_dns("a.example.com", addr_infos): + assert ( + http_request_tool.urllib3_connection.create_connection + is http_request_tool._pinned_create_connection + ) + with http_request_tool._pin_dns("b.example.com", addr_infos): + assert ( + http_request_tool.urllib3_connection.create_connection + is http_request_tool._pinned_create_connection + ) + + assert http_request_tool.urllib3_connection.create_connection is sentinel_original + assert http_request_tool._install_count == 0 + assert http_request_tool._original_create_connection is None + + +def test_pinned_connection_propagates_timeout_and_socket_options(monkeypatch) -> None: + """urllib3 calls create_connection with a positional timeout and keyword + socket_options; the pinned wrapper must forward both to the underlying socket + so connect timeouts and TCP options aren't silently dropped. + """ + hostname = "pinned.example.com" + public_addr = "93.184.216.34" + addr_infos = [_addr_info(public_addr)] + + _FakeSocket.instances = [] + monkeypatch.setattr(http_request_tool.socket, "socket", _FakeSocket) + + sock_opts = [(real_socket.IPPROTO_TCP, real_socket.TCP_NODELAY, 1)] + + with http_request_tool._pin_dns(hostname, addr_infos): + # Match how urllib3.connection calls create_connection: + # positional timeout, keyword source_address + socket_options. + http_request_tool._pinned_create_connection( + (hostname, 80), + 7.5, + source_address=None, + socket_options=sock_opts, + ) + + assert len(_FakeSocket.instances) == 1 + sock = _FakeSocket.instances[0] + assert sock.connected_to == (public_addr, 80) + assert sock.timeout == 7.5, f"connect timeout was dropped: {sock.timeout!r}" + assert sock.sockopts == sock_opts, f"socket_options were dropped: {sock.sockopts!r}"