mirror of
https://github.com/Sea-Haven-Industries/open-swe.git
synced 2026-09-30 10:23:14 +00:00
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:
parent
51bdde93b9
commit
a9331e78e6
2 changed files with 357 additions and 19 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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}"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue