Merge pull request #25983 from BerriAI/litellm_yj_apr17

[Infra] Merge dev branch
This commit is contained in:
yuneng-jiang 2026-04-18 14:43:04 -07:00 • committed by GitHub
commit e69051916e
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
78 changed files with 2613 additions and 218 deletions

View file

@ -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
)

View file

@ -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:

View file

@ -48,7 +48,6 @@ _supported_callback_params = [
"braintrust_host",
"slack_webhook_url",
"lunary_public_key",
"turn_off_message_logging",
]

View file

@ -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

View file

@ -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)

View file

@ -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")

View file

@ -1019,6 +1019,7 @@ class HTTPHandler:
url,
params=params,
headers=headers,
follow_redirects=_follow_redirects,
)
return response

View file

@ -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()

File diff suppressed because one or more lines are too long

File diff suppressed because one or more lines are too long

File diff suppressed because one or more lines are too long

File diff suppressed because one or more lines are too long

File diff suppressed because one or more lines are too long

File diff suppressed because one or more lines are too long

File diff suppressed because one or more lines are too long

File diff suppressed because one or more lines are too long

File diff suppressed because one or more lines are too long

File diff suppressed because one or more lines are too long

File diff suppressed because one or more lines are too long

File diff suppressed because one or more lines are too long

File diff suppressed because one or more lines are too long

File diff suppressed because one or more lines are too long

File diff suppressed because one or more lines are too long

File diff suppressed because one or more lines are too long

File diff suppressed because one or more lines are too long

File diff suppressed because one or more lines are too long

File diff suppressed because one or more lines are too long

File diff suppressed because one or more lines are too long

File diff suppressed because one or more lines are too long

File diff suppressed because one or more lines are too long

File diff suppressed because one or more lines are too long

File diff suppressed because one or more lines are too long

File diff suppressed because one or more lines are too long

File diff suppressed because one or more lines are too long

File diff suppressed because one or more lines are too long

File diff suppressed because one or more lines are too long

File diff suppressed because one or more lines are too long

File diff suppressed because one or more lines are too long

File diff suppressed because one or more lines are too long

File diff suppressed because one or more lines are too long

File diff suppressed because one or more lines are too long

File diff suppressed because one or more lines are too long

View file

@ -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)

View file

@ -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")

View file

@ -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:

View file

@ -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,

View file

@ -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)

View file

@ -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.

View file

@ -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)

View file

@ -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] = []

View file

@ -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")

View file

