refac: match the web fetch filter list against resolved addresses

The filter list is now evaluated against the hostname together with the addresses it resolves to, at URL validation and on each connection, on both transports. An IPv6 address is also matched by the IPv4 address it carries.
This commit is contained in:
Classic298 2026-08-19 18:17:35 +02:00
parent e84a2b5d26
commit ee4774b545
2 changed files with 57 additions and 33 deletions

View file

@ -63,7 +63,7 @@ from open_webui.retrieval.loaders.external_web import ExternalWebLoader
from open_webui.retrieval.loaders.microsoft_web_iq import MicrosoftWebIQLoader
from open_webui.retrieval.loaders.tavily import TavilyLoader
from open_webui.retrieval.web.firecrawl import scrape_firecrawl_url
from open_webui.utils.misc import is_host_allowed
from open_webui.utils.misc import is_host_allowed, is_host_blocked
log = logging.getLogger(__name__)
@ -79,12 +79,14 @@ def resolve_hostname(hostname):
return ipv4_addresses, ipv6_addresses
def _is_global_addr(ip: str) -> bool:
_NAT64_PREFIX_48 = b'\x00\x64\xff\x9b\x00\x01'
def _embedded_ipv4(ip: str) -> list[ipaddress.IPv4Address]:
"""The IPv4 addresses an IPv6 address carries inside it: mapped, 6to4, teredo and NAT64."""
addr = ipaddress.ip_address(ip)
if not addr.is_global:
return False
if not isinstance(addr, ipaddress.IPv6Address):
return True
return []
embedded = []
if addr.ipv4_mapped:
@ -95,16 +97,22 @@ def _is_global_addr(ip: str) -> bool:
embedded.extend(addr.teredo)
b = addr.packed
if b[:12] == b'\x00' * 12:
if b[:12] == b'\x00' * 12 or b[:12] == b'\x00\x64\xff\x9b' + b'\x00' * 8:
embedded.append(ipaddress.IPv4Address(b[12:]))
elif b[:12] == b'\x00\x64\xff\x9b' + b'\x00' * 8:
embedded.append(ipaddress.IPv4Address(b[12:]))
elif b[:6] == b'\x00\x64\xff\x9b\x00\x01':
if b[8] != 0:
return False
elif b[:6] == _NAT64_PREFIX_48:
embedded.append(ipaddress.IPv4Address(bytes((b[6], b[7], b[9], b[10]))))
return all(ip.is_global for ip in embedded)
return embedded
def _is_global_addr(ip: str) -> bool:
addr = ipaddress.ip_address(ip)
if not addr.is_global:
return False
# The NAT64 /48 prefix reserves the u-octet, so a non-zero one is malformed.
if isinstance(addr, ipaddress.IPv6Address) and addr.packed[:6] == _NAT64_PREFIX_48 and addr.packed[8] != 0:
return False
return all(embedded.is_global for embedded in _embedded_ipv4(ip))
def _assert_host_allowed(host: str | None) -> None:
@ -113,6 +121,18 @@ def _assert_host_allowed(host: str | None) -> None:
raise ValueError(ERROR_MESSAGES.INVALID_URL)
def _assert_addresses_allowed(addresses: Sequence[str]) -> None:
# An IPv6 address can carry a blocked IPv4 address inside it, so match both spellings.
candidates = [*addresses, *(str(ipv4) for address in addresses for ipv4 in _embedded_ipv4(address))]
if is_host_blocked(candidates, WEB_FETCH_FILTER_LIST):
log.warning(f'Blocked by filter list: {", ".join(addresses)}')
raise ValueError(ERROR_MESSAGES.INVALID_URL)
if not ENABLE_LOCAL_WEB_FETCH:
for address in addresses:
if not _is_global_addr(address):
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):
@ -137,16 +157,18 @@ def validate_url(url: Union[str, Sequence[str]]):
# 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.
parsed_url = urllib.parse.urlparse(url)
# Get IPv4 and IPv6 addresses
try:
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 _SSRFSafeConnector / _SSRFSafeAdapter
for ip in ipv4_addresses + ipv6_addresses:
if not _is_global_addr(ip):
raise ValueError(ERROR_MESSAGES.INVALID_URL)
except (socket.gaierror, UnicodeError) as e:
# With local fetch on, a proxied deployment can carry names only the proxy resolves.
if not ENABLE_LOCAL_WEB_FETCH:
log.warning(f'Could not resolve host {parsed_url.hostname}: {e}')
raise ValueError(ERROR_MESSAGES.INVALID_URL) from None
ipv4_addresses, ipv6_addresses = [], []
# A hostname match alone lets a DNS record point at a blocked address.
# DNS rebinding is mitigated at the connection layer; see _SSRFSafeConnector / _SSRFSafeAdapter
_assert_addresses_allowed(ipv4_addresses + ipv6_addresses)
return True
elif isinstance(url, Sequence):
return all(validate_url(u) for u in url)
@ -178,10 +200,7 @@ def _ssrf_safe_new_conn(self):
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_LOCAL_WEB_FETCH:
for _, _, _, _, sa in infos:
if not _is_global_addr(sa[0]):
raise ValueError(ERROR_MESSAGES.INVALID_URL)
_assert_addresses_allowed([sa[0] for _, _, _, _, sa in infos])
err = None
for fam, typ, proto, _, sa in infos:
sock = None
@ -250,10 +269,7 @@ class _SSRFSafeConnector(aiohttp.TCPConnector):
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']):
raise ValueError(ERROR_MESSAGES.INVALID_URL)
_assert_addresses_allowed([entry['host'] for entry in results])
return results

View file

@ -112,17 +112,25 @@ def is_host_allowed(host: Union[str, Sequence[str]], filter_list: list[str | Non
if not filter_list:
return True
allow_list, block_list = get_allow_block_lists(filter_list)
allow_list, _ = get_allow_block_lists(filter_list)
hosts = [host] if isinstance(host, str) else list(host or [])
if allow_list:
if not any(_host_matches_pattern(h, allowed) for h in hosts for allowed in allow_list):
return False
if any(_host_matches_pattern(h, blocked) for h in hosts for blocked in block_list):
return False
return not is_host_blocked(hosts, filter_list)
return True
def is_host_blocked(host: Union[str, Sequence[str]], filter_list: list[str | None] = None) -> bool:
"""Whether a host or resolved address matches a block entry, ignoring any allow entries.
For addresses, where an allow entry cannot apply: it names a host, and the address at hand
may belong to a forward proxy rather than to the host the request is actually for.
"""
_, block_list = get_allow_block_lists(filter_list)
hosts = [host] if isinstance(host, str) else list(host or [])
return any(_host_matches_pattern(h, blocked) for h in hosts for blocked in block_list)
def get_message_list(messages_map, message_id):