mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
chore(providers): allowlist URL model destinations
This commit is contained in:
parent
87849b74b9
commit
7784b7f4ad
9 changed files with 278 additions and 28 deletions
|
|
@ -280,7 +280,7 @@ ssl_security_level: Optional[str] = None
|
|||
ssl_certificate: Optional[str] = None
|
||||
user_url_validation: bool = True
|
||||
user_url_allowed_hosts: List[str] = []
|
||||
reject_url_model_destinations: bool = True
|
||||
provider_url_destination_allowed_hosts: List[str] = []
|
||||
ssl_ecdh_curve: Optional[str] = (
|
||||
None # Set to 'X25519' to disable PQC and improve performance
|
||||
)
|
||||
|
|
|
|||
|
|
@ -21,7 +21,7 @@ Admins can opt out via two ``litellm`` globals (wired from proxy config):
|
|||
|
||||
import socket
|
||||
from ipaddress import ip_address, ip_network
|
||||
from typing import Any, List, Set, Tuple
|
||||
from typing import Any, List, Optional, Set, Tuple
|
||||
from urllib.parse import urlparse, urlunparse
|
||||
|
||||
import httpx
|
||||
|
|
@ -70,6 +70,85 @@ def _normalize_host(host: str) -> str:
|
|||
return host.lower().rstrip(".")
|
||||
|
||||
|
||||
def _default_port_for_scheme(scheme: str) -> int:
|
||||
return 443 if scheme == "https" else 80
|
||||
|
||||
|
||||
def _parse_url_destination_allowlist_entry(
|
||||
entry: str,
|
||||
) -> Optional[Tuple[str, Optional[str], Optional[int]]]:
|
||||
"""Parse an admin allowlist entry into host, optional scheme, optional port.
|
||||
|
||||
Entries may be bare hosts (``api.example.com``), host+port
|
||||
(``api.example.com:8443``), or origins (``https://api.example.com``).
|
||||
URL paths are intentionally ignored so admins can paste an api_base value.
|
||||
"""
|
||||
entry = entry.strip()
|
||||
if not entry:
|
||||
return None
|
||||
|
||||
has_scheme = "://" in entry
|
||||
parsed = urlparse(entry if has_scheme else f"//{entry}")
|
||||
if has_scheme and parsed.scheme not in _ALLOWED_SCHEMES:
|
||||
return None
|
||||
if parsed.username is not None or parsed.password is not None:
|
||||
return None
|
||||
if not parsed.hostname:
|
||||
return None
|
||||
|
||||
try:
|
||||
port = parsed.port
|
||||
except ValueError:
|
||||
return None
|
||||
|
||||
scheme: Optional[str] = parsed.scheme if has_scheme else None
|
||||
if scheme is not None and port is None:
|
||||
port = _default_port_for_scheme(scheme)
|
||||
|
||||
return _normalize_host(parsed.hostname), scheme, port
|
||||
|
||||
|
||||
def is_url_destination_allowed_by_host(url: str, allowed_hosts: List[str]) -> bool:
|
||||
"""Return True when a credential-bearing provider URL is admin-allowlisted.
|
||||
|
||||
This does not fetch, resolve, or rewrite URLs. It only answers whether the
|
||||
destination origin is explicitly trusted by configuration. Use ``safe_get``
|
||||
for user-controlled content fetches that require SSRF protection.
|
||||
"""
|
||||
parsed = urlparse(url)
|
||||
if parsed.scheme not in _ALLOWED_SCHEMES:
|
||||
return False
|
||||
if parsed.username is not None or parsed.password is not None:
|
||||
return False
|
||||
if not parsed.hostname:
|
||||
return False
|
||||
|
||||
try:
|
||||
effective_port = parsed.port or _default_port_for_scheme(parsed.scheme)
|
||||
except ValueError:
|
||||
return False
|
||||
|
||||
normalized_host = _normalize_host(parsed.hostname)
|
||||
configured_entries = (
|
||||
[allowed_hosts] if isinstance(allowed_hosts, str) else allowed_hosts
|
||||
)
|
||||
for entry in configured_entries or []:
|
||||
if not isinstance(entry, str):
|
||||
continue
|
||||
parsed_entry = _parse_url_destination_allowlist_entry(entry)
|
||||
if parsed_entry is None:
|
||||
continue
|
||||
allowed_host, allowed_scheme, allowed_port = parsed_entry
|
||||
if allowed_host != normalized_host:
|
||||
continue
|
||||
if allowed_scheme is not None and allowed_scheme != parsed.scheme:
|
||||
continue
|
||||
if allowed_port is not None and allowed_port != effective_port:
|
||||
continue
|
||||
return True
|
||||
return False
|
||||
|
||||
|
||||
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
|
||||
|
|
@ -145,7 +224,7 @@ def validate_url(url: str) -> Tuple[str, str]:
|
|||
raise SSRFError("URL has no hostname")
|
||||
|
||||
port = parsed.port
|
||||
default_port = 443 if parsed.scheme == "https" else 80
|
||||
default_port = _default_port_for_scheme(parsed.scheme)
|
||||
effective_port = port if port is not None else default_port
|
||||
host_header = _format_host_header(hostname, effective_port, default_port)
|
||||
|
||||
|
|
|
|||
|
|
@ -13,6 +13,7 @@ from openai.types.file_deleted import FileDeleted
|
|||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.litellm_core_utils.prompt_templates.common_utils import extract_file_data
|
||||
from litellm.litellm_core_utils.url_utils import is_url_destination_allowed_by_host
|
||||
from litellm.llms.base_llm.files.transformation import (
|
||||
BaseFilesConfig,
|
||||
LiteLLMLoggingObj,
|
||||
|
|
@ -29,8 +30,6 @@ from litellm.types.utils import LlmProviders
|
|||
|
||||
from ..common_utils import GeminiModelInfo
|
||||
|
||||
_GEMINI_FILES_HOST = "generativelanguage.googleapis.com"
|
||||
|
||||
|
||||
class GoogleAIStudioFilesHandler(GeminiModelInfo, BaseFilesConfig):
|
||||
def __init__(self):
|
||||
|
|
@ -226,20 +225,21 @@ class GoogleAIStudioFilesHandler(GeminiModelInfo, BaseFilesConfig):
|
|||
if not api_key:
|
||||
raise ValueError("api_key is required")
|
||||
|
||||
file_part = self._normalize_gemini_file_id(file_id)
|
||||
|
||||
api_base = (
|
||||
self.get_api_base(litellm_params.get("api_base"))
|
||||
or "https://generativelanguage.googleapis.com"
|
||||
)
|
||||
api_base = api_base.rstrip("/")
|
||||
file_part = self._normalize_gemini_file_id(file_id, api_base=api_base)
|
||||
|
||||
url = f"{api_base}/v1beta/{file_part}"
|
||||
|
||||
# API key is passed via x-goog-api-key header (set in validate_environment)
|
||||
return url, {}
|
||||
|
||||
def _normalize_gemini_file_id(self, file_id: str) -> str:
|
||||
def _normalize_gemini_file_id(
|
||||
self, file_id: str, api_base: Optional[str] = None
|
||||
) -> str:
|
||||
"""
|
||||
Normalize file identifier into `files/{id}` form.
|
||||
|
||||
|
|
@ -251,10 +251,11 @@ class GoogleAIStudioFilesHandler(GeminiModelInfo, BaseFilesConfig):
|
|||
if file_id.startswith(("http://", "https://")):
|
||||
parsed = urlparse(file_id)
|
||||
if (
|
||||
parsed.scheme != "https"
|
||||
or parsed.hostname != _GEMINI_FILES_HOST
|
||||
or parsed.username is not None
|
||||
parsed.username is not None
|
||||
or parsed.password is not None
|
||||
or not self._is_allowed_gemini_file_url(
|
||||
file_url=file_id, api_base=api_base
|
||||
)
|
||||
):
|
||||
raise ValueError("Invalid Gemini file URL")
|
||||
path = parsed.path.lstrip("/")
|
||||
|
|
@ -273,6 +274,20 @@ class GoogleAIStudioFilesHandler(GeminiModelInfo, BaseFilesConfig):
|
|||
|
||||
return normalized_file_id
|
||||
|
||||
@staticmethod
|
||||
def _is_allowed_gemini_file_url(
|
||||
file_url: str, api_base: Optional[str] = None
|
||||
) -> bool:
|
||||
import litellm
|
||||
|
||||
allowed_hosts: List[str] = []
|
||||
if api_base:
|
||||
allowed_hosts.append(api_base)
|
||||
allowed_hosts.extend(
|
||||
getattr(litellm, "provider_url_destination_allowed_hosts", []) or []
|
||||
)
|
||||
return is_url_destination_allowed_by_host(file_url, allowed_hosts)
|
||||
|
||||
@staticmethod
|
||||
def _validate_gemini_file_name(file_name: str) -> None:
|
||||
parts = file_name.split("/")
|
||||
|
|
@ -367,7 +382,8 @@ class GoogleAIStudioFilesHandler(GeminiModelInfo, BaseFilesConfig):
|
|||
if not api_key:
|
||||
raise ValueError("api_key is required")
|
||||
|
||||
file_name = self._normalize_gemini_file_id(file_id)
|
||||
api_base = api_base.rstrip("/")
|
||||
file_name = self._normalize_gemini_file_id(file_id, api_base=api_base)
|
||||
|
||||
# Construct the delete URL
|
||||
url = f"{api_base}/v1beta/{file_name}"
|
||||
|
|
|
|||
|
|
@ -4,6 +4,7 @@ from typing import Literal, Optional, Union
|
|||
|
||||
import httpx
|
||||
|
||||
from litellm.litellm_core_utils.url_utils import is_url_destination_allowed_by_host
|
||||
from litellm.llms.base_llm.chat.transformation import BaseLLMException
|
||||
|
||||
HF_HUB_URL = "https://huggingface.co"
|
||||
|
|
@ -13,25 +14,27 @@ def is_url_model_destination(model: str) -> bool:
|
|||
return model.startswith(("http://", "https://"))
|
||||
|
||||
|
||||
def _should_reject_url_model_destinations() -> bool:
|
||||
def _is_url_model_destination_allowed(model: str) -> bool:
|
||||
import litellm
|
||||
|
||||
return getattr(litellm, "reject_url_model_destinations", True) is True
|
||||
allowed_hosts = getattr(litellm, "provider_url_destination_allowed_hosts", []) or []
|
||||
return is_url_destination_allowed_by_host(model, allowed_hosts)
|
||||
|
||||
|
||||
def validate_huggingface_model_identifier(model: str) -> None:
|
||||
"""Reject URL-valued model identifiers before provider credentials are added."""
|
||||
if "://" not in model:
|
||||
return
|
||||
if is_url_model_destination(model) and not _should_reject_url_model_destinations():
|
||||
if is_url_model_destination(model) and _is_url_model_destination_allowed(model):
|
||||
return
|
||||
raise HuggingFaceError(
|
||||
status_code=400,
|
||||
message=(
|
||||
"Invalid Hugging Face model identifier. Configure custom endpoints with "
|
||||
"api_base or HF_API_BASE/HUGGINGFACE_API_BASE instead of passing a URL "
|
||||
"as the model. To keep legacy URL-valued models for trusted inputs, set "
|
||||
"litellm.reject_url_model_destinations=False."
|
||||
"as the model. To keep legacy URL-valued models for trusted endpoints, "
|
||||
"add the destination host or origin to "
|
||||
"`provider_url_destination_allowed_hosts` in litellm_settings."
|
||||
),
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -2,6 +2,7 @@ from typing import Optional, Union
|
|||
|
||||
import httpx
|
||||
|
||||
from litellm.litellm_core_utils.url_utils import is_url_destination_allowed_by_host
|
||||
from litellm.llms.base_llm.chat.transformation import BaseLLMException
|
||||
|
||||
|
||||
|
|
@ -19,24 +20,26 @@ def is_url_model_destination(model: str) -> bool:
|
|||
return model.startswith(("http://", "https://"))
|
||||
|
||||
|
||||
def _should_reject_url_model_destinations() -> bool:
|
||||
def _is_url_model_destination_allowed(model: str) -> bool:
|
||||
import litellm
|
||||
|
||||
return getattr(litellm, "reject_url_model_destinations", True) is True
|
||||
allowed_hosts = getattr(litellm, "provider_url_destination_allowed_hosts", []) or []
|
||||
return is_url_destination_allowed_by_host(model, allowed_hosts)
|
||||
|
||||
|
||||
def validate_oobabooga_model_identifier(model: str) -> None:
|
||||
"""Oobabooga endpoints must be configured with api_base, not model URLs."""
|
||||
if "://" not in model:
|
||||
return
|
||||
if is_url_model_destination(model) and not _should_reject_url_model_destinations():
|
||||
if is_url_model_destination(model) and _is_url_model_destination_allowed(model):
|
||||
return
|
||||
raise OobaboogaError(
|
||||
status_code=400,
|
||||
message=(
|
||||
"Invalid Oobabooga model identifier. Configure the endpoint with "
|
||||
"api_base instead of passing a URL as the model. To keep legacy "
|
||||
"URL-valued models for trusted inputs, set "
|
||||
"litellm.reject_url_model_destinations=False."
|
||||
"URL-valued models for trusted endpoints, add the destination host "
|
||||
"or origin to `provider_url_destination_allowed_hosts` in "
|
||||
"litellm_settings."
|
||||
),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -394,3 +394,45 @@ class TestHostAllowlist:
|
|||
|
||||
monkeypatch.setattr(url_utils.socket, "getaddrinfo", fake)
|
||||
validate_url("http://internal.corp/")
|
||||
|
||||
|
||||
class TestProviderUrlDestinationAllowlist:
|
||||
def test_host_entry_matches_any_scheme_and_port(self):
|
||||
assert url_utils.is_url_destination_allowed_by_host(
|
||||
"https://trusted.example/v1/chat/completions",
|
||||
["trusted.example"],
|
||||
)
|
||||
assert url_utils.is_url_destination_allowed_by_host(
|
||||
"http://trusted.example:8080/v1/chat/completions",
|
||||
["trusted.example"],
|
||||
)
|
||||
|
||||
def test_origin_entry_matches_scheme_and_default_port(self):
|
||||
assert url_utils.is_url_destination_allowed_by_host(
|
||||
"https://trusted.example/v1/chat/completions",
|
||||
["https://trusted.example"],
|
||||
)
|
||||
assert not url_utils.is_url_destination_allowed_by_host(
|
||||
"http://trusted.example/v1/chat/completions",
|
||||
["https://trusted.example"],
|
||||
)
|
||||
|
||||
def test_port_entry_only_matches_same_effective_port(self):
|
||||
assert url_utils.is_url_destination_allowed_by_host(
|
||||
"https://trusted.example/v1/chat/completions",
|
||||
["trusted.example:443"],
|
||||
)
|
||||
assert not url_utils.is_url_destination_allowed_by_host(
|
||||
"https://trusted.example:8443/v1/chat/completions",
|
||||
["trusted.example:443"],
|
||||
)
|
||||
|
||||
def test_rejects_userinfo_and_invalid_port(self):
|
||||
assert not url_utils.is_url_destination_allowed_by_host(
|
||||
"https://user:pass@trusted.example/v1/chat/completions",
|
||||
["trusted.example"],
|
||||
)
|
||||
assert not url_utils.is_url_destination_allowed_by_host(
|
||||
"https://trusted.example:99999/v1/chat/completions",
|
||||
["trusted.example"],
|
||||
)
|
||||
|
|
|
|||
|
|
@ -7,6 +7,7 @@ from unittest.mock import Mock, patch
|
|||
import httpx
|
||||
import pytest
|
||||
|
||||
import litellm
|
||||
from litellm.llms.gemini.files.transformation import GoogleAIStudioFilesHandler
|
||||
from litellm.types.llms.openai import OpenAIFileObject
|
||||
|
||||
|
|
@ -103,6 +104,43 @@ class TestGoogleAIStudioFilesTransformation:
|
|||
litellm_params=litellm_params,
|
||||
)
|
||||
|
||||
def test_transform_retrieve_file_request_allows_full_url_matching_api_base(self):
|
||||
file_id = "https://custom-gemini.example/v1beta/files/test123"
|
||||
litellm_params = {
|
||||
"api_key": "test-api-key",
|
||||
"api_base": "https://custom-gemini.example",
|
||||
}
|
||||
|
||||
url, params = self.handler.transform_retrieve_file_request(
|
||||
file_id=file_id,
|
||||
optional_params={},
|
||||
litellm_params=litellm_params,
|
||||
)
|
||||
|
||||
assert url == "https://custom-gemini.example/v1beta/files/test123"
|
||||
assert params == {}
|
||||
|
||||
def test_transform_retrieve_file_request_allows_full_url_when_host_allowlisted(
|
||||
self,
|
||||
monkeypatch,
|
||||
):
|
||||
monkeypatch.setattr(
|
||||
litellm,
|
||||
"provider_url_destination_allowed_hosts",
|
||||
["trusted-gemini.example"],
|
||||
)
|
||||
file_id = "https://trusted-gemini.example/v1beta/files/test123"
|
||||
litellm_params = {"api_key": "test-api-key"}
|
||||
|
||||
url, params = self.handler.transform_retrieve_file_request(
|
||||
file_id=file_id,
|
||||
optional_params={},
|
||||
litellm_params=litellm_params,
|
||||
)
|
||||
|
||||
assert url == "https://generativelanguage.googleapis.com/v1beta/files/test123"
|
||||
assert params == {}
|
||||
|
||||
def test_transform_retrieve_file_request_rejects_traversal_name(self):
|
||||
litellm_params = {"api_key": "test-api-key"}
|
||||
|
||||
|
|
@ -386,6 +424,21 @@ class TestGoogleAIStudioFilesTransformation:
|
|||
litellm_params=litellm_params,
|
||||
)
|
||||
|
||||
def test_transform_delete_file_request_allows_full_url_matching_api_base(self):
|
||||
litellm_params = {
|
||||
"api_key": "test-api-key",
|
||||
"api_base": "https://custom-gemini.example",
|
||||
}
|
||||
|
||||
url, params = self.handler.transform_delete_file_request(
|
||||
file_id="https://custom-gemini.example/v1beta/files/test123",
|
||||
optional_params={},
|
||||
litellm_params=litellm_params,
|
||||
)
|
||||
|
||||
assert url == "https://custom-gemini.example/v1beta/files/test123"
|
||||
assert params == {}
|
||||
|
||||
def test_transform_delete_file_request_rejects_encoded_traversal_url(self):
|
||||
litellm_params = {
|
||||
"api_key": "test-api-key",
|
||||
|
|
|
|||
|
|
@ -26,10 +26,12 @@ def test_huggingface_chat_rejects_url_valued_model():
|
|||
assert exc_info.value.status_code == 400
|
||||
|
||||
|
||||
def test_huggingface_chat_allows_legacy_url_model_when_rejection_disabled(
|
||||
def test_huggingface_chat_allows_legacy_url_model_when_host_is_allowlisted(
|
||||
monkeypatch,
|
||||
):
|
||||
monkeypatch.setattr(litellm, "reject_url_model_destinations", False)
|
||||
monkeypatch.setattr(
|
||||
litellm, "provider_url_destination_allowed_hosts", ["trusted.example"]
|
||||
)
|
||||
config = HuggingFaceChatConfig()
|
||||
|
||||
complete_url = config.get_complete_url(
|
||||
|
|
@ -43,6 +45,26 @@ def test_huggingface_chat_allows_legacy_url_model_when_rejection_disabled(
|
|||
assert complete_url == "https://trusted.example/v1/chat/completions"
|
||||
|
||||
|
||||
def test_huggingface_chat_rejects_url_model_when_host_is_not_allowlisted(
|
||||
monkeypatch,
|
||||
):
|
||||
monkeypatch.setattr(
|
||||
litellm, "provider_url_destination_allowed_hosts", ["trusted.example"]
|
||||
)
|
||||
config = HuggingFaceChatConfig()
|
||||
|
||||
with pytest.raises(HuggingFaceError) as exc_info:
|
||||
config.get_complete_url(
|
||||
api_base=None,
|
||||
api_key="hf-secret",
|
||||
model="https://other.example",
|
||||
optional_params={},
|
||||
litellm_params={},
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 400
|
||||
|
||||
|
||||
def test_huggingface_chat_keeps_explicit_api_base_for_custom_endpoints():
|
||||
config = HuggingFaceChatConfig()
|
||||
|
||||
|
|
@ -69,10 +91,14 @@ def test_huggingface_embedding_config_rejects_url_valued_model():
|
|||
assert exc_info.value.status_code == 400
|
||||
|
||||
|
||||
def test_huggingface_embedding_config_allows_legacy_url_model_when_rejection_disabled(
|
||||
def test_huggingface_embedding_config_allows_legacy_url_model_when_origin_is_allowlisted(
|
||||
monkeypatch,
|
||||
):
|
||||
monkeypatch.setattr(litellm, "reject_url_model_destinations", False)
|
||||
monkeypatch.setattr(
|
||||
litellm,
|
||||
"provider_url_destination_allowed_hosts",
|
||||
["https://trusted.example"],
|
||||
)
|
||||
config = HuggingFaceEmbeddingConfig()
|
||||
|
||||
api_base = config.get_api_base(
|
||||
|
|
|
|||
|
|
@ -27,10 +27,12 @@ def test_oobabooga_completion_rejects_url_valued_model_before_request():
|
|||
mock_get.assert_not_called()
|
||||
|
||||
|
||||
def test_oobabooga_completion_allows_legacy_url_model_when_rejection_disabled(
|
||||
def test_oobabooga_completion_allows_legacy_url_model_when_host_is_allowlisted(
|
||||
monkeypatch,
|
||||
):
|
||||
monkeypatch.setattr(litellm, "reject_url_model_destinations", False)
|
||||
monkeypatch.setattr(
|
||||
litellm, "provider_url_destination_allowed_hosts", ["trusted.example"]
|
||||
)
|
||||
response = MagicMock()
|
||||
client = MagicMock()
|
||||
client.post.return_value = response
|
||||
|
|
@ -63,6 +65,32 @@ def test_oobabooga_completion_allows_legacy_url_model_when_rejection_disabled(
|
|||
)
|
||||
|
||||
|
||||
def test_oobabooga_completion_rejects_url_model_when_host_is_not_allowlisted(
|
||||
monkeypatch,
|
||||
):
|
||||
monkeypatch.setattr(
|
||||
litellm, "provider_url_destination_allowed_hosts", ["trusted.example"]
|
||||
)
|
||||
|
||||
with patch("litellm.llms.oobabooga.chat.oobabooga._get_httpx_client") as mock_get:
|
||||
with pytest.raises(OobaboogaError) as exc_info:
|
||||
completion(
|
||||
model="https://other.example",
|
||||
messages=[],
|
||||
api_base=None,
|
||||
model_response=MagicMock(),
|
||||
print_verbose=MagicMock(),
|
||||
encoding=MagicMock(),
|
||||
api_key="ooba-secret",
|
||||
logging_obj=MagicMock(),
|
||||
optional_params={},
|
||||
litellm_params={},
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 400
|
||||
mock_get.assert_not_called()
|
||||
|
||||
|
||||
def test_oobabooga_embedding_rejects_url_valued_model_before_request():
|
||||
with patch(
|
||||
"litellm.llms.oobabooga.chat.oobabooga.litellm.module_level_client.post"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue