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 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 '<missing>'}"
return False, f"Unsupported URL scheme: {parsed.scheme or '<missing>'}", 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

View file

@ -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}"