@ -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(

View file

@ -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(

View file

@ -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

View file

@ -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(

View file

@ -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(

View file

@ -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

View file

@ -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

View file

@ -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"

View file

@ -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)

View file

@ -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

View file

@ -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

View file

@ -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")

View file

@ -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

View file

@ -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

View file

@ -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

View file

@ -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))

View file

@ -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():

View file

@ -0,0 +1,396 @@
import socket
import pytest
import litellm
from litellm.litellm_core_utils import url_utils
from litellm.litellm_core_utils.url_utils import SSRFError, _is_blocked_ip, validate_url
@pytest.fixture
def mock_dns_public(monkeypatch):
"""Resolve any hostname to 93.184.216.34 (public)."""
def fake_getaddrinfo(host, port, *args, **kwargs):
return [
(socket.AF_INET, socket.SOCK_STREAM, 6, "", ("93.184.216.34", port or 80))
]
monkeypatch.setattr(url_utils.socket, "getaddrinfo", fake_getaddrinfo)
@pytest.fixture
def mock_dns_failure(monkeypatch):
"""Make every DNS lookup raise gaierror."""
def fake_getaddrinfo(host, port, *args, **kwargs):
raise socket.gaierror("Name or service not known")
monkeypatch.setattr(url_utils.socket, "getaddrinfo", fake_getaddrinfo)
class TestIsBlockedIp:
def test_blocks_private(self):
assert _is_blocked_ip("10.0.0.1") is True
def test_allows_public(self):
assert _is_blocked_ip("8.8.8.8") is False
def test_unparseable_is_blocked(self):
assert _is_blocked_ip("not-an-ip") is True
# Coverage delta picked up by switching to `not ip.is_global` (RFC 6890)
# over the old hand-maintained CIDR list.
def test_blocks_cgnat_alibaba_metadata(self):
"""100.100.100.200 is Alibaba Cloud metadata; lives in CGNAT."""
assert _is_blocked_ip("100.100.100.200") is True
def test_blocks_ietf_protocol_assignments_old_oracle_metadata(self):
"""192.0.0.192 was the legacy Oracle Cloud metadata IP."""
assert _is_blocked_ip("192.0.0.192") is True
def test_blocks_documentation_ranges(self):
assert _is_blocked_ip("192.0.2.1") is True
assert _is_blocked_ip("198.51.100.1") is True
assert _is_blocked_ip("203.0.113.1") is True
def test_blocks_multicast(self):
assert _is_blocked_ip("224.0.0.1") is True
def test_blocks_reserved_future_use(self):
assert _is_blocked_ip("240.0.0.1") is True
def test_blocks_broadcast(self):
assert _is_blocked_ip("255.255.255.255") is True
def test_blocks_azure_wire_server(self):
"""168.63.129.16 is globally routable but cloud-internal — explicit exception."""
assert _is_blocked_ip("168.63.129.16") is True
def test_blocks_aws_ipv6_imds(self):
"""fd00:ec2::254 is AWS's IPv6 IMDS, in IPv6 ULA (fc00::/7)."""
assert _is_blocked_ip("fd00:ec2::254") is True
def test_blocks_ipv4_mapped_private(self):
"""::ffff:10.0.0.1 must be unwrapped and blocked as 10.0.0.1."""
assert _is_blocked_ip("::ffff:10.0.0.1") is True
def test_blocks_ipv4_mapped_azure_wire_server(self):
"""::ffff:168.63.129.16 must be unwrapped and blocked via the exception list."""
assert _is_blocked_ip("::ffff:168.63.129.16") is True
class TestValidateUrl:
def test_blocks_loopback(self):
with pytest.raises(SSRFError):
validate_url("http://127.0.0.1/test")
def test_blocks_imds(self):
with pytest.raises(SSRFError):
validate_url("http://169.254.169.254/latest/meta-data/")
def test_blocks_rfc1918_class_a(self):
with pytest.raises(SSRFError):
validate_url("http://10.0.1.5:8080/v1/completions")
def test_blocks_rfc1918_class_b(self):
with pytest.raises(SSRFError):
validate_url("http://172.16.0.1/")
def test_blocks_rfc1918_class_c(self):
with pytest.raises(SSRFError):
validate_url("http://192.168.1.1/")
def test_blocks_file_scheme(self):
with pytest.raises(SSRFError):
validate_url("file:///etc/passwd")
def test_blocks_ftp_scheme(self):
with pytest.raises(SSRFError):
validate_url("ftp://internal.host/data")
def test_blocks_no_hostname(self):
with pytest.raises(SSRFError):
validate_url("http:///path")
def test_allows_public_https(self, mock_dns_public):
rewritten, host = validate_url("https://example.com/image.png")
assert host == "example.com"
assert rewritten == "https://example.com/image.png"
def test_rewrites_public_http_to_ip(self, mock_dns_public):
rewritten, host = validate_url("http://example.com/image.png")
assert host == "example.com"
assert "example.com" not in rewritten
def test_preserves_path_and_query(self, mock_dns_public):
rewritten, host = validate_url("http://example.com/path?key=value")
assert "/path" in rewritten
assert "key=value" in rewritten
def test_dns_failure_raises(self, mock_dns_failure):
with pytest.raises(SSRFError, match="DNS resolution failed"):
validate_url("http://this-domain-does-not-exist-xyz123.invalid/test")
def test_blocks_localhost_hostname(self, monkeypatch):
def fake(host, port, *a, **kw):
return [
(socket.AF_INET, socket.SOCK_STREAM, 6, "", ("127.0.0.1", port or 80))
]
monkeypatch.setattr(url_utils.socket, "getaddrinfo", fake)
with pytest.raises(SSRFError):
validate_url("http://localhost/")
def test_blocks_ipv6_loopback(self):
with pytest.raises(SSRFError):
validate_url("http://[::1]/")
def test_https_rewrites_when_ssl_verify_disabled(
self, monkeypatch, mock_dns_public
):
monkeypatch.setattr(litellm, "ssl_verify", False)
rewritten, host = validate_url("https://example.com/image.png")
assert host == "example.com"
assert "example.com" not in rewritten # rewritten to IP
def test_https_not_rewritten_when_ssl_verify_enabled(
self, monkeypatch, mock_dns_public
):
monkeypatch.setattr(litellm, "ssl_verify", True)
rewritten, host = validate_url("https://example.com/image.png")
assert rewritten == "https://example.com/image.png"
class TestHostHeaderFormatting:
"""RFC 7230 §5.4: IPv6 literals must be bracketed in the Host header."""
def test_ipv4_no_port(self, monkeypatch):
def fake(host, port, *a, **kw):
return [(socket.AF_INET, socket.SOCK_STREAM, 6, "", ("1.2.3.4", port))]
monkeypatch.setattr(url_utils.socket, "getaddrinfo", fake)
_, host = validate_url("http://example.com/")
assert host == "example.com"
def test_ipv4_with_explicit_nondefault_port(self, monkeypatch):
def fake(host, port, *a, **kw):
return [(socket.AF_INET, socket.SOCK_STREAM, 6, "", ("1.2.3.4", port))]
monkeypatch.setattr(url_utils.socket, "getaddrinfo", fake)
_, host = validate_url("http://example.com:8080/")
assert host == "example.com:8080"
def test_ipv4_with_explicit_default_port_strips_port(self, monkeypatch):
def fake(host, port, *a, **kw):
return [(socket.AF_INET, socket.SOCK_STREAM, 6, "", ("1.2.3.4", port))]
monkeypatch.setattr(url_utils.socket, "getaddrinfo", fake)
_, host = validate_url("http://example.com:80/")
assert host == "example.com"
def test_ipv6_literal_is_bracketed_with_port(self, monkeypatch):
"""Regression: IPv6 + port produced ambiguous `Host: 2001:db8::1:8080`."""
monkeypatch.setattr(litellm, "user_url_allowed_hosts", ["[2001:db8::1]"])
def fake(host, port, *a, **kw):
return [
(
socket.AF_INET6,
socket.SOCK_STREAM,
6,
"",
("2001:db8::1", port, 0, 0),
)
]
monkeypatch.setattr(url_utils.socket, "getaddrinfo", fake)
_, host = validate_url("http://[2001:db8::1]:8080/")
assert host == "[2001:db8::1]:8080"
def test_ipv6_literal_is_bracketed_without_port(self, monkeypatch):
monkeypatch.setattr(litellm, "user_url_allowed_hosts", ["[2001:db8::1]"])
def fake(host, port, *a, **kw):
return [
(
socket.AF_INET6,
socket.SOCK_STREAM,
6,
"",
("2001:db8::1", port, 0, 0),
)
]
monkeypatch.setattr(url_utils.socket, "getaddrinfo", fake)
_, host = validate_url("http://[2001:db8::1]/")
assert host == "[2001:db8::1]"
class TestRedirectHostnamePreservation:
"""Relative-location redirects must keep the original hostname, not the
rewritten IP, so the next hop's Host header still identifies the site."""
def test_relative_redirect_preserves_hostname_for_next_hop(self, monkeypatch):
def fake(host, port, *a, **kw):
return [
(socket.AF_INET, socket.SOCK_STREAM, 6, "", ("93.184.216.34", port))
]
monkeypatch.setattr(url_utils.socket, "getaddrinfo", fake)
class FakeResponse:
def __init__(self, status, location=None):
self.status_code = status
self.headers = {"location": location} if location else {}
self.is_redirect = 300 <= status < 400
hops = []
class FakeClient:
def __init__(self):
self._n = 0
def get(self, url, headers=None, follow_redirects=False, **kw):
hops.append({"url": url, "host": (headers or {}).get("Host")})
self._n += 1
if self._n == 1:
return FakeResponse(302, "/redirected")
return FakeResponse(200)
url_utils.safe_get(FakeClient(), "http://example.com/initial")
assert len(hops) == 2
# Both hops must carry the ORIGINAL hostname in the Host header.
assert hops[0]["host"] == "example.com"
assert hops[1]["host"] == "example.com"
# Both outbound URLs go to the resolved IP (rewritten), not the hostname.
assert "93.184.216.34" in hops[0]["url"]
assert "93.184.216.34" in hops[1]["url"]
# The second hop resolved /redirected relative to the original, not the IP.
assert hops[1]["url"].endswith("/redirected")
class TestValidationMasterSwitch:
def test_disabled_bypasses_fetch_in_safe_get(self, monkeypatch):
"""When user_url_validation is False, safe_get delegates to client.get without validation."""
monkeypatch.setattr(litellm, "user_url_validation", False)
calls = []
class FakeClient:
def get(self, url, **kwargs):
calls.append((url, kwargs))
class R:
is_redirect = False
return R()
url_utils.safe_get(FakeClient(), "http://127.0.0.1/internal")
assert calls and calls[0][0] == "http://127.0.0.1/internal"
assert calls[0][1].get("follow_redirects") is True
def test_enabled_still_blocks(self, monkeypatch):
monkeypatch.setattr(litellm, "user_url_validation", True)
with pytest.raises(SSRFError):
validate_url("http://127.0.0.1/")
class TestHostAllowlist:
def test_allowlisted_hostname_permits_private_ip(self, monkeypatch):
monkeypatch.setattr(litellm, "user_url_allowed_hosts", ["internal.corp"])
def fake(host, port, *a, **kw):
return [(socket.AF_INET, socket.SOCK_STREAM, 6, "", ("10.0.1.5", port))]
monkeypatch.setattr(url_utils.socket, "getaddrinfo", fake)
rewritten, host = validate_url("http://internal.corp/path")
assert host == "internal.corp"
assert "10.0.1.5" in rewritten
def test_non_allowlisted_hostname_still_blocked(self, monkeypatch):
monkeypatch.setattr(litellm, "user_url_allowed_hosts", ["internal.corp"])
def fake(host, port, *a, **kw):
return [(socket.AF_INET, socket.SOCK_STREAM, 6, "", ("10.0.1.5", port))]
monkeypatch.setattr(url_utils.socket, "getaddrinfo", fake)
with pytest.raises(SSRFError):
validate_url("http://other.corp/")
def test_allowlist_case_insensitive(self, monkeypatch):
monkeypatch.setattr(litellm, "user_url_allowed_hosts", ["Internal.Corp"])
def fake(host, port, *a, **kw):
return [(socket.AF_INET, socket.SOCK_STREAM, 6, "", ("10.0.1.5", port))]
monkeypatch.setattr(url_utils.socket, "getaddrinfo", fake)
rewritten, _ = validate_url("http://internal.corp/")
assert "10.0.1.5" in rewritten
def test_allowlist_with_port_matches_explicit_port(self, monkeypatch):
monkeypatch.setattr(litellm, "user_url_allowed_hosts", ["internal.corp:8080"])
def fake(host, port, *a, **kw):
return [(socket.AF_INET, socket.SOCK_STREAM, 6, "", ("10.0.1.5", port))]
monkeypatch.setattr(url_utils.socket, "getaddrinfo", fake)
rewritten, host = validate_url("http://internal.corp:8080/")
assert host == "internal.corp:8080"
assert "10.0.1.5" in rewritten
def test_allowlist_with_port_matches_default_port(self, monkeypatch):
"""Admin entry `host:443` matches `https://host/` (port=None, default 443)."""
monkeypatch.setattr(litellm, "user_url_allowed_hosts", ["internal.corp:443"])
def fake(host, port, *a, **kw):
return [(socket.AF_INET, socket.SOCK_STREAM, 6, "", ("10.0.1.5", port))]
monkeypatch.setattr(url_utils.socket, "getaddrinfo", fake)
# Should succeed — no SSRFError raised
validate_url("https://internal.corp/")
def test_allowlist_port_specific_does_not_match_other_port(self, monkeypatch):
monkeypatch.setattr(litellm, "user_url_allowed_hosts", ["internal.corp:8080"])
def fake(host, port, *a, **kw):
return [(socket.AF_INET, socket.SOCK_STREAM, 6, "", ("10.0.1.5", port))]
monkeypatch.setattr(url_utils.socket, "getaddrinfo", fake)
with pytest.raises(SSRFError):
validate_url("http://internal.corp:9090/")
def test_allowlist_host_entry_matches_any_port(self, monkeypatch):
monkeypatch.setattr(litellm, "user_url_allowed_hosts", ["internal.corp"])
def fake(host, port, *a, **kw):
return [(socket.AF_INET, socket.SOCK_STREAM, 6, "", ("10.0.1.5", port))]
monkeypatch.setattr(url_utils.socket, "getaddrinfo", fake)
validate_url("http://internal.corp:9090/")
validate_url("https://internal.corp:8443/")
def test_allowlist_permits_loopback(self, monkeypatch):
"""Admin may opt into loopback if they explicitly configure it."""
monkeypatch.setattr(litellm, "user_url_allowed_hosts", ["localhost"])
def fake(host, port, *a, **kw):
return [(socket.AF_INET, socket.SOCK_STREAM, 6, "", ("127.0.0.1", port))]
monkeypatch.setattr(url_utils.socket, "getaddrinfo", fake)
rewritten, host = validate_url("http://localhost:8080/")
assert host == "localhost:8080"
def test_empty_allowlist_retains_default_deny(self, monkeypatch):
monkeypatch.setattr(litellm, "user_url_allowed_hosts", [])
with pytest.raises(SSRFError):
validate_url("http://127.0.0.1/")
def test_allowlist_strips_trailing_dot(self, monkeypatch):
monkeypatch.setattr(litellm, "user_url_allowed_hosts", ["internal.corp."])
def fake(host, port, *a, **kw):
return [(socket.AF_INET, socket.SOCK_STREAM, 6, "", ("10.0.1.5", port))]
monkeypatch.setattr(url_utils.socket, "getaddrinfo", fake)
validate_url("http://internal.corp/")

View file

@ -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"'})

View file

@ -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():
"""

View file

@ -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"]

View file

@ -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."""

View file

@ -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": "<caller>", "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
# =====================================================================

View file

@ -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
# ============================================================================

View file

@ -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"

View file

@ -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"

View file

@ -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,