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] <open-swe@users.noreply.github.com>
Co-authored-by: Johannes du Plessis <51395795+johannes117@users.noreply.github.com>
Co-authored-by: Johannes du Plessis <johannes@langchain.dev>
This commit is contained in:
open-swe[bot] 2026-05-08 22:06:25 +00:00 • committed by GitHub
parent 51bdde93b9
commit a9331e78e6
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 357 additions and 19 deletions

View file

@ -1,42 +1,170 @@
import contextlib
import ipaddress import ipaddress
import socket import socket
import threading
from collections.abc import Iterator
from typing import Any from typing import Any
from urllib.parse import urljoin, urlparse from urllib.parse import urljoin, urlparse
import requests import requests
from urllib3.util import connection as urllib3_connection
_MAX_REDIRECTS = 5 _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: try:
parsed = urlparse(url) parsed = urlparse(url)
if parsed.scheme not in {"http", "https"}: if parsed.scheme not in {"http", "https"}:
return False, f"Unsupported URL scheme: {parsed.scheme or '<missing>'}" return False, f"Unsupported URL scheme: {parsed.scheme or '<missing>'}", None, None
hostname = parsed.hostname hostname = parsed.hostname
if not hostname: if not hostname:
return False, "Could not parse hostname from URL" return False, "Could not parse hostname from URL", None, None
try: try:
addr_infos = socket.getaddrinfo(hostname, None) addr_infos = socket.getaddrinfo(hostname, None)
except socket.gaierror: 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: for addr_info in addr_infos:
ip_str = addr_info[4][0] ip_str = addr_info[4][0]
try: try:
ip = ipaddress.ip_address(ip_str) ip = ipaddress.ip_address(ip_str)
except ValueError: 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: 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 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]: def _blocked_response(url: str, reason: str) -> dict[str, Any]:
@ -56,23 +184,30 @@ def _request_with_safe_redirects(
timeout: int, timeout: int,
**kwargs: Any, **kwargs: Any,
) -> tuple[requests.Response | None, dict[str, Any] | None]: ) -> 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_method = method.upper()
current_url = url current_url = url
request_kwargs = dict(kwargs) request_kwargs = dict(kwargs)
for redirect_count in range(_MAX_REDIRECTS + 1): for redirect_count in range(_MAX_REDIRECTS + 1):
is_safe, reason = _is_url_safe(current_url) is_safe, reason, hostname, addr_infos = _resolve_and_validate(current_url)
if not is_safe: if not is_safe or hostname is None or addr_infos is None:
return None, _blocked_response(current_url, reason) return None, _blocked_response(current_url, reason)
response = requests.request( with _pin_dns(hostname, addr_infos):
current_method, response = requests.request(
current_url, current_method,
timeout=timeout, current_url,
allow_redirects=False, timeout=timeout,
**request_kwargs, allow_redirects=False,
) **request_kwargs,
)
if not response.is_redirect and not response.is_permanent_redirect: if not response.is_redirect and not response.is_permanent_redirect:
return response, None return response, None

View file

@ -1,6 +1,7 @@
from __future__ import annotations from __future__ import annotations
import importlib import importlib
import socket as real_socket
import sys import sys
import types import types
@ -20,6 +21,16 @@ _PERMANENT_REDIRECT_CODES = {301, 308}
_NO_JSON = object() _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: class FakeResponse:
def __init__( def __init__(
self, 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: def test_fetch_url_blocks_redirects_to_private_ips(monkeypatch) -> None:
calls: list[tuple[str, str, bool]] = [] 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( def fake_request(
method: str, url: str, *, timeout: int, allow_redirects: bool, **kwargs method: str, url: str, *, timeout: int, allow_redirects: bool, **kwargs
) -> FakeResponse: # type: ignore[no-untyped-def] ) -> 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["status_code"] == 0
assert result["url"] == "http://169.254.169.254/latest/meta-data" assert result["url"] == "http://169.254.169.254/latest/meta-data"
assert "Request blocked" in result["error"] 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}"