diff --git a/litellm/__init__.py b/litellm/__init__.py index 7734b60d9cb..fd4e90c7447 100644 --- a/litellm/__init__.py +++ b/litellm/__init__.py @@ -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 ) diff --git a/litellm/integrations/custom_guardrail.py b/litellm/integrations/custom_guardrail.py index 853e5c7151a..abf010e0d65 100644 --- a/litellm/integrations/custom_guardrail.py +++ b/litellm/integrations/custom_guardrail.py @@ -255,26 +255,44 @@ class CustomGuardrail(CustomLogger): f"Event hook {event_hook} is not in the supported event hooks {supported_event_hooks}" ) + @staticmethod + def _get_admin_metadata(data: dict) -> dict: + """Return merged admin-configured key and team metadata from the request data. + + The proxy may inject admin metadata (user_api_key_metadata, + user_api_key_team_metadata) into either ``metadata`` or + ``litellm_metadata`` depending on endpoint. Check both so a caller + cannot shadow admin config by pre-populating the other key. + Key-level settings override team-level. + """ + team_meta: dict = {} + key_meta: dict = {} + for key in ("metadata", "litellm_metadata"): + # Defensive: an unparsed JSON-string metadata could leak past the + # proxy's normal parse path; don't AttributeError on .get(). + meta = data.get(key) + if not isinstance(meta, dict): + continue + team_meta = meta.get("user_api_key_team_metadata") or team_meta + key_meta = meta.get("user_api_key_metadata") or key_meta + return {**team_meta, **key_meta} + def get_disable_global_guardrail(self, data: dict) -> Optional[bool]: """ - Returns True if the global guardrail should be disabled + Returns True if the global guardrail should be disabled. + + Reads from admin-configured key/team metadata only, not from + the request body, to prevent callers from disabling guardrails. """ - if "disable_global_guardrails" in data: - return data["disable_global_guardrails"] - metadata = data.get("litellm_metadata") or data.get("metadata", {}) - if "disable_global_guardrails" in metadata: - return metadata["disable_global_guardrails"] - return False + return self._get_admin_metadata(data).get("disable_global_guardrails", False) def get_opted_out_global_guardrails_from_metadata(self, data: dict) -> List[str]: """ Returns the list of global guardrail names the team/key has opted out of. + + Reads from admin-configured key/team metadata only. """ - if "opted_out_global_guardrails" in data: - value = data["opted_out_global_guardrails"] - return value if isinstance(value, list) else [] - metadata = data.get("litellm_metadata") or data.get("metadata", {}) - value = metadata.get("opted_out_global_guardrails") + value = self._get_admin_metadata(data).get("opted_out_global_guardrails") return value if isinstance(value, list) else [] def _is_valid_response_type(self, result: Any) -> bool: diff --git a/litellm/litellm_core_utils/initialize_dynamic_callback_params.py b/litellm/litellm_core_utils/initialize_dynamic_callback_params.py index 92c97a59924..563609af1e4 100644 --- a/litellm/litellm_core_utils/initialize_dynamic_callback_params.py +++ b/litellm/litellm_core_utils/initialize_dynamic_callback_params.py @@ -48,7 +48,6 @@ _supported_callback_params = [ "braintrust_host", "slack_webhook_url", "lunary_public_key", - "turn_off_message_logging", ] diff --git a/litellm/litellm_core_utils/prompt_templates/image_handling.py b/litellm/litellm_core_utils/prompt_templates/image_handling.py index eaf78b7bcf5..fd38bc9388d 100644 --- a/litellm/litellm_core_utils/prompt_templates/image_handling.py +++ b/litellm/litellm_core_utils/prompt_templates/image_handling.py @@ -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 diff --git a/litellm/litellm_core_utils/token_counter.py b/litellm/litellm_core_utils/token_counter.py index 09c62f2eb55..01e5dc39a34 100644 --- a/litellm/litellm_core_utils/token_counter.py +++ b/litellm/litellm_core_utils/token_counter.py @@ -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) diff --git a/litellm/litellm_core_utils/url_utils.py b/litellm/litellm_core_utils/url_utils.py new file mode 100644 index 00000000000..a65d0892aa2 --- /dev/null +++ b/litellm/litellm_core_utils/url_utils.py @@ -0,0 +1,274 @@ +""" +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 _sockaddr_host(sockaddr: Any) -> str: + """Return the host element of a ``getaddrinfo`` sockaddr as ``str``. + + ``getaddrinfo`` with ``IPPROTO_TCP`` returns AF_INET / AF_INET6 sockaddrs + whose first element is always a host string. mypy types it as + ``str | int`` (since sockaddrs for other families can hold ints), so we + narrow at the boundary. Fail closed if the stdlib ever returns something + unexpected — a non-string here would mean we have no IP to check against + the SSRF blocklist. + """ + host = sockaddr[0] + if not isinstance(host, str): + raise SSRFError(f"getaddrinfo returned non-string host: {host!r}") + return host + + +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: + resolved_ip = _sockaddr_host(sockaddr) + if _is_blocked_ip(resolved_ip): + raise SSRFError( + f"URL targets a blocked address ({resolved_ip}). " + "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 = _sockaddr_host(addrinfo[0][4]) + 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") diff --git a/litellm/llms/custom_httpx/http_handler.py b/litellm/llms/custom_httpx/http_handler.py index 489a56daf8e..03d2af72329 100644 --- a/litellm/llms/custom_httpx/http_handler.py +++ b/litellm/llms/custom_httpx/http_handler.py @@ -1019,6 +1019,7 @@ class HTTPHandler: url, params=params, headers=headers, + follow_redirects=_follow_redirects, ) return response diff --git a/litellm/proxy/_experimental/mcp_server/openapi_to_mcp_generator.py b/litellm/proxy/_experimental/mcp_server/openapi_to_mcp_generator.py index 4b4818892bb..3b2fa097b70 100644 --- a/litellm/proxy/_experimental/mcp_server/openapi_to_mcp_generator.py +++ b/litellm/proxy/_experimental/mcp_server/openapi_to_mcp_generator.py @@ -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() diff --git a/litellm/proxy/_experimental/out/404/index.html b/litellm/proxy/_experimental/out/404/index.html new file mode 100644 index 00000000000..344481d3aed --- /dev/null +++ b/litellm/proxy/_experimental/out/404/index.html @@ -0,0 +1 @@ +404: This page could not be found.LiteLLM Dashboard

404

This page could not be found.

\ No newline at end of file diff --git a/litellm/proxy/_experimental/out/_not-found/index.html b/litellm/proxy/_experimental/out/_not-found/index.html new file mode 100644 index 00000000000..344481d3aed --- /dev/null +++ b/litellm/proxy/_experimental/out/_not-found/index.html @@ -0,0 +1 @@ +404: This page could not be found.LiteLLM Dashboard

404

This page could not be found.

