diff --git a/backend/open_webui/retrieval/web/utils.py b/backend/open_webui/retrieval/web/utils.py index 535e1f1b8f..f653a60131 100644 --- a/backend/open_webui/retrieval/web/utils.py +++ b/backend/open_webui/retrieval/web/utils.py @@ -20,7 +20,6 @@ from typing import ( ) import aiohttp -import aiohttp.resolver import certifi import urllib3.connection import urllib3.connectionpool @@ -103,6 +102,12 @@ def _is_global_addr(ip: str) -> bool: return all(ip.is_global for ip in embedded) +def _assert_host_allowed(host: str | None) -> None: + if WEB_FETCH_FILTER_LIST and not is_host_allowed(host, WEB_FETCH_FILTER_LIST): + log.warning(f'Blocked by filter list: {host}') + raise ValueError(ERROR_MESSAGES.INVALID_URL) + + def validate_url(url: Union[str, Sequence[str]]): if isinstance(url, str): if isinstance(validators.url(url), validators.ValidationError): @@ -123,13 +128,9 @@ def validate_url(url: Union[str, Sequence[str]]): log.warning(f'Blocked non-HTTP(S) protocol: {parsed_url.scheme} in URL: {url}') raise ValueError(ERROR_MESSAGES.INVALID_URL) - # Blocklist check using unified filtering logic - if WEB_FETCH_FILTER_LIST: - # Match on the parsed hostname, not the full URL: a path component would - # otherwise let any URL slip past a hostname-based block/allow entry. - if not is_host_allowed(parsed_url.hostname, WEB_FETCH_FILTER_LIST): - log.warning(f'URL blocked by filter list: {url}') - raise ValueError(ERROR_MESSAGES.INVALID_URL) + # Match on the parsed hostname, not the full URL: a path component would + # otherwise let any URL slip past a hostname-based block/allow entry. + _assert_host_allowed(parsed_url.hostname) if not ENABLE_LOCAL_WEB_FETCH: # Local web fetch is disabled, filter out URLs that resolve to non-global IP addresses. @@ -137,7 +138,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 - # DNS rebinding is mitigated at the connection layer; see _SSRFSafeResolver / _SSRFSafeAdapter + # DNS rebinding is mitigated at the connection layer; see _SSRFSafeConnector / _SSRFSafeAdapter for ip in ipv4_addresses + ipv6_addresses: if not _is_global_addr(ip): raise ValueError(ERROR_MESSAGES.INVALID_URL) @@ -218,7 +219,7 @@ class _SafeHTTPSPool(urllib3.connectionpool.HTTPSConnectionPool): class _SSRFSafeAdapter(HTTPAdapter): - """requests transport adapter that validates resolved IPs at connect time.""" + """requests adapter that rejects filter-listed request targets and non-global IPs at connect time.""" def init_poolmanager(self, *args, **kwargs): super().init_poolmanager(*args, **kwargs) @@ -227,12 +228,23 @@ class _SSRFSafeAdapter(HTTPAdapter): 'https': _SafeHTTPSPool, } + def send(self, request, *args, **kwargs): + # Per request, not per connection: the connection layer sees the proxy. + _assert_host_allowed(urllib.parse.urlparse(request.url).hostname) + return super().send(request, *args, **kwargs) -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) +class _SSRFSafeConnector(aiohttp.TCPConnector): + """Rejects filter-listed request targets, and non-global IPs on each new connection.""" + + async def connect(self, req, traces, timeout): + # Per request, not per connection: _resolve_host sees the proxy and pooled reuse skips it. + _assert_host_allowed(req.url.host) + return await super().connect(req, traces, timeout) + + async def _resolve_host(self, host, port, traces=None): + # aiohttp answers IP-literal hosts itself without consulting a resolver. + results = await super()._resolve_host(host, port, traces=traces) if not ENABLE_LOCAL_WEB_FETCH: for entry in results: if not _is_global_addr(entry['host']): @@ -241,13 +253,13 @@ class _SSRFSafeResolver(aiohttp.resolver.DefaultResolver): def get_ssrf_safe_session() -> aiohttp.ClientSession: - """A one-off aiohttp session that re-validates the connect-time IP via _SSRFSafeResolver, + """A one-off aiohttp session that re-validates every connection via _SSRFSafeConnector, defeating DNS rebinding. Use for validate_url-gated fetches of user-supplied URLs that must not use the shared (rebinding-vulnerable) pool. Use as a context manager so it is closed: ``async with get_ssrf_safe_session() as session: ...``. """ return aiohttp.ClientSession( - connector=aiohttp.TCPConnector(resolver=_SSRFSafeResolver()), + connector=_SSRFSafeConnector(), timeout=aiohttp.ClientTimeout(total=AIOHTTP_CLIENT_TIMEOUT), trust_env=True, ) @@ -793,7 +805,7 @@ class SafeWebBaseLoader(WebBaseLoader): self.session.mount('https://', _SSRFSafeAdapter()) async def _fetch(self, url: str, retries: int = 3, cooldown: int = 2, backoff: float = 1.5) -> str: - connector = aiohttp.TCPConnector(resolver=_SSRFSafeResolver()) + connector = _SSRFSafeConnector() async with aiohttp.ClientSession(trust_env=self.trust_env, connector=connector) as session: for i in range(retries): try: