fix: mitigate DNS rebinding in web loader fetch paths (#24759)
validate_url() resolves DNS to check IPs but discards the result; the HTTP client resolves again independently. Between those two lookups an attacker can swap the DNS record from a public IP to an internal one (DNS rebinding). Push the IP-is-global check into the actual connection layer so the validated resolution is the one used for the TCP connect: - aiohttp (_fetch): _SSRFSafeResolver wraps DefaultResolver and rejects non-global IPs at resolve time (zero TOCTOU window). - requests (_scrape): _SSRFSafeAdapter mounts custom urllib3 connection classes whose _new_conn resolves, validates, and connects to the validated IP in one shot (zero TOCTOU window). Both paths respect ENABLE_RAG_LOCAL_WEB_FETCH (skip validation when on). Co-authored-by: Claude Opus 4.6 <noreply@anthropic.com>
This commit is contained in:
co-authored by
Claude Opus 4.6
parent
cfa6908d57
commit
854440f703
@@ -19,9 +19,13 @@ from typing import (
|
||||
)
|
||||
|
||||
import aiohttp
|
||||
import aiohttp.resolver
|
||||
import certifi
|
||||
import requests
|
||||
import urllib3.connection
|
||||
import urllib3.connectionpool
|
||||
import validators
|
||||
from requests.adapters import HTTPAdapter
|
||||
from fastapi.concurrency import run_in_threadpool
|
||||
from langchain_community.document_loaders import PlaywrightURLLoader, WebBaseLoader
|
||||
from langchain_community.document_loaders.base import BaseLoader
|
||||
@@ -94,7 +98,7 @@ def validate_url(url: Union[str, Sequence[str]]):
|
||||
# Get IPv4 and IPv6 addresses
|
||||
ipv4_addresses, ipv6_addresses = resolve_hostname(parsed_url.hostname)
|
||||
# Check if any of the resolved addresses are private
|
||||
# This is technically still vulnerable to DNS rebinding attacks, as we don't control WebBaseLoader
|
||||
# DNS rebinding is mitigated at the connection layer; see _SSRFSafeResolver / _SSRFSafeAdapter
|
||||
for ip in ipv4_addresses + ipv6_addresses:
|
||||
addr = ipaddress.ip_address(ip)
|
||||
if not addr.is_global:
|
||||
@@ -118,6 +122,81 @@ def safe_validate_urls(url: Sequence[str]) -> Sequence[str]:
|
||||
return valid_urls
|
||||
|
||||
|
||||
def _ssrf_safe_new_conn(self):
|
||||
"""Resolve DNS, validate all IPs are global, connect to validated IP.
|
||||
|
||||
Replaces urllib3's _new_conn so the DNS lookup that feeds the actual TCP
|
||||
connect is the same one we validate — no second resolution, no rebinding
|
||||
window.
|
||||
"""
|
||||
host = getattr(self, '_dns_host', self.host)
|
||||
port = self.port
|
||||
infos = socket.getaddrinfo(host, port, 0, socket.SOCK_STREAM)
|
||||
if not infos:
|
||||
raise OSError(f'getaddrinfo for {host!r} returned empty list')
|
||||
if not ENABLE_RAG_LOCAL_WEB_FETCH:
|
||||
for _, _, _, _, sa in infos:
|
||||
if not ipaddress.ip_address(sa[0]).is_global:
|
||||
raise ValueError(ERROR_MESSAGES.INVALID_URL)
|
||||
err = None
|
||||
for fam, typ, proto, _, sa in infos:
|
||||
sock = None
|
||||
try:
|
||||
sock = socket.socket(fam, typ, proto)
|
||||
if self.timeout is not socket._GLOBAL_DEFAULT_TIMEOUT:
|
||||
sock.settimeout(self.timeout)
|
||||
if getattr(self, 'source_address', None):
|
||||
sock.bind(self.source_address)
|
||||
for opt in getattr(self, 'socket_options', None) or ():
|
||||
sock.setsockopt(*opt)
|
||||
sock.connect(sa)
|
||||
return sock
|
||||
except OSError as exc:
|
||||
err = exc
|
||||
if sock is not None:
|
||||
sock.close()
|
||||
raise err or OSError(f'connect to {host!r}:{port} failed')
|
||||
|
||||
|
||||
class _SafeHTTPConn(urllib3.connection.HTTPConnection):
|
||||
_new_conn = _ssrf_safe_new_conn
|
||||
|
||||
|
||||
class _SafeHTTPSConn(urllib3.connection.HTTPSConnection):
|
||||
_new_conn = _ssrf_safe_new_conn
|
||||
|
||||
|
||||
class _SafeHTTPPool(urllib3.connectionpool.HTTPConnectionPool):
|
||||
ConnectionCls = _SafeHTTPConn
|
||||
|
||||
|
||||
class _SafeHTTPSPool(urllib3.connectionpool.HTTPSConnectionPool):
|
||||
ConnectionCls = _SafeHTTPSConn
|
||||
|
||||
|
||||
class _SSRFSafeAdapter(HTTPAdapter):
|
||||
"""requests transport adapter that validates resolved IPs at connect time."""
|
||||
|
||||
def init_poolmanager(self, *args, **kwargs):
|
||||
super().init_poolmanager(*args, **kwargs)
|
||||
self.poolmanager.pool_classes_by_scheme = {
|
||||
'http': _SafeHTTPPool,
|
||||
'https': _SafeHTTPSPool,
|
||||
}
|
||||
|
||||
|
||||
class _SSRFSafeResolver(aiohttp.resolver.DefaultResolver):
|
||||
"""aiohttp resolver that rejects non-global IPs unless local fetch is on."""
|
||||
|
||||
async def resolve(self, host, port=0, family=socket.AF_INET):
|
||||
results = await super().resolve(host, port, family)
|
||||
if not ENABLE_RAG_LOCAL_WEB_FETCH:
|
||||
for entry in results:
|
||||
if not ipaddress.ip_address(entry['host']).is_global:
|
||||
raise ValueError(ERROR_MESSAGES.INVALID_URL)
|
||||
return results
|
||||
|
||||
|
||||
def extract_metadata(soup, url):
|
||||
metadata = {'source': url}
|
||||
if title := soup.find('title'):
|
||||
@@ -570,8 +649,12 @@ class SafeWebBaseLoader(WebBaseLoader):
|
||||
'allow_redirects': AIOHTTP_CLIENT_ALLOW_REDIRECTS,
|
||||
}
|
||||
|
||||
self.session.mount('http://', _SSRFSafeAdapter())
|
||||
self.session.mount('https://', _SSRFSafeAdapter())
|
||||
|
||||
async def _fetch(self, url: str, retries: int = 3, cooldown: int = 2, backoff: float = 1.5) -> str:
|
||||
async with aiohttp.ClientSession(trust_env=self.trust_env) as session:
|
||||
connector = aiohttp.TCPConnector(resolver=_SSRFSafeResolver())
|
||||
async with aiohttp.ClientSession(trust_env=self.trust_env, connector=connector) as session:
|
||||
for i in range(retries):
|
||||
try:
|
||||
kwargs: Dict = dict(
|
||||
|
||||
Reference in New Issue
Block a user