mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-12 23:01:41 +00:00
Merge pull request #25906 from stuxf/fix/ssrf-url-validation
fix(proxy): add URL validation for user-supplied URLs
This commit is contained in:
commit
7c66edbf4a
44 changed files with 745 additions and 28 deletions
|
|
@ -274,6 +274,8 @@ use_client: bool = False
|
|||
ssl_verify: Union[str, bool] = True
|
||||
ssl_security_level: Optional[str] = None
|
||||
ssl_certificate: Optional[str] = None
|
||||
user_url_validation: bool = True
|
||||
user_url_allowed_hosts: List[str] = []
|
||||
ssl_ecdh_curve: Optional[
|
||||
str
|
||||
] = None # Set to 'X25519' to disable PQC and improve performance
|
||||
|
|
|
|||
|
|
@ -10,6 +10,7 @@ import litellm
|
|||
from litellm import verbose_logger
|
||||
from litellm.caching.caching import InMemoryCache
|
||||
from litellm.constants import MAX_IMAGE_URL_DOWNLOAD_SIZE_MB
|
||||
from litellm.litellm_core_utils.url_utils import async_safe_get, safe_get
|
||||
|
||||
MAX_IMGS_IN_MEMORY = 10
|
||||
|
||||
|
|
@ -84,7 +85,7 @@ async def async_convert_url_to_base64(url: str) -> str:
|
|||
client = litellm.module_level_aclient
|
||||
for _ in range(3):
|
||||
try:
|
||||
response = await client.get(url, follow_redirects=True)
|
||||
response = await async_safe_get(client, url)
|
||||
return _process_image_response(response, url)
|
||||
except litellm.ImageFetchError:
|
||||
raise
|
||||
|
|
@ -109,7 +110,7 @@ def convert_url_to_base64(url: str) -> str:
|
|||
client = litellm.module_level_client
|
||||
for _ in range(3):
|
||||
try:
|
||||
response = client.get(url, follow_redirects=True)
|
||||
response = safe_get(client, url)
|
||||
return _process_image_response(response, url)
|
||||
except litellm.ImageFetchError:
|
||||
raise
|
||||
|
|
|
|||
|
|
@ -30,6 +30,7 @@ from litellm.constants import (
|
|||
)
|
||||
from litellm.litellm_core_utils.default_encoding import encoding as default_encoding
|
||||
from litellm.llms.custom_httpx.http_handler import _get_httpx_client
|
||||
from litellm.litellm_core_utils.url_utils import safe_get
|
||||
from litellm.types.llms.anthropic import (
|
||||
AnthropicMessagesToolResultParam,
|
||||
AnthropicMessagesToolUseParam,
|
||||
|
|
@ -210,13 +211,15 @@ def get_image_dimensions(
|
|||
Tuple[int, int]: The width and height of the image.
|
||||
"""
|
||||
img_data = None
|
||||
try:
|
||||
# Try to open as URL
|
||||
client = _get_httpx_client()
|
||||
response = client.get(data)
|
||||
img_data = response.read()
|
||||
except Exception:
|
||||
# If not URL, assume it's base64
|
||||
if data.startswith(("http://", "https://")):
|
||||
try:
|
||||
client = _get_httpx_client()
|
||||
response = safe_get(client, data)
|
||||
img_data = response.read()
|
||||
except Exception:
|
||||
pass
|
||||
if img_data is None:
|
||||
# Not a URL or fetch failed — assume base64
|
||||
_header, encoded = data.split(",", 1)
|
||||
img_data = base64.b64decode(encoded)
|
||||
|
||||
|
|
|
|||
257
litellm/litellm_core_utils/url_utils.py
Normal file
257
litellm/litellm_core_utils/url_utils.py
Normal file
|
|
@ -0,0 +1,257 @@
|
|||
"""
|
||||
URL validation for user-controlled URLs.
|
||||
|
||||
Use validate_url() before fetching any URL that originates from user
|
||||
input (image_url, file_url, spec_path, etc.) to prevent SSRF attacks.
|
||||
|
||||
validate_url() resolves DNS once, validates all IPs, and rewrites the
|
||||
URL to connect to the validated IP directly — no TOCTOU gap, no DNS
|
||||
rebinding. Redirects are followed manually with validation at each hop.
|
||||
|
||||
Admins can opt out via two ``litellm`` globals (wired from proxy config):
|
||||
|
||||
- ``litellm.user_url_validation`` (bool, default True): master switch.
|
||||
When False, ``safe_get``/``async_safe_get`` perform a plain fetch with
|
||||
no DNS check, no block list, and no rewrite.
|
||||
- ``litellm.user_url_allowed_hosts`` (List[str], default []): per-host
|
||||
allowlist. Entries are ``hostname`` or ``hostname:port`` (IPv6 hosts as
|
||||
``[addr]`` / ``[addr]:port``). Matching hosts skip the blocked-networks
|
||||
check but still resolve DNS and still rewrite HTTP to the resolved IP.
|
||||
"""
|
||||
|
||||
import socket
|
||||
from ipaddress import ip_address, ip_network
|
||||
from typing import Any, List, Set, Tuple
|
||||
from urllib.parse import urlparse, urlunparse
|
||||
|
||||
import httpx
|
||||
|
||||
import litellm
|
||||
|
||||
# Globally-routable IPs that are cloud-internal. Everything else
|
||||
# non-public is caught by ``not ip.is_global`` (RFC 6890, as implemented by
|
||||
# Python's ``ipaddress`` module). This list only holds IPs that are
|
||||
# publicly routable *and* point to cloud-fabric services reachable from
|
||||
# inside a VM via special in-fabric routing.
|
||||
_CLOUD_METADATA_EXCEPTIONS = [
|
||||
ip_network("168.63.129.16/32"), # Azure Wire Server
|
||||
]
|
||||
|
||||
_ALLOWED_SCHEMES = ("http", "https")
|
||||
|
||||
|
||||
class SSRFError(ValueError):
|
||||
"""Raised when a URL targets a blocked network."""
|
||||
|
||||
pass
|
||||
|
||||
|
||||
def _is_blocked_ip(addr: str) -> bool:
|
||||
"""Return True for any IP not safe to reach from a user-supplied URL.
|
||||
|
||||
Policy: default-deny via ``ip.is_global`` (RFC 6890), plus an explicit
|
||||
exception list for globally-routable cloud-fabric IPs that are still
|
||||
dangerous from inside a cloud VM (currently just Azure Wire Server).
|
||||
Unparseable addresses fail closed.
|
||||
"""
|
||||
try:
|
||||
ip = ip_address(addr)
|
||||
except ValueError:
|
||||
return True # fail-closed: unparseable addresses are blocked
|
||||
if ip.version == 6 and hasattr(ip, "ipv4_mapped") and ip.ipv4_mapped:
|
||||
ip = ip.ipv4_mapped
|
||||
if not ip.is_global or ip.is_multicast:
|
||||
return True
|
||||
return any(ip in net for net in _CLOUD_METADATA_EXCEPTIONS)
|
||||
|
||||
|
||||
def _normalize_host(host: str) -> str:
|
||||
"""Lowercase and strip a trailing dot from a hostname."""
|
||||
return host.lower().rstrip(".")
|
||||
|
||||
|
||||
def _format_host_header(hostname: str, port: int, default_port: int) -> str:
|
||||
"""Build an RFC 7230 Host header value, bracketing IPv6 literals."""
|
||||
bracketed = f"[{hostname}]" if ":" in hostname else hostname
|
||||
if port == default_port:
|
||||
return bracketed
|
||||
return f"{bracketed}:{port}"
|
||||
|
||||
|
||||
def _is_host_allowlisted(hostname: str, effective_port: int) -> bool:
|
||||
"""Check whether a host is in the admin-configured allowlist.
|
||||
|
||||
Admin entries may be ``hostname`` (any port) or ``hostname:port``. IPv6
|
||||
literals are written bracketed (``[::1]`` / ``[::1]:8080``). Matching
|
||||
is case-insensitive on the hostname.
|
||||
"""
|
||||
configured: List[str] = getattr(litellm, "user_url_allowed_hosts", []) or []
|
||||
if not configured:
|
||||
return False
|
||||
normalized_host = _normalize_host(hostname)
|
||||
host_repr = f"[{normalized_host}]" if ":" in normalized_host else normalized_host
|
||||
candidates: Set[str] = {host_repr, f"{host_repr}:{effective_port}"}
|
||||
allowlist: Set[str] = {_normalize_host(entry) for entry in configured if entry}
|
||||
return bool(candidates & allowlist)
|
||||
|
||||
|
||||
def validate_url(url: str) -> Tuple[str, str]:
|
||||
"""
|
||||
Validate a user-supplied URL and rewrite it to connect to a validated IP.
|
||||
|
||||
Resolves the hostname, checks all resolved IPs against blocked networks,
|
||||
then returns a rewritten URL that points to the validated IP along with
|
||||
the original hostname (for use in the Host header).
|
||||
|
||||
This eliminates DNS rebinding because the caller connects to the IP we
|
||||
validated, not the hostname that could rebind. Callers should also disable
|
||||
follow_redirects to prevent redirect-based SSRF bypasses.
|
||||
|
||||
Args:
|
||||
url: The user-supplied URL to validate.
|
||||
|
||||
Returns:
|
||||
Tuple of (rewritten_url, host_header).
|
||||
The rewritten URL has the hostname replaced with the validated IP.
|
||||
The host_header value should be sent as the Host header.
|
||||
|
||||
Raises:
|
||||
SSRFError: If the URL scheme is invalid or the hostname resolves
|
||||
to a private/internal IP address.
|
||||
"""
|
||||
parsed = urlparse(url)
|
||||
|
||||
if parsed.scheme not in _ALLOWED_SCHEMES:
|
||||
raise SSRFError(f"URL scheme '{parsed.scheme}' is not allowed")
|
||||
|
||||
hostname = parsed.hostname
|
||||
if not hostname:
|
||||
raise SSRFError("URL has no hostname")
|
||||
|
||||
port = parsed.port
|
||||
default_port = 443 if parsed.scheme == "https" else 80
|
||||
effective_port = port if port is not None else default_port
|
||||
host_header = _format_host_header(hostname, effective_port, default_port)
|
||||
|
||||
is_allowlisted = _is_host_allowlisted(hostname, effective_port)
|
||||
|
||||
# Resolve hostname and validate ALL addresses
|
||||
try:
|
||||
addrinfo = socket.getaddrinfo(
|
||||
hostname, effective_port, proto=socket.IPPROTO_TCP
|
||||
)
|
||||
except socket.gaierror as e:
|
||||
raise SSRFError(f"DNS resolution failed for '{hostname}': {e}")
|
||||
|
||||
if not addrinfo:
|
||||
raise SSRFError(f"No addresses found for '{hostname}'")
|
||||
|
||||
if not is_allowlisted:
|
||||
for family, type_, proto, canonname, sockaddr in addrinfo:
|
||||
if _is_blocked_ip(sockaddr[0]):
|
||||
raise SSRFError(
|
||||
f"URL targets a blocked address ({sockaddr[0]}). "
|
||||
"If this is a legitimate internal service, add the host "
|
||||
"to `user_url_allowed_hosts` in general_settings."
|
||||
)
|
||||
|
||||
# For HTTPS with SSL verification enabled, TLS certificate validation
|
||||
# binds the connection to the hostname — DNS rebinding can't redirect
|
||||
# to a different server because the cert wouldn't match.
|
||||
# When SSL verification is disabled, this defense doesn't apply, so
|
||||
# we rewrite to the validated IP like HTTP.
|
||||
ssl_verify = getattr(litellm, "ssl_verify", True)
|
||||
if parsed.scheme == "https" and ssl_verify is not False:
|
||||
return url, host_header
|
||||
|
||||
# For HTTP, rewrite URL to connect to the validated IP directly
|
||||
# to prevent DNS rebinding (no TLS to bind the connection).
|
||||
validated_ip = addrinfo[0][4][0]
|
||||
is_ipv6 = addrinfo[0][0] == socket.AF_INET6
|
||||
ip_host = f"[{validated_ip}]" if is_ipv6 else validated_ip
|
||||
|
||||
if port is not None:
|
||||
new_netloc = f"{ip_host}:{port}"
|
||||
else:
|
||||
new_netloc = ip_host
|
||||
|
||||
rewritten = urlunparse(
|
||||
(parsed.scheme, new_netloc, parsed.path, parsed.params, parsed.query, "")
|
||||
)
|
||||
|
||||
return rewritten, host_header
|
||||
|
||||
|
||||
_MAX_REDIRECTS = 10
|
||||
|
||||
|
||||
def _extract_redirect_url(response: Any, request_url: str) -> str:
|
||||
"""Extract and resolve the redirect target from a response's Location header."""
|
||||
location = response.headers.get("location")
|
||||
if not location:
|
||||
raise SSRFError("Redirect response has no Location header")
|
||||
# Resolve relative URLs against the request URL
|
||||
return str(httpx.URL(request_url).join(location))
|
||||
|
||||
|
||||
def safe_get(client: Any, url: str, **kwargs: Any) -> Any:
|
||||
"""
|
||||
Fetch a user-supplied URL with SSRF protection on every redirect hop.
|
||||
|
||||
Validates the initial URL and each redirect target before making the
|
||||
request. No DNS rebinding (resolve-and-rewrite). No redirect bypass
|
||||
(each hop validated). No breaking change for legitimate CDN redirects.
|
||||
|
||||
When ``litellm.user_url_validation`` is False, validation is bypassed
|
||||
and this function delegates to ``client.get(url, follow_redirects=True)``.
|
||||
|
||||
Args:
|
||||
client: An httpx.Client (sync).
|
||||
url: The user-supplied URL.
|
||||
**kwargs: Additional kwargs passed to client.get().
|
||||
|
||||
Returns:
|
||||
The final httpx.Response.
|
||||
"""
|
||||
if not getattr(litellm, "user_url_validation", True):
|
||||
kwargs.setdefault("follow_redirects", True)
|
||||
return client.get(url, **kwargs)
|
||||
kwargs.pop("follow_redirects", None)
|
||||
caller_headers = kwargs.pop("headers", {})
|
||||
for _ in range(_MAX_REDIRECTS):
|
||||
validated_url, original_host = validate_url(url)
|
||||
response = client.get(
|
||||
validated_url,
|
||||
headers={**caller_headers, "Host": original_host},
|
||||
follow_redirects=False,
|
||||
**kwargs,
|
||||
)
|
||||
if not response.is_redirect:
|
||||
return response
|
||||
# Resolve the next hop against the ORIGINAL (pre-rewrite) URL so
|
||||
# relative Location headers keep the original hostname.
|
||||
url = _extract_redirect_url(response, url)
|
||||
raise SSRFError("Too many redirects")
|
||||
|
||||
|
||||
async def async_safe_get(client: Any, url: str, **kwargs: Any) -> Any:
|
||||
"""Async version of safe_get."""
|
||||
if not getattr(litellm, "user_url_validation", True):
|
||||
kwargs.setdefault("follow_redirects", True)
|
||||
return await client.get(url, **kwargs)
|
||||
kwargs.pop("follow_redirects", None)
|
||||
caller_headers = kwargs.pop("headers", {})
|
||||
for _ in range(_MAX_REDIRECTS):
|
||||
validated_url, original_host = validate_url(url)
|
||||
response = await client.get(
|
||||
validated_url,
|
||||
headers={**caller_headers, "Host": original_host},
|
||||
follow_redirects=False,
|
||||
**kwargs,
|
||||
)
|
||||
if not response.is_redirect:
|
||||
return response
|
||||
# Resolve the next hop against the ORIGINAL (pre-rewrite) URL so
|
||||
# relative Location headers keep the original hostname.
|
||||
url = _extract_redirect_url(response, url)
|
||||
raise SSRFError("Too many redirects")
|
||||
|
|
@ -1019,6 +1019,7 @@ class HTTPHandler:
|
|||
url,
|
||||
params=params,
|
||||
headers=headers,
|
||||
follow_redirects=_follow_redirects,
|
||||
)
|
||||
|
||||
return response
|
||||
|
|
|
|||
|
|
@ -15,6 +15,7 @@ from litellm.llms.custom_httpx.http_handler import (
|
|||
get_async_httpx_client,
|
||||
httpxSpecialProvider,
|
||||
)
|
||||
from litellm.litellm_core_utils.url_utils import async_safe_get
|
||||
from litellm.proxy._experimental.mcp_server.tool_registry import (
|
||||
global_mcp_tool_registry,
|
||||
)
|
||||
|
|
@ -75,9 +76,7 @@ def load_openapi_spec(filepath: str) -> Dict[str, Any]:
|
|||
async def load_openapi_spec_async(filepath: str) -> Dict[str, Any]:
|
||||
if filepath.startswith("http://") or filepath.startswith("https://"):
|
||||
client = get_async_httpx_client(llm_provider=httpxSpecialProvider.MCP)
|
||||
# NOTE: do not close shared client if get_async_httpx_client returns a shared singleton.
|
||||
# If it returns a new client each time, consider wrapping it in an async context manager.
|
||||
r = await client.get(filepath)
|
||||
r = await async_safe_get(client, filepath)
|
||||
r.raise_for_status()
|
||||
return r.json()
|
||||
|
||||
|
|
|
|||
1
litellm/proxy/_experimental/out/404/index.html
Normal file
1
litellm/proxy/_experimental/out/404/index.html
Normal file
File diff suppressed because one or more lines are too long
1
litellm/proxy/_experimental/out/_not-found/index.html
Normal file
1
litellm/proxy/_experimental/out/_not-found/index.html
Normal file
File diff suppressed because one or more lines are too long
1
litellm/proxy/_experimental/out/api-reference/index.html
Normal file
1
litellm/proxy/_experimental/out/api-reference/index.html
Normal file
File diff suppressed because one or more lines are too long
1
litellm/proxy/_experimental/out/chat/index.html
Normal file
1
litellm/proxy/_experimental/out/chat/index.html
Normal file
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
1
litellm/proxy/_experimental/out/guardrails/index.html
Normal file
1
litellm/proxy/_experimental/out/guardrails/index.html
Normal file
File diff suppressed because one or more lines are too long
1
litellm/proxy/_experimental/out/login/index.html
Normal file
1
litellm/proxy/_experimental/out/login/index.html
Normal file
File diff suppressed because one or more lines are too long
1
litellm/proxy/_experimental/out/logs/index.html
Normal file
1
litellm/proxy/_experimental/out/logs/index.html
Normal file
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
1
litellm/proxy/_experimental/out/model-hub/index.html
Normal file
1
litellm/proxy/_experimental/out/model-hub/index.html
Normal file
File diff suppressed because one or more lines are too long
1
litellm/proxy/_experimental/out/model_hub/index.html
Normal file
1
litellm/proxy/_experimental/out/model_hub/index.html
Normal file
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
1
litellm/proxy/_experimental/out/onboarding/index.html
Normal file
1
litellm/proxy/_experimental/out/onboarding/index.html
Normal file
File diff suppressed because one or more lines are too long
1
litellm/proxy/_experimental/out/organizations/index.html
Normal file
1
litellm/proxy/_experimental/out/organizations/index.html
Normal file
File diff suppressed because one or more lines are too long
1
litellm/proxy/_experimental/out/playground/index.html
Normal file
1
litellm/proxy/_experimental/out/playground/index.html
Normal file
File diff suppressed because one or more lines are too long
1
litellm/proxy/_experimental/out/policies/index.html
Normal file
1
litellm/proxy/_experimental/out/policies/index.html
Normal file
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
1
litellm/proxy/_experimental/out/teams/index.html
Normal file
1
litellm/proxy/_experimental/out/teams/index.html
Normal file
File diff suppressed because one or more lines are too long
1
litellm/proxy/_experimental/out/test-key/index.html
Normal file
1
litellm/proxy/_experimental/out/test-key/index.html
Normal file
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
File diff suppressed because one or more lines are too long
1
litellm/proxy/_experimental/out/usage/index.html
Normal file
1
litellm/proxy/_experimental/out/usage/index.html
Normal file
File diff suppressed because one or more lines are too long
1
litellm/proxy/_experimental/out/users/index.html
Normal file
1
litellm/proxy/_experimental/out/users/index.html
Normal file
File diff suppressed because one or more lines are too long
1
litellm/proxy/_experimental/out/virtual-keys/index.html
Normal file
1
litellm/proxy/_experimental/out/virtual-keys/index.html
Normal file
File diff suppressed because one or more lines are too long
|
|
@ -24,6 +24,7 @@ from litellm.llms.custom_httpx.http_handler import (
|
|||
get_async_httpx_client,
|
||||
httpxSpecialProvider,
|
||||
)
|
||||
from litellm.litellm_core_utils.url_utils import async_safe_get
|
||||
from litellm.rag.ingestion.file_parsers import extract_text_from_pdf
|
||||
from litellm.rag.text_splitters import RecursiveCharacterTextSplitter
|
||||
from litellm.types.rag import RAGIngestOptions, RAGIngestResponse
|
||||
|
|
@ -112,7 +113,7 @@ class BaseRAGIngestion(ABC):
|
|||
|
||||
if file_url:
|
||||
http_client = get_async_httpx_client(llm_provider=httpxSpecialProvider.RAG)
|
||||
response = await http_client.get(file_url)
|
||||
response = await async_safe_get(http_client, file_url)
|
||||
response.raise_for_status()
|
||||
file_content = response.content
|
||||
filename = file_url.split("/")[-1] or "document"
|
||||
|
|
|
|||
|
|
@ -55,6 +55,11 @@ def test_load_openapi_spec_supports_http_url(monkeypatch: pytest.MonkeyPatch) ->
|
|||
# Ensure shared/custom client path is used
|
||||
monkeypatch.setattr(gen, "get_async_httpx_client", fake_get_async_httpx_client)
|
||||
|
||||
# Bypass SSRF validation in test (example.local doesn't resolve)
|
||||
monkeypatch.setattr(
|
||||
gen, "async_safe_get", lambda client, url, **kw: client.get(url)
|
||||
)
|
||||
|
||||
# Fail loudly if someone reintroduces direct httpx.get()
|
||||
def boom(*args, **kwargs):
|
||||
raise AssertionError("Direct httpx.get() must not be used for URL spec loading")
|
||||
|
|
@ -68,7 +73,9 @@ def test_load_openapi_spec_supports_http_url(monkeypatch: pytest.MonkeyPatch) ->
|
|||
assert handler_holder["handler"].calls == 1
|
||||
|
||||
|
||||
def test_load_openapi_spec_supports_local_file_path(tmp_path, monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
def test_load_openapi_spec_supports_local_file_path(
|
||||
tmp_path, monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
expected: Dict[str, Any] = {
|
||||
"openapi": "3.0.0",
|
||||
"info": {"title": "Local API", "version": "1.0.0"},
|
||||
|
|
@ -83,10 +90,11 @@ def test_load_openapi_spec_supports_local_file_path(tmp_path, monkeypatch: pytes
|
|||
|
||||
# For local files, shared client must NOT be used.
|
||||
def boom_client(*args, **kwargs):
|
||||
raise AssertionError("get_async_httpx_client() must not be called for local file paths")
|
||||
raise AssertionError(
|
||||
"get_async_httpx_client() must not be called for local file paths"
|
||||
)
|
||||
|
||||
monkeypatch.setattr(gen, "get_async_httpx_client", boom_client)
|
||||
|
||||
spec = gen.load_openapi_spec(str(p))
|
||||
assert spec == expected
|
||||
|
||||
|
|
|
|||
|
|
@ -5,11 +5,22 @@ from httpx import Request, Response
|
|||
|
||||
import litellm
|
||||
from litellm import constants
|
||||
from litellm.litellm_core_utils.prompt_templates import image_handling
|
||||
from litellm.litellm_core_utils.prompt_templates.image_handling import (
|
||||
convert_url_to_base64,
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _bypass_ssrf(monkeypatch):
|
||||
"""Bypass SSRF validation in image handling tests — tests use fake URLs."""
|
||||
monkeypatch.setattr(
|
||||
image_handling,
|
||||
"safe_get",
|
||||
lambda client, url, **kw: client.get(url, follow_redirects=True),
|
||||
)
|
||||
|
||||
|
||||
class DummyClient:
|
||||
def get(self, url, follow_redirects=True):
|
||||
return Response(status_code=404, request=Request("GET", url))
|
||||
|
|
@ -37,9 +48,7 @@ def test_completion_with_invalid_image_url(monkeypatch):
|
|||
}
|
||||
]
|
||||
with pytest.raises(litellm.ImageFetchError) as excinfo:
|
||||
litellm.completion(
|
||||
model="gemini/gemini-pro", messages=messages, api_key="test"
|
||||
)
|
||||
litellm.completion(model="gemini/gemini-pro", messages=messages, api_key="test")
|
||||
assert excinfo.value.status_code == 400
|
||||
assert "Unable to fetch image" in str(excinfo.value)
|
||||
|
||||
|
|
@ -81,7 +90,7 @@ class StreamingLargeImageClient:
|
|||
headers = {"Content-Type": "image/jpeg"}
|
||||
if self.include_content_length:
|
||||
headers["Content-Length"] = str(size_bytes)
|
||||
|
||||
|
||||
# Create a generator that yields chunks without creating the whole file in memory
|
||||
def generate_chunks(total_size, chunk_size=8192):
|
||||
bytes_sent = 0
|
||||
|
|
@ -89,7 +98,7 @@ class StreamingLargeImageClient:
|
|||
chunk = b"x" * min(chunk_size, total_size - bytes_sent)
|
||||
bytes_sent += len(chunk)
|
||||
yield chunk
|
||||
|
||||
|
||||
# Create response with streaming content
|
||||
response = Response(
|
||||
status_code=200,
|
||||
|
|
@ -97,7 +106,9 @@ class StreamingLargeImageClient:
|
|||
request=Request("GET", url),
|
||||
)
|
||||
# Mock the iter_bytes method to return our generator
|
||||
response.iter_bytes = lambda chunk_size=8192: generate_chunks(size_bytes, chunk_size)
|
||||
response.iter_bytes = lambda chunk_size=8192: generate_chunks(
|
||||
size_bytes, chunk_size
|
||||
)
|
||||
return response
|
||||
|
||||
|
||||
|
|
@ -121,7 +132,9 @@ def test_image_exceeds_size_limit_without_content_length(monkeypatch):
|
|||
This uses the old non-streaming mock for backward compatibility.
|
||||
"""
|
||||
monkeypatch.setattr(
|
||||
litellm, "module_level_client", LargeImageClient(size_mb=100, include_content_length=False)
|
||||
litellm,
|
||||
"module_level_client",
|
||||
LargeImageClient(size_mb=100, include_content_length=False),
|
||||
)
|
||||
|
||||
with pytest.raises(litellm.ImageFetchError) as excinfo:
|
||||
|
|
@ -134,7 +147,7 @@ def test_streaming_download_protects_against_huge_files(monkeypatch):
|
|||
"""
|
||||
Test that streaming download aborts early when file exceeds size limit,
|
||||
preventing memory exhaustion from huge files (e.g., petabyte-sized files).
|
||||
|
||||
|
||||
This test verifies that the streaming implementation doesn't download the entire
|
||||
file into memory before checking size. Instead, it should abort as soon as the
|
||||
limit is exceeded during streaming.
|
||||
|
|
@ -148,7 +161,7 @@ def test_streaming_download_protects_against_huge_files(monkeypatch):
|
|||
|
||||
# Verify the error message shows it was caught during streaming
|
||||
assert "exceeds maximum allowed size" in str(excinfo.value)
|
||||
|
||||
|
||||
# The error should be raised after downloading just slightly more than the limit
|
||||
# not after downloading the full 1GB
|
||||
|
||||
|
|
@ -187,13 +200,15 @@ def test_streaming_download_handles_petabyte_file(monkeypatch):
|
|||
"""
|
||||
Test that streaming download can handle extremely large file URLs (e.g., petabyte-sized)
|
||||
without attempting to download the entire file or causing memory exhaustion.
|
||||
|
||||
|
||||
This simulates what happens if a malicious actor or misconfiguration provides
|
||||
a URL to an extremely large file.
|
||||
"""
|
||||
# Simulate a 1 petabyte file (1,000,000 GB)
|
||||
# Without streaming protection, this would cause OOM or hang indefinitely
|
||||
client = StreamingLargeImageClient(size_mb=1_000_000_000, include_content_length=False)
|
||||
client = StreamingLargeImageClient(
|
||||
size_mb=1_000_000_000, include_content_length=False
|
||||
)
|
||||
monkeypatch.setattr(litellm, "module_level_client", client)
|
||||
|
||||
with pytest.raises(litellm.ImageFetchError) as excinfo:
|
||||
|
|
@ -214,6 +229,6 @@ def test_image_size_limit_disabled(monkeypatch):
|
|||
|
||||
with pytest.raises(litellm.ImageFetchError) as excinfo:
|
||||
convert_url_to_base64("https://example.com/image.jpg")
|
||||
|
||||
|
||||
assert "Image URL download is disabled" in str(excinfo.value)
|
||||
assert "MAX_IMAGE_URL_DOWNLOAD_SIZE_MB=0" in str(excinfo.value)
|
||||
|
|
|
|||
396
tests/test_litellm/litellm_core_utils/test_url_utils.py
Normal file
396
tests/test_litellm/litellm_core_utils/test_url_utils.py
Normal file
|
|
@ -0,0 +1,396 @@
|
|||
import socket
|
||||
|
||||
import pytest
|
||||
|
||||
import litellm
|
||||
from litellm.litellm_core_utils import url_utils
|
||||
from litellm.litellm_core_utils.url_utils import SSRFError, _is_blocked_ip, validate_url
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_dns_public(monkeypatch):
|
||||
"""Resolve any hostname to 93.184.216.34 (public)."""
|
||||
|
||||
def fake_getaddrinfo(host, port, *args, **kwargs):
|
||||
return [
|
||||
(socket.AF_INET, socket.SOCK_STREAM, 6, "", ("93.184.216.34", port or 80))
|
||||
]
|
||||
|
||||
monkeypatch.setattr(url_utils.socket, "getaddrinfo", fake_getaddrinfo)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_dns_failure(monkeypatch):
|
||||
"""Make every DNS lookup raise gaierror."""
|
||||
|
||||
def fake_getaddrinfo(host, port, *args, **kwargs):
|
||||
raise socket.gaierror("Name or service not known")
|
||||
|
||||
monkeypatch.setattr(url_utils.socket, "getaddrinfo", fake_getaddrinfo)
|
||||
|
||||
|
||||
class TestIsBlockedIp:
|
||||
def test_blocks_private(self):
|
||||
assert _is_blocked_ip("10.0.0.1") is True
|
||||
|
||||
def test_allows_public(self):
|
||||
assert _is_blocked_ip("8.8.8.8") is False
|
||||
|
||||
def test_unparseable_is_blocked(self):
|
||||
assert _is_blocked_ip("not-an-ip") is True
|
||||
|
||||
# Coverage delta picked up by switching to `not ip.is_global` (RFC 6890)
|
||||
# over the old hand-maintained CIDR list.
|
||||
def test_blocks_cgnat_alibaba_metadata(self):
|
||||
"""100.100.100.200 is Alibaba Cloud metadata; lives in CGNAT."""
|
||||
assert _is_blocked_ip("100.100.100.200") is True
|
||||
|
||||
def test_blocks_ietf_protocol_assignments_old_oracle_metadata(self):
|
||||
"""192.0.0.192 was the legacy Oracle Cloud metadata IP."""
|
||||
assert _is_blocked_ip("192.0.0.192") is True
|
||||
|
||||
def test_blocks_documentation_ranges(self):
|
||||
assert _is_blocked_ip("192.0.2.1") is True
|
||||
assert _is_blocked_ip("198.51.100.1") is True
|
||||
assert _is_blocked_ip("203.0.113.1") is True
|
||||
|
||||
def test_blocks_multicast(self):
|
||||
assert _is_blocked_ip("224.0.0.1") is True
|
||||
|
||||
def test_blocks_reserved_future_use(self):
|
||||
assert _is_blocked_ip("240.0.0.1") is True
|
||||
|
||||
def test_blocks_broadcast(self):
|
||||
assert _is_blocked_ip("255.255.255.255") is True
|
||||
|
||||
def test_blocks_azure_wire_server(self):
|
||||
"""168.63.129.16 is globally routable but cloud-internal — explicit exception."""
|
||||
assert _is_blocked_ip("168.63.129.16") is True
|
||||
|
||||
def test_blocks_aws_ipv6_imds(self):
|
||||
"""fd00:ec2::254 is AWS's IPv6 IMDS, in IPv6 ULA (fc00::/7)."""
|
||||
assert _is_blocked_ip("fd00:ec2::254") is True
|
||||
|
||||
def test_blocks_ipv4_mapped_private(self):
|
||||
"""::ffff:10.0.0.1 must be unwrapped and blocked as 10.0.0.1."""
|
||||
assert _is_blocked_ip("::ffff:10.0.0.1") is True
|
||||
|
||||
def test_blocks_ipv4_mapped_azure_wire_server(self):
|
||||
"""::ffff:168.63.129.16 must be unwrapped and blocked via the exception list."""
|
||||
assert _is_blocked_ip("::ffff:168.63.129.16") is True
|
||||
|
||||
|
||||
class TestValidateUrl:
|
||||
def test_blocks_loopback(self):
|
||||
with pytest.raises(SSRFError):
|
||||
validate_url("http://127.0.0.1/test")
|
||||
|
||||
def test_blocks_imds(self):
|
||||
with pytest.raises(SSRFError):
|
||||
validate_url("http://169.254.169.254/latest/meta-data/")
|
||||
|
||||
def test_blocks_rfc1918_class_a(self):
|
||||
with pytest.raises(SSRFError):
|
||||
validate_url("http://10.0.1.5:8080/v1/completions")
|
||||
|
||||
def test_blocks_rfc1918_class_b(self):
|
||||
with pytest.raises(SSRFError):
|
||||
validate_url("http://172.16.0.1/")
|
||||
|
||||
def test_blocks_rfc1918_class_c(self):
|
||||
with pytest.raises(SSRFError):
|
||||
validate_url("http://192.168.1.1/")
|
||||
|
||||
def test_blocks_file_scheme(self):
|
||||
with pytest.raises(SSRFError):
|
||||
validate_url("file:///etc/passwd")
|
||||
|
||||
def test_blocks_ftp_scheme(self):
|
||||
with pytest.raises(SSRFError):
|
||||
validate_url("ftp://internal.host/data")
|
||||
|
||||
def test_blocks_no_hostname(self):
|
||||
with pytest.raises(SSRFError):
|
||||
validate_url("http:///path")
|
||||
|
||||
def test_allows_public_https(self, mock_dns_public):
|
||||
rewritten, host = validate_url("https://example.com/image.png")
|
||||
assert host == "example.com"
|
||||
assert rewritten == "https://example.com/image.png"
|
||||
|
||||
def test_rewrites_public_http_to_ip(self, mock_dns_public):
|
||||
rewritten, host = validate_url("http://example.com/image.png")
|
||||
assert host == "example.com"
|
||||
assert "example.com" not in rewritten
|
||||
|
||||
def test_preserves_path_and_query(self, mock_dns_public):
|
||||
rewritten, host = validate_url("http://example.com/path?key=value")
|
||||
assert "/path" in rewritten
|
||||
assert "key=value" in rewritten
|
||||
|
||||
def test_dns_failure_raises(self, mock_dns_failure):
|
||||
with pytest.raises(SSRFError, match="DNS resolution failed"):
|
||||
validate_url("http://this-domain-does-not-exist-xyz123.invalid/test")
|
||||
|
||||
def test_blocks_localhost_hostname(self, monkeypatch):
|
||||
def fake(host, port, *a, **kw):
|
||||
return [
|
||||
(socket.AF_INET, socket.SOCK_STREAM, 6, "", ("127.0.0.1", port or 80))
|
||||
]
|
||||
|
||||
monkeypatch.setattr(url_utils.socket, "getaddrinfo", fake)
|
||||
with pytest.raises(SSRFError):
|
||||
validate_url("http://localhost/")
|
||||
|
||||
def test_blocks_ipv6_loopback(self):
|
||||
with pytest.raises(SSRFError):
|
||||
validate_url("http://[::1]/")
|
||||
|
||||
def test_https_rewrites_when_ssl_verify_disabled(
|
||||
self, monkeypatch, mock_dns_public
|
||||
):
|
||||
monkeypatch.setattr(litellm, "ssl_verify", False)
|
||||
rewritten, host = validate_url("https://example.com/image.png")
|
||||
assert host == "example.com"
|
||||
assert "example.com" not in rewritten # rewritten to IP
|
||||
|
||||
def test_https_not_rewritten_when_ssl_verify_enabled(
|
||||
self, monkeypatch, mock_dns_public
|
||||
):
|
||||
monkeypatch.setattr(litellm, "ssl_verify", True)
|
||||
rewritten, host = validate_url("https://example.com/image.png")
|
||||
assert rewritten == "https://example.com/image.png"
|
||||
|
||||
|
||||
class TestHostHeaderFormatting:
|
||||
"""RFC 7230 §5.4: IPv6 literals must be bracketed in the Host header."""
|
||||
|
||||
def test_ipv4_no_port(self, monkeypatch):
|
||||
def fake(host, port, *a, **kw):
|
||||
return [(socket.AF_INET, socket.SOCK_STREAM, 6, "", ("1.2.3.4", port))]
|
||||
|
||||
monkeypatch.setattr(url_utils.socket, "getaddrinfo", fake)
|
||||
_, host = validate_url("http://example.com/")
|
||||
assert host == "example.com"
|
||||
|
||||
def test_ipv4_with_explicit_nondefault_port(self, monkeypatch):
|
||||
def fake(host, port, *a, **kw):
|
||||
return [(socket.AF_INET, socket.SOCK_STREAM, 6, "", ("1.2.3.4", port))]
|
||||
|
||||
monkeypatch.setattr(url_utils.socket, "getaddrinfo", fake)
|
||||
_, host = validate_url("http://example.com:8080/")
|
||||
assert host == "example.com:8080"
|
||||
|
||||
def test_ipv4_with_explicit_default_port_strips_port(self, monkeypatch):
|
||||
def fake(host, port, *a, **kw):
|
||||
return [(socket.AF_INET, socket.SOCK_STREAM, 6, "", ("1.2.3.4", port))]
|
||||
|
||||
monkeypatch.setattr(url_utils.socket, "getaddrinfo", fake)
|
||||
_, host = validate_url("http://example.com:80/")
|
||||
assert host == "example.com"
|
||||
|
||||
def test_ipv6_literal_is_bracketed_with_port(self, monkeypatch):
|
||||
"""Regression: IPv6 + port produced ambiguous `Host: 2001:db8::1:8080`."""
|
||||
monkeypatch.setattr(litellm, "user_url_allowed_hosts", ["[2001:db8::1]"])
|
||||
|
||||
def fake(host, port, *a, **kw):
|
||||
return [
|
||||
(
|
||||
socket.AF_INET6,
|
||||
socket.SOCK_STREAM,
|
||||
6,
|
||||
"",
|
||||
("2001:db8::1", port, 0, 0),
|
||||
)
|
||||
]
|
||||
|
||||
monkeypatch.setattr(url_utils.socket, "getaddrinfo", fake)
|
||||
_, host = validate_url("http://[2001:db8::1]:8080/")
|
||||
assert host == "[2001:db8::1]:8080"
|
||||
|
||||
def test_ipv6_literal_is_bracketed_without_port(self, monkeypatch):
|
||||
monkeypatch.setattr(litellm, "user_url_allowed_hosts", ["[2001:db8::1]"])
|
||||
|
||||
def fake(host, port, *a, **kw):
|
||||
return [
|
||||
(
|
||||
socket.AF_INET6,
|
||||
socket.SOCK_STREAM,
|
||||
6,
|
||||
"",
|
||||
("2001:db8::1", port, 0, 0),
|
||||
)
|
||||
]
|
||||
|
||||
monkeypatch.setattr(url_utils.socket, "getaddrinfo", fake)
|
||||
_, host = validate_url("http://[2001:db8::1]/")
|
||||
assert host == "[2001:db8::1]"
|
||||
|
||||
|
||||
class TestRedirectHostnamePreservation:
|
||||
"""Relative-location redirects must keep the original hostname, not the
|
||||
rewritten IP, so the next hop's Host header still identifies the site."""
|
||||
|
||||
def test_relative_redirect_preserves_hostname_for_next_hop(self, monkeypatch):
|
||||
def fake(host, port, *a, **kw):
|
||||
return [
|
||||
(socket.AF_INET, socket.SOCK_STREAM, 6, "", ("93.184.216.34", port))
|
||||
]
|
||||
|
||||
monkeypatch.setattr(url_utils.socket, "getaddrinfo", fake)
|
||||
|
||||
class FakeResponse:
|
||||
def __init__(self, status, location=None):
|
||||
self.status_code = status
|
||||
self.headers = {"location": location} if location else {}
|
||||
self.is_redirect = 300 <= status < 400
|
||||
|
||||
hops = []
|
||||
|
||||
class FakeClient:
|
||||
def __init__(self):
|
||||
self._n = 0
|
||||
|
||||
def get(self, url, headers=None, follow_redirects=False, **kw):
|
||||
hops.append({"url": url, "host": (headers or {}).get("Host")})
|
||||
self._n += 1
|
||||
if self._n == 1:
|
||||
return FakeResponse(302, "/redirected")
|
||||
return FakeResponse(200)
|
||||
|
||||
url_utils.safe_get(FakeClient(), "http://example.com/initial")
|
||||
assert len(hops) == 2
|
||||
# Both hops must carry the ORIGINAL hostname in the Host header.
|
||||
assert hops[0]["host"] == "example.com"
|
||||
assert hops[1]["host"] == "example.com"
|
||||
# Both outbound URLs go to the resolved IP (rewritten), not the hostname.
|
||||
assert "93.184.216.34" in hops[0]["url"]
|
||||
assert "93.184.216.34" in hops[1]["url"]
|
||||
# The second hop resolved /redirected relative to the original, not the IP.
|
||||
assert hops[1]["url"].endswith("/redirected")
|
||||
|
||||
|
||||
class TestValidationMasterSwitch:
|
||||
def test_disabled_bypasses_fetch_in_safe_get(self, monkeypatch):
|
||||
"""When user_url_validation is False, safe_get delegates to client.get without validation."""
|
||||
monkeypatch.setattr(litellm, "user_url_validation", False)
|
||||
|
||||
calls = []
|
||||
|
||||
class FakeClient:
|
||||
def get(self, url, **kwargs):
|
||||
calls.append((url, kwargs))
|
||||
|
||||
class R:
|
||||
is_redirect = False
|
||||
|
||||
return R()
|
||||
|
||||
url_utils.safe_get(FakeClient(), "http://127.0.0.1/internal")
|
||||
assert calls and calls[0][0] == "http://127.0.0.1/internal"
|
||||
assert calls[0][1].get("follow_redirects") is True
|
||||
|
||||
def test_enabled_still_blocks(self, monkeypatch):
|
||||
monkeypatch.setattr(litellm, "user_url_validation", True)
|
||||
with pytest.raises(SSRFError):
|
||||
validate_url("http://127.0.0.1/")
|
||||
|
||||
|
||||
class TestHostAllowlist:
|
||||
def test_allowlisted_hostname_permits_private_ip(self, monkeypatch):
|
||||
monkeypatch.setattr(litellm, "user_url_allowed_hosts", ["internal.corp"])
|
||||
|
||||
def fake(host, port, *a, **kw):
|
||||
return [(socket.AF_INET, socket.SOCK_STREAM, 6, "", ("10.0.1.5", port))]
|
||||
|
||||
monkeypatch.setattr(url_utils.socket, "getaddrinfo", fake)
|
||||
rewritten, host = validate_url("http://internal.corp/path")
|
||||
assert host == "internal.corp"
|
||||
assert "10.0.1.5" in rewritten
|
||||
|
||||
def test_non_allowlisted_hostname_still_blocked(self, monkeypatch):
|
||||
monkeypatch.setattr(litellm, "user_url_allowed_hosts", ["internal.corp"])
|
||||
|
||||
def fake(host, port, *a, **kw):
|
||||
return [(socket.AF_INET, socket.SOCK_STREAM, 6, "", ("10.0.1.5", port))]
|
||||
|
||||
monkeypatch.setattr(url_utils.socket, "getaddrinfo", fake)
|
||||
with pytest.raises(SSRFError):
|
||||
validate_url("http://other.corp/")
|
||||
|
||||
def test_allowlist_case_insensitive(self, monkeypatch):
|
||||
monkeypatch.setattr(litellm, "user_url_allowed_hosts", ["Internal.Corp"])
|
||||
|
||||
def fake(host, port, *a, **kw):
|
||||
return [(socket.AF_INET, socket.SOCK_STREAM, 6, "", ("10.0.1.5", port))]
|
||||
|
||||
monkeypatch.setattr(url_utils.socket, "getaddrinfo", fake)
|
||||
rewritten, _ = validate_url("http://internal.corp/")
|
||||
assert "10.0.1.5" in rewritten
|
||||
|
||||
def test_allowlist_with_port_matches_explicit_port(self, monkeypatch):
|
||||
monkeypatch.setattr(litellm, "user_url_allowed_hosts", ["internal.corp:8080"])
|
||||
|
||||
def fake(host, port, *a, **kw):
|
||||
return [(socket.AF_INET, socket.SOCK_STREAM, 6, "", ("10.0.1.5", port))]
|
||||
|
||||
monkeypatch.setattr(url_utils.socket, "getaddrinfo", fake)
|
||||
rewritten, host = validate_url("http://internal.corp:8080/")
|
||||
assert host == "internal.corp:8080"
|
||||
assert "10.0.1.5" in rewritten
|
||||
|
||||
def test_allowlist_with_port_matches_default_port(self, monkeypatch):
|
||||
"""Admin entry `host:443` matches `https://host/` (port=None, default 443)."""
|
||||
monkeypatch.setattr(litellm, "user_url_allowed_hosts", ["internal.corp:443"])
|
||||
|
||||
def fake(host, port, *a, **kw):
|
||||
return [(socket.AF_INET, socket.SOCK_STREAM, 6, "", ("10.0.1.5", port))]
|
||||
|
||||
monkeypatch.setattr(url_utils.socket, "getaddrinfo", fake)
|
||||
# Should succeed — no SSRFError raised
|
||||
validate_url("https://internal.corp/")
|
||||
|
||||
def test_allowlist_port_specific_does_not_match_other_port(self, monkeypatch):
|
||||
monkeypatch.setattr(litellm, "user_url_allowed_hosts", ["internal.corp:8080"])
|
||||
|
||||
def fake(host, port, *a, **kw):
|
||||
return [(socket.AF_INET, socket.SOCK_STREAM, 6, "", ("10.0.1.5", port))]
|
||||
|
||||
monkeypatch.setattr(url_utils.socket, "getaddrinfo", fake)
|
||||
with pytest.raises(SSRFError):
|
||||
validate_url("http://internal.corp:9090/")
|
||||
|
||||
def test_allowlist_host_entry_matches_any_port(self, monkeypatch):
|
||||
monkeypatch.setattr(litellm, "user_url_allowed_hosts", ["internal.corp"])
|
||||
|
||||
def fake(host, port, *a, **kw):
|
||||
return [(socket.AF_INET, socket.SOCK_STREAM, 6, "", ("10.0.1.5", port))]
|
||||
|
||||
monkeypatch.setattr(url_utils.socket, "getaddrinfo", fake)
|
||||
validate_url("http://internal.corp:9090/")
|
||||
validate_url("https://internal.corp:8443/")
|
||||
|
||||
def test_allowlist_permits_loopback(self, monkeypatch):
|
||||
"""Admin may opt into loopback if they explicitly configure it."""
|
||||
monkeypatch.setattr(litellm, "user_url_allowed_hosts", ["localhost"])
|
||||
|
||||
def fake(host, port, *a, **kw):
|
||||
return [(socket.AF_INET, socket.SOCK_STREAM, 6, "", ("127.0.0.1", port))]
|
||||
|
||||
monkeypatch.setattr(url_utils.socket, "getaddrinfo", fake)
|
||||
rewritten, host = validate_url("http://localhost:8080/")
|
||||
assert host == "localhost:8080"
|
||||
|
||||
def test_empty_allowlist_retains_default_deny(self, monkeypatch):
|
||||
monkeypatch.setattr(litellm, "user_url_allowed_hosts", [])
|
||||
with pytest.raises(SSRFError):
|
||||
validate_url("http://127.0.0.1/")
|
||||
|
||||
def test_allowlist_strips_trailing_dot(self, monkeypatch):
|
||||
monkeypatch.setattr(litellm, "user_url_allowed_hosts", ["internal.corp."])
|
||||
|
||||
def fake(host, port, *a, **kw):
|
||||
return [(socket.AF_INET, socket.SOCK_STREAM, 6, "", ("10.0.1.5", port))]
|
||||
|
||||
monkeypatch.setattr(url_utils.socket, "getaddrinfo", fake)
|
||||
validate_url("http://internal.corp/")
|
||||
Loading…
Add table
Reference in a new issue