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:
Classic298
2026-05-19 23:57:12 +04:00
committed by GitHub
co-authored by Claude Opus 4.6
parent cfa6908d57
commit 854440f703
+85 -2
View File
@@ -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(