\ No newline at end of file diff --git a/litellm/proxy/_experimental/out/api-reference/index.html b/litellm/proxy/_experimental/out/api-reference/index.html new file mode 100644 index 00000000000..b636faba290 --- /dev/null +++ b/litellm/proxy/_experimental/out/api-reference/index.html @@ -0,0 +1 @@ +LiteLLM Dashboard
Loading...
\ No newline at end of file diff --git a/litellm/proxy/_experimental/out/chat/index.html b/litellm/proxy/_experimental/out/chat/index.html new file mode 100644 index 00000000000..0d684c66cb5 --- /dev/null +++ b/litellm/proxy/_experimental/out/chat/index.html @@ -0,0 +1 @@ +LiteLLM Dashboard \ No newline at end of file diff --git a/litellm/proxy/_experimental/out/experimental/api-playground/index.html b/litellm/proxy/_experimental/out/experimental/api-playground/index.html new file mode 100644 index 00000000000..5268cc3d9ca --- /dev/null +++ b/litellm/proxy/_experimental/out/experimental/api-playground/index.html @@ -0,0 +1 @@ +LiteLLM Dashboard
Loading...
\ No newline at end of file diff --git a/litellm/proxy/_experimental/out/experimental/budgets/index.html b/litellm/proxy/_experimental/out/experimental/budgets/index.html new file mode 100644 index 00000000000..f463b2d5df3 --- /dev/null +++ b/litellm/proxy/_experimental/out/experimental/budgets/index.html @@ -0,0 +1 @@ +LiteLLM Dashboard
Loading...
\ No newline at end of file diff --git a/litellm/proxy/_experimental/out/experimental/caching/index.html b/litellm/proxy/_experimental/out/experimental/caching/index.html new file mode 100644 index 00000000000..cf2a1aa14a0 --- /dev/null +++ b/litellm/proxy/_experimental/out/experimental/caching/index.html @@ -0,0 +1 @@ +LiteLLM Dashboard
Loading...
\ No newline at end of file diff --git a/litellm/proxy/_experimental/out/experimental/claude-code-plugins/index.html b/litellm/proxy/_experimental/out/experimental/claude-code-plugins/index.html new file mode 100644 index 00000000000..069f97b082a --- /dev/null +++ b/litellm/proxy/_experimental/out/experimental/claude-code-plugins/index.html @@ -0,0 +1 @@ +LiteLLM Dashboard
Loading...
\ No newline at end of file diff --git a/litellm/proxy/_experimental/out/experimental/old-usage/index.html b/litellm/proxy/_experimental/out/experimental/old-usage/index.html new file mode 100644 index 00000000000..53540d126c4 --- /dev/null +++ b/litellm/proxy/_experimental/out/experimental/old-usage/index.html @@ -0,0 +1 @@ +LiteLLM Dashboard
Loading...
\ No newline at end of file diff --git a/litellm/proxy/_experimental/out/experimental/prompts/index.html b/litellm/proxy/_experimental/out/experimental/prompts/index.html new file mode 100644 index 00000000000..615f06b8166 --- /dev/null +++ b/litellm/proxy/_experimental/out/experimental/prompts/index.html @@ -0,0 +1 @@ +LiteLLM Dashboard
Loading...
\ No newline at end of file diff --git a/litellm/proxy/_experimental/out/experimental/tag-management/index.html b/litellm/proxy/_experimental/out/experimental/tag-management/index.html new file mode 100644 index 00000000000..e7d0631c339 --- /dev/null +++ b/litellm/proxy/_experimental/out/experimental/tag-management/index.html @@ -0,0 +1 @@ +LiteLLM Dashboard
Loading...
\ No newline at end of file diff --git a/litellm/proxy/_experimental/out/guardrails/index.html b/litellm/proxy/_experimental/out/guardrails/index.html new file mode 100644 index 00000000000..ebbe174662b --- /dev/null +++ b/litellm/proxy/_experimental/out/guardrails/index.html @@ -0,0 +1 @@ +LiteLLM Dashboard
Loading...
\ No newline at end of file diff --git a/litellm/proxy/_experimental/out/login/index.html b/litellm/proxy/_experimental/out/login/index.html new file mode 100644 index 00000000000..54472c6cc11 --- /dev/null +++ b/litellm/proxy/_experimental/out/login/index.html @@ -0,0 +1 @@ +LiteLLM Dashboard
🚅 LiteLLM
Loading...
\ No newline at end of file diff --git a/litellm/proxy/_experimental/out/logs/index.html b/litellm/proxy/_experimental/out/logs/index.html new file mode 100644 index 00000000000..ec43b677a2f --- /dev/null +++ b/litellm/proxy/_experimental/out/logs/index.html @@ -0,0 +1 @@ +LiteLLM Dashboard
Loading...
\ No newline at end of file diff --git a/litellm/proxy/_experimental/out/mcp/oauth/callback/index.html b/litellm/proxy/_experimental/out/mcp/oauth/callback/index.html new file mode 100644 index 00000000000..830060c7aa2 --- /dev/null +++ b/litellm/proxy/_experimental/out/mcp/oauth/callback/index.html @@ -0,0 +1 @@ +LiteLLM Dashboard
Loading...
\ No newline at end of file diff --git a/litellm/proxy/_experimental/out/model-hub/index.html b/litellm/proxy/_experimental/out/model-hub/index.html new file mode 100644 index 00000000000..506c3695285 --- /dev/null +++ b/litellm/proxy/_experimental/out/model-hub/index.html @@ -0,0 +1 @@ +LiteLLM Dashboard
Loading...
\ No newline at end of file diff --git a/litellm/proxy/_experimental/out/model_hub/index.html b/litellm/proxy/_experimental/out/model_hub/index.html new file mode 100644 index 00000000000..27bac5cde7d --- /dev/null +++ b/litellm/proxy/_experimental/out/model_hub/index.html @@ -0,0 +1 @@ +LiteLLM Dashboard
Loading...
\ No newline at end of file diff --git a/litellm/proxy/_experimental/out/model_hub_table/index.html b/litellm/proxy/_experimental/out/model_hub_table/index.html new file mode 100644 index 00000000000..db5d0e6a718 --- /dev/null +++ b/litellm/proxy/_experimental/out/model_hub_table/index.html @@ -0,0 +1 @@ +LiteLLM Dashboard
Loading...
\ No newline at end of file diff --git a/litellm/proxy/_experimental/out/models-and-endpoints/index.html b/litellm/proxy/_experimental/out/models-and-endpoints/index.html new file mode 100644 index 00000000000..96c1a43a7c0 --- /dev/null +++ b/litellm/proxy/_experimental/out/models-and-endpoints/index.html @@ -0,0 +1 @@ +LiteLLM Dashboard
Loading...
\ No newline at end of file diff --git a/litellm/proxy/_experimental/out/onboarding/index.html b/litellm/proxy/_experimental/out/onboarding/index.html new file mode 100644 index 00000000000..5c2121443f9 --- /dev/null +++ b/litellm/proxy/_experimental/out/onboarding/index.html @@ -0,0 +1 @@ +LiteLLM Dashboard
Loading...
\ No newline at end of file diff --git a/litellm/proxy/_experimental/out/organizations/index.html b/litellm/proxy/_experimental/out/organizations/index.html new file mode 100644 index 00000000000..51dd7d1c764 --- /dev/null +++ b/litellm/proxy/_experimental/out/organizations/index.html @@ -0,0 +1 @@ +LiteLLM Dashboard
Loading...
\ No newline at end of file diff --git a/litellm/proxy/_experimental/out/playground/index.html b/litellm/proxy/_experimental/out/playground/index.html new file mode 100644 index 00000000000..41ef863e95b --- /dev/null +++ b/litellm/proxy/_experimental/out/playground/index.html @@ -0,0 +1 @@ +LiteLLM Dashboard
Loading...
\ No newline at end of file diff --git a/litellm/proxy/_experimental/out/policies/index.html b/litellm/proxy/_experimental/out/policies/index.html new file mode 100644 index 00000000000..a452ae4c4aa --- /dev/null +++ b/litellm/proxy/_experimental/out/policies/index.html @@ -0,0 +1 @@ +LiteLLM Dashboard
Loading...
\ No newline at end of file diff --git a/litellm/proxy/_experimental/out/settings/admin-settings/index.html b/litellm/proxy/_experimental/out/settings/admin-settings/index.html new file mode 100644 index 00000000000..b29b4856b0d --- /dev/null +++ b/litellm/proxy/_experimental/out/settings/admin-settings/index.html @@ -0,0 +1 @@ +LiteLLM Dashboard
Loading...
\ No newline at end of file diff --git a/litellm/proxy/_experimental/out/settings/logging-and-alerts/index.html b/litellm/proxy/_experimental/out/settings/logging-and-alerts/index.html new file mode 100644 index 00000000000..7d5d218fda4 --- /dev/null +++ b/litellm/proxy/_experimental/out/settings/logging-and-alerts/index.html @@ -0,0 +1 @@ +LiteLLM Dashboard
Loading...
\ No newline at end of file diff --git a/litellm/proxy/_experimental/out/settings/router-settings/index.html b/litellm/proxy/_experimental/out/settings/router-settings/index.html new file mode 100644 index 00000000000..eb3fd3fde00 --- /dev/null +++ b/litellm/proxy/_experimental/out/settings/router-settings/index.html @@ -0,0 +1 @@ +LiteLLM Dashboard
Loading...
\ No newline at end of file diff --git a/litellm/proxy/_experimental/out/settings/ui-theme/index.html b/litellm/proxy/_experimental/out/settings/ui-theme/index.html new file mode 100644 index 00000000000..17d352321c5 --- /dev/null +++ b/litellm/proxy/_experimental/out/settings/ui-theme/index.html @@ -0,0 +1 @@ +LiteLLM Dashboard
Loading...
\ No newline at end of file diff --git a/litellm/proxy/_experimental/out/teams/index.html b/litellm/proxy/_experimental/out/teams/index.html new file mode 100644 index 00000000000..781441c0732 --- /dev/null +++ b/litellm/proxy/_experimental/out/teams/index.html @@ -0,0 +1 @@ +LiteLLM Dashboard
Loading...
\ No newline at end of file diff --git a/litellm/proxy/_experimental/out/test-key/index.html b/litellm/proxy/_experimental/out/test-key/index.html new file mode 100644 index 00000000000..22c06d24381 --- /dev/null +++ b/litellm/proxy/_experimental/out/test-key/index.html @@ -0,0 +1 @@ +LiteLLM Dashboard
Loading...
\ No newline at end of file diff --git a/litellm/proxy/_experimental/out/tools/mcp-servers/index.html b/litellm/proxy/_experimental/out/tools/mcp-servers/index.html new file mode 100644 index 00000000000..64b747528e0 --- /dev/null +++ b/litellm/proxy/_experimental/out/tools/mcp-servers/index.html @@ -0,0 +1 @@ +LiteLLM Dashboard
Loading...
\ No newline at end of file diff --git a/litellm/proxy/_experimental/out/tools/vector-stores/index.html b/litellm/proxy/_experimental/out/tools/vector-stores/index.html new file mode 100644 index 00000000000..098b0d212c6 --- /dev/null +++ b/litellm/proxy/_experimental/out/tools/vector-stores/index.html @@ -0,0 +1 @@ +LiteLLM Dashboard
Loading...
\ No newline at end of file diff --git a/litellm/proxy/_experimental/out/usage/index.html b/litellm/proxy/_experimental/out/usage/index.html new file mode 100644 index 00000000000..ed6ac2eba97 --- /dev/null +++ b/litellm/proxy/_experimental/out/usage/index.html @@ -0,0 +1 @@ +LiteLLM Dashboard
Loading...
\ No newline at end of file diff --git a/litellm/proxy/_experimental/out/users/index.html b/litellm/proxy/_experimental/out/users/index.html new file mode 100644 index 00000000000..247dda941bd --- /dev/null +++ b/litellm/proxy/_experimental/out/users/index.html @@ -0,0 +1 @@ +LiteLLM Dashboard
Loading...
\ No newline at end of file diff --git a/litellm/proxy/_experimental/out/virtual-keys/index.html b/litellm/proxy/_experimental/out/virtual-keys/index.html new file mode 100644 index 00000000000..b17ef6de095 --- /dev/null +++ b/litellm/proxy/_experimental/out/virtual-keys/index.html @@ -0,0 +1 @@ +LiteLLM Dashboard
Loading...
\ No newline at end of file diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index a592a81ed86..d3bf3f1998a 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -691,6 +691,13 @@ class LiteLLMRoutes(enum.Enum): "/organization/delete", "/organization/member_add", "/organization/member_update", + # member_delete is equally destructive as member_add / member_update + # and must be scoped the same way — otherwise it falls through to + # the management_routes / self_managed_routes path and lets any + # non-PROXY_ADMIN caller that reaches the route delete arbitrary + # org memberships without the organization_role_based_access_check + # that member_add / member_update trigger. + "/organization/member_delete", ] # Routes accessible by Admin Viewer (read-only admin access) diff --git a/litellm/proxy/agent_endpoints/a2a_routing.py b/litellm/proxy/agent_endpoints/a2a_routing.py index 8f951414994..a9b75de9807 100644 --- a/litellm/proxy/agent_endpoints/a2a_routing.py +++ b/litellm/proxy/agent_endpoints/a2a_routing.py @@ -8,10 +8,17 @@ Looks up agents in the registry and injects their API base URL. from typing import Any, Optional import litellm +from fastapi import HTTPException + from litellm._logging import verbose_proxy_logger +from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth -def route_a2a_agent_request(data: dict, route_type: str) -> Optional[Any]: +async def route_a2a_agent_request( + data: dict, + route_type: str, + user_api_key_dict: Optional[UserAPIKeyAuth] = None, +) -> Optional[Any]: """ Route A2A agent requests directly to litellm with injected API base. @@ -19,6 +26,9 @@ def route_a2a_agent_request(data: dict, route_type: str) -> Optional[Any]: """ # Import here to avoid circular imports from litellm.proxy.agent_endpoints.agent_registry import global_agent_registry + from litellm.proxy.agent_endpoints.auth.agent_permission_handler import ( + AgentRequestHandler, + ) from litellm.proxy.route_llm_request import ( ROUTE_ENDPOINT_MAPPING, ProxyModelNotFoundError, @@ -40,6 +50,22 @@ def route_a2a_agent_request(data: dict, route_type: str) -> Optional[Any]: route_name = ROUTE_ENDPOINT_MAPPING.get(route_type, route_type) raise ProxyModelNotFoundError(route=route_name, model_name=model_name) + # Verify the caller is permitted to use this agent (admins bypass the check) + is_admin = user_api_key_dict is not None and ( + user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN + or user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN.value + ) + if not is_admin: + is_allowed = await AgentRequestHandler.is_agent_allowed( + agent_id=agent.agent_id, + user_api_key_auth=user_api_key_dict, + ) + if not is_allowed: + raise HTTPException( + status_code=403, + detail=f"Agent '{agent_name}' is not allowed for your key/team. Contact proxy admin for access.", + ) + # Get API base URL from agent config if not agent.agent_card_params or "url" not in agent.agent_card_params: verbose_proxy_logger.error(f"[A2A] Agent '{agent_name}' has no URL configured") diff --git a/litellm/proxy/agent_endpoints/endpoints.py b/litellm/proxy/agent_endpoints/endpoints.py index 063610ac5ef..c351d3cfecb 100644 --- a/litellm/proxy/agent_endpoints/endpoints.py +++ b/litellm/proxy/agent_endpoints/endpoints.py @@ -392,6 +392,24 @@ async def get_agent_by_id( """ await check_feature_access_for_user(user_api_key_dict, "agents") + is_admin = ( + user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN + or user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN.value + ) + if not is_admin: + from litellm.proxy.agent_endpoints.auth.agent_permission_handler import ( + AgentRequestHandler, + ) + + is_allowed = await AgentRequestHandler.is_agent_allowed( + agent_id=agent_id, user_api_key_auth=user_api_key_dict + ) + if not is_allowed: + raise HTTPException( + status_code=403, + detail=f"Agent '{agent_id}' is not allowed for your key/team. Contact proxy admin for access.", + ) + from litellm.proxy.proxy_server import prisma_client if prisma_client is None: diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 290b34be836..ddee8115c34 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -32,6 +32,7 @@ from litellm.constants import ( ) from litellm.litellm_core_utils.dd_tracing import tracer from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider +from litellm.litellm_core_utils.safe_json_loads import safe_json_loads from litellm.proxy._types import ( RBAC_ROLES, CallInfo, @@ -328,15 +329,63 @@ def _global_proxy_budget_check( ) +_GUARDRAIL_MODIFICATION_KEYS: tuple = ( + "guardrails", + "disable_global_guardrails", + "disable_global_guardrail", + "opted_out_global_guardrails", +) + + def _guardrail_modification_check( request_body: dict, team_object: Optional[LiteLLM_TeamTable] ) -> None: - _request_metadata: dict = request_body.get("metadata", {}) or {} - if not _request_metadata.get("guardrails"): - return + """ + Reject user-supplied metadata flags that would modify guardrail behavior + unless the team has explicit permission. Checked keys include the plural + ``guardrails`` list plus the per-request toggles that influence whether + default-on guardrails run (``disable_global_guardrails``, + ``disable_global_guardrail`` singular, and ``opted_out_global_guardrails``). + User-supplied values for the bypass toggles are also silently ignored by + ``_get_admin_metadata`` at read time; this check adds defense in depth by + failing loudly at the auth layer so operators see an explicit 403 instead + of a confusing silent-ignore. + """ from litellm.proxy.guardrails.guardrail_helpers import can_modify_guardrails + def _coerce_to_dict(container: Any) -> Optional[dict]: + """Accept dict or JSON-string (from multipart/form-data or extra_body). + + Without this, an attacker can smuggle guardrail keys past the check by + sending ``{"metadata": "{\\"disable_global_guardrails\\": true}"}`` — + ``isinstance(dict)`` on the string returns False, the check returns + no-modification, and ``add_litellm_data_to_request`` parses the string + to a dict downstream. + """ + if isinstance(container, dict): + return container + if isinstance(container, str): + parsed = safe_json_loads(container) + return parsed if isinstance(parsed, dict) else None + return None + + def _user_requested_modification(container: Any) -> bool: + coerced = _coerce_to_dict(container) + if coerced is None: + return False + return any(coerced.get(key) for key in _GUARDRAIL_MODIFICATION_KEYS) + + # Check both metadata keys — callers can populate either depending on the + # endpoint. Cover the top-level too so root-level injection is rejected. + modifies = ( + _user_requested_modification(request_body.get("metadata")) + or _user_requested_modification(request_body.get("litellm_metadata")) + or _user_requested_modification(request_body) + ) + if not modifies: + return + if not can_modify_guardrails(team_object): raise HTTPException( status_code=403, diff --git a/litellm/proxy/auth/auth_checks_organization.py b/litellm/proxy/auth/auth_checks_organization.py index 50efe137209..d89afcffa9a 100644 --- a/litellm/proxy/auth/auth_checks_organization.py +++ b/litellm/proxy/auth/auth_checks_organization.py @@ -144,7 +144,7 @@ def _user_is_org_admin( user_object: Optional[LiteLLM_UserTable] = None, ) -> bool: """ - Helper function to check if user is an org admin for any of the passed organizations. + Helper function to check if user is an org admin for all of the passed organizations. Checks both: - `organization_id` (singular string) — legacy callers @@ -168,9 +168,13 @@ def _user_is_org_admin( if not candidate_org_ids: return False - for _membership in user_object.organization_memberships: - if _membership.organization_id in candidate_org_ids: - if _membership.user_role == LitellmUserRoles.ORG_ADMIN.value: - return True + # Build set of orgs where user is admin + admin_org_ids = { + _membership.organization_id + for _membership in user_object.organization_memberships + if _membership.user_role == LitellmUserRoles.ORG_ADMIN.value + and _membership.organization_id is not None + } - return False + # User must be admin of ALL requested orgs, not just any one + return all(org_id in admin_org_ids for org_id in candidate_org_ids) diff --git a/litellm/proxy/auth/auth_utils.py b/litellm/proxy/auth/auth_utils.py index e15b388ccea..18aea48e96b 100644 --- a/litellm/proxy/auth/auth_utils.py +++ b/litellm/proxy/auth/auth_utils.py @@ -842,19 +842,30 @@ def get_end_user_id_from_request_body( user_from_body_user_field = request_body["user"] return str(user_from_body_user_field) + def _as_dict(value: Any) -> dict: + # metadata / litellm_metadata can arrive as JSON strings from + # multipart/form-data or extra_body; coerce so string-encoded + # payloads can't evade end-user attribution. + if isinstance(value, dict): + return value + if isinstance(value, str): + from litellm.litellm_core_utils.safe_json_loads import safe_json_loads + + parsed = safe_json_loads(value) + return parsed if isinstance(parsed, dict) else {} + return {} + # Check 4: 'litellm_metadata.user' in request_body (commonly Anthropic) - litellm_metadata = request_body.get("litellm_metadata") - if isinstance(litellm_metadata, dict): - user_from_litellm_metadata = litellm_metadata.get("user") - if user_from_litellm_metadata is not None: - return str(user_from_litellm_metadata) + litellm_metadata = _as_dict(request_body.get("litellm_metadata")) + user_from_litellm_metadata = litellm_metadata.get("user") + if user_from_litellm_metadata is not None: + return str(user_from_litellm_metadata) # Check 5: 'metadata.user_id' in request_body (another common pattern) - metadata_dict = request_body.get("metadata") - if isinstance(metadata_dict, dict): - user_id_from_metadata_field = metadata_dict.get("user_id") - if user_id_from_metadata_field is not None: - return str(user_id_from_metadata_field) + metadata_dict = _as_dict(request_body.get("metadata")) + user_id_from_metadata_field = metadata_dict.get("user_id") + if user_id_from_metadata_field is not None: + return str(user_id_from_metadata_field) # Check 6: 'safety_identifier' in request body (OpenAI Responses API parameter) # SECURITY NOTE: safety_identifier can be set by any caller in the request body. diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index 3b89f8e1d25..97801baaf0c 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -1056,6 +1056,7 @@ class ProxyBaseLLMRequestProcessing: route_type=route_type, llm_router=llm_router, user_model=user_model, + user_api_key_dict=user_api_key_dict, ) tasks.append(llm_call) diff --git a/litellm/proxy/common_utils/http_parsing_utils.py b/litellm/proxy/common_utils/http_parsing_utils.py index 88c3a6be9da..71abdfa5e9e 100644 --- a/litellm/proxy/common_utils/http_parsing_utils.py +++ b/litellm/proxy/common_utils/http_parsing_utils.py @@ -426,7 +426,16 @@ def get_tags_from_request_body(request_body: dict) -> List[str]: List of tag names (strings), empty list if no valid tags found """ metadata_variable_name = get_metadata_variable_name_from_kwargs(request_body) - metadata = request_body.get(metadata_variable_name) or {} + metadata = request_body.get(metadata_variable_name) + # metadata can arrive as a JSON string from multipart/form-data or extra_body; + # coerce defensively so .get() below never raises AttributeError. + if isinstance(metadata, str): + from litellm.litellm_core_utils.safe_json_loads import safe_json_loads + + parsed = safe_json_loads(metadata) + metadata = parsed if isinstance(parsed, dict) else {} + elif not isinstance(metadata, dict): + metadata = {} tags_in_metadata: Any = metadata.get("tags", []) tags_in_request_body: Any = request_body.get("tags", []) combined_tags: List[str] = [] diff --git a/litellm/proxy/db/create_views.py b/litellm/proxy/db/create_views.py index 3598045545b..d84cebcf05a 100644 --- a/litellm/proxy/db/create_views.py +++ b/litellm/proxy/db/create_views.py @@ -251,7 +251,7 @@ async def should_create_missing_views(db: _db) -> bool: and len(result) > 0 and isinstance(result[0], dict) and "reltuples" in result[0] - and result[0]["reltuples"] + and result[0]["reltuples"] is not None and (result[0]["reltuples"] == 0 or result[0]["reltuples"] == -1) ): verbose_logger.debug("Should create views") diff --git a/litellm/proxy/litellm_pre_call_utils.py b/litellm/proxy/litellm_pre_call_utils.py index b0adf7aa6ee..7467bbae232 100644 --- a/litellm/proxy/litellm_pre_call_utils.py +++ b/litellm/proxy/litellm_pre_call_utils.py @@ -981,13 +981,16 @@ async def add_litellm_data_to_request( # noqa: PLR0915 # Init - Proxy Server Request # we do this as soon as entering so we track the original request ########################################################## - # Track arrival time for queue time metric + # Track arrival time for queue time metric. The body snapshot is filled + # in after the admin-injection strip below so the audit / spend-tracking + # consumers of proxy_server_request["body"] see the cleaned metadata + # rather than attacker-forged user_api_key_* fields. arrival_time = time.time() data["proxy_server_request"] = { "url": str(request.url), "method": request.method, "headers": _headers, - "body": copy.copy(data), # use copy instead of deepcopy + "body": None, # filled in post-strip; see below "arrival_time": arrival_time, # Track when request arrived at proxy } @@ -1069,9 +1072,10 @@ async def add_litellm_data_to_request( # noqa: PLR0915 verbose_proxy_logger.warning( f"Failed to parse 'metadata' as JSON dict. Received value: {data['metadata']}" ) - data[_metadata_variable_name]["requester_metadata"] = copy.deepcopy( - data["metadata"] - ) + # requester_metadata is snapshotted AFTER the strip below so + # downstream consumers (e.g. PANW guardrail reading user_ip / + # profile_id) don't see attacker-injected admin slots preserved in + # the deepcopy. # Parse litellm_metadata if it's a string (e.g., from multipart/form-data or extra_body) if "litellm_metadata" in data and data["litellm_metadata"] is not None: @@ -1083,11 +1087,89 @@ async def add_litellm_data_to_request( # noqa: PLR0915 ) else: data["litellm_metadata"] = parsed_litellm_metadata - # Merge litellm_metadata into the metadata variable (preserving existing values) - if isinstance(data["litellm_metadata"], dict): - for key, value in data["litellm_metadata"].items(): - if key not in data[_metadata_variable_name]: - data[_metadata_variable_name][key] = value + + # Strip internal pipeline state and admin-injection slots from user input. + # Runs AFTER the string-to-dict parse above so JSON-string metadata (sent + # via multipart/form-data or extra_body) cannot smuggle admin fields past + # the isinstance(dict) guard. + # + # The proxy populates a family of ``user_api_key_*`` fields below + # (user_api_key_metadata, user_api_key_user_id, user_api_key_alias, + # user_api_key_spend, user_api_key_team_metadata, …) into + # data[_metadata_variable_name]. Because the proxy only writes to ONE of + # the two metadata dicts, a caller pre-populating any of these keys on + # the OTHER metadata dict would have their forged values surface in + # guardrails, spend tracking, audit logs, and identity resolution. Strip + # by prefix so new ``user_api_key_*`` fields added in the future are + # covered without per-key maintenance. + for _meta_key in ("metadata", "litellm_metadata"): + _user_meta = data.get(_meta_key) + if isinstance(_user_meta, dict): + _user_meta.pop("_pipeline_managed_guardrails", None) + for _k in [k for k in _user_meta if k.startswith("user_api_key_")]: + _user_meta.pop(_k, None) + + # Strip caller-supplied routing/budget tags unless the admin has opted + # this key or team in via metadata.allow_client_tags=True. Tags drive + # tag-based routing and tag budget attribution — accepting them from + # untrusted callers lets an attacker reach restricted deployments or + # misattribute spend to a victim team's tag. + _admin_allow_client_tags = False + for _admin_meta in ( + user_api_key_dict.metadata, + user_api_key_dict.team_metadata, + ): + if ( + isinstance(_admin_meta, dict) + and _admin_meta.get("allow_client_tags") is True + ): + _admin_allow_client_tags = True + break + if not _admin_allow_client_tags: + _stripped_from: List[str] = [] + for _meta_key in ("metadata", "litellm_metadata"): + _user_meta = data.get(_meta_key) + if isinstance(_user_meta, dict) and "tags" in _user_meta: + _user_meta.pop("tags", None) + _stripped_from.append(_meta_key) + # Also strip the root-level `tags` field. get_tags_from_request_body + # reads request_body["tags"] directly and feeds it to the policy + # engine, so leaving it in place here would let the strip-in-metadata + # above be trivially bypassed by moving the tags to the body root. + if "tags" in data: + data.pop("tags", None) + _stripped_from.append("tags (root)") + if _stripped_from: + verbose_proxy_logger.warning( + "Stripped caller-supplied tags from %s: this key/team does " + "not have `allow_client_tags: true` in its metadata. Set it " + "to opt into client-supplied routing/budget tags.", + ", ".join(_stripped_from), + ) + + # Fill in the proxy_server_request body snapshot now that metadata has + # been parsed and stripped. Consumers (standard_logging_payload, lago, + # spend_tracking_utils, streaming_iterator) read `body` to audit the + # request; taking the snapshot here ensures they see cleaned metadata. + data["proxy_server_request"]["body"] = copy.copy(data) + + # Snapshot the (now-cleaned) requester-supplied metadata for downstream + # consumers. Taking the deepcopy AFTER the strip prevents attacker- + # injected admin slots (user_api_key_*, tags without opt-in, + # _pipeline_managed_guardrails) from surviving in requester_metadata + # where guardrails and audit paths may read from it. + if "metadata" in data and isinstance(data["metadata"], dict): + data[_metadata_variable_name]["requester_metadata"] = copy.deepcopy( + data["metadata"] + ) + + # Now merge litellm_metadata into the metadata variable (preserving existing + # values) — runs AFTER the strip so attacker injections in litellm_metadata + # cannot cross-contaminate the admin-authoritative metadata dict. + if "litellm_metadata" in data and isinstance(data["litellm_metadata"], dict): + for key, value in data["litellm_metadata"].items(): + if key not in data[_metadata_variable_name]: + data[_metadata_variable_name][key] = value data = LiteLLMProxyRequestSetup.add_user_api_key_auth_to_request_metadata( data=data, @@ -1250,15 +1332,24 @@ async def add_litellm_data_to_request( # noqa: PLR0915 user_agent = request.headers["user-agent"] data[_metadata_variable_name]["user_agent"] = user_agent - # Check if using tag based routing + # Check if using tag based routing. The helper reads caller-controlled + # sources (x-litellm-tags header, data["tags"] root-level), so its result + # is still gated by the same allow_client_tags flag that gated the + # body-metadata tag strip above. Otherwise the strip is trivially + # bypassed by sending tags via header or at the root of the body. tags = LiteLLMProxyRequestSetup.add_request_tag_to_metadata( llm_router=llm_router, headers=_headers, data=data, ) - if tags is not None: + if tags is not None and _admin_allow_client_tags: data[_metadata_variable_name]["tags"] = tags + elif tags is not None: + verbose_proxy_logger.warning( + "Ignored caller-supplied tags from header/root body: this " + "key/team does not have `allow_client_tags: true` in its metadata." + ) # Team Callbacks controls callback_settings_obj = _get_dynamic_logging_metadata( diff --git a/litellm/proxy/management_endpoints/internal_user_endpoints.py b/litellm/proxy/management_endpoints/internal_user_endpoints.py index ff98a38313a..8474f026111 100644 --- a/litellm/proxy/management_endpoints/internal_user_endpoints.py +++ b/litellm/proxy/management_endpoints/internal_user_endpoints.py @@ -1181,6 +1181,23 @@ async def _update_single_user_helper( "error": "User does not have permission to update this user. Only PROXY_ADMIN can update other users." }, ) + else: + # Silent-create guard: if the target user doesn't exist, the update + # path falls through to an upsert that creates a new user with + # caller-supplied fields (models, metadata, budgets, …). Only + # PROXY_ADMIN is allowed to create users this way; otherwise an org + # admin could spawn arbitrary users attached to nothing by supplying + # a fresh email, bypassing the /user/new org/team-scoping checks. + if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN.value: + raise HTTPException( + status_code=404, + detail={ + "error": ( + "User not found. Only PROXY_ADMIN can create users " + "via /user/update; use /user/new instead." + ) + }, + ) existing_metadata = ( cast(Dict, getattr(existing_user_row, "metadata", {}) or {}) @@ -2059,6 +2076,54 @@ async def delete_user( if data.user_ids is None: raise HTTPException(status_code=400, detail={"error": "No user id passed in"}) + # Per-target authorization: the route-level gate accepts this call when + # the caller is PROXY_ADMIN or an ORG_ADMIN of *any* org named in + # request_data["organization_id"]/["organizations"]. That gate does NOT + # cross-check data.user_ids against the caller's scope, so without this + # loop an org-admin of org-A could delete users in org-B by supplying + # {"user_ids": [victim_in_org_B], "organization_id": "org-A"}. + caller_is_proxy_admin = ( + user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN.value + ) + caller_admin_org_ids: set = set() + if not caller_is_proxy_admin: + caller_memberships = ( + await prisma_client.db.litellm_organizationmembership.find_many( + where={ + "user_id": user_api_key_dict.user_id, + "user_role": LitellmUserRoles.ORG_ADMIN.value, + } + ) + if user_api_key_dict.user_id + else [] + ) + caller_admin_org_ids = { + m.organization_id for m in caller_memberships if m.organization_id + } + if not caller_admin_org_ids: + raise HTTPException( + status_code=403, + detail={ + "error": "Only PROXY_ADMIN or ORG_ADMIN users may delete users." + }, + ) + + # Batch-fetch target memberships once before the per-user loop. Avoids + # an N+1 DB call when delete_user is called with a large user_ids list. + target_org_ids_by_user: Dict[str, set] = {} + if not caller_is_proxy_admin: + all_target_memberships = ( + await prisma_client.db.litellm_organizationmembership.find_many( + where={"user_id": {"in": data.user_ids}} + ) + ) + for m in all_target_memberships: + if not m.organization_id: + continue + target_org_ids_by_user.setdefault(m.user_id, set()).add( + m.organization_id + ) + # check that all teams passed exist for user_id in data.user_ids: user_row = await prisma_client.db.litellm_usertable.find_unique( @@ -2070,30 +2135,49 @@ async def delete_user( status_code=404, detail={"error": f"User not found, passed user_id={user_id}"}, ) - else: - # Enterprise Feature - Audit Logging. Enable with litellm.store_audit_logs = True - # we do this after the first for loop, since first for loop is for validation. we only want this inserted after validation passes - if litellm.store_audit_logs is True: - # make an audit log for each team deleted - _user_row = user_row.json(exclude_none=True) - asyncio.create_task( - create_audit_log_for_update( - request_data=LiteLLM_AuditLogs( - id=str(uuid.uuid4()), - updated_at=datetime.now(timezone.utc), - changed_by=litellm_changed_by - or user_api_key_dict.user_id - or litellm_proxy_admin_name, - changed_by_api_key=user_api_key_dict.api_key, - table_name=LitellmTableNames.USER_TABLE_NAME, - object_id=user_id, - action="deleted", - updated_values="{}", - before_value=_user_row, + if not caller_is_proxy_admin: + target_org_ids = target_org_ids_by_user.get(user_id, set()) + # Org-admin may only delete users whose entire org membership is + # within their admin scope. A target with ANY org outside the + # caller's scope (or no org at all) requires PROXY_ADMIN. + if not target_org_ids or not target_org_ids.issubset( + caller_admin_org_ids + ): + raise HTTPException( + status_code=403, + detail={ + "error": ( + f"User {user_id} is not within your admin scope. " + "Only PROXY_ADMIN may delete users outside your " + "administered organizations." ) + }, + ) + + # Enterprise Feature - Audit Logging. Enable with litellm.store_audit_logs = True + # we do this after the first for loop, since first for loop is for validation. we only want this inserted after validation passes + if litellm.store_audit_logs is True: + # make an audit log for each team deleted + _user_row = user_row.json(exclude_none=True) + + asyncio.create_task( + create_audit_log_for_update( + request_data=LiteLLM_AuditLogs( + id=str(uuid.uuid4()), + updated_at=datetime.now(timezone.utc), + changed_by=litellm_changed_by + or user_api_key_dict.user_id + or litellm_proxy_admin_name, + changed_by_api_key=user_api_key_dict.api_key, + table_name=LitellmTableNames.USER_TABLE_NAME, + object_id=user_id, + action="deleted", + updated_values="{}", + before_value=_user_row, ) ) + ) ## CLEANUP MEMBERS_WITH_ROLES fetch_all_teams = await prisma_client.db.litellm_teamtable.find_many( diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index 028b765e021..b3cbf95401f 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -1975,22 +1975,57 @@ async def _validate_update_key_data( user_api_key_cache=user_api_key_cache, ) - # Admin-only: only proxy admins, team admins, or org admins can modify max_budget or spend - if ( + # Cross-key authorization. Previously only gated on max_budget/spend + # changes, which let a non-admin blanket-rewrite any OTHER field on + # any key (models, alias, metadata, tpm_limit, rpm_limit, + # allowed_routes, guardrails, blocked, duration, permissions, …) as + # long as they avoided budget/spend. + # + # Policy: + # - Key owner (same user_id): may update non-budget fields on their + # own key without the admin check. + # - Team member with /key/update grant (on a team key): may update + # non-budget fields. Team membership + permission is already + # enforced by can_team_member_execute_key_management_endpoint + # above, which raises 401 for non-members or members without the + # grant — so reaching this point on a team key means the caller + # was authorized via member_permissions. This preserves the + # documented member_permissions feature while still blocking the + # cross-org attack (an outside org admin is not a member of the + # victim team and gets rejected at the earlier check). + # - Anyone else (non-PROXY_ADMIN, not the owner, not a team member + # on a team key): must pass _check_key_admin_access (PROXY_ADMIN + # / key-owner / team-admin / org-admin of the key). + # - max_budget / spend: always require the admin check, even for the + # key owner or a team member (matches the existing admin-only + # budget semantics). + is_key_owner = ( + user_api_key_dict.user_id is not None + and existing_key_row.user_id == user_api_key_dict.user_id + ) + _is_budget_change = ( data.max_budget is not None and data.max_budget != existing_key_row.max_budget ) or ( data.spend is not None and data.spend != getattr(existing_key_row, "spend", None) + ) + is_team_key = existing_key_row.team_id is not None + can_skip_admin_check_for_non_budget = is_key_owner or is_team_key + if ( + (not _is_proxy_admin) + and prisma_client is not None + and (_is_budget_change or not can_skip_admin_check_for_non_budget) ): - if prisma_client is not None: - hashed_key = existing_key_row.token - await _check_key_admin_access( - user_api_key_dict=user_api_key_dict, - hashed_token=hashed_key, - prisma_client=prisma_client, - user_api_key_cache=user_api_key_cache, - route="/key/update (max_budget/spend)", - ) + hashed_key = existing_key_row.token + await _check_key_admin_access( + user_api_key_dict=user_api_key_dict, + hashed_token=hashed_key, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + route=( + "/key/update (max_budget/spend)" if _is_budget_change else "/key/update" + ), + ) # Check team limits if key has a team_id (from request or existing key) team_obj: Optional[LiteLLM_TeamTableCachedObj] = None diff --git a/litellm/proxy/management_endpoints/organization_endpoints.py b/litellm/proxy/management_endpoints/organization_endpoints.py index a712088dec0..a6a1af971e5 100644 --- a/litellm/proxy/management_endpoints/organization_endpoints.py +++ b/litellm/proxy/management_endpoints/organization_endpoints.py @@ -1065,6 +1065,33 @@ async def organization_member_update( }, ) + # Reject attempts to change the role of a global PROXY_ADMIN via + # org-scoped operations. An org-admin of any org could otherwise + # alter a PROXY_ADMIN user's per-org role, which has downstream + # effects on admin UI filtering and scope derivation. + target_user_row = await prisma_client.db.litellm_usertable.find_unique( + where={"user_id": data.user_id} + ) + if target_user_row is not None and getattr( + target_user_row, "user_role", None + ) in ( + LitellmUserRoles.PROXY_ADMIN.value, + LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY.value, + ): + if ( + user_api_key_dict.user_role + != LitellmUserRoles.PROXY_ADMIN.value + ): + raise HTTPException( + status_code=403, + detail={ + "error": ( + "Only PROXY_ADMIN may modify the organization " + "role of a user who is a global PROXY_ADMIN." + ) + }, + ) + # Update member role if data.role is not None: await prisma_client.db.litellm_organizationmembership.update( diff --git a/litellm/proxy/management_endpoints/team_endpoints.py b/litellm/proxy/management_endpoints/team_endpoints.py index 377adb216c2..8e21b851857 100644 --- a/litellm/proxy/management_endpoints/team_endpoints.py +++ b/litellm/proxy/management_endpoints/team_endpoints.py @@ -1561,6 +1561,42 @@ async def update_team( # noqa: PLR0915 if ( data.organization_id is not None and len(data.organization_id) > 0 ): # allow unsetting the organization_id + # If the caller is relocating the team to a different org, they + # must also be PROXY_ADMIN or an org-admin of the DESTINATION org. + # _verify_team_access above only checked the team's CURRENT org, + # so without this gate an org-admin could hand their team to any + # other org (or capture a team from another org they once + # administered into a new destination). + current_org_id = getattr(existing_team_row, "organization_id", None) + if ( + data.organization_id != current_org_id + and user_api_key_dict.user_role + != LitellmUserRoles.PROXY_ADMIN.value + ): + # Is the caller org_admin of the destination org? + caller_memberships = ( + await prisma_client.db.litellm_organizationmembership.find_many( + where={ + "user_id": user_api_key_dict.user_id, + "organization_id": data.organization_id, + "user_role": LitellmUserRoles.ORG_ADMIN.value, + } + ) + if user_api_key_dict.user_id + else [] + ) + if not caller_memberships: + raise HTTPException( + status_code=403, + detail={ + "error": ( + "Relocating a team to a different organization " + "requires PROXY_ADMIN or org-admin of the " + "destination org." + ) + }, + ) + await fetch_and_validate_organization( organization_id=data.organization_id, existing_team_row=existing_team_row, @@ -2720,6 +2756,20 @@ async def bulk_team_member_add( ) if data.all_users: + # `all_users=True` pulls every user in the database into this team, + # regardless of org. Any team admin could use it to capture every + # user across every org into a team they control. Restrict to + # PROXY_ADMIN. + if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN.value: + raise HTTPException( + status_code=403, + detail={ + "error": ( + "`all_users=true` is restricted to PROXY_ADMIN. " + "Org/team admins must specify explicit member lists." + ) + }, + ) # get all users from the database all_users_in_db = await prisma_client.db.litellm_usertable.find_many( order={"created_at": "desc"} @@ -3448,6 +3498,7 @@ async def _get_org_admin_org_ids( m.organization_id for m in (caller_user.organization_memberships or []) if m.user_role == LitellmUserRoles.ORG_ADMIN.value + and m.organization_id is not None ] return org_ids if org_ids else None @@ -3482,13 +3533,8 @@ async def _build_team_list_where_conditions( if organization_id: where_conditions["organization_id"] = organization_id - elif org_admin_org_ids is not None and not user_id: - # Org admin without explicit org or user filter: scope to their orgs. - # NOTE: when user_id is provided, no org filter is applied — the - # query returns all teams the target user belongs to across all - # organisations. This matches the legacy /team/list behaviour in - # _authorize_and_filter_teams which fetches direct-membership teams - # without an org constraint. + elif org_admin_org_ids is not None: + # Org admin: always scope to their orgs, even when filtering by user_id. where_conditions["organization_id"] = {"in": org_admin_org_ids} if user_id: @@ -3858,7 +3904,7 @@ async def _authorize_and_filter_teams( Authorize the /team/list request and return filtered teams. - Proxy admins: all teams (or filtered by user_id if provided). - - Org admins: teams from their orgs + teams they are direct members of. + - Org admins: teams from their orgs (scoped to user_id if provided). - Own query (user_id matches caller): teams the user is a member of. - Others: 401. """ @@ -3886,6 +3932,7 @@ async def _authorize_and_filter_teams( m.organization_id for m in (caller_user.organization_memberships or []) if m.user_role == LitellmUserRoles.ORG_ADMIN.value + and m.organization_id is not None ] if not allowed_org_ids: allowed_org_ids = None @@ -3908,20 +3955,13 @@ async def _authorize_and_filter_teams( ) if not user_id: return list(org_teams) - # Also include teams the user is a direct member of (outside their orgs) - seen_team_ids = {team.team_id for team in org_teams} - all_teams = list(org_teams) - # Prisma doesn't support filtering JSON array fields, so we fetch by membership separately - member_teams = await prisma_client.db.litellm_teamtable.find_many( - where={"team_id": {"not_in": list(seen_team_ids)}} if seen_team_ids else {}, - include={"litellm_model_table": True}, - ) - for team in member_teams: - if team.members_with_roles and any( - m.get("user_id") == user_id for m in team.members_with_roles - ): - all_teams.append(team) - return all_teams + # Filter org teams to only those where the target user is a member + return [ + team + for team in org_teams + if team.members_with_roles + and any(m.get("user_id") == user_id for m in team.members_with_roles) + ] elif user_id: # Regular user: fetch all and filter by membership (Prisma can't filter JSON arrays) response = await prisma_client.db.litellm_teamtable.find_many( diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index f3ff3f5e23f..d8354a798b1 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -7199,7 +7199,10 @@ async def chat_completion( # noqa: PLR0915 global user_temperature, user_request_timeout, user_max_tokens, user_api_base data = await _read_request_body(request=request) if user_api_key_dict is not None: - if data.get("metadata") is None: + if not isinstance(data.get("metadata"), dict): + # Covers both missing and JSON-string metadata (multipart / + # extra_body); otherwise `data["metadata"][k] = v` below raises + # TypeError on a string value and 500s the request. data["metadata"] = {} if ( hasattr(user_api_key_dict, "user_id") @@ -11443,7 +11446,9 @@ async def async_queue_request( # if users are using user_api_key_auth, set `user` in `data` data["user"] = user_api_key_dict.user_id - if "metadata" not in data: + if not isinstance(data.get("metadata"), dict): + # Covers both missing and JSON-string metadata (multipart / + # extra_body); see above for the same guard upstream. data["metadata"] = {} data["metadata"]["user_api_key"] = user_api_key_dict.api_key data["metadata"]["user_api_key_metadata"] = user_api_key_dict.metadata diff --git a/litellm/proxy/route_llm_request.py b/litellm/proxy/route_llm_request.py index f1590b16c24..17cc4374560 100644 --- a/litellm/proxy/route_llm_request.py +++ b/litellm/proxy/route_llm_request.py @@ -4,6 +4,7 @@ from typing import TYPE_CHECKING, Any, Literal, Optional from fastapi import HTTPException, status import litellm +from litellm.proxy._types import UserAPIKeyAuth if TYPE_CHECKING: from litellm.router import Router as _Router @@ -314,6 +315,7 @@ async def route_request( # noqa: PLR0915 - Complex routing function, refactorin "acancel_run", "adelete_run", ], + user_api_key_dict: Optional[UserAPIKeyAuth] = None, ): """ Common helper to route the request @@ -548,7 +550,9 @@ async def route_request( # noqa: PLR0915 - Complex routing function, refactorin route_a2a_agent_request, ) - result = route_a2a_agent_request(data, route_type) + result = await route_a2a_agent_request( + data, route_type, user_api_key_dict=user_api_key_dict + ) if result is not None: return result # Fall through to raise exception below if result is None diff --git a/litellm/rag/ingestion/base_ingestion.py b/litellm/rag/ingestion/base_ingestion.py index 0d12bdfffc1..6a4eb89d0fd 100644 --- a/litellm/rag/ingestion/base_ingestion.py +++ b/litellm/rag/ingestion/base_ingestion.py @@ -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" diff --git a/litellm/router_strategy/budget_limiter.py b/litellm/router_strategy/budget_limiter.py index 94231d13df4..be27b852478 100644 --- a/litellm/router_strategy/budget_limiter.py +++ b/litellm/router_strategy/budget_limiter.py @@ -29,6 +29,9 @@ from litellm.caching.redis_cache import RedisPipelineIncrementOperation from litellm.integrations.custom_logger import CustomLogger, Span from litellm.litellm_core_utils.duration_parser import duration_in_seconds from litellm.router_strategy.tag_based_routing import _get_tags_from_request_kwargs +from litellm.litellm_core_utils.core_helpers import ( + get_metadata_variable_name_from_kwargs, +) from litellm.router_utils.cooldown_callbacks import ( _get_prometheus_logger_from_callbacks, ) @@ -175,7 +178,10 @@ class RouterBudgetLimiting(CustomLogger): spend_map=spend_map, potential_deployments=potential_deployments, request_tags=_get_tags_from_request_kwargs( - request_kwargs=request_kwargs + request_kwargs=request_kwargs, + metadata_variable_name=get_metadata_variable_name_from_kwargs( + request_kwargs or {} + ), ), ) @@ -304,6 +310,16 @@ class RouterBudgetLimiting(CustomLogger): deployment_configs: Dict[str, GenericBudgetInfo] = {} deployment_providers: List[Optional[str]] = [] + # Resolve tags once before the loop (loop-invariant) + _request_tags: List[str] = [] + if self.tag_budget_config: + _request_tags = _get_tags_from_request_kwargs( + request_kwargs=request_kwargs, + metadata_variable_name=get_metadata_variable_name_from_kwargs( + request_kwargs or {} + ), + ) + for deployment in healthy_deployments: # Check provider budgets if self.provider_budget_config: @@ -330,17 +346,14 @@ class RouterBudgetLimiting(CustomLogger): cache_keys.append( f"deployment_spend:{model_id}:{budget_config.budget_duration}" ) - # Check tag budgets - if self.tag_budget_config: - request_tags = _get_tags_from_request_kwargs( - request_kwargs=request_kwargs + + # Check tag budgets (outside loop — tags are per-request, not per-deployment) + for _tag in _request_tags: + _tag_budget_config = self._get_budget_config_for_tag(_tag) + if _tag_budget_config: + cache_keys.append( + f"tag_spend:{_tag}:{_tag_budget_config.budget_duration}" ) - for _tag in request_tags: - _tag_budget_config = self._get_budget_config_for_tag(_tag) - if _tag_budget_config: - cache_keys.append( - f"tag_spend:{_tag}:{_tag_budget_config.budget_duration}" - ) return ( cache_keys, provider_configs, @@ -459,7 +472,10 @@ class RouterBudgetLimiting(CustomLogger): response_cost=response_cost, ) - request_tags = _get_tags_from_request_kwargs(kwargs) + request_tags = _get_tags_from_request_kwargs( + kwargs, + metadata_variable_name=get_metadata_variable_name_from_kwargs(kwargs or {}), + ) if len(request_tags) > 0: for _tag in request_tags: _tag_budget_config = self._get_budget_config_for_tag(_tag) diff --git a/litellm/router_strategy/tag_based_routing.py b/litellm/router_strategy/tag_based_routing.py index 1188ce9d592..0163f3bbd4f 100644 --- a/litellm/router_strategy/tag_based_routing.py +++ b/litellm/router_strategy/tag_based_routing.py @@ -102,13 +102,11 @@ def _match_deployment( return {"matched_via": "tags", "matched_value": matched_value} # 2. Regex match against request headers. - # When match_any=False and the deployment has both plain tags and tag_regex, - # the strict tag check has already failed (step 1 returned None). Allow - # the regex to fire only when the deployment has NO plain tags, so we never - # use regex as a backdoor around the operator's strict-tag policy. - strict_tag_check_failed = ( - not match_any and bool(deployment_tags) and bool(request_tags) - ) + # When match_any=False and the deployment has plain tags, the strict tag + # check either didn't run (no request tags) or failed (step 1 returned + # None). Block the regex path so it cannot circumvent the operator's + # strict-tag policy. + strict_tag_check_failed = not match_any and bool(deployment_tags) if deployment_tag_regex and header_strings and not strict_tag_check_failed: regex_match = _is_valid_deployment_tag_regex( deployment_tag_regex, header_strings diff --git a/tests/logging_callback_tests/test_logging_redaction_e2e_test.py b/tests/logging_callback_tests/test_logging_redaction_e2e_test.py index 1ca8b599d7c..63ef4bafbb8 100644 --- a/tests/logging_callback_tests/test_logging_redaction_e2e_test.py +++ b/tests/logging_callback_tests/test_logging_redaction_e2e_test.py @@ -56,7 +56,13 @@ async def test_global_redaction_on(): @pytest.mark.parametrize("turn_off_message_logging", [True, False]) @pytest.mark.asyncio -async def test_global_redaction_with_dynamic_params(turn_off_message_logging): +async def test_global_redaction_ignores_dynamic_param(turn_off_message_logging): + """ + Request-body `turn_off_message_logging` is no longer honored as a dynamic + callback param — global setting (or admin-configured key/team config) wins. + With global redaction ON, the caller cannot disable redaction via the + request body. + """ litellm.turn_off_message_logging = True test_custom_logger = TestCustomLogger() litellm.callbacks = [test_custom_logger] @@ -75,23 +81,20 @@ async def test_global_redaction_with_dynamic_params(turn_off_message_logging): json.dumps(standard_logging_payload, indent=2), ) - if turn_off_message_logging is True: - response = standard_logging_payload["response"] - assert response["choices"][0]["message"]["content"] == "redacted-by-litellm" - assert ( - standard_logging_payload["messages"][0]["content"] == "redacted-by-litellm" - ) - else: - assert ( - standard_logging_payload["response"]["choices"][0]["message"]["content"] - == "hello" - ) - assert standard_logging_payload["messages"][0]["content"] == "hi" + response = standard_logging_payload["response"] + assert response["choices"][0]["message"]["content"] == "redacted-by-litellm" + assert standard_logging_payload["messages"][0]["content"] == "redacted-by-litellm" @pytest.mark.parametrize("turn_off_message_logging", [True, False]) @pytest.mark.asyncio -async def test_global_redaction_off_with_dynamic_params(turn_off_message_logging): +async def test_global_redaction_off_ignores_dynamic_param(turn_off_message_logging): + """ + Request-body `turn_off_message_logging` is no longer honored as a dynamic + callback param — global setting (or admin-configured key/team config) wins. + With global redaction OFF, the caller cannot enable redaction via the + request body. + """ litellm.turn_off_message_logging = False test_custom_logger = TestCustomLogger() litellm.callbacks = [test_custom_logger] @@ -109,18 +112,11 @@ async def test_global_redaction_off_with_dynamic_params(turn_off_message_logging "logged standard logging payload", json.dumps(standard_logging_payload, indent=2), ) - if turn_off_message_logging is True: - response = standard_logging_payload["response"] - assert response["choices"][0]["message"]["content"] == "redacted-by-litellm" - assert ( - standard_logging_payload["messages"][0]["content"] == "redacted-by-litellm" - ) - else: - assert ( - standard_logging_payload["response"]["choices"][0]["message"]["content"] - == "hello" - ) - assert standard_logging_payload["messages"][0]["content"] == "hi" + assert ( + standard_logging_payload["response"]["choices"][0]["message"]["content"] + == "hello" + ) + assert standard_logging_payload["messages"][0]["content"] == "hi" @pytest.mark.asyncio diff --git a/tests/mcp_tests/test_openapi_spec_path_url.py b/tests/mcp_tests/test_openapi_spec_path_url.py index 1a3c16f12a4..17a0022046e 100644 --- a/tests/mcp_tests/test_openapi_spec_path_url.py +++ b/tests/mcp_tests/test_openapi_spec_path_url.py @@ -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") diff --git a/tests/pass_through_tests/test_vertex_ai.py b/tests/pass_through_tests/test_vertex_ai.py index ba27a4cc460..73bf03c5000 100644 --- a/tests/pass_through_tests/test_vertex_ai.py +++ b/tests/pass_through_tests/test_vertex_ai.py @@ -72,6 +72,10 @@ async def call_spend_logs_endpoint(): response = requests.get(url, headers=headers) print("response from call_spend_logs_endpoint", response) + if response.status_code != 200: + print(f"spend logs endpoint returned {response.status_code}: {response.text}") + return None + json_response = response.json() # get spend for today diff --git a/tests/proxy_unit_tests/test_proxy_utils.py b/tests/proxy_unit_tests/test_proxy_utils.py index 1de0ab450d7..8703331ea83 100644 --- a/tests/proxy_unit_tests/test_proxy_utils.py +++ b/tests/proxy_unit_tests/test_proxy_utils.py @@ -167,9 +167,13 @@ async def test_add_key_or_team_level_spend_logs_metadata_to_request( print(f"team_sl_metadata: {team_sl_metadata}") mock_request.url.path = "/chat/completions" + # Opt the key into client-supplied tags so request_tags are preserved + # and merged with admin-configured key/team tags. Without this flag, + # request_tags would be stripped by add_litellm_data_to_request. key_metadata = { "tags": key_tags, "spend_logs_metadata": key_sl_metadata, + "allow_client_tags": True, } team_metadata = { "tags": team_tags, @@ -859,12 +863,13 @@ async def test_add_litellm_data_to_request_duplicate_tags( mock_request.headers = {} mock_request.state = State() - # Setup key with tags in metadata + # Setup key with tags in metadata. Opt into client-supplied tags so the + # request_tags are preserved for the merge under test. user_api_key_dict = UserAPIKeyAuth( api_key="test_api_key", user_id="test_user_id", org_id="test_org_id", - metadata={"tags": key_tags}, + metadata={"tags": key_tags, "allow_client_tags": True}, ) # Setup request data with tags diff --git a/tests/test_litellm/integrations/test_custom_guardrail.py b/tests/test_litellm/integrations/test_custom_guardrail.py index 1aab3dbac93..0c904e9df50 100644 --- a/tests/test_litellm/integrations/test_custom_guardrail.py +++ b/tests/test_litellm/integrations/test_custom_guardrail.py @@ -173,17 +173,16 @@ class TestCustomGuardrailShouldRunGuardrail: assert result is False def test_should_run_guardrail_with_disable_global_guardrail(self): - """Test that disable_global_guardrail disables a global guardrail when set to True""" + """Test that disable_global_guardrails only works from admin metadata""" from litellm.types.guardrails import GuardrailEventHooks - # Create a guardrail with default_on=True (global guardrail) custom_guardrail = CustomGuardrail( guardrail_name="global_guardrail", default_on=True, event_hook=GuardrailEventHooks.pre_call, ) - # Test 1: Global guardrail runs by default when default_on=True + # Test 1: Global guardrail runs by default data = { "model": "gpt-3.5-turbo", "messages": [{"role": "user", "content": "test"}], @@ -193,7 +192,7 @@ class TestCustomGuardrailShouldRunGuardrail: ) assert result is True, "Global guardrail should run when default_on=True" - # Test 2: Global guardrail is disabled when disable_global_guardrail=True at root level + # Test 2: User-injected disable at root level is IGNORED data_with_disable_root = { "model": "gpt-3.5-turbo", "messages": [{"role": "user", "content": "test"}], @@ -203,23 +202,10 @@ class TestCustomGuardrailShouldRunGuardrail: data=data_with_disable_root, event_type=GuardrailEventHooks.pre_call ) assert ( - result is False - ), "Global guardrail should be disabled when disable_global_guardrail=True" + result is True + ), "User-injected disable_global_guardrails should be ignored" - # Test 3: Global guardrail is disabled when disable_global_guardrail=True in litellm_metadata - data_with_disable_litellm = { - "model": "gpt-3.5-turbo", - "messages": [{"role": "user", "content": "test"}], - "litellm_metadata": {"disable_global_guardrails": True}, - } - result = custom_guardrail.should_run_guardrail( - data=data_with_disable_litellm, event_type=GuardrailEventHooks.pre_call - ) - assert ( - result is False - ), "Global guardrail should be disabled when disable_global_guardrail=True in litellm_metadata" - - # Test 4: Global guardrail is disabled when disable_global_guardrail=True in metadata + # Test 3: User-injected disable in metadata is IGNORED data_with_disable_metadata = { "model": "gpt-3.5-turbo", "messages": [{"role": "user", "content": "test"}], @@ -228,25 +214,51 @@ class TestCustomGuardrailShouldRunGuardrail: result = custom_guardrail.should_run_guardrail( data=data_with_disable_metadata, event_type=GuardrailEventHooks.pre_call ) - assert ( - result is False - ), "Global guardrail should be disabled when disable_global_guardrail=True in metadata" + assert result is True, "User-injected metadata disable should be ignored" - # Test 5: Global guardrail runs when disable_global_guardrail=False - data_with_disable_false = { + # Test 4: Admin-configured disable via user_api_key_metadata IS respected + data_with_admin_disable = { "model": "gpt-3.5-turbo", "messages": [{"role": "user", "content": "test"}], - "disable_global_guardrails": False, + "metadata": {"user_api_key_metadata": {"disable_global_guardrails": True}}, } result = custom_guardrail.should_run_guardrail( - data=data_with_disable_false, event_type=GuardrailEventHooks.pre_call + data=data_with_admin_disable, event_type=GuardrailEventHooks.pre_call + ) + assert result is False, "Admin-configured disable should be respected" + + # Test 5: Admin config in metadata isn't shadowed by user-supplied litellm_metadata + data_cross_key = { + "model": "gpt-3.5-turbo", + "messages": [{"role": "user", "content": "test"}], + "metadata": {"user_api_key_metadata": {"disable_global_guardrails": True}}, + "litellm_metadata": {"request_tags": ["user-supplied"]}, + } + result = custom_guardrail.should_run_guardrail( + data=data_cross_key, event_type=GuardrailEventHooks.pre_call ) assert ( - result is True - ), "Global guardrail should still run when disable_global_guardrail=False" + result is False + ), "Admin config in metadata must not be shadowed by user-supplied litellm_metadata" + + # Test 6: After the pre-call strip runs, user-injected + # user_api_key_metadata in the non-authoritative metadata key is gone. + # _get_admin_metadata must then surface admin config unchanged. + data_post_strip = { + "model": "gpt-3.5-turbo", + "messages": [{"role": "user", "content": "test"}], + "metadata": {"user_api_key_metadata": {"disable_global_guardrails": True}}, + "litellm_metadata": {}, # post-strip: attacker payload removed + } + result = custom_guardrail.should_run_guardrail( + data=data_post_strip, event_type=GuardrailEventHooks.pre_call + ) + assert ( + result is False + ), "Admin config in metadata must be respected when other metadata key is empty" def test_should_run_guardrail_with_opted_out_global_guardrails(self): - """Test the per-guardrail opt-out list for global (default_on=True) guardrails""" + """Test that per-guardrail opt-out only works from admin metadata""" from litellm.types.guardrails import GuardrailEventHooks custom_guardrail = CustomGuardrail( @@ -255,7 +267,7 @@ class TestCustomGuardrailShouldRunGuardrail: event_hook=GuardrailEventHooks.pre_call, ) - # Test 1: guardrail in the opt-out list at root level → skipped + # Test 1: User-injected opt-out at root level is IGNORED data_root = { "model": "gpt-3.5-turbo", "messages": [{"role": "user", "content": "test"}], @@ -265,23 +277,10 @@ class TestCustomGuardrailShouldRunGuardrail: custom_guardrail.should_run_guardrail( data=data_root, event_type=GuardrailEventHooks.pre_call ) - is False + is True ) - # Test 2: guardrail in the opt-out list inside litellm_metadata → skipped - data_litellm = { - "model": "gpt-3.5-turbo", - "messages": [{"role": "user", "content": "test"}], - "litellm_metadata": {"opted_out_global_guardrails": ["global_guardrail"]}, - } - assert ( - custom_guardrail.should_run_guardrail( - data=data_litellm, event_type=GuardrailEventHooks.pre_call - ) - is False - ) - - # Test 3: guardrail in the opt-out list inside metadata → skipped + # Test 2: User-injected opt-out in metadata is IGNORED data_metadata = { "model": "gpt-3.5-turbo", "messages": [{"role": "user", "content": "test"}], @@ -291,7 +290,7 @@ class TestCustomGuardrailShouldRunGuardrail: custom_guardrail.should_run_guardrail( data=data_metadata, event_type=GuardrailEventHooks.pre_call ) - is False + is True ) # Test 4: a different guardrail in the opt-out list → still runs diff --git a/tests/test_litellm/litellm_core_utils/test_image_handling.py b/tests/test_litellm/litellm_core_utils/test_image_handling.py index e1fc2f775b8..cc13e816dde 100644 --- a/tests/test_litellm/litellm_core_utils/test_image_handling.py +++ b/tests/test_litellm/litellm_core_utils/test_image_handling.py @@ -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)) diff --git a/tests/test_litellm/litellm_core_utils/test_initialize_dynamic_callback_params.py b/tests/test_litellm/litellm_core_utils/test_initialize_dynamic_callback_params.py index 91969a2b8e2..55f3c2ba3aa 100644 --- a/tests/test_litellm/litellm_core_utils/test_initialize_dynamic_callback_params.py +++ b/tests/test_litellm/litellm_core_utils/test_initialize_dynamic_callback_params.py @@ -82,13 +82,18 @@ def test_env_reference_in_litellm_params_metadata_raises(): def test_non_string_values_are_not_flagged(): kwargs = { "langsmith_sampling_rate": 0.5, - "turn_off_message_logging": True, } params = initialize_standard_callback_dynamic_params(kwargs) assert params.get("langsmith_sampling_rate") == 0.5 - assert params.get("turn_off_message_logging") is True + + +def test_turn_off_message_logging_not_extracted_from_request(): + """turn_off_message_logging is admin-only — must not be settable via request.""" + kwargs = {"turn_off_message_logging": True} + params = initialize_standard_callback_dynamic_params(kwargs) + assert params.get("turn_off_message_logging") is None def test_empty_kwargs_returns_empty_params(): diff --git a/tests/test_litellm/litellm_core_utils/test_url_utils.py b/tests/test_litellm/litellm_core_utils/test_url_utils.py new file mode 100644 index 00000000000..4579c203218 --- /dev/null +++ b/tests/test_litellm/litellm_core_utils/test_url_utils.py @@ -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/") diff --git a/tests/test_litellm/proxy/auth/test_auth_checks.py b/tests/test_litellm/proxy/auth/test_auth_checks.py index 982784b07a3..0391e224084 100644 --- a/tests/test_litellm/proxy/auth/test_auth_checks.py +++ b/tests/test_litellm/proxy/auth/test_auth_checks.py @@ -1795,3 +1795,142 @@ async def test_team_member_budget_check_reads_from_spend_counter(): proxy_logging_obj=proxy_logging_obj, ) assert exc_info.value.current_cost == 1.5 + + +class TestGuardrailModificationCheck: + """Defense-in-depth: `_guardrail_modification_check` must 403 when the + caller's metadata attempts to modify any guardrail-related key and the + team lacks the `modify_guardrails` permission. Checks both the + historically-covered `guardrails` list and the bypass toggles that + `_get_admin_metadata` silently ignores at read time. + """ + + def _call(self, request_body): + from litellm.proxy.auth.auth_checks import _guardrail_modification_check + + team_object = MagicMock() + team_object.metadata = {} # no permission + return _guardrail_modification_check( + request_body=request_body, team_object=team_object + ) + + def test_noop_when_no_guardrail_keys_present(self): + # no-op — should return silently + self._call({"metadata": {"unrelated": "value"}}) + + def test_rejects_guardrails_list(self): + from fastapi import HTTPException + + with patch( + "litellm.proxy.guardrails.guardrail_helpers.can_modify_guardrails", + return_value=False, + ): + with pytest.raises(HTTPException) as exc: + self._call({"metadata": {"guardrails": ["custom"]}}) + assert exc.value.status_code == 403 + + def test_rejects_disable_global_guardrails_plural(self): + from fastapi import HTTPException + + with patch( + "litellm.proxy.guardrails.guardrail_helpers.can_modify_guardrails", + return_value=False, + ): + with pytest.raises(HTTPException) as exc: + self._call({"metadata": {"disable_global_guardrails": True}}) + assert exc.value.status_code == 403 + + def test_rejects_disable_global_guardrail_singular(self): + """VERIA-28's originally-reported singular-key typo variant.""" + from fastapi import HTTPException + + with patch( + "litellm.proxy.guardrails.guardrail_helpers.can_modify_guardrails", + return_value=False, + ): + with pytest.raises(HTTPException) as exc: + self._call({"metadata": {"disable_global_guardrail": True}}) + assert exc.value.status_code == 403 + + def test_rejects_opted_out_global_guardrails(self): + from fastapi import HTTPException + + with patch( + "litellm.proxy.guardrails.guardrail_helpers.can_modify_guardrails", + return_value=False, + ): + with pytest.raises(HTTPException) as exc: + self._call( + {"metadata": {"opted_out_global_guardrails": ["some_guardrail"]}} + ) + assert exc.value.status_code == 403 + + def test_rejects_injection_via_litellm_metadata_key(self): + """Caller can populate the OTHER metadata key; that must also 403.""" + from fastapi import HTTPException + + with patch( + "litellm.proxy.guardrails.guardrail_helpers.can_modify_guardrails", + return_value=False, + ): + with pytest.raises(HTTPException) as exc: + self._call({"litellm_metadata": {"disable_global_guardrails": True}}) + assert exc.value.status_code == 403 + + def test_rejects_root_level_injection(self): + """Top-level injection (`request_body["disable_global_guardrails"]`) + was VERIA-28's easiest variant to hit — keep it rejected.""" + from fastapi import HTTPException + + with patch( + "litellm.proxy.guardrails.guardrail_helpers.can_modify_guardrails", + return_value=False, + ): + with pytest.raises(HTTPException) as exc: + self._call({"disable_global_guardrails": True}) + assert exc.value.status_code == 403 + + def test_allows_when_team_has_permission(self): + with patch( + "litellm.proxy.guardrails.guardrail_helpers.can_modify_guardrails", + return_value=True, + ): + # no-op, should not raise + self._call({"metadata": {"disable_global_guardrails": True}}) + + def test_rejects_string_encoded_metadata_bypass(self): + """Regression: attacker sends metadata as JSON string to bypass the + isinstance(dict) guard. The check must coerce the string to dict + and evaluate guardrail modification keys inside it.""" + import json as _json + + from fastapi import HTTPException + + attacker_payload = {"disable_global_guardrails": True} + with patch( + "litellm.proxy.guardrails.guardrail_helpers.can_modify_guardrails", + return_value=False, + ): + with pytest.raises(HTTPException) as exc: + self._call({"metadata": _json.dumps(attacker_payload)}) + assert exc.value.status_code == 403 + + def test_rejects_string_encoded_litellm_metadata_bypass(self): + """Same bypass via the litellm_metadata key.""" + import json as _json + + from fastapi import HTTPException + + attacker_payload = {"guardrails": ["evaded"]} + with patch( + "litellm.proxy.guardrails.guardrail_helpers.can_modify_guardrails", + return_value=False, + ): + with pytest.raises(HTTPException) as exc: + self._call({"litellm_metadata": _json.dumps(attacker_payload)}) + assert exc.value.status_code == 403 + + def test_noop_when_string_is_not_json_object(self): + """Unparseable strings should not trigger a 403 — they have no keys.""" + self._call({"metadata": "not-json"}) + self._call({"metadata": '"just a string"'}) diff --git a/tests/test_litellm/proxy/auth/test_route_checks.py b/tests/test_litellm/proxy/auth/test_route_checks.py index 6703ed678ef..bca6b9e78d9 100644 --- a/tests/test_litellm/proxy/auth/test_route_checks.py +++ b/tests/test_litellm/proxy/auth/test_route_checks.py @@ -1391,6 +1391,38 @@ def test_non_org_admin_with_organizations_list(): assert _user_is_org_admin({"organizations": ["org-1"]}, user_obj) is False +def test_org_admin_cannot_escalate_to_other_org(): + """Regression: admin of org-A requesting [org-A, org-B] must be rejected.""" + user_obj = _make_org_admin_user("org-A") + assert _user_is_org_admin({"organizations": ["org-A", "org-B"]}, user_obj) is False + + +def test_org_admin_of_multiple_orgs_can_operate_on_both(): + """Admin of both org-A and org-B can operate on both.""" + memberships = [ + LiteLLM_OrganizationMembershipTable( + user_id="multi-admin", + organization_id="org-A", + user_role=LitellmUserRoles.ORG_ADMIN.value, + created_at=datetime(2024, 1, 1), + updated_at=datetime(2024, 1, 1), + ), + LiteLLM_OrganizationMembershipTable( + user_id="multi-admin", + organization_id="org-B", + user_role=LitellmUserRoles.ORG_ADMIN.value, + created_at=datetime(2024, 1, 1), + updated_at=datetime(2024, 1, 1), + ), + ] + user_obj = LiteLLM_UserTable( + user_id="multi-admin", + user_role=LitellmUserRoles.INTERNAL_USER.value, + organization_memberships=memberships, + ) + assert _user_is_org_admin({"organizations": ["org-A", "org-B"]}, user_obj) is True + + @pytest.mark.asyncio async def test_initialize_pass_through_registers_wildcard_for_auth_subpath(): """ diff --git a/tests/test_litellm/proxy/common_utils/test_http_parsing_utils.py b/tests/test_litellm/proxy/common_utils/test_http_parsing_utils.py index 0370e465627..b4343f6b2e1 100644 --- a/tests/test_litellm/proxy/common_utils/test_http_parsing_utils.py +++ b/tests/test_litellm/proxy/common_utils/test_http_parsing_utils.py @@ -815,3 +815,41 @@ def test_safe_get_request_headers_state_unavailable(): result = _safe_get_request_headers(mock_request) assert result == {"content-type": "application/json"} + + +class TestGetTagsFromRequestBodyStringCoerce: + """Regression: the auth-time tag helper used `metadata.get("tags", ...)` + directly, which raised AttributeError when metadata arrived as a JSON + string (multipart/form-data or extra_body). That turned into a DoS at + auth time and potentially bypassed tag-based RBAC if the caller caught + the exception and fell through with empty tags. + """ + + def test_json_string_metadata_is_coerced_to_dict(self): + from litellm.proxy.common_utils.http_parsing_utils import ( + get_tags_from_request_body, + ) + + metadata_json = json.dumps({"tags": ["a", "b"]}) + # Must not raise + tags = get_tags_from_request_body({"metadata": metadata_json}) + assert tags == ["a", "b"] + + def test_unparseable_string_metadata_is_ignored(self): + from litellm.proxy.common_utils.http_parsing_utils import ( + get_tags_from_request_body, + ) + + # Must not raise; must yield no metadata tags but keep root tags + tags = get_tags_from_request_body( + {"metadata": "not-json", "tags": ["root-only"]} + ) + assert tags == ["root-only"] + + def test_dict_metadata_still_works(self): + from litellm.proxy.common_utils.http_parsing_utils import ( + get_tags_from_request_body, + ) + + tags = get_tags_from_request_body({"metadata": {"tags": ["x"]}}) + assert tags == ["x"] diff --git a/tests/test_litellm/proxy/db/test_create_views.py b/tests/test_litellm/proxy/db/test_create_views.py index 1a90b4c204d..c0c09d0137b 100644 --- a/tests/test_litellm/proxy/db/test_create_views.py +++ b/tests/test_litellm/proxy/db/test_create_views.py @@ -130,6 +130,42 @@ async def test_create_views_reraises_undefined_function_error(): mock_db.execute_raw.assert_not_called() +@pytest.mark.asyncio +async def test_should_create_missing_views_reltuples_zero(): + """should return True when reltuples is 0 (fresh empty table).""" + from litellm.proxy.db.create_views import should_create_missing_views + + mock_db = MagicMock() + mock_db.query_raw = AsyncMock(return_value=[{"reltuples": 0}]) + + result = await should_create_missing_views(mock_db) + assert result is True + + +@pytest.mark.asyncio +async def test_should_create_missing_views_reltuples_negative_one(): + """should return True when reltuples is -1 (table created, no ANALYZE yet).""" + from litellm.proxy.db.create_views import should_create_missing_views + + mock_db = MagicMock() + mock_db.query_raw = AsyncMock(return_value=[{"reltuples": -1}]) + + result = await should_create_missing_views(mock_db) + assert result is True + + +@pytest.mark.asyncio +async def test_should_create_missing_views_reltuples_positive(): + """should return False when reltuples > 0 (table has data).""" + from litellm.proxy.db.create_views import should_create_missing_views + + mock_db = MagicMock() + mock_db.query_raw = AsyncMock(return_value=[{"reltuples": 1000}]) + + result = await should_create_missing_views(mock_db) + assert result is False + + @pytest.mark.asyncio async def test_create_views_creates_view_on_undefined_table_error(): """should treat 'undefined table' as a missing-view signal and attempt creation.""" diff --git a/tests/test_litellm/proxy/management_endpoints/test_internal_user_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_internal_user_endpoints.py index 0f90d236aed..a1ba7ecd677 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_internal_user_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_internal_user_endpoints.py @@ -1878,6 +1878,121 @@ async def test_delete_user_cleans_up_created_by_invitation_links(mocker): assert condition[field] == {"in": ["admin-creator"]} +@pytest.mark.asyncio +async def test_delete_user_rejects_org_admin_deleting_outside_scope(mocker): + """Regression: an org admin of org-A must not be able to delete a user + whose org memberships include org-B. + + Route-level gate accepts the request when the caller supplies an + `organization_id` they administer; without per-user org authorization + the handler would cascade-delete the victim's keys, memberships, and + user row regardless of where the victim actually belongs. + """ + from fastapi import HTTPException + + from litellm.proxy._types import DeleteUserRequest, UserAPIKeyAuth + from litellm.proxy.management_endpoints.internal_user_endpoints import delete_user + + mock_prisma_client = mocker.MagicMock() + + # Target user exists and is a member of org-B only. + mock_target_user = mocker.MagicMock() + mock_target_user.user_id = "victim" + mock_target_user.user_email = "victim@example.com" + mock_target_user.teams = [] + mock_target_user.json.return_value = "{}" + + async def mock_find_unique(*args, **kwargs): + return mock_target_user + + mock_prisma_client.db.litellm_usertable.find_unique = mocker.AsyncMock( + side_effect=mock_find_unique + ) + + # Caller (org_admin_user) administers org-A. + caller_membership = mocker.MagicMock() + caller_membership.organization_id = "org-A" + + # Target user is a member of org-B (outside caller's scope). + target_membership = mocker.MagicMock() + target_membership.organization_id = "org-B" + + async def mock_find_memberships(*args, **kwargs): + where = kwargs.get("where") or (args[0] if args else {}) + user_id_filter = where.get("user_id") + # Batched lookup: {"user_id": {"in": [...]}} returns target memberships. + # Caller role lookup: {"user_id": "", "user_role": ...}. + if isinstance(user_id_filter, dict) and "in" in user_id_filter: + if "victim" in user_id_filter["in"]: + # Attach user_id on the mock so the caller can build its + # per-user dict from the batch result. + target_membership.user_id = "victim" + return [target_membership] + return [] + if user_id_filter == "org_admin_user": + return [caller_membership] + return [] + + mock_prisma_client.db.litellm_organizationmembership.find_many = mocker.AsyncMock( + side_effect=mock_find_memberships + ) + + mocker.patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) + + data = DeleteUserRequest(user_ids=["victim"]) + user_api_key_dict = UserAPIKeyAuth( + user_id="org_admin_user", user_role=LitellmUserRoles.ORG_ADMIN + ) + + with pytest.raises(HTTPException) as exc: + await delete_user(data=data, user_api_key_dict=user_api_key_dict) + assert exc.value.status_code == 403 + + # Critical: no delete_many calls should have executed. + assert not hasattr( + mock_prisma_client.db.litellm_verificationtoken.delete_many, "mock_calls" + ) or len( + mock_prisma_client.db.litellm_verificationtoken.delete_many.mock_calls + ) == 0 + + +@pytest.mark.asyncio +async def test_user_update_rejects_silent_create_for_non_proxy_admin(mocker): + """Regression: `/user/update` with an unknown user_email used to fall + through to an INSERT, silently creating a new user with caller-supplied + budget, models, and metadata. An org admin could use this to spawn + arbitrary users outside the /user/new authorization flow.""" + from fastapi import HTTPException + + from litellm.proxy._types import UpdateUserRequest, UserAPIKeyAuth + from litellm.proxy.management_endpoints.internal_user_endpoints import ( + _update_single_user_helper, + ) + + mock_prisma_client = mocker.MagicMock() + # user_email lookup yields None → would silently create pre-fix. + mock_prisma_client.db.litellm_usertable.find_first = mocker.AsyncMock( + return_value=None + ) + mocker.patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) + + user_request = UpdateUserRequest( + user_email="newcomer@example.com", + max_budget=1_000_000, + models=["gpt-4"], + ) + org_admin = UserAPIKeyAuth( + user_id="org-admin", + user_role=LitellmUserRoles.ORG_ADMIN, + ) + + with pytest.raises(HTTPException) as exc: + await _update_single_user_helper( + user_request=user_request, user_api_key_dict=org_admin + ) + assert exc.value.status_code == 404 + + # ===================================================================== # /v2/user/info endpoint tests # ===================================================================== diff --git a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py index 6f56b5ca18b..8e6fd78a6f6 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py @@ -8076,6 +8076,277 @@ async def test_update_key_non_budget_fields_allowed_for_internal_user(monkeypatc assert result is not None +@pytest.mark.asyncio +async def test_update_key_non_budget_rejects_cross_user_modification(monkeypatch): + """Regression: previously _check_key_admin_access was gated on + max_budget/spend changes only, so an internal user could rewrite any + OTHER field (alias, models, tpm_limit, blocked, metadata, …) on any + key they weren't admin of as long as they avoided budget/spend. This + confirms that a non-admin user updating a key that belongs to another + user fails with 403 even for non-budget fields.""" + from litellm.proxy.management_endpoints.key_management_endpoints import ( + update_key_fn, + ) + + mock_prisma_client = AsyncMock() + test_hashed_token = ( + "cafebabe" * 8 + ) + + mock_existing_key = MagicMock() + mock_existing_key.token = test_hashed_token + mock_existing_key.user_id = "victim_user" # owned by someone else + mock_existing_key.team_id = None + mock_existing_key.project_id = None + mock_existing_key.max_budget = 10.0 + mock_existing_key.key_alias = "original" + mock_existing_key.models = [] + mock_existing_key.model_dump.return_value = { + "token": test_hashed_token, + "user_id": "victim_user", + "team_id": None, + "max_budget": 10.0, + } + + mock_prisma_client.get_data = AsyncMock(return_value=mock_existing_key) + mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock( + return_value=mock_existing_key + ) + + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) + monkeypatch.setattr("litellm.proxy.proxy_server.user_api_key_cache", AsyncMock()) + monkeypatch.setattr("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock()) + monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", None) + monkeypatch.setattr("litellm.proxy.proxy_server.premium_user", True) + monkeypatch.setattr( + "litellm.proxy.proxy_server.hash_token", lambda t: test_hashed_token + ) + + mock_request = MagicMock() + mock_request.query_params = {} + attacker = UserAPIKeyAuth( + user_role=LitellmUserRoles.INTERNAL_USER, + api_key="sk-attacker", + user_id="attacker_user", # NOT the owner + ) + + # Trying to blanket-rewrite a non-budget field on someone else's key + # must now fail. + with pytest.raises(ProxyException) as exc: + await update_key_fn( + request=mock_request, + data=UpdateKeyRequest( + key=test_hashed_token, key_alias="pwned", blocked=True + ), + user_api_key_dict=attacker, + litellm_changed_by=None, + ) + assert str(exc.value.code) == "403" + + +@pytest.mark.asyncio +async def test_update_key_team_member_with_permission_can_update_non_budget( + monkeypatch, +): + """A team member whose team grants /key/update in member_permissions can + update non-budget fields on a team key even though they are not a team + admin. Regression: the cross-key admin check was over-broad and rejected + this documented path.""" + from litellm.proxy.management_endpoints.key_management_endpoints import ( + update_key_fn, + ) + + test_hashed_token = "deadbeef" * 8 + team_id = "team-with-update-grant" + member_user_id = "team-member-user" + + mock_existing_key = MagicMock() + mock_existing_key.token = test_hashed_token + mock_existing_key.user_id = None # team-scoped key (no owning user) + mock_existing_key.team_id = team_id + mock_existing_key.project_id = None + mock_existing_key.max_budget = 10.0 + mock_existing_key.key_alias = "original" + mock_existing_key.models = [] + mock_existing_key.model_dump.return_value = { + "token": test_hashed_token, + "user_id": None, + "team_id": team_id, + "max_budget": 10.0, + } + + team_table = LiteLLM_TeamTableCachedObj( + team_id=team_id, + team_alias="test-team", + tpm_limit=None, + rpm_limit=None, + max_budget=None, + spend=0.0, + models=[], + blocked=False, + members_with_roles=[ + Member(user_id="some-team-admin", role="admin"), + Member(user_id=member_user_id, role="user"), + ], + team_member_permissions=["/key/update", "/key/info"], + ) + + mock_updated_key = MagicMock() + mock_updated_key.token = test_hashed_token + mock_updated_key.key_alias = "renamed-by-member" + + mock_prisma_client = AsyncMock() + mock_prisma_client.get_data = AsyncMock(return_value=mock_existing_key) + mock_prisma_client.update_data = AsyncMock(return_value=mock_updated_key) + mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock( + return_value=mock_existing_key + ) + + async def mock_get_team_object(*args, **kwargs): + return team_table + + async def mock_enforce_unique_key_alias(**kwargs): + pass + + async def mock_delete_cache_key_object(**kwargs): + pass + + monkeypatch.setattr( + "litellm.proxy.management_endpoints.key_management_endpoints.get_team_object", + mock_get_team_object, + ) + monkeypatch.setattr( + "litellm.proxy.management_helpers.team_member_permission_checks.get_team_object", + mock_get_team_object, + ) + monkeypatch.setattr( + "litellm.proxy.management_endpoints.key_management_endpoints._enforce_unique_key_alias", + mock_enforce_unique_key_alias, + ) + monkeypatch.setattr( + "litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object", + mock_delete_cache_key_object, + ) + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) + monkeypatch.setattr("litellm.proxy.proxy_server.user_api_key_cache", AsyncMock()) + monkeypatch.setattr("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock()) + monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", None) + monkeypatch.setattr("litellm.proxy.proxy_server.premium_user", True) + monkeypatch.setattr("litellm.store_audit_logs", False) + monkeypatch.setattr( + "litellm.proxy.proxy_server.hash_token", lambda t: test_hashed_token + ) + + mock_request = MagicMock() + mock_request.query_params = {} + team_member = UserAPIKeyAuth( + user_role=LitellmUserRoles.INTERNAL_USER, + api_key="sk-team-member", + user_id=member_user_id, + team_id=team_id, + ) + + # Non-budget update on a team key by a team member with /key/update + # permission should succeed. + result = await update_key_fn( + request=mock_request, + data=UpdateKeyRequest(key=test_hashed_token, key_alias="renamed-by-member"), + user_api_key_dict=team_member, + litellm_changed_by=None, + ) + + assert result is not None + + +@pytest.mark.asyncio +async def test_update_key_team_member_cannot_change_budget(monkeypatch): + """A team member with /key/update in member_permissions still cannot + change max_budget — budget/spend changes require team/org admin. The + member_permissions bypass only applies to non-budget fields.""" + from litellm.proxy.management_endpoints.key_management_endpoints import ( + update_key_fn, + ) + + test_hashed_token = "feedface" * 8 + team_id = "team-with-update-grant" + member_user_id = "team-member-user" + + mock_existing_key = MagicMock() + mock_existing_key.token = test_hashed_token + mock_existing_key.user_id = None # team-scoped key (no owning user) + mock_existing_key.team_id = team_id + mock_existing_key.project_id = None + mock_existing_key.max_budget = 10.0 + mock_existing_key.key_alias = "original" + mock_existing_key.models = [] + mock_existing_key.model_dump.return_value = { + "token": test_hashed_token, + "user_id": None, + "team_id": team_id, + "max_budget": 10.0, + } + + team_table = LiteLLM_TeamTableCachedObj( + team_id=team_id, + team_alias="test-team", + tpm_limit=None, + rpm_limit=None, + max_budget=None, + spend=0.0, + models=[], + blocked=False, + members_with_roles=[ + Member(user_id="some-team-admin", role="admin"), + Member(user_id=member_user_id, role="user"), + ], + team_member_permissions=["/key/update", "/key/info"], + ) + + mock_prisma_client = AsyncMock() + mock_prisma_client.get_data = AsyncMock(return_value=mock_existing_key) + mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock( + return_value=mock_existing_key + ) + + async def mock_get_team_object(*args, **kwargs): + return team_table + + monkeypatch.setattr( + "litellm.proxy.management_endpoints.key_management_endpoints.get_team_object", + mock_get_team_object, + ) + monkeypatch.setattr( + "litellm.proxy.management_helpers.team_member_permission_checks.get_team_object", + mock_get_team_object, + ) + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) + monkeypatch.setattr("litellm.proxy.proxy_server.user_api_key_cache", AsyncMock()) + monkeypatch.setattr("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock()) + monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", None) + monkeypatch.setattr("litellm.proxy.proxy_server.premium_user", True) + monkeypatch.setattr( + "litellm.proxy.proxy_server.hash_token", lambda t: test_hashed_token + ) + + mock_request = MagicMock() + mock_request.query_params = {} + team_member = UserAPIKeyAuth( + user_role=LitellmUserRoles.INTERNAL_USER, + api_key="sk-team-member", + user_id=member_user_id, + team_id=team_id, + ) + + with pytest.raises(ProxyException) as exc: + await update_key_fn( + request=mock_request, + data=UpdateKeyRequest(key=test_hashed_token, max_budget=500.0), + user_api_key_dict=team_member, + litellm_changed_by=None, + ) + assert str(exc.value.code) == "403" + + # ============================================================================ # LIT-1884: Internal users cannot create invalid keys # ============================================================================ diff --git a/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py index f1e077b5e4a..8da0ef19f81 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py @@ -2742,10 +2742,10 @@ async def test_list_team_v2_org_admin_with_user_id_returns_user_teams(): assert result["total"] == 1 - # Verify the where clause filters by user's teams, not org scope + # Verify the where clause filters by user's teams AND org scope where = mock_db.litellm_teamtable.find_many.call_args.kwargs["where"] assert where["team_id"] == {"in": ["team_X", "team_Y"]} - assert "organization_id" not in where + assert where["organization_id"] == {"in": ["org_A"]} @pytest.mark.asyncio @@ -5134,6 +5134,16 @@ async def test_update_team_guardrails_with_org_id(): return_value=mock_org ) + # Destination-org guard in update_team queries for the caller's + # ORG_ADMIN membership on the destination org. Return a match so + # the guardrails-update path (the subject under test) proceeds. + mock_org_admin_membership = MagicMock() + mock_org_admin_membership.user_id = "org-admin-guardrails-test" + mock_org_admin_membership.organization_id = "test-org-guardrails" + mock_prisma.db.litellm_organizationmembership.find_many = AsyncMock( + return_value=[mock_org_admin_membership] + ) + # Mock team update mock_updated_team = MagicMock(spec=LiteLLM_TeamTable) mock_updated_team.team_id = "team-guardrails-123" diff --git a/tests/test_litellm/proxy/openai_files_endpoint/test_files_endpoint.py b/tests/test_litellm/proxy/openai_files_endpoint/test_files_endpoint.py index f7cdb95bd54..5fc36b71f2b 100644 --- a/tests/test_litellm/proxy/openai_files_endpoint/test_files_endpoint.py +++ b/tests/test_litellm/proxy/openai_files_endpoint/test_files_endpoint.py @@ -1295,7 +1295,6 @@ def test_create_file_with_nested_litellm_metadata( "target_model_names": "gpt-3.5-turbo", "litellm_metadata[spend_logs_metadata][owner]": "john_doe", "litellm_metadata[spend_logs_metadata][team]": "engineering", - "litellm_metadata[tags]": "production", "litellm_metadata[environment]": "prod", }, headers={"Authorization": "Bearer test-key"}, @@ -1306,11 +1305,12 @@ def test_create_file_with_nested_litellm_metadata( result = response.json() assert result["id"] == "file-test-123" - # Verify nested metadata was correctly parsed + # Verify nested metadata was correctly parsed. + # Note: caller-supplied `tags` is stripped by default; test removed + # to keep the parsing test focused on parser correctness. assert "spend_logs_metadata" in captured_litellm_metadata assert captured_litellm_metadata["spend_logs_metadata"]["owner"] == "john_doe" assert captured_litellm_metadata["spend_logs_metadata"]["team"] == "engineering" - assert captured_litellm_metadata["tags"] == "production" assert captured_litellm_metadata["environment"] == "prod" diff --git a/tests/test_litellm/proxy/test_litellm_pre_call_utils.py b/tests/test_litellm/proxy/test_litellm_pre_call_utils.py index 545adedfce8..34d3c203377 100644 --- a/tests/test_litellm/proxy/test_litellm_pre_call_utils.py +++ b/tests/test_litellm/proxy/test_litellm_pre_call_utils.py @@ -209,6 +209,580 @@ async def test_add_litellm_data_to_request_parses_string_metadata(): assert updated_data["metadata"]["generation_name"] == "gen123" +@pytest.mark.asyncio +async def test_add_litellm_data_to_request_strips_admin_injection_slots(): + """User-supplied user_api_key_metadata / user_api_key_team_metadata / + _pipeline_managed_guardrails must be stripped from both metadata keys + before the proxy writes its own admin-populated values. Otherwise a + caller can shadow admin config via the non-`_metadata_variable_name` + metadata key (e.g. litellm_metadata while the proxy writes to metadata). + """ + from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request + + request_mock = MagicMock(spec=Request) + request_mock.url.path = "/v1/chat/completions" + request_mock.url = MagicMock() + request_mock.url.__str__.return_value = "http://localhost/v1/chat/completions" + request_mock.method = "POST" + request_mock.query_params = {} + request_mock.headers = {"Content-Type": "application/json"} + request_mock.client = MagicMock() + request_mock.client.host = "127.0.0.1" + + # Caller tries to inject admin config into BOTH metadata keys + attacker_admin_payload = {"disable_global_guardrails": True} + data = { + "model": "gpt-3.5-turbo", + "metadata": { + "user_api_key_metadata": attacker_admin_payload, + "user_api_key_team_metadata": attacker_admin_payload, + "_pipeline_managed_guardrails": ["evaded"], + }, + "litellm_metadata": { + "user_api_key_metadata": attacker_admin_payload, + "user_api_key_team_metadata": attacker_admin_payload, + "_pipeline_managed_guardrails": ["evaded"], + }, + } + + real_admin_metadata = {"admin_flag": "from_proxy"} + user_api_key_dict = UserAPIKeyAuth( + api_key="hashed-key", + metadata=real_admin_metadata, + team_metadata=real_admin_metadata, + spend=0.0, + max_budget=100.0, + model_max_budget={}, + team_spend=0.0, + team_max_budget=200.0, + ) + + updated = await add_litellm_data_to_request( + data=data, + request=request_mock, + user_api_key_dict=user_api_key_dict, + proxy_config=MagicMock(), + general_settings={}, + version="test-version", + ) + + # The key that matches `_metadata_variable_name` gets proxy-populated + # with the real admin payload; the OTHER key must not retain the + # attacker's injection. + populated = updated["metadata"] + assert populated["user_api_key_metadata"] == real_admin_metadata + assert populated["user_api_key_team_metadata"] == real_admin_metadata + assert "_pipeline_managed_guardrails" not in populated or populated[ + "_pipeline_managed_guardrails" + ] != ["evaded"] + + other = updated.get("litellm_metadata") or {} + assert other.get("user_api_key_metadata") in (None, {}, real_admin_metadata) + assert other.get("user_api_key_team_metadata") in (None, {}, real_admin_metadata) + assert "_pipeline_managed_guardrails" not in other + + +@pytest.mark.asyncio +async def test_add_litellm_data_to_request_strips_all_user_api_key_prefix_keys(): + """Strip must cover the full user_api_key_* family, not a hand-maintained + list of 2-3 names. Proxy writes a dozen such fields (user_id, alias, + spend, team_id, request_route, …) and an attacker populating any of them + in the non-authoritative metadata key would otherwise forge identity / + spend in audit logs and guardrails.""" + from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request + + request_mock = MagicMock(spec=Request) + request_mock.url.path = "/v1/chat/completions" + request_mock.url = MagicMock() + request_mock.url.__str__.return_value = "http://localhost/v1/chat/completions" + request_mock.method = "POST" + request_mock.query_params = {} + request_mock.headers = {"Content-Type": "application/json"} + request_mock.client = MagicMock() + request_mock.client.host = "127.0.0.1" + + attacker_injected = { + "user_api_key_user_id": "victim", + "user_api_key_alias": "admin-key", + "user_api_key_spend": 0.0, + "user_api_key_team_id": "victim-team", + "user_api_key_end_user_id": "victim-user", + "user_api_key_request_route": "/fake/route", + "user_api_key_hash": "fake-hash", + } + data = { + "model": "gpt-3.5-turbo", + "metadata": {**attacker_injected}, + "litellm_metadata": {**attacker_injected}, + } + + user_api_key_dict = UserAPIKeyAuth( + api_key="hashed-key", + user_id="real-user", + metadata={}, + team_metadata={}, + spend=42.0, + max_budget=100.0, + model_max_budget={}, + team_spend=0.0, + team_max_budget=200.0, + ) + + updated = await add_litellm_data_to_request( + data=data, + request=request_mock, + user_api_key_dict=user_api_key_dict, + proxy_config=MagicMock(), + general_settings={}, + version="test-version", + ) + + # The non-authoritative metadata dict must not retain ANY attacker-injected + # user_api_key_* key. + other = updated.get("litellm_metadata") or {} + attacker_leaks = [k for k in other if k.startswith("user_api_key_")] + assert attacker_leaks == [], f"Unexpected leaked keys: {attacker_leaks}" + + +@pytest.mark.asyncio +async def test_add_litellm_data_to_request_string_metadata_does_not_crash(): + """Regression: pre-strip code that pre-populated data['metadata'][k]=v + before the string-to-dict parse would crash on JSON-string metadata. + The snapshot / strip / admin-population pipeline must survive metadata + arriving as a string.""" + import json as _json + + from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request + + request_mock = MagicMock(spec=Request) + request_mock.url.path = "/v1/chat/completions" + request_mock.url = MagicMock() + request_mock.url.__str__.return_value = "http://localhost/v1/chat/completions" + request_mock.method = "POST" + request_mock.query_params = {} + request_mock.headers = {"Content-Type": "multipart/form-data"} + request_mock.client = MagicMock() + request_mock.client.host = "127.0.0.1" + + data = { + "model": "gpt-3.5-turbo", + "metadata": _json.dumps({"generation_name": "test"}), + } + + user_api_key_dict = UserAPIKeyAuth( + api_key="hashed-key", + metadata={}, + team_metadata={}, + spend=0.0, + max_budget=100.0, + model_max_budget={}, + team_spend=0.0, + team_max_budget=200.0, + ) + + # Must not raise TypeError / AttributeError. + updated = await add_litellm_data_to_request( + data=data, + request=request_mock, + user_api_key_dict=user_api_key_dict, + proxy_config=MagicMock(), + general_settings={}, + version="test-version", + ) + + # The parsed metadata should be a dict and the proxy snapshot body + # should have been taken AFTER the strip (so no leaked user_api_key_* + # from a raw string snapshot). + assert isinstance(updated["metadata"], dict) + assert updated["metadata"].get("generation_name") == "test" + + +@pytest.mark.asyncio +async def test_add_litellm_data_to_request_proxy_server_request_body_is_post_strip(): + """Regression: proxy_server_request['body'] used to be snapshotted before + the admin-slot strip, so standard_logging_object and spend-tracking + readers saw attacker-injected payload. Snapshot must now be post-strip.""" + from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request + + request_mock = MagicMock(spec=Request) + request_mock.url.path = "/v1/chat/completions" + request_mock.url = MagicMock() + request_mock.url.__str__.return_value = "http://localhost/v1/chat/completions" + request_mock.method = "POST" + request_mock.query_params = {} + request_mock.headers = {"Content-Type": "application/json"} + request_mock.client = MagicMock() + request_mock.client.host = "127.0.0.1" + + data = { + "model": "gpt-3.5-turbo", + "metadata": {"user_api_key_user_id": "victim"}, + } + + user_api_key_dict = UserAPIKeyAuth( + api_key="hashed-key", + user_id="real-user", + metadata={}, + team_metadata={}, + spend=0.0, + max_budget=100.0, + model_max_budget={}, + team_spend=0.0, + team_max_budget=200.0, + ) + + updated = await add_litellm_data_to_request( + data=data, + request=request_mock, + user_api_key_dict=user_api_key_dict, + proxy_config=MagicMock(), + general_settings={}, + version="test-version", + ) + + snapshot_body = updated["proxy_server_request"]["body"] + assert snapshot_body is not None + snapshot_metadata = snapshot_body.get("metadata") or {} + assert "user_api_key_user_id" not in snapshot_metadata or ( + snapshot_metadata["user_api_key_user_id"] != "victim" + ) + + +@pytest.mark.asyncio +async def test_add_litellm_data_to_request_strips_string_encoded_admin_injection(): + """Regression: metadata arriving as a JSON string (multipart/form-data or + extra_body) must not bypass the admin-injection strip. The parse happens + AFTER receipt, so the strip has to run after the parse, not before. + """ + from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request + + request_mock = MagicMock(spec=Request) + request_mock.url.path = "/v1/chat/completions" + request_mock.url = MagicMock() + request_mock.url.__str__.return_value = "http://localhost/v1/chat/completions" + request_mock.method = "POST" + request_mock.query_params = {} + request_mock.headers = {"Content-Type": "multipart/form-data"} + request_mock.client = MagicMock() + request_mock.client.host = "127.0.0.1" + + # Attacker encodes an admin-injection payload inside a JSON string. + attacker_payload = { + "user_api_key_metadata": {"disable_global_guardrails": True}, + "user_api_key_team_metadata": {"disable_global_guardrails": True}, + "_pipeline_managed_guardrails": ["evaded"], + } + data = { + "model": "gpt-3.5-turbo", + "metadata": json.dumps(attacker_payload), + "litellm_metadata": json.dumps(attacker_payload), + } + + real_admin_metadata = {"admin_flag": "from_proxy"} + user_api_key_dict = UserAPIKeyAuth( + api_key="hashed-key", + metadata=real_admin_metadata, + team_metadata=real_admin_metadata, + spend=0.0, + max_budget=100.0, + model_max_budget={}, + team_spend=0.0, + team_max_budget=200.0, + ) + + updated = await add_litellm_data_to_request( + data=data, + request=request_mock, + user_api_key_dict=user_api_key_dict, + proxy_config=MagicMock(), + general_settings={}, + version="test-version", + ) + + populated = updated["metadata"] + # The real admin payload from user_api_key_dict wins. + assert populated["user_api_key_metadata"] == real_admin_metadata + assert populated["user_api_key_team_metadata"] == real_admin_metadata + assert populated.get("_pipeline_managed_guardrails") != ["evaded"] + + other = updated.get("litellm_metadata") or {} + # After the strip, litellm_metadata has no admin-injection slots. + assert "user_api_key_metadata" not in other + assert "user_api_key_team_metadata" not in other + assert "_pipeline_managed_guardrails" not in other + + +@pytest.mark.asyncio +async def test_add_litellm_data_to_request_ignores_x_litellm_tags_header_without_permission(): + """Regression: the `x-litellm-tags` header bypassed the body-metadata + tag strip. Header tags must also be gated by `allow_client_tags`.""" + from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request + + request_mock = MagicMock(spec=Request) + request_mock.url.path = "/v1/chat/completions" + request_mock.url = MagicMock() + request_mock.url.__str__.return_value = "http://localhost/v1/chat/completions" + request_mock.method = "POST" + request_mock.query_params = {} + request_mock.headers = { + "Content-Type": "application/json", + "x-litellm-tags": "restricted-tier,victim-team", + } + request_mock.client = MagicMock() + request_mock.client.host = "127.0.0.1" + + data = {"model": "gpt-3.5-turbo"} + + user_api_key_dict = UserAPIKeyAuth( + api_key="hashed-key", + metadata={}, + team_metadata={}, + spend=0.0, + max_budget=100.0, + model_max_budget={}, + team_spend=0.0, + team_max_budget=200.0, + ) + + updated = await add_litellm_data_to_request( + data=data, + request=request_mock, + user_api_key_dict=user_api_key_dict, + proxy_config=MagicMock(), + general_settings={}, + version="test-version", + ) + + assert "tags" not in (updated.get("metadata") or {}) + + +@pytest.mark.asyncio +async def test_add_litellm_data_to_request_ignores_root_level_tags_without_permission(): + """Regression: root-level `data["tags"]` bypassed the body-metadata + tag strip. Root-level tags must also be gated by `allow_client_tags`.""" + from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request + + request_mock = MagicMock(spec=Request) + request_mock.url.path = "/v1/chat/completions" + request_mock.url = MagicMock() + request_mock.url.__str__.return_value = "http://localhost/v1/chat/completions" + request_mock.method = "POST" + request_mock.query_params = {} + request_mock.headers = {"Content-Type": "application/json"} + request_mock.client = MagicMock() + request_mock.client.host = "127.0.0.1" + + data = { + "model": "gpt-3.5-turbo", + "tags": ["restricted-tier", "victim-team"], + } + + user_api_key_dict = UserAPIKeyAuth( + api_key="hashed-key", + metadata={}, + team_metadata={}, + spend=0.0, + max_budget=100.0, + model_max_budget={}, + team_spend=0.0, + team_max_budget=200.0, + ) + + updated = await add_litellm_data_to_request( + data=data, + request=request_mock, + user_api_key_dict=user_api_key_dict, + proxy_config=MagicMock(), + general_settings={}, + version="test-version", + ) + + assert "tags" not in (updated.get("metadata") or {}) + # Also ensure the root-level tags are removed. get_tags_from_request_body + # reads request_body["tags"] directly, so leaving it in place would let + # the policy engine see caller-supplied tags even after the metadata + # strip. + assert "tags" not in updated + + +@pytest.mark.asyncio +async def test_add_litellm_data_to_request_honors_header_tags_when_opted_in(): + """When allow_client_tags=True, header-supplied tags flow through.""" + from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request + + request_mock = MagicMock(spec=Request) + request_mock.url.path = "/v1/chat/completions" + request_mock.url = MagicMock() + request_mock.url.__str__.return_value = "http://localhost/v1/chat/completions" + request_mock.method = "POST" + request_mock.query_params = {} + request_mock.headers = { + "Content-Type": "application/json", + "x-litellm-tags": "production,ab-test", + } + request_mock.client = MagicMock() + request_mock.client.host = "127.0.0.1" + + data = {"model": "gpt-3.5-turbo"} + + user_api_key_dict = UserAPIKeyAuth( + api_key="hashed-key", + metadata={"allow_client_tags": True}, + team_metadata={}, + spend=0.0, + max_budget=100.0, + model_max_budget={}, + team_spend=0.0, + team_max_budget=200.0, + ) + + updated = await add_litellm_data_to_request( + data=data, + request=request_mock, + user_api_key_dict=user_api_key_dict, + proxy_config=MagicMock(), + general_settings={}, + version="test-version", + ) + + assert updated["metadata"].get("tags") == ["production", "ab-test"] + + +@pytest.mark.asyncio +async def test_add_litellm_data_to_request_strips_user_tags_without_permission(): + """Caller-supplied metadata.tags must be stripped when the key/team + metadata does not opt in via allow_client_tags=True. Otherwise an + attacker can reach restricted tag-routed deployments or attribute + spend to a victim team's tag.""" + from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request + + request_mock = MagicMock(spec=Request) + request_mock.url.path = "/v1/chat/completions" + request_mock.url = MagicMock() + request_mock.url.__str__.return_value = "http://localhost/v1/chat/completions" + request_mock.method = "POST" + request_mock.query_params = {} + request_mock.headers = {"Content-Type": "application/json"} + request_mock.client = MagicMock() + request_mock.client.host = "127.0.0.1" + + data = { + "model": "gpt-3.5-turbo", + "metadata": {"tags": ["restricted-tier", "victim-team"]}, + "litellm_metadata": {"tags": ["also-stripped"]}, + } + + user_api_key_dict = UserAPIKeyAuth( + api_key="hashed-key", + metadata={}, + team_metadata={}, + spend=0.0, + max_budget=100.0, + model_max_budget={}, + team_spend=0.0, + team_max_budget=200.0, + ) + + updated = await add_litellm_data_to_request( + data=data, + request=request_mock, + user_api_key_dict=user_api_key_dict, + proxy_config=MagicMock(), + general_settings={}, + version="test-version", + ) + + assert "tags" not in (updated.get("metadata") or {}) + assert "tags" not in (updated.get("litellm_metadata") or {}) + + +@pytest.mark.asyncio +async def test_add_litellm_data_to_request_preserves_user_tags_when_key_opts_in(): + """When key.metadata.allow_client_tags=True, caller-supplied tags are + preserved and reach the router.""" + from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request + + request_mock = MagicMock(spec=Request) + request_mock.url.path = "/v1/chat/completions" + request_mock.url = MagicMock() + request_mock.url.__str__.return_value = "http://localhost/v1/chat/completions" + request_mock.method = "POST" + request_mock.query_params = {} + request_mock.headers = {"Content-Type": "application/json"} + request_mock.client = MagicMock() + request_mock.client.host = "127.0.0.1" + + data = { + "model": "gpt-3.5-turbo", + "metadata": {"tags": ["opted-in-tag"]}, + } + + user_api_key_dict = UserAPIKeyAuth( + api_key="hashed-key", + metadata={"allow_client_tags": True}, + team_metadata={}, + spend=0.0, + max_budget=100.0, + model_max_budget={}, + team_spend=0.0, + team_max_budget=200.0, + ) + + updated = await add_litellm_data_to_request( + data=data, + request=request_mock, + user_api_key_dict=user_api_key_dict, + proxy_config=MagicMock(), + general_settings={}, + version="test-version", + ) + + assert updated["metadata"].get("tags") == ["opted-in-tag"] + + +@pytest.mark.asyncio +async def test_add_litellm_data_to_request_preserves_user_tags_when_team_opts_in(): + """Team-level allow_client_tags is also honored (not just key-level).""" + from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request + + request_mock = MagicMock(spec=Request) + request_mock.url.path = "/v1/chat/completions" + request_mock.url = MagicMock() + request_mock.url.__str__.return_value = "http://localhost/v1/chat/completions" + request_mock.method = "POST" + request_mock.query_params = {} + request_mock.headers = {"Content-Type": "application/json"} + request_mock.client = MagicMock() + request_mock.client.host = "127.0.0.1" + + data = { + "model": "gpt-3.5-turbo", + "metadata": {"tags": ["team-allowed"]}, + } + + user_api_key_dict = UserAPIKeyAuth( + api_key="hashed-key", + metadata={}, + team_metadata={"allow_client_tags": True}, + spend=0.0, + max_budget=100.0, + model_max_budget={}, + team_spend=0.0, + team_max_budget=200.0, + ) + + updated = await add_litellm_data_to_request( + data=data, + request=request_mock, + user_api_key_dict=user_api_key_dict, + proxy_config=MagicMock(), + general_settings={}, + version="test-version", + ) + + assert updated["metadata"].get("tags") == ["team-allowed"] + + @pytest.mark.asyncio async def test_add_litellm_data_to_request_user_spend_and_budget(): from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request @@ -280,9 +854,11 @@ async def test_add_litellm_data_to_request_audio_transcription_multipart(): "file": b"Fake audio bytes", } + # Opt the key in to client-supplied tags so the parsed tags from the + # JSON-string multipart body aren't stripped by the admin-injection strip. user_api_key_dict = UserAPIKeyAuth( api_key="hashed-key", - metadata={}, + metadata={"allow_client_tags": True}, team_metadata={}, spend=0.0, max_budget=100.0,