diff --git a/.github/workflows/create-release.yml b/.github/workflows/create-release.yml index 39d078267f6..a726a921a2b 100644 --- a/.github/workflows/create-release.yml +++ b/.github/workflows/create-release.yml @@ -4,7 +4,7 @@ on: workflow_dispatch: inputs: tag: - description: "Release tag (e.g. 1.84.0, 1.84.0rc1, 1.84.0.dev42, 1.84.0.post1; legacy v1.83.10-stable still accepted)" + description: "Release tag (e.g. 1.84.0, 1.84.0rc1, 1.84.0.dev42, 1.84.0-dev.2, 1.84.0.post1; legacy v1.83.10-stable still accepted)" required: true type: string commit_hash: @@ -46,9 +46,11 @@ jobs: const commitHash = process.env.COMMIT_HASH; // Mark RC / dev / nightly / alpha / beta tags as GitHub pre-releases. + // Accept both PEP 440 (`.dev`) and SemVer (`-dev`) separators so tags + // like `1.84.0.dev2` and `1.84.0-dev.2` are both detected. // PEP 440 post-releases (e.g. `1.84.0.post1`) and legacy `-stable[.patch.N]` // are stable maintenance releases, not pre-releases. - const isPrerelease = /(?:rc|nightly|alpha|beta|\.dev)/i.test(tag); + const isPrerelease = /(?:rc|nightly|alpha|beta|[-.]dev)/i.test(tag); const cosignSection = [ `## Verify Docker Image Signature`, diff --git a/Dockerfile b/Dockerfile index 03779d6c884..915daff999c 100644 --- a/Dockerfile +++ b/Dockerfile @@ -1,9 +1,9 @@ # Base image for building -ARG LITELLM_BUILD_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:f26d42a15d09d9a643b231df929fa3cf609bedc58a728eb445be89a9d8d1da9f +ARG LITELLM_BUILD_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:3258be472764337fd13095bcbb3182da170243b5819fd67ad4c0754590588b31 # Runtime image -ARG LITELLM_RUNTIME_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:f26d42a15d09d9a643b231df929fa3cf609bedc58a728eb445be89a9d8d1da9f -ARG UV_IMAGE=ghcr.io/astral-sh/uv:0.11.7@sha256:733b4042187702f832f7fdecb3aff14a61b288c4ca37af188bb5715c1caebaf8 +ARG LITELLM_RUNTIME_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:3258be472764337fd13095bcbb3182da170243b5819fd67ad4c0754590588b31 +ARG UV_IMAGE=ghcr.io/astral-sh/uv:0.11.7@sha256:240fb85ab0f263ef12f492d8476aa3a2e4e1e333f7d67fbdd923d00a506a516a FROM $UV_IMAGE AS uvbin diff --git a/docker/Dockerfile.alpine b/docker/Dockerfile.alpine index 2cfc5ef03ff..5de588cf4e4 100644 --- a/docker/Dockerfile.alpine +++ b/docker/Dockerfile.alpine @@ -3,7 +3,7 @@ ARG LITELLM_BUILD_IMAGE=python:3.11-alpine@sha256:f07e2ace46f560f09a6eeec7b4913b # Runtime image ARG LITELLM_RUNTIME_IMAGE=python:3.11-alpine@sha256:f07e2ace46f560f09a6eeec7b4913b80ee99546e749ef82342a419a326620856 -ARG UV_IMAGE=ghcr.io/astral-sh/uv:0.11.7@sha256:733b4042187702f832f7fdecb3aff14a61b288c4ca37af188bb5715c1caebaf8 +ARG UV_IMAGE=ghcr.io/astral-sh/uv:0.11.7@sha256:240fb85ab0f263ef12f492d8476aa3a2e4e1e333f7d67fbdd923d00a506a516a FROM $UV_IMAGE AS uvbin diff --git a/docker/Dockerfile.database b/docker/Dockerfile.database index e3edcece61c..671f374ca27 100644 --- a/docker/Dockerfile.database +++ b/docker/Dockerfile.database @@ -1,9 +1,9 @@ # Base image for building -ARG LITELLM_BUILD_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:f26d42a15d09d9a643b231df929fa3cf609bedc58a728eb445be89a9d8d1da9f +ARG LITELLM_BUILD_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:3258be472764337fd13095bcbb3182da170243b5819fd67ad4c0754590588b31 # Runtime image -ARG LITELLM_RUNTIME_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:f26d42a15d09d9a643b231df929fa3cf609bedc58a728eb445be89a9d8d1da9f -ARG UV_IMAGE=ghcr.io/astral-sh/uv:0.11.7@sha256:733b4042187702f832f7fdecb3aff14a61b288c4ca37af188bb5715c1caebaf8 +ARG LITELLM_RUNTIME_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:3258be472764337fd13095bcbb3182da170243b5819fd67ad4c0754590588b31 +ARG UV_IMAGE=ghcr.io/astral-sh/uv:0.11.7@sha256:240fb85ab0f263ef12f492d8476aa3a2e4e1e333f7d67fbdd923d00a506a516a FROM $UV_IMAGE AS uvbin diff --git a/docker/Dockerfile.dev b/docker/Dockerfile.dev index e2dc1857835..ebc92a22d50 100644 --- a/docker/Dockerfile.dev +++ b/docker/Dockerfile.dev @@ -3,7 +3,7 @@ ARG LITELLM_BUILD_IMAGE=python:3.13-slim@sha256:739e7213785e88c0f702dcdc12c0973a # Runtime image ARG LITELLM_RUNTIME_IMAGE=python:3.13-slim@sha256:739e7213785e88c0f702dcdc12c0973afcbd606dbf021a589cab77d6b00b579d -ARG UV_IMAGE=ghcr.io/astral-sh/uv:0.11.7@sha256:733b4042187702f832f7fdecb3aff14a61b288c4ca37af188bb5715c1caebaf8 +ARG UV_IMAGE=ghcr.io/astral-sh/uv:0.11.7@sha256:240fb85ab0f263ef12f492d8476aa3a2e4e1e333f7d67fbdd923d00a506a516a FROM $UV_IMAGE AS uvbin diff --git a/docker/Dockerfile.health_check b/docker/Dockerfile.health_check index 07d35b5e291..a2e5cb9f71f 100644 --- a/docker/Dockerfile.health_check +++ b/docker/Dockerfile.health_check @@ -1,4 +1,4 @@ -ARG UV_IMAGE=ghcr.io/astral-sh/uv:0.11.7@sha256:733b4042187702f832f7fdecb3aff14a61b288c4ca37af188bb5715c1caebaf8 +ARG UV_IMAGE=ghcr.io/astral-sh/uv:0.11.7@sha256:240fb85ab0f263ef12f492d8476aa3a2e4e1e333f7d67fbdd923d00a506a516a FROM $UV_IMAGE AS uvbin FROM python:3.13-slim@sha256:739e7213785e88c0f702dcdc12c0973afcbd606dbf021a589cab77d6b00b579d diff --git a/docker/Dockerfile.non_root b/docker/Dockerfile.non_root index 8512ae8ad92..ab40ee138e3 100644 --- a/docker/Dockerfile.non_root +++ b/docker/Dockerfile.non_root @@ -1,8 +1,8 @@ # Base images -ARG LITELLM_BUILD_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:f26d42a15d09d9a643b231df929fa3cf609bedc58a728eb445be89a9d8d1da9f -ARG LITELLM_RUNTIME_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:f26d42a15d09d9a643b231df929fa3cf609bedc58a728eb445be89a9d8d1da9f +ARG LITELLM_BUILD_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:3258be472764337fd13095bcbb3182da170243b5819fd67ad4c0754590588b31 +ARG LITELLM_RUNTIME_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:3258be472764337fd13095bcbb3182da170243b5819fd67ad4c0754590588b31 ARG PROXY_EXTRAS_SOURCE=published -ARG UV_IMAGE=ghcr.io/astral-sh/uv:0.11.7@sha256:733b4042187702f832f7fdecb3aff14a61b288c4ca37af188bb5715c1caebaf8 +ARG UV_IMAGE=ghcr.io/astral-sh/uv:0.11.7@sha256:240fb85ab0f263ef12f492d8476aa3a2e4e1e333f7d67fbdd923d00a506a516a FROM $UV_IMAGE AS uvbin diff --git a/docs/my-website/docs/providers/crusoe.md b/docs/my-website/docs/providers/crusoe.md new file mode 100644 index 00000000000..aa737cbdcd8 --- /dev/null +++ b/docs/my-website/docs/providers/crusoe.md @@ -0,0 +1,196 @@ +import Tabs from '@theme/Tabs'; +import TabItem from '@theme/TabItem'; + +# Crusoe + +## Overview + +| Property | Details | +|-------|-------| +| Description | Crusoe Cloud provides GPU-accelerated inference for open-source large language models, optimized for performance and cost efficiency. | +| Provider Route on LiteLLM | `crusoe/` | +| Link to Provider Doc | [Crusoe Managed Inference Documentation ↗](https://docs.crusoecloud.com/managed-inference/overview/index.html) | +| Base URL | `https://managed-inference-api-proxy.crusoecloud.com/v1` | +| Supported Operations | [`/chat/completions`](#sample-usage) | + +
+
+ +**We support ALL Crusoe models, just set `crusoe/` as a prefix when sending completion requests** + +## Available Models + +| Model | Description | Context Window | +|-------|-------------|----------------| +| `crusoe/deepseek-ai/DeepSeek-R1-0528` | DeepSeek R1 reasoning model (May 2025) | 163,840 tokens | +| `crusoe/deepseek-ai/DeepSeek-V3-0324` | DeepSeek V3 chat model (March 2025) | 163,840 tokens | +| `crusoe/google/gemma-3-12b-it` | Google Gemma 3 12B instruction-tuned | 131,072 tokens | +| `crusoe/meta-llama/Llama-3.3-70B-Instruct` | Llama 3.3 70B instruction-tuned | 131,072 tokens | +| `crusoe/moonshotai/Kimi-K2-Thinking` | Kimi K2 extended thinking model | 262,144 tokens | +| `crusoe/openai/gpt-oss-120b` | OpenAI 120B open-source model | 131,072 tokens | +| `crusoe/Qwen/Qwen3-235B-A22B-Instruct-2507` | Qwen3 235B MoE instruction-tuned | 262,144 tokens | + +## Required Variables + +```python showLineNumbers title="Environment Variables" +os.environ["CRUSOE_API_KEY"] = "" # your Crusoe API key +``` + +## Usage - LiteLLM Python SDK + +### Non-streaming + +```python showLineNumbers title="Crusoe Non-streaming Completion" +import os +import litellm +from litellm import completion + +os.environ["CRUSOE_API_KEY"] = "" # your Crusoe API key + +messages = [{"content": "Hello, how are you?", "role": "user"}] + +# Crusoe call +response = completion( + model="crusoe/meta-llama/Llama-3.3-70B-Instruct", + messages=messages +) + +print(response) +``` + +### Streaming + +```python showLineNumbers title="Crusoe Streaming Completion" +import os +import litellm +from litellm import completion + +os.environ["CRUSOE_API_KEY"] = "" # your Crusoe API key + +messages = [{"content": "Write a short story about AI", "role": "user"}] + +# Crusoe call with streaming +response = completion( + model="crusoe/meta-llama/Llama-3.3-70B-Instruct", + messages=messages, + stream=True +) + +for chunk in response: + print(chunk) +``` + +### Function Calling + +```python showLineNumbers title="Crusoe Function Calling" +import os +import litellm +from litellm import completion + +os.environ["CRUSOE_API_KEY"] = "" # your Crusoe API key + +tools = [{ + "type": "function", + "function": { + "name": "get_weather", + "description": "Get the current weather in a location", + "parameters": { + "type": "object", + "properties": { + "location": { + "type": "string", + "description": "The city and state, e.g. San Francisco, CA" + } + }, + "required": ["location"] + } + } +}] + +messages = [{"role": "user", "content": "What's the weather in Boston?"}] + +response = completion( + model="crusoe/meta-llama/Llama-3.3-70B-Instruct", + messages=messages, + tools=tools, + tool_choice="auto" +) + +print(response) +``` + +## Usage - LiteLLM Proxy Server + +```yaml showLineNumbers title="config.yaml" +model_list: + - model_name: llama-3.3-70b + litellm_params: + model: crusoe/meta-llama/Llama-3.3-70B-Instruct + api_key: os.environ/CRUSOE_API_KEY + - model_name: deepseek-r1 + litellm_params: + model: crusoe/deepseek-ai/DeepSeek-R1-0528 + api_key: os.environ/CRUSOE_API_KEY + - model_name: deepseek-v3 + litellm_params: + model: crusoe/deepseek-ai/DeepSeek-V3-0324 + api_key: os.environ/CRUSOE_API_KEY + - model_name: qwen3-235b + litellm_params: + model: crusoe/Qwen/Qwen3-235B-A22B-Instruct-2507 + api_key: os.environ/CRUSOE_API_KEY + - model_name: kimi-k2 + litellm_params: + model: crusoe/moonshotai/Kimi-K2-Thinking + api_key: os.environ/CRUSOE_API_KEY +``` + +## Custom API Base + +**Option 1: Environment variable** + +```python showLineNumbers title="Custom API Base via env var" +import os +from litellm import completion + +os.environ["CRUSOE_API_BASE"] = "https://custom.crusoecloud.com/v1" +os.environ["CRUSOE_API_KEY"] = "" # your API key + +response = completion( + model="crusoe/meta-llama/Llama-3.3-70B-Instruct", + messages=[{"content": "Hello!", "role": "user"}], +) +``` + +**Option 2: Pass directly** + +```python showLineNumbers title="Custom API Base via parameter" +from litellm import completion + +response = completion( + model="crusoe/meta-llama/Llama-3.3-70B-Instruct", + messages=[{"content": "Hello!", "role": "user"}], + api_base="https://custom.crusoecloud.com/v1", + api_key="your-api-key", +) +``` + +## Supported OpenAI Parameters + +- `temperature` +- `max_tokens` +- `max_completion_tokens` +- `top_p` +- `frequency_penalty` +- `presence_penalty` +- `stop` +- `n` +- `stream` +- `tools` +- `tool_choice` +- `response_format` +- `seed` +- `user` +- `logit_bias` +- `logprobs` +- `top_logprobs` diff --git a/enterprise/litellm_enterprise/proxy/auth/custom_sso_handler.py b/enterprise/litellm_enterprise/proxy/auth/custom_sso_handler.py index a3682320387..e8f104c2625 100644 --- a/enterprise/litellm_enterprise/proxy/auth/custom_sso_handler.py +++ b/enterprise/litellm_enterprise/proxy/auth/custom_sso_handler.py @@ -10,28 +10,21 @@ has already authenticated the user) and you need to extract user information fro custom headers or other request attributes. """ -from typing import TYPE_CHECKING, Dict, Optional, Union, cast +from typing import cast from fastapi import Request from fastapi.responses import RedirectResponse -if TYPE_CHECKING: - from fastapi_sso.sso.base import OpenID -else: - from typing import Any as OpenID - -from litellm.proxy.management_endpoints.types import CustomOpenID - class EnterpriseCustomSSOHandler: """ Enterprise Custom SSO Handler for LiteLLM Proxy - + This class provides methods for handling custom SSO authentication flows where users can implement their own authentication logic by processing request headers and returning user information in OpenID format. """ - + @staticmethod async def handle_custom_ui_sso_sign_in( request: Request, @@ -40,16 +33,16 @@ class EnterpriseCustomSSOHandler: Allow a user to execute their custom code to parse incoming request headers and return a OpenID object Use this when you have an OAuth proxy in front of LiteLLM (where the OAuth proxy has already authenticated the user) - + Args: request: The FastAPI request object containing headers and other request data - + Returns: RedirectResponse: Redirect response that sends the user to the LiteLLM UI with authentication token - + Raises: ValueError: If custom_ui_sso_sign_in_handler is not configured - + Example: This method is typically called when a user has already been authenticated by an external OAuth proxy and the proxy has added custom headers containing user information. @@ -60,27 +53,44 @@ class EnterpriseCustomSSOHandler: from litellm.integrations.custom_sso_handler import CustomSSOLoginHandler from litellm.proxy.proxy_server import ( CommonProxyErrors, + general_settings, premium_user, user_custom_ui_sso_sign_in_handler, ) + from litellm.proxy.auth.trusted_proxy_utils import ( + require_trusted_proxy_request, + ) + if premium_user is not True: raise ValueError(CommonProxyErrors.not_premium_user.value) - + if user_custom_ui_sso_sign_in_handler is None: - raise ValueError("custom_ui_sso_sign_in_handler is not configured. Please set it in general_settings.") - - custom_sso_login_handler = cast(CustomSSOLoginHandler, user_custom_ui_sso_sign_in_handler) - openid_response: OpenID = await custom_sso_login_handler.handle_custom_ui_sso_sign_in( + raise ValueError( + "custom_ui_sso_sign_in_handler is not configured. Please set it in general_settings." + ) + + require_trusted_proxy_request( request=request, + general_settings=general_settings, + feature_name="Custom UI SSO", ) - + + custom_sso_login_handler = cast( + CustomSSOLoginHandler, user_custom_ui_sso_sign_in_handler + ) + openid_response: OpenID = ( + await custom_sso_login_handler.handle_custom_ui_sso_sign_in( + request=request, + ) + ) + # Import here to avoid circular imports from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler - + return await SSOAuthenticationHandler.get_redirect_response_from_openid( result=openid_response, request=request, received_response=None, generic_client_id=None, ui_access_mode=None, - ) \ No newline at end of file + ) diff --git a/enterprise/litellm_enterprise/proxy/management_endpoints/project_endpoints.py b/enterprise/litellm_enterprise/proxy/management_endpoints/project_endpoints.py index 4bfe9d31874..75229bacc8f 100644 --- a/enterprise/litellm_enterprise/proxy/management_endpoints/project_endpoints.py +++ b/enterprise/litellm_enterprise/proxy/management_endpoints/project_endpoints.py @@ -588,24 +588,21 @@ async def update_project( # noqa: PLR0915 param="project_id", ) - # Validate team exists and get team object for limit + permission checks - team_id_to_check = data.team_id or existing_project.team_id - team_obj_for_checks = None - if team_id_to_check is not None: - team_obj_for_checks = await _validate_team_exists( - team_id=team_id_to_check, prisma_client=prisma_client + # Permission to *edit* the project must be evaluated against the + # project's CURRENT team. Sourcing the team from `data.team_id` + # would let an admin of any team pass the check by supplying their + # own team_id, hijacking the project (VERIA-55). + target_team_id = data.team_id or existing_project.team_id + target_team_obj = None + if target_team_id is not None: + target_team_obj = await _validate_team_exists( + team_id=target_team_id, prisma_client=prisma_client ) - # Check if user has permission to update this project has_permission = await _check_user_permission_for_project( user_api_key_dict=user_api_key_dict, team_id=existing_project.team_id, prisma_client=prisma_client, - team_object=( - LiteLLM_TeamTable(**team_obj_for_checks.model_dump()) - if team_obj_for_checks - else None - ), ) if not has_permission: @@ -614,10 +611,32 @@ async def update_project( # noqa: PLR0915 detail={"error": "Only admins or team admins can update projects"}, ) + # Reassigning to a different team also requires admin rights on the + # destination team — otherwise a team admin could shed projects into + # an unsuspecting team's namespace. + if data.team_id is not None and data.team_id != existing_project.team_id: + can_assign_to_target = await _check_user_permission_for_project( + user_api_key_dict=user_api_key_dict, + team_id=data.team_id, + prisma_client=prisma_client, + team_object=( + LiteLLM_TeamTable(**target_team_obj.model_dump()) + if target_team_obj + else None + ), + ) + if not can_assign_to_target: + raise HTTPException( + status_code=403, + detail={ + "error": "Cannot reassign project to a team you are not an admin of" + }, + ) + # Validate project limits against team limits - if team_obj_for_checks is not None: + if target_team_obj is not None: _check_team_project_limits( - team_object=LiteLLM_TeamTable(**team_obj_for_checks.model_dump()), + team_object=LiteLLM_TeamTable(**target_team_obj.model_dump()), data=data, ) diff --git a/litellm/__init__.py b/litellm/__init__.py index 77fa48625d9..28112f5c12a 100644 --- a/litellm/__init__.py +++ b/litellm/__init__.py @@ -288,6 +288,7 @@ disable_token_counter: bool = False disable_add_transform_inline_image_block: bool = False disable_add_user_agent_to_request_tags: bool = False disable_anthropic_gemini_context_caching_transform: bool = False +disable_vertex_batch_output_transformation: bool = False extra_spend_tag_headers: Optional[List[str]] = None in_memory_llm_clients_cache: "LLMClientCache" safe_memory_mode: bool = False @@ -330,6 +331,9 @@ enable_model_config_credential_overrides: bool = False enable_key_alias_format_validation: bool = ( False # opt-in validation of key_alias format on /key/generate and /key/update ) +enable_gemini_default_thinking_level_low: bool = ( + False # opt-in: force thinkingLevel low/minimal for Gemini 3 thinking param mapping +) #################### logging: bool = True enable_loadbalancing_on_batch_endpoints: Optional[bool] = None diff --git a/litellm/_logging.py b/litellm/_logging.py index d072cc549d0..5ddafd6c6af 100644 --- a/litellm/_logging.py +++ b/litellm/_logging.py @@ -1,12 +1,12 @@ import ast import logging import os -import re import sys from datetime import datetime from logging import Formatter -from typing import Any, Dict, List, Optional +from typing import Any, Dict, Optional +from litellm.litellm_core_utils.secret_redaction import redact_string from litellm.litellm_core_utils.safe_json_dumps import safe_dumps from litellm.litellm_core_utils.safe_json_loads import safe_json_loads @@ -21,74 +21,11 @@ _ENABLE_SECRET_REDACTION = ( os.getenv("LITELLM_DISABLE_REDACT_SECRETS", "").lower() != "true" ) -_REDACTED = "REDACTED" - - -def _build_secret_patterns() -> re.Pattern: - patterns: List[str] = [ - # ── PEM private key / certificate blocks ── - r"-----BEGIN[A-Z \-]*PRIVATE KEY-----[\s\S]*?-----END[A-Z \-]*PRIVATE KEY-----", - # ── GCP OAuth2 access tokens (ya29.*) ── - r"\bya29\.[A-Za-z0-9_.~+/-]+", - # ── Credential %s formatting (space separator, no key= prefix) ── - r"(?:client_secret|azure_password|azure_username)\s+[^\s,'\"})\]{}>]+", - # AWS access key IDs - r"(?:AKIA|ASIA)[0-9A-Z]{16}", - # AWS secrets / session tokens / access key IDs (key=value) - r"(?:aws_secret_access_key|aws_session_token|aws_access_key_id)" - r"\s*[:=]\s*[A-Za-z0-9/+=]{20,}", - # Bearer tokens (OAuth, JWT, etc.) - r"Bearer\s+[A-Za-z0-9\-._~+/]{10,}=*", - # Basic auth headers - r"Basic\s+[A-Za-z0-9+/]{10,}={0,2}", - # OpenAI / Anthropic sk- prefixed keys - r"sk-[A-Za-z0-9\-_]{20,}", - # Generic api_key / api-key / apikey (handles 'key': 'value' dict repr) - r"(?:api[_-]?key)['\"]?\s*[:=]\s*['\"]?[^\s,'\"})\]{}>]{8,}", - # x-api-key / api-key header values (handles 'key': 'value' dict repr) - r"(?:x-api-key|api-key)['\"]?\s*[:=]\s*['\"]?[^\s,'\"})\]{}>]+", - # Anthropic internal header keys - r"x-ak-[A-Za-z0-9\-_]{20,}", - # Google API keys - r"AIza[0-9A-Za-z\-_]{35}", - # Password / secret params (handles key=value and 'key': 'value') - # Word boundary prevents O(n^2) backtracking on long word-char runs. - r"(?:^|(?<=\W))\w*(?:password|passwd|client_secret|secret_key|_secret)" - r"['\"]?\s*[:=]\s*['\"]?[^\s,'\"})\]{}>]+", - # Database connection string credentials (scheme://user:pass@host) - r"(?<=://)[^\s'\"]*:[^\s'\"@]+(?=@)", - # Databricks personal access tokens - r"dapi[0-9a-f]{32}", - # ── Key-name-based redaction ── - # Catches secrets inside dicts/config dumps by matching on the KEY name - # regardless of what the value looks like. - # e.g. 'master_key': 'any-value-here', "database_url": "postgres://..." - # private_key with PEM-aware value capture - r"""private_key['\"]?\s*[:=]\s*['\"]?(?:-----BEGIN[A-Z \-]*PRIVATE KEY-----[\s\S]*?-----END[A-Z \-]*PRIVATE KEY-----|[^\s,'\"})\]{}>]+)""", - r"(?:master_key|database_url|db_url|connection_string|" - r"signing_key|encryption_key|" - r"auth_token|access_token|refresh_token|" - r"slack_webhook_url|webhook_url|" - r"database_connection_string|" - r"huggingface_token|jwt_secret)" - r"""['\"]?\s*[:=]\s*['\"]?[^\s,'\"})\]{}>]+""", - # ── Raw JWTs (without Bearer prefix) ── - r"\beyJ[A-Za-z0-9_-]{10,}\.[A-Za-z0-9_-]+\.[A-Za-z0-9_-]*", - # ── Azure SAS tokens in URLs ── - r"[?&]sig=[A-Za-z0-9%+/=]+", - # ── Full JSON service-account blobs (single-line and multi-line) ── - r'\{[^{}]*"type"\s*:\s*"service_account"[^{}]*(?:\{[^{}]*\}[^{}]*)*\}', - ] - return re.compile("|".join(patterns), re.IGNORECASE) - - -_SECRET_RE = _build_secret_patterns() - def _redact_string(value: str) -> str: if not _ENABLE_SECRET_REDACTION: return value - return _SECRET_RE.sub(_REDACTED, value) + return redact_string(value) def redact_secrets(value: str) -> str: diff --git a/litellm/batches/batch_utils.py b/litellm/batches/batch_utils.py index 4b965d4e635..aaf083e75d6 100644 --- a/litellm/batches/batch_utils.py +++ b/litellm/batches/batch_utils.py @@ -387,6 +387,27 @@ def _get_batch_job_total_usage_from_file_content( ) +def _get_models_from_batch_input_file_content( + file_content_dictionary: List[dict], +) -> List[str]: + """Extract the distinct ``body.model`` values from a batch *input* file. + + Used by the proxy's batch pre-call hook to enforce that the caller is + authorized for every model named inside the JSONL — not just the one + on the outer request — so the proxy's per-key model allowlist isn't + bypassed by smuggling expensive models into the batch file. + """ + models: List[str] = [] + seen: set = set() + for _item in file_content_dictionary: + body = _item.get("body") or {} + model = body.get("model") + if model and model not in seen: + seen.add(model) + models.append(model) + return models + + def _get_batch_job_input_file_usage( file_content_dictionary: List[dict], custom_llm_provider: Literal["openai", "azure", "vertex_ai"] = "openai", @@ -403,11 +424,25 @@ def _get_batch_job_input_file_usage( for _item in file_content_dictionary: body = _item.get("body", {}) model = body.get("model", model_name or "") - messages = body.get("messages", []) + # Chat completion payloads. + messages = body.get("messages") if messages: - item_tokens = token_counter(model=model, messages=messages) - prompt_tokens += item_tokens + prompt_tokens += token_counter(model=model, messages=messages) + continue + + # Text completion payloads (`prompt`). + prompt = body.get("prompt") + if prompt: + prompt_tokens += _count_prompt_or_input_tokens(model=model, value=prompt) + continue + + # Embedding payloads (`input`). + input_data = body.get("input") + if input_data: + prompt_tokens += _count_prompt_or_input_tokens( + model=model, value=input_data + ) return Usage( total_tokens=prompt_tokens + completion_tokens, @@ -416,6 +451,43 @@ def _get_batch_job_input_file_usage( ) +def _count_prompt_or_input_tokens(model: str, value: Any) -> int: + """Token-count a ``prompt`` / ``input`` field that the OpenAI batch + schema allows in four shapes: + + - ``str``: a single text prompt. + - ``list[str]``: multiple text prompts. + - ``list[int]``: a pre-tokenized prompt (each int counts as 1 token). + - ``list[list[int]]``: multiple pre-tokenized prompts. + + Pre-fix only the string shapes were counted, so a caller could send + a large ``list[list[int]]`` payload and slip past TPM rate limits + with a recorded cost of zero tokens. + """ + if isinstance(value, str): + return token_counter(model=model, text=value) + if isinstance(value, list): + total = 0 + for chunk in value: + if isinstance(chunk, str): + total += token_counter(model=model, text=chunk) + elif isinstance(chunk, int): + # Single pre-tokenized prompt at the top level: each + # int counts as one token. + total += 1 + elif isinstance(chunk, list): + # Nested pre-tokenized prompt: every int contributes a + # token. Mixed string/int items still count. + total += sum(1 if isinstance(t, int) else 0 for t in chunk) + total += sum( + token_counter(model=model, text=t) + for t in chunk + if isinstance(t, str) + ) + return total + return 0 + + def _get_batch_job_usage_from_response_body(response_body: dict) -> Usage: """ Get the tokens of a batch job from the response body diff --git a/litellm/constants.py b/litellm/constants.py index 6c889a317b8..334ef8d48a4 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -419,9 +419,6 @@ CACHED_STREAMING_CHUNK_DELAY = float(os.getenv("CACHED_STREAMING_CHUNK_DELAY", 0 AUDIO_SPEECH_CHUNK_SIZE = int( os.getenv("AUDIO_SPEECH_CHUNK_SIZE", 8192) ) # chunk_size for audio speech streaming. Balance between latency and memory usage -MAX_SIZE_PER_ITEM_IN_MEMORY_CACHE_IN_KB = int( - os.getenv("MAX_SIZE_PER_ITEM_IN_MEMORY_CACHE_IN_KB", 512) -) DEFAULT_MAX_TOKENS_FOR_TRITON = int(os.getenv("DEFAULT_MAX_TOKENS_FOR_TRITON", 2000)) #### Networking settings #### # Sentinel used when `REQUEST_TIMEOUT` is unset: `litellm.request_timeout` keeps this diff --git a/litellm/cost_calculator.py b/litellm/cost_calculator.py index 8a68d74be5b..9b4dd80265c 100644 --- a/litellm/cost_calculator.py +++ b/litellm/cost_calculator.py @@ -513,7 +513,10 @@ def cost_per_token( # noqa: PLR0915 return fireworks_ai_cost_per_token(model=model, usage=usage_block) elif custom_llm_provider == "azure": return azure_openai_cost_per_token( - model=model, usage=usage_block, response_time_ms=response_time_ms + model=model, + usage=usage_block, + response_time_ms=response_time_ms, + service_tier=service_tier, ) elif custom_llm_provider == "gemini": return gemini_cost_per_token( @@ -539,6 +542,7 @@ def cost_per_token( # noqa: PLR0915 usage=usage_block, response_time_ms=response_time_ms, request_model=request_model, + service_tier=service_tier, ) else: model_info = _cached_get_model_info_helper( diff --git a/litellm/integrations/arize/_utils.py b/litellm/integrations/arize/_utils.py index 8dfaa8b1425..a1bf65141c9 100644 --- a/litellm/integrations/arize/_utils.py +++ b/litellm/integrations/arize/_utils.py @@ -220,23 +220,57 @@ def _set_structured_outputs(span: "Span", response_obj, msg_attrs, span_attrs): safe_set_attribute(span, f"{prefix}.{msg_attrs.MESSAGE_ROLE}", message_role) +def _safe_get(obj, key, default=None): + """Read ``key`` from a dict-like or Pydantic-model-like object. + + The arize/langfuse_otel logger receives ``usage`` objects from many sources: + plain dicts, litellm ``Usage`` (which exposes ``.get``), and raw OpenAI + Pydantic models (e.g. ``openai.types.completion_usage.CompletionUsage`` and + nested ``CompletionTokensDetails`` / ``OutputTokensDetails``) which do NOT + expose ``.get``. Calling ``.get`` on the latter raised ``AttributeError`` — + see https://github.com/BerriAI/litellm/issues/13672. + """ + if obj is None: + return default + getter = getattr(obj, "get", None) + if callable(getter): + try: + return getter(key, default) + except TypeError: + # Some objects expose `.get` with a different signature + pass + return getattr(obj, key, default) + + def _set_usage_outputs(span: "Span", response_obj, span_attrs): usage = response_obj and response_obj.get("usage") if not usage: return safe_set_attribute( - span, span_attrs.LLM_TOKEN_COUNT_TOTAL, usage.get("total_tokens") + span, span_attrs.LLM_TOKEN_COUNT_TOTAL, _safe_get(usage, "total_tokens") + ) + completion_tokens = _safe_get(usage, "completion_tokens") or _safe_get( + usage, "output_tokens" ) - completion_tokens = usage.get("completion_tokens") or usage.get("output_tokens") if completion_tokens: safe_set_attribute( span, span_attrs.LLM_TOKEN_COUNT_COMPLETION, completion_tokens ) - prompt_tokens = usage.get("prompt_tokens") or usage.get("input_tokens") + prompt_tokens = _safe_get(usage, "prompt_tokens") or _safe_get( + usage, "input_tokens" + ) if prompt_tokens: safe_set_attribute(span, span_attrs.LLM_TOKEN_COUNT_PROMPT, prompt_tokens) - reasoning_tokens = usage.get("output_tokens_details", {}).get("reasoning_tokens") + + # Reasoning tokens live in `completion_tokens_details` for Chat Completions + # API (Usage) and in `output_tokens_details` for Responses API + # (ResponseAPIUsage). Both nested objects may be plain Pydantic models + # without `.get`. + token_details = _safe_get(usage, "completion_tokens_details") or _safe_get( + usage, "output_tokens_details" + ) + reasoning_tokens = _safe_get(token_details, "reasoning_tokens") if reasoning_tokens: safe_set_attribute( span, diff --git a/litellm/integrations/custom_sso_handler.py b/litellm/integrations/custom_sso_handler.py index 7f60decabc3..202e488e0e4 100644 --- a/litellm/integrations/custom_sso_handler.py +++ b/litellm/integrations/custom_sso_handler.py @@ -18,6 +18,17 @@ class CustomSSOLoginHandler(CustomLogger): self, request: Request, ) -> OpenID: + from litellm.proxy.auth.trusted_proxy_utils import ( + require_trusted_proxy_request, + ) + from litellm.proxy.proxy_server import general_settings + + require_trusted_proxy_request( + request=request, + general_settings=general_settings, + feature_name="Custom UI SSO", + ) + request_headers_dict = dict(request.headers) return OpenID( id=request_headers_dict.get("x-litellm-user-id"), diff --git a/litellm/integrations/langfuse/langfuse.py b/litellm/integrations/langfuse/langfuse.py index e691c490c85..0efc7d66876 100644 --- a/litellm/integrations/langfuse/langfuse.py +++ b/litellm/integrations/langfuse/langfuse.py @@ -90,6 +90,29 @@ def _extract_cache_read_input_tokens(usage_obj) -> int: return cache_read_input_tokens +def resolve_langfuse_credentials( + langfuse_public_key=None, + langfuse_secret=None, + langfuse_secret_key=None, + langfuse_host=None, + allow_env_credentials: bool = True, +): + if allow_env_credentials is False and langfuse_host is not None: + secret_key = langfuse_secret or langfuse_secret_key + public_key = langfuse_public_key + else: + secret_key = ( + langfuse_secret or langfuse_secret_key or os.getenv("LANGFUSE_SECRET_KEY") + ) + public_key = langfuse_public_key or os.getenv("LANGFUSE_PUBLIC_KEY") + + resolved_host = langfuse_host or os.getenv( + "LANGFUSE_HOST", "https://cloud.langfuse.com" + ) + + return public_key, secret_key, resolved_host + + class LangFuseLogger: # Class variables or attributes def __init__( @@ -98,6 +121,7 @@ class LangFuseLogger: langfuse_secret=None, langfuse_host=None, flush_interval=1, + allow_env_credentials: bool = True, ): try: import langfuse @@ -106,11 +130,13 @@ class LangFuseLogger: raise Exception( f"\033[91mLangfuse not installed, try running 'pip install langfuse' to fix this error: {e}\n{traceback.format_exc()}\033[0m" ) - # Instance variables - self.secret_key = langfuse_secret or os.getenv("LANGFUSE_SECRET_KEY") - self.public_key = langfuse_public_key or os.getenv("LANGFUSE_PUBLIC_KEY") - self.langfuse_host = langfuse_host or os.getenv( - "LANGFUSE_HOST", "https://cloud.langfuse.com" + self.public_key, self.secret_key, self.langfuse_host = ( + resolve_langfuse_credentials( + langfuse_public_key=langfuse_public_key, + langfuse_secret=langfuse_secret, + langfuse_host=langfuse_host, + allow_env_credentials=allow_env_credentials, + ) ) if not ( self.langfuse_host.startswith("http://") @@ -160,9 +186,10 @@ class LangFuseLogger: project_id = None if os.getenv("UPSTREAM_LANGFUSE_SECRET_KEY") is not None: + upstream_langfuse_debug_env = os.getenv("UPSTREAM_LANGFUSE_DEBUG") upstream_langfuse_debug = ( - str_to_bool(self.upstream_langfuse_debug) - if self.upstream_langfuse_debug is not None + str_to_bool(upstream_langfuse_debug_env) + if upstream_langfuse_debug_env is not None else None ) self.upstream_langfuse_secret_key = os.getenv( @@ -173,7 +200,7 @@ class LangFuseLogger: ) self.upstream_langfuse_host = os.getenv("UPSTREAM_LANGFUSE_HOST") self.upstream_langfuse_release = os.getenv("UPSTREAM_LANGFUSE_RELEASE") - self.upstream_langfuse_debug = os.getenv("UPSTREAM_LANGFUSE_DEBUG") + self.upstream_langfuse_debug = upstream_langfuse_debug_env self.upstream_langfuse = Langfuse( public_key=self.upstream_langfuse_public_key, secret_key=self.upstream_langfuse_secret_key, diff --git a/litellm/integrations/langfuse/langfuse_handler.py b/litellm/integrations/langfuse/langfuse_handler.py index fbadf1a2fc7..4a809726424 100644 --- a/litellm/integrations/langfuse/langfuse_handler.py +++ b/litellm/integrations/langfuse/langfuse_handler.py @@ -115,8 +115,10 @@ class LangFuseHandler: langfuse_logger = LangFuseLogger( langfuse_public_key=credentials.get("langfuse_public_key"), - langfuse_secret=credentials.get("langfuse_secret"), + langfuse_secret=credentials.get("langfuse_secret") + or credentials.get("langfuse_secret_key"), langfuse_host=credentials.get("langfuse_host"), + allow_env_credentials=credentials.get("langfuse_host") is None, ) in_memory_dynamic_logger_cache.set_cache( credentials=credentials, diff --git a/litellm/integrations/langfuse/langfuse_prompt_management.py b/litellm/integrations/langfuse/langfuse_prompt_management.py index 5f4ced3a5cb..b7a565512c6 100644 --- a/litellm/integrations/langfuse/langfuse_prompt_management.py +++ b/litellm/integrations/langfuse/langfuse_prompt_management.py @@ -20,7 +20,7 @@ from ...litellm_core_utils.specialty_caches.dynamic_logging_cache import ( DynamicLoggingCache, ) from ..prompt_management_base import PromptManagementBase -from .langfuse import LangFuseLogger +from .langfuse import LangFuseLogger, resolve_langfuse_credentials from .langfuse_handler import LangFuseHandler if TYPE_CHECKING: @@ -46,6 +46,7 @@ def langfuse_client_init( langfuse_secret_key=None, langfuse_host=None, flush_interval=1, + allow_env_credentials: bool = True, ) -> LangfuseClass: """ Initialize Langfuse client with caching to prevent multiple initializations. @@ -70,14 +71,12 @@ def langfuse_client_init( f"\033[91mLangfuse not installed, try running 'pip install langfuse' to fix this error: {e}\n\033[0m" ) - # Instance variables - - secret_key = ( - langfuse_secret or langfuse_secret_key or os.getenv("LANGFUSE_SECRET_KEY") - ) - public_key = langfuse_public_key or os.getenv("LANGFUSE_PUBLIC_KEY") - langfuse_host = langfuse_host or os.getenv( - "LANGFUSE_HOST", "https://cloud.langfuse.com" + public_key, secret_key, langfuse_host = resolve_langfuse_credentials( + langfuse_public_key=langfuse_public_key, + langfuse_secret=langfuse_secret, + langfuse_secret_key=langfuse_secret_key, + langfuse_host=langfuse_host, + allow_env_credentials=allow_env_credentials, ) if not ( @@ -222,6 +221,7 @@ class LangfusePromptManagement(LangFuseLogger, PromptManagementBase, CustomLogge langfuse_secret=dynamic_callback_params.get("langfuse_secret"), langfuse_secret_key=dynamic_callback_params.get("langfuse_secret_key"), langfuse_host=dynamic_callback_params.get("langfuse_host"), + allow_env_credentials=dynamic_callback_params.get("langfuse_host") is None, ) langfuse_prompt_client = self._get_prompt_from_id( langfuse_prompt_id=prompt_id, @@ -246,6 +246,7 @@ class LangfusePromptManagement(LangFuseLogger, PromptManagementBase, CustomLogge langfuse_secret=dynamic_callback_params.get("langfuse_secret"), langfuse_secret_key=dynamic_callback_params.get("langfuse_secret_key"), langfuse_host=dynamic_callback_params.get("langfuse_host"), + allow_env_credentials=dynamic_callback_params.get("langfuse_host") is None, ) langfuse_prompt_client = self._get_prompt_from_id( langfuse_prompt_id=prompt_id, diff --git a/litellm/integrations/langsmith.py b/litellm/integrations/langsmith.py index 3d4fd39ebe1..81570e462c4 100644 --- a/litellm/integrations/langsmith.py +++ b/litellm/integrations/langsmith.py @@ -19,6 +19,7 @@ from litellm.integrations.langsmith_mock_client import ( create_mock_langsmith_client, should_use_langsmith_mock, ) +from litellm.litellm_core_utils.redact_messages import redact_user_api_key_info from litellm.llms.custom_httpx.http_handler import ( get_async_httpx_client, httpxSpecialProvider, @@ -112,17 +113,28 @@ class LangsmithLogger(CustomBatchLogger): langsmith_project: Optional[str] = None, langsmith_base_url: Optional[str] = None, langsmith_tenant_id: Optional[str] = None, + allow_env_credentials: bool = True, ) -> LangsmithCredentialsObject: - _credentials_api_key = langsmith_api_key or os.getenv("LANGSMITH_API_KEY") - _credentials_project = ( - langsmith_project or os.getenv("LANGSMITH_PROJECT") or "litellm-completion" - ) - _credentials_base_url = ( - langsmith_base_url - or os.getenv("LANGSMITH_BASE_URL") - or "https://api.smith.langchain.com" - ) - _credentials_tenant_id = langsmith_tenant_id or os.getenv("LANGSMITH_TENANT_ID") + if allow_env_credentials is False and langsmith_base_url is not None: + _credentials_api_key = langsmith_api_key + _credentials_project = langsmith_project or "litellm-completion" + _credentials_base_url = langsmith_base_url + _credentials_tenant_id = langsmith_tenant_id + else: + _credentials_api_key = langsmith_api_key or os.getenv("LANGSMITH_API_KEY") + _credentials_project = ( + langsmith_project + or os.getenv("LANGSMITH_PROJECT") + or "litellm-completion" + ) + _credentials_base_url = ( + langsmith_base_url + or os.getenv("LANGSMITH_BASE_URL") + or "https://api.smith.langchain.com" + ) + _credentials_tenant_id = langsmith_tenant_id or os.getenv( + "LANGSMITH_TENANT_ID" + ) return LangsmithCredentialsObject( LANGSMITH_API_KEY=_credentials_api_key, @@ -153,6 +165,15 @@ class LangsmithLogger(CustomBatchLogger): for key in ("session_id", "thread_id", "conversation_id"): if key in requester_metadata and key not in extra_metadata: extra_metadata[key] = requester_metadata[key] + + # helper is shallow; also scrub nested requester_metadata since + # LangSmith forwards the whole dict into `extra` + extra_metadata = redact_user_api_key_info(metadata=extra_metadata) + nested = extra_metadata.get("requester_metadata") + if isinstance(nested, dict): + extra_metadata["requester_metadata"] = redact_user_api_key_info( + metadata=nested + ) return extra_metadata def _build_outputs_with_usage( @@ -540,6 +561,10 @@ class LangsmithLogger(CustomBatchLogger): langsmith_tenant_id=standard_callback_dynamic_params.get( "langsmith_tenant_id", None ), + allow_env_credentials=standard_callback_dynamic_params.get( + "langsmith_base_url", None + ) + is None, ) else: credentials = self.default_credentials diff --git a/litellm/integrations/opentelemetry.py b/litellm/integrations/opentelemetry.py index b6d91d0b76d..77833e5de0f 100644 --- a/litellm/integrations/opentelemetry.py +++ b/litellm/integrations/opentelemetry.py @@ -69,6 +69,8 @@ class OpenTelemetryConfig: deployment_environment: Optional[str] = None model_id: Optional[str] = None ignore_context_propagation: Optional[bool] = None + # When True, create a private TracerProvider instead of reusing or setting the global one. + skip_set_global: bool = False def __post_init__(self) -> None: # If endpoint is specified but exporter is still the default "console", @@ -259,16 +261,21 @@ class OpenTelemetry(CustomLogger): try: existing_provider = get_existing_provider_fn() - # If a real SDK provider exists (set by another SDK like Langfuse), use it - # This uses a positive check for SDK providers instead of a negative check for proxy providers if isinstance(existing_provider, sdk_provider_class): - verbose_logger.debug( - "OpenTelemetry: Using existing %s: %s", - provider_name, - type(existing_provider).__name__, - ) - provider = existing_provider - # Don't call set_provider to preserve existing context + if skip_set_global: + verbose_logger.debug( + "OpenTelemetry: existing %s found but skip_set_global=True; creating private %s for isolation", + provider_name, + provider_name, + ) + provider = create_new_provider_fn() + else: + verbose_logger.debug( + "OpenTelemetry: Using existing %s: %s", + provider_name, + type(existing_provider).__name__, + ) + provider = existing_provider else: # Default proxy provider or unknown type, create our own verbose_logger.debug("OpenTelemetry: Creating new %s", provider_name) @@ -293,6 +300,12 @@ class OpenTelemetry(CustomLogger): return provider + def _skip_set_global(self) -> bool: + # langfuse_otel relies on the Langfuse SDK's providers; don't overwrite them. + return self.config.skip_set_global or ( + hasattr(self, "callback_name") and self.callback_name == "langfuse_otel" + ) + def _init_tracing(self, tracer_provider): from opentelemetry import trace from opentelemetry.sdk.trace import TracerProvider @@ -303,11 +316,6 @@ class OpenTelemetry(CustomLogger): provider.add_span_processor(self._get_span_processor()) return provider - # CRITICAL FIX: For Langfuse OTEL, skip setting global provider to prevent interference - skip_global = ( - hasattr(self, "callback_name") and self.callback_name == "langfuse_otel" - ) - tracer_provider = self._get_or_create_provider( provider=tracer_provider, provider_name="TracerProvider", @@ -315,16 +323,18 @@ class OpenTelemetry(CustomLogger): sdk_provider_class=TracerProvider, create_new_provider_fn=create_tracer_provider, set_provider_fn=trace.set_tracer_provider, - skip_set_global=skip_global, + skip_set_global=self._skip_set_global(), ) # Grab our tracer from the TracerProvider (not from global context) # This ensures we use the provided TracerProvider (e.g., for testing) self.tracer = tracer_provider.get_tracer(LITELLM_TRACER_NAME) + self._tracer_provider = tracer_provider self.span_kind = SpanKind def _init_metrics(self, meter_provider): if not self.config.enable_metrics: + self._meter_provider = None self._operation_duration_histogram = None self._token_usage_histogram = None self._cost_histogram = None @@ -350,7 +360,9 @@ class OpenTelemetry(CustomLogger): sdk_provider_class=MeterProvider, create_new_provider_fn=create_meter_provider, set_provider_fn=metrics.set_meter_provider, + skip_set_global=self._skip_set_global(), ) + self._meter_provider = meter_provider meter = meter_provider.get_meter(__name__) @@ -388,6 +400,7 @@ class OpenTelemetry(CustomLogger): def _init_logs(self, logger_provider): # nothing to do if events disabled if not self.config.enable_events: + self._logger_provider = None return from opentelemetry._logs import get_logger_provider, set_logger_provider @@ -404,13 +417,14 @@ class OpenTelemetry(CustomLogger): ) return provider - self._get_or_create_provider( + self._logger_provider = self._get_or_create_provider( provider=logger_provider, provider_name="LoggerProvider", get_existing_provider_fn=get_logger_provider, sdk_provider_class=OTLoggerProvider, create_new_provider_fn=create_logger_provider, set_provider_fn=set_logger_provider, + skip_set_global=self._skip_set_global(), ) def log_success_event(self, kwargs, response_obj, start_time, end_time): @@ -1073,7 +1087,7 @@ class OpenTelemetry(CustomLogger): # See: https://github.com/open-telemetry/opentelemetry-python/pull/4676 # TODO: Refactor to use the proper OTEL Logs API instead of directly creating SDK LogRecords - from opentelemetry._logs import SeverityNumber, get_logger + from opentelemetry._logs import SeverityNumber try: from opentelemetry.sdk._logs import ( # type: ignore[attr-defined] # OTEL < 1.39.0 @@ -1084,7 +1098,10 @@ class OpenTelemetry(CustomLogger): LogRecord as SdkLogRecord, # type: ignore[attr-defined] # OTEL >= 1.39.0 ) - otel_logger = get_logger(LITELLM_LOGGER_NAME) + # Resolve through the handler's own LoggerProvider (which may be a + # private one when skip_set_global=True) rather than the module-level + # get_logger() which always goes through the global provider. + otel_logger = self._logger_provider.get_logger(LITELLM_LOGGER_NAME) parent_ctx = span.get_span_context() provider = (kwargs.get("litellm_params") or {}).get( diff --git a/litellm/integrations/prometheus_helpers/prometheus_api.py b/litellm/integrations/prometheus_helpers/prometheus_api.py index b25da577237..0901d7b6801 100644 --- a/litellm/integrations/prometheus_helpers/prometheus_api.py +++ b/litellm/integrations/prometheus_helpers/prometheus_api.py @@ -2,6 +2,7 @@ Helper functions to query prometheus API """ +import json import time from datetime import datetime, timedelta from typing import Optional @@ -81,6 +82,24 @@ def is_prometheus_connected() -> bool: return False +def _quote_promql_string_literal(value: str) -> str: + """Render ``value`` as a PromQL double-quoted string literal. + + PromQL string literals follow Go's escape rules + (https://prometheus.io/docs/prometheus/latest/querying/basics/): a + backslash begins an escape sequence and a bare ``"`` ends the literal. + Without escaping, callers that accept arbitrary user-supplied values + (like the ``api_key`` filter on ``/global/spend/logs``) can inject extra + label matchers or selectors and read cross-tenant metrics. + + JSON's quoting rules are a strict subset of Go's, so ``json.dumps`` of + a Python string produces a literal Prometheus accepts: ``\\``, ``\\"``, + and the standard ``\\n`` / ``\\t`` / ``\\uNNNN`` control-character + escapes. The returned value already includes the surrounding quotes. + """ + return json.dumps(value, ensure_ascii=False) + + async def get_daily_spend_from_prometheus(api_key: Optional[str]): """ Expected Response Format: @@ -109,8 +128,11 @@ async def get_daily_spend_from_prometheus(api_key: Optional[str]): if api_key is None: query = "sum(delta(litellm_spend_metric_total[1d]))" else: + quoted_api_key = _quote_promql_string_literal(api_key) query = ( - f'sum(delta(litellm_spend_metric_total{{hashed_api_key="{api_key}"}}[1d]))' + "sum(delta(litellm_spend_metric_total{" + f"hashed_api_key={quoted_api_key}" + "}[1d]))" ) params = { diff --git a/litellm/litellm_core_utils/exception_mapping_utils.py b/litellm/litellm_core_utils/exception_mapping_utils.py index 5a7d4e33b6d..2c1d92920af 100644 --- a/litellm/litellm_core_utils/exception_mapping_utils.py +++ b/litellm/litellm_core_utils/exception_mapping_utils.py @@ -6,7 +6,8 @@ from typing import Any, Optional import httpx import litellm -from litellm._logging import _redact_string, verbose_logger +from litellm._logging import _ENABLE_SECRET_REDACTION, _redact_string, verbose_logger +from litellm.litellm_core_utils.secret_redaction import redact_string from litellm.types.utils import LlmProviders from ..exceptions import ( @@ -261,10 +262,18 @@ def exception_type( # type: ignore # noqa: PLR0915 original_exception=original_exception ) try: - error_str = str(original_exception) + error_str = ( + redact_string(str(original_exception)) + if _ENABLE_SECRET_REDACTION + else str(original_exception) + ) if model: if hasattr(original_exception, "message"): - error_str = str(original_exception.message) + error_str = ( + redact_string(str(original_exception.message)) + if _ENABLE_SECRET_REDACTION + else str(original_exception.message) + ) if isinstance(original_exception, BaseException): exception_type = type(original_exception).__name__ else: @@ -2431,7 +2440,8 @@ def exception_type( # type: ignore # noqa: PLR0915 else: raise APIConnectionError( message="{}\n{}".format( - str(original_exception), _redact_string(traceback.format_exc()) + str(original_exception), + _redact_string(traceback.format_exc()), ), llm_provider=custom_llm_provider, model=model, @@ -2461,7 +2471,8 @@ def exception_type( # type: ignore # noqa: PLR0915 raise e # it's already mapped raised_exc = APIConnectionError( message="{}\n{}".format( - original_exception, _redact_string(traceback.format_exc()) + original_exception, + _redact_string(traceback.format_exc()), ), llm_provider="", model="", diff --git a/litellm/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py index 829c1c9ca07..a815442c2f9 100644 --- a/litellm/litellm_core_utils/litellm_logging.py +++ b/litellm/litellm_core_utils/litellm_logging.py @@ -3242,10 +3242,15 @@ class Logging(LiteLLMLoggingBaseClass): ), langfuse_secret=self.standard_callback_dynamic_params.get( "langfuse_secret" - ), + ) + or self.standard_callback_dynamic_params.get("langfuse_secret_key"), langfuse_host=self.standard_callback_dynamic_params.get( "langfuse_host" ), + allow_env_credentials=self.standard_callback_dynamic_params.get( + "langfuse_host" + ) + is None, ) return langFuseLogger @@ -4720,7 +4725,7 @@ class StandardLoggingPayloadSetup: ): for key, value in litellm_params["metadata"].items(): # Skip non-serializable objects like UserAPIKeyAuth - if key == "user_api_key_auth": + if key in {"user_api_key_auth", "user_api_key_budget_reservation"}: continue merged_metadata[key] = value diff --git a/litellm/litellm_core_utils/secret_redaction.py b/litellm/litellm_core_utils/secret_redaction.py new file mode 100644 index 00000000000..5c4e3e3dacf --- /dev/null +++ b/litellm/litellm_core_utils/secret_redaction.py @@ -0,0 +1,81 @@ +""" +Credential/secret redaction utilities. + +This module owns the compiled regex and the public `redact_string` helper so +that any part of the codebase (logging, exception mapping, etc.) can scrub +secrets from strings without depending on the logging-configuration module. +""" + +import re +from typing import List + +_REDACTED = "REDACTED" + + +def _build_secret_patterns() -> "re.Pattern[str]": + patterns: List[str] = [ + # PEM private key / certificate blocks + r"-----BEGIN[A-Z \-]*PRIVATE KEY-----[\s\S]*?-----END[A-Z \-]*PRIVATE KEY-----", + # GCP OAuth2 access tokens (ya29.*) + r"\bya29\.[A-Za-z0-9_.~+/-]+", + # Credential %s formatting (space separator, no key= prefix) + r"(?:client_secret|azure_password|azure_username)\s+[^\s,'\"})\]{}>]+", + # AWS access key IDs + r"(?:AKIA|ASIA)[0-9A-Z]{16}", + # AWS secrets / session tokens / access key IDs (key=value) + r"(?:aws_secret_access_key|aws_session_token|aws_access_key_id)" + r"\s*[:=]\s*[A-Za-z0-9/+=]{20,}", + # Bearer tokens (OAuth, JWT, etc.) + r"Bearer\s+[A-Za-z0-9\-._~+/]{10,}=*", + # Basic auth headers + r"Basic\s+[A-Za-z0-9+/]{10,}={0,2}", + # OpenAI / Anthropic sk- prefixed keys + r"sk-[A-Za-z0-9\-_]{20,}", + # Generic api_key / api-key / apikey (handles 'key': 'value' dict repr) + r"(?:api[_-]?key)['\"]?\s*[:=]\s*['\"]?[^\s,'\"})\]{}>]{8,}", + # x-api-key / api-key header values (handles 'key': 'value' dict repr) + r"(?:x-api-key|api-key)['\"]?\s*[:=]\s*['\"]?[^\s,'\"})\]{}>]+", + # Anthropic internal header keys + r"x-ak-[A-Za-z0-9\-_]{20,}", + # Google API keys (bare key value) + r"AIza[0-9A-Za-z\-_]{35}", + # URL query-param key=VALUE (e.g. ?key=AIza... or &key=...) — catches the + # full "key=" fragment so the value is redacted regardless of format. + r"(?<=[?&])key=[^\s&'\"]{8,}", + # Password / secret params (handles key=value and 'key': 'value') + # Word boundary prevents O(n^2) backtracking on long word-char runs. + r"(?:^|(?<=\W))\w*(?:password|passwd|client_secret|secret_key|_secret)" + r"['\"]?\s*[:=]\s*['\"]?[^\s,'\"})\]{}>]+", + # Database connection string credentials (scheme://user:pass@host) + r"(?<=://)[^\s'\"]*:[^\s'\"@]+(?=@)", + # Databricks personal access tokens + r"dapi[0-9a-f]{32}", + # ── Key-name-based redaction ── + # Catches secrets inside dicts/config dumps by matching on the KEY name + # regardless of what the value looks like. + # e.g. 'master_key': 'any-value-here', "database_url": "postgres://..." + # private_key with PEM-aware value capture + r"""private_key['\"]?\s*[:=]\s*['\"]?(?:-----BEGIN[A-Z \-]*PRIVATE KEY-----[\s\S]*?-----END[A-Z \-]*PRIVATE KEY-----|[^\s,'\"})\]{}>]+)""", + r"(?:master_key|database_url|db_url|connection_string|" + r"signing_key|encryption_key|" + r"auth_token|access_token|refresh_token|" + r"slack_webhook_url|webhook_url|" + r"database_connection_string|" + r"huggingface_token|jwt_secret)" + r"""['\"]?\s*[:=]\s*['\"]?[^\s,'\"})\]{}>]+""", + # Raw JWTs (without Bearer prefix) + r"\beyJ[A-Za-z0-9_-]{10,}\.[A-Za-z0-9_-]+\.[A-Za-z0-9_-]*", + # Azure SAS tokens in URLs + r"[?&]sig=[A-Za-z0-9%+/=]+", + # Full JSON service-account blobs (single-line and multi-line) + r'\{[^{}]*"type"\s*:\s*"service_account"[^{}]*(?:\{[^{}]*\}[^{}]*)*\}', + ] + return re.compile("|".join(patterns), re.IGNORECASE) + + +_SECRET_RE = _build_secret_patterns() + + +def redact_string(value: str) -> str: + """Scrub known secret/credential patterns from *value* and return the result.""" + return _SECRET_RE.sub(_REDACTED, value) diff --git a/litellm/llms/anthropic/chat/transformation.py b/litellm/llms/anthropic/chat/transformation.py index 61ddf801a25..35624c93b37 100644 --- a/litellm/llms/anthropic/chat/transformation.py +++ b/litellm/llms/anthropic/chat/transformation.py @@ -1088,24 +1088,29 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): elif param == "thinking": optional_params["thinking"] = value elif param == "reasoning_effort" and isinstance(value, str): - optional_params["thinking"] = AnthropicConfig._map_reasoning_effort( + mapped_thinking = AnthropicConfig._map_reasoning_effort( reasoning_effort=value, model=model ) - # For Claude 4.6+ models, effort is controlled via output_config, - # not thinking budget_tokens. Map reasoning_effort to output_config. - if AnthropicConfig._is_claude_4_6_model( - model - ) or AnthropicConfig._is_claude_4_7_model(model): - effort_map = { - "low": "low", - "minimal": "low", - "medium": "medium", - "high": "high", - "xhigh": "xhigh", - "max": "max", - } - mapped_effort = effort_map.get(value, value) - optional_params["output_config"] = {"effort": mapped_effort} + if mapped_thinking is None: + optional_params.pop("thinking", None) + optional_params.pop("output_config", None) + else: + optional_params["thinking"] = mapped_thinking + # For Claude 4.6+ models, effort is controlled via output_config, + # not thinking budget_tokens. Map reasoning_effort to output_config. + if AnthropicConfig._is_claude_4_6_model( + model + ) or AnthropicConfig._is_claude_4_7_model(model): + effort_map = { + "low": "low", + "minimal": "low", + "medium": "medium", + "high": "high", + "xhigh": "xhigh", + "max": "max", + } + mapped_effort = effort_map.get(value, value) + optional_params["output_config"] = {"effort": mapped_effort} elif param == "web_search_options" and isinstance(value, dict): hosted_web_search_tool = self.map_web_search_tool( cast(OpenAIWebSearchOptions, value) diff --git a/litellm/llms/anthropic/experimental_pass_through/adapters/handler.py b/litellm/llms/anthropic/experimental_pass_through/adapters/handler.py index 829ce14d69d..8ed6126d2eb 100644 --- a/litellm/llms/anthropic/experimental_pass_through/adapters/handler.py +++ b/litellm/llms/anthropic/experimental_pass_through/adapters/handler.py @@ -27,6 +27,16 @@ from litellm.utils import get_model_info if TYPE_CHECKING: pass + +# Anthropic-only fields that the translator above already maps into the +# OpenAI-format completion_kwargs (output_config → reasoning_effort / +# response_format, etc.). They must be filtered out of the raw +# extra_kwargs re-merge below or non-Anthropic backends reject the call +# with 400 "Extra inputs are not permitted". Add new entries here when +# extending AnthropicMessagesRequestOptionalParams with another Anthropic- +# specific key. +ANTHROPIC_ONLY_REQUEST_KEYS: frozenset[str] = frozenset({"output_config"}) + ######################################################## # init adapter ANTHROPIC_ADAPTER = AnthropicAdapter() @@ -202,8 +212,12 @@ class LiteLLMMessagesToCompletionTransformationHandler: request_data["output_format"] = output_format # Extract output_config from extra_kwargs so the translator can use it - # (e.g. output_config.effort for adaptive thinking → reasoning_effort) - extra_kwargs = extra_kwargs or {} + # (e.g. output_config.effort for adaptive thinking → reasoning_effort, + # output_config.format → response_format for structured outputs). + # Use explicit None check rather than `or {}` so an explicit empty dict + # caller-passed argument is preserved (matters for tests that drive + # the fallback inference path). + extra_kwargs = extra_kwargs if extra_kwargs is not None else {} if "output_config" in extra_kwargs: request_data["output_config"] = extra_kwargs["output_config"] @@ -225,8 +239,23 @@ class LiteLLMMessagesToCompletionTransformationHandler: "include_usage": True, } - excluded_keys = {"anthropic_messages"} - extra_kwargs = extra_kwargs or {} + # Keys that must NOT be forwarded as raw extras into the OpenAI-format + # ``completion_kwargs`` after translation. The translator above has + # already consumed the meaningful parts of these inputs (e.g. + # ``output_config.format`` → ``response_format``, ``output_config.effort`` + # → ``reasoning_effort`` for non-Claude targets). Re-adding the raw + # Anthropic-shaped key here causes 400 "Extra inputs are not permitted" + # on non-Anthropic backends (Azure OpenAI, Fireworks, Bedrock Nova, + # etc.) and is silently lossy on Anthropic-family targets, which would + # see the translated key ``response_format`` AND a duplicate, conflicting + # ``output_config``. + # + # Maintainability: when adding a new Anthropic-only request param to + # ``AnthropicMessagesRequestOptionalParams``, also extend + # ``ANTHROPIC_ONLY_REQUEST_KEYS`` here so it doesn't silently leak. + excluded_keys = ANTHROPIC_ONLY_REQUEST_KEYS | {"anthropic_messages"} + # NOTE: extra_kwargs was already coerced from None to {} at the top of + # this method (line ~220). It is guaranteed to be a dict here. for key, value in extra_kwargs.items(): if ( key == "litellm_logging_obj" diff --git a/litellm/llms/anthropic/experimental_pass_through/adapters/transformation.py b/litellm/llms/anthropic/experimental_pass_through/adapters/transformation.py index 08797889192..fe8e694efe5 100644 --- a/litellm/llms/anthropic/experimental_pass_through/adapters/transformation.py +++ b/litellm/llms/anthropic/experimental_pass_through/adapters/transformation.py @@ -667,7 +667,7 @@ class LiteLLMAnthropicMessagesAdapter: @staticmethod def translate_anthropic_thinking_to_reasoning_effort( - thinking: Dict[str, Any] + thinking: Dict[str, Any], ) -> Optional[str]: """ Translate Anthropic's thinking parameter to OpenAI's reasoning_effort. @@ -1084,10 +1084,23 @@ class LiteLLMAnthropicMessagesAdapter: anthropic_message_request: AnthropicMessagesRequest, new_kwargs: ChatCompletionRequest, ) -> None: - """Translate output_format to response_format when applicable.""" - if "output_format" not in anthropic_message_request: - return - output_format = anthropic_message_request["output_format"] + """Translate Anthropic structured-output config to OpenAI ``response_format``. + + Accepts either the legacy top-level ``output_format`` field OR the + newer ``output_config.format`` (sub-key on ``output_config``) so that + both shapes flow through to non-Anthropic backends as + ``response_format``. Without the ``output_config.format`` branch, + callers using the new Anthropic Structured Outputs API would have + their schema silently dropped on the adapter path — only the legacy + top-level ``output_format`` was being mapped. + + ``output_format`` takes precedence when both are provided. + """ + output_format: Any = anthropic_message_request.get("output_format") + if not output_format: + output_config = anthropic_message_request.get("output_config") + if isinstance(output_config, dict): + output_format = output_config.get("format") if not output_format: return response_format = self.translate_anthropic_output_format_to_openai( diff --git a/litellm/llms/azure/azure.py b/litellm/llms/azure/azure.py index 877a7d3c84a..c0e070b6c1f 100644 --- a/litellm/llms/azure/azure.py +++ b/litellm/llms/azure/azure.py @@ -793,6 +793,7 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM): client=client, litellm_params=litellm_params, api_base=api_base, + api_version=api_version, ) azure_client = self.get_azure_openai_client( api_version=api_version, diff --git a/litellm/llms/azure/cost_calculation.py b/litellm/llms/azure/cost_calculation.py index 5b411095ea1..2a20c55a6ce 100644 --- a/litellm/llms/azure/cost_calculation.py +++ b/litellm/llms/azure/cost_calculation.py @@ -12,7 +12,10 @@ from litellm.utils import get_model_info def cost_per_token( - model: str, usage: Usage, response_time_ms: Optional[float] = 0.0 + model: str, + usage: Usage, + response_time_ms: Optional[float] = 0.0, + service_tier: Optional[str] = None, ) -> Tuple[float, float]: """ Calculates the cost per token for a given model, prompt tokens, and completion tokens. @@ -47,4 +50,5 @@ def cost_per_token( model=model, usage=usage, custom_llm_provider="azure", + service_tier=service_tier, ) diff --git a/litellm/llms/azure_ai/cost_calculator.py b/litellm/llms/azure_ai/cost_calculator.py index 067181b946a..755d44fdef7 100644 --- a/litellm/llms/azure_ai/cost_calculator.py +++ b/litellm/llms/azure_ai/cost_calculator.py @@ -65,6 +65,7 @@ def cost_per_token( usage: Usage, response_time_ms: Optional[float] = 0.0, request_model: Optional[str] = None, + service_tier: Optional[str] = None, ) -> Tuple[float, float]: """ Calculate the cost per token for Azure AI models. @@ -102,6 +103,7 @@ def cost_per_token( model=model, usage=usage, custom_llm_provider="azure_ai", + service_tier=service_tier, ) except Exception as e: # For Model Router, the model name (e.g., "azure-model-router") may not be in the cost map diff --git a/litellm/llms/bedrock/chat/converse_transformation.py b/litellm/llms/bedrock/chat/converse_transformation.py index 61a7d4c08db..8b8500ac060 100644 --- a/litellm/llms/bedrock/chat/converse_transformation.py +++ b/litellm/llms/bedrock/chat/converse_transformation.py @@ -449,9 +449,13 @@ class AmazonConverseConfig(BaseConfig): optional_params.update(reasoning_config) else: # Anthropic and other models: convert to thinking parameter - optional_params["thinking"] = AnthropicConfig._map_reasoning_effort( + mapped_thinking = AnthropicConfig._map_reasoning_effort( reasoning_effort=reasoning_effort, model=model ) + if mapped_thinking is None: + optional_params.pop("thinking", None) + else: + optional_params["thinking"] = mapped_thinking @staticmethod def _clamp_thinking_budget_tokens(optional_params: dict) -> None: diff --git a/litellm/llms/cloudflare/chat/transformation.py b/litellm/llms/cloudflare/chat/transformation.py index b9e219f5cbc..66e253f304d 100644 --- a/litellm/llms/cloudflare/chat/transformation.py +++ b/litellm/llms/cloudflare/chat/transformation.py @@ -149,9 +149,9 @@ class CloudflareChatConfig(BaseConfig): ) -> ModelResponse: completion_response = raw_response.json() - model_response.choices[0].message.content = completion_response["result"][ # type: ignore - "response" - ] + # Support both "response" and "response_text" keys (newer models like Nemotron use "response_text") + result = completion_response["result"] + model_response.choices[0].message.content = result.get("response") if result.get("response") is not None else result.get("response_text", "") # type: ignore prompt_tokens = litellm.utils.get_token_count(messages=messages, model=model) completion_tokens = len( @@ -201,8 +201,10 @@ class CloudflareChatResponseIterator(BaseModelResponseIterator): index = int(chunk.get("index", 0)) - if "response" in chunk: + if "response" in chunk and chunk["response"] is not None: text = chunk["response"] + elif "response_text" in chunk and chunk["response_text"] is not None: + text = chunk["response_text"] returned_chunk = GenericStreamingChunk( text=text, diff --git a/litellm/llms/hosted_vllm/embedding/README.md b/litellm/llms/hosted_vllm/embedding/README.md index f82b3c77a6e..2c58e16fc23 100644 --- a/litellm/llms/hosted_vllm/embedding/README.md +++ b/litellm/llms/hosted_vllm/embedding/README.md @@ -2,4 +2,15 @@ No transformation is required for hosted_vllm embedding. VLLM is a superset of OpenAI's `embedding` endpoint. -To pass provider-specific parameters, see [this](https://docs.litellm.ai/docs/completion/provider_specific_params) \ No newline at end of file +## `encoding_format` + +For OpenAI-compatible embedding calls (including `openai/...` with a custom `api_base` pointing at vLLM), LiteLLM resolves `encoding_format` when it is not set on the request: + +1. Explicit value on the embedding call (`encoding_format=...`). +2. Model config (`litellm_params.encoding_format` on the proxy `model_list` entry). +3. Environment variable `LITELLM_DEFAULT_EMBEDDING_ENCODING_FORMAT` (e.g. in `.env` or container env). +4. Default **`float`**. + +That avoids forwarding `encoding_format=None` to the provider/SDK where some servers behave poorly. + +To pass provider-specific parameters, see [provider-specific params](https://docs.litellm.ai/docs/completion/provider_specific_params). \ No newline at end of file diff --git a/litellm/llms/openai/chat/gpt_5_transformation.py b/litellm/llms/openai/chat/gpt_5_transformation.py index 34941a545eb..4e34d10b187 100644 --- a/litellm/llms/openai/chat/gpt_5_transformation.py +++ b/litellm/llms/openai/chat/gpt_5_transformation.py @@ -244,9 +244,11 @@ class OpenAIGPT5Config(OpenAIGPTConfig): ), status_code=400, ) - elif effective_effort == "minimal": - # minimal is opt-out: unknown models pass through; only block when - # the model map explicitly sets supports_minimal_reasoning_effort=false. + elif effective_effort in ("minimal", "low"): + # minimal/low are opt-out: unknown models pass through; only block when + # the model map explicitly sets supports_{level}_reasoning_effort=false. + # Example: gpt-5.5-pro only accepts {medium, high, xhigh}, so it sets + # supports_low_reasoning_effort=false (and supports_minimal=false). if self._is_reasoning_effort_level_explicitly_disabled( model, effective_effort ): diff --git a/litellm/llms/openai_like/providers.json b/litellm/llms/openai_like/providers.json index 5dd1247001e..b5e5aa4ea28 100644 --- a/litellm/llms/openai_like/providers.json +++ b/litellm/llms/openai_like/providers.json @@ -106,5 +106,13 @@ "base_url": "https://aihubmix.com/v1", "api_key_env": "AIHUBMIX_API_KEY", "api_base_env": "AIHUBMIX_API_BASE" + }, + "crusoe": { + "base_url": "https://managed-inference-api-proxy.crusoecloud.com/v1", + "api_key_env": "CRUSOE_API_KEY", + "api_base_env": "CRUSOE_API_BASE", + "param_mappings": { + "max_completion_tokens": "max_tokens" + } } } diff --git a/litellm/llms/vertex_ai/common_utils.py b/litellm/llms/vertex_ai/common_utils.py index c72160f7d0a..e6e39651109 100644 --- a/litellm/llms/vertex_ai/common_utils.py +++ b/litellm/llms/vertex_ai/common_utils.py @@ -656,6 +656,8 @@ def process_items(schema, depth=0): and ("items" not in schema or schema.get("items") == {}) ): schema["items"] = {"type": "object"} + elif schema.get("type") == "array" and "items" not in schema: + schema["items"] = {"type": "object"} for key, value in schema.items(): if isinstance(value, dict): process_items(value, depth + 1) diff --git a/litellm/llms/vertex_ai/files/handler.py b/litellm/llms/vertex_ai/files/handler.py index 1a92d521065..c31bfde69e7 100644 --- a/litellm/llms/vertex_ai/files/handler.py +++ b/litellm/llms/vertex_ai/files/handler.py @@ -1,4 +1,5 @@ import asyncio +import time from urllib.parse import unquote from typing import Any, Coroutine, Optional, Tuple, Union @@ -21,9 +22,10 @@ from litellm.types.llms.openai import ( HttpxBinaryResponseContent, OpenAIFileObject, ) +from litellm.litellm_core_utils.litellm_logging import Logging from litellm.types.llms.vertex_ai import VERTEX_CREDENTIALS_TYPES -from .transformation import VertexAIJsonlFilesTransformation +from .transformation import VertexAIFilesConfig, VertexAIJsonlFilesTransformation vertex_ai_files_transformation = VertexAIJsonlFilesTransformation() @@ -198,11 +200,30 @@ class VertexAIFilesHandler(GCSBucketBase): mock_response = httpx.Response( status_code=200, content=file_content, - headers={"content-type": "application/octet-stream"}, + headers={ + "content-type": "application/octet-stream", + "content-length": str(len(file_content)), + }, request=httpx.Request(method="GET", url=decoded_file_id), ) - return HttpxBinaryResponseContent(response=mock_response) + # Apply transformation to convert Vertex AI batch outputs to OpenAI format + config = VertexAIFilesConfig() + + # Create a logging object for transformation + logging_obj = Logging( + model="", + messages=[], + stream=False, + call_type="afile_content", + start_time=time.time(), + litellm_call_id="", + function_id="", + ) + + return config.transform_file_content_response( + raw_response=mock_response, logging_obj=logging_obj, litellm_params={} + ) def file_content( self, diff --git a/litellm/llms/vertex_ai/files/transformation.py b/litellm/llms/vertex_ai/files/transformation.py index 85663bd551c..f30518bc7ca 100644 --- a/litellm/llms/vertex_ai/files/transformation.py +++ b/litellm/llms/vertex_ai/files/transformation.py @@ -1,10 +1,15 @@ +import base64 import json import os -from typing import Any, Dict, List, Optional, Tuple, Union +import re +import time +from typing import Any, Callable, Dict, List, Optional, Tuple, Union +import httpx from httpx import Headers, Response from openai.types.file_deleted import FileDeleted +import litellm from litellm._uuid import uuid from litellm.files.utils import FilesAPIUtils from litellm.litellm_core_utils.cloud_storage_security import ( @@ -16,6 +21,7 @@ from litellm.litellm_core_utils.cloud_storage_security import ( split_configured_cloud_bucket_name, validate_managed_cloud_file_id, ) +from litellm.litellm_core_utils.litellm_logging import Logging from litellm.litellm_core_utils.prompt_templates.common_utils import extract_file_data from litellm.llms.base_llm.chat.transformation import BaseLLMException from litellm.llms.base_llm.files.transformation import ( @@ -39,11 +45,135 @@ from litellm.types.llms.openai import ( PathLike, ) from litellm.types.llms.vertex_ai import GcsBucketResponse -from litellm.types.utils import ExtractedFileData, LlmProviders +from litellm.types.utils import ExtractedFileData, LlmProviders, ModelResponse from ..common_utils import VertexAIError from ..vertex_llm_base import VertexBase +_GCP_LABEL_VALUE_MAX_LEN = 63 +_CUSTOM_ID_RAW_LABEL_PREFIX = "b32_" + + +def _sanitize_gcp_label_value(value: str) -> str: + """ + Sanitize a string to meet GCP label value constraints. + + GCP label values must: + - Be lowercase + - Contain only letters, numbers, underscores, and hyphens + - Be max 63 characters + + Args: + value: The string to sanitize + + Returns: + A sanitized string that meets GCP label constraints + """ + sanitized = re.sub(r"[^a-z0-9_-]", "_", value.lower()) + return sanitized[:_GCP_LABEL_VALUE_MAX_LEN] + + +def _encode_gcp_label_value_chunks(value: str) -> List[str]: + """Encode arbitrary text across one or more GCP-label-safe values.""" + max_encoded_len = _GCP_LABEL_VALUE_MAX_LEN - len(_CUSTOM_ID_RAW_LABEL_PREFIX) + encoded = ( + base64.b32encode(value.encode("utf-8")).decode("ascii").rstrip("=").lower() + ) + return [ + f"{_CUSTOM_ID_RAW_LABEL_PREFIX}{encoded[i : i + max_encoded_len]}" + for i in range(0, len(encoded), max_encoded_len) + ] or [_CUSTOM_ID_RAW_LABEL_PREFIX] + + +def _decode_gcp_label_value_chunks(values: List[str]) -> Optional[str]: + """Decode values produced by _encode_gcp_label_value_chunks.""" + encoded_parts = [] + for value in values: + if not value.startswith(_CUSTOM_ID_RAW_LABEL_PREFIX): + return None + encoded_parts.append(value[len(_CUSTOM_ID_RAW_LABEL_PREFIX) :]) + encoded = "".join(encoded_parts).upper() + padding = "=" * (-len(encoded) % 8) + try: + return base64.b32decode(encoded + padding).decode("utf-8") + except Exception: + return None + + +def _set_litellm_batch_custom_id_labels(labels: Dict[str, str], custom_id: Any) -> None: + """ + Store OpenAI batch custom_id for Vertex batch correlation. + + ``litellm_custom_id`` is GCP-label-safe (may alter casing and characters). + ``litellm_custom_id_raw`` encodes the original string for + round-trip correlation in batch output transforms. + """ + custom_id_str = str(custom_id) + labels["litellm_custom_id"] = _sanitize_gcp_label_value(custom_id_str) + raw_label_chunks = _encode_gcp_label_value_chunks(custom_id_str) + labels["litellm_custom_id_raw"] = raw_label_chunks[0] + for index, raw_label_chunk in enumerate(raw_label_chunks[1:], start=1): + labels[f"litellm_custom_id_raw_{index}"] = raw_label_chunk + + +def _get_litellm_batch_custom_id_from_labels(labels: Dict[str, Any]) -> str: + """Prefer encoded custom_id when present (see _set_litellm_batch_custom_id_labels).""" + raw = labels.get("litellm_custom_id_raw") + if raw: + raw_chunks = [str(raw)] + chunk_prefix = "litellm_custom_id_raw_" + indexed_chunks = [] + for key, value in labels.items(): + if key.startswith(chunk_prefix) and key[len(chunk_prefix) :].isdigit(): + indexed_chunks.append((int(key[len(chunk_prefix) :]), str(value))) + raw_chunks.extend( + raw_label_chunk + for _, raw_label_chunk in sorted(indexed_chunks, key=lambda item: item[0]) + ) + decoded = _decode_gcp_label_value_chunks(raw_chunks) + if decoded is not None: + return decoded + return str(raw) + return str(labels.get("litellm_custom_id", "unknown")) + + +def _openai_batch_jsonl_entries_to_vertex_wrapped_requests( + openai_jsonl_content: List[Dict[str, Any]], + map_openai_to_vertex_params: Callable[[Dict[str, Any]], Dict[str, Any]], +) -> List[Dict[str, Any]]: + """ + Transforms OpenAI JSONL batch entries to Vertex AI JSONL lines. + + jsonl body for vertex is {"request": } + Example Vertex jsonl + {"request":{"contents": [{"role": "user", "parts": [{"text": "What is the relation between the following video and image samples?"}, {"fileData": {"fileUri": "gs://cloud-samples-data/generative-ai/video/animals.mp4", "mimeType": "video/mp4"}}, {"fileData": {"fileUri": "gs://cloud-samples-data/generative-ai/image/cricket.jpeg", "mimeType": "image/jpeg"}}]}]}} + {"request":{"contents": [{"role": "user", "parts": [{"text": "Describe what is happening in this video."}, {"fileData": {"fileUri": "gs://cloud-samples-data/generative-ai/video/another_video.mov", "mimeType": "video/mov"}}]}]}} + """ + + vertex_jsonl_content = [] + for _openai_jsonl_content in openai_jsonl_content: + openai_request_body = _openai_jsonl_content.get("body") or {} + vertex_request_body = _transform_request_body( + messages=openai_request_body.get("messages", []), + model=openai_request_body.get("model", ""), + optional_params=map_openai_to_vertex_params(openai_request_body), + custom_llm_provider="vertex_ai", + litellm_params={}, + cached_content=None, + ) + + # Add custom_id as a label for correlation in batch outputs + custom_id = _openai_jsonl_content.get("custom_id") + if custom_id is not None: + if "labels" not in vertex_request_body: + vertex_request_body["labels"] = {} + _set_litellm_batch_custom_id_labels( + vertex_request_body["labels"], custom_id + ) + + vertex_jsonl_content.append({"request": vertex_request_body}) + return vertex_jsonl_content + class VertexAIFilesConfig(VertexBase, BaseFilesConfig): """ @@ -239,28 +369,10 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig): def _transform_openai_jsonl_content_to_vertex_ai_jsonl_content( self, openai_jsonl_content: List[Dict[str, Any]] ) -> List[Dict[str, Any]]: - """ - Transforms OpenAI JSONL content to VertexAI JSONL content - - jsonl body for vertex is {"request": } - Example Vertex jsonl - {"request":{"contents": [{"role": "user", "parts": [{"text": "What is the relation between the following video and image samples?"}, {"fileData": {"fileUri": "gs://cloud-samples-data/generative-ai/video/animals.mp4", "mimeType": "video/mp4"}}, {"fileData": {"fileUri": "gs://cloud-samples-data/generative-ai/image/cricket.jpeg", "mimeType": "image/jpeg"}}]}]}} - {"request":{"contents": [{"role": "user", "parts": [{"text": "Describe what is happening in this video."}, {"fileData": {"fileUri": "gs://cloud-samples-data/generative-ai/video/another_video.mov", "mimeType": "video/mov"}}]}]}} - """ - - vertex_jsonl_content = [] - for _openai_jsonl_content in openai_jsonl_content: - openai_request_body = _openai_jsonl_content.get("body") or {} - vertex_request_body = _transform_request_body( - messages=openai_request_body.get("messages", []), - model=openai_request_body.get("model", ""), - optional_params=self._map_openai_to_vertex_params(openai_request_body), - custom_llm_provider="vertex_ai", - litellm_params={}, - cached_content=None, - ) - vertex_jsonl_content.append({"request": vertex_request_body}) - return vertex_jsonl_content + return _openai_batch_jsonl_entries_to_vertex_wrapped_requests( + openai_jsonl_content=openai_jsonl_content, + map_openai_to_vertex_params=self._map_openai_to_vertex_params, + ) def transform_create_file_request( self, @@ -461,8 +573,229 @@ class VertexAIFilesConfig(VertexBase, BaseFilesConfig): logging_obj: LiteLLMLoggingObj, litellm_params: dict, ) -> HttpxBinaryResponseContent: + """ + Transform file content response, converting Vertex AI batch output to OpenAI format if applicable. + + This method automatically detects and transforms Vertex AI batch prediction outputs + (predictions.jsonl files) into OpenAI-compatible batch response format. + + If the file is not a batch output or transformation fails, the original content + is returned as-is to maintain backward compatibility. + """ + try: + # Allow users to opt out of automatic Vertex batch output -> OpenAI + # transformation, e.g. if they consume raw `predictions.jsonl` directly. + if getattr(litellm, "disable_vertex_batch_output_transformation", False): + return HttpxBinaryResponseContent(response=raw_response) + + # Try to transform batch output if it's a JSONL file + content = raw_response.content + if content: + transformed_content = self._try_transform_vertex_batch_output_to_openai( + content=content, + logging_obj=logging_obj, + ) + if transformed_content != content: + # Create a new response with transformed content and updated Content-Length + # Update headers with correct Content-Length + new_headers = dict(raw_response.headers) + new_headers["content-length"] = str(len(transformed_content)) + + mock_response = httpx.Response( + status_code=raw_response.status_code, + content=transformed_content, + headers=new_headers, + request=raw_response.request, + ) + return HttpxBinaryResponseContent(response=mock_response) + except Exception: + # If transformation fails, return as-is + pass + return HttpxBinaryResponseContent(response=raw_response) + def _try_transform_vertex_batch_output_to_openai( + self, content: bytes, logging_obj: Optional[LiteLLMLoggingObj] = None + ) -> bytes: + """ + Try to transform Vertex AI batch output to OpenAI format. + If conversion fails at any point, return the original content as-is. + + Vertex AI batch output format (predictions.jsonl): + { + "request": {"contents": [...], "labels": {"litellm_custom_id": "request-1", "litellm_custom_id_raw": "..."}}, + "status": "", + "response": {"candidates": [...], "modelVersion": "gemini-2.5-flash", ...}, + "processed_time": "2026-04-13T10:18:18.102004+00:00" + } + + OpenAI batch output format: + { + "id": "batch_req_...", + "custom_id": "request-1", + "response": { + "status_code": 200, + "request_id": "chatcmpl-...", + "body": {} + }, + "error": null + } + """ + try: + # Decode content + content_str = content.decode("utf-8") + + # Check if it's JSONL (multiple lines) + lines = content_str.strip().split("\n") + if not lines: + return content + + # Try to parse the first line to see if it's Vertex AI batch output + first_line = json.loads(lines[0]) + + # Check if it has Vertex AI batch output structure with discriminating fields + # Must have request, response, and processed_time + # Plus either candidates (success) or status (error) + has_base_structure = ( + "response" in first_line + and "request" in first_line + and "processed_time" in first_line + ) + has_success_or_error = ( + "candidates" in first_line.get("response", {}) + or "promptFeedback" in first_line.get("response", {}) + or bool(first_line.get("status")) + ) + + if not (has_base_structure and has_success_or_error): + # Not a Vertex AI batch output, return as-is + return content + + vertex_gemini_config = VertexGeminiConfig() + # Always use a fresh local Logging object for the per-line transformation + # so we never mutate the caller's logging_obj (which already went through + # pre_call and has its own model/start_time/optional_params set). + batch_transform_logging_obj = Logging( + model="", + messages=[], + stream=False, + call_type="batch_transform", + start_time=time.time(), + litellm_call_id="", + function_id="", + ) + batch_transform_logging_obj.optional_params = {} + mock_httpx_response = httpx.Response( + status_code=200, + headers={"content-type": "application/json"}, + request=httpx.Request(method="POST", url="https://example.com"), + ) + + # Transform all lines + transformed_lines = [] + for line in lines: + if not line.strip(): + continue + + try: + vertex_output = json.loads(line) + openai_output = ( + self._transform_single_vertex_batch_output_to_openai( + vertex_output=vertex_output, + vertex_gemini_config=vertex_gemini_config, + logging_obj=batch_transform_logging_obj, + mock_httpx_response=mock_httpx_response, + ) + ) + transformed_lines.append(json.dumps(openai_output)) + except Exception: + # If any line fails, return original content + return content + + # Return transformed content + return "\n".join(transformed_lines).encode("utf-8") + + except Exception: + # If anything fails, return original content + return content + + def _transform_single_vertex_batch_output_to_openai( + self, + vertex_output: Dict[str, Any], + vertex_gemini_config: VertexGeminiConfig, + logging_obj: Logging, + mock_httpx_response: httpx.Response, + ) -> Dict[str, Any]: + """ + Transform a single Vertex AI batch output line to OpenAI format. + Uses the existing VertexGeminiConfig transformation for the response. + """ + # Extract custom_id from request labels (prefer raw for OpenAI round-trip) + request_data = vertex_output.get("request", {}) + labels = request_data.get("labels", {}) or {} + custom_id = _get_litellm_batch_custom_id_from_labels(labels) + + # Check if there's an error + status = vertex_output.get("status", "") + has_error = bool(status) + + if has_error: + return { + "id": f"batch_req_{uuid.uuid4()}", + "custom_id": custom_id, + "response": None, + "error": { + "code": "vertex_ai_error", + "message": status, + }, + } + + # Transform successful response using existing transformation + vertex_response = vertex_output.get("response", {}) + + # Extract model from response + model = vertex_response.get("modelVersion", "gemini-1.5-flash-001") + if "@" in model: + model = model.split("@")[0] + + try: + # Use existing VertexGeminiConfig transformation + model_response = ModelResponse() + + transformed_response = vertex_gemini_config._transform_google_generate_content_to_openai_model_response( + completion_response=vertex_response, + model_response=model_response, + model=model, + logging_obj=logging_obj, + raw_response=mock_httpx_response, + ) + + # Convert ModelResponse to dict + response_dict = transformed_response.model_dump() + + # Return in OpenAI batch format + return { + "id": f"batch_req_{uuid.uuid4()}", + "custom_id": custom_id, + "response": { + "status_code": 200, + "request_id": response_dict.get("id", ""), + "body": response_dict, + }, + "error": None, + } + + except Exception as e: + return { + "id": f"batch_req_{uuid.uuid4()}", + "custom_id": custom_id, + "response": None, + "error": { + "code": "transformation_error", + "message": f"Failed to transform response: {str(e)}", + }, + } + class VertexAIJsonlFilesTransformation(VertexGeminiConfig): """ @@ -500,29 +833,11 @@ class VertexAIJsonlFilesTransformation(VertexGeminiConfig): def _transform_openai_jsonl_content_to_vertex_ai_jsonl_content( self, openai_jsonl_content: List[Dict[str, Any]] - ): - """ - Transforms OpenAI JSONL content to VertexAI JSONL content - - jsonl body for vertex is {"request": } - Example Vertex jsonl - {"request":{"contents": [{"role": "user", "parts": [{"text": "What is the relation between the following video and image samples?"}, {"fileData": {"fileUri": "gs://cloud-samples-data/generative-ai/video/animals.mp4", "mimeType": "video/mp4"}}, {"fileData": {"fileUri": "gs://cloud-samples-data/generative-ai/image/cricket.jpeg", "mimeType": "image/jpeg"}}]}]}} - {"request":{"contents": [{"role": "user", "parts": [{"text": "Describe what is happening in this video."}, {"fileData": {"fileUri": "gs://cloud-samples-data/generative-ai/video/another_video.mov", "mimeType": "video/mov"}}]}]}} - """ - - vertex_jsonl_content = [] - for _openai_jsonl_content in openai_jsonl_content: - openai_request_body = _openai_jsonl_content.get("body") or {} - vertex_request_body = _transform_request_body( - messages=openai_request_body.get("messages", []), - model=openai_request_body.get("model", ""), - optional_params=self._map_openai_to_vertex_params(openai_request_body), - custom_llm_provider="vertex_ai", - litellm_params={}, - cached_content=None, - ) - vertex_jsonl_content.append({"request": vertex_request_body}) - return vertex_jsonl_content + ) -> List[Dict[str, Any]]: + return _openai_batch_jsonl_entries_to_vertex_wrapped_requests( + openai_jsonl_content=openai_jsonl_content, + map_openai_to_vertex_params=self._map_openai_to_vertex_params, + ) def _get_gcs_object_name( self, diff --git a/litellm/llms/vertex_ai/gemini/transformation.py b/litellm/llms/vertex_ai/gemini/transformation.py index 87bd4843822..9afa5dec465 100644 --- a/litellm/llms/vertex_ai/gemini/transformation.py +++ b/litellm/llms/vertex_ai/gemini/transformation.py @@ -212,6 +212,22 @@ def _process_gemini_media( return _apply_gemini_metadata( part, model, media_resolution_enum, video_metadata ) + elif image_url.startswith( + "https://generativelanguage.googleapis.com/v1beta/files/" + ): + # Gemini Files API URIs — the file is already uploaded to Google's + # servers; pass the URI through as file_data without fetching it. + # These URLs return 403 when accessed directly, so we must not try + # to resolve their MIME type via HTTP. + if format: + file_data = FileDataType(mime_type=format, file_uri=image_url) + else: + # Gemini Files API references can be passed through as URI-only. + file_data = cast(FileDataType, {"file_uri": image_url}) + part = {"file_data": file_data} + return _apply_gemini_metadata( + part, model, media_resolution_enum, video_metadata + ) elif ( "https://" in image_url and (image_type := format or _get_image_mime_type_from_url(image_url)) @@ -743,16 +759,22 @@ def _transform_request_body( # noqa: PLR0915 ] data = RequestBody(contents=content) - if system_instructions is not None: - data["system_instruction"] = system_instructions - if tools is not None: - data["tools"] = tools - if tool_choice is not None: - data["toolConfig"] = tool_choice - if include_server_side_tool_invocations: - if "toolConfig" not in data: - data["toolConfig"] = {} - data["toolConfig"]["includeServerSideToolInvocations"] = True + # Vertex rejects system_instruction/tools/toolConfig alongside cachedContent. + # Treat dropping these fields as a request mutation guarded by modify_params. + can_send_cache_incompatible_fields = ( + cached_content is None or litellm.modify_params is False + ) + if can_send_cache_incompatible_fields: + if system_instructions is not None: + data["system_instruction"] = system_instructions + if tools is not None: + data["tools"] = tools + if tool_choice is not None: + data["toolConfig"] = tool_choice + if include_server_side_tool_invocations: + if "toolConfig" not in data: + data["toolConfig"] = {} + data["toolConfig"]["includeServerSideToolInvocations"] = True if safety_settings is not None: data["safetySettings"] = safety_settings if generation_config is not None and len(generation_config) > 0: diff --git a/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py b/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py index 474ddb402a1..6278de662f8 100644 --- a/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py +++ b/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py @@ -979,15 +979,8 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): params["includeThoughts"] = False else: params["includeThoughts"] = True - if thinking_budget >= 10000: - is_gemini3flash = ( - "gemini-3-flash-preview" in model.lower() - or "gemini-3-flash" in model.lower() - ) - params["thinkingLevel"] = ( - "minimal" if is_gemini3flash else "low" - ) - else: + # Follow provider defaults unless explicitly opted into legacy behavior. + if litellm.enable_gemini_default_thinking_level_low is True: is_gemini3flash = ( "gemini-3-flash-preview" in model.lower() or "gemini-3-flash" in model.lower() diff --git a/litellm/llms/vertex_ai/vertex_ai_partner_models/anthropic/experimental_pass_through/transformation.py b/litellm/llms/vertex_ai/vertex_ai_partner_models/anthropic/experimental_pass_through/transformation.py index 5c3bbf61ee2..d450f7a4635 100644 --- a/litellm/llms/vertex_ai/vertex_ai_partner_models/anthropic/experimental_pass_through/transformation.py +++ b/litellm/llms/vertex_ai/vertex_ai_partner_models/anthropic/experimental_pass_through/transformation.py @@ -13,6 +13,7 @@ from litellm.types.llms.vertex_ai import VertexPartnerProvider from litellm.types.router import GenericLiteLLMParams from ....vertex_llm_base import VertexBase +from ..output_params_utils import sanitize_vertex_anthropic_output_params class VertexAIPartnerModelsAnthropicMessagesConfig(AnthropicMessagesConfig, VertexBase): @@ -158,12 +159,10 @@ class VertexAIPartnerModelsAnthropicMessagesConfig(AnthropicMessagesConfig, Vert "model", None ) # do not pass model in request body to vertex ai - anthropic_messages_request.pop( - "output_format", None - ) # do not pass output_format in request body to vertex ai - vertex ai does not support output_format as yet - - anthropic_messages_request.pop( - "output_config", None - ) # do not pass output_config in request body to vertex ai - vertex ai does not support output_config + # Vertex AI Claude accepts ``output_config.format`` (structured outputs) + # and ``output_format``, but rejects ``output_config.effort`` with 400 + # "Extra inputs are not permitted". Sanitize in place so the supported + # bits flow through. + sanitize_vertex_anthropic_output_params(anthropic_messages_request) return anthropic_messages_request diff --git a/litellm/llms/vertex_ai/vertex_ai_partner_models/anthropic/output_params_utils.py b/litellm/llms/vertex_ai/vertex_ai_partner_models/anthropic/output_params_utils.py new file mode 100644 index 00000000000..982d8edbf20 --- /dev/null +++ b/litellm/llms/vertex_ai/vertex_ai_partner_models/anthropic/output_params_utils.py @@ -0,0 +1,50 @@ +""" +Shared sanitization for ``output_config`` / ``output_format`` on Vertex AI +Claude. Lives in its own module so both the chat-completion transformation +(``transformation.py``) and the Messages pass-through transformation +(``experimental_pass_through/transformation.py``) can import it without +forming a cycle through the parent module's heavier imports. + +CodeQL flagged the ``..transformation`` import path as a potential cyclic +import; extracting the helper into a leaf module resolves the warning and +keeps the parent module's import surface narrow. +""" + +# Keys inside ``output_config`` that Vertex AI Claude does not accept. +# Today only ``effort`` triggers "Extra inputs are not permitted"; add new +# entries here as Vertex parity drifts. Keep this list narrow — anything +# Vertex DOES accept (e.g. ``format`` for structured outputs) must be +# preserved so callers can rely on Anthropic-native features. +VERTEX_UNSUPPORTED_OUTPUT_CONFIG_KEYS: frozenset = frozenset({"effort"}) + + +def sanitize_vertex_anthropic_output_params(data: dict) -> None: + """ + Strip Vertex-unsupported keys from ``output_config`` / + ``output_format`` in-place; forward whatever remains. + + Behavior: + * ``output_config`` containing only unsupported keys (e.g. ``effort`` + alone) is removed entirely so the request body has no empty dict. + * ``output_config`` containing a mix of supported + unsupported keys + has the unsupported subset filtered out and the rest forwarded. + * ``output_config`` that is supported in full passes through unchanged. + * ``output_format`` is forwarded as-is (Vertex AI Claude accepts it). + * Non-dict values for ``output_config`` are dropped to avoid sending + malformed payloads downstream. + """ + output_config = data.get("output_config") + if output_config is None: + return + if not isinstance(output_config, dict): + data.pop("output_config", None) + return + sanitized = { + k: v + for k, v in output_config.items() + if k not in VERTEX_UNSUPPORTED_OUTPUT_CONFIG_KEYS + } + if sanitized: + data["output_config"] = sanitized + else: + data.pop("output_config", None) diff --git a/litellm/llms/vertex_ai/vertex_ai_partner_models/anthropic/transformation.py b/litellm/llms/vertex_ai/vertex_ai_partner_models/anthropic/transformation.py index 504914c4796..914c7e92e5e 100644 --- a/litellm/llms/vertex_ai/vertex_ai_partner_models/anthropic/transformation.py +++ b/litellm/llms/vertex_ai/vertex_ai_partner_models/anthropic/transformation.py @@ -10,6 +10,7 @@ from litellm.types.llms.openai import AllMessageValues from litellm.types.utils import ModelResponse from ....anthropic.chat.transformation import AnthropicConfig +from .output_params_utils import sanitize_vertex_anthropic_output_params class VertexAIError(Exception): @@ -105,11 +106,12 @@ class VertexAIAnthropicConfig(AnthropicConfig): data.pop("model", None) # vertex anthropic doesn't accept 'model' parameter - # VertexAI doesn't support output_format parameter, remove it if present - data.pop("output_format", None) - - # VertexAI doesn't support output_config parameter, remove it if present - data.pop("output_config", None) + # Vertex AI Claude accepts ``output_config.format`` (structured outputs / + # JSON Schema) but NOT ``output_config.effort`` — sending ``effort`` to + # Vertex returns 400 "Extra inputs are not permitted". Sanitize in place: + # forward the structured-output bits, drop the unsupported keys. + # Same treatment for the legacy top-level ``output_format`` field. + sanitize_vertex_anthropic_output_params(data) tools = optional_params.get("tools") tool_search_used = self.is_tool_search_used(tools) diff --git a/litellm/main.py b/litellm/main.py index 0079bd750cf..0553cf9d422 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -4923,8 +4923,17 @@ def embedding( # noqa: PLR0915 if encoding_format is not None: optional_params["encoding_format"] = encoding_format else: - # Omiting causes openai sdk to add default value of "float" - optional_params["encoding_format"] = None + env_fmt = get_secret_str("LITELLM_DEFAULT_EMBEDDING_ENCODING_FORMAT") + if env_fmt is not None and env_fmt.strip().lower() == "none": + optional_params.pop("encoding_format", None) + else: + _default_fmt = ( + optional_params.get("encoding_format") or env_fmt or "float" + ) + if _default_fmt.strip().lower() == "none": + optional_params.pop("encoding_format", None) + else: + optional_params["encoding_format"] = _default_fmt api_version = None diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index a1e3e42a9c5..b49d97dc4e3 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -19928,7 +19928,7 @@ "supports_web_search": true, "supports_none_reasoning_effort": true, "supports_xhigh_reasoning_effort": true, - "supports_minimal_reasoning_effort": true + "supports_minimal_reasoning_effort": false }, "gpt-5.5-2026-04-23": { "cache_read_input_token_cost": 5e-07, @@ -19976,7 +19976,7 @@ "supports_web_search": true, "supports_none_reasoning_effort": true, "supports_xhigh_reasoning_effort": true, - "supports_minimal_reasoning_effort": true + "supports_minimal_reasoning_effort": false }, "gpt-5.5-pro": { "cache_read_input_token_cost": 3e-06, @@ -20019,7 +20019,8 @@ "supports_web_search": true, "supports_none_reasoning_effort": false, "supports_xhigh_reasoning_effort": true, - "supports_minimal_reasoning_effort": true + "supports_minimal_reasoning_effort": false, + "supports_low_reasoning_effort": false }, "gpt-5.5-pro-2026-04-23": { "cache_read_input_token_cost": 3e-06, @@ -20062,7 +20063,8 @@ "supports_web_search": true, "supports_none_reasoning_effort": false, "supports_xhigh_reasoning_effort": true, - "supports_minimal_reasoning_effort": true + "supports_minimal_reasoning_effort": false, + "supports_low_reasoning_effort": false }, "gpt-5.4": { "cache_read_input_token_cost": 2.5e-07, @@ -22061,6 +22063,98 @@ "tool_use_system_prompt_tokens": 346, "supports_native_structured_output": true }, + "crusoe/deepseek-ai/DeepSeek-R1-0528": { + "input_cost_per_token": 3e-06, + "litellm_provider": "crusoe", + "max_input_tokens": 163840, + "max_output_tokens": 163840, + "max_tokens": 163840, + "mode": "chat", + "output_cost_per_token": 7e-06, + "supports_function_calling": false, + "supports_reasoning": true, + "supports_system_messages": true, + "supports_tool_choice": false + }, + "crusoe/deepseek-ai/DeepSeek-V3-0324": { + "input_cost_per_token": 1.5e-06, + "litellm_provider": "crusoe", + "max_input_tokens": 163840, + "max_output_tokens": 163840, + "max_tokens": 163840, + "mode": "chat", + "output_cost_per_token": 1.5e-06, + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_system_messages": true, + "supports_tool_choice": true + }, + "crusoe/google/gemma-3-12b-it": { + "input_cost_per_token": 1e-07, + "litellm_provider": "crusoe", + "max_input_tokens": 131072, + "max_output_tokens": 131072, + "max_tokens": 131072, + "mode": "chat", + "output_cost_per_token": 1e-07, + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_vision": true + }, + "crusoe/meta-llama/Llama-3.3-70B-Instruct": { + "input_cost_per_token": 2e-07, + "litellm_provider": "crusoe", + "max_input_tokens": 131072, + "max_output_tokens": 131072, + "max_tokens": 131072, + "mode": "chat", + "output_cost_per_token": 2e-07, + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_system_messages": true, + "supports_tool_choice": true + }, + "crusoe/moonshotai/Kimi-K2-Thinking": { + "input_cost_per_token": 2.5e-06, + "litellm_provider": "crusoe", + "max_input_tokens": 262144, + "max_output_tokens": 262144, + "max_tokens": 262144, + "mode": "chat", + "output_cost_per_token": 2.5e-06, + "supports_function_calling": false, + "supports_reasoning": true, + "supports_system_messages": true, + "supports_tool_choice": false + }, + "crusoe/openai/gpt-oss-120b": { + "input_cost_per_token": 8e-07, + "litellm_provider": "crusoe", + "max_input_tokens": 131072, + "max_output_tokens": 131072, + "max_tokens": 131072, + "mode": "chat", + "output_cost_per_token": 8e-07, + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_system_messages": true, + "supports_tool_choice": true + }, + "crusoe/Qwen/Qwen3-235B-A22B-Instruct-2507": { + "input_cost_per_token": 3e-06, + "litellm_provider": "crusoe", + "max_input_tokens": 262144, + "max_output_tokens": 262144, + "max_tokens": 262144, + "mode": "chat", + "output_cost_per_token": 3e-06, + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_system_messages": true, + "supports_tool_choice": true + }, "lambda_ai/deepseek-llama3.3-70b": { "input_cost_per_token": 2e-07, "litellm_provider": "lambda_ai", diff --git a/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py b/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py index 756b2ed91d7..a05af66118c 100644 --- a/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py +++ b/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py @@ -409,9 +409,12 @@ class MCPRequestHandler: Permission hierarchy (all rules are intersections): 1. Get allowed servers from key permissions - 2. Get allowed servers from team permissions - 3. Get allowed servers from end_user permissions - 4. Final result = intersection of key/team AND end_user (if end_user has permissions set) + 2. Get allowed servers from team permissions (key inherits from team, or intersection) + 3. Get allowed servers from end_user permissions (intersected if set) + 4. Get allowed servers from agent permissions (intersected if set) + 5. Get allowed servers from org permissions — org acts as a ceiling: if the org + has an explicit MCP server list, the combined key/team/end_user/agent result is + capped to that list. If the org has no list, no extra restriction is applied. Returns: List[str]: List of allowed MCP servers by server id @@ -435,6 +438,10 @@ class MCPRequestHandler: # Calculate key/team allowed servers using inheritance and intersection logic ######################################################### allowed_mcp_servers: List[str] = [] + has_lower_level_mcp_restrictions = ( + len(allowed_mcp_servers_for_key) > 0 + or len(allowed_mcp_servers_for_team) > 0 + ) if len(allowed_mcp_servers_for_team) > 0: if len(allowed_mcp_servers_for_key) > 0: # Key has its own MCP permissions - use intersection with team permissions @@ -459,6 +466,7 @@ class MCPRequestHandler: # If end_user has explicit MCP server permissions, apply intersection if len(allowed_mcp_servers_for_end_user) > 0: + has_lower_level_mcp_restrictions = True verbose_logger.debug( f"End user {user_api_key_auth.end_user_id} has explicit MCP permissions: {allowed_mcp_servers_for_end_user}" ) @@ -490,6 +498,7 @@ class MCPRequestHandler: ) ) if len(allowed_mcp_servers_for_agent) > 0: + has_lower_level_mcp_restrictions = True # Intersect: agent can only use servers allowed by BOTH key/team AND agent config allowed_mcp_servers = [ s @@ -500,6 +509,30 @@ class MCPRequestHandler: f"Applied agent intersection filter. Final allowed servers: {allowed_mcp_servers}" ) + ######################################################### + # Apply org-level ceiling if org_id is set + ######################################################### + if user_api_key_auth and user_api_key_auth.org_id: + allowed_mcp_servers_for_org = ( + await MCPRequestHandler._get_allowed_mcp_servers_for_org( + user_api_key_auth + ) + ) + if len(allowed_mcp_servers_for_org) > 0: + if has_lower_level_mcp_restrictions: + # Lower-level restrictions exist, so org can only cap them. + allowed_mcp_servers = [ + s + for s in allowed_mcp_servers + if s in allowed_mcp_servers_for_org + ] + else: + # No lower-level restrictions → org list becomes the ceiling + allowed_mcp_servers = allowed_mcp_servers_for_org + verbose_logger.debug( + f"Applied org ceiling filter. Final allowed servers: {allowed_mcp_servers}" + ) + return list(set(allowed_mcp_servers)) except Exception as e: verbose_logger.warning(f"Failed to get allowed MCP servers: {str(e)}") @@ -638,6 +671,27 @@ class MCPRequestHandler: allowed_tools = list(set(allowed_tools) & set(agent_tools)) else: allowed_tools = agent_tools + + # Apply org-level tool ceiling if org_id is set + if user_api_key_auth.org_id: + # _get_org_object_permission uses user_api_key_cache, so this is not a + # fresh DB round-trip when get_allowed_mcp_servers was already called. + org_obj_perm = await MCPRequestHandler._get_org_object_permission( + user_api_key_auth + ) + org_tools = ( + global_mcp_server_manager.expand_tool_permissions( + org_obj_perm.mcp_tool_permissions + ).get(server_id) + if org_obj_perm and org_obj_perm.mcp_tool_permissions + else None + ) + if org_tools is not None: + if allowed_tools is not None: + allowed_tools = list(set(allowed_tools) & set(org_tools)) + else: + allowed_tools = list(org_tools) + return allowed_tools except Exception as e: @@ -805,6 +859,120 @@ class MCPRequestHandler: ) return [] + # Sentinel stored in cache when an org has no object_permission, so we + # don't re-query the DB on every MCP request for that org. + _ORG_NO_PERMISSION_SENTINEL = "__org_no_mcp_permission__" + + @staticmethod + async def _get_org_object_permission( + user_api_key_auth: Optional[UserAPIKeyAuth] = None, + ): + """ + Get org object_permission, using user_api_key_cache to avoid DB hits on every request. + + Caches both positive results and the absence of an object_permission so that orgs + with no MCP permissions configured (the common default) do not trigger a DB query + on every request. + """ + from litellm.proxy.proxy_server import prisma_client, user_api_key_cache + + if not user_api_key_auth or not user_api_key_auth.org_id: + return None + + if prisma_client is None: + verbose_logger.debug("prisma_client is None") + return None + + org_id = user_api_key_auth.org_id + cache_key = f"org_object_permission:{org_id}" + + from litellm.proxy._types import LiteLLM_ObjectPermissionTable + + try: + cached = await user_api_key_cache.async_get_cache(key=cache_key) + if cached is not None: + # Sentinel means the DB confirmed no object_permission for this org + if cached == MCPRequestHandler._ORG_NO_PERMISSION_SENTINEL: + return None + # Redis deserialises to a plain dict; reconstruct the Pydantic model + # so callers can access .mcp_servers / .mcp_tool_permissions as attrs. + if isinstance(cached, dict): + return LiteLLM_ObjectPermissionTable(**cached) + return cached + + org_row = await prisma_client.db.litellm_organizationtable.find_unique( + where={"organization_id": org_id}, + include={"object_permission": True}, + ) + + if org_row is None or org_row.object_permission is None: + # Cache the negative result so subsequent calls skip the DB + await user_api_key_cache.async_set_cache( + key=cache_key, + value=MCPRequestHandler._ORG_NO_PERMISSION_SENTINEL, + ) + return None + + # Convert raw Prisma model → Pydantic before caching. Caching the + # Pydantic .dict() ensures the value survives a Redis JSON round-trip + # as a plain dict that we can reconstruct above (same pattern used by + # get_end_user_object / get_team_object in auth_checks.py). + obj_perm = LiteLLM_ObjectPermissionTable(**org_row.object_permission.dict()) + await user_api_key_cache.async_set_cache( + key=cache_key, value=obj_perm.dict() + ) + return obj_perm + except Exception as e: + verbose_logger.warning(f"Failed to get org object permission: {str(e)}") + return None + + @staticmethod + async def _get_allowed_mcp_servers_for_org( + user_api_key_auth: Optional[UserAPIKeyAuth] = None, + ) -> List[str]: + """ + Get allowed MCP servers for an organization. + + Returns the MCP servers from the org's object_permission. + An empty result means the org places no restriction (allow-all from this level). + """ + try: + object_permissions = await MCPRequestHandler._get_org_object_permission( + user_api_key_auth + ) + + if object_permissions is None: + return [] + + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + global_mcp_server_manager, + ) + + # Expand names/aliases to canonical server IDs (consistent with key/team/end-user path) + direct_mcp_servers = global_mcp_server_manager.expand_permission_list( + object_permissions.mcp_servers or [] + ) + + access_group_servers = ( + await MCPRequestHandler._get_mcp_servers_from_access_groups( + object_permissions.mcp_access_groups or [] + ) + ) + + tool_perm_servers = list( + global_mcp_server_manager.expand_tool_permissions( + object_permissions.mcp_tool_permissions + ).keys() + ) + + all_servers = direct_mcp_servers + access_group_servers + tool_perm_servers + return list(set(all_servers)) + except Exception as e: + verbose_logger.warning( + f"Failed to get allowed MCP servers for org: {str(e)}" + ) + return [] + @staticmethod async def _get_allowed_mcp_servers_for_end_user( user_api_key_auth: Optional[UserAPIKeyAuth] = None, diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index abb4b5cfa6f..54d9bbe6e28 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -2138,6 +2138,47 @@ if MCP_AVAILABLE: ######################################################### local_tool = global_mcp_tool_registry.get_tool(name) if local_tool: + # OpenAPI-backed tools used to bypass `pre_call_tool_check` — + # only the managed path ran allowed/banned-tool checks, key/team + # tool permissions, and parameter validation. Run the same checks + # before dispatching to the local registry. Refuse the call if + # we cannot resolve a server: tools registered via + # openapi_to_mcp_generator are always tied to a server, so a + # missing mcp_server here means the tool->server mapping has + # not finished initializing or the registry entry is orphaned. + # Skipping the check would re-open the same authorization gap. + if mcp_server is None: + raise HTTPException( + status_code=503, + detail=( + f"MCP server for tool '{name}' is not available; " + "refusing to dispatch without authorization checks. " + "Retry once the server is registered." + ), + ) + + # `pre_call_tool_check` calls into `proxy_logging_obj` for the + # pre-call guardrail hooks, so source it from the canonical + # `proxy_server` module the same way `_handle_managed_mcp_tool` + # does. `kwargs.get("proxy_logging_obj")` is None on the MCP + # entry path and would crash with AttributeError after the + # security checks pass. + from litellm.proxy.proxy_server import proxy_logging_obj + + hook_result = await global_mcp_server_manager.pre_call_tool_check( + name=original_tool_name, + arguments=arguments or {}, + server_name=server_name or mcp_server.name, + user_api_key_auth=user_api_key_auth, + proxy_logging_obj=proxy_logging_obj, + server=mcp_server, + raw_headers=raw_headers, + ) + # `pre_call_tool_check` may return guardrail-modified + # arguments; honor them on the local path too. + if isinstance(hook_result, dict) and "arguments" in hook_result: + arguments = hook_result["arguments"] + verbose_logger.debug(f"Executing local registry tool: {name}") # For BYOK servers the credential must be injected via a ContextVar # because the tool function has headers baked into its closure. diff --git a/litellm/proxy/_lazy_openapi_snapshot.json b/litellm/proxy/_lazy_openapi_snapshot.json index 46a514c0870..eb35dd6cb3e 100644 --- a/litellm/proxy/_lazy_openapi_snapshot.json +++ b/litellm/proxy/_lazy_openapi_snapshot.json @@ -3616,7 +3616,7 @@ }, "get": { "description": "[Docs](https://docs.litellm.ai/docs/pass_through/anthropic_completion)", - "operationId": "anthropic_proxy_route_anthropic__endpoint__get", + "operationId": "anthropic_proxy_route_anthropic__endpoint__delete", "parameters": [ { "in": "path", @@ -3660,7 +3660,7 @@ }, "patch": { "description": "[Docs](https://docs.litellm.ai/docs/pass_through/anthropic_completion)", - "operationId": "anthropic_proxy_route_anthropic__endpoint__patch", + "operationId": "anthropic_proxy_route_anthropic__endpoint__delete", "parameters": [ { "in": "path", @@ -3704,7 +3704,7 @@ }, "post": { "description": "[Docs](https://docs.litellm.ai/docs/pass_through/anthropic_completion)", - "operationId": "anthropic_proxy_route_anthropic__endpoint__post", + "operationId": "anthropic_proxy_route_anthropic__endpoint__delete", "parameters": [ { "in": "path", @@ -3748,7 +3748,7 @@ }, "put": { "description": "[Docs](https://docs.litellm.ai/docs/pass_through/anthropic_completion)", - "operationId": "anthropic_proxy_route_anthropic__endpoint__put", + "operationId": "anthropic_proxy_route_anthropic__endpoint__delete", "parameters": [ { "in": "path", @@ -13299,7 +13299,7 @@ }, "get": { "description": "Call Langfuse via LiteLLM proxy. Works with Langfuse SDK.\n\n[Docs](https://docs.litellm.ai/docs/pass_through/langfuse)", - "operationId": "langfuse_proxy_route_langfuse__endpoint__get", + "operationId": "langfuse_proxy_route_langfuse__endpoint__delete", "parameters": [ { "in": "path", @@ -13338,7 +13338,7 @@ }, "patch": { "description": "Call Langfuse via LiteLLM proxy. Works with Langfuse SDK.\n\n[Docs](https://docs.litellm.ai/docs/pass_through/langfuse)", - "operationId": "langfuse_proxy_route_langfuse__endpoint__patch", + "operationId": "langfuse_proxy_route_langfuse__endpoint__delete", "parameters": [ { "in": "path", @@ -13377,7 +13377,7 @@ }, "post": { "description": "Call Langfuse via LiteLLM proxy. Works with Langfuse SDK.\n\n[Docs](https://docs.litellm.ai/docs/pass_through/langfuse)", - "operationId": "langfuse_proxy_route_langfuse__endpoint__post", + "operationId": "langfuse_proxy_route_langfuse__endpoint__delete", "parameters": [ { "in": "path", @@ -13416,7 +13416,7 @@ }, "put": { "description": "Call Langfuse via LiteLLM proxy. Works with Langfuse SDK.\n\n[Docs](https://docs.litellm.ai/docs/pass_through/langfuse)", - "operationId": "langfuse_proxy_route_langfuse__endpoint__put", + "operationId": "langfuse_proxy_route_langfuse__endpoint__delete", "parameters": [ { "in": "path", @@ -14008,7 +14008,7 @@ "/mcp-rest/test/connection": { "post": { "description": "Test if we can connect to the provided MCP server before adding it", - "operationId": "test_connection_mcp_rest_test_connection_post", + "operationId": "test_connection_mcp_rest_test_connection_post_2", "requestBody": { "content": { "application/json": { @@ -14053,7 +14053,7 @@ "/mcp-rest/test/tools/list": { "post": { "description": "Preview tools available from MCP server before adding it", - "operationId": "test_tools_list_mcp_rest_test_tools_list_post", + "operationId": "test_tools_list_mcp_rest_test_tools_list_post_2", "requestBody": { "content": { "application/json": { @@ -14098,7 +14098,7 @@ "/mcp-rest/tools/call": { "post": { "description": "REST API to call a specific MCP tool with the provided arguments", - "operationId": "call_tool_rest_api_mcp_rest_tools_call_post", + "operationId": "call_tool_rest_api_mcp_rest_tools_call_post_2", "responses": { "200": { "content": { @@ -14123,7 +14123,7 @@ "/mcp-rest/tools/list": { "get": { "description": "List all available tools with information about the server they belong to.\n\nExample response:\n{\n \"tools\": [\n {\n \"name\": \"create_zap\",\n \"description\": \"Create a new zap\",\n \"inputSchema\": \"tool_input_schema\",\n \"mcp_info\": {\n \"server_name\": \"zapier\",\n \"logo_url\": \"https://www.zapier.com/logo.png\",\n }\n }\n ],\n \"error\": null,\n \"message\": \"Successfully retrieved tools\"\n}", - "operationId": "list_tool_rest_api_mcp_rest_tools_list_get", + "operationId": "list_tool_rest_api_mcp_rest_tools_list_get_2", "parameters": [ { "description": "The server id to list tools for", @@ -21896,7 +21896,7 @@ "/policies/usage/overview": { "get": { "description": "Return policy performance overview for the dashboard.", - "operationId": "policies_usage_overview_policies_usage_overview_get", + "operationId": "policies_usage_overview_policies_usage_overview_get_2", "parameters": [ { "description": "YYYY-MM-DD", @@ -22521,7 +22521,7 @@ "/policies/attachments/estimate-impact": { "post": { "description": "Estimate how many keys and teams would be affected by a policy attachment.\n\nUse this before creating an attachment to preview the blast radius.\n\nExample Request:\n```bash\ncurl -X POST \"http://localhost:4000/policies/attachments/estimate-impact\" \\\n -H \"Authorization: Bearer \" \\\n -H \"Content-Type: application/json\" \\\n -d '{\n \"policy_name\": \"hipaa-compliance\",\n \"tags\": [\"healthcare\", \"health-*\"]\n }'\n```", - "operationId": "estimate_attachment_impact_policies_attachments_estimate_impact_post", + "operationId": "estimate_attachment_impact_policies_attachments_estimate_impact_post_2", "requestBody": { "content": { "application/json": { @@ -22568,7 +22568,7 @@ "/policies/resolve": { "post": { "description": "Resolve which policies and guardrails apply for a given context.\n\nUse this endpoint to debug \"what guardrails would apply to a request\nwith this team/key/model/tags combination?\"\n\nExample Request:\n```bash\ncurl -X POST \"http://localhost:4000/policies/resolve\" \\\n -H \"Authorization: Bearer \" \\\n -H \"Content-Type: application/json\" \\\n -d '{\n \"tags\": [\"healthcare\"],\n \"model\": \"gpt-4\"\n }'\n```", - "operationId": "resolve_policies_for_context_policies_resolve_post", + "operationId": "resolve_policies_for_context_policies_resolve_post_2", "parameters": [ { "description": "Force a DB sync before resolving. Default uses in-memory cache.", @@ -26922,7 +26922,7 @@ }, "get": { "description": "Namespace a toolset as its own MCP endpoint.\n\nConnecting to /toolset//mcp exposes exactly the tools defined in\nthe toolset. Access is enforced: non-admin API keys must have the toolset\nlisted in their object_permission.mcp_toolsets grant list, or the request\nwill be rejected with a 403.", - "operationId": "toolset_mcp_route_toolset__toolset_name__mcp_get", + "operationId": "toolset_mcp_route_toolset__toolset_name__mcp_delete", "parameters": [ { "in": "path", @@ -26961,7 +26961,7 @@ }, "head": { "description": "Namespace a toolset as its own MCP endpoint.\n\nConnecting to /toolset//mcp exposes exactly the tools defined in\nthe toolset. Access is enforced: non-admin API keys must have the toolset\nlisted in their object_permission.mcp_toolsets grant list, or the request\nwill be rejected with a 403.", - "operationId": "toolset_mcp_route_toolset__toolset_name__mcp_head", + "operationId": "toolset_mcp_route_toolset__toolset_name__mcp_delete", "parameters": [ { "in": "path", @@ -27000,7 +27000,7 @@ }, "options": { "description": "Namespace a toolset as its own MCP endpoint.\n\nConnecting to /toolset//mcp exposes exactly the tools defined in\nthe toolset. Access is enforced: non-admin API keys must have the toolset\nlisted in their object_permission.mcp_toolsets grant list, or the request\nwill be rejected with a 403.", - "operationId": "toolset_mcp_route_toolset__toolset_name__mcp_options", + "operationId": "toolset_mcp_route_toolset__toolset_name__mcp_delete", "parameters": [ { "in": "path", @@ -27039,7 +27039,7 @@ }, "patch": { "description": "Namespace a toolset as its own MCP endpoint.\n\nConnecting to /toolset//mcp exposes exactly the tools defined in\nthe toolset. Access is enforced: non-admin API keys must have the toolset\nlisted in their object_permission.mcp_toolsets grant list, or the request\nwill be rejected with a 403.", - "operationId": "toolset_mcp_route_toolset__toolset_name__mcp_patch", + "operationId": "toolset_mcp_route_toolset__toolset_name__mcp_delete", "parameters": [ { "in": "path", @@ -27078,7 +27078,7 @@ }, "post": { "description": "Namespace a toolset as its own MCP endpoint.\n\nConnecting to /toolset//mcp exposes exactly the tools defined in\nthe toolset. Access is enforced: non-admin API keys must have the toolset\nlisted in their object_permission.mcp_toolsets grant list, or the request\nwill be rejected with a 403.", - "operationId": "toolset_mcp_route_toolset__toolset_name__mcp_post", + "operationId": "toolset_mcp_route_toolset__toolset_name__mcp_delete", "parameters": [ { "in": "path", @@ -27117,7 +27117,7 @@ }, "put": { "description": "Namespace a toolset as its own MCP endpoint.\n\nConnecting to /toolset//mcp exposes exactly the tools defined in\nthe toolset. Access is enforced: non-admin API keys must have the toolset\nlisted in their object_permission.mcp_toolsets grant list, or the request\nwill be rejected with a 403.", - "operationId": "toolset_mcp_route_toolset__toolset_name__mcp_put", + "operationId": "toolset_mcp_route_toolset__toolset_name__mcp_delete", "parameters": [ { "in": "path", @@ -28329,7 +28329,7 @@ "/v1/vector_stores": { "get": { "description": "List vector stores.\n\nAPI Reference:\nhttps://platform.openai.com/docs/api-reference/vector-stores/list", - "operationId": "vector_store_list_v1_vector_stores_get", + "operationId": "vector_store_list_v1_vector_stores_get_2", "parameters": [ { "in": "query", @@ -28430,7 +28430,7 @@ }, "post": { "description": "Create a vector store.\n\nAPI Reference:\nhttps://platform.openai.com/docs/api-reference/vector-stores/create\n\nSupports target_model_names parameter for creating vector stores across multiple models:\n```json\n{\n \"name\": \"my-vector-store\",\n \"target_model_names\": \"gpt-4,gemini-2.0\"\n}\n```", - "operationId": "vector_store_create_v1_vector_stores_post", + "operationId": "vector_store_create_v1_vector_stores_post_2", "responses": { "200": { "content": { @@ -28455,7 +28455,7 @@ "/v1/vector_stores/{vector_store_id}": { "delete": { "description": "Delete a vector store.\n\nAPI Reference:\nhttps://platform.openai.com/docs/api-reference/vector-stores/delete", - "operationId": "vector_store_delete_v1_vector_stores__vector_store_id__delete", + "operationId": "vector_store_delete_v1_vector_stores__vector_store_id__delete_2", "parameters": [ { "in": "path", @@ -28499,7 +28499,7 @@ }, "get": { "description": "Retrieve a vector store.\n\nAPI Reference:\nhttps://platform.openai.com/docs/api-reference/vector-stores/retrieve", - "operationId": "vector_store_retrieve_v1_vector_stores__vector_store_id__get", + "operationId": "vector_store_retrieve_v1_vector_stores__vector_store_id__get_2", "parameters": [ { "in": "path", @@ -28543,7 +28543,7 @@ }, "post": { "description": "Update a vector store.\n\nAPI Reference:\nhttps://platform.openai.com/docs/api-reference/vector-stores/modify", - "operationId": "vector_store_update_v1_vector_stores__vector_store_id__post", + "operationId": "vector_store_update_v1_vector_stores__vector_store_id__post_2", "parameters": [ { "in": "path", @@ -28588,7 +28588,7 @@ }, "/v1/vector_stores/{vector_store_id}/files": { "get": { - "operationId": "vector_store_file_list_v1_vector_stores__vector_store_id__files_get", + "operationId": "vector_store_file_list_v1_vector_stores__vector_store_id__files_get_2", "parameters": [ { "in": "path", @@ -28631,7 +28631,7 @@ ] }, "post": { - "operationId": "vector_store_file_create_v1_vector_stores__vector_store_id__files_post", + "operationId": "vector_store_file_create_v1_vector_stores__vector_store_id__files_post_2", "parameters": [ { "in": "path", @@ -28676,7 +28676,7 @@ }, "/v1/vector_stores/{vector_store_id}/files/{file_id}": { "delete": { - "operationId": "vector_store_file_delete_v1_vector_stores__vector_store_id__files__file_id__delete", + "operationId": "vector_store_file_delete_v1_vector_stores__vector_store_id__files__file_id__delete_2", "parameters": [ { "in": "path", @@ -28728,7 +28728,7 @@ ] }, "get": { - "operationId": "vector_store_file_retrieve_v1_vector_stores__vector_store_id__files__file_id__get", + "operationId": "vector_store_file_retrieve_v1_vector_stores__vector_store_id__files__file_id__get_2", "parameters": [ { "in": "path", @@ -28780,7 +28780,7 @@ ] }, "post": { - "operationId": "vector_store_file_update_v1_vector_stores__vector_store_id__files__file_id__post", + "operationId": "vector_store_file_update_v1_vector_stores__vector_store_id__files__file_id__post_2", "parameters": [ { "in": "path", @@ -28834,7 +28834,7 @@ }, "/v1/vector_stores/{vector_store_id}/files/{file_id}/content": { "get": { - "operationId": "vector_store_file_content_v1_vector_stores__vector_store_id__files__file_id__content_get", + "operationId": "vector_store_file_content_v1_vector_stores__vector_store_id__files__file_id__content_get_2", "parameters": [ { "in": "path", @@ -28889,7 +28889,7 @@ "/v1/vector_stores/{vector_store_id}/search": { "post": { "description": "Search a vector store.\n\nAPI Reference:\nhttps://platform.openai.com/docs/api-reference/vector-stores/search", - "operationId": "vector_store_search_v1_vector_stores__vector_store_id__search_post", + "operationId": "vector_store_search_v1_vector_stores__vector_store_id__search_post_2", "parameters": [ { "in": "path", @@ -28935,7 +28935,7 @@ "/vector_stores": { "get": { "description": "List vector stores.\n\nAPI Reference:\nhttps://platform.openai.com/docs/api-reference/vector-stores/list", - "operationId": "vector_store_list_vector_stores_get", + "operationId": "vector_store_list_vector_stores_get_2", "parameters": [ { "in": "query", @@ -29036,7 +29036,7 @@ }, "post": { "description": "Create a vector store.\n\nAPI Reference:\nhttps://platform.openai.com/docs/api-reference/vector-stores/create\n\nSupports target_model_names parameter for creating vector stores across multiple models:\n```json\n{\n \"name\": \"my-vector-store\",\n \"target_model_names\": \"gpt-4,gemini-2.0\"\n}\n```", - "operationId": "vector_store_create_vector_stores_post", + "operationId": "vector_store_create_vector_stores_post_2", "responses": { "200": { "content": { @@ -29061,7 +29061,7 @@ "/vector_stores/{vector_store_id}": { "delete": { "description": "Delete a vector store.\n\nAPI Reference:\nhttps://platform.openai.com/docs/api-reference/vector-stores/delete", - "operationId": "vector_store_delete_vector_stores__vector_store_id__delete", + "operationId": "vector_store_delete_vector_stores__vector_store_id__delete_2", "parameters": [ { "in": "path", @@ -29105,7 +29105,7 @@ }, "get": { "description": "Retrieve a vector store.\n\nAPI Reference:\nhttps://platform.openai.com/docs/api-reference/vector-stores/retrieve", - "operationId": "vector_store_retrieve_vector_stores__vector_store_id__get", + "operationId": "vector_store_retrieve_vector_stores__vector_store_id__get_2", "parameters": [ { "in": "path", @@ -29149,7 +29149,7 @@ }, "post": { "description": "Update a vector store.\n\nAPI Reference:\nhttps://platform.openai.com/docs/api-reference/vector-stores/modify", - "operationId": "vector_store_update_vector_stores__vector_store_id__post", + "operationId": "vector_store_update_vector_stores__vector_store_id__post_2", "parameters": [ { "in": "path", @@ -29194,7 +29194,7 @@ }, "/vector_stores/{vector_store_id}/files": { "get": { - "operationId": "vector_store_file_list_vector_stores__vector_store_id__files_get", + "operationId": "vector_store_file_list_vector_stores__vector_store_id__files_get_2", "parameters": [ { "in": "path", @@ -29237,7 +29237,7 @@ ] }, "post": { - "operationId": "vector_store_file_create_vector_stores__vector_store_id__files_post", + "operationId": "vector_store_file_create_vector_stores__vector_store_id__files_post_2", "parameters": [ { "in": "path", @@ -29282,7 +29282,7 @@ }, "/vector_stores/{vector_store_id}/files/{file_id}": { "delete": { - "operationId": "vector_store_file_delete_vector_stores__vector_store_id__files__file_id__delete", + "operationId": "vector_store_file_delete_vector_stores__vector_store_id__files__file_id__delete_2", "parameters": [ { "in": "path", @@ -29334,7 +29334,7 @@ ] }, "get": { - "operationId": "vector_store_file_retrieve_vector_stores__vector_store_id__files__file_id__get", + "operationId": "vector_store_file_retrieve_vector_stores__vector_store_id__files__file_id__get_2", "parameters": [ { "in": "path", @@ -29386,7 +29386,7 @@ ] }, "post": { - "operationId": "vector_store_file_update_vector_stores__vector_store_id__files__file_id__post", + "operationId": "vector_store_file_update_vector_stores__vector_store_id__files__file_id__post_2", "parameters": [ { "in": "path", @@ -29440,7 +29440,7 @@ }, "/vector_stores/{vector_store_id}/files/{file_id}/content": { "get": { - "operationId": "vector_store_file_content_vector_stores__vector_store_id__files__file_id__content_get", + "operationId": "vector_store_file_content_vector_stores__vector_store_id__files__file_id__content_get_2", "parameters": [ { "in": "path", @@ -29495,7 +29495,7 @@ "/vector_stores/{vector_store_id}/search": { "post": { "description": "Search a vector store.\n\nAPI Reference:\nhttps://platform.openai.com/docs/api-reference/vector-stores/search", - "operationId": "vector_store_search_vector_stores__vector_store_id__search_post", + "operationId": "vector_store_search_vector_stores__vector_store_id__search_post_2", "parameters": [ { "in": "path", diff --git a/litellm/proxy/_lazy_openapi_snapshot.py b/litellm/proxy/_lazy_openapi_snapshot.py index 309a0276aac..c63ff8d0733 100644 --- a/litellm/proxy/_lazy_openapi_snapshot.py +++ b/litellm/proxy/_lazy_openapi_snapshot.py @@ -8,12 +8,35 @@ any drift as a neutral check. """ import json +import re import sys from pathlib import Path -from typing import Dict, Optional +from typing import Dict, Optional, Set SNAPSHOT_FILE = Path(__file__).parent / "_lazy_openapi_snapshot.json" -HTTP_METHODS = {"delete", "get", "head", "options", "patch", "post", "put"} +HTTP_METHOD_SUFFIXES = { + "delete", + "get", + "head", + "options", + "patch", + "post", + "put", + "trace", +} + + +def _stabilize_multi_method_route_ids(routes) -> None: + """FastAPI derives route IDs from a set of methods; make snapshots stable.""" + + for route in routes: + methods = sorted(getattr(route, "methods", None) or []) + if len(methods) <= 1 or not getattr(route, "path_format", None): + continue + + operation_id = f"{route.name}{route.path_format}" + operation_id = re.sub(r"\W", "_", operation_id) + route.unique_id = f"{operation_id}_{methods[0].lower()}" def load_snapshot() -> Optional[Dict[str, Dict]]: @@ -38,12 +61,12 @@ def _normalize_operation_ids(paths: Dict[str, Dict]) -> None: if not isinstance(path_ops, dict): continue - methods = {method for method in path_ops if method in HTTP_METHODS} + methods = {method for method in path_ops if method in HTTP_METHOD_SUFFIXES} if not methods: continue for method, operation in path_ops.items(): - if method not in HTTP_METHODS or not isinstance(operation, dict): + if method not in HTTP_METHOD_SUFFIXES or not isinstance(operation, dict): continue operation_id = operation.get("operationId") @@ -65,7 +88,7 @@ def generate_snapshot() -> Dict[str, Dict]: from fastapi.openapi.utils import get_openapi from litellm.proxy._lazy_features import LAZY_FEATURES - from litellm.proxy.proxy_server import app + from litellm.proxy.proxy_server import app, ensure_unique_openapi_operation_ids for feat in LAZY_FEATURES: if feat.module_path in sys.modules: @@ -77,6 +100,7 @@ def generate_snapshot() -> Dict[str, Dict]: sys.stderr.write(f"warning: skip {feat.name}: {exc}\n") fragments: Dict[str, Dict] = {} + used_operation_ids: Set[str] = set() for feat in LAZY_FEATURES: feat_routes = [ r @@ -85,14 +109,24 @@ def generate_snapshot() -> Dict[str, Dict]: ] if not feat_routes: continue + _stabilize_multi_method_route_ids(feat_routes) full = get_openapi(title=app.title, version=app.version, routes=feat_routes) paths = full.get("paths", {}) _normalize_operation_ids(paths) # Group all of a feature's routes under one tag. - for path_ops in paths.values(): - for op in path_ops.values(): + for path_ops in full.get("paths", {}).values(): + for method, op in path_ops.items(): if isinstance(op, dict): + operation_id = op.get("operationId") + if isinstance(operation_id, str): + for suffix in HTTP_METHOD_SUFFIXES: + if operation_id.endswith(f"_{suffix}"): + op["operationId"] = ( + operation_id[: -len(suffix)] + method + ) + break op["tags"] = [feat.name] + full = ensure_unique_openapi_operation_ids(full, used_operation_ids) fragments[feat.name] = { "paths": paths, "components": {"schemas": full.get("components", {}).get("schemas", {})}, diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 8520e03f834..3cca23f07ab 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -724,21 +724,73 @@ class LiteLLMRoutes(enum.Enum): "/organization/member_delete", ] - # Routes accessible by Admin Viewer (read-only admin access) - admin_viewer_routes = [ - "/user/list", - "/user/available_users", - "/user/available_roles", - "/user/daily/activity", - "/team/daily/activity", - "/tag/daily/activity", - "/tag/list", - "/audit", - "/audit/{id}", - "/global/activity", - "/global/activity/model", - "/global/activity/cache_hits", - ] + info_routes + # Routes accessible by Admin Viewer (read-only admin access). + # + # Admin Viewer follows a read-parity-with-Proxy-Admin rule: anything Proxy + # Admin can read/list/get, Admin Viewer can too (no writes, no cost-incurring + # actions). + # + # NOTE: This list is no longer the primary mechanism for granting access — + # `_check_proxy_admin_viewer_access()` in route_checks.py default-allows + # any safe HTTP method (GET/HEAD/OPTIONS) on non-inference routes. This + # list now matters only for non-GET routes that are semantically reads + # (e.g. POST /spend/calculate). Adding a new GET endpoint does not require + # updating this list — the default-allow behavior covers it automatically. + admin_viewer_routes = ( + [ + "/user/list", + "/user/available_users", + "/user/available_roles", + "/user/daily/activity", + "/team/daily/activity", + "/tag/daily/activity", + "/tag/list", + "/audit", + "/audit/{id}", + "/global/activity", + "/global/activity/model", + "/global/activity/cache_hits", + # Customer / end-user listing (handlers already gate on + # PROXY_ADMIN_VIEW_ONLY — the route gate must match). + "/customer/list", + "/customer/info", + # UI Logs page detail drawer (single + session). The list endpoint + # `/spend/logs/ui` is covered via spend_tracking_routes below. + "/spend/logs/ui/{logId}", + "/spend/logs/session/ui", + # Settings / observability read endpoints exposed in admin-only + # sidebar groups (Logging & Alerts, Admin Settings, Budgets, + # Invitations). + "/callbacks/list", + "/callbacks/configs", + "/get/config/callbacks", + "/alerting/settings", + "/config/list", + "/config/field/info", + "/budget/list", + "/budget/settings", + # Invitation viewing (admin viewer cannot create/delete; can read). + "/invitation/info", + # Guardrails / Policies pages (read-only views). + "/guardrails/list", + "/v2/guardrails/list", + "/guardrails/submissions", + "/guardrails/submissions/{guardrail_id}", + "/guardrails/usage/overview", + "/policies/attachments/list", + # MCP semantic filter settings (read). + "/get/mcp_semantic_filter_settings", + # Model cost map maintenance views (read-only status / source). + "/schedule/model_cost_map_reload/status", + "/model/cost_map/source", + ] + # Spend tracking reads (/spend/logs, /spend/logs/ui, /spend/keys, + # /spend/users, /spend/tags, /spend/calculate, /cost/estimate). Admin + # Viewer can already read /global/spend/* via global_spend_tracking_routes; + # the per-tenant /spend/* views were the missing peer. + + spend_tracking_routes + + info_routes + ) # All routes accesible by an Org Admin org_admin_allowed_routes = ( @@ -2386,6 +2438,10 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase): None, description="CIDR ranges of trusted reverse proxies. When set, X-Forwarded-For headers are only trusted from these IPs.", ) + trusted_proxy_ranges: Optional[List[str]] = Field( + None, + description="CIDR ranges of trusted reverse proxies allowed to provide identity headers for header-based auth paths such as enable_oauth2_proxy_auth and custom_ui_sso_sign_in_handler.", + ) store_model_in_db: Optional[bool] = Field( None, description="If True, models and config are stored in and loaded from the database. Default is False.", @@ -2579,6 +2635,7 @@ class UserAPIKeyAuth( user_spend: Optional[float] = None user_max_budget: Optional[float] = None request_route: Optional[str] = None + budget_reservation: Optional[Dict[str, Any]] = Field(default=None, exclude=True) user: Optional[Any] = None # Expanded user object when expand=user is used created_by_user: Optional[Any] = ( None # Expanded created_by user when expand=user is used diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 113a8f538c0..754488367e2 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -60,6 +60,10 @@ from litellm.proxy._types import ( UserAPIKeyAuth, ) from litellm.proxy.auth.route_checks import RouteChecks +from litellm.proxy.common_utils.http_parsing_utils import ( + _safe_get_request_headers, + _safe_get_request_query_params, +) from litellm.proxy.db.exception_handler import PrismaDBExceptionHandler from litellm.proxy.guardrails.tool_name_extraction import ( TOOL_CAPABLE_CALL_TYPES, @@ -486,7 +490,10 @@ async def common_checks( # noqa: PLR0915 from litellm.proxy.proxy_server import prisma_client, user_api_key_cache _model: Optional[Union[str, List[str]]] = get_model_from_request( - request_body, route + request_data=request_body, + route=route, + request_headers=_safe_get_request_headers(request=request), + request_query_params=_safe_get_request_query_params(request=request), ) # 1. If team is blocked @@ -495,23 +502,28 @@ async def common_checks( # noqa: PLR0915 f"Team={team_object.team_id} is blocked. Update via `/team/unblock` if you're an admin." ) - # 2. If team can call model + # 2. If team can call model (or key's access_group_ids grant it) if _model and team_object: with tracer.trace("litellm.proxy.auth.common_checks.can_team_access_model"): - if not await can_team_access_model( - model=_model, - team_object=team_object, - llm_router=llm_router, - team_model_aliases=( - valid_token.team_model_aliases if valid_token else None - ), - ): - raise ProxyException( - message=f"Team not allowed to access model. Team={team_object.team_id}, Model={_model}. Allowed team models = {team_object.models}", - type=ProxyErrorTypes.team_model_access_denied, - param="model", - code=status.HTTP_401_UNAUTHORIZED, + try: + await can_team_access_model( + model=_model, + team_object=team_object, + llm_router=llm_router, + team_model_aliases=( + valid_token.team_model_aliases if valid_token else None + ), ) + except ProxyException as team_denial: + if team_denial.type != ProxyErrorTypes.team_model_access_denied: + raise + if not await _key_access_group_grants_model( + model=_model, + valid_token=valid_token, + team_object=team_object, + llm_router=llm_router, + ): + raise # 2.2. If team member has per-member model scope, enforce it if _model and team_object and valid_token and valid_token.user_id: @@ -656,13 +668,7 @@ async def common_checks( # noqa: PLR0915 end_user_object is not None and end_user_object.litellm_budget_table is not None ): - end_user_budget = end_user_object.litellm_budget_table.max_budget - if end_user_budget is not None and end_user_object.spend > end_user_budget: - raise litellm.BudgetExceededError( - current_cost=end_user_object.spend, - max_budget=end_user_budget, - message=f"ExceededBudget: End User={end_user_object.user_id} over budget. Spend={end_user_object.spend}, Budget={end_user_budget}", - ) + await _check_end_user_budget(end_user_obj=end_user_object, route=route) _enforce_user_param_check(general_settings, request, request_body, route) _reject_clientside_metadata_tags_check(general_settings, request_body, route) @@ -1012,7 +1018,7 @@ async def _apply_default_budget_to_end_user( return end_user_obj -def _check_end_user_budget( +async def _check_end_user_budget( end_user_obj: LiteLLM_EndUserTable, route: str, ) -> None: @@ -1033,11 +1039,20 @@ def _check_end_user_budget( return end_user_budget = end_user_obj.litellm_budget_table.max_budget - if end_user_budget is not None and end_user_obj.spend > end_user_budget: + if end_user_budget is None: + return + + from litellm.proxy.proxy_server import get_current_spend + + end_user_spend = await get_current_spend( + counter_key=f"spend:end_user:{end_user_obj.user_id}", + fallback_spend=end_user_obj.spend or 0.0, + ) + if end_user_spend > end_user_budget: raise litellm.BudgetExceededError( - current_cost=end_user_obj.spend, + current_cost=end_user_spend, max_budget=end_user_budget, - message=f"ExceededBudget: End User={end_user_obj.user_id} over budget. Spend={end_user_obj.spend}, Budget={end_user_budget}", + message=f"ExceededBudget: End User={end_user_obj.user_id} over budget. Spend={end_user_spend}, Budget={end_user_budget}", ) @@ -1091,7 +1106,7 @@ async def get_end_user_object( ) # Check budget limits - _check_end_user_budget(end_user_obj=return_obj, route=route) + await _check_end_user_budget(end_user_obj=return_obj, route=route) return return_obj @@ -1124,7 +1139,7 @@ async def get_end_user_object( ) # Check budget limits - _check_end_user_budget(end_user_obj=_response, route=route) + await _check_end_user_budget(end_user_obj=_response, route=route) return _response @@ -1616,9 +1631,12 @@ async def _cache_key_object( ## CACHE REFRESH TIME user_api_key_obj.last_refreshed_at = time.time() + cached_key_obj = _copy_user_api_key_auth_for_cache( + user_api_key_obj=user_api_key_obj + ) await _cache_management_object( key=key, - value=user_api_key_obj, + value=cached_key_obj, user_api_key_cache=user_api_key_cache, proxy_logging_obj=proxy_logging_obj, model_type=UserAPIKeyAuth, @@ -2348,7 +2366,7 @@ async def get_key_object( model_type=UserAPIKeyAuth, ) if user_api_key_auth is not None: - return user_api_key_auth + return _copy_user_api_key_auth_for_cache(user_api_key_obj=user_api_key_auth) if check_cache_only: raise Exception( @@ -2401,6 +2419,16 @@ async def get_key_object( return _response +def _copy_user_api_key_auth_for_cache( + user_api_key_obj: UserAPIKeyAuth, +) -> UserAPIKeyAuth: + copied_key_obj = user_api_key_obj.model_copy() + copied_key_obj.budget_reservation = None + copied_key_obj.parent_otel_span = None + copied_key_obj.request_route = None + return copied_key_obj + + @log_db_metrics async def get_object_permission( object_permission_id: str, @@ -2952,6 +2980,77 @@ async def can_team_access_model( raise +async def _key_access_group_grants_model( + model: Union[str, List[str]], + valid_token: Optional[UserAPIKeyAuth], + team_object: Optional[LiteLLM_TeamTable], + llm_router: Optional[Router], +) -> bool: + """ + Returns True if the key's `access_group_ids` expand to models that grant + access to `model`. Used to let a key's access group override a team's + model restriction in `common_checks`. + + A key's access group only counts if the access group itself authorizes the + caller as an owner — that is, the group's `assigned_team_ids` includes the + key's `team_id`, or the group's `assigned_key_ids` includes the key's + token. This preserves the team-as-owner boundary (a team member cannot + escalate by naming a group assigned to a different team) while still + letting a group reach the key without first being added to the team's + `access_group_ids` list. + """ + if valid_token is None: + return False + key_access_group_ids = list(valid_token.access_group_ids or []) + if not key_access_group_ids: + return False + + from litellm.proxy.proxy_server import prisma_client as _prisma_client + from litellm.proxy.proxy_server import proxy_logging_obj as _proxy_logging_obj + from litellm.proxy.proxy_server import user_api_key_cache as _user_api_key_cache + + if _prisma_client is None or _user_api_key_cache is None: + return False + + key_team_id = valid_token.team_id or ( + team_object.team_id if team_object is not None else None + ) + key_token = valid_token.token + + authorized_models: List[str] = [] + for ag_id in key_access_group_ids: + try: + ag = await get_access_object( + access_group_id=ag_id, + prisma_client=_prisma_client, + user_api_key_cache=_user_api_key_cache, + proxy_logging_obj=_proxy_logging_obj, + ) + except Exception: + continue + team_authorized = bool( + key_team_id and key_team_id in (ag.assigned_team_ids or []) + ) + key_authorized = bool(key_token and key_token in (ag.assigned_key_ids or [])) + if team_authorized or key_authorized: + authorized_models.extend(ag.access_model_names or []) + + if not authorized_models: + return False + try: + _can_object_call_model( + model=model, + llm_router=llm_router, + models=list(set(authorized_models)), + team_model_aliases=valid_token.team_model_aliases, + team_id=valid_token.team_id, + object_type="key", + ) + return True + except ProxyException: + return False + + def can_project_access_model( model: Union[str, List[str]], project_object: LiteLLM_ProjectTableCachedObj, @@ -3967,13 +4066,19 @@ async def _tag_max_budget_check( if ( tag_object.litellm_budget_table is not None and tag_object.litellm_budget_table.max_budget is not None - and tag_object.spend is not None - and tag_object.spend > tag_object.litellm_budget_table.max_budget ): + from litellm.proxy.proxy_server import get_current_spend + + tag_spend = await get_current_spend( + counter_key=f"spend:tag:{tag_name}", + fallback_spend=tag_object.spend or 0.0, + ) + if tag_spend <= tag_object.litellm_budget_table.max_budget: + continue raise litellm.BudgetExceededError( - current_cost=tag_object.spend, + current_cost=tag_spend, max_budget=tag_object.litellm_budget_table.max_budget, - message=f"Budget has been exceeded! Tag={tag_name} Current cost: {tag_object.spend}, Max budget: {tag_object.litellm_budget_table.max_budget}", + message=f"Budget has been exceeded! Tag={tag_name} Current cost: {tag_spend}, Max budget: {tag_object.litellm_budget_table.max_budget}", ) diff --git a/litellm/proxy/auth/auth_utils.py b/litellm/proxy/auth/auth_utils.py index cbed34adacf..51108827f6b 100644 --- a/litellm/proxy/auth/auth_utils.py +++ b/litellm/proxy/auth/auth_utils.py @@ -2,7 +2,7 @@ import os import re import sys from functools import lru_cache -from typing import Any, List, Optional, Tuple +from typing import Any, Dict, List, Mapping, Optional, Tuple, Union from fastapi import HTTPException, Request, status @@ -976,20 +976,257 @@ def get_end_user_id_from_request_body( return None -def get_model_from_request( - request_data: dict, route: str -) -> Optional[Union[str, List[str]]]: - # First try to get model from request_data - model = request_data.get("model") or request_data.get("target_model_names") +MODEL_ROUTING_HEADER_NAME = "x-litellm-model" +_MODEL_ROUTING_ROUTE_MARKERS = ( + "/files", + "/batches", + "/vector_stores", + "/skills", + "/evals", + "/fine_tuning", + "/videos", +) +_MODEL_ROUTING_HEADER_OR_QUERY_ROUTE_MARKERS = ( + "/files", + "/batches", + "/skills", + "/evals", +) +_MODEL_ROUTING_QUERY_TARGET_MODEL_ROUTE_MARKERS = ( + "/files", + "/batches", + "/fine_tuning", +) +_MODEL_ROUTING_BODY_TARGET_MODEL_ROUTE_MARKERS = ( + "/files", + "/batches", + "/vector_stores", +) +_MODEL_ROUTING_COMPLETION_MODEL_ROUTE_MARKERS = ("/evals",) +_MODEL_ROUTING_ID_FIELDS = ( + "file_id", + "input_file_id", + "output_file_id", + "error_file_id", + "batch_id", + "fine_tuning_job_id", + "training_file", + "validation_file", + "vector_store_id", + "video_id", + "character_id", +) - if model is not None: - model_names = model.split(",") - if len(model_names) == 1: - model = model_names[0].strip() + +def _append_model_candidates(candidates: List[str], value: Any) -> None: + if value is None: + return + + values = value if isinstance(value, (list, tuple, set)) else [value] + for item in values: + if item is None: + continue + if isinstance(item, str): + model_names = [model.strip() for model in item.split(",")] else: - model = [m.strip() for m in model_names] + model_names = [str(item).strip()] + candidates.extend(model for model in model_names if model) - # If model not in request_data, try to extract from route + +def _dedupe_model_candidates(candidates: List[str]) -> List[str]: + deduped: List[str] = [] + for model in candidates: + if model not in deduped: + deduped.append(model) + return deduped + + +def _get_case_insensitive_mapping_value( + mapping: Optional[Mapping[str, Any]], key: str +) -> Any: + if not mapping: + return None + if key in mapping: + return mapping[key] + key_lower = key.lower() + for mapping_key, value in mapping.items(): + if str(mapping_key).lower() == key_lower: + return value + return None + + +def _route_matches_any_marker(route: str, markers: Tuple[str, ...]) -> bool: + normalized_route = route.lower() + return any(marker in normalized_route for marker in markers) + + +def _route_uses_model_routing_sources(route: str) -> bool: + return _route_matches_any_marker(route=route, markers=_MODEL_ROUTING_ROUTE_MARKERS) + + +def _extract_models_from_managed_resource_id( + resource_id: Any, resource_id_field: Optional[str] = None +) -> List[str]: + if not isinstance(resource_id, str) or not resource_id: + return [] + + candidates: List[str] = [] + + try: + from litellm.proxy.openai_files_endpoints.common_utils import ( + _is_base64_encoded_unified_file_id, + decode_model_from_file_id, + get_model_id_from_unified_batch_id, + get_models_from_unified_file_id, + ) + + _append_model_candidates( + candidates=candidates, value=decode_model_from_file_id(resource_id) + ) + unified_file_id = _is_base64_encoded_unified_file_id(resource_id) + if unified_file_id: + _append_model_candidates( + candidates=candidates, + value=get_models_from_unified_file_id(unified_file_id), + ) + _append_model_candidates( + candidates=candidates, + value=get_model_id_from_unified_batch_id(unified_file_id), + ) + except Exception as e: + verbose_proxy_logger.debug( + "Unable to extract model from managed file/batch ID: %s", str(e) + ) + + try: + from litellm.llms.base_llm.managed_resources.utils import parse_unified_id + + parsed_id = parse_unified_id(resource_id) + if parsed_id: + _append_model_candidates( + candidates=candidates, value=parsed_id.get("model_id") + ) + _append_model_candidates( + candidates=candidates, value=parsed_id.get("target_model_names") + ) + except Exception as e: + verbose_proxy_logger.debug( + "Unable to extract model from unified managed resource ID: %s", str(e) + ) + + if resource_id_field in ("video_id", "character_id"): + try: + from litellm.types.videos.utils import ( + decode_character_id_with_provider, + decode_video_id_with_provider, + ) + + if resource_id_field == "video_id": + _append_model_candidates( + candidates=candidates, + value=decode_video_id_with_provider(resource_id).get("model_id"), + ) + else: + _append_model_candidates( + candidates=candidates, + value=decode_character_id_with_provider(resource_id).get( + "model_id" + ), + ) + except Exception as e: + verbose_proxy_logger.debug( + "Unable to extract model from managed video/character ID: %s", str(e) + ) + + return _dedupe_model_candidates(candidates) + + +def _extract_model_candidates_from_request( + request_data: dict, + route: str, + request_headers: Optional[Mapping[str, Any]] = None, + request_query_params: Optional[Mapping[str, Any]] = None, +) -> List[str]: + candidates: List[str] = [] + uses_model_routing_sources = _route_uses_model_routing_sources(route=route) + uses_header_or_query_model_sources = _route_matches_any_marker( + route=route, markers=_MODEL_ROUTING_HEADER_OR_QUERY_ROUTE_MARKERS + ) + uses_query_target_model_sources = _route_matches_any_marker( + route=route, markers=_MODEL_ROUTING_QUERY_TARGET_MODEL_ROUTE_MARKERS + ) + uses_body_target_model_sources = _route_matches_any_marker( + route=route, markers=_MODEL_ROUTING_BODY_TARGET_MODEL_ROUTE_MARKERS + ) + uses_completion_model_sources = _route_matches_any_marker( + route=route, markers=_MODEL_ROUTING_COMPLETION_MODEL_ROUTE_MARKERS + ) + + body_model = request_data.get("model") + _append_model_candidates(candidates, body_model) + if uses_body_target_model_sources or not body_model: + _append_model_candidates(candidates, request_data.get("target_model_names")) + if uses_completion_model_sources and isinstance( + request_data.get("completion"), dict + ): + _append_model_candidates(candidates, request_data["completion"].get("model")) + + if uses_model_routing_sources: + if uses_header_or_query_model_sources: + _append_model_candidates( + candidates, + _get_case_insensitive_mapping_value(request_query_params, "model"), + ) + _append_model_candidates( + candidates, + _get_case_insensitive_mapping_value( + request_headers, MODEL_ROUTING_HEADER_NAME + ), + ) + if uses_query_target_model_sources: + _append_model_candidates( + candidates, + _get_case_insensitive_mapping_value( + request_query_params, "target_model_names" + ), + ) + + for field in _MODEL_ROUTING_ID_FIELDS: + _append_model_candidates( + candidates, + _extract_models_from_managed_resource_id( + request_data.get(field), resource_id_field=field + ), + ) + + return _dedupe_model_candidates(candidates) + + +def _format_model_candidates( + candidates: List[str], +) -> Optional[Union[str, List[str]]]: + if not candidates: + return None + if len(candidates) == 1: + return candidates[0] + return candidates + + +def get_model_from_request( + request_data: dict, + route: str, + request_headers: Optional[Mapping[str, Any]] = None, + request_query_params: Optional[Mapping[str, Any]] = None, +) -> Optional[Union[str, List[str]]]: + candidates = _extract_model_candidates_from_request( + request_data=request_data, + route=route, + request_headers=request_headers, + request_query_params=request_query_params, + ) + model = _format_model_candidates(candidates) + + # If no explicit model was found, try to extract from route if model is None: # Parse model from route that follows the pattern /openai/deployments/{model}/* match = re.match(r"/openai/deployments/([^/]+)", route) diff --git a/litellm/proxy/auth/handle_jwt.py b/litellm/proxy/auth/handle_jwt.py index 71411bed7fd..d1fd5818f35 100644 --- a/litellm/proxy/auth/handle_jwt.py +++ b/litellm/proxy/auth/handle_jwt.py @@ -707,11 +707,48 @@ class JWTHandler: verbose_proxy_logger.error(f"Error fetching OIDC UserInfo: {str(e)}") raise Exception(f"Failed to fetch OIDC UserInfo: {str(e)}") - async def auth_jwt(self, token: str) -> dict: + _unscoped_jwt_warning_emitted = False + + @classmethod + def _build_decode_kwargs(cls) -> dict: + """Build the audience/issuer/options kwargs for ``jwt.decode``. + + Setting ``JWT_AUDIENCE`` (and optionally ``JWT_ISSUER``) turns on the + corresponding PyJWT verifications, blocking cross-tenant tokens + minted by other applications that share the same IdP signing keys. + When both are unset PyJWT only checks the signature and expiry, which + is preserved for backward compatibility but logged once as a warning. + """ audience = os.getenv("JWT_AUDIENCE") - decode_options = None + issuer = os.getenv("JWT_ISSUER") + + if ( + audience is None + and issuer is None + and not cls._unscoped_jwt_warning_emitted + ): + verbose_proxy_logger.warning( + "JWT auth is enabled but neither JWT_AUDIENCE nor JWT_ISSUER " + "is configured. Tokens minted by any application that shares " + "the same IdP signing keys will be accepted. Set JWT_AUDIENCE " + "(and ideally JWT_ISSUER) to scope this proxy." + ) + cls._unscoped_jwt_warning_emitted = True + + options: dict = {} if audience is None: - decode_options = {"verify_aud": False} + options["verify_aud"] = False + if issuer is None: + options["verify_iss"] = False + + return { + "audience": audience, + "issuer": issuer, + "options": options or None, + } + + async def auth_jwt(self, token: str) -> dict: + decode_kwargs = self._build_decode_kwargs() header = jwt.get_unverified_header(token) @@ -747,9 +784,8 @@ class JWTHandler: token, public_key_obj, # type: ignore algorithms=self.SUPPORTED_JWT_ALGORITHMS, - options=decode_options, # type: ignore[arg-type] - audience=audience, leeway=self.leeway, # allow testing of expired tokens + **decode_kwargs, ) return payload @@ -775,8 +811,7 @@ class JWTHandler: token, key, algorithms=self.SUPPORTED_JWT_ALGORITHMS, - audience=audience, - options=decode_options, + **decode_kwargs, ) return payload diff --git a/litellm/proxy/auth/oauth2_proxy_hook.py b/litellm/proxy/auth/oauth2_proxy_hook.py index 0dc696bc455..9fc4c4fb531 100644 --- a/litellm/proxy/auth/oauth2_proxy_hook.py +++ b/litellm/proxy/auth/oauth2_proxy_hook.py @@ -1,19 +1,69 @@ -from typing import Any, Dict +from typing import Any, Dict, FrozenSet from fastapi import Request from litellm._logging import verbose_proxy_logger from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.auth.trusted_proxy_utils import require_trusted_proxy_request + +# OAuth2-proxy header trust is for **identity assertion** from a trusted +# upstream auth proxy (oauth2-proxy, Authelia, etc.). The allowlist below +# is the only safe surface — anything else (``user_role``, ``api_key``, +# ``permissions``, ``max_budget``, ``user_max_budget``, +# ``team_tpm_limit``, ``end_user_max_budget``, ``allowed_model_region``, +# and dozens of similar policy fields scattered across the +# ``LiteLLM_VerificationTokenView`` hierarchy) is a privilege grant that +# would let a caller forge their own enforcement parameters by sending +# the matching header. +# +# A denylist of "privileged fields" is unmaintainable in this codebase: +# the auth model has ~50 budget/spend/limit/permission fields and gains +# more with each release. An allowlist scoped to identity assertion is +# default-secure — new fields are blocked automatically. +# +# Operators who need a trusted upstream to assert anything beyond +# identity should switch to JWT authentication, which validates a +# signature on the assertion rather than blindly trusting headers. +ALLOWED_OAUTH2_PROXY_FIELDS: FrozenSet[str] = frozenset( + { + "user_id", + "user_email", + "team_id", + "team_alias", + "org_id", + "models", + } +) async def handle_oauth2_proxy_request(request: Request) -> UserAPIKeyAuth: """ - Handle request from oauth2 proxy. + Resolve a ``UserAPIKeyAuth`` from request headers per the admin-set + ``oauth2_config_mappings``. + + The auth model assumes the proxy is deployed behind a trusted OAuth2 + reverse proxy that injects authenticated identity headers (e.g. + oauth2-proxy, Authelia). + + **Identity-only allowlist.** ``oauth2_config_mappings`` maps header + names to ``UserAPIKeyAuth`` fields. Without an allowlist, an admin + who maps the wrong header to ``user_role`` lets any caller send + ``X-User-Role: proxy_admin`` and gain full admin privileges + (Pydantic coerces the string into the enum). Only fields in + ``ALLOWED_OAUTH2_PROXY_FIELDS`` (identity assertion only — see the + constant's comment) may be mapped; any other mapping is rejected at + request time so the misconfiguration surfaces loudly rather than as + a silent privesc. """ from litellm.proxy.proxy_server import general_settings verbose_proxy_logger.debug("Handling oauth2 proxy request") - # Define the OAuth2 config mappings + require_trusted_proxy_request( + request=request, + general_settings=general_settings, + feature_name="OAuth2 proxy auth", + ) + oauth2_config_mappings: Dict[str, str] = ( general_settings.get("oauth2_config_mappings") or {} ) @@ -21,21 +71,32 @@ async def handle_oauth2_proxy_request(request: Request) -> UserAPIKeyAuth: if not oauth2_config_mappings: raise ValueError("Oauth2 config mappings not found in general_settings") - # Initialize a dictionary to store the mapped values - auth_data: Dict[str, Any] = {} - # Extract values from headers based on the mappings + disallowed = sorted( + set(oauth2_config_mappings.keys()) - ALLOWED_OAUTH2_PROXY_FIELDS + ) + if disallowed: + raise ValueError( + "Oauth2 proxy auth refuses to map non-identity UserAPIKeyAuth " + f"fields from request headers: {disallowed}. Only identity " + f"fields are accepted ({sorted(ALLOWED_OAUTH2_PROXY_FIELDS)}); " + "anything else (privileges, budgets, rate limits, metadata) " + "would let a caller forge enforcement parameters by spoofing " + "the matching header. If you need a trusted upstream to " + "assert anything beyond identity, use JWT auth " + "(signature-validated) instead of header-trust." + ) + + auth_data: Dict[str, Any] = {} for key, header in oauth2_config_mappings.items(): value = request.headers.get(header) - if value: - # Convert max_budget to float if present - if key == "max_budget": - auth_data[key] = float(value) - # Convert models to list if present - elif key == "models": - auth_data[key] = [model.strip() for model in value.split(",")] - else: - auth_data[key] = value + if not value: + continue + if key == "models": + auth_data[key] = [model.strip() for model in value.split(",")] + else: + auth_data[key] = value + verbose_proxy_logger.debug( "Auth data before creating UserAPIKeyAuth object: keys=%s", list(auth_data.keys()), @@ -45,5 +106,4 @@ async def handle_oauth2_proxy_request(request: Request) -> UserAPIKeyAuth: "UserAPIKeyAuth object created with keys: %s", list(user_api_key_auth.__fields_set__), ) - # Create and return UserAPIKeyAuth object return user_api_key_auth diff --git a/litellm/proxy/auth/route_checks.py b/litellm/proxy/auth/route_checks.py index 6417307f691..dba29f84133 100644 --- a/litellm/proxy/auth/route_checks.py +++ b/litellm/proxy/auth/route_checks.py @@ -202,6 +202,7 @@ class RouteChecks: route=route, _user_role=_user_role, request_data=request_data, + request=request, ) elif ( _user_role == LitellmUserRoles.INTERNAL_USER.value @@ -596,14 +597,66 @@ class RouteChecks: return True return False + # HTTP methods that are intrinsically read-only and therefore safe to + # default-allow for PROXY_ADMIN_VIEW_ONLY. Anything else (POST/PUT/PATCH/ + # DELETE) is treated as a write attempt and goes through the explicit + # write-allowlist below. + _SAFE_HTTP_METHODS = frozenset({"GET", "HEAD", "OPTIONS"}) + + # Explicit write routes that PROXY_ADMIN_VIEW_ONLY must NEVER call. The + # role-principle is "no writes, ever" — the management_routes list is the + # authoritative source for which non-llm routes are writes; we just need + # to filter out the read endpoints (info / list) that share the prefix. + # A cleaner approach is to denylist by HTTP verb (POST/PUT/PATCH/DELETE); + # this block stays as a backstop in case a write is implemented as GET. + _ADMIN_VIEWER_BLOCKED_WRITE_ROUTES = frozenset( + [ + "/user/new", + "/user/delete", + "/user/bulk_update", + "/team/new", + "/team/update", + "/team/delete", + "/model/new", + "/model/update", + "/model/delete", + "/key/generate", + "/key/delete", + "/key/update", + "/key/regenerate", + "/key/service-account/generate", + "/key/block", + "/key/unblock", + ] + ) + @staticmethod def _check_proxy_admin_viewer_access( route: str, _user_role: str, request_data: dict, + request: Optional[Request] = None, ) -> None: """ - Check access for PROXY_ADMIN_VIEW_ONLY role + Check access for PROXY_ADMIN_VIEW_ONLY role. + + Admin Viewer follows a read-parity-with-Proxy-Admin rule: anything Proxy + Admin can read/list/get, Admin Viewer can read/list/get. The only + exclusions are cost-incurring inference routes (Playground, /chat/ + completions, etc.) and any state-mutating request. + + Implementation: + 1. LLM/inference routes → 403 (cost-incurring). + 2. Safe HTTP method (GET/HEAD/OPTIONS) → allow by default. This is + the read-parity guarantee — every new GET endpoint added anywhere + in the codebase is automatically readable by Admin Viewer + without needing to remember to add it to an allowlist. + 3. Unsafe HTTP method (POST/PUT/PATCH/DELETE): + - Allow `/user/update` only when restricted to user_email/password. + - Block all explicit writes in `_ADMIN_VIEWER_BLOCKED_WRITE_ROUTES`. + - Otherwise allow only if the route is in admin_viewer_routes / + global_spend_tracking_routes (legacy explicit-allow set). + - Else 403. """ if RouteChecks.is_llm_api_route(route=route): raise HTTPException( @@ -611,65 +664,60 @@ class RouteChecks: detail=f"user not allowed to access this OpenAI routes, role= {_user_role}", ) - # Check if this is a write operation on management routes - if RouteChecks.check_route_access( - route=route, allowed_routes=LiteLLMRoutes.management_routes.value - ): - # For management routes, only allow read operations or specific allowed updates - if route == "/user/update": - # Check the Request params are valid for PROXY_ADMIN_VIEW_ONLY - if request_data is not None and isinstance(request_data, dict): - _params_updated = request_data.keys() - for param in _params_updated: - if param not in ["user_email", "password"]: - raise HTTPException( - status_code=status.HTTP_403_FORBIDDEN, - detail=f"user not allowed to access this route, role= {_user_role}. Trying to access: {route} and updating invalid param: {param}. only user_email and password can be updated", - ) - elif ( - route - in [ - "/user/new", - "/user/delete", - "/user/bulk_update", - "/team/new", - "/team/update", - "/team/delete", - "/model/new", - "/model/update", - "/model/delete", - "/key/generate", - "/key/delete", - "/key/update", - "/key/regenerate", - "/key/service-account/generate", - "/key/block", - "/key/unblock", - ] - or route.startswith("/key/") - and route.endswith("/regenerate") - ): - # Block write operations for PROXY_ADMIN_VIEW_ONLY - raise HTTPException( - status_code=status.HTTP_403_FORBIDDEN, - detail=f"user not allowed to access this route, role= {_user_role}. Trying to access: {route}", - ) - # Allow read operations on management routes (like /user/info, /team/info, /model/info) + method = request.method.upper() if request is not None else "GET" + is_safe_method = method in RouteChecks._SAFE_HTTP_METHODS + + # ── Safe HTTP method: default-allow ────────────────────────────── + if is_safe_method: return - elif RouteChecks.check_route_access( - route=route, allowed_routes=LiteLLMRoutes.admin_viewer_routes.value - ): - # Allow access to admin viewer routes (read-only admin endpoints) + + # ── Unsafe HTTP method: explicit checks ────────────────────────── + # Allow `/user/update` for self-service email / password change. + if route == "/user/update": + if request_data is not None and isinstance(request_data, dict): + for param in request_data.keys(): + if param not in ["user_email", "password"]: + raise HTTPException( + status_code=status.HTTP_403_FORBIDDEN, + detail=( + f"user not allowed to access this route, role= {_user_role}. " + f"Trying to access: {route} and updating invalid param: {param}. " + "only user_email and password can be updated" + ), + ) return - elif RouteChecks.check_route_access( - route=route, allowed_routes=LiteLLMRoutes.global_spend_tracking_routes.value + + # Hard-block known write routes regardless of HTTP method (defensive + # — these are POSTs in practice, but pinning them here protects + # against future GET-shaped writes). + if route in RouteChecks._ADMIN_VIEWER_BLOCKED_WRITE_ROUTES or ( + route.startswith("/key/") and route.endswith("/regenerate") ): - # Allow access to global spend tracking routes (read-only spend endpoints) - # proxy_admin_viewer role description: "view all keys, view all spend" - return - else: - # For other routes, block access raise HTTPException( status_code=status.HTTP_403_FORBIDDEN, detail=f"user not allowed to access this route, role= {_user_role}. Trying to access: {route}", ) + + # Legacy explicit-allow sets (kept for routes that are POST but + # semantically read-only, e.g. /spend/calculate). Both admin_viewer_routes + # and global_spend_tracking_routes are reads/listings. + if RouteChecks.check_route_access( + route=route, allowed_routes=LiteLLMRoutes.admin_viewer_routes.value + ): + return + if RouteChecks.check_route_access( + route=route, allowed_routes=LiteLLMRoutes.global_spend_tracking_routes.value + ): + return + + # NOTE: We intentionally do NOT fall back to allowing all + # `management_routes`. That set is a mix of reads (info/list — handled + # via the safe-method branch above) and writes (`/team/block`, + # `/team/permissions_update`, `/jwt/key/mapping/{new,update,delete}`, + # `/key/bulk_update`, `/key/{id}/reset_spend`). A blanket allow would + # let Admin Viewer POST these write endpoints — violating the + # "no writes, ever" rule. Default-deny instead. + raise HTTPException( + status_code=status.HTTP_403_FORBIDDEN, + detail=f"user not allowed to access this route, role= {_user_role}. Trying to access: {route}", + ) diff --git a/litellm/proxy/auth/trusted_proxy_utils.py b/litellm/proxy/auth/trusted_proxy_utils.py new file mode 100644 index 00000000000..df7b3080f28 --- /dev/null +++ b/litellm/proxy/auth/trusted_proxy_utils.py @@ -0,0 +1,118 @@ +import ipaddress +from typing import Any, Dict, List, Optional, Union + +from fastapi import Request + +from litellm._logging import verbose_proxy_logger + +TRUSTED_PROXY_RANGES_KEY = "trusted_proxy_ranges" +TrustedProxyNetwork = Union[ipaddress.IPv4Network, ipaddress.IPv6Network] + + +def _get_proxy_general_settings() -> Dict[str, Any]: + try: + from litellm.proxy.proxy_server import general_settings + + return general_settings or {} + except ImportError: + return {} + + +def _normalize_cidr_ranges(configured_ranges: Any, *, setting_name: str) -> List[str]: + if not configured_ranges: + return [] + if isinstance(configured_ranges, str): + return [ + raw_range.strip() + for raw_range in configured_ranges.split(",") + if raw_range.strip() + ] + if isinstance(configured_ranges, (list, tuple, set)): + return [ + str(raw_range).strip() + for raw_range in configured_ranges + if str(raw_range).strip() + ] + verbose_proxy_logger.warning( + "Invalid %s value: expected a list of CIDR ranges, got %s", + setting_name, + type(configured_ranges).__name__, + ) + return [] + + +def parse_trusted_proxy_ranges( + configured_ranges: Any, + *, + setting_name: str = TRUSTED_PROXY_RANGES_KEY, +) -> List[TrustedProxyNetwork]: + networks: List[TrustedProxyNetwork] = [] + for cidr in _normalize_cidr_ranges(configured_ranges, setting_name=setting_name): + try: + networks.append(ipaddress.ip_network(cidr, strict=False)) + except ValueError: + verbose_proxy_logger.warning( + "Invalid CIDR in %s: %s, skipping", setting_name, cidr + ) + return networks + + +def _get_direct_client_ip(request: Request) -> Optional[str]: + client = getattr(request, "client", None) + client_host = getattr(client, "host", None) + if isinstance(client_host, str): + return client_host + return None + + +def _is_ip_in_networks( + client_ip: Optional[str], networks: List[TrustedProxyNetwork] +) -> bool: + if not client_ip or not networks: + return False + try: + addr = ipaddress.ip_address(client_ip.strip()) + except ValueError: + return False + return any(addr in network for network in networks) + + +def require_trusted_proxy_request( + *, + request: Request, + general_settings: Optional[Dict[str, Any]] = None, + feature_name: str, + setting_name: str = TRUSTED_PROXY_RANGES_KEY, +) -> None: + """ + Fail closed unless the direct TCP peer is one of the configured + trusted reverse proxies. + + Header-based auth paths must validate the direct peer, not + X-Forwarded-For, because the direct peer is the actor supplying the + identity headers. + """ + if general_settings is None: + general_settings = _get_proxy_general_settings() + + trusted_networks = parse_trusted_proxy_ranges( + general_settings.get(setting_name), setting_name=setting_name + ) + if not trusted_networks: + raise ValueError( + f"{feature_name} requires general_settings.{setting_name} before " + "trusting identity headers from an upstream proxy." + ) + + direct_client_ip = _get_direct_client_ip(request) + if not _is_ip_in_networks(direct_client_ip, trusted_networks): + verbose_proxy_logger.warning( + "%s rejected identity headers from untrusted direct client IP %r", + feature_name, + direct_client_ip, + ) + raise ValueError( + f"{feature_name} only accepts identity headers from configured " + f"trusted proxy ranges. Direct client IP {direct_client_ip!r} " + "is not trusted." + ) diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index bfd1f2e0b3a..9159a8ff9da 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -11,7 +11,7 @@ import asyncio import re import secrets from datetime import datetime, timezone -from typing import Any, List, Optional, Tuple, cast +from typing import Any, Iterator, List, Optional, Tuple, Union, cast import fastapi from fastapi import HTTPException, Request, WebSocket, status @@ -63,6 +63,7 @@ from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache from litellm.proxy.common_utils.http_parsing_utils import ( _read_request_body, _safe_get_request_headers, + _safe_get_request_query_params, populate_request_with_path_params, ) from litellm.proxy.common_utils.realtime_utils import _realtime_request_body @@ -118,6 +119,29 @@ azure_apim_header = APIKeyHeader( ) +def _get_model_from_request_context( + request_data: dict, + route: str, + request: Optional[Request], +) -> Optional[Union[str, List[str]]]: + return get_model_from_request( + request_data=request_data, + route=route, + request_headers=_safe_get_request_headers(request=request), + request_query_params=_safe_get_request_query_params(request=request), + ) + + +def _get_model_names_for_budget_checks( + model: Optional[Union[str, List[str]]], +) -> List[str]: + if model is None: + return [] + if isinstance(model, str): + return [model] + return model + + def _get_bearer_token_or_received_api_key(api_key: str) -> str: if api_key.startswith("Bearer "): # ensure Bearer token passed in api_key = api_key.replace("Bearer ", "") # extract the token @@ -884,7 +908,11 @@ async def _user_api_key_auth_builder( # noqa: PLR0915 ) # Check if model has zero cost - if so, skip all budget checks - model = get_model_from_request(request_data, route) + model = _get_model_from_request_context( + request_data=request_data, + route=route, + request=request, + ) skip_budget_checks = False if model is not None and llm_router is not None: from litellm.proxy.auth.auth_checks import _is_model_cost_zero @@ -1254,6 +1282,7 @@ async def _user_api_key_auth_builder( # noqa: PLR0915 valid_token=valid_token, request_data=request_data, route=route, + request=request, llm_model_list=llm_model_list, llm_router=llm_router, ) @@ -1279,7 +1308,11 @@ async def _user_api_key_auth_builder( # noqa: PLR0915 user_obj = None # Check 2a. Check if model has zero cost - if so, skip all budget checks - model = get_model_from_request(request_data, route) + model = _get_model_from_request_context( + request_data=request_data, + route=route, + request=request, + ) skip_budget_checks = False if model is not None and llm_router is not None: from litellm.proxy.auth.auth_checks import _is_model_cost_zero @@ -1403,21 +1436,29 @@ async def _user_api_key_auth_builder( # noqa: PLR0915 # Check 5. Token Model Spend is under Model budget max_budget_per_model = valid_token.model_max_budget - current_model = request_data.get("model", None) + current_model = _get_model_from_request_context( + request_data=request_data, + route=route, + request=request, + ) + current_models = _get_model_names_for_budget_checks( + model=current_model + ) if ( max_budget_per_model is not None and isinstance(max_budget_per_model, dict) and len(max_budget_per_model) > 0 and prisma_client is not None - and current_model is not None + and current_models and valid_token.token is not None ): ## GET THE SPEND FOR THIS MODEL - await model_max_budget_limiter.is_key_within_model_budget( - user_api_key_dict=valid_token, - model=current_model, - ) + for model_name in current_models: + await model_max_budget_limiter.is_key_within_model_budget( + user_api_key_dict=valid_token, + model=model_name, + ) # Check 5b. End-user model max budget end_user_mmb = valid_token.end_user_model_max_budget @@ -1425,14 +1466,15 @@ async def _user_api_key_auth_builder( # noqa: PLR0915 end_user_mmb is not None and isinstance(end_user_mmb, dict) and len(end_user_mmb) > 0 - and current_model is not None + and current_models and valid_token.end_user_id is not None ): - await model_max_budget_limiter.is_end_user_within_model_budget( - end_user_id=valid_token.end_user_id, - end_user_model_max_budget=end_user_mmb, - model=current_model, - ) + for model_name in current_models: + await model_max_budget_limiter.is_end_user_within_model_budget( + end_user_id=valid_token.end_user_id, + end_user_model_max_budget=end_user_mmb, + model=model_name, + ) # Check 6: Additional Common Checks across jwt + key auth if valid_token.team_id is not None: @@ -1862,10 +1904,12 @@ async def _run_centralized_common_checks( user_api_key_auth_obj.project_metadata = project_object.metadata user_api_key_auth_obj.project_alias = project_object.project_alias - skip_budget_checks = False - model = get_model_from_request(request_data, route) - if model is not None and llm_router is not None: - skip_budget_checks = _is_model_cost_zero(model=model, llm_router=llm_router) + skip_budget_checks = _should_skip_budget_checks( + request_data=request_data, + route=route, + request=request, + llm_router=llm_router, + ) _ = await common_checks( request=request, @@ -1883,6 +1927,21 @@ async def _run_centralized_common_checks( project_object=project_object, ) + await _reserve_budget_after_common_checks( + user_api_key_auth_obj=user_api_key_auth_obj, + request_data=request_data, + route=route, + llm_router=llm_router, + team_object=team_object, + user_object=user_object, + end_user_id=end_user_id, + end_user_object=end_user_object, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + proxy_logging_obj=proxy_logging_obj, + skip_budget_checks=skip_budget_checks, + ) + async def _noop_none() -> None: """Sentinel coroutine for asyncio.gather when a fetch is unnecessary @@ -1890,6 +1949,59 @@ async def _noop_none() -> None: return None +async def _reserve_budget_after_common_checks( + user_api_key_auth_obj: UserAPIKeyAuth, + request_data: dict, + route: str, + llm_router: Optional[Any], + team_object: Optional[LiteLLM_TeamTableCachedObj], + user_object: Optional[LiteLLM_UserTable], + prisma_client: Optional[PrismaClient], + user_api_key_cache: UserApiKeyCache, + proxy_logging_obj: ProxyLogging, + skip_budget_checks: bool, + end_user_id: Optional[str] = None, + end_user_object: Optional[LiteLLM_EndUserTable] = None, +) -> None: + user_api_key_auth_obj.budget_reservation = None + if skip_budget_checks: + return + + from litellm.proxy.spend_tracking.budget_reservation import ( + reserve_budget_for_request, + ) + + user_api_key_auth_obj.budget_reservation = await reserve_budget_for_request( + request_body=request_data, + route=route, + llm_router=llm_router, + valid_token=user_api_key_auth_obj, + team_object=team_object, + user_object=user_object, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + proxy_logging_obj=proxy_logging_obj, + end_user_id=end_user_id, + end_user_object=end_user_object, + ) + + +def _should_skip_budget_checks( + request_data: dict, + route: str, + request: Optional[Request], + llm_router: Optional[Any], +) -> bool: + model = _get_model_from_request_context( + request_data=request_data, + route=route, + request=request, + ) + if model is not None and llm_router is not None: + return _is_model_cost_zero(model=model, llm_router=llm_router) + return False + + @tracer.wrap() async def user_api_key_auth( request: Request, @@ -1927,6 +2039,7 @@ async def user_api_key_auth( request_data=request_data, custom_litellm_key_header=custom_litellm_key_header, ) + user_api_key_auth_obj.budget_reservation = None ## ENSURE DISABLE ROUTE WORKS ACROSS ALL USER AUTH FLOWS ## RouteChecks.should_call_route(route=route, valid_token=user_api_key_auth_obj) @@ -2134,6 +2247,7 @@ async def _enforce_key_and_fallback_model_access( valid_token: UserAPIKeyAuth, request_data: dict, route: str, + request: Optional[Request], llm_model_list: Optional[list], llm_router: Optional[Any], ) -> None: @@ -2152,10 +2266,10 @@ async def _enforce_key_and_fallback_model_access( ): pass else: - model = get_model_from_request(request_data, route) - fallback_models = cast( - Optional[List[ALL_FALLBACK_MODEL_VALUES]], - request_data.get("fallbacks", None), + model = _get_model_from_request_context( + request_data=request_data, + route=route, + request=request, ) if model is not None: @@ -2166,20 +2280,69 @@ async def _enforce_key_and_fallback_model_access( llm_router=llm_router, ) - if fallback_models is not None: - for m in fallback_models: - await can_key_call_model( - model=m["model"] if isinstance(m, dict) else m, - llm_model_list=llm_model_list, - valid_token=valid_token, - llm_router=llm_router, - ) - await is_valid_fallback_model( - model=m["model"] if isinstance(m, dict) else m, - llm_router=llm_router, - user_model=None, + # Validate every fallback model name reachable by this request. + # All three fields (``fallbacks``, ``context_window_fallbacks``, + # ``content_policy_fallbacks``) are forwarded to the router as + # per-request kwargs whether they appear at the top level of + # ``request_data`` or nested under ``router_settings_override``. + # Both surfaces must be validated against the API key's model + # allowlist or a caller can smuggle a restricted model. VERIA-44. + fallback_names: List[str] = [] + override_settings = request_data.get("router_settings_override") + for _fb_key in ROUTER_FALLBACK_FIELDS: + fallback_names.extend( + iter_router_fallback_model_names(request_data.get(_fb_key)) + ) + if isinstance(override_settings, dict): + fallback_names.extend( + iter_router_fallback_model_names(override_settings.get(_fb_key)) ) + for _name in dict.fromkeys(fallback_names): # dedupe, preserve order + await can_key_call_model( + model=_name, + llm_model_list=llm_model_list, + valid_token=valid_token, + llm_router=llm_router, + ) + await is_valid_fallback_model( + model=_name, + llm_router=llm_router, + user_model=None, + ) + + +ROUTER_FALLBACK_FIELDS: Tuple[str, ...] = ( + "fallbacks", + "context_window_fallbacks", + "content_policy_fallbacks", +) + + +def iter_router_fallback_model_names(fallbacks: Any) -> Iterator[str]: + """Yield leaf model names from any of the supported fallbacks shapes. + + Handles the simple top-level shape (``str`` or ``{"model": str}``) and + the nested router-config shape (``[{primary: [fallback_list]}]``). + """ + if not isinstance(fallbacks, list): + return + for entry in fallbacks: + if isinstance(entry, str): + yield entry + elif isinstance(entry, dict): + if isinstance(entry.get("model"), str): + yield entry["model"] + continue + for fallback_list in entry.values(): + if not isinstance(fallback_list, list): + continue + for m in fallback_list: + if isinstance(m, str): + yield m + elif isinstance(m, dict) and isinstance(m.get("model"), str): + yield m["model"] + async def _run_post_custom_auth_checks( valid_token: UserAPIKeyAuth, @@ -2239,11 +2402,17 @@ async def _run_post_custom_auth_checks( valid_token=valid_token, request_data=request_data, route=route, + request=request, llm_model_list=llm_model_list, llm_router=llm_router, ) - current_model = request_data.get("model", None) + current_model = _get_model_from_request_context( + request_data=request_data, + route=route, + request=request, + ) + current_models = _get_model_names_for_budget_checks(model=current_model) # 3. Check key-level model_max_budget max_budget_per_model = valid_token.model_max_budget @@ -2251,13 +2420,14 @@ async def _run_post_custom_auth_checks( max_budget_per_model is not None and isinstance(max_budget_per_model, dict) and len(max_budget_per_model) > 0 - and current_model is not None + and current_models and valid_token.token is not None ): - await model_max_budget_limiter.is_key_within_model_budget( - user_api_key_dict=valid_token, - model=current_model, - ) + for model_name in current_models: + await model_max_budget_limiter.is_key_within_model_budget( + user_api_key_dict=valid_token, + model=model_name, + ) # 4. Check end-user model_max_budget end_user_mmb = valid_token.end_user_model_max_budget @@ -2265,14 +2435,15 @@ async def _run_post_custom_auth_checks( end_user_mmb is not None and isinstance(end_user_mmb, dict) and len(end_user_mmb) > 0 - and current_model is not None + and current_models and valid_token.end_user_id is not None ): - await model_max_budget_limiter.is_end_user_within_model_budget( - end_user_id=valid_token.end_user_id, - end_user_model_max_budget=end_user_mmb, - model=current_model, - ) + for model_name in current_models: + await model_max_budget_limiter.is_end_user_within_model_budget( + end_user_id=valid_token.end_user_id, + end_user_model_max_budget=end_user_mmb, + model=model_name, + ) # team / user / end_user / project context objects are fetched by # the centralized common_checks gate in user_api_key_auth after diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index f3138f10dac..baa08537003 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -97,6 +97,55 @@ def _serialize_http_exception_detail( return str(detail), None +def _collect_response_file_search_vector_store_ids(data: Dict[str, Any]) -> set[str]: + vector_store_ids: set[str] = set() + tools = data.get("tools") + if not isinstance(tools, list): + return vector_store_ids + + for tool in tools: + if not isinstance(tool, dict) or tool.get("type") != "file_search": + continue + ids = tool.get("vector_store_ids") or [] + if not isinstance(ids, list): + raise HTTPException( + status_code=400, + detail={ + "error": "file_search.vector_store_ids must be a list of strings" + }, + ) + for vector_store_id in ids: + if not isinstance(vector_store_id, str) or not vector_store_id: + raise HTTPException( + status_code=400, + detail={ + "error": "file_search.vector_store_ids must be a list of strings" + }, + ) + vector_store_ids.add(vector_store_id) + + return vector_store_ids + + +async def _authorize_response_file_search_vector_stores( + data: Dict[str, Any], + user_api_key_dict: UserAPIKeyAuth, +) -> None: + vector_store_ids = _collect_response_file_search_vector_store_ids(data) + if not vector_store_ids: + return + + from litellm.proxy.vector_store_endpoints.utils import ( + assert_user_can_access_vector_store_id, + ) + + for vector_store_id in sorted(vector_store_ids): + await assert_user_can_access_vector_store_id( + vector_store_id=vector_store_id, + user_api_key_dict=user_api_key_dict, + ) + + async def _parse_event_data_for_error(event_line: Union[str, bytes]) -> Optional[int]: """Parses an event line and returns an error code if present, else None.""" event_line = ( @@ -791,6 +840,11 @@ class ProxyBaseLLMRequestProcessing: version=version, proxy_config=proxy_config, ) + if route_type in {"aresponses", "_aresponses_websocket"}: + await _authorize_response_file_search_vector_stores( + data=self.data, + user_api_key_dict=user_api_key_dict, + ) # Calculate request queue time after add_litellm_data_to_request # which sets arrival_time in proxy_server_request @@ -1604,6 +1658,12 @@ class ProxyBaseLLMRequestProcessing: # here would duplicate the guardrail API call # (e.g. double OpenAI Moderation charges). continue + if "async_post_call_streaming_iterator_hook" in type(cb).__dict__: + # Skip — the guardrail already scanned the assembled + # response via its own streaming iterator hook in the + # streaming pipeline. re running this function async_post_call_success_hook + # here would duplicate the scan and can spuriously block the guardrail that already passed / failed. + continue else: guardrail_result = await cb.async_post_call_success_hook( user_api_key_dict=captured_user_api_key_dict, diff --git a/litellm/proxy/db/db_spend_update_writer.py b/litellm/proxy/db/db_spend_update_writer.py index c06e1850d9f..3697e498bb4 100644 --- a/litellm/proxy/db/db_spend_update_writer.py +++ b/litellm/proxy/db/db_spend_update_writer.py @@ -60,6 +60,7 @@ from litellm.proxy.db.db_transaction_queue.tool_discovery_queue import ( ToolDiscoveryQueue, ) from litellm.proxy.route_llm_request import ROUTE_ENDPOINT_MAPPING +from litellm.proxy.spend_tracking.spend_log_error_logger import spend_log_error if TYPE_CHECKING: from litellm.proxy.utils import PrismaClient, ProxyLogging @@ -192,17 +193,16 @@ class DBSpendUpdateWriter: verbose_proxy_logger.debug("Runs spend update on all tables") except Exception: - verbose_proxy_logger.error( + spend_log_error( "Spend tracking - update_database failed. Spend log insertion or daily transaction enqueue " "may not have completed for this request. " - "response_cost=%s, token=%s, user_id=%s, team_id=%s, org_id=%s, end_user_id=%s - %s", + "response_cost=%s, token=%s, user_id=%s, team_id=%s, org_id=%s, end_user_id=%s", response_cost, token, user_id, team_id, org_id, end_user_id, - traceback.format_exc(), ) def _enqueue_tool_registry_upsert( @@ -491,9 +491,7 @@ class DBSpendUpdateWriter: ) ) except Exception as e: - verbose_proxy_logger.exception( - f"Update Key DB Call failed to execute - {str(e)}" - ) + spend_log_error("Update Key DB Call failed to execute - %s", str(e), exc=e) raise e async def _update_user_db( @@ -540,14 +538,14 @@ class DBSpendUpdateWriter: ) ) except Exception as e: - verbose_proxy_logger.error( + spend_log_error( "Spend tracking - failed to enqueue user spend update. " - "user_id=%s, end_user_id=%s, response_cost=%s - %s\n%s", + "user_id=%s, end_user_id=%s, response_cost=%s - %s", user_id, end_user_id, response_cost, str(e), - traceback.format_exc(), + exc=e, ) async def _update_team_db( @@ -585,23 +583,23 @@ class DBSpendUpdateWriter: ) ) except Exception as e: - verbose_proxy_logger.error( + spend_log_error( "Spend tracking - failed to enqueue team member spend update. " - "team_id=%s, user_id=%s, response_cost=%s - %s\n%s", + "team_id=%s, user_id=%s, response_cost=%s - %s", team_id, user_id, response_cost, str(e), - traceback.format_exc(), + exc=e, ) except Exception as e: - verbose_proxy_logger.error( + spend_log_error( "Spend tracking - failed to enqueue team spend update. " - "team_id=%s, response_cost=%s - %s\n%s", + "team_id=%s, response_cost=%s - %s", team_id, response_cost, str(e), - traceback.format_exc(), + exc=e, ) raise e @@ -626,13 +624,13 @@ class DBSpendUpdateWriter: ) ) except Exception as e: - verbose_proxy_logger.error( + spend_log_error( "Spend tracking - failed to enqueue org spend update. " - "org_id=%s, response_cost=%s - %s\n%s", + "org_id=%s, response_cost=%s - %s", org_id, response_cost, str(e), - traceback.format_exc(), + exc=e, ) raise e @@ -654,13 +652,13 @@ class DBSpendUpdateWriter: ) ) except Exception as e: - verbose_proxy_logger.error( + spend_log_error( "Spend tracking - failed to enqueue agent spend update. " - "agent_id=%s, response_cost=%s - %s\n%s", + "agent_id=%s, response_cost=%s - %s", agent_id, response_cost, str(e), - traceback.format_exc(), + exc=e, ) raise e @@ -707,13 +705,13 @@ class DBSpendUpdateWriter: ) ) except Exception as e: - verbose_proxy_logger.error( + spend_log_error( "Spend tracking - failed to enqueue tag spend update. " - "request_tags=%s, response_cost=%s - %s\n%s", + "request_tags=%s, response_cost=%s - %s", request_tags, response_cost, str(e), - traceback.format_exc(), + exc=e, ) raise e @@ -906,11 +904,11 @@ class DBSpendUpdateWriter: daily_spend_transactions=daily_agent_spend_update_transactions, ) except Exception as e: - verbose_proxy_logger.error( + spend_log_error( "Spend tracking - failed to commit spend updates from Redis to DB. " - "Data already popped from Redis may be lost. Error: %s\n%s", + "Data already popped from Redis may be lost. Error: %s", str(e), - traceback.format_exc(), + exc=e, ) finally: await self.pod_lock_manager.release_lock( @@ -1074,11 +1072,11 @@ class DBSpendUpdateWriter: daily_spend_transactions=daily_tag_spend_update_transactions, ) except Exception as e: - verbose_proxy_logger.error( + spend_log_error( "Spend tracking - failed to commit daily tag spend updates from Redis to DB. " - "Data already popped from Redis may be lost. Error: %s\n%s", + "Data already popped from Redis may be lost. Error: %s", str(e), - traceback.format_exc(), + exc=e, ) finally: await self.pod_lock_manager.release_lock( @@ -1736,11 +1734,15 @@ class DBSpendUpdateWriter: except Exception as batch_error: # Log detailed error information for debugging batch upsert failures # This helps diagnose issues like unique constraint violations - verbose_proxy_logger.exception( - f"Daily {entity_type} spend batch upsert failed. " - f"Table: {table_name}, Constraint: {unique_constraint_name}, " - f"Batch size: {len(transactions_to_process)}, " - f"Error: {str(batch_error)}" + spend_log_error( + "Daily %s spend batch upsert failed. " + "Table: %s, Constraint: %s, Batch size: %d, Error: %s", + entity_type, + table_name, + unique_constraint_name, + len(transactions_to_process), + str(batch_error), + exc=batch_error, ) raise diff --git a/litellm/proxy/db/spend_counter_reseed.py b/litellm/proxy/db/spend_counter_reseed.py index a979471dc8e..19ec6699390 100644 --- a/litellm/proxy/db/spend_counter_reseed.py +++ b/litellm/proxy/db/spend_counter_reseed.py @@ -14,10 +14,12 @@ memory in long-lived deployments. import asyncio from collections import OrderedDict +from datetime import datetime from typing import TYPE_CHECKING, ClassVar, Optional from litellm._logging import verbose_proxy_logger from litellm.constants import SPEND_COUNTER_RESEED_LOCKS_MAX_SIZE +from litellm.litellm_core_utils.duration_parser import duration_in_seconds if TYPE_CHECKING: from litellm.caching.dual_cache import DualCache @@ -35,6 +37,10 @@ class SpendCounterReseed: spend:team_member:{uid}:{tid} -> LiteLLM_TeamMembership.spend spend:user:{user_id} -> LiteLLM_UserTable.spend spend:org:{org_id} -> LiteLLM_OrganizationTable.spend + + End-user and tag spend counters intentionally do not reseed here. Their + auth paths already load the corresponding objects via get_end_user_object() + and get_tag_objects_batch(); callers pass those values as fallback_spend. """ _locks: ClassVar["OrderedDict[str, asyncio.Lock]"] = OrderedDict() @@ -69,9 +75,10 @@ class SpendCounterReseed: """ if prisma_client is None: return None - # Per-window counters share prefixes with primary counters but - # don't correspond to a DB row. - if ":window:" in counter_key: + # Per-window key/team counters share prefixes with primary counters + # but don't correspond to a DB row. Do not reject arbitrary entity IDs + # or tag names that merely contain ":window:". + if SpendCounterReseed._is_key_or_team_window_counter(counter_key): return None try: if counter_key.startswith("spend:key:"): @@ -97,6 +104,10 @@ class SpendCounterReseed: row = await prisma_client.db.litellm_usertable.find_unique( where={"user_id": user_id} ) + elif counter_key.startswith("spend:end_user:"): + return None + elif counter_key.startswith("spend:tag:"): + return None elif counter_key.startswith("spend:org:"): org_id = counter_key[len("spend:org:") :] row = await prisma_client.db.litellm_organizationtable.find_unique( @@ -113,11 +124,27 @@ class SpendCounterReseed: return None return float(getattr(row, "spend", 0.0) or 0.0) + @staticmethod + def _is_key_or_team_window_counter(counter_key: str) -> bool: + for prefix in ("spend:key:", "spend:team:"): + if not counter_key.startswith(prefix): + continue + _, separator, duration = counter_key.rpartition(":window:") + if not separator or not duration: + return False + try: + duration_in_seconds(duration) + except Exception: + return False + return True + return False + @staticmethod async def coalesced( prisma_client: Optional["PrismaClient"], spend_counter_cache: "DualCache", counter_key: str, + require_cache_warm: bool = False, ) -> Optional[float]: """ Reseed a cold spend counter from the DB and warm the cache, @@ -152,12 +179,156 @@ class SpendCounterReseed: return None # Warm even when 0 so subsequent reads hit cache, not DB. try: - await spend_counter_cache.async_increment_cache( - key=counter_key, value=db_spend, refresh_ttl=True - ) + if spend_counter_cache.redis_cache is not None: + current_value = ( + await spend_counter_cache.redis_cache.async_increment( + key=counter_key, + value=db_spend, + refresh_ttl=True, + ) + ) + spend_counter_cache.in_memory_cache.set_cache( + key=counter_key, + value=current_value, + ) + else: + await spend_counter_cache.async_increment_cache( + key=counter_key, value=db_spend, refresh_ttl=True + ) except Exception: verbose_proxy_logger.exception( "SpendCounterReseed.coalesced: failed to warm counter %s", counter_key, ) + if require_cache_warm: + raise return db_spend + + @staticmethod + async def window_from_spend_logs( + prisma_client: Optional["PrismaClient"], + entity_type: str, + entity_id: str, + window_start: datetime, + ) -> Optional[float]: + if prisma_client is None: + return None + + if entity_type == "Key": + group_field = "api_key" + where = { + "api_key": entity_id, + "startTime": {"gte": window_start}, + } + elif entity_type == "Team": + group_field = "team_id" + where = { + "team_id": entity_id, + "startTime": {"gte": window_start}, + } + else: + return None + + try: + response = await prisma_client.db.litellm_spendlogs.group_by( + by=[group_field], + where=where, # type: ignore[arg-type] + sum={"spend": True}, + ) + except Exception: + verbose_proxy_logger.exception( + "SpendCounterReseed.window_from_spend_logs: failed for %s=%s", + entity_type, + entity_id, + ) + return None + + if not response: + return 0.0 + first_row = response[0] + sum_row = ( + first_row.get("_sum") + if isinstance(first_row, dict) + else getattr(first_row, "_sum", None) + ) + spend = ( + sum_row.get("spend") + if isinstance(sum_row, dict) + else getattr(sum_row, "spend", None) + ) + return float(spend or 0.0) + + @staticmethod + async def coalesced_window( + prisma_client: Optional["PrismaClient"], + spend_counter_cache: "DualCache", + counter_key: str, + entity_type: str, + entity_id: str, + window_start: datetime, + ) -> Optional[float]: + lock = await SpendCounterReseed._get_lock(counter_key) + async with lock: + redis_clean_miss = False + if spend_counter_cache.redis_cache is not None: + try: + val = await spend_counter_cache.redis_cache.async_get_cache( + key=counter_key + ) + if val is not None: + return float(val) + redis_clean_miss = True + except Exception: + pass + if not redis_clean_miss: + val = spend_counter_cache.in_memory_cache.get_cache(key=counter_key) + if val is not None: + return float(val) + + window_spend = await SpendCounterReseed.window_from_spend_logs( + prisma_client=prisma_client, + entity_type=entity_type, + entity_id=entity_id, + window_start=window_start, + ) + if window_spend is None: + return None + try: + if spend_counter_cache.redis_cache is not None: + seeded = await spend_counter_cache.redis_cache.async_set_cache( + key=counter_key, + value=window_spend, + nx=True, + ) + if seeded: + current_value = window_spend + else: + current_cached_value = ( + await spend_counter_cache.redis_cache.async_get_cache( + key=counter_key + ) + ) + if current_cached_value is None: + current_value = ( + await spend_counter_cache.redis_cache.async_increment( + key=counter_key, + value=window_spend, + ) + ) + else: + current_value = float(current_cached_value) + spend_counter_cache.in_memory_cache.set_cache( + key=counter_key, + value=current_value, + ) + else: + await spend_counter_cache.async_increment_cache( + key=counter_key, value=window_spend + ) + except Exception: + verbose_proxy_logger.exception( + "SpendCounterReseed.coalesced_window: failed to warm counter %s", + counter_key, + ) + raise + return window_spend diff --git a/litellm/proxy/guardrails/guardrail_endpoints.py b/litellm/proxy/guardrails/guardrail_endpoints.py index ac487fb06d0..5351391e5e1 100644 --- a/litellm/proxy/guardrails/guardrail_endpoints.py +++ b/litellm/proxy/guardrails/guardrail_endpoints.py @@ -21,6 +21,7 @@ from litellm.integrations.custom_guardrail import CustomGuardrail from litellm.litellm_core_utils.safe_json_dumps import safe_dumps from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth from litellm.proxy.auth.user_api_key_auth import user_api_key_auth +from litellm.proxy.management_endpoints.common_utils import _user_has_admin_view from litellm.proxy.guardrails.guardrail_hooks.custom_code.sandbox import ( build_sandbox_globals, compile_sandboxed, @@ -842,7 +843,10 @@ async def list_guardrail_submissions( if prisma_client is None: raise HTTPException(status_code=500, detail="Prisma client not initialized") - is_admin = user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN + # Admin Viewer follows the read-parity rule: see all submissions like a + # Proxy Admin would (no writes — registration / approval still gated + # elsewhere by their own per-action checks). + is_admin = _user_has_admin_view(user_api_key_dict) visible_team_ids: Optional[List[str]] = None if not is_admin: visible_team_ids = await _get_user_team_ids(user_api_key_dict) diff --git a/litellm/proxy/guardrails/guardrail_hooks/presidio.py b/litellm/proxy/guardrails/guardrail_hooks/presidio.py index 03a754fa5df..fc414ab7b54 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/presidio.py +++ b/litellm/proxy/guardrails/guardrail_hooks/presidio.py @@ -1160,14 +1160,38 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): from litellm.types.utils import ModelResponse all_chunks: List[ModelResponseStream] = [] + passthrough_due_to_unknown_stream_shape = False try: async for chunk in response: if isinstance(chunk, ModelResponseStream): - all_chunks.append(chunk) + if passthrough_due_to_unknown_stream_shape: + yield chunk + else: + all_chunks.append(chunk) elif isinstance(chunk, bytes): yield chunk # type: ignore[misc] continue - + else: + if all_chunks: + # Flush buffered chunks and switch to transparent passthrough for this stream shape. + # NOTE: these buffered chunks are emitted unmasked because this + # stream mixed chunk types and cannot be safely reconstructed. + verbose_proxy_logger.warning( + "Presidio apply_to_output: mixed stream detected (ModelResponseStream + unknown event). " + "Flushing %d buffered chunks without PII masking and switching to transparent passthrough.", + len(all_chunks), + ) + for buffered_chunk in all_chunks: + yield buffered_chunk + all_chunks = [] + passthrough_due_to_unknown_stream_shape = True + yield chunk + if passthrough_due_to_unknown_stream_shape: + verbose_proxy_logger.warning( + "Presidio apply_to_output: streaming response contained unknown event objects " + "(e.g. /v1/responses events). Output PII masking was skipped for this response." + ) + return if not all_chunks: verbose_proxy_logger.warning( "Presidio apply_to_output: streaming response contained only " diff --git a/litellm/proxy/guardrails/guardrail_hooks/qohash/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/qohash/__init__.py new file mode 100644 index 00000000000..465f52db3d0 --- /dev/null +++ b/litellm/proxy/guardrails/guardrail_hooks/qohash/__init__.py @@ -0,0 +1,35 @@ +from typing import TYPE_CHECKING + +from litellm.types.guardrails import SupportedGuardrailIntegrations + +from .qohash import QostodianNexus + +if TYPE_CHECKING: + from litellm.types.guardrails import Guardrail, LitellmParams + + +def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"): + import litellm + + _instance = QostodianNexus( + api_base=litellm_params.api_base, + guardrail_name=guardrail.get("guardrail_name", ""), + event_hook=litellm_params.mode, + default_on=litellm_params.default_on, + additional_provider_specific_params=litellm_params.additional_provider_specific_params, + extra_headers=getattr(litellm_params, "extra_headers", None), + ) + + litellm.logging_callback_manager.add_litellm_callback(_instance) + + return _instance + + +guardrail_initializer_registry = { + SupportedGuardrailIntegrations.QOSTODIAN_NEXUS.value: initialize_guardrail, +} + + +guardrail_class_registry = { + SupportedGuardrailIntegrations.QOSTODIAN_NEXUS.value: QostodianNexus, +} diff --git a/litellm/proxy/guardrails/guardrail_hooks/qohash/qohash.py b/litellm/proxy/guardrails/guardrail_hooks/qohash/qohash.py new file mode 100644 index 00000000000..a1bab6dbac9 --- /dev/null +++ b/litellm/proxy/guardrails/guardrail_hooks/qohash/qohash.py @@ -0,0 +1,81 @@ +""" +Qostodian Nexus (by Qohash) — LiteLLM guardrail integration. +""" + +import os +from typing import TYPE_CHECKING, Literal, Optional, Type + +from litellm.integrations.custom_guardrail import log_guardrail_information +from litellm.proxy.guardrails.guardrail_hooks.generic_guardrail_api.generic_guardrail_api import ( + GenericGuardrailAPI, +) +from litellm.types.proxy.guardrails.guardrail_hooks.qohash import ( + QostodianNexusConfigModel, +) +from litellm.types.utils import GenericGuardrailAPIInputs + +if TYPE_CHECKING: + from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj + +GUARDRAIL_NAME = "qostodian_nexus" + + +class QostodianNexus(GenericGuardrailAPI): + def __init__( + self, + api_base: Optional[str] = None, + **kwargs, + ): + api_base = api_base or os.environ.get( + "QOSTODIAN_NEXUS_API_BASE", "http://nexus:8800" + ) + + kwargs["guardrail_name"] = kwargs.get("guardrail_name", GUARDRAIL_NAME) + + # Merge built-in Qostodian Nexus identifier headers with any caller-supplied extra_headers + nexus_headers = [ + "x-qostodian-nexus-identifiers-trace", + "x-qostodian-nexus-identifiers-source", + "x-qostodian-nexus-identifiers-container", + "x-qostodian-nexus-identifiers-identity", + ] + + existing = kwargs.get("extra_headers") or [] + kwargs["extra_headers"] = nexus_headers + [ + h for h in existing if h not in nexus_headers + ] + + super().__init__( + api_base=api_base, + **kwargs, + ) + + @log_guardrail_information + async def apply_guardrail( + self, + inputs: GenericGuardrailAPIInputs, + request_data: dict, + input_type: Literal["request", "response"], + logging_obj: Optional["LiteLLMLoggingObj"] = None, + ) -> GenericGuardrailAPIInputs: + """ + Apply Qostodian Nexus to the given inputs. + + NOTE: This override is intentionally a pass-through. It must be present + directly in this class's __dict__ so that LiteLLM's unified guardrail + routing check (`"apply_guardrail" in type(callback).__dict__` in + litellm/proxy/utils.py) routes calls correctly. Do not remove. + """ + return await super().apply_guardrail( + inputs=inputs, + request_data=request_data, + input_type=input_type, + logging_obj=logging_obj, + ) + + @classmethod + def get_config_model(cls) -> Optional[Type[QostodianNexusConfigModel]]: + """ + Returns the config model for Qostodian Nexus. + """ + return QostodianNexusConfigModel diff --git a/litellm/proxy/hooks/batch_rate_limiter.py b/litellm/proxy/hooks/batch_rate_limiter.py index 2ee0588f19e..f740d5dd40c 100644 --- a/litellm/proxy/hooks/batch_rate_limiter.py +++ b/litellm/proxy/hooks/batch_rate_limiter.py @@ -27,6 +27,7 @@ from litellm._logging import verbose_proxy_logger from litellm.batches.batch_utils import ( _get_batch_job_input_file_usage, _get_file_content_as_dictionary, + _get_models_from_batch_input_file_content, ) from litellm.integrations.custom_logger import CustomLogger from litellm.proxy._types import UserAPIKeyAuth @@ -246,6 +247,17 @@ class _PROXY_BatchRateLimiter(CustomLogger): file_content_as_dict = _get_file_content_as_dictionary(file_content.content) + # Validate every model named in the batch JSONL against the + # caller's per-key model allowlist. Without this, a caller + # could smuggle restricted/expensive models inside the file + # and the upstream provider would execute the batch under + # the proxy's shared API key. + if user_api_key_dict is not None: + await self._enforce_batch_file_model_access( + user_api_key_dict=user_api_key_dict, + file_content_as_dict=file_content_as_dict, + ) + input_file_usage = _get_batch_job_input_file_usage( file_content_dictionary=file_content_as_dict, custom_llm_provider=custom_llm_provider, @@ -256,12 +268,69 @@ class _PROXY_BatchRateLimiter(CustomLogger): request_count=request_count, ) + except HTTPException as e: + # Distinguish intentional 403s from `_enforce_batch_file_model_access` + # from genuine I/O failures so security-relevant rejections show up + # in the access log instead of getting buried in error noise. + if e.status_code == 403: + verbose_proxy_logger.warning( + f"Batch rejected: caller not authorized for a model named in {file_id}: {e.detail}" + ) + else: + verbose_proxy_logger.error( + f"Batch input file rejected for {file_id}: status={e.status_code} detail={e.detail}" + ) + raise except Exception as e: verbose_proxy_logger.error( f"Error counting input file usage for {file_id}: {str(e)}" ) raise + async def _enforce_batch_file_model_access( + self, + user_api_key_dict: UserAPIKeyAuth, + file_content_as_dict: List[dict], + ) -> None: + """Reject the batch if the caller is not authorized for every + ``body.model`` named inside the JSONL. + + Reuses ``can_key_call_model`` so the same allowlist semantics + (wildcards, access groups, ``all-proxy-models``, team aliases) + the proxy enforces on `/chat/completions` apply here. + """ + from litellm.proxy.auth.auth_checks import can_key_call_model + from litellm.proxy.proxy_server import llm_router + + models = _get_models_from_batch_input_file_content(file_content_as_dict) + if not models: + return + + llm_model_list = llm_router.model_list if llm_router is not None else None + for model in models: + try: + await can_key_call_model( + model=model, + llm_model_list=llm_model_list, + valid_token=user_api_key_dict, + llm_router=llm_router, + ) + except HTTPException: + raise + except Exception as e: + # `can_key_call_model` raises ProxyException on denial; + # re-shape to a 403 so the batch endpoint returns a + # consistent rejection without leaking internal types. + raise HTTPException( + status_code=403, + detail={ + "error": ( + "Batch input file references a model the caller is " + f"not authorized to use: model={model}, reason={str(e)}" + ) + }, + ) + async def _fetch_managed_file_content( self, file_id: str, diff --git a/litellm/proxy/hooks/max_budget_limiter.py b/litellm/proxy/hooks/max_budget_limiter.py index 7789fa6a349..9a7e5117945 100644 --- a/litellm/proxy/hooks/max_budget_limiter.py +++ b/litellm/proxy/hooks/max_budget_limiter.py @@ -32,10 +32,25 @@ class _PROXY_MaxBudgetLimiter(CustomLogger): if user_api_key_dict.team_id is not None: return + # The reservation path admits at the strict-`<` boundary and + # atomically pre-fills the same counter we'd read here. Re-checking + # with `>=` would reject a request the reservation already admitted + # when the reservation fills the counter to exactly max_budget. + # Imported lazily to avoid a circular import via proxy.utils. + from litellm.proxy.spend_tracking.budget_reservation import ( + get_reserved_counter_keys, + ) + + user_counter_key = f"spend:user:{user_id}" + if user_counter_key in get_reserved_counter_keys( + user_api_key_dict.budget_reservation + ): + return + from litellm.proxy.proxy_server import get_current_spend curr_spend = await get_current_spend( - counter_key=f"spend:user:{user_id}", + counter_key=user_counter_key, fallback_spend=user_api_key_dict.user_spend or 0.0, ) diff --git a/litellm/proxy/hooks/proxy_track_cost_callback.py b/litellm/proxy/hooks/proxy_track_cost_callback.py index c9946f4e26f..67f702e31a2 100644 --- a/litellm/proxy/hooks/proxy_track_cost_callback.py +++ b/litellm/proxy/hooks/proxy_track_cost_callback.py @@ -19,6 +19,10 @@ from litellm.proxy.auth.auth_checks import ( ) from litellm.proxy.auth.route_checks import RouteChecks from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup +from litellm.proxy.spend_tracking.spend_log_error_logger import ( + should_suppress_spend_log_tracebacks, + spend_log_error, +) from litellm.proxy.utils import ProxyUpdateSpend from litellm.types.utils import StandardLoggingPayload from litellm.utils import get_end_user_id_for_cost_tracking @@ -30,16 +34,35 @@ class _ProxyDBLogger(CustomLogger): kwargs, response_obj, start_time, end_time ) - async def async_post_call_failure_hook( - self, - request_data: dict, - original_exception: Exception, - user_api_key_dict: UserAPIKeyAuth, - traceback_str: Optional[str] = None, - ): - request_route = user_api_key_dict.request_route - if _ProxyDBLogger._should_track_errors_in_db() is False: - return + async def async_post_call_failure_hook( + self, + request_data: dict, + original_exception: Exception, + user_api_key_dict: UserAPIKeyAuth, + traceback_str: Optional[str] = None, + ): + try: + await _release_budget_reservation( + budget_reservation=user_api_key_dict.budget_reservation + ) + except Exception: + verbose_proxy_logger.exception( + "Failed to release budget reservation during failure handling" + ) + try: + await _invalidate_budget_reservation_counters( + budget_reservation=user_api_key_dict.budget_reservation + ) + if user_api_key_dict.budget_reservation is not None: + user_api_key_dict.budget_reservation["finalized"] = True + except Exception: + verbose_proxy_logger.exception( + "Failed to invalidate budget reservation counters after failure release failed" + ) + + request_route = user_api_key_dict.request_route + if _ProxyDBLogger._should_track_errors_in_db() is False: + return elif request_route is not None and not ( RouteChecks.is_llm_api_route(route=request_route) or RouteChecks.is_info_route(route=request_route) @@ -55,12 +78,18 @@ class _ProxyDBLogger(CustomLogger): ) _metadata["user_api_key"] = user_api_key_dict.api_key _metadata["status"] = "failure" - _metadata["error_information"] = ( - StandardLoggingPayloadSetup.get_error_information( - original_exception=original_exception, - traceback_str=traceback_str, - ) + _error_information = StandardLoggingPayloadSetup.get_error_information( + original_exception=original_exception, + traceback_str=traceback_str, ) + if should_suppress_spend_log_tracebacks(): + # Drop the traceback key entirely so the per-row Metadata pane in + # the UI (which renders the JSON blob verbatim) doesn't show a + # noisy ``"traceback": ""`` line. Downstream consumers all use + # ``.get("traceback")`` / truthy checks, and the TypedDict marks + # the field as optional, so omitting is type-safe. + _error_information.pop("traceback", None) + _metadata["error_information"] = _error_information _metadata = await _ProxyDBLogger._enrich_failure_metadata_with_key_info( metadata=_metadata, @@ -155,66 +184,64 @@ class _ProxyDBLogger(CustomLogger): f"kwargs stream: {kwargs.get('stream', None)} + complete streaming response: {kwargs.get('complete_streaming_response', None)}" ) parent_otel_span = _get_parent_otel_span_from_kwargs(kwargs=kwargs) - litellm_params = kwargs.get("litellm_params", {}) or {} - end_user_id = get_end_user_id_for_cost_tracking(litellm_params) - metadata = get_litellm_metadata_from_kwargs(kwargs=kwargs) - user_id = cast(Optional[str], metadata.get("user_api_key_user_id", None)) - team_id = cast(Optional[str], metadata.get("user_api_key_team_id", None)) - org_id = cast(Optional[str], metadata.get("user_api_key_org_id", None)) + litellm_params = kwargs.get("litellm_params", {}) or {} + end_user_id = get_end_user_id_for_cost_tracking(litellm_params) + metadata = get_litellm_metadata_from_kwargs(kwargs=kwargs) + budget_reservation = _get_budget_reservation_from_metadata( + metadata=metadata + ) + user_id = cast(Optional[str], metadata.get("user_api_key_user_id", None)) + team_id = cast(Optional[str], metadata.get("user_api_key_team_id", None)) + org_id = cast(Optional[str], metadata.get("user_api_key_org_id", None)) key_alias = cast(Optional[str], metadata.get("user_api_key_alias", None)) end_user_max_budget = metadata.get("user_api_end_user_max_budget", None) sl_object: Optional[StandardLoggingPayload] = kwargs.get( "standard_logging_object", None ) - response_cost = ( - sl_object.get("response_cost", None) - if sl_object is not None - else kwargs.get("response_cost", None) - ) - tags: Optional[List[str]] = ( - sl_object.get("request_tags", None) if sl_object is not None else None - ) - - if response_cost is not None: - user_api_key = metadata.get("user_api_key", None) + response_cost = ( + sl_object.get("response_cost", None) + if sl_object is not None + else kwargs.get("response_cost", None) + ) + tags = _get_request_tags_for_cost_tracking( + sl_object=sl_object, + metadata=metadata, + ) + + if response_cost is not None: + user_api_key = metadata.get("user_api_key", None) if kwargs.get("cache_hit", False) is True: response_cost = 0.0 verbose_proxy_logger.debug( f"Cache Hit: response_cost {response_cost}, for user_id {user_id}" ) - verbose_proxy_logger.debug( - f"user_api_key {user_api_key}, user_id {user_id}, team_id {team_id}, end_user_id {end_user_id}" - ) - if _should_track_cost_callback( - user_api_key=user_api_key, + verbose_proxy_logger.debug( + f"user_api_key {user_api_key}, user_id {user_id}, team_id {team_id}, end_user_id {end_user_id}" + ) + if _should_track_cost_callback( + user_api_key=user_api_key, user_id=user_id, team_id=team_id, - end_user_id=end_user_id, - ): - ## UPDATE DATABASE - await proxy_logging_obj.db_spend_update_writer.update_database( - token=user_api_key, - response_cost=response_cost, - user_id=user_id, - end_user_id=end_user_id, - team_id=team_id, - kwargs=kwargs, - completion_response=completion_response, - start_time=start_time, - end_time=end_time, - org_id=org_id, - ) - - # Atomically update spend counters (in-memory + Redis) - # for cross-pod budget enforcement. - await increment_spend_counters( - token=user_api_key, - team_id=team_id, - user_id=user_id, - response_cost=response_cost, - org_id=org_id, - ) + end_user_id=end_user_id, + ): + ## UPDATE DATABASE + await _update_database_and_spend_counters( + proxy_logging_obj=proxy_logging_obj, + increment_spend_counters=increment_spend_counters, + user_api_key=user_api_key, + user_id=user_id, + end_user_id=end_user_id, + team_id=team_id, + org_id=org_id, + kwargs=kwargs, + completion_response=completion_response, + start_time=start_time, + end_time=end_time, + response_cost=response_cost, + budget_reservation=budget_reservation, + request_tags=tags, + ) # update cache (fire-and-forget for backward compat: # cached object fields, soft budget alerts, etc.) @@ -234,10 +261,15 @@ class _ProxyDBLogger(CustomLogger): token=user_api_key, key_alias=key_alias, end_user_id=end_user_id, - response_cost=response_cost, - max_budget=end_user_max_budget, - ) + response_cost=response_cost, + max_budget=end_user_max_budget, + ) + elif budget_reservation is not None: + await _release_budget_reservation( + budget_reservation=budget_reservation + ) else: + await _release_budget_reservation(budget_reservation=budget_reservation) # Non-model call types (health checks, afile_delete) have no model or standard_logging_object. # Use .get() for "stream" to avoid KeyError on health checks. if sl_object is None and not kwargs.get("model"): @@ -280,9 +312,7 @@ class _ProxyDBLogger(CustomLogger): ) ) - verbose_proxy_logger.exception( - "Error in tracking cost callback - %s", str(e) - ) + spend_log_error("Error in tracking cost callback - %s", str(e), exc=e) @staticmethod async def _enrich_failure_metadata_with_key_info(metadata: dict) -> dict: @@ -366,7 +396,7 @@ class _ProxyDBLogger(CustomLogger): return -def _should_track_cost_callback( +def _should_track_cost_callback( user_api_key: Optional[str], user_id: Optional[str], team_id: Optional[str], @@ -387,4 +417,135 @@ def _should_track_cost_callback( or end_user_id is not None ): return True - return False + return False + + +def _get_budget_reservation_from_metadata(metadata: dict) -> Optional[dict]: + metadata_budget_reservation = metadata.get("user_api_key_budget_reservation") + if isinstance(metadata_budget_reservation, dict): + return metadata_budget_reservation + + user_api_key_auth_obj = metadata.get("user_api_key_auth") + if user_api_key_auth_obj is None: + return None + if isinstance(user_api_key_auth_obj, dict): + budget_reservation = user_api_key_auth_obj.get("budget_reservation") + return budget_reservation if isinstance(budget_reservation, dict) else None + return getattr(user_api_key_auth_obj, "budget_reservation", None) + + +def _get_request_tags_for_cost_tracking( + sl_object: Optional[StandardLoggingPayload], + metadata: dict, +) -> Optional[List[str]]: + if sl_object is not None: + request_tags = sl_object.get("request_tags", None) + if isinstance(request_tags, list): + return request_tags + + metadata_tags = metadata.get("tags", None) + if isinstance(metadata_tags, list): + return metadata_tags + + return None + + +async def _update_database_and_spend_counters( + proxy_logging_obj: Any, + increment_spend_counters: Any, + user_api_key: Optional[str], + user_id: Optional[str], + end_user_id: Optional[str], + team_id: Optional[str], + org_id: Optional[str], + kwargs: dict, + completion_response: Optional[Union[litellm.ModelResponse, Any]], + start_time: Any, + end_time: Any, + response_cost: float, + budget_reservation: Optional[dict], + request_tags: Optional[List[str]] = None, +) -> None: + try: + await proxy_logging_obj.db_spend_update_writer.update_database( + token=user_api_key, + response_cost=response_cost, + user_id=user_id, + end_user_id=end_user_id, + team_id=team_id, + kwargs=kwargs, + completion_response=completion_response, + start_time=start_time, + end_time=end_time, + org_id=org_id, + ) + except Exception: + if budget_reservation is not None: + try: + await _release_budget_reservation(budget_reservation=budget_reservation) + except Exception: + verbose_proxy_logger.exception( + "Failed to release budget reservation after database update failed" + ) + try: + await _invalidate_budget_reservation_counters( + budget_reservation=budget_reservation + ) + except Exception: + verbose_proxy_logger.exception( + "Failed to invalidate budget reservation counters after release failed" + ) + raise + + try: + await increment_spend_counters( + token=user_api_key, + team_id=team_id, + user_id=user_id, + response_cost=response_cost, + org_id=org_id, + budget_reservation=budget_reservation, + end_user_id=end_user_id, + tags=request_tags, + ) + except Exception: + if budget_reservation is not None: + try: + await _invalidate_budget_reservation_counters( + budget_reservation=budget_reservation + ) + except Exception: + verbose_proxy_logger.exception( + "Failed to invalidate budget reservation counters after spend counter update failed" + ) + finally: + budget_reservation["finalized"] = True + raise + + +async def _release_budget_reservation(budget_reservation: Optional[dict]) -> None: + if budget_reservation is None: + return + + from litellm.proxy.spend_tracking.budget_reservation import ( + release_budget_reservation, + ) + + await release_budget_reservation( + budget_reservation=budget_reservation, + ) + + +async def _invalidate_budget_reservation_counters( + budget_reservation: Optional[dict], +) -> None: + if budget_reservation is None: + return + + from litellm.proxy.spend_tracking.budget_reservation import ( + invalidate_budget_reservation_counters, + ) + + await invalidate_budget_reservation_counters( + budget_reservation=budget_reservation, + ) diff --git a/litellm/proxy/litellm_pre_call_utils.py b/litellm/proxy/litellm_pre_call_utils.py index 3077efe1167..853c56856fc 100644 --- a/litellm/proxy/litellm_pre_call_utils.py +++ b/litellm/proxy/litellm_pre_call_utils.py @@ -893,6 +893,10 @@ class LiteLLMProxyRequestSetup: data[_metadata_variable_name]["user_api_end_user_max_budget"] = getattr( user_api_key_dict, "end_user_max_budget", None ) + if user_api_key_dict.budget_reservation is not None: + data[_metadata_variable_name][ + "user_api_key_budget_reservation" + ] = user_api_key_dict.budget_reservation # Add the full UserAPIKeyAuth object for MCP server access control data[_metadata_variable_name]["user_api_key_auth"] = user_api_key_dict return data diff --git a/litellm/proxy/management_endpoints/access_group_endpoints.py b/litellm/proxy/management_endpoints/access_group_endpoints.py index ceaef20a8d0..62a770f46ae 100644 --- a/litellm/proxy/management_endpoints/access_group_endpoints.py +++ b/litellm/proxy/management_endpoints/access_group_endpoints.py @@ -38,6 +38,17 @@ def _require_proxy_admin(user_api_key_dict: UserAPIKeyAuth) -> None: ) +def _require_admin_view(user_api_key_dict: UserAPIKeyAuth) -> None: + """Admin Viewer parity: PROXY_ADMIN or PROXY_ADMIN_VIEW_ONLY may read.""" + from litellm.proxy.management_endpoints.common_utils import _user_has_admin_view + + if not _user_has_admin_view(user_api_key_dict): + raise HTTPException( + status_code=status.HTTP_403_FORBIDDEN, + detail={"error": CommonProxyErrors.not_allowed_access.value}, + ) + + def _record_to_response(record) -> AccessGroupResponse: return AccessGroupResponse( access_group_id=record.access_group_id, @@ -370,7 +381,7 @@ async def create_access_group( async def list_access_groups( user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), ) -> List[AccessGroupResponse]: - _require_proxy_admin(user_api_key_dict) + _require_admin_view(user_api_key_dict) prisma_client = get_prisma_client_or_throw( CommonProxyErrors.db_not_connected_error.value ) @@ -389,7 +400,7 @@ async def get_access_group( access_group_id: str, user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), ) -> AccessGroupResponse: - _require_proxy_admin(user_api_key_dict) + _require_admin_view(user_api_key_dict) prisma_client = get_prisma_client_or_throw( CommonProxyErrors.db_not_connected_error.value ) diff --git a/litellm/proxy/management_endpoints/budget_management_endpoints.py b/litellm/proxy/management_endpoints/budget_management_endpoints.py index 90c0d02d1e0..81b133e6c81 100644 --- a/litellm/proxy/management_endpoints/budget_management_endpoints.py +++ b/litellm/proxy/management_endpoints/budget_management_endpoints.py @@ -17,6 +17,7 @@ from fastapi import APIRouter, Depends, HTTPException from litellm.proxy.common_utils.timezone_utils import get_budget_reset_time from litellm.proxy._types import * from litellm.proxy.auth.user_api_key_auth import user_api_key_auth +from litellm.proxy.management_endpoints.common_utils import _user_has_admin_view from litellm.proxy.utils import jsonify_object router = APIRouter() @@ -238,7 +239,7 @@ async def budget_settings( detail={"error": CommonProxyErrors.db_not_connected_error.value}, ) - if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN: + if not _user_has_admin_view(user_api_key_dict): raise HTTPException( status_code=400, detail={ @@ -305,7 +306,7 @@ async def list_budget( detail={"error": CommonProxyErrors.db_not_connected_error.value}, ) - if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN: + if not _user_has_admin_view(user_api_key_dict): raise HTTPException( status_code=400, detail={ diff --git a/litellm/proxy/management_endpoints/common_utils.py b/litellm/proxy/management_endpoints/common_utils.py index b0ea6b41ac5..8ad44b53007 100644 --- a/litellm/proxy/management_endpoints/common_utils.py +++ b/litellm/proxy/management_endpoints/common_utils.py @@ -1,5 +1,7 @@ from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union +from fastapi import HTTPException, status + from litellm._logging import verbose_proxy_logger from litellm.caching import DualCache from litellm.proxy._types import ( @@ -29,6 +31,34 @@ def _user_has_admin_view(user_api_key_dict: UserAPIKeyAuth) -> bool: ) +def require_caller_user_id_for_non_admin( + user_api_key_dict: UserAPIKeyAuth, +) -> str: + """Return the caller's user_id, or raise 403 if missing. + + Non-admin analytics endpoints scope queries by the caller's own user_id. + Service-account keys are deliberately created with user_id=None + (key_management_endpoints.py forces ``data.user_id = None`` at key + creation). Without this guard, that None value flows through to the + daily-activity builder, which treats ``entity_id is None`` as "no filter" + and returns every tenant's data. + + Callers must check is_admin first; this helper is only valid on the + non-admin scoping branch. + """ + if user_api_key_dict.user_id is None: + raise HTTPException( + status_code=status.HTTP_403_FORBIDDEN, + detail={ + "error": ( + "Service-account keys cannot query user analytics. " + "Use a user-bound key, or call as a proxy admin." + ) + }, + ) + return user_api_key_dict.user_id + + def _is_user_team_admin( user_api_key_dict: UserAPIKeyAuth, team_obj: LiteLLM_TeamTable ) -> bool: diff --git a/litellm/proxy/management_endpoints/config_override_endpoints.py b/litellm/proxy/management_endpoints/config_override_endpoints.py index d78c5526e66..b736ba1081e 100644 --- a/litellm/proxy/management_endpoints/config_override_endpoints.py +++ b/litellm/proxy/management_endpoints/config_override_endpoints.py @@ -267,9 +267,11 @@ async def get_hashicorp_vault_config( Get current Hashicorp Vault configuration. Returns decrypted values from DB, or falls back to current env vars. """ + from litellm.proxy.management_endpoints.common_utils import _user_has_admin_view from litellm.proxy.proxy_server import prisma_client, proxy_config - if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN: + # Admin Viewer follows the read-parity rule. + if not _user_has_admin_view(user_api_key_dict): raise HTTPException( status_code=403, detail="Only admin users can view config overrides", diff --git a/litellm/proxy/management_endpoints/internal_user_endpoints.py b/litellm/proxy/management_endpoints/internal_user_endpoints.py index 921d24da043..6f73c6a632d 100644 --- a/litellm/proxy/management_endpoints/internal_user_endpoints.py +++ b/litellm/proxy/management_endpoints/internal_user_endpoints.py @@ -35,6 +35,7 @@ from litellm.proxy.management_endpoints.common_daily_activity import ( from litellm.proxy.management_endpoints.common_utils import ( _is_user_team_admin, _user_has_admin_view, + require_caller_user_id_for_non_admin, ) from litellm.proxy.management_endpoints.key_management_endpoints import ( generate_key_helper_fn, @@ -618,6 +619,40 @@ def _normalize_user_info_user_id( return user_id +def _enforce_user_info_access( + user_id: Optional[str], user_api_key_dict: UserAPIKeyAuth +) -> None: + """Re-validate that the caller may read the resolved ``user_id`` after + URL-decoding has been finalized. + + The route-level check in ``RouteChecks.non_proxy_admin_allowed_routes_check`` + runs against ``request.query_params``, which decodes a literal ``+`` to a + space. ``_normalize_user_info_user_id`` then re-parses the raw query with + ``unquote`` so the endpoint can return rows for user_ids that contain ``+`` + (e.g. plus-addressed emails). That asymmetry let an attacker who registered + a username with a literal space pass the route check and then read another + user's row by sending the encoded ``+`` form. Re-checking ownership here + closes the gap without changing the supported user_id grammar. + """ + if user_id is None: + return + # Only true proxy admin bypasses ownership. PROXY_ADMIN_VIEW_ONLY is + # subject to the same `user_id == valid_token.user_id` rule that + # `RouteChecks.non_proxy_admin_allowed_routes_check` applies upstream + # for the `/user/info` route. + if user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN: + return + if user_id == user_api_key_dict.user_id: + return + raise HTTPException( + status_code=status.HTTP_403_FORBIDDEN, + detail=( + f"key not allowed to access this user's info. user_id={user_id}, " + f"key's user_id={user_api_key_dict.user_id}" + ), + ) + + async def _get_user_info_teams( prisma_client: Any, user_id: Optional[str], @@ -732,6 +767,7 @@ async def user_info( # noqa: PLR0915 try: user_id = _normalize_user_info_user_id(request=request, user_id=user_id) + _enforce_user_info_access(user_id=user_id, user_api_key_dict=user_api_key_dict) if prisma_client is None: raise Exception( @@ -2587,9 +2623,10 @@ async def get_user_daily_activity( if is_admin: entity_id = user_id # None means global view, otherwise filter by user else: + caller_user_id = require_caller_user_id_for_non_admin(user_api_key_dict) if user_id is None: - user_id = user_api_key_dict.user_id - if user_id != user_api_key_dict.user_id: + user_id = caller_user_id + if user_id != caller_user_id: raise HTTPException( status_code=status.HTTP_403_FORBIDDEN, detail={ @@ -2684,9 +2721,10 @@ async def get_user_daily_activity_aggregated( if is_admin: entity_id = user_id # None means global view, otherwise filter by user else: + caller_user_id = require_caller_user_id_for_non_admin(user_api_key_dict) if user_id is None: - user_id = user_api_key_dict.user_id - if user_id != user_api_key_dict.user_id: + user_id = caller_user_id + if user_id != caller_user_id: raise HTTPException( status_code=status.HTTP_403_FORBIDDEN, detail={ diff --git a/litellm/proxy/management_endpoints/jwt_key_mapping_endpoints.py b/litellm/proxy/management_endpoints/jwt_key_mapping_endpoints.py index e474cb7d155..1ee5bfb0226 100644 --- a/litellm/proxy/management_endpoints/jwt_key_mapping_endpoints.py +++ b/litellm/proxy/management_endpoints/jwt_key_mapping_endpoints.py @@ -10,6 +10,7 @@ from litellm.proxy._types import ( hash_token, ) from litellm.proxy.auth.user_api_key_auth import user_api_key_auth +from litellm.proxy.management_endpoints.common_utils import _user_has_admin_view router = APIRouter() @@ -194,7 +195,8 @@ async def list_jwt_key_mappings( ): from litellm.proxy.proxy_server import prisma_client - if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN: + # Admin Viewer follows the read-parity rule. + if not _user_has_admin_view(user_api_key_dict): raise HTTPException( status_code=403, detail="Only proxy admins can list JWT key mappings" ) @@ -233,7 +235,8 @@ async def info_jwt_key_mapping( ): from litellm.proxy.proxy_server import prisma_client - if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN: + # Admin Viewer follows the read-parity rule. + if not _user_has_admin_view(user_api_key_dict): raise HTTPException( status_code=403, detail="Only proxy admins can get JWT key mapping info" ) diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index a01f5e63211..b112af1fe20 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -57,6 +57,7 @@ from litellm.proxy.management_endpoints.common_utils import ( _is_user_org_admin_for_team, _is_user_team_admin, _set_object_metadata_field, + _team_member_has_permission, ) from litellm.proxy.management_endpoints.model_management_endpoints import ( _add_model_to_db, @@ -809,6 +810,20 @@ async def _common_key_generation_helper( # noqa: PLR0915 from litellm.proxy.proxy_server import prisma_client, user_api_key_cache if prisma_client: + # Mirror the membership rule applied to /key/update: when the + # caller specifies an organization_id, require that they are a + # member of (or proxy admin over) the target organization. + _is_proxy_admin = ( + user_api_key_dict.user_role is not None + and user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN.value + ) + if not _is_proxy_admin: + await _validate_caller_can_assign_key_org( + user_api_key_dict=user_api_key_dict, + organization_id=data.organization_id, + prisma_client=prisma_client, + ) + org_table = await get_org_object( org_id=data.organization_id, user_api_key_cache=user_api_key_cache, @@ -1168,6 +1183,42 @@ def check_org_key_rpm_tpm_limits( ) +async def _validate_caller_can_assign_key_org( + user_api_key_dict: UserAPIKeyAuth, + organization_id: str, + prisma_client: PrismaClient, +) -> None: + """Reject ``/key/update`` requests that point a key at an organization + the caller does not belong to. + + Mirrors the org-membership rule already enforced on ``/key/list`` in + ``validate_key_list_check``. Proxy admins are checked at the call site. + """ + if user_api_key_dict.user_id is None: + raise HTTPException( + status_code=status.HTTP_403_FORBIDDEN, + detail="Cannot assign a key to an organization without a user_id on the caller's token", + ) + + user_row = await prisma_client.db.litellm_usertable.find_unique( + where={"user_id": user_api_key_dict.user_id}, + include={"organization_memberships": True}, + ) + memberships = ( + getattr(user_row, "organization_memberships", None) if user_row else None + ) + member_org_ids = { + membership.organization_id + for membership in (memberships or []) + if membership.organization_id is not None + } + if organization_id not in member_org_ids: + raise HTTPException( + status_code=status.HTTP_403_FORBIDDEN, + detail=f"Caller is not a member of organization_id={organization_id}", + ) + + async def _check_org_key_limits( org_table: LiteLLM_OrganizationTable, data: Union[GenerateKeyRequest, UpdateKeyRequest], @@ -2168,10 +2219,26 @@ async def _validate_update_key_data( user_api_key_cache=user_api_key_cache, ) + # When the caller asks to change the key's organization_id, require that + # they are a member of (or a proxy admin over) the target organization. + # Without this gate, any caller could assign their key to an arbitrary + # organization_id by passing it in the request body — VERIA-55 secondary + # IDOR. The check mirrors the membership rule already used on the + # `/key/list` filter path in `validate_key_list_check`. + _existing_org_id = getattr(existing_key_row, "organization_id", None) + if ( + data.organization_id is not None + and data.organization_id != _existing_org_id + and not _is_proxy_admin + ): + await _validate_caller_can_assign_key_org( + user_api_key_dict=user_api_key_dict, + organization_id=data.organization_id, + prisma_client=prisma_client, + ) + # Check org key limits only when throughput-related fields or organization_id change - _org_id_to_check = data.organization_id or getattr( - existing_key_row, "organization_id", None - ) + _org_id_to_check = data.organization_id or _existing_org_id _throughput_fields_changed = ( data.organization_id is not None or data.tpm_limit is not None @@ -3868,6 +3935,22 @@ async def _execute_virtual_key_regeneration( """Generate new token, update DB, invalidate cache, and return response.""" from litellm.proxy.proxy_server import hash_token + # Apply the same membership rule used on /key/update: when the caller + # asks to point the regenerated key at a different organization_id, + # require they are a member of (or proxy admin over) the target org. + if data is not None and data.organization_id is not None: + _existing_org_id = getattr(key_in_db, "organization_id", None) + _is_proxy_admin = ( + user_api_key_dict.user_role is not None + and user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN.value + ) + if data.organization_id != _existing_org_id and not _is_proxy_admin: + await _validate_caller_can_assign_key_org( + user_api_key_dict=user_api_key_dict, + organization_id=data.organization_id, + prisma_client=prisma_client, + ) + new_token = await get_new_token(data=data) new_token_hash = hash_token(new_token) new_token_key_name = f"sk-...{new_token[-4:]}" @@ -4436,6 +4519,26 @@ def _get_admin_team_ids_from_objects( ] +def _get_team_ids_with_key_list_permission_from_objects( + user_api_key_dict: UserAPIKeyAuth, + team_objects: List[LiteLLM_TeamTable], +) -> List[str]: + """Filter team objects to non-admin teams where the caller has /key/list + permission via team_member_permissions. These teams should grant the + caller full key visibility (same as a team admin), so other members' + keys and service account keys (user_id=NULL) are returned.""" + return [ + team.team_id + for team in team_objects + if not _is_user_team_admin(user_api_key_dict=user_api_key_dict, team_obj=team) + and _team_member_has_permission( + user_api_key_dict=user_api_key_dict, + team_obj=team, + permission=KeyManagementRoutes.KEY_LIST.value, + ) + ] + + def _get_member_team_ids_from_objects( user_api_key_dict: UserAPIKeyAuth, team_objects: List[LiteLLM_TeamTable], @@ -4589,6 +4692,17 @@ async def list_keys( user_api_key_dict=user_api_key_dict, team_objects=team_objects, ) + # Non-admin members with /key/list permission get full team-key + # visibility for that team — matching the UI contract that + # granting this permission lets them see all keys within the team. + list_permission_team_ids = ( + _get_team_ids_with_key_list_permission_from_objects( + user_api_key_dict=user_api_key_dict, + team_objects=team_objects, + ) + ) + if list_permission_team_ids: + admin_team_ids = list({*admin_team_ids, *list_permission_team_ids}) else: admin_team_ids = None diff --git a/litellm/proxy/management_endpoints/mcp_management_endpoints.py b/litellm/proxy/management_endpoints/mcp_management_endpoints.py index 9c510a568e6..729493e1df8 100644 --- a/litellm/proxy/management_endpoints/mcp_management_endpoints.py +++ b/litellm/proxy/management_endpoints/mcp_management_endpoints.py @@ -2120,7 +2120,8 @@ if MCP_AVAILABLE: Used by the UI to show a discovery grid when adding new MCP servers. """ - if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN: + # Admin Viewer follows the read-parity rule. + if not _user_has_admin_view(user_api_key_dict): raise HTTPException( status_code=403, detail={ @@ -2177,7 +2178,8 @@ if MCP_AVAILABLE: async def get_openapi_registry( user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), ): - if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN: + # Admin Viewer follows the read-parity rule. + if not _user_has_admin_view(user_api_key_dict): raise HTTPException( status_code=403, detail={ diff --git a/litellm/proxy/management_endpoints/team_endpoints.py b/litellm/proxy/management_endpoints/team_endpoints.py index 61247ce5deb..466ce47a6fd 100644 --- a/litellm/proxy/management_endpoints/team_endpoints.py +++ b/litellm/proxy/management_endpoints/team_endpoints.py @@ -4683,9 +4683,11 @@ async def team_member_permissions( complete_team_data = LiteLLM_TeamTable(**existing_team_row.model_dump()) + # Admin Viewer follows the read-parity rule: see team permissions like + # a Proxy Admin would. Team / org admins keep their existing scope. if ( hasattr(user_api_key_dict, "user_role") - and user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN.value + and not _user_has_admin_view(user_api_key_dict) and not _is_user_team_admin( user_api_key_dict=user_api_key_dict, team_obj=complete_team_data ) diff --git a/litellm/proxy/management_endpoints/ui_sso.py b/litellm/proxy/management_endpoints/ui_sso.py index 9dfc67370fe..74ee7c7220d 100644 --- a/litellm/proxy/management_endpoints/ui_sso.py +++ b/litellm/proxy/management_endpoints/ui_sso.py @@ -678,6 +678,7 @@ async def google_login( google_client_id=google_client_id, generic_client_id=generic_client_id, state=cli_state, + request=request, ) if return_to is not None and sso_redirect is not None: if SSOAuthenticationHandler._validate_return_to(return_to): @@ -1159,6 +1160,30 @@ async def get_generic_sso_response( authorization_code = request.query_params.get("code") if code_verifier: + # State-to-session-cookie binding. The non-PKCE branch below + # delegates to fastapi-sso's ``verify_and_process``, which + # performs its own session-cookie check. The PKCE branch + # bypasses that helper, so we validate the URL ``state`` + # against the ``litellm_oauth_state`` cookie set on the + # redirect response — without this an attacker can pre-mint + # a state + cached PKCE verifier and hijack a victim's auth + # code (Login-CSRF / token theft). + url_state = request.query_params.get("state") + cookie_state = request.cookies.get("litellm_oauth_state") + if ( + not url_state + or not cookie_state + or not secrets.compare_digest(url_state, cookie_state) + ): + raise ProxyException( + message=( + "Invalid OAuth state parameter — does not match " + "the browser-bound state cookie." + ), + type=ProxyErrorTypes.auth_error, + param="state", + code=status.HTTP_400_BAD_REQUEST, + ) if not authorization_code: raise ProxyException( message="Missing authorization code in callback", @@ -2147,6 +2172,7 @@ class SSOAuthenticationHandler: microsoft_client_id: Optional[str] = None, generic_client_id: Optional[str] = None, state: Optional[str] = None, + request: Optional[Request] = None, ) -> Optional[RedirectResponse]: """ Step 1. Call Get Login Redirect for the SSO provider. Send the redirect response to `redirect_url` @@ -2156,6 +2182,8 @@ class SSOAuthenticationHandler: google_client_id (Optional[str], optional): The Google Client ID. Defaults to None. microsoft_client_id (Optional[str], optional): The Microsoft Client ID. Defaults to None. generic_client_id (Optional[str], optional): The Generic Client ID. Defaults to None. + request: Optional FastAPI request, used to drive the ``Secure`` + attribute on the ``litellm_oauth_state`` CSRF cookie. Returns: RedirectResponse: The redirect response from the SSO provider. @@ -2266,6 +2294,7 @@ class SSOAuthenticationHandler: generic_sso=generic_sso, state=state, generic_authorization_endpoint=generic_authorization_endpoint, + request=request, ) raise ValueError( "Unknown SSO provider. Please setup SSO with client IDs https://docs.litellm.ai/docs/proxy/admin_ui_sso" @@ -2276,6 +2305,7 @@ class SSOAuthenticationHandler: generic_sso: Any, state: Optional[str] = None, generic_authorization_endpoint: Optional[str] = None, + request: Optional[Request] = None, ) -> Optional[RedirectResponse]: """ Get the redirect response for Generic SSO @@ -2285,10 +2315,13 @@ class SSOAuthenticationHandler: from litellm.proxy.proxy_server import redis_usage_cache, user_api_key_cache with generic_sso: - # TODO: state should be a random string and added to the user session with cookie - # or a cryptographicly signed state that we can verify stateless - # For simplification we are using a static state, this is not perfect but some - # SSO providers do not allow stateless verification + # State is bound to the caller's browser via a ``litellm_oauth_state`` + # HttpOnly cookie set on the redirect response below; the SSO + # callback validates the URL ``state`` against that cookie before + # completing the PKCE token exchange. Without this binding, an + # attacker who pre-mints a state + a cached PKCE verifier can hand + # the link to a victim and capture the resulting access token + # (Login CSRF / token theft). ( redirect_params, code_verifier, @@ -2355,6 +2388,31 @@ class SSOAuthenticationHandler: # Update the redirect response redirect_response.headers["location"] = new_url + + # Bind state to the user's browser session. The /callback + # handler validates the URL ``state`` against this cookie via + # ``secrets.compare_digest`` before exchanging the PKCE + # code_verifier. Only set the cookie when PKCE is in use + # (i.e. inside this ``code_verifier`` branch) so two + # concurrent SSO sessions — one PKCE, one plain — cannot + # overwrite each other's state cookie. + state_value = redirect_params.get("state") + if state_value and redirect_response is not None: + # Production-safe default: require HTTPS for the + # CSRF-protection cookie unless we can prove the + # incoming request is HTTP (local dev). Without + # ``Secure`` the cookie is sent over plain HTTP, + # letting a network observer read and replay the + # state value and bypass this protection. + secure_flag = request is None or request.url.scheme == "https" + redirect_response.set_cookie( + key="litellm_oauth_state", + value=state_value, + max_age=600, + httponly=True, + samesite="lax", + secure=secure_flag, + ) return redirect_response @staticmethod @@ -3972,6 +4030,7 @@ async def debug_sso_login(request: Request): microsoft_client_id=microsoft_client_id, google_client_id=google_client_id, generic_client_id=generic_client_id, + request=request, ) diff --git a/litellm/proxy/management_endpoints/usage_endpoints/ai_usage_chat.py b/litellm/proxy/management_endpoints/usage_endpoints/ai_usage_chat.py index 4de29e04092..a50ce1d3c48 100644 --- a/litellm/proxy/management_endpoints/usage_endpoints/ai_usage_chat.py +++ b/litellm/proxy/management_endpoints/usage_endpoints/ai_usage_chat.py @@ -440,6 +440,15 @@ def _resolve_fetch_kwargs( kwargs: Dict[str, Any] = {"start_date": start_date, "end_date": end_date} if fn_name == "get_usage_data": if not is_admin: + if user_id is None: + # Defense-in-depth: the endpoint guard in usage_endpoints/endpoints.py + # should have already rejected this. If we ever reach here it means + # a future caller invoked the helper without scoping — fail loudly + # rather than issuing an unfiltered global query. + raise ValueError( + "Non-admin caller has user_id=None; refusing to issue an " + "unscoped query. Endpoint-level guard missing." + ) kwargs["user_id"] = user_id elif fn_args.get("user_id"): kwargs["user_id"] = fn_args["user_id"] diff --git a/litellm/proxy/management_endpoints/usage_endpoints/endpoints.py b/litellm/proxy/management_endpoints/usage_endpoints/endpoints.py index 0dbe518afb7..d0df80fed0d 100644 --- a/litellm/proxy/management_endpoints/usage_endpoints/endpoints.py +++ b/litellm/proxy/management_endpoints/usage_endpoints/endpoints.py @@ -44,13 +44,17 @@ async def usage_ai_chat( """ from litellm.proxy.management_endpoints.common_utils import ( _user_has_admin_view, + require_caller_user_id_for_non_admin, ) from litellm.proxy.management_endpoints.usage_endpoints.ai_usage_chat import ( stream_usage_ai_chat, ) is_admin = _user_has_admin_view(user_api_key_dict) - user_id = user_api_key_dict.user_id + if is_admin: + user_id = user_api_key_dict.user_id + else: + user_id = require_caller_user_id_for_non_admin(user_api_key_dict) messages = [{"role": m.role, "content": m.content} for m in data.messages] return StreamingResponse( diff --git a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py index 6521abffb85..ce103f806e1 100644 --- a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py @@ -47,6 +47,8 @@ from litellm.types.passthrough_endpoints.pass_through_endpoints import ( ) from litellm.proxy.utils import is_known_model from litellm.proxy.vector_store_endpoints.utils import ( + assert_user_can_access_vector_store, + get_litellm_managed_vector_store, is_allowed_to_call_vector_store_endpoint, ) from litellm.secret_managers.main import get_secret_str @@ -533,6 +535,10 @@ async def milvus_proxy_route( ) if vector_store is None: raise Exception(f"Vector store not found for {vector_store_name}") + await assert_user_can_access_vector_store( + vector_store=vector_store, + user_api_key_dict=user_api_key_dict, + ) litellm_params = vector_store.get("litellm_params") or {} auth_credentials = provider_config.get_auth_credentials( litellm_params=litellm_params @@ -1438,6 +1444,10 @@ async def azure_proxy_route( ) if vector_store is None: raise Exception(f"Vector store not found for {vector_store_name}") + await assert_user_can_access_vector_store( + vector_store=vector_store, + user_api_key_dict=user_api_key_dict, + ) litellm_params = vector_store.get("litellm_params") or {} auth_credentials = provider_config.get_auth_credentials( litellm_params=litellm_params @@ -1777,6 +1787,11 @@ async def _base_vertex_proxy_route( request=request, api_key=api_key_to_use, ) + if router_credentials is not None: + await assert_user_can_access_vector_store( + vector_store=router_credentials, + user_api_key_dict=user_api_key_dict, + ) vertex_project: Optional[str] = get_vertex_project_id_from_url(endpoint) vertex_location: Optional[str] = get_vertex_location_from_url(endpoint) @@ -1913,11 +1928,11 @@ async def vertex_discovery_proxy_route( "Extracted vector store ID from endpoint: %s", vector_store_id ) - # Retrieve vector store credentials from the registry - vector_store_credentials = ( - passthrough_endpoint_router.get_vector_store_credentials( - vector_store_id=vector_store_id - ) + # Retrieve LiteLLM-managed vector store credentials if the datastore id + # is registered with LiteLLM. Unknown datastore ids keep the existing + # direct Vertex pass-through behavior. + vector_store_credentials = await get_litellm_managed_vector_store( + vector_store_id=vector_store_id ) if vector_store_credentials: @@ -1925,7 +1940,7 @@ async def vertex_discovery_proxy_route( "Found vector store credentials for ID: %s", vector_store_id ) else: - verbose_proxy_logger.warning( + verbose_proxy_logger.debug( "Vector store ID %s found in endpoint but no credentials found in registry", vector_store_id, ) diff --git a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py index 714b5f3c7b6..cc6c26fdf90 100644 --- a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py @@ -2324,14 +2324,10 @@ async def _register_pass_through_endpoint( dependencies = None if auth is not None and str(auth).lower() == "true": - # Authentication on a pass-through endpoint used to be enterprise- - # only — which left the OSS tier with no safe configuration: the - # default was ``auth=False`` (unauthenticated forwarder) and the - # safe ``auth=True`` raised at startup unless the operator had a - # license. The default is now ``True`` (safe-by-default), and - # turning it on no longer requires a license: an unauthenticated - # forwarder is a deployment choice the operator should be allowed - # to make explicitly, but the safe option must always be free. + # Authentication on a pass-through endpoint used to be enterprise-only. + # That left OSS with no safe configuration: auth=True raised at startup + # unless the operator had a license. The safe option must always be free, + # and unauthenticated forwarding should require explicit opt-in. dependencies = [Depends(user_api_key_auth)] if path not in LiteLLMRoutes.openai_routes.value: LiteLLMRoutes.openai_routes.value.append(path) diff --git a/litellm/proxy/policy_engine/attachment_registry.py b/litellm/proxy/policy_engine/attachment_registry.py index 6d1096d5ee9..8d5d8116919 100644 --- a/litellm/proxy/policy_engine/attachment_registry.py +++ b/litellm/proxy/policy_engine/attachment_registry.py @@ -220,6 +220,7 @@ class AttachmentRegistry: attachment: PolicyAttachment object to add """ self._attachments.append(attachment) + self._initialized = True verbose_proxy_logger.debug(f"Added attachment for policy: {attachment.policy}") def remove_attachments_for_policy(self, policy_name: str) -> int: diff --git a/litellm/proxy/policy_engine/policy_registry.py b/litellm/proxy/policy_engine/policy_registry.py index d3df16afde6..75017c46603 100644 --- a/litellm/proxy/policy_engine/policy_registry.py +++ b/litellm/proxy/policy_engine/policy_registry.py @@ -226,6 +226,7 @@ class PolicyRegistry: policy: Policy object to add """ self._policies[policy_name] = policy + self._initialized = True verbose_proxy_logger.debug(f"Added/updated policy: {policy_name}") def remove_policy(self, policy_name: str) -> bool: diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index a2bdc3ab3cf..3db13e8ada6 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -6,6 +6,7 @@ import inspect import io import os import random +import re import secrets import shutil import subprocess @@ -334,6 +335,7 @@ from litellm.proxy.management_endpoints.callback_management_endpoints import ( ) from litellm.proxy.management_endpoints.common_utils import ( _user_has_admin_privileges, + _user_has_admin_view, admin_can_invite_user, ) from litellm.proxy.management_endpoints.cost_tracking_settings import ( @@ -955,6 +957,85 @@ async def proxy_startup_event(app: FastAPI): # noqa: PLR0915 await proxy_shutdown_event() # type: ignore[reportGeneralTypeIssues] +def _generate_stable_operation_id(route: Any) -> str: + operation_id = re.sub(r"\W", "_", f"{route.name}{route.path_format}") + route_methods = sorted(route.methods or []) + if len(route_methods) == 1: + operation_id = f"{operation_id}_{route_methods[0].lower()}" + return operation_id + + +_OPENAPI_HTTP_METHODS = { + "delete", + "get", + "head", + "options", + "patch", + "post", + "put", + "trace", +} + + +def _strip_operation_id_method_suffix(operation_id: str) -> str: + base, separator, suffix = operation_id.rpartition("_") + if separator and suffix in _OPENAPI_HTTP_METHODS: + return base + return operation_id + + +def ensure_unique_openapi_operation_ids( + openapi_schema: Dict[str, Any], + reserved_operation_ids: Optional[Set[str]] = None, +) -> Dict[str, Any]: + operation_entries = [] + operation_id_counts: Dict[str, int] = {} + for path_item in openapi_schema.get("paths", {}).values(): + if not isinstance(path_item, dict): + continue + for method, operation in path_item.items(): + if method not in _OPENAPI_HTTP_METHODS or not isinstance(operation, dict): + continue + operation_id = operation.get("operationId") + if not isinstance(operation_id, str): + continue + operation_entries.append((method, operation, operation_id)) + operation_id_counts[operation_id] = ( + operation_id_counts.get(operation_id, 0) + 1 + ) + + used_operation_ids = set(reserved_operation_ids or set()) + seen_operation_ids: Set[str] = set() + for method, operation, operation_id in operation_entries: + should_rewrite = ( + operation_id_counts[operation_id] > 1 + or operation_id in used_operation_ids + or operation_id in seen_operation_ids + ) + if not should_rewrite: + seen_operation_ids.add(operation_id) + used_operation_ids.add(operation_id) + continue + + base_operation_id = _strip_operation_id_method_suffix(operation_id) + new_operation_id = f"{base_operation_id}_{method}" + suffix = 2 + while ( + new_operation_id in used_operation_ids + or new_operation_id in seen_operation_ids + ): + new_operation_id = f"{base_operation_id}_{method}_{suffix}" + suffix += 1 + operation["operationId"] = new_operation_id + seen_operation_ids.add(new_operation_id) + used_operation_ids.add(new_operation_id) + + if reserved_operation_ids is not None: + reserved_operation_ids.update(used_operation_ids) + + return openapi_schema + + app = FastAPI( docs_url=_get_docs_url(), redoc_url=_get_redoc_url(), @@ -964,6 +1045,7 @@ app = FastAPI( version=version, root_path=server_root_path, lifespan=proxy_startup_event, # type: ignore[reportGeneralTypeIssues] + generate_unique_id_function=_generate_stable_operation_id, ) vertex_live_passthrough_vertex_base = VertexBase() @@ -1043,6 +1125,7 @@ def get_openapi_schema(): from litellm.proxy._lazy_features import inject_lazy_stubs openapi_schema = inject_lazy_stubs(openapi_schema) + openapi_schema = ensure_unique_openapi_operation_ids(openapi_schema) # Fix Swagger UI execute path error when server_root_path is set if server_root_path: @@ -1074,6 +1157,7 @@ def custom_openapi(): from litellm.proxy._lazy_features import inject_lazy_stubs openapi_schema = inject_lazy_stubs(openapi_schema) + openapi_schema = ensure_unique_openapi_operation_ids(openapi_schema) # Fix Swagger UI execute path error when server_root_path is set if server_root_path: @@ -1845,6 +1929,9 @@ async def increment_spend_counters( user_id: Optional[str], response_cost: Optional[float], org_id: Optional[str] = None, + budget_reservation: Optional[dict] = None, + end_user_id: Optional[str] = None, + tags: Optional[List[str]] = None, ): """ Atomically increment spend counters for budget enforcement. @@ -1856,7 +1943,14 @@ async def increment_spend_counters( Awaited (not create_task) in the cost callback, so the counter is updated before the next request's auth check runs. """ + reserved_counter_keys = await _reconcile_budget_reservation_for_counter_update( + budget_reservation=budget_reservation, + response_cost=response_cost, + ) + if response_cost is None or response_cost == 0: + if budget_reservation is not None: + budget_reservation["finalized"] = True return if token is not None: @@ -1871,11 +1965,13 @@ async def increment_spend_counters( if isinstance(token, str) and token.startswith("sk-") else token ) - await _init_and_increment_spend_counter( - counter_key=f"spend:key:{hashed_token}", - source_cache_key=hashed_token, - increment=response_cost, - ) + key_counter_key = f"spend:key:{hashed_token}" + if key_counter_key not in reserved_counter_keys: + await _init_and_increment_spend_counter( + counter_key=key_counter_key, + source_cache_key=hashed_token, + increment=response_cost, + ) # Increment per-window budget counters for multi-budget keys key_obj = await user_api_key_cache.async_get_cache(key=hashed_token) @@ -1892,17 +1988,28 @@ async def increment_spend_counters( if isinstance(window, dict) else window.budget_duration ) - await spend_counter_cache.async_increment_cache( - key=f"spend:key:{hashed_token}:window:{duration}", - value=response_cost, - ) + key_window_counter = f"spend:key:{hashed_token}:window:{duration}" + if key_window_counter not in reserved_counter_keys: + from litellm.proxy.spend_tracking.budget_reservation import ( + get_budget_window_start, + ) + + await _init_and_increment_window_spend_counter( + counter_key=key_window_counter, + entity_type="Key", + entity_id=hashed_token, + window_start=get_budget_window_start(window), + increment=response_cost, + ) if team_id is not None: - await _init_and_increment_spend_counter( - counter_key=f"spend:team:{team_id}", - source_cache_key=f"team_id:{team_id}", - increment=response_cost, - ) + team_counter_key = f"spend:team:{team_id}" + if team_counter_key not in reserved_counter_keys: + await _init_and_increment_spend_counter( + counter_key=team_counter_key, + source_cache_key=f"team_id:{team_id}", + increment=response_cost, + ) # Increment per-window budget counters for multi-budget teams team_obj = await user_api_key_cache.async_get_cache(key=f"team_id:{team_id}") @@ -1919,36 +2026,157 @@ async def increment_spend_counters( if isinstance(window, dict) else window.budget_duration ) - await spend_counter_cache.async_increment_cache( - key=f"spend:team:{team_id}:window:{duration}", - value=response_cost, - ) + team_window_counter = f"spend:team:{team_id}:window:{duration}" + if team_window_counter not in reserved_counter_keys: + from litellm.proxy.spend_tracking.budget_reservation import ( + get_budget_window_start, + ) + + await _init_and_increment_window_spend_counter( + counter_key=team_window_counter, + entity_type="Team", + entity_id=team_id, + window_start=get_budget_window_start(window), + increment=response_cost, + ) if user_id is not None and team_id is not None: - await _init_and_increment_spend_counter( - counter_key=f"spend:team_member:{user_id}:{team_id}", - source_cache_key=f"team_membership:{user_id}:{team_id}", - increment=response_cost, - ) + team_member_counter_key = f"spend:team_member:{user_id}:{team_id}" + if team_member_counter_key not in reserved_counter_keys: + await _init_and_increment_spend_counter( + counter_key=team_member_counter_key, + source_cache_key=f"team_membership:{user_id}:{team_id}", + increment=response_cost, + ) if user_id is not None: - await _init_and_increment_spend_counter( - counter_key=f"spend:user:{user_id}", - source_cache_key=user_id, + user_counter_key = f"spend:user:{user_id}" + if user_counter_key not in reserved_counter_keys: + await _init_and_increment_spend_counter( + counter_key=user_counter_key, + source_cache_key=user_id, + increment=response_cost, + ) + + await _increment_end_user_and_tag_spend_counters( + end_user_id=end_user_id, + tags=tags, + response_cost=response_cost, + reserved_counter_keys=reserved_counter_keys, + ) + + await _increment_org_spend_counter( + org_id=org_id, + response_cost=response_cost, + reserved_counter_keys=reserved_counter_keys, + ) + if budget_reservation is not None: + budget_reservation["finalized"] = True + + +async def _reconcile_budget_reservation_for_counter_update( + budget_reservation: Optional[dict], + response_cost: Optional[float], +) -> Set[str]: + if budget_reservation is None: + return set() + + from litellm.proxy.spend_tracking.budget_reservation import ( + get_reserved_counter_keys, + invalidate_budget_reservation_counters, + reconcile_budget_reservation, + ) + + reserved_counter_keys = get_reserved_counter_keys( + budget_reservation=budget_reservation + ) + try: + await reconcile_budget_reservation( + budget_reservation=budget_reservation, + actual_cost=response_cost or 0.0, + finalize=False, + ) + except Exception: + verbose_proxy_logger.warning( + "Failed to reconcile budget reservation after persisted spend; invalidating reserved counters and continuing", + exc_info=True, + ) + try: + await invalidate_budget_reservation_counters( + budget_reservation=budget_reservation + ) + except Exception: + verbose_proxy_logger.exception( + "Failed to invalidate reserved counters after reservation reconciliation failed" + ) + return reserved_counter_keys + + +async def _increment_end_user_and_tag_spend_counters( + end_user_id: Optional[str], + tags: Optional[List[str]], + response_cost: float, + reserved_counter_keys: Set[str], +) -> None: + if end_user_id is not None: + await _init_and_increment_unreserved_spend_counter( + counter_key=f"spend:end_user:{end_user_id}", + source_cache_key=f"end_user_id:{end_user_id}", increment=response_cost, + reserved_counter_keys=reserved_counter_keys, ) - if org_id is not None: - await _init_and_increment_spend_counter( - counter_key=f"spend:org:{org_id}", - source_cache_key=f"org_id:{org_id}", + if tags is None: + return + + seen_tags: Set[str] = set() + for tag_name in tags: + if not tag_name or not isinstance(tag_name, str) or tag_name in seen_tags: + continue + seen_tags.add(tag_name) + await _init_and_increment_unreserved_spend_counter( + counter_key=f"spend:tag:{tag_name}", + source_cache_key=f"tag:{tag_name}", increment=response_cost, + reserved_counter_keys=reserved_counter_keys, ) +async def _increment_org_spend_counter( + org_id: Optional[str], + response_cost: float, + reserved_counter_keys: Set[str], +) -> None: + if org_id is None: + return + + await _init_and_increment_unreserved_spend_counter( + counter_key=f"spend:org:{org_id}", + source_cache_key=[f"org_id:{org_id}:with_budget", f"org_id:{org_id}"], + increment=response_cost, + reserved_counter_keys=reserved_counter_keys, + ) + + +async def _init_and_increment_unreserved_spend_counter( + counter_key: str, + source_cache_key: Union[str, List[str]], + increment: float, + reserved_counter_keys: Set[str], +) -> None: + if counter_key in reserved_counter_keys: + return + + await _init_and_increment_spend_counter( + counter_key=counter_key, + source_cache_key=source_cache_key, + increment=increment, + ) + + async def _init_and_increment_spend_counter( counter_key: str, - source_cache_key: str, + source_cache_key: Union[str, List[str]], increment: float, ): """ @@ -1967,31 +2195,163 @@ async def _init_and_increment_spend_counter( under-counting (would allow overspend). 4. Increment atomically (both in-memory + Redis) """ - current = await spend_counter_cache.async_get_cache(key=counter_key) - if current is None: + await _ensure_spend_counter_initialized( + counter_key=counter_key, + source_cache_key=source_cache_key, + ) + await _increment_spend_counter_cache(counter_key=counter_key, increment=increment) + + +async def _init_and_increment_window_spend_counter( + counter_key: str, + entity_type: str, + entity_id: str, + window_start: Optional[datetime], + increment: float, +): + if window_start is None: + verbose_proxy_logger.warning( + "Skipping spend counter increment for invalid budget window %s", + counter_key, + ) + return + + initialized = await _ensure_window_spend_counter_initialized( + counter_key=counter_key, + entity_type=entity_type, + entity_id=entity_id, + window_start=window_start, + ) + if initialized is False: + return + await _increment_spend_counter_cache(counter_key=counter_key, increment=increment) + + +async def _ensure_spend_counter_initialized( + counter_key: str, + source_cache_key: Union[str, List[str]], +): + is_warm = await _is_spend_counter_cache_warm(counter_key=counter_key) + if is_warm is False: # Shares the per-counter lock with get_current_spend. db_spend = await SpendCounterReseed.coalesced( prisma_client=prisma_client, spend_counter_cache=spend_counter_cache, counter_key=counter_key, + require_cache_warm=True, ) if db_spend is None: # DB unavailable - fall back to in-process cache (may be stale). - source = await user_api_key_cache.async_get_cache(key=source_cache_key) - base_spend: float = 0.0 - if source is not None: - if isinstance(source, dict): - base_spend = source.get("spend", 0.0) or 0.0 - else: - base_spend = getattr(source, "spend", 0.0) or 0.0 + base_spend = await _get_source_cache_base_spend( + source_cache_key=source_cache_key + ) if base_spend > 0: - await spend_counter_cache.async_increment_cache( - key=counter_key, value=base_spend, refresh_ttl=True + await _increment_spend_counter_cache( + counter_key=counter_key, increment=base_spend ) - await spend_counter_cache.async_increment_cache( - key=counter_key, value=increment, refresh_ttl=True + +async def _get_source_cache_base_spend( + source_cache_key: Union[str, List[str]], +) -> float: + source_cache_keys = ( + [source_cache_key] if isinstance(source_cache_key, str) else source_cache_key ) + for cache_key in source_cache_keys: + source = await user_api_key_cache.async_get_cache(key=cache_key) + if source is None: + continue + if isinstance(source, dict): + return float(source.get("spend", 0.0) or 0.0) + return float(getattr(source, "spend", 0.0) or 0.0) + return 0.0 + + +async def _ensure_window_spend_counter_initialized( + counter_key: str, + entity_type: str, + entity_id: str, + window_start: datetime, +) -> bool: + is_warm = await _is_spend_counter_cache_warm(counter_key=counter_key) + if is_warm is True: + return True + + window_spend = await SpendCounterReseed.coalesced_window( + prisma_client=prisma_client, + spend_counter_cache=spend_counter_cache, + counter_key=counter_key, + entity_type=entity_type, + entity_id=entity_id, + window_start=window_start, + ) + if window_spend is None: + verbose_proxy_logger.warning( + "Skipping cold spend counter seed for %s because window spend could not be loaded", + counter_key, + ) + return False + return True + + +async def _is_spend_counter_cache_warm(counter_key: str) -> bool: + if spend_counter_cache.redis_cache is not None: + try: + current_value = await spend_counter_cache.redis_cache.async_get_cache( + key=counter_key, + ) + if current_value is None: + return False + spend_counter_cache.in_memory_cache.set_cache( + key=counter_key, + value=current_value, + ) + return True + except Exception as e: + verbose_proxy_logger.debug( + "Unable to read Redis spend counter %s before initialization, falling back to in-memory: %s", + counter_key, + e, + ) + + return spend_counter_cache.in_memory_cache.get_cache(key=counter_key) is not None + + +async def _increment_spend_counter_cache(counter_key: str, increment: float): + if spend_counter_cache.redis_cache is not None: + try: + current_value = await spend_counter_cache.redis_cache.async_increment( + key=counter_key, + value=increment, + refresh_ttl=True, + ) + except Exception: + await _invalidate_spend_counter(counter_key=counter_key) + raise + spend_counter_cache.in_memory_cache.set_cache( + key=counter_key, + value=current_value, + ) + return current_value + + return await spend_counter_cache.async_increment_cache( + key=counter_key, + value=increment, + refresh_ttl=True, + ) + + +async def _invalidate_spend_counter(counter_key: str): + spend_counter_cache.in_memory_cache.delete_cache(key=counter_key) + if spend_counter_cache.redis_cache is not None: + try: + await spend_counter_cache.redis_cache.async_delete_cache(key=counter_key) + except Exception: + verbose_proxy_logger.debug( + "Unable to delete stale spend counter %s after increment failure", + counter_key, + exc_info=True, + ) async def update_cache( # noqa: PLR0915 @@ -5889,10 +6249,15 @@ async def initialize( # noqa: PLR0915 if litellm_log_setting.upper() == "INFO": import logging - from litellm._logging import verbose_proxy_logger, verbose_router_logger + from litellm._logging import ( + verbose_logger, + verbose_proxy_logger, + verbose_router_logger, + ) # this must ALWAYS remain logging.INFO, DO NOT MODIFY THIS + verbose_logger.setLevel(level=logging.INFO) # set package log to info verbose_router_logger.setLevel( level=logging.INFO ) # set router logs to info @@ -11603,7 +11968,7 @@ async def alerting_settings( detail={"error": CommonProxyErrors.db_not_connected_error.value}, ) - if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN: + if not _user_has_admin_view(user_api_key_dict): raise HTTPException( status_code=400, detail={ @@ -12715,7 +13080,7 @@ async def invitation_info( detail={"error": CommonProxyErrors.db_not_connected_error.value}, ) - if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN: + if not _user_has_admin_view(user_api_key_dict): raise HTTPException( status_code=400, detail={ @@ -13137,7 +13502,7 @@ async def get_config_general_settings( detail={"error": CommonProxyErrors.db_not_connected_error.value}, ) - if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN: + if not _user_has_admin_view(user_api_key_dict): raise HTTPException( status_code=400, detail={"error": CommonProxyErrors.not_allowed_access.value}, @@ -13201,7 +13566,7 @@ async def get_config_list( detail={"error": CommonProxyErrors.db_not_connected_error.value}, ) - if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN: + if not _user_has_admin_view(user_api_key_dict): raise HTTPException( status_code=400, detail={ @@ -13913,8 +14278,8 @@ async def get_model_cost_map_reload_status( Get the status of the scheduled model cost map reload job. """ - # Check if user is admin - if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN: + # Read-only status check — admin viewers can read. + if not _user_has_admin_view(user_api_key_dict): raise HTTPException( status_code=403, detail=f"Access denied. Admin role required. Current role: {user_api_key_dict.user_role}", @@ -14016,7 +14381,8 @@ async def get_model_cost_map_source( - fallback_reason: human-readable reason why remote failed (null on success) - model_count: number of models in the currently loaded cost map """ - if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN: + # Read-only source info — admin viewers can read. + if not _user_has_admin_view(user_api_key_dict): raise HTTPException( status_code=403, detail=f"Access denied. Admin role required. Current role: {user_api_key_dict.user_role}", @@ -14273,8 +14639,8 @@ async def get_anthropic_beta_headers_reload_status( Get the status of the scheduled Anthropic beta headers reload job. """ - # Check if user is admin - if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN: + # Read-only status — admin viewers can read. + if not _user_has_admin_view(user_api_key_dict): raise HTTPException( status_code=403, detail=f"Access denied. Admin role required. Current role: {user_api_key_dict.user_role}", @@ -14380,7 +14746,8 @@ async def get_adaptive_router_state( adaptive-router deployment. Each snapshot's `router_name` field identifies which deployment it came from. """ - if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN: + # Read-only state — admin viewers can read. + if not _user_has_admin_view(user_api_key_dict): raise HTTPException( status_code=403, detail={"error": CommonProxyErrors.not_allowed_access.value}, diff --git a/litellm/proxy/rag_endpoints/endpoints.py b/litellm/proxy/rag_endpoints/endpoints.py index 95ca51612fc..498d77f7535 100644 --- a/litellm/proxy/rag_endpoints/endpoints.py +++ b/litellm/proxy/rag_endpoints/endpoints.py @@ -15,6 +15,7 @@ from fastapi.responses import ORJSONResponse import litellm from litellm._logging import verbose_proxy_logger +from litellm.constants import DEFAULT_MAX_RECURSE_DEPTH from litellm.proxy._types import * from litellm.proxy.auth.user_api_key_auth import UserAPIKeyAuth, user_api_key_auth from litellm.proxy.common_utils.http_parsing_utils import ( @@ -22,10 +23,88 @@ from litellm.proxy.common_utils.http_parsing_utils import ( _safe_get_request_headers, get_form_data, ) +from litellm.proxy.vector_store_endpoints.utils import ( + assert_user_can_access_vector_store_id, +) router = APIRouter() +def _raise_vector_store_scan_depth_exceeded() -> None: + raise HTTPException( + status_code=400, + detail={ + "error": f"Max depth of {DEFAULT_MAX_RECURSE_DEPTH} exceeded while scanning vector_store_id values" + }, + ) + + +def _append_payload_to_scan_stack( + payload_stack: list[tuple[Any, int]], + value: Any, + next_depth: int, +) -> None: + if isinstance(value, dict): + if next_depth > DEFAULT_MAX_RECURSE_DEPTH: + _raise_vector_store_scan_depth_exceeded() + payload_stack.append((value, next_depth)) + elif isinstance(value, list): + if next_depth > DEFAULT_MAX_RECURSE_DEPTH: + if any(isinstance(item, (dict, list)) for item in value): + _raise_vector_store_scan_depth_exceeded() + return + payload_stack.append((value, next_depth)) + + +def _collect_vector_store_ids_from_payload(payload: Any) -> set[str]: + vector_store_ids: set[str] = set() + payload_stack = [(payload, 0)] + + while payload_stack: + current_payload, depth = payload_stack.pop() + if depth > DEFAULT_MAX_RECURSE_DEPTH: + _raise_vector_store_scan_depth_exceeded() + + if isinstance(current_payload, dict): + for key, value in current_payload.items(): + if key == "vector_store_id": + if not isinstance(value, str) or not value: + raise HTTPException( + status_code=400, + detail={ + "error": "vector_store_id must be a non-empty string" + }, + ) + vector_store_ids.add(value) + continue + if isinstance(value, (dict, list)): + _append_payload_to_scan_stack( + payload_stack=payload_stack, + value=value, + next_depth=depth + 1, + ) + elif isinstance(current_payload, list): + for item in current_payload: + _append_payload_to_scan_stack( + payload_stack=payload_stack, + value=item, + next_depth=depth + 1, + ) + + return vector_store_ids + + +async def _authorize_nested_vector_store_ids( + payload: Any, + user_api_key_dict: UserAPIKeyAuth, +) -> None: + for vector_store_id in sorted(_collect_vector_store_ids_from_payload(payload)): + await assert_user_can_access_vector_store_id( + vector_store_id=vector_store_id, + user_api_key_dict=user_api_key_dict, + ) + + def _build_file_metadata_entry( response: Any, file_data: Optional[Tuple[str, bytes, str]] = None, @@ -385,6 +464,11 @@ async def rag_ingest( }, ) + await _authorize_nested_vector_store_ids( + payload=ingest_options, + user_api_key_dict=user_api_key_dict, + ) + # Add litellm data request_data: Dict[str, Any] = {} request_data = await add_litellm_data_to_request( @@ -537,11 +621,20 @@ async def rag_query( status_code=400, detail={"error": "retrieval_config is required"}, ) + if not isinstance(retrieval_config, dict): + raise HTTPException( + status_code=400, + detail={"error": "retrieval_config must be an object"}, + ) if "vector_store_id" not in retrieval_config: raise HTTPException( status_code=400, detail={"error": "retrieval_config must contain 'vector_store_id'"}, ) + await _authorize_nested_vector_store_ids( + payload=retrieval_config, + user_api_key_dict=user_api_key_dict, + ) # Add litellm data request_data: Dict[str, Any] = {} diff --git a/litellm/proxy/route_llm_request.py b/litellm/proxy/route_llm_request.py index 17cc4374560..bfe6b8484fa 100644 --- a/litellm/proxy/route_llm_request.py +++ b/litellm/proxy/route_llm_request.py @@ -6,6 +6,18 @@ from fastapi import HTTPException, status import litellm from litellm.proxy._types import UserAPIKeyAuth +# Router-internal mock_testing_* flag names — kept in sync with +# ``litellm.types.router.MockRouterTestingParams`` by the test +# ``test_mock_testing_kwarg_names_matches_dataclass``. Hardcoding (rather +# than deriving via ``dataclasses.fields(MockRouterTestingParams)`` at +# import time) avoids a cyclic import: ``litellm.types.router`` imports +# back into proxy modules before this module finishes loading. +_MOCK_TESTING_KWARG_NAMES: tuple = ( + "mock_testing_fallbacks", + "mock_testing_context_fallbacks", + "mock_testing_content_policy_fallbacks", +) + if TYPE_CHECKING: from litellm.router import Router as _Router @@ -322,6 +334,13 @@ async def route_request( # noqa: PLR0915 - Complex routing function, refactorin """ await add_shared_session_to_data(data) + # Strip router-internal mock_testing_* flags. Combined with an + # unauthorized fallback in ``router_settings_override`` they let a + # caller deterministically execute requests against restricted + # models. VERIA-44. + for _key in _MOCK_TESTING_KWARG_NAMES: + data.pop(_key, None) + team_id = get_team_id_from_data(data) router_model_names = llm_router.model_names if llm_router is not None else [] diff --git a/litellm/proxy/spend_tracking/budget_reservation.py b/litellm/proxy/spend_tracking/budget_reservation.py new file mode 100644 index 00000000000..1d296611bfc --- /dev/null +++ b/litellm/proxy/spend_tracking/budget_reservation.py @@ -0,0 +1,1029 @@ +from __future__ import annotations + +import json +from dataclasses import dataclass +from datetime import datetime, timedelta, timezone +from typing import Any, Dict, List, Optional, Sequence, cast + +import litellm +from litellm._logging import verbose_proxy_logger +from litellm.caching import DualCache +from litellm.litellm_core_utils.duration_parser import duration_in_seconds +from litellm.proxy._types import ( + LiteLLM_TeamMembership, + LiteLLM_TeamTable, + LiteLLM_UserTable, + UserAPIKeyAuth, +) +from litellm.proxy.auth.auth_utils import get_model_from_request +from litellm.proxy.auth.route_checks import RouteChecks +from litellm.proxy.utils import PrismaClient, ProxyLogging +from litellm.router import Router + + +@dataclass +class _BudgetCounter: + counter_key: str + max_budget: float + fallback_spend: float + entity_type: str + entity_id: str + source_cache_key: Optional[str] = None + spend_log_entity_id: Optional[str] = None + window_start: Optional[datetime] = None + + +class _CounterReservationUnavailable(Exception): + def __init__( + self, + touched_counter: bool = False, + counter_invalidated: bool = False, + ) -> None: + self.touched_counter = touched_counter + self.counter_invalidated = counter_invalidated + super().__init__("Counter reservation unavailable") + + +def get_reserved_counter_keys(budget_reservation: Optional[dict]) -> set: + if not budget_reservation: + return set() + entries = budget_reservation.get("entries") or [] + return { + entry["counter_key"] + for entry in entries + if isinstance(entry, dict) and entry.get("counter_key") is not None + } + + +async def reserve_budget_for_request( + request_body: dict, + route: str, + llm_router: Optional[Router], + valid_token: Optional[UserAPIKeyAuth], + team_object: Optional[LiteLLM_TeamTable], + user_object: Optional[LiteLLM_UserTable], + prisma_client: Optional[PrismaClient], + user_api_key_cache: DualCache, + proxy_logging_obj: ProxyLogging, + end_user_id: Optional[str] = None, + end_user_object: Optional[Any] = None, +) -> Optional[dict]: + if valid_token is None or not RouteChecks.is_llm_api_route(route=route): + return None + if route in {"/models", "/v1/models", "/utils/token_counter"}: + return None + if get_model_from_request(request_body, route) is None: + return None + + counters = await _get_budget_counters( + request_body=request_body, + valid_token=valid_token, + team_object=team_object, + user_object=user_object, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + proxy_logging_obj=proxy_logging_obj, + end_user_id=end_user_id, + end_user_object=end_user_object, + ) + if not counters: + return None + + current_spend_by_counter_key: Dict[str, float] = {} + reservation_cost = estimate_request_max_cost( + request_body=request_body, + route=route, + llm_router=llm_router, + ) + if reservation_cost is None: + reservation_cost = await _get_smallest_remaining_budget( + counters=counters, + current_spend_by_counter_key=current_spend_by_counter_key, + ) + if reservation_cost is None or reservation_cost <= 0: + return None + + applied_entries: List[Dict[str, Any]] = [] + try: + for counter in counters: + entry = _counter_to_reservation_entry( + counter=counter, + reserved_cost=reservation_cost, + ) + applied_entries.append(entry) + try: + reserved_value = await _reserve_counter( + counter=counter, + reservation_cost=reservation_cost, + ) + except _CounterReservationUnavailable as exc: + if exc.touched_counter and not exc.counter_invalidated: + await _release_applied_entries_best_effort( + entries=[entry], + default_reserved_cost=reservation_cost, + ) + applied_entries.remove(entry) + continue + + if reserved_value is not None: + current_spend = reserved_value + else: + cached_spend = current_spend_by_counter_key.get(counter.counter_key) + if cached_spend is None: + cached_spend = await _get_current_counter_value(counter=counter) + current_spend = cached_spend + reservation_cost + if current_spend > counter.max_budget: + remaining_before_reservation = counter.max_budget - ( + current_spend - reservation_cost + ) + if remaining_before_reservation > 1e-12: + await _resize_applied_reservation( + entries=applied_entries, + current_reserved_cost=reservation_cost, + new_reserved_cost=remaining_before_reservation, + ) + reservation_cost = remaining_before_reservation + continue + raise litellm.BudgetExceededError( + current_cost=current_spend, + max_budget=counter.max_budget, + message=( + "Budget has been exceeded! " + f"{counter.entity_type}={counter.entity_id} " + f"Current cost: {current_spend}, " + f"Max budget: {counter.max_budget}" + ), + ) + except Exception: + await _release_applied_entries_best_effort( + entries=applied_entries, + default_reserved_cost=reservation_cost, + ) + raise + + if not applied_entries: + return None + + return { + "reserved_cost": reservation_cost, + "entries": applied_entries, + "finalized": False, + } + + +async def reconcile_budget_reservation( + budget_reservation: Optional[dict], + actual_cost: Optional[float], + finalize: bool = True, +) -> None: + if not budget_reservation or budget_reservation.get("finalized") is True: + return + + reserved_cost = float(budget_reservation.get("reserved_cost") or 0.0) + actual = float(actual_cost or 0.0) + await _set_reserved_entries_actual_cost( + entries=budget_reservation.get("entries") or [], + actual_cost=actual, + default_reserved_cost=reserved_cost, + ) + if finalize: + budget_reservation["finalized"] = True + + +async def release_budget_reservation(budget_reservation: Optional[dict]) -> None: + await reconcile_budget_reservation( + budget_reservation=budget_reservation, + actual_cost=0.0, + ) + + +async def invalidate_budget_reservation_counters( + budget_reservation: Optional[dict], +) -> None: + if budget_reservation is None: + return + + from litellm.proxy.proxy_server import _invalidate_spend_counter + + for counter_key in get_reserved_counter_keys(budget_reservation=budget_reservation): + await _invalidate_spend_counter(counter_key=counter_key) + + +async def _get_budget_counters( + request_body: dict, + valid_token: UserAPIKeyAuth, + team_object: Optional[LiteLLM_TeamTable], + user_object: Optional[LiteLLM_UserTable], + prisma_client: Optional[PrismaClient], + user_api_key_cache: DualCache, + proxy_logging_obj: ProxyLogging, + end_user_id: Optional[str] = None, + end_user_object: Optional[Any] = None, +) -> List[_BudgetCounter]: + counters: List[_BudgetCounter] = [] + + if valid_token.token is not None: + if valid_token.max_budget is not None and valid_token.max_budget > 0: + counters.append( + _BudgetCounter( + counter_key=f"spend:key:{valid_token.token}", + source_cache_key=valid_token.token, + max_budget=float(valid_token.max_budget), + fallback_spend=float(valid_token.spend or 0.0), + entity_type="Key", + entity_id=valid_token.token, + ) + ) + counters.extend( + _get_budget_limit_counters( + entity_prefix=f"spend:key:{valid_token.token}", + entity_type="Key", + entity_id=valid_token.token, + budget_limits=valid_token.budget_limits, + fallback_spend=float(valid_token.spend or 0.0), + ) + ) + + if team_object is not None and team_object.team_id is not None: + team_id = team_object.team_id + if team_object.max_budget is not None and team_object.max_budget > 0: + counters.append( + _BudgetCounter( + counter_key=f"spend:team:{team_id}", + source_cache_key=f"team_id:{team_id}", + max_budget=float(team_object.max_budget), + fallback_spend=float(team_object.spend or 0.0), + entity_type="Team", + entity_id=team_id, + ) + ) + counters.extend( + _get_budget_limit_counters( + entity_prefix=f"spend:team:{team_id}", + entity_type="Team", + entity_id=team_id, + budget_limits=team_object.budget_limits, + fallback_spend=float(team_object.spend or 0.0), + ) + ) + + if ( + (team_object is None or team_object.team_id is None) + and user_object is not None + and user_object.user_id is not None + and user_object.max_budget is not None + and user_object.max_budget > 0 + ): + counters.append( + _BudgetCounter( + counter_key=f"spend:user:{user_object.user_id}", + source_cache_key=user_object.user_id, + max_budget=float(user_object.max_budget), + fallback_spend=float(user_object.spend or 0.0), + entity_type="User", + entity_id=user_object.user_id, + ) + ) + + end_user_counter = await _get_end_user_budget_counter( + valid_token=valid_token, + end_user_id=end_user_id, + end_user_object=end_user_object, + ) + if end_user_counter is not None: + counters.append(end_user_counter) + + counters.extend( + await _get_tag_budget_counters( + request_body=request_body, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + proxy_logging_obj=proxy_logging_obj, + ) + ) + + team_member_counter = await _get_team_member_budget_counter( + valid_token=valid_token, + team_object=team_object, + user_object=user_object, + user_api_key_cache=user_api_key_cache, + ) + if team_member_counter is not None: + counters.append(team_member_counter) + + org_counter = await _get_org_budget_counter( + valid_token=valid_token, + team_object=team_object, + user_api_key_cache=user_api_key_cache, + ) + if org_counter is not None: + counters.append(org_counter) + + return counters + + +async def _get_end_user_budget_counter( + valid_token: UserAPIKeyAuth, + end_user_id: Optional[str], + end_user_object: Optional[Any], +) -> Optional[_BudgetCounter]: + end_user_id = end_user_id or valid_token.end_user_id + if end_user_id is None: + return None + + source_cache_key = f"end_user_id:{end_user_id}" + max_budget = _to_float(valid_token.end_user_max_budget) + fallback_spend = 0.0 + if end_user_object is not None: + fallback_spend = _to_float(_get_value(end_user_object, "spend")) or 0.0 + if max_budget is None: + budget_table = _get_value(end_user_object, "litellm_budget_table") + max_budget = _to_float(_get_value(budget_table, "max_budget")) + + if max_budget is None or max_budget <= 0: + return None + + return _BudgetCounter( + counter_key=f"spend:end_user:{end_user_id}", + source_cache_key=source_cache_key, + max_budget=max_budget, + fallback_spend=fallback_spend, + entity_type="EndUser", + entity_id=end_user_id, + ) + + +async def _get_tag_budget_counters( + request_body: dict, + prisma_client: Optional[PrismaClient], + user_api_key_cache: DualCache, + proxy_logging_obj: ProxyLogging, +) -> List[_BudgetCounter]: + from litellm.proxy.common_utils.http_parsing_utils import get_tags_from_request_body + from litellm.proxy.auth.auth_checks import get_tag_objects_batch + + tag_names = _dedupe_tags(get_tags_from_request_body(request_body=request_body)) + if not tag_names: + return [] + + tag_objects = await get_tag_objects_batch( + tag_names=tag_names, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + proxy_logging_obj=proxy_logging_obj, + ) + + counters: List[_BudgetCounter] = [] + for tag_name in tag_names: + tag_object = tag_objects.get(tag_name) + if tag_object is None: + continue + budget_table = _get_value(tag_object, "litellm_budget_table") + max_budget = _to_float(_get_value(budget_table, "max_budget")) + if max_budget is None or max_budget <= 0: + continue + counters.append( + _BudgetCounter( + counter_key=f"spend:tag:{tag_name}", + source_cache_key=f"tag:{tag_name}", + max_budget=max_budget, + fallback_spend=_to_float(_get_value(tag_object, "spend")) or 0.0, + entity_type="Tag", + entity_id=tag_name, + ) + ) + return counters + + +def _dedupe_tags(tags: List[str]) -> List[str]: + seen = set() + deduped_tags = [] + for tag in tags: + if tag in seen: + continue + seen.add(tag) + deduped_tags.append(tag) + return deduped_tags + + +async def _get_team_member_budget_counter( + valid_token: UserAPIKeyAuth, + team_object: Optional[LiteLLM_TeamTable], + user_object: Optional[LiteLLM_UserTable], + user_api_key_cache: DualCache, +) -> Optional[_BudgetCounter]: + if ( + team_object is None + or team_object.team_id is None + or user_object is None + or valid_token.user_id is None + ): + return None + + membership_cache_key = ( + f"team_membership:{valid_token.user_id}:{team_object.team_id}" + ) + cached_team_membership = await user_api_key_cache.async_get_cache( + key=membership_cache_key + ) + team_membership: Optional[LiteLLM_TeamMembership] = None + if isinstance(cached_team_membership, LiteLLM_TeamMembership): + team_membership = cached_team_membership + elif isinstance(cached_team_membership, dict): + team_membership = LiteLLM_TeamMembership(**cached_team_membership) + + team_member_budget: Optional[float] = None + if team_membership is not None and team_membership.litellm_budget_table is not None: + team_member_budget = team_membership.litellm_budget_table.max_budget + else: + default_budget_id = (team_object.metadata or {}).get("team_member_budget_id") + if isinstance(default_budget_id, str): + default_budget = await user_api_key_cache.async_get_cache( + key=f"team_member_default_budget:{default_budget_id}", + ) + team_member_budget = _to_float(_get_value(default_budget, "max_budget")) + + if team_member_budget is None or team_member_budget <= 0: + return None + + team_member_spend = ( + cast(LiteLLM_TeamMembership, team_membership).spend + if team_membership is not None + else 0.0 + ) + return _BudgetCounter( + counter_key=f"spend:team_member:{valid_token.user_id}:{team_object.team_id}", + source_cache_key=membership_cache_key, + max_budget=float(team_member_budget), + fallback_spend=float(team_member_spend or 0.0), + entity_type="TeamMember", + entity_id=f"{valid_token.user_id}:{team_object.team_id}", + ) + + +async def _get_org_budget_counter( + valid_token: UserAPIKeyAuth, + team_object: Optional[LiteLLM_TeamTable], + user_api_key_cache: DualCache, +) -> Optional[_BudgetCounter]: + org_id: Optional[str] = None + if valid_token.org_id is not None: + org_id = valid_token.org_id + elif team_object is not None and team_object.organization_id is not None: + org_id = team_object.organization_id + if org_id is None: + return None + + org_table = await user_api_key_cache.async_get_cache( + key=f"org_id:{org_id}:with_budget", + ) + if org_table is None: + return None + + org_budget_table = _get_value(org_table, "litellm_budget_table") + if org_budget_table is None: + return None + + org_max_budget = _to_float(_get_value(org_budget_table, "max_budget")) + if org_max_budget is None or org_max_budget <= 0: + return None + + org_spend = _to_float(_get_value(org_table, "spend")) or 0.0 + return _BudgetCounter( + counter_key=f"spend:org:{org_id}", + source_cache_key=f"org_id:{org_id}:with_budget", + max_budget=org_max_budget, + fallback_spend=org_spend, + entity_type="Organization", + entity_id=org_id, + ) + + +def _get_budget_limit_counters( + entity_prefix: str, + entity_type: str, + entity_id: str, + budget_limits: Optional[Sequence[Any]], + fallback_spend: float, +) -> List[_BudgetCounter]: + counters: List[_BudgetCounter] = [] + if not budget_limits: + return counters + + for window in budget_limits: + window_dict = _coerce_window(window) + budget_duration = window_dict.get("budget_duration") + max_budget = window_dict.get("max_budget") + if not budget_duration or max_budget is None or max_budget <= 0: + continue + window_start = get_budget_window_start(window_dict) + if window_start is None: + verbose_proxy_logger.warning( + "Skipping budget window with invalid duration for %s=%s: %s", + entity_type, + entity_id, + budget_duration, + ) + continue + counters.append( + _BudgetCounter( + counter_key=f"{entity_prefix}:window:{budget_duration}", + max_budget=float(max_budget), + fallback_spend=0.0, + entity_type=entity_type, + entity_id=f"{entity_id}:{budget_duration}", + spend_log_entity_id=entity_id, + window_start=window_start, + ) + ) + return counters + + +def _coerce_window(window: Any) -> dict: + if isinstance(window, dict): + return window + if isinstance(window, str): + try: + parsed = json.loads(window) + return parsed if isinstance(parsed, dict) else {} + except Exception: + return {} + if hasattr(window, "model_dump"): + return window.model_dump() + return {} + + +async def _get_smallest_remaining_budget( + counters: List[_BudgetCounter], + current_spend_by_counter_key: Dict[str, float], +) -> Optional[float]: + remaining_budget: Optional[float] = None + for counter in counters: + current_spend = await _get_current_counter_value(counter=counter) + current_spend_by_counter_key[counter.counter_key] = current_spend + remaining = counter.max_budget - current_spend + if remaining <= 0: + raise litellm.BudgetExceededError( + current_cost=current_spend, + max_budget=counter.max_budget, + message=( + "Budget has been exceeded! " + f"{counter.entity_type}={counter.entity_id} " + f"Current cost: {current_spend}, " + f"Max budget: {counter.max_budget}" + ), + ) + remaining_budget = ( + remaining if remaining_budget is None else min(remaining_budget, remaining) + ) + return remaining_budget + + +async def _reserve_counter( + counter: _BudgetCounter, + reservation_cost: float, +) -> Optional[float]: + from litellm.proxy.proxy_server import ( + _ensure_spend_counter_initialized, + _ensure_window_spend_counter_initialized, + _invalidate_spend_counter, + _increment_spend_counter_cache, + ) + + attempted_increment = False + try: + if counter.source_cache_key is not None: + await _ensure_spend_counter_initialized( + counter_key=counter.counter_key, + source_cache_key=counter.source_cache_key, + ) + elif ( + counter.spend_log_entity_id is not None and counter.window_start is not None + ): + initialized = await _ensure_window_spend_counter_initialized( + counter_key=counter.counter_key, + entity_type=counter.entity_type, + entity_id=counter.spend_log_entity_id, + window_start=counter.window_start, + ) + if initialized is False: + verbose_proxy_logger.warning( + "Skipping budget reservation for %s because window spend could not be loaded", + counter.counter_key, + ) + raise _CounterReservationUnavailable + + attempted_increment = True + reserved_value = await _increment_spend_counter_cache( + counter_key=counter.counter_key, + increment=reservation_cost, + ) + return float(reserved_value) if reserved_value is not None else None + except _CounterReservationUnavailable: + raise + except Exception: + verbose_proxy_logger.warning( + "Skipping budget reservation for %s because spend counter reservation failed", + counter.counter_key, + exc_info=True, + ) + counter_invalidated = False + try: + await _invalidate_spend_counter(counter_key=counter.counter_key) + counter_invalidated = True + except Exception: + verbose_proxy_logger.warning( + "Failed to invalidate spend counter after budget reservation failure for %s", + counter.counter_key, + exc_info=True, + ) + raise _CounterReservationUnavailable( + touched_counter=attempted_increment, + counter_invalidated=counter_invalidated, + ) + + +async def _get_current_counter_value(counter: _BudgetCounter) -> float: + from litellm.proxy.proxy_server import get_current_spend + + return await get_current_spend( + counter_key=counter.counter_key, + fallback_spend=counter.fallback_spend, + ) + + +async def _set_reserved_entries_actual_cost( + entries: List[dict], + actual_cost: float, + default_reserved_cost: float, +) -> None: + for entry in entries: + await _set_reserved_entry_actual_cost( + entry=entry, + actual_cost=actual_cost, + default_reserved_cost=default_reserved_cost, + ) + + +async def _set_reserved_entry_actual_cost( + entry: dict, + actual_cost: float, + default_reserved_cost: float, +) -> None: + from litellm.proxy.proxy_server import _increment_spend_counter_cache + + counter_key = entry.get("counter_key") + if counter_key is None: + return + reserved_cost = _get_entry_reserved_cost( + entry=entry, + default_reserved_cost=default_reserved_cost, + ) + target_adjustment = actual_cost - reserved_cost + applied_adjustment = float(entry.get("applied_adjustment") or 0.0) + adjustment = target_adjustment - applied_adjustment + if adjustment == 0: + return + await _ensure_counter_can_apply_adjustment( + counter_key=counter_key, + adjustment=adjustment, + ) + await _increment_spend_counter_cache( + counter_key=counter_key, + increment=adjustment, + ) + entry["applied_adjustment"] = target_adjustment + + +async def _ensure_counter_can_apply_adjustment( + counter_key: str, + adjustment: float, +) -> None: + from litellm.proxy.proxy_server import ( + _invalidate_spend_counter, + spend_counter_cache, + ) + + current_value = await spend_counter_cache.async_get_cache(key=counter_key) + if current_value is None: + await _invalidate_spend_counter(counter_key=counter_key) + raise RuntimeError( + f"Cannot apply budget reservation adjustment to missing counter {counter_key}" + ) + + try: + current_float = float(current_value) + except (TypeError, ValueError): + await _invalidate_spend_counter(counter_key=counter_key) + raise RuntimeError( + f"Cannot apply budget reservation adjustment to non-numeric counter {counter_key}" + ) + + if adjustment < 0 and current_float + adjustment < -1e-12: + await _invalidate_spend_counter(counter_key=counter_key) + raise RuntimeError( + f"Budget reservation adjustment would make counter negative {counter_key}" + ) + + +async def _release_applied_entries_best_effort( + entries: List[dict], + default_reserved_cost: float, +) -> None: + for entry in entries: + try: + await _set_reserved_entry_actual_cost( + entry=entry, + actual_cost=0.0, + default_reserved_cost=default_reserved_cost, + ) + except Exception: + counter_key = entry.get("counter_key") + verbose_proxy_logger.exception( + "Failed to release partial budget reservation during exception cleanup" + ) + if counter_key is None: + continue + try: + from litellm.proxy.proxy_server import _invalidate_spend_counter + + await _invalidate_spend_counter(counter_key=counter_key) + except Exception: + verbose_proxy_logger.exception( + "Failed to invalidate partial budget reservation counter during exception cleanup" + ) + + +async def _resize_applied_reservation( + entries: List[dict], + current_reserved_cost: float, + new_reserved_cost: float, +) -> None: + await _set_reserved_entries_actual_cost( + entries=entries, + actual_cost=new_reserved_cost, + default_reserved_cost=current_reserved_cost, + ) + for entry in entries: + entry["reserved_cost"] = new_reserved_cost + entry["applied_adjustment"] = 0.0 + + +def _counter_to_reservation_entry( + counter: _BudgetCounter, + reserved_cost: float, +) -> Dict[str, Any]: + return { + "counter_key": counter.counter_key, + "entity_type": counter.entity_type, + "entity_id": counter.entity_id, + "reserved_cost": reserved_cost, + "applied_adjustment": 0.0, + } + + +def _get_entry_reserved_cost(entry: dict, default_reserved_cost: float) -> float: + try: + return float(entry.get("reserved_cost", default_reserved_cost) or 0.0) + except (TypeError, ValueError): + return default_reserved_cost + + +def get_budget_window_start(window: Any) -> Optional[datetime]: + window_dict = _coerce_window(window) + budget_duration = window_dict.get("budget_duration") + if budget_duration is None: + return None + try: + duration_seconds = duration_in_seconds(str(budget_duration)) + except Exception: + return None + + reset_at = _coerce_datetime(window_dict.get("reset_at")) + if reset_at is None: + return datetime.now(timezone.utc) - timedelta(seconds=duration_seconds) + if reset_at.tzinfo is None: + reset_at = reset_at.replace(tzinfo=timezone.utc) + return reset_at - timedelta(seconds=duration_seconds) + + +def _coerce_datetime(value: Any) -> Optional[datetime]: + if value is None: + return None + if isinstance(value, datetime): + return value + if isinstance(value, str): + try: + return datetime.fromisoformat(value.replace("Z", "+00:00")) + except ValueError: + return None + return None + + +def estimate_request_max_cost( + request_body: dict, + route: str, + llm_router: Optional[Router], +) -> Optional[float]: + model = get_model_from_request(request_body, route) + if model is None: + return None + + models = [model] if isinstance(model, str) else model + estimates = [ + _estimate_request_max_cost_for_model( + request_body=request_body, + route=route, + model=model_name, + llm_router=llm_router, + ) + for model_name in models + ] + estimates = [estimate for estimate in estimates if estimate is not None] + if not estimates: + return None + return max(cast(List[float], estimates)) + + +def _estimate_request_max_cost_for_model( + request_body: dict, + route: str, + model: str, + llm_router: Optional[Router], +) -> Optional[float]: + model_info = _get_model_cost_info(model=model, llm_router=llm_router) + if model_info is None: + return None + + input_cost_per_token = _to_float(model_info.get("input_cost_per_token")) + output_cost_per_token = _to_float(model_info.get("output_cost_per_token")) + input_tokens = _estimate_input_tokens( + request_body=request_body, + route=route, + model=model, + model_info=model_info, + ) + output_tokens = _estimate_output_tokens( + request_body=request_body, + route=route, + model_info=model_info, + ) + if input_tokens is None or output_tokens is None: + return None + + cost = 0.0 + if input_cost_per_token is not None: + cost += input_tokens * input_cost_per_token + elif input_tokens > 0: + return None + + output_multiplier = _get_output_multiplier(request_body=request_body) + if output_cost_per_token is not None: + cost += output_tokens * output_multiplier * output_cost_per_token + elif output_tokens > 0: + return None + + return cost + + +def _get_model_cost_info( + model: str, + llm_router: Optional[Router], +) -> Optional[Dict[str, Any]]: + if llm_router is not None: + try: + model_group_info = llm_router.get_model_group_info(model_group=model) + if model_group_info is not None: + return model_group_info.model_dump() + except Exception: + verbose_proxy_logger.debug( + "Unable to load router model group info for budget reservation", + exc_info=True, + ) + + try: + return dict(litellm.get_model_info(model=model)) + except Exception: + return None + + +def _estimate_input_tokens( + request_body: dict, + route: str, + model: str, + model_info: Dict[str, Any], +) -> Optional[int]: + try: + if "messages" in request_body: + return litellm.token_counter( + model=model, + messages=request_body.get("messages") or [], + tools=request_body.get("tools"), + tool_choice=request_body.get("tool_choice"), + ) + if "prompt" in request_body: + return _count_text_tokens(model=model, text=request_body.get("prompt")) + if "input" in request_body: + return _count_text_tokens(model=model, text=request_body.get("input")) + if "query" in request_body or "documents" in request_body: + query_tokens = _count_text_tokens( + model=model, text=request_body.get("query") + ) + document_tokens = _count_text_tokens( + model=model, + text=request_body.get("documents"), + ) + return query_tokens + document_tokens + except Exception: + verbose_proxy_logger.debug( + "Unable to count input tokens for budget reservation", exc_info=True + ) + + max_input_tokens = _to_int(model_info.get("max_input_tokens")) + if max_input_tokens is not None: + return max_input_tokens + + return None + + +def _estimate_output_tokens( + request_body: dict, + route: str, + model_info: Dict[str, Any], +) -> Optional[int]: + if _is_input_only_route(route=route): + return 0 + + for key in ("max_completion_tokens", "max_tokens", "max_output_tokens"): + max_tokens = _to_int(request_body.get(key)) + if max_tokens is not None: + return max_tokens + + # If the caller did not cap output tokens, avoid reserving a model's + # theoretical maximum context. The caller can still admit one request by + # reserving the smallest remaining budget in reserve_budget_for_request(). + return None + + +def _count_text_tokens(model: str, text: Any) -> int: + if text is None: + return 0 + + token_count = 0 + stack = [text] + while stack: + item = stack.pop() + if item is None: + continue + if isinstance(item, list): + stack.extend(item) + continue + if isinstance(item, dict): + token_count += litellm.token_counter(model=model, text=json.dumps(item)) + continue + token_count += litellm.token_counter(model=model, text=str(item)) + return token_count + + +def _get_output_multiplier(request_body: dict) -> int: + output_multiplier = 1 + for key in ("n", "best_of"): + value = _to_int(request_body.get(key)) + if value is not None: + output_multiplier = max(output_multiplier, value) + return output_multiplier + + +def _is_input_only_route(route: str) -> bool: + return any( + route_part in route + for route_part in ( + "embeddings", + "rerank", + "moderations", + ) + ) + + +def _to_float(value: Any) -> Optional[float]: + if value is None: + return None + try: + return float(value) + except (TypeError, ValueError): + return None + + +def _to_int(value: Any) -> Optional[int]: + if value is None: + return None + try: + return int(value) + except (TypeError, ValueError): + return None + + +def _get_value(obj: Any, key: str) -> Any: + if isinstance(obj, dict): + return obj.get(key) + return getattr(obj, key, None) diff --git a/litellm/proxy/spend_tracking/cloudzero_endpoints.py b/litellm/proxy/spend_tracking/cloudzero_endpoints.py index c7bff7ec642..1f551d5ffea 100644 --- a/litellm/proxy/spend_tracking/cloudzero_endpoints.py +++ b/litellm/proxy/spend_tracking/cloudzero_endpoints.py @@ -6,6 +6,7 @@ from litellm._logging import verbose_proxy_logger from litellm.litellm_core_utils.sensitive_data_masker import SensitiveDataMasker from litellm.proxy._types import CommonProxyErrors, LitellmUserRoles, UserAPIKeyAuth from litellm.proxy.auth.user_api_key_auth import user_api_key_auth +from litellm.proxy.management_endpoints.common_utils import _user_has_admin_view from litellm.proxy.common_utils.encrypt_decrypt_utils import ( decrypt_value_helper, encrypt_value_helper, @@ -127,10 +128,10 @@ async def get_cloudzero_settings( Only the first 4 and last 4 characters of the API key are shown. Returns null/empty values when settings are not configured (consistent with other settings endpoints). - Only admin users can view CloudZero settings. + Only admin users (Proxy Admin or Admin Viewer) can view CloudZero settings. """ - # Validation - if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN: + # Validation — Admin Viewer follows the read-parity rule. + if not _user_has_admin_view(user_api_key_dict): raise HTTPException( status_code=403, detail={"error": CommonProxyErrors.not_allowed_access.value}, diff --git a/litellm/proxy/spend_tracking/spend_log_error_logger.py b/litellm/proxy/spend_tracking/spend_log_error_logger.py new file mode 100644 index 00000000000..cfe647b6600 --- /dev/null +++ b/litellm/proxy/spend_tracking/spend_log_error_logger.py @@ -0,0 +1,85 @@ +""" +Logging helpers for spend-tracking error paths. + +Proxy operators have asked for a way to keep both their downstream log sinks +and the SpendLogs UI free of the stack traces that the spend-tracking +machinery emits when it hits 4xx/5xx or transient DB errors. The errors still +need to be logged (and still flow to Sentry via +``proxy_logging_obj.failure_handler``), but the multi-line stack traces +dominate log volume and clutter the per-row Metadata pane in the UI. + +The opt-in is a single env var, ``LITELLM_SUPPRESS_SPEND_LOG_TRACEBACKS=true``, +gated by ``should_suppress_spend_log_tracebacks``. When it returns ``True``: + * ``spend_log_error`` drops the traceback from the console / structured log + record (this module), and + * the failure callback in ``proxy_track_cost_callback`` drops the + ``error_information.traceback`` field from the SpendLogs row before it is + persisted, so the UI's per-row Metadata pane (which renders the metadata + JSON verbatim) stays clean. The key is omitted entirely rather than set + to ``""`` — ``StandardLoggingPayloadErrorInformation`` marks the field + optional and every downstream consumer uses ``.get("traceback")``. + +At DEBUG the full traceback is always preserved so operators can still +troubleshoot. The UI suppression follows the same gate. +""" + +import logging +import os +from typing import Any, Optional + +from litellm._logging import verbose_proxy_logger +from litellm.secret_managers.main import str_to_bool + +SUPPRESS_SPEND_LOG_TRACEBACKS_ENV = "LITELLM_SUPPRESS_SPEND_LOG_TRACEBACKS" + + +def _is_suppression_env_enabled() -> bool: + """Read the opt-in env var fresh each call so dynamic flips are honored. + + Kept separate from ``should_suppress_spend_log_tracebacks`` so tests and + other call sites can introspect just the env-var state without also + consulting the live logger level. + """ + return str_to_bool(os.getenv(SUPPRESS_SPEND_LOG_TRACEBACKS_ENV)) is True + + +def should_suppress_spend_log_tracebacks() -> bool: + """Return ``True`` when spend-log traceback suppression should apply. + + Suppression only kicks in when both: + * the operator opted in via the env var, and + * the proxy logger is at INFO or above (i.e. not DEBUG) — at DEBUG we + still want full tracebacks for troubleshooting. + """ + if not _is_suppression_env_enabled(): + return False + return not verbose_proxy_logger.isEnabledFor(logging.DEBUG) + + +def spend_log_error( + message: str, + *args: Any, + exc: Optional[BaseException] = None, +) -> None: + """Log a spend-tracking error, with the traceback gated on the env var. + + By default this behaves like ``verbose_proxy_logger.exception`` — the + active exception (or ``exc`` if supplied) is attached so the formatter + renders its traceback. When ``LITELLM_SUPPRESS_SPEND_LOG_TRACEBACKS`` is + truthy and the logger is at INFO or above, the traceback is dropped and + only ``message % args`` is emitted. + + Sentry / ``proxy_logging_obj.failure_handler`` is NOT invoked here — call + sites still own the alerting path. This helper is purely about console / + structured-log output volume. + """ + if should_suppress_spend_log_tracebacks(): + verbose_proxy_logger.error(message, *args) + return + + if exc is not None: + verbose_proxy_logger.error( + message, *args, exc_info=(type(exc), exc, exc.__traceback__) + ) + else: + verbose_proxy_logger.error(message, *args, exc_info=True) diff --git a/litellm/proxy/spend_tracking/spend_tracking_utils.py b/litellm/proxy/spend_tracking/spend_tracking_utils.py index ec6245f47e9..5421700ef58 100644 --- a/litellm/proxy/spend_tracking/spend_tracking_utils.py +++ b/litellm/proxy/spend_tracking/spend_tracking_utils.py @@ -25,6 +25,7 @@ from litellm.litellm_core_utils.core_helpers import ( ) from litellm.litellm_core_utils.safe_json_dumps import safe_dumps from litellm.proxy._types import SpendLogsMetadata, SpendLogsPayload +from litellm.proxy.spend_tracking.spend_log_error_logger import spend_log_error from litellm.proxy.utils import PrismaClient, hash_token from litellm.types.utils import ( CostBreakdown, @@ -471,9 +472,7 @@ def get_logging_payload( # noqa: PLR0915 return payload except Exception as e: - verbose_proxy_logger.exception( - "Error creating spendlogs object - {}".format(str(e)) - ) + spend_log_error("Error creating spendlogs object - %s", str(e), exc=e) raise e diff --git a/litellm/proxy/spend_tracking/vantage_endpoints.py b/litellm/proxy/spend_tracking/vantage_endpoints.py index 7d8fbf74615..60e54d005b3 100644 --- a/litellm/proxy/spend_tracking/vantage_endpoints.py +++ b/litellm/proxy/spend_tracking/vantage_endpoints.py @@ -7,6 +7,7 @@ from litellm._logging import verbose_proxy_logger from litellm.litellm_core_utils.sensitive_data_masker import SensitiveDataMasker from litellm.proxy._types import CommonProxyErrors, LitellmUserRoles, UserAPIKeyAuth from litellm.proxy.auth.user_api_key_auth import user_api_key_auth +from litellm.proxy.management_endpoints.common_utils import _user_has_admin_view from litellm.proxy.common_utils.encrypt_decrypt_utils import ( decrypt_value_helper, encrypt_value_helper, @@ -140,9 +141,10 @@ async def get_vantage_settings( View current Vantage settings. Returns the current Vantage configuration with the API key masked for security. - Only admin users can view Vantage settings. + Only admin users (Proxy Admin or Admin Viewer) can view Vantage settings. """ - if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN: + # Admin Viewer follows the read-parity rule. + if not _user_has_admin_view(user_api_key_dict): raise HTTPException( status_code=403, detail={"error": CommonProxyErrors.not_allowed_access.value}, diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index 8c5fce84099..e7f5f4ee396 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -36,6 +36,7 @@ from litellm.proxy._types import ( SpendLogsMetadata, SpendLogsPayload, ) +from litellm.proxy.spend_tracking.spend_log_error_logger import spend_log_error from litellm.types.guardrails import GuardrailEventHooks from litellm.types.utils import CallTypes, CallTypesLiteral @@ -3188,6 +3189,8 @@ class PrismaClient: t.organization_id as org_id, p.project_alias AS project_alias, tm.spend AS team_member_spend, + b_tm.tpm_limit AS team_member_tpm_limit, + b_tm.rpm_limit AS team_member_rpm_limit, m.aliases AS team_model_aliases, -- Added comma to separate b.* columns b.max_budget AS litellm_budget_table_max_budget, @@ -3203,6 +3206,7 @@ class PrismaClient: FROM "LiteLLM_VerificationToken" AS v LEFT JOIN "LiteLLM_TeamTable" AS t ON v.team_id = t.team_id LEFT JOIN "LiteLLM_TeamMembership" AS tm ON v.team_id = tm.team_id AND tm.user_id = v.user_id + LEFT JOIN "LiteLLM_BudgetTable" AS b_tm ON tm.budget_id = b_tm.budget_id LEFT JOIN "LiteLLM_ModelTable" m ON t.model_id = m.id LEFT JOIN "LiteLLM_BudgetTable" AS b ON v.budget_id = b.budget_id LEFT JOIN "LiteLLM_ProjectTable" AS p ON v.project_id = p.project_id @@ -5103,6 +5107,11 @@ async def update_daily_tag_spend( proxy_logging_obj=proxy_logging_obj, ) except Exception as e: + # NOTE: keep this as a plain ``error`` (no traceback) to match the + # historical behavior of this site. ``spend_log_error`` would attach + # the active exception's traceback whenever the suppression env var + # is unset, which would be a regression for operators who never saw + # one here before. verbose_proxy_logger.error(f"Error updating daily tag spend: {e}") @@ -5235,9 +5244,7 @@ async def _monitor_spend_logs_queue( await asyncio.sleep(current_interval) except Exception as e: - verbose_proxy_logger.error( - f"Error in spend logs queue monitor: {str(e)}\n{traceback.format_exc()}" - ) + spend_log_error("Error in spend logs queue monitor: %s", str(e), exc=e) # Continue monitoring even if there's an error, with exponential backoff current_interval = min(current_interval * backoff_multiplier, max_backoff) await asyncio.sleep(current_interval) diff --git a/litellm/proxy/vector_store_endpoints/endpoints.py b/litellm/proxy/vector_store_endpoints/endpoints.py index 1fdfad8c96c..86e316e7f40 100644 --- a/litellm/proxy/vector_store_endpoints/endpoints.py +++ b/litellm/proxy/vector_store_endpoints/endpoints.py @@ -1,8 +1,6 @@ from typing import Any, Dict, Optional from fastapi import APIRouter, Depends, HTTPException, Request, Response - -import litellm from litellm.integrations.vector_store_integrations.vector_store_pre_call_hook import ( LiteLLM_ManagedVectorStore, ) @@ -10,7 +8,10 @@ from litellm.proxy._types import CommonProxyErrors, UserAPIKeyAuth from litellm.proxy.auth.user_api_key_auth import user_api_key_auth from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing from litellm.proxy.utils import jsonify_object -from litellm.proxy.vector_store_endpoints.utils import can_user_access_vector_store +from litellm.proxy.vector_store_endpoints.utils import ( + assert_user_can_access_vector_store, + get_litellm_managed_vector_store, +) from litellm.types.vector_stores import IndexCreateRequest router = APIRouter() @@ -19,24 +20,6 @@ router = APIRouter() ######################################################## -async def _check_vector_store_access( - vector_store: LiteLLM_ManagedVectorStore, - user_api_key_dict: UserAPIKeyAuth, -) -> bool: - """ - Check if the user has access to the vector store. - - Delegates to :func:`can_user_access_vector_store`, which honors: - - PROXY_ADMIN bypass - - legacy vector stores with no team_id - - key-level and team-level ``object_permission.vector_stores`` allowlists - - team_id match between key and store - """ - return await can_user_access_vector_store( - vector_store=vector_store, user_api_key_dict=user_api_key_dict - ) - - async def _update_request_data_with_litellm_managed_vector_store_registry( data: Dict, vector_store_id: str, @@ -53,35 +36,27 @@ async def _update_request_data_with_litellm_managed_vector_store_registry( Raises: HTTPException: If user doesn't have access to the vector store """ - if litellm.vector_store_registry is not None: - vector_store_to_run: Optional[LiteLLM_ManagedVectorStore] = ( - litellm.vector_store_registry.get_litellm_managed_vector_store_from_registry( - vector_store_id=vector_store_id + vector_store_to_run: Optional[LiteLLM_ManagedVectorStore] = ( + await get_litellm_managed_vector_store(vector_store_id=vector_store_id) + ) + if vector_store_to_run is not None: + if user_api_key_dict is not None: + await assert_user_can_access_vector_store( + vector_store=vector_store_to_run, + user_api_key_dict=user_api_key_dict, ) - ) - if vector_store_to_run is not None: - if user_api_key_dict is not None: - if not await _check_vector_store_access( - vector_store_to_run, user_api_key_dict - ): - raise HTTPException( - status_code=403, - detail="Access denied: You do not have permission to access this vector store", - ) - if "custom_llm_provider" in vector_store_to_run: - data["custom_llm_provider"] = vector_store_to_run.get( - "custom_llm_provider" - ) + if "custom_llm_provider" in vector_store_to_run: + data["custom_llm_provider"] = vector_store_to_run.get("custom_llm_provider") - if "litellm_credential_name" in vector_store_to_run: - data["litellm_credential_name"] = vector_store_to_run.get( - "litellm_credential_name" - ) + if "litellm_credential_name" in vector_store_to_run: + data["litellm_credential_name"] = vector_store_to_run.get( + "litellm_credential_name" + ) - if "litellm_params" in vector_store_to_run: - litellm_params = vector_store_to_run.get("litellm_params", {}) or {} - data.update(litellm_params) + if "litellm_params" in vector_store_to_run: + litellm_params = vector_store_to_run.get("litellm_params", {}) or {} + data.update(litellm_params) return data @@ -121,8 +96,7 @@ async def vector_store_search( ) data = await _read_request_body(request=request) - if "vector_store_id" not in data: - data["vector_store_id"] = vector_store_id + data["vector_store_id"] = vector_store_id # Check for legacy vector store registry (non-managed vector stores) data = await _update_request_data_with_litellm_managed_vector_store_registry( diff --git a/litellm/proxy/vector_store_endpoints/utils.py b/litellm/proxy/vector_store_endpoints/utils.py index 061a8aaa240..657b520b271 100644 --- a/litellm/proxy/vector_store_endpoints/utils.py +++ b/litellm/proxy/vector_store_endpoints/utils.py @@ -1,7 +1,9 @@ +import json from typing import Any, Dict, Literal, Optional from fastapi import HTTPException, Request +import litellm from litellm._logging import verbose_proxy_logger from litellm.proxy._types import ( LiteLLM_ObjectPermissionTable, @@ -13,6 +15,21 @@ from litellm.types.vector_stores import LiteLLM_ManagedVectorStore from litellm.utils import ProviderConfigManager +def _normalize_litellm_params( + vector_store: LiteLLM_ManagedVectorStore, +) -> LiteLLM_ManagedVectorStore: + litellm_params = vector_store.get("litellm_params") + if isinstance(litellm_params, str): + normalized = LiteLLM_ManagedVectorStore(**dict(vector_store)) + try: + parsed = json.loads(litellm_params) + normalized["litellm_params"] = parsed if isinstance(parsed, dict) else {} + except (TypeError, ValueError): + normalized["litellm_params"] = {} + return normalized + return vector_store + + def _is_proxy_admin(user_api_key_dict: UserAPIKeyAuth) -> bool: return ( user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN @@ -120,6 +137,104 @@ async def can_user_access_vector_store( return False +async def get_litellm_managed_vector_store( + vector_store_id: str, +) -> Optional[LiteLLM_ManagedVectorStore]: + """ + Resolve a LiteLLM-managed vector store from the registry or shared cache. + + Provider-native vector store IDs will not be present in either location and + return None, preserving direct provider behavior while still protecting + LiteLLM-managed multi-tenant stores. + """ + if not vector_store_id: + return None + + if litellm.vector_store_registry is not None: + try: + vector_store = litellm.vector_store_registry.get_litellm_managed_vector_store_from_registry( + vector_store_id=vector_store_id + ) + if vector_store is not None: + return _normalize_litellm_params(vector_store) + except Exception as e: + verbose_proxy_logger.warning( + "Failed to resolve vector store id=%s from registry: %s", + vector_store_id, + e, + ) + raise HTTPException( + status_code=500, + detail="Unable to validate vector store access", + ) from e + + try: + from litellm.proxy.auth.auth_checks import ( + get_managed_vector_store_rows_by_uuids, + ) + from litellm.proxy.proxy_server import ( + prisma_client, + proxy_logging_obj, + user_api_key_cache, + ) + + if prisma_client is None: + return None + rows = await get_managed_vector_store_rows_by_uuids( + uuids=[vector_store_id], + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + proxy_logging_obj=proxy_logging_obj, + ) + if not rows: + return None + return _normalize_litellm_params( + LiteLLM_ManagedVectorStore(**rows[0].model_dump()) + ) + except Exception as e: + verbose_proxy_logger.warning( + "Failed to resolve vector store id=%s from shared cache: %s", + vector_store_id, + e, + ) + raise HTTPException( + status_code=500, + detail="Unable to validate vector store access", + ) from e + + +async def assert_user_can_access_vector_store( + vector_store: LiteLLM_ManagedVectorStore, + user_api_key_dict: UserAPIKeyAuth, + detail: str = "Access denied: You do not have permission to access this vector store", +) -> None: + """Raise 403 unless the caller can access the resolved vector store.""" + if not await can_user_access_vector_store(vector_store, user_api_key_dict): + raise HTTPException(status_code=403, detail=detail) + + +async def assert_user_can_access_vector_store_id( + vector_store_id: str, + user_api_key_dict: UserAPIKeyAuth, + detail: str = "Access denied: You do not have permission to access this vector store", +) -> Optional[LiteLLM_ManagedVectorStore]: + """ + Resolve a managed vector store id and enforce ownership if it exists. + + Unknown ids are treated as provider-native ids and are not rejected here. + """ + vector_store = await get_litellm_managed_vector_store( + vector_store_id=vector_store_id + ) + if vector_store is not None: + await assert_user_can_access_vector_store( + vector_store=vector_store, + user_api_key_dict=user_api_key_dict, + detail=detail, + ) + return vector_store + + def _does_endpoint_match(endpoint_path: str, request_path: str) -> bool: if endpoint_path in request_path: return True diff --git a/litellm/proxy/vector_store_files_endpoints/endpoints.py b/litellm/proxy/vector_store_files_endpoints/endpoints.py index 7cdf865692b..346a847c5dd 100644 --- a/litellm/proxy/vector_store_files_endpoints/endpoints.py +++ b/litellm/proxy/vector_store_files_endpoints/endpoints.py @@ -17,9 +17,11 @@ from litellm.proxy.openai_files_endpoints.common_utils import ( prepare_data_with_credentials, ) from litellm.proxy.vector_store_endpoints.utils import ( + assert_user_can_access_vector_store_id, is_allowed_to_call_vector_store_files_endpoint, ) from litellm.types.utils import LlmProviders +from litellm.types.vector_stores import LiteLLM_ManagedVectorStore if TYPE_CHECKING: from litellm.router import Router @@ -193,6 +195,8 @@ def _update_request_data_with_litellm_managed_vector_store_registry( data: Dict, vector_store_id: str, llm_router: Optional["Router"] = None, + managed_vector_store: Optional[LiteLLM_ManagedVectorStore] = None, + should_lookup_registry: bool = True, ) -> Dict: """ Update request data with model routing information from managed vector store. @@ -262,23 +266,27 @@ def _update_request_data_with_litellm_managed_vector_store_registry( return data - # Legacy path: Check vector store registry for non-managed vector stores - if litellm.vector_store_registry is not None: + # Legacy path: Check vector store registry for non-managed vector stores. + vector_store_to_run = managed_vector_store + if ( + vector_store_to_run is None + and should_lookup_registry + and litellm.vector_store_registry is not None + ): vector_store_to_run = litellm.vector_store_registry.get_litellm_managed_vector_store_from_registry( vector_store_id=vector_store_id ) - if vector_store_to_run is not None: - if "custom_llm_provider" in vector_store_to_run: - data["custom_llm_provider"] = vector_store_to_run.get( - "custom_llm_provider" - ) - if "litellm_credential_name" in vector_store_to_run: - data["litellm_credential_name"] = vector_store_to_run.get( - "litellm_credential_name" - ) - if "litellm_params" in vector_store_to_run: - litellm_params = vector_store_to_run.get("litellm_params", {}) or {} - data.update(litellm_params) + + if vector_store_to_run is not None: + if "custom_llm_provider" in vector_store_to_run: + data["custom_llm_provider"] = vector_store_to_run.get("custom_llm_provider") + if "litellm_credential_name" in vector_store_to_run: + data["litellm_credential_name"] = vector_store_to_run.get( + "litellm_credential_name" + ) + if "litellm_params" in vector_store_to_run: + litellm_params = vector_store_to_run.get("litellm_params", {}) or {} + data.update(litellm_params) return data @@ -363,8 +371,11 @@ async def vector_store_file_create( ) data = await _read_request_body(request=request) - if "vector_store_id" not in data: - data["vector_store_id"] = vector_store_id + data["vector_store_id"] = vector_store_id + managed_vector_store = await assert_user_can_access_vector_store_id( + vector_store_id=vector_store_id, + user_api_key_dict=user_api_key_dict, + ) # Handle managed file IDs if present in request body original_managed_file_id = None @@ -375,7 +386,11 @@ async def vector_store_file_create( # Then handle managed vector store IDs data = _update_request_data_with_litellm_managed_vector_store_registry( - data=data, vector_store_id=vector_store_id, llm_router=llm_router + data=data, + vector_store_id=vector_store_id, + llm_router=llm_router, + managed_vector_store=managed_vector_store, + should_lookup_registry=False, ) provider_enum = await _resolve_provider(data=data, request=request) @@ -459,9 +474,18 @@ async def vector_store_file_list( query_params = dict(request.query_params) data: Dict[str, Optional[str]] = {"vector_store_id": vector_store_id} data.update(query_params) + data["vector_store_id"] = vector_store_id + managed_vector_store = await assert_user_can_access_vector_store_id( + vector_store_id=vector_store_id, + user_api_key_dict=user_api_key_dict, + ) data = _update_request_data_with_litellm_managed_vector_store_registry( - data=data, vector_store_id=vector_store_id, llm_router=llm_router + data=data, + vector_store_id=vector_store_id, + llm_router=llm_router, + managed_vector_store=managed_vector_store, + should_lookup_registry=False, ) provider_enum = await _resolve_provider(data=data, request=request) @@ -541,6 +565,10 @@ async def vector_store_file_retrieve( "vector_store_id": vector_store_id, "file_id": file_id, } + managed_vector_store = await assert_user_can_access_vector_store_id( + vector_store_id=vector_store_id, + user_api_key_dict=user_api_key_dict, + ) # Handle managed file IDs first data, original_managed_file_id = _update_request_data_with_managed_file_id( @@ -549,7 +577,11 @@ async def vector_store_file_retrieve( # Then handle managed vector store IDs data = _update_request_data_with_litellm_managed_vector_store_registry( - data=data, vector_store_id=vector_store_id, llm_router=llm_router + data=data, + vector_store_id=vector_store_id, + llm_router=llm_router, + managed_vector_store=managed_vector_store, + should_lookup_registry=False, ) provider_enum = await _resolve_provider(data=data, request=request) @@ -635,6 +667,10 @@ async def vector_store_file_content( "vector_store_id": vector_store_id, "file_id": file_id, } + managed_vector_store = await assert_user_can_access_vector_store_id( + vector_store_id=vector_store_id, + user_api_key_dict=user_api_key_dict, + ) # Handle managed file IDs first data, original_managed_file_id = _update_request_data_with_managed_file_id( @@ -643,7 +679,11 @@ async def vector_store_file_content( # Then handle managed vector store IDs data = _update_request_data_with_litellm_managed_vector_store_registry( - data=data, vector_store_id=vector_store_id, llm_router=llm_router + data=data, + vector_store_id=vector_store_id, + llm_router=llm_router, + managed_vector_store=managed_vector_store, + should_lookup_registry=False, ) provider_enum = await _resolve_provider(data=data, request=request) @@ -729,6 +769,10 @@ async def vector_store_file_update( data = await _read_request_body(request=request) data["vector_store_id"] = vector_store_id data["file_id"] = file_id + managed_vector_store = await assert_user_can_access_vector_store_id( + vector_store_id=vector_store_id, + user_api_key_dict=user_api_key_dict, + ) # Handle managed file IDs first data, original_managed_file_id = _update_request_data_with_managed_file_id( @@ -737,7 +781,11 @@ async def vector_store_file_update( # Then handle managed vector store IDs data = _update_request_data_with_litellm_managed_vector_store_registry( - data=data, vector_store_id=vector_store_id, llm_router=llm_router + data=data, + vector_store_id=vector_store_id, + llm_router=llm_router, + managed_vector_store=managed_vector_store, + should_lookup_registry=False, ) provider_enum = await _resolve_provider(data=data, request=request) @@ -823,6 +871,10 @@ async def vector_store_file_delete( "vector_store_id": vector_store_id, "file_id": file_id, } + managed_vector_store = await assert_user_can_access_vector_store_id( + vector_store_id=vector_store_id, + user_api_key_dict=user_api_key_dict, + ) # Handle managed file IDs first data, original_managed_file_id = _update_request_data_with_managed_file_id( @@ -831,7 +883,11 @@ async def vector_store_file_delete( # Then handle managed vector store IDs data = _update_request_data_with_litellm_managed_vector_store_registry( - data=data, vector_store_id=vector_store_id, llm_router=llm_router + data=data, + vector_store_id=vector_store_id, + llm_router=llm_router, + managed_vector_store=managed_vector_store, + should_lookup_registry=False, ) provider_enum = await _resolve_provider(data=data, request=request) diff --git a/litellm/proxy/vertex_ai_endpoints/langfuse_endpoints.py b/litellm/proxy/vertex_ai_endpoints/langfuse_endpoints.py index b6454bf077b..8ce1bedcf90 100644 --- a/litellm/proxy/vertex_ai_endpoints/langfuse_endpoints.py +++ b/litellm/proxy/vertex_ai_endpoints/langfuse_endpoints.py @@ -12,11 +12,13 @@ import base64 import os from base64 import b64encode from typing import Optional +from urllib.parse import unquote import httpx -from fastapi import APIRouter, Request, Response +from fastapi import APIRouter, HTTPException, Request, Response, status import litellm +from litellm.litellm_core_utils.url_utils import SSRFError, validate_url from litellm.proxy._types import * from litellm.proxy.auth.user_api_key_auth import user_api_key_auth from litellm.proxy.common_utils.http_parsing_utils import _safe_get_request_headers @@ -27,6 +29,7 @@ from litellm.proxy.pass_through_endpoints.pass_through_endpoints import ( router = APIRouter() default_vertex_config = None +_DEFAULT_LANGFUSE_HOST = "https://cloud.langfuse.com" def create_request_copy(request: Request): @@ -39,6 +42,116 @@ def create_request_copy(request: Request): } +def _decode_to_convergence(value: str) -> str: + previous = value + while True: + decoded = unquote(previous) + if decoded == previous: + return decoded + previous = decoded + + +def _normalize_langfuse_base_url(base_target_url: str) -> str: + if not ( + base_target_url.startswith("http://") or base_target_url.startswith("https://") + ): + # Existing behavior allows host-only Langfuse settings. + base_target_url = "http://" + base_target_url + + try: + base_url = httpx.URL(base_target_url) + except Exception as e: + raise HTTPException( + status_code=status.HTTP_400_BAD_REQUEST, + detail={"error": f"Invalid Langfuse host: {str(e)}"}, + ) + + if base_url.scheme not in ("http", "https") or not base_url.host: + raise HTTPException( + status_code=status.HTTP_400_BAD_REQUEST, + detail={"error": "Invalid Langfuse host"}, + ) + + if base_url.userinfo: + raise HTTPException( + status_code=status.HTTP_400_BAD_REQUEST, + detail={"error": "Langfuse host must not include credentials"}, + ) + + return str(base_url) + + +def _validate_langfuse_proxy_path(endpoint: str) -> str: + decoded_endpoint = _decode_to_convergence(endpoint) + if any(ord(char) < 32 for char in decoded_endpoint): + raise HTTPException( + status_code=status.HTTP_400_BAD_REQUEST, + detail={"error": "Invalid Langfuse endpoint path"}, + ) + if "\\" in decoded_endpoint or decoded_endpoint.startswith("//"): + raise HTTPException( + status_code=status.HTTP_400_BAD_REQUEST, + detail={"error": "Invalid Langfuse endpoint path"}, + ) + + endpoint_path = "/" + decoded_endpoint.lstrip("/") + if any(segment in (".", "..") for segment in endpoint_path.split("/")): + raise HTTPException( + status_code=status.HTTP_400_BAD_REQUEST, + detail={"error": "Invalid Langfuse endpoint path"}, + ) + return endpoint_path + + +def _get_langfuse_proxy_credentials( + *, + dynamic_host_supplied: bool, + dynamic_langfuse_public_key: Optional[str], + dynamic_langfuse_secret_key: Optional[str], +): + if dynamic_host_supplied: + if not dynamic_langfuse_public_key or not dynamic_langfuse_secret_key: + raise HTTPException( + status_code=status.HTTP_400_BAD_REQUEST, + detail={ + "error": "Dynamic Langfuse hosts must include dynamic Langfuse credentials" + }, + ) + return dynamic_langfuse_public_key, dynamic_langfuse_secret_key + + return ( + dynamic_langfuse_public_key + or litellm.utils.get_secret(secret_name="LANGFUSE_PUBLIC_KEY"), + dynamic_langfuse_secret_key + or litellm.utils.get_secret(secret_name="LANGFUSE_SECRET_KEY"), + ) + + +def _build_langfuse_proxy_target( + *, + endpoint: str, + base_target_url: str, + dynamic_host_supplied: bool, +): + endpoint_path = _validate_langfuse_proxy_path(endpoint) + base_url = httpx.URL(_normalize_langfuse_base_url(base_target_url)) + updated_url = base_url.copy_with(path=endpoint_path) + custom_headers = {} + + if dynamic_host_supplied and getattr(litellm, "user_url_validation", True): + try: + target_url, host_header = validate_url(str(updated_url)) + except SSRFError as e: + raise HTTPException( + status_code=status.HTTP_400_BAD_REQUEST, + detail={"error": f"Invalid Langfuse host: {str(e)}"}, + ) + custom_headers["Host"] = host_header + return target_url, custom_headers + + return str(updated_url), custom_headers + + @router.api_route( "/langfuse/{endpoint:path}", methods=["GET", "POST", "PUT", "DELETE", "PATCH"], @@ -91,44 +204,33 @@ async def langfuse_proxy_route( elif k == "langfuse_host": dynamic_langfuse_host = v + dynamic_host_supplied = dynamic_langfuse_host is not None base_target_url: str = ( dynamic_langfuse_host - or os.getenv("LANGFUSE_HOST", "https://cloud.langfuse.com") - or "https://cloud.langfuse.com" + or os.getenv("LANGFUSE_HOST", _DEFAULT_LANGFUSE_HOST) + or _DEFAULT_LANGFUSE_HOST ) - if not ( - base_target_url.startswith("http://") or base_target_url.startswith("https://") - ): - # add http:// if unset, assume communicating over private network - e.g. render - base_target_url = "http://" + base_target_url - - encoded_endpoint = httpx.URL(endpoint).path - - # Ensure endpoint starts with '/' for proper URL construction - if not encoded_endpoint.startswith("/"): - encoded_endpoint = "/" + encoded_endpoint - - # Construct the full target URL using httpx - base_url = httpx.URL(base_target_url) - updated_url = base_url.copy_with(path=encoded_endpoint) - - # Add or update query parameters - langfuse_public_key = dynamic_langfuse_public_key or litellm.utils.get_secret( - secret_name="LANGFUSE_PUBLIC_KEY" + langfuse_public_key, langfuse_secret_key = _get_langfuse_proxy_credentials( + dynamic_host_supplied=dynamic_host_supplied, + dynamic_langfuse_public_key=dynamic_langfuse_public_key, + dynamic_langfuse_secret_key=dynamic_langfuse_secret_key, ) - langfuse_secret_key = dynamic_langfuse_secret_key or litellm.utils.get_secret( - secret_name="LANGFUSE_SECRET_KEY" + target_url, target_headers = _build_langfuse_proxy_target( + endpoint=endpoint, + base_target_url=base_target_url, + dynamic_host_supplied=dynamic_host_supplied, ) langfuse_combined_key = "Basic " + b64encode( f"{langfuse_public_key}:{langfuse_secret_key}".encode("utf-8") ).decode("ascii") + target_headers["Authorization"] = langfuse_combined_key ## CREATE PASS-THROUGH endpoint_func = create_pass_through_route( endpoint=endpoint, - target=str(updated_url), - custom_headers={"Authorization": langfuse_combined_key}, + target=target_url, + custom_headers=target_headers, query_params=dict(request.query_params), # type: ignore ) # dynamically construct pass-through endpoint based on incoming path received_value = await endpoint_func( diff --git a/litellm/router.py b/litellm/router.py index 50fd7eaed0b..8e29f8cfc11 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -9351,6 +9351,52 @@ class Router: """ return [m for m in self.model_list if m["litellm_params"]["model"] == model] + def _try_early_resolve_deployments_for_model_not_in_names( + self, model: str, request_team_id: Optional[str] + ) -> Optional[Tuple[str, Union[List, Dict]]]: + """ + When ``model`` is not in ``self.model_names``, try team routes, pattern routes, + team pattern routers, then default deployment. Returns None if none apply. + """ + if model in self.model_names: + return None + # Check for team-specific deployments by team_public_model_name. + # This intentionally takes priority over team pattern routers below, + # so that named team deployments shadow wildcard/pattern routes. + if request_team_id is not None: + team_deployments = self._get_all_deployments( + model_name=model, team_id=request_team_id + ) + if team_deployments: + return model, team_deployments + + pattern_deployments = self.pattern_router.get_deployments_by_pattern( + model=model, + ) + + if pattern_deployments: + return model, pattern_deployments + + if request_team_id is not None and request_team_id in self.team_pattern_routers: + pattern_deployments = self.team_pattern_routers[ + request_team_id + ].get_deployments_by_pattern( + model=model, + ) + if pattern_deployments: + return model, pattern_deployments + + if self.default_deployment is not None: + # Shallow copy with nested litellm_params copy (100x+ faster than deepcopy) + updated_deployment = self.default_deployment.copy() + updated_deployment["litellm_params"] = self.default_deployment[ + "litellm_params" + ].copy() + updated_deployment["litellm_params"]["model"] = model + return model, updated_deployment + + return None + def _common_checks_available_deployment( self, model: str, @@ -9393,56 +9439,52 @@ class Router: if _model_from_alias is not None: model = _model_from_alias - if model not in self.model_names: - # Check for team-specific deployments by team_public_model_name. - # This intentionally takes priority over team pattern routers below, - # so that named team deployments shadow wildcard/pattern routes. - if request_team_id is not None: - team_deployments = self._get_all_deployments( - model_name=model, team_id=request_team_id - ) - if team_deployments: - return model, team_deployments - - # check if provider/ specific wildcard routing use pattern matching - pattern_deployments = self.pattern_router.get_deployments_by_pattern( - model=model, - ) - - if pattern_deployments: - return model, pattern_deployments - - if ( - request_team_id is not None - and request_team_id in self.team_pattern_routers - ): - pattern_deployments = self.team_pattern_routers[ - request_team_id - ].get_deployments_by_pattern( - model=model, - ) - if pattern_deployments: - return model, pattern_deployments - - # check if default deployment is set - if self.default_deployment is not None: - # Shallow copy with nested litellm_params copy (100x+ faster than deepcopy) - updated_deployment = self.default_deployment.copy() - updated_deployment["litellm_params"] = self.default_deployment[ - "litellm_params" - ].copy() - updated_deployment["litellm_params"]["model"] = model - return model, updated_deployment + early = self._try_early_resolve_deployments_for_model_not_in_names( + model=model, request_team_id=request_team_id + ) + if early is not None: + return early ## get healthy deployments ### get all deployments healthy_deployments = self._get_all_deployments( model_name=model, team_id=request_team_id ) + _pre_model_access_group_filter_len = len(healthy_deployments) + healthy_deployments = self._filter_deployments_by_model_access_groups( + model=model, + healthy_deployments=healthy_deployments, + request_kwargs=request_kwargs, + request_team_id=request_team_id, + ) + _access_group_filter_emptied_candidates = ( + _pre_model_access_group_filter_len > 0 and len(healthy_deployments) == 0 + ) if len(healthy_deployments) == 0: # check if the user sent in a deployment name instead - healthy_deployments = self._get_deployment_by_litellm_model(model=model) + # Do not fall back when access-group filtering removed every candidate; + # _get_deployment_by_litellm_model does not re-apply that filter. + if _pre_model_access_group_filter_len == 0: + _litellm_model_deployments = self._get_deployment_by_litellm_model( + model=model + ) + healthy_deployments = self._filter_deployments_by_model_access_groups( + model=model, + healthy_deployments=_litellm_model_deployments, + request_kwargs=request_kwargs, + request_team_id=request_team_id, + ) + # If the litellm-model lookup produced candidates that access-group + # filtering then removed, treat this the same as the by-name path + # being emptied: prevent default-model fallback from bypassing the + # restriction (the fallback model may have no access_groups and + # would short-circuit the filter). + if ( + len(_litellm_model_deployments) > 0 + and len(healthy_deployments) == 0 + ): + _access_group_filter_emptied_candidates = True if verbose_router_logger.isEnabledFor(logging.DEBUG): verbose_router_logger.debug( @@ -9451,7 +9493,13 @@ class Router: if len(healthy_deployments) == 0: # Check for default fallbacks if no deployments are found for the requested model - if self._has_default_fallbacks(): + # Do not fall back to another model when access-group filtering removed every + # candidate for the requested name: re-filtering the fallback model can be a + # no-op when it has no access_groups, incorrectly serving a different model. + if ( + self._has_default_fallbacks() + and not _access_group_filter_emptied_candidates + ): fallback_model = self._get_first_default_fallback() if fallback_model: verbose_router_logger.info( @@ -9462,6 +9510,14 @@ class Router: healthy_deployments = self._get_all_deployments( model_name=model, team_id=request_team_id ) + healthy_deployments = ( + self._filter_deployments_by_model_access_groups( + model=model, + healthy_deployments=healthy_deployments, + request_kwargs=request_kwargs, + request_team_id=request_team_id, + ) + ) # If still no deployments after checking for fallbacks, raise an error if len(healthy_deployments) == 0: @@ -9487,6 +9543,70 @@ class Router: return model, healthy_deployments + def _filter_deployments_by_model_access_groups( + self, + model: str, + healthy_deployments: List, + request_kwargs: Optional[Dict], + request_team_id: Optional[str], + ) -> List: + """ + Restrict candidate deployments to caller-authorized model access groups. + + This is only applied when: + - request metadata includes `user_api_key_auth`, and + - caller permissions for this model are access-group-only + (no explicit model, wildcard, or all-proxy grants). + """ + if not healthy_deployments or request_kwargs is None: + return healthy_deployments + + metadata = request_kwargs.get("metadata") or {} + litellm_metadata = request_kwargs.get("litellm_metadata") or {} + user_api_key_auth = metadata.get("user_api_key_auth") or litellm_metadata.get( + "user_api_key_auth" + ) + if user_api_key_auth is None: + return healthy_deployments + + object_models = set(getattr(user_api_key_auth, "models", []) or []) + object_team_models = set(getattr(user_api_key_auth, "team_models", []) or []) + allowed_models = object_models | object_team_models + if not allowed_models: + return healthy_deployments + + # If caller has direct model/wildcard/all-proxy access, do not constrain + # deployment choice by access group. + if ( + model in allowed_models + or "*" in allowed_models + or "all-proxy-models" in allowed_models + ): + return healthy_deployments + + access_groups_for_model = self.get_model_access_groups( + model_name=model, team_id=request_team_id + ) + if len(access_groups_for_model) == 0: + return healthy_deployments + + allowed_access_groups = set(access_groups_for_model.keys()) & allowed_models + if not allowed_access_groups: + # No overlap means this request was not authorized via model access + # group membership for this model, so do not force group filtering. + return healthy_deployments + + filtered_deployments = [] + for deployment in healthy_deployments: + deployment_model_info = deployment.get("model_info") or {} + deployment_access_groups = set( + deployment_model_info.get("access_groups", []) or [] + ) + if deployment_access_groups & allowed_access_groups: + filtered_deployments.append(deployment) + + return filtered_deployments + async def async_get_healthy_deployments( self, model: str, @@ -10007,6 +10127,7 @@ class Router: messages=messages, input=input, specific_deployment=specific_deployment, + request_kwargs=request_kwargs, ) if isinstance(healthy_deployments, dict): diff --git a/litellm/types/guardrails.py b/litellm/types/guardrails.py index a98f9d666ae..04347aebe3b 100644 --- a/litellm/types/guardrails.py +++ b/litellm/types/guardrails.py @@ -38,6 +38,9 @@ from litellm.types.proxy.guardrails.guardrail_hooks.tool_permission import ( from litellm.types.proxy.guardrails.guardrail_hooks.hiddenlayer import ( HiddenlayerGuardrailConfigModel, ) +from litellm.types.proxy.guardrails.guardrail_hooks.qohash import ( + QostodianNexusConfigModel, +) """ Pydantic object defining how to set guardrails on litellm proxy @@ -96,6 +99,7 @@ class SupportedGuardrailIntegrations(Enum): AKTO = "akto" MCP_JWT_SIGNER = "mcp_jwt_signer" LLM_AS_A_JUDGE = "llm_as_a_judge" + QOSTODIAN_NEXUS = "qostodian_nexus" class Role(Enum): @@ -773,6 +777,7 @@ class LitellmParams( QualifireGuardrailConfigModel, BlockCodeExecutionGuardrailConfigModel, HiddenlayerGuardrailConfigModel, + QostodianNexusConfigModel, ): guardrail: str = Field(description="The type of guardrail integration to use") mode: Union[str, List[str], Mode] = Field( diff --git a/litellm/types/proxy/guardrails/guardrail_hooks/qohash.py b/litellm/types/proxy/guardrails/guardrail_hooks/qohash.py new file mode 100644 index 00000000000..5abfa69148e --- /dev/null +++ b/litellm/types/proxy/guardrails/guardrail_hooks/qohash.py @@ -0,0 +1,16 @@ +from typing import Optional + +from pydantic import Field + +from .base import GuardrailConfigModel + + +class QostodianNexusConfigModel(GuardrailConfigModel): + api_base: Optional[str] = Field( + default=None, + description="The API base URL for Qostodian Nexus. If not provided, the `QOSTODIAN_NEXUS_API_BASE` environment variable is checked. Defaults to http://nexus:8800.", + ) + + @staticmethod + def ui_friendly_name() -> str: + return "Qostodian Nexus" diff --git a/litellm/types/utils.py b/litellm/types/utils.py index ed29d49fc29..c05c46e0d45 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -140,6 +140,7 @@ class ProviderSpecificModelInfo(TypedDict, total=False): supports_url_context: Optional[bool] supports_none_reasoning_effort: Optional[bool] supports_minimal_reasoning_effort: Optional[bool] + supports_low_reasoning_effort: Optional[bool] supports_xhigh_reasoning_effort: Optional[bool] supports_max_reasoning_effort: Optional[bool] diff --git a/litellm/utils.py b/litellm/utils.py index 027c9fedced..8c1b6452ced 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -5896,6 +5896,9 @@ def _get_model_info_helper( # noqa: PLR0915 supports_minimal_reasoning_effort=_model_info.get( "supports_minimal_reasoning_effort", None ), + supports_low_reasoning_effort=_model_info.get( + "supports_low_reasoning_effort", None + ), supports_xhigh_reasoning_effort=_model_info.get( "supports_xhigh_reasoning_effort", None ), diff --git a/litellm/vector_stores/main.py b/litellm/vector_stores/main.py index 6d28d670979..13f2f27d3fa 100644 --- a/litellm/vector_stores/main.py +++ b/litellm/vector_stores/main.py @@ -377,15 +377,11 @@ def search( _is_async = kwargs.pop("asearch", False) is True # pull credentials from registry if available - vector_store_id_for_credentials = kwargs.get("vector_store_id", vector_store_id) - if ( - litellm.vector_store_registry is not None - and vector_store_id_for_credentials is not None - ): + if litellm.vector_store_registry is not None and vector_store_id is not None: try: registry_credentials = ( litellm.vector_store_registry.get_credentials_for_vector_store( - vector_store_id_for_credentials + vector_store_id ) ) kwargs.update(registry_credentials) diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index 8391fdb48f9..76dcfc55c10 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -19942,7 +19942,7 @@ "supports_web_search": true, "supports_none_reasoning_effort": true, "supports_xhigh_reasoning_effort": true, - "supports_minimal_reasoning_effort": true + "supports_minimal_reasoning_effort": false }, "gpt-5.5-2026-04-23": { "cache_read_input_token_cost": 5e-07, @@ -19990,7 +19990,7 @@ "supports_web_search": true, "supports_none_reasoning_effort": true, "supports_xhigh_reasoning_effort": true, - "supports_minimal_reasoning_effort": true + "supports_minimal_reasoning_effort": false }, "gpt-5.5-pro": { "cache_read_input_token_cost": 3e-06, @@ -20033,7 +20033,8 @@ "supports_web_search": true, "supports_none_reasoning_effort": false, "supports_xhigh_reasoning_effort": true, - "supports_minimal_reasoning_effort": true + "supports_minimal_reasoning_effort": false, + "supports_low_reasoning_effort": false }, "gpt-5.5-pro-2026-04-23": { "cache_read_input_token_cost": 3e-06, @@ -20076,7 +20077,8 @@ "supports_web_search": true, "supports_none_reasoning_effort": false, "supports_xhigh_reasoning_effort": true, - "supports_minimal_reasoning_effort": true + "supports_minimal_reasoning_effort": false, + "supports_low_reasoning_effort": false }, "gpt-5.4": { "cache_read_input_token_cost": 2.5e-07, @@ -22109,6 +22111,98 @@ "tool_use_system_prompt_tokens": 346, "supports_native_structured_output": true }, + "crusoe/deepseek-ai/DeepSeek-R1-0528": { + "input_cost_per_token": 3e-06, + "litellm_provider": "crusoe", + "max_input_tokens": 163840, + "max_output_tokens": 163840, + "max_tokens": 163840, + "mode": "chat", + "output_cost_per_token": 7e-06, + "supports_function_calling": false, + "supports_reasoning": true, + "supports_system_messages": true, + "supports_tool_choice": false + }, + "crusoe/deepseek-ai/DeepSeek-V3-0324": { + "input_cost_per_token": 1.5e-06, + "litellm_provider": "crusoe", + "max_input_tokens": 163840, + "max_output_tokens": 163840, + "max_tokens": 163840, + "mode": "chat", + "output_cost_per_token": 1.5e-06, + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_system_messages": true, + "supports_tool_choice": true + }, + "crusoe/google/gemma-3-12b-it": { + "input_cost_per_token": 1e-07, + "litellm_provider": "crusoe", + "max_input_tokens": 131072, + "max_output_tokens": 131072, + "max_tokens": 131072, + "mode": "chat", + "output_cost_per_token": 1e-07, + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_vision": true + }, + "crusoe/meta-llama/Llama-3.3-70B-Instruct": { + "input_cost_per_token": 2e-07, + "litellm_provider": "crusoe", + "max_input_tokens": 131072, + "max_output_tokens": 131072, + "max_tokens": 131072, + "mode": "chat", + "output_cost_per_token": 2e-07, + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_system_messages": true, + "supports_tool_choice": true + }, + "crusoe/moonshotai/Kimi-K2-Thinking": { + "input_cost_per_token": 2.5e-06, + "litellm_provider": "crusoe", + "max_input_tokens": 262144, + "max_output_tokens": 262144, + "max_tokens": 262144, + "mode": "chat", + "output_cost_per_token": 2.5e-06, + "supports_function_calling": false, + "supports_reasoning": true, + "supports_system_messages": true, + "supports_tool_choice": false + }, + "crusoe/openai/gpt-oss-120b": { + "input_cost_per_token": 8e-07, + "litellm_provider": "crusoe", + "max_input_tokens": 131072, + "max_output_tokens": 131072, + "max_tokens": 131072, + "mode": "chat", + "output_cost_per_token": 8e-07, + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_system_messages": true, + "supports_tool_choice": true + }, + "crusoe/Qwen/Qwen3-235B-A22B-Instruct-2507": { + "input_cost_per_token": 3e-06, + "litellm_provider": "crusoe", + "max_input_tokens": 262144, + "max_output_tokens": 262144, + "max_tokens": 262144, + "mode": "chat", + "output_cost_per_token": 3e-06, + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_system_messages": true, + "supports_tool_choice": true + }, "lambda_ai/deepseek-llama3.3-70b": { "input_cost_per_token": 2e-07, "litellm_provider": "lambda_ai", diff --git a/provider_endpoints_support.json b/provider_endpoints_support.json index ed49c146210..3fc7cd43187 100644 --- a/provider_endpoints_support.json +++ b/provider_endpoints_support.json @@ -635,6 +635,24 @@ "interactions": true } }, + "crusoe": { + "display_name": "Crusoe (`crusoe`)", + "url": "https://docs.litellm.ai/docs/providers/crusoe", + "endpoints": { + "chat_completions": true, + "messages": true, + "responses": true, + "embeddings": false, + "image_generations": false, + "audio_transcriptions": false, + "audio_speech": false, + "moderations": false, + "batches": false, + "rerank": false, + "a2a": true, + "interactions": true + } + }, "custom": { "display_name": "Custom (`custom`)", "url": "https://docs.litellm.ai/docs/providers/custom_llm_server", diff --git a/tests/_vcr_redis_persister.py b/tests/_vcr_redis_persister.py index 4d72a1142bb..a6ed448f1cb 100644 --- a/tests/_vcr_redis_persister.py +++ b/tests/_vcr_redis_persister.py @@ -13,6 +13,8 @@ CASSETTE_REDIS_URL_ENV = "CASSETTE_REDIS_URL" VCR_VERBOSE_ENV = "LITELLM_VCR_VERBOSE" MAX_EPISODES_PER_CASSETTE = 50 +_REPO_ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) + _log = logging.getLogger(__name__) _passed_by_cassette_key: dict[str, bool] = {} @@ -22,7 +24,11 @@ def mark_test_outcome_for_cassette(cassette_path: str, passed: bool) -> None: def redis_key_for(cassette_path: str) -> str: - rel = os.path.relpath(str(cassette_path)) + abs_path = os.path.abspath(str(cassette_path)) + try: + rel = os.path.relpath(abs_path, start=_REPO_ROOT) + except ValueError: + rel = os.path.basename(abs_path) if rel.endswith(".yaml"): rel = rel[: -len(".yaml")] rel = rel.replace("/cassettes/", "/").lstrip("./") diff --git a/tests/litellm/llms/azure/__init__.py b/tests/litellm/llms/azure/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/litellm/llms/azure/test_azure_embedding.py b/tests/litellm/llms/azure/test_azure_embedding.py new file mode 100644 index 00000000000..22ee503ef0d --- /dev/null +++ b/tests/litellm/llms/azure/test_azure_embedding.py @@ -0,0 +1,94 @@ +import os +import sys +from unittest.mock import MagicMock, patch + +import pytest + +sys.path.insert( + 0, os.path.abspath(os.path.join(os.path.dirname(__file__), "../../../../..")) +) + +from litellm.llms.azure.azure import AzureChatCompletion +from litellm.types.utils import EmbeddingResponse, Usage + + +def _make_embedding_response() -> EmbeddingResponse: + return EmbeddingResponse( + model="text-embedding-3-large", + usage=Usage(prompt_tokens=3, completion_tokens=0, total_tokens=3), + data=[{"embedding": [0.1, 0.2, 0.3], "index": 0, "object": "embedding"}], + ) + + +def _make_logging_obj() -> MagicMock: + return MagicMock() + + +class TestAzureV1AsyncEmbedding: + def test_aembedding_receives_api_version(self): + """Regression: api_version must be forwarded to aembedding() when aembedding=True. + Without the fix, it was silently dropped, causing AsyncAzureOpenAI to be used + instead of AsyncOpenAI for Azure AI Foundry (v1) endpoints. Fixes #24848.""" + handler = AzureChatCompletion() + + with patch.object(handler, "aembedding") as mock_aembedding: + handler.embedding( + model="text-embedding-3-large", + input=["hello world"], + api_base="https://my-endpoint.openai.azure.com", + api_version="v1", + timeout=60.0, + logging_obj=_make_logging_obj(), + model_response=_make_embedding_response(), + optional_params={}, + api_key="fake-key", + aembedding=True, + litellm_params={}, + ) + + mock_aembedding.assert_called_once() + _, kwargs = mock_aembedding.call_args + assert kwargs.get("api_version") == "v1" + + def test_get_azure_openai_client_returns_async_openai_for_v1(self): + from openai import AsyncAzureOpenAI, AsyncOpenAI + + handler = AzureChatCompletion() + client = handler.get_azure_openai_client( + api_key="fake-key", + api_base="https://my-endpoint.openai.azure.com", + api_version="v1", + _is_async=True, + litellm_params={}, + ) + + assert isinstance(client, AsyncOpenAI) + assert not isinstance(client, AsyncAzureOpenAI) + + def test_get_azure_openai_client_uses_v1_base_url(self): + handler = AzureChatCompletion() + client = handler.get_azure_openai_client( + api_key="fake-key", + api_base="https://my-endpoint.openai.azure.com", + api_version="v1", + _is_async=True, + litellm_params={}, + ) + + assert client is not None + assert "/openai/v1/" in str(client.base_url) + + @pytest.mark.parametrize("api_version", ["v1", "latest", "preview"]) + def test_all_v1_variants_use_openai_client(self, api_version: str): + from openai import AsyncOpenAI + + handler = AzureChatCompletion() + client = handler.get_azure_openai_client( + api_key="fake-key", + api_base="https://my-endpoint.openai.azure.com", + api_version=api_version, + _is_async=True, + litellm_params={}, + ) + + assert isinstance(client, AsyncOpenAI) diff --git a/tests/litellm_utils_tests/test_anthropic_token_counter.py b/tests/litellm_utils_tests/test_anthropic_token_counter.py index d099eb4f8e9..028586203a5 100644 --- a/tests/litellm_utils_tests/test_anthropic_token_counter.py +++ b/tests/litellm_utils_tests/test_anthropic_token_counter.py @@ -26,7 +26,7 @@ class TestAnthropicTokenCounter(BaseTokenCounterTest): return AnthropicTokenCounter() def get_test_model(self) -> str: - return "claude-sonnet-4-20250514" + return "claude-haiku-4-5-20251001" def get_test_messages(self) -> List[Dict[str, Any]]: return [{"role": "user", "content": "Hello, how are you today?"}] diff --git a/tests/litellm_utils_tests/test_bedrock_token_counter.py b/tests/litellm_utils_tests/test_bedrock_token_counter.py index 62a595baaaf..9fb2463e8b5 100644 --- a/tests/litellm_utils_tests/test_bedrock_token_counter.py +++ b/tests/litellm_utils_tests/test_bedrock_token_counter.py @@ -130,7 +130,7 @@ class TestBedrockCountTokensEndpoint: ) assert ( url - == "https://bedrock-runtime.us-east-1.amazonaws.com/model/amazon.nova-lite-v1:0/count-tokens" + == "https://bedrock-runtime.us-east-1.amazonaws.com/model/amazon.nova-lite-v1%3A0/count-tokens" ) def test_api_base_overrides_default(self): @@ -141,7 +141,7 @@ class TestBedrockCountTokensEndpoint: aws_region_name="us-east-1", api_base=custom_base, ) - assert url == f"{custom_base}/model/amazon.nova-lite-v1:0/count-tokens" + assert url == f"{custom_base}/model/amazon.nova-lite-v1%3A0/count-tokens" def test_aws_bedrock_runtime_endpoint_overrides_default(self): handler = self._make_handler() @@ -153,7 +153,7 @@ class TestBedrockCountTokensEndpoint: aws_region_name="eu-west-1", aws_bedrock_runtime_endpoint=custom_endpoint, ) - assert url == f"{custom_endpoint}/model/amazon.nova-lite-v1:0/count-tokens" + assert url == f"{custom_endpoint}/model/amazon.nova-lite-v1%3A0/count-tokens" def test_api_base_takes_priority_over_aws_bedrock_runtime_endpoint(self): handler = self._make_handler() @@ -165,7 +165,7 @@ class TestBedrockCountTokensEndpoint: api_base=api_base, aws_bedrock_runtime_endpoint=runtime_endpoint, ) - assert url == f"{api_base}/model/amazon.nova-lite-v1:0/count-tokens" + assert url == f"{api_base}/model/amazon.nova-lite-v1%3A0/count-tokens" def test_env_var_overrides_default(self, monkeypatch): monkeypatch.setenv( diff --git a/tests/llm_responses_api_testing/test_anthropic_responses_api.py b/tests/llm_responses_api_testing/test_anthropic_responses_api.py index 575e1af21a3..6537f67acb9 100644 --- a/tests/llm_responses_api_testing/test_anthropic_responses_api.py +++ b/tests/llm_responses_api_testing/test_anthropic_responses_api.py @@ -89,7 +89,7 @@ def test_multiturn_tool_calls(): "type": "message", } ], - model="anthropic/claude-4-sonnet-20250514", + model="anthropic/claude-haiku-4-5-20251001", instructions="You are a helpful coding assistant.", tools=[shell_tool], ) @@ -115,7 +115,7 @@ def test_multiturn_tool_calls(): # Use await with asyncio.run for the async function follow_up_response = litellm.responses( - model="anthropic/claude-4-sonnet-20250514", + model="anthropic/claude-haiku-4-5-20251001", previous_response_id=response_id, input=[ { diff --git a/tests/llm_translation/test_cloudflare.py b/tests/llm_translation/test_cloudflare.py index 0c799b4f399..5a6a0008398 100644 --- a/tests/llm_translation/test_cloudflare.py +++ b/tests/llm_translation/test_cloudflare.py @@ -43,6 +43,14 @@ def _streaming_chunks() -> list[str]: ] +def _streaming_chunks_response_text() -> list[str]: + return [ + json.dumps({"response_text": "I am"}), + json.dumps({"response_text": " a language"}), + json.dumps({"response_text": " model."}), + ] + + @pytest.mark.parametrize("sync_mode", [True, False]) def test_completion_cloudflare(sync_mode): messages = [{"role": "user", "content": "what llm are you"}] @@ -145,3 +153,76 @@ def test_completion_cloudflare_stream(sync_mode): if c.choices[0].delta.content ) assert "language" in content.lower() + + +@pytest.mark.parametrize("sync_mode", [True, False]) +def test_completion_cloudflare_stream_response_text(sync_mode): + """Newer Cloudflare Workers AI models (e.g. Nemotron) emit `response_text` + instead of `response` in streamed chunks. The iterator must surface that + text so streaming output is not silently empty. + """ + messages = [{"role": "user", "content": "what llm are you"}] + raw_chunks = _streaming_chunks_response_text() + + if sync_mode: + + def _iter_lines(): + for chunk in raw_chunks: + yield f"data: {chunk}" + yield "data: [DONE]" + + mock_resp = MagicMock() + mock_resp.iter_lines.return_value = _iter_lines() + mock_resp.status_code = 200 + mock_resp.headers = {"content-type": "text/event-stream"} + + with patch.object(HTTPHandler, "post", return_value=mock_resp) as mock_post: + response = completion( + model="cloudflare/@cf/nvidia/nemotron-mini-4b-instruct", + messages=messages, + max_tokens=15, + stream=True, + api_base=FAKE_API_BASE, + api_key=FAKE_API_KEY, + ) + chunks_received = list(response) + mock_post.assert_called_once() + else: + + async def _aiter_lines(): + for chunk in raw_chunks: + yield f"data: {chunk}" + yield "data: [DONE]" + + mock_resp = MagicMock() + mock_resp.aiter_lines.return_value = _aiter_lines() + mock_resp.status_code = 200 + mock_resp.headers = {"content-type": "text/event-stream"} + + async def _run(): + with patch.object( + AsyncHTTPHandler, "post", new_callable=AsyncMock, return_value=mock_resp + ) as mock_post: + resp = await acompletion( + model="cloudflare/@cf/nvidia/nemotron-mini-4b-instruct", + messages=messages, + max_tokens=15, + stream=True, + api_base=FAKE_API_BASE, + api_key=FAKE_API_KEY, + ) + received = [] + async for chunk in resp: + received.append(chunk) + mock_post.assert_called_once() + return received + + chunks_received = asyncio.run(_run()) + + assert len(chunks_received) > 0 + content = "".join( + c.choices[0].delta.content + for c in chunks_received + if c.choices[0].delta.content + ) + assert "language" in content.lower() diff --git a/tests/llm_translation/test_crusoe.py b/tests/llm_translation/test_crusoe.py new file mode 100644 index 00000000000..56aa4e4cd42 --- /dev/null +++ b/tests/llm_translation/test_crusoe.py @@ -0,0 +1,108 @@ +""" +Tests for Crusoe provider integration +""" +import os +from unittest import mock + +import litellm + +CRUSOE_API_BASE = "https://managed-inference-api-proxy.crusoecloud.com/v1" + + +def test_crusoe_json_registry(): + """Test CrusoeChatConfig is loaded from JSON provider registry""" + from litellm.llms.openai_like.json_loader import JSONProviderRegistry + + assert JSONProviderRegistry.exists("crusoe") + config = JSONProviderRegistry.get("crusoe") + assert config is not None + assert config.base_url == CRUSOE_API_BASE + assert config.api_key_env == "CRUSOE_API_KEY" + assert config.api_base_env == "CRUSOE_API_BASE" + + +def test_crusoe_get_openai_compatible_provider_info(): + """Test Crusoe provider info retrieval""" + from litellm.llms.openai_like.dynamic_config import create_config_class + from litellm.llms.openai_like.json_loader import JSONProviderRegistry + + config = create_config_class(JSONProviderRegistry.get("crusoe"))() + + # Test with default values (no env vars set) + with mock.patch.dict(os.environ, {}, clear=True): + api_base, api_key = config._get_openai_compatible_provider_info(None, None) + assert api_base == CRUSOE_API_BASE + assert api_key is None + + # Test with environment variables + with mock.patch.dict( + os.environ, + { + "CRUSOE_API_KEY": "test-key", + "CRUSOE_API_BASE": "https://custom.crusoecloud.com/v1", + }, + ): + api_base, api_key = config._get_openai_compatible_provider_info(None, None) + assert api_base == "https://custom.crusoecloud.com/v1" + assert api_key == "test-key" + + # Test with explicit parameters (should override env vars) + with mock.patch.dict( + os.environ, + { + "CRUSOE_API_KEY": "env-key", + "CRUSOE_API_BASE": "https://env.crusoecloud.com/v1", + }, + ): + api_base, api_key = config._get_openai_compatible_provider_info( + "https://param.crusoecloud.com/v1", "param-key" + ) + assert api_base == "https://param.crusoecloud.com/v1" + assert api_key == "param-key" + + +def test_get_llm_provider_crusoe(): + """Test that get_llm_provider correctly identifies Crusoe""" + from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider + + # Test with crusoe/model-name format + model, provider, api_key, api_base = get_llm_provider( + "crusoe/meta-llama/Llama-3.3-70B-Instruct" + ) + assert model == "meta-llama/Llama-3.3-70B-Instruct" + assert provider == "crusoe" + + +def test_crusoe_models_configuration(): + """Test that Crusoe models are configured correctly""" + from litellm import get_model_info + + original_model_cost = litellm.model_cost + original_env = os.environ.get("LITELLM_LOCAL_MODEL_COST_MAP") + try: + os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True" + litellm.model_cost = litellm.get_model_cost_map(url="") + + crusoe_models = [ + "crusoe/meta-llama/Llama-3.3-70B-Instruct", + "crusoe/deepseek-ai/DeepSeek-R1-0528", + "crusoe/deepseek-ai/DeepSeek-V3-0324", + "crusoe/Qwen/Qwen3-235B-A22B-Instruct-2507", + "crusoe/moonshotai/Kimi-K2-Thinking", + "crusoe/openai/gpt-oss-120b", + "crusoe/google/gemma-3-12b-it", + ] + + for model in crusoe_models: + model_info = get_model_info(model) + assert model_info is not None, f"Model info not found for {model}" + assert model_info.get("litellm_provider") == "crusoe", ( + f"{model} should have crusoe as provider" + ) + assert model_info.get("mode") == "chat", f"{model} should be in chat mode" + finally: + litellm.model_cost = original_model_cost + if original_env is None: + os.environ.pop("LITELLM_LOCAL_MODEL_COST_MAP", None) + else: + os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = original_env diff --git a/tests/llm_translation/test_gemini.py b/tests/llm_translation/test_gemini.py index a945a5c1ea1..97b0aaee86b 100644 --- a/tests/llm_translation/test_gemini.py +++ b/tests/llm_translation/test_gemini.py @@ -1331,14 +1331,14 @@ def test_gemini_function_args_preserve_unicode(): assert "José" in arguments_str -def test_anthropic_thinking_param_to_gemini_3_thinkingLevel(): +def test_anthropic_thinking_param_to_gemini_3_provider_defaults(): """ - Test that Anthropic thinking parameters are correctly transformed to Gemini 3 thinkingLevel - instead of thinkingBudget. + Test that Anthropic thinking parameters for Gemini 3+ follow provider defaults + unless force-low behavior is explicitly enabled. For Gemini 3+ models (gemini-3-flash, gemini-3-pro, gemini-3-flash-preview): - - Should use thinkingLevel instead of thinkingBudget - - budget_tokens should map to thinkingLevel + - Should not force thinkingLevel by default + - Should still set includeThoughts correctly Related issue: https://github.com/BerriAI/litellm/issues/XXXX """ @@ -1347,72 +1347,102 @@ def test_anthropic_thinking_param_to_gemini_3_thinkingLevel(): ) from litellm.types.llms.anthropic import AnthropicThinkingParam + original_force_low_flag = litellm.enable_gemini_default_thinking_level_low + litellm.enable_gemini_default_thinking_level_low = False + # Test 1: Anthropic thinking enabled with budget_tokens for Gemini 3 model thinking_param: AnthropicThinkingParam = { "type": "enabled", "budget_tokens": 10000, } + try: + result = VertexGeminiConfig._map_thinking_param( + thinking_param=thinking_param, + model="gemini-3-flash", + ) - result = VertexGeminiConfig._map_thinking_param( - thinking_param=thinking_param, - model="gemini-3-flash", + # For Gemini 3, should not force thinkingLevel by default + assert "thinkingLevel" not in result, "Should not force thinkingLevel for Gemini 3" + assert "thinkingBudget" not in result, "Should NOT have thinkingBudget for Gemini 3" + assert result["includeThoughts"] is True + + # Test 2: Anthropic thinking disabled for Gemini 3 + thinking_param_disabled: AnthropicThinkingParam = { + "type": "disabled", + "budget_tokens": None, + } + + result_disabled = VertexGeminiConfig._map_thinking_param( + thinking_param=thinking_param_disabled, + model="gemini-3-pro-preview", + ) + + assert result_disabled.get("includeThoughts") is False + assert ( + "thinkingLevel" not in result_disabled + or result_disabled.get("thinkingLevel") is None + ) + + # Test 3: Budget tokens = 0 for Gemini 3 + thinking_param_zero: AnthropicThinkingParam = { + "type": "enabled", + "budget_tokens": 0, + } + + result_zero = VertexGeminiConfig._map_thinking_param( + thinking_param=thinking_param_zero, + model="gemini-3-flash", + ) + + assert result_zero["includeThoughts"] is False + assert "thinkingLevel" not in result_zero or result_zero.get("thinkingLevel") is None + + # Test 4: Gemini 3 flash-preview should also follow provider defaults by default + result_gemini3flashpreview = VertexGeminiConfig._map_thinking_param( + thinking_param=thinking_param, + model="gemini-3-flash-preview", + ) + + assert "thinkingLevel" not in result_gemini3flashpreview + assert "thinkingBudget" not in result_gemini3flashpreview + assert result_gemini3flashpreview["includeThoughts"] is True + finally: + litellm.enable_gemini_default_thinking_level_low = original_force_low_flag + + +def test_anthropic_thinking_param_to_gemini_3_force_low_feature_flag(): + """ + Test that Gemini 3 thinkingLevel forced mapping is available behind a feature flag. + """ + from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import ( + VertexGeminiConfig, ) + from litellm.types.llms.anthropic import AnthropicThinkingParam - # For Gemini 3, should use thinkingLevel, not thinkingBudget - assert "thinkingLevel" in result, "Should have thinkingLevel for Gemini 3" - assert "thinkingBudget" not in result, "Should NOT have thinkingBudget for Gemini 3" - assert result["includeThoughts"] is True - assert result["thinkingLevel"] in [ - "minimal", - "low", - ], "thinkingLevel should be 'minimal' or 'low'" + original_force_low_flag = litellm.enable_gemini_default_thinking_level_low + litellm.enable_gemini_default_thinking_level_low = True - # Test 2: Anthropic thinking disabled for Gemini 3 - thinking_param_disabled: AnthropicThinkingParam = { - "type": "disabled", - "budget_tokens": None, - } - - result_disabled = VertexGeminiConfig._map_thinking_param( - thinking_param=thinking_param_disabled, - model="gemini-3-pro-preview", - ) - - assert result_disabled.get("includeThoughts") is False - assert ( - "thinkingLevel" not in result_disabled - or result_disabled.get("thinkingLevel") is None - ) - - # Test 3: Budget tokens = 0 for Gemini 3 - thinking_param_zero: AnthropicThinkingParam = { + thinking_param: AnthropicThinkingParam = { "type": "enabled", - "budget_tokens": 0, + "budget_tokens": 10000, } - result_zero = VertexGeminiConfig._map_thinking_param( - thinking_param=thinking_param_zero, - model="gemini-3-flash", - ) + try: + result_flash = VertexGeminiConfig._map_thinking_param( + thinking_param=thinking_param, + model="gemini-3-flash", + ) + assert result_flash["thinkingLevel"] == "minimal" + assert result_flash["includeThoughts"] is True - assert result_zero["includeThoughts"] is False - assert ( - "thinkingLevel" not in result_zero or result_zero.get("thinkingLevel") is None - ) - - # Test 4: Fiercefalcon model (Gemini 3 Flash checkpoint) should use thinkingLevel - result_gemini3flashpreview = VertexGeminiConfig._map_thinking_param( - thinking_param=thinking_param, - model="gemini-3-flash-preview", - ) - - assert ( - "thinkingLevel" in result_gemini3flashpreview - ), "Should have thinkingLevel for gemini-3-flash-preview" - assert ( - "thinkingBudget" not in result_gemini3flashpreview - ), "Should NOT have thinkingBudget for gemini-3-flash-preview" - assert result_gemini3flashpreview["includeThoughts"] is True + result_pro = VertexGeminiConfig._map_thinking_param( + thinking_param=thinking_param, + model="gemini-3-pro-preview", + ) + assert result_pro["thinkingLevel"] == "low" + assert result_pro["includeThoughts"] is True + finally: + litellm.enable_gemini_default_thinking_level_low = original_force_low_flag def test_anthropic_thinking_param_to_gemini_2_thinkingBudget(): @@ -1465,7 +1495,7 @@ def test_anthropic_thinking_param_to_gemini_2_thinkingBudget(): def test_anthropic_thinking_param_via_map_openai_params(): """ Test that the thinking parameter is correctly transformed through the full map_openai_params flow - for Gemini 3 models, resulting in thinkingConfig with thinkingLevel. + for Gemini 3 models, without forcing thinkingLevel by default. This tests the full integration from Anthropic API format to Gemini format. """ @@ -1492,13 +1522,11 @@ def test_anthropic_thinking_param_via_map_openai_params(): drop_params=False, ) - # Check that thinkingConfig was created with thinkingLevel + # Check that thinkingConfig was created without forced thinkingLevel assert "thinkingConfig" in result, "Should have thinkingConfig in optional_params" thinking_config = result["thinkingConfig"] - assert "thinkingLevel" in thinking_config, "Should have thinkingLevel for Gemini 3" - assert ( - "thinkingBudget" not in thinking_config - ), "Should NOT have thinkingBudget for Gemini 3" + assert "thinkingLevel" not in thinking_config, "Should not force thinkingLevel for Gemini 3 by default" + assert "thinkingBudget" not in thinking_config, "Should NOT have thinkingBudget for Gemini 3" assert thinking_config["includeThoughts"] is True # Test with Gemini 2 model diff --git a/tests/llm_translation/test_vcr_redis_persister.py b/tests/llm_translation/test_vcr_redis_persister.py index 853558150c1..6e62e4491cb 100644 --- a/tests/llm_translation/test_vcr_redis_persister.py +++ b/tests/llm_translation/test_vcr_redis_persister.py @@ -81,6 +81,31 @@ def test_redis_key_normalizes_path_passed_by_pytest_recording(): ) +def test_redis_key_is_stable_across_working_directories(tmp_path, monkeypatch): + repo_root = os.path.dirname( + os.path.dirname(os.path.dirname(os.path.abspath(__file__))) + ) + abs_cassette = os.path.join( + repo_root, + "tests/llm_translation/cassettes/test_anthropic/test_streaming.yaml", + ) + + monkeypatch.chdir(repo_root) + key_from_root = redis_key_for(abs_cassette) + + monkeypatch.chdir(os.path.join(repo_root, "tests", "llm_translation")) + key_from_subdir = redis_key_for(abs_cassette) + + monkeypatch.chdir(tmp_path) + key_from_tmp = redis_key_for(abs_cassette) + + assert key_from_root == key_from_subdir == key_from_tmp + assert ( + key_from_root + == "litellm:vcr:cassette:tests/llm_translation/test_anthropic/test_streaming" + ) + + class _FlakyRedis: def __init__(self, inner, fail_on: str): self._inner = inner diff --git a/tests/local_testing/test_embedding.py b/tests/local_testing/test_embedding.py index 3f1a397ebcb..82510b6f4fd 100644 --- a/tests/local_testing/test_embedding.py +++ b/tests/local_testing/test_embedding.py @@ -1257,22 +1257,13 @@ def test_jina_ai_img_embeddings(input_data, expected_payload_input): assert sent_data["input"] == expected_payload_input -def test_encoding_format_none_not_omitted_from_openai_sdk(): +def test_encoding_format_defaults_to_float_for_openai_sdk(monkeypatch): """ - Test that encoding_format=None is explicitly sent to OpenAI SDK. + When encoding_format is not provided, LiteLLM sends `float` for OpenAI-path embeddings. - This test verifies that when encoding_format is not provided by the user, - liteLLM explicitly sets it to None rather than omitting it. This prevents - the OpenAI SDK from adding its default value of 'base64'. - - Without this fix: - - OpenAI SDK adds encoding_format='base64' as default when parameter is missing - - This causes issues with providers that don't support encoding_format (like Gemini) - - With this fix: - - encoding_format=None is explicitly passed - - OpenAI SDK respects the explicit None and doesn't add defaults + Optional global override: `LITELLM_DEFAULT_EMBEDDING_ENCODING_FORMAT`. """ + monkeypatch.delenv("LITELLM_DEFAULT_EMBEDDING_ENCODING_FORMAT", raising=False) with patch( "litellm.llms.openai.openai.OpenAIChatCompletion._get_openai_client" ) as mock_get_client: @@ -1310,17 +1301,12 @@ def test_encoding_format_none_not_omitted_from_openai_sdk(): call_kwargs = call_args[1] # Get kwargs - # The key assertion: encoding_format should be in the request with value None - # This prevents OpenAI SDK from adding its default 'base64' value - assert "encoding_format" in call_kwargs, ( - "encoding_format should be explicitly passed to OpenAI SDK " - "(even if None) to prevent SDK from adding default value" - ) + assert "encoding_format" in call_kwargs assert ( - call_kwargs["encoding_format"] is None - ), "encoding_format should be None when not provided by user" + call_kwargs["encoding_format"] == "float" + ), "encoding_format should default to float when not provided by user" - print("✅ PASS: encoding_format=None is correctly passed to OpenAI SDK") + print("✅ PASS: encoding_format='float' is correctly passed to OpenAI SDK") def test_encoding_format_explicit_value_preserved(): diff --git a/tests/local_testing/test_function_calling.py b/tests/local_testing/test_function_calling.py index 02affa1d57c..6fd253ee294 100644 --- a/tests/local_testing/test_function_calling.py +++ b/tests/local_testing/test_function_calling.py @@ -158,7 +158,7 @@ def test_aaparallel_function_call(model): @pytest.mark.parametrize( "model", [ - "anthropic/claude-4-sonnet-20250514", + "anthropic/claude-haiku-4-5-20251001", "bedrock/us.anthropic.claude-sonnet-4-5-20250929-v1:0", ], ) diff --git a/tests/local_testing/test_streaming.py b/tests/local_testing/test_streaming.py index bc3c34b1a5d..7fe42ed5461 100644 --- a/tests/local_testing/test_streaming.py +++ b/tests/local_testing/test_streaming.py @@ -1270,7 +1270,7 @@ def test_bedrock_claude_3_streaming(): @pytest.mark.parametrize( "model", [ - "claude-4-sonnet-20250514", + "claude-haiku-4-5-20251001", "cohere.command-r-plus-v1:0", # bedrock "gpt-3.5-turbo", ], @@ -2696,7 +2696,7 @@ def test_completion_claude_3_function_call_with_streaming(): try: # test without max tokens response = completion( - model="claude-4-sonnet-20250514", + model="claude-haiku-4-5-20251001", messages=messages, tools=tools, tool_choice="required", diff --git a/tests/logging_callback_tests/test_langfuse_dynamic_credentials.py b/tests/logging_callback_tests/test_langfuse_dynamic_credentials.py new file mode 100644 index 00000000000..1b198623381 --- /dev/null +++ b/tests/logging_callback_tests/test_langfuse_dynamic_credentials.py @@ -0,0 +1,129 @@ +import sys +from types import ModuleType, SimpleNamespace + +import litellm +from litellm.integrations.langfuse.langfuse import resolve_langfuse_credentials +from litellm.integrations.langfuse.langfuse_handler import LangFuseHandler + + +def test_resolve_langfuse_credentials_does_not_use_env_for_dynamic_host(monkeypatch): + monkeypatch.setenv("LANGFUSE_PUBLIC_KEY", "global-public") + monkeypatch.setenv("LANGFUSE_SECRET_KEY", "global-secret") + + public_key, secret_key, host = resolve_langfuse_credentials( + langfuse_host="https://attacker.example", + allow_env_credentials=False, + ) + + assert public_key is None + assert secret_key is None + assert host == "https://attacker.example" + + +def test_resolve_langfuse_credentials_accepts_secret_key_alias_for_dynamic_host( + monkeypatch, +): + monkeypatch.setenv("LANGFUSE_SECRET_KEY", "global-secret") + + public_key, secret_key, host = resolve_langfuse_credentials( + langfuse_public_key="dynamic-public", + langfuse_secret_key="dynamic-secret", + langfuse_host="https://team-langfuse.example", + allow_env_credentials=False, + ) + + assert public_key == "dynamic-public" + assert secret_key == "dynamic-secret" + assert host == "https://team-langfuse.example" + + +def test_resolve_langfuse_credentials_keeps_env_for_global_config(monkeypatch): + monkeypatch.setenv("LANGFUSE_PUBLIC_KEY", "global-public") + monkeypatch.setenv("LANGFUSE_SECRET_KEY", "global-secret") + + public_key, secret_key, host = resolve_langfuse_credentials( + langfuse_host="https://admin-configured.example", + allow_env_credentials=True, + ) + + assert public_key == "global-public" + assert secret_key == "global-secret" + assert host == "https://admin-configured.example" + + +def test_upstream_langfuse_debug_env_is_passed(monkeypatch): + from litellm.integrations.langfuse.langfuse import LangFuseLogger + + class FakeLangfuse: + instances = [] + + def __init__(self, **kwargs): + self.kwargs = kwargs + FakeLangfuse.instances.append(self) + + fake_langfuse_module = ModuleType("langfuse") + fake_langfuse_module.Langfuse = FakeLangfuse + fake_langfuse_module.version = SimpleNamespace(__version__="2.6.0") + + monkeypatch.setitem(sys.modules, "langfuse", fake_langfuse_module) + monkeypatch.setattr(litellm, "initialized_langfuse_clients", 0) + monkeypatch.setenv("LANGFUSE_MOCK", "true") + monkeypatch.setenv("UPSTREAM_LANGFUSE_SECRET_KEY", "upstream-secret") + monkeypatch.setenv("UPSTREAM_LANGFUSE_PUBLIC_KEY", "upstream-public") + monkeypatch.setenv("UPSTREAM_LANGFUSE_HOST", "https://upstream.example") + monkeypatch.setenv("UPSTREAM_LANGFUSE_RELEASE", "release") + monkeypatch.setenv("UPSTREAM_LANGFUSE_DEBUG", "true") + + logger = LangFuseLogger( + langfuse_public_key="public", + langfuse_secret="secret", + langfuse_host="https://langfuse.example", + ) + + assert logger.upstream_langfuse_debug == "true" + assert FakeLangfuse.instances[-1].kwargs["debug"] is True + + +def test_langfuse_handler_accepts_secret_key_alias(monkeypatch): + captured = {} + + class FakeLangFuseLogger: + def __init__( + self, + *, + langfuse_public_key=None, + langfuse_secret=None, + langfuse_host=None, + allow_env_credentials=True, + ): + captured["langfuse_public_key"] = langfuse_public_key + captured["langfuse_secret"] = langfuse_secret + captured["langfuse_host"] = langfuse_host + captured["allow_env_credentials"] = allow_env_credentials + + class FakeDynamicLoggingCache: + def set_cache(self, *, credentials, service_name, logging_obj): + captured["cached_credentials"] = credentials + captured["cached_service_name"] = service_name + captured["cached_logging_obj"] = logging_obj + + monkeypatch.setattr( + "litellm.integrations.langfuse.langfuse_handler.LangFuseLogger", + FakeLangFuseLogger, + ) + + logger = LangFuseHandler._create_langfuse_logger_from_credentials( + credentials={ + "langfuse_public_key": "dynamic-public", + "langfuse_secret_key": "dynamic-secret", + "langfuse_host": "https://langfuse.example", + }, + in_memory_dynamic_logger_cache=FakeDynamicLoggingCache(), + ) + + assert captured["langfuse_public_key"] == "dynamic-public" + assert captured["langfuse_secret"] == "dynamic-secret" + assert captured["langfuse_host"] == "https://langfuse.example" + assert captured["allow_env_credentials"] is False + assert captured["cached_service_name"] == "langfuse" + assert captured["cached_logging_obj"] is logger diff --git a/tests/logging_callback_tests/test_langsmith_dynamic_credentials.py b/tests/logging_callback_tests/test_langsmith_dynamic_credentials.py new file mode 100644 index 00000000000..f1912c58464 --- /dev/null +++ b/tests/logging_callback_tests/test_langsmith_dynamic_credentials.py @@ -0,0 +1,50 @@ +import pytest + +from litellm.integrations.langsmith import LangsmithLogger + + +@pytest.mark.asyncio +async def test_get_credentials_from_env_does_not_use_env_for_dynamic_base_url( + monkeypatch, +): + monkeypatch.setenv("LANGSMITH_API_KEY", "global-key") + monkeypatch.setenv("LANGSMITH_PROJECT", "global-project") + monkeypatch.setenv("LANGSMITH_TENANT_ID", "global-tenant") + logger = LangsmithLogger( + langsmith_api_key="default-key", + langsmith_project="default-project", + langsmith_base_url="https://default.example", + ) + + credentials = logger.get_credentials_from_env( + langsmith_base_url="https://attacker.example", + allow_env_credentials=False, + ) + + assert credentials["LANGSMITH_API_KEY"] is None + assert credentials["LANGSMITH_PROJECT"] == "litellm-completion" + assert credentials["LANGSMITH_BASE_URL"] == "https://attacker.example" + assert credentials["LANGSMITH_TENANT_ID"] is None + + +@pytest.mark.asyncio +async def test_dynamic_langsmith_base_url_does_not_inherit_default_api_key( + monkeypatch, +): + monkeypatch.setenv("LANGSMITH_API_KEY", "global-key") + logger = LangsmithLogger( + langsmith_api_key="default-key", + langsmith_project="default-project", + langsmith_base_url="https://default.example", + ) + + credentials = logger._get_credentials_to_use_for_request( + kwargs={ + "standard_callback_dynamic_params": { + "langsmith_base_url": "https://attacker.example" + } + } + ) + + assert credentials["LANGSMITH_API_KEY"] is None + assert credentials["LANGSMITH_BASE_URL"] == "https://attacker.example" diff --git a/tests/pass_through_tests/base_anthropic_messages_test.py b/tests/pass_through_tests/base_anthropic_messages_test.py index e86e58de33b..95f709e3880 100644 --- a/tests/pass_through_tests/base_anthropic_messages_test.py +++ b/tests/pass_through_tests/base_anthropic_messages_test.py @@ -54,7 +54,7 @@ class BaseAnthropicMessagesTest(ABC): print("making request to anthropic passthrough with thinking") client = self.get_client() response = client.messages.create( - model="claude-4-sonnet-20250514", + model="claude-haiku-4-5-20251001", max_tokens=20000, thinking={"type": "enabled", "budget_tokens": 16000}, messages=[ @@ -75,7 +75,7 @@ class BaseAnthropicMessagesTest(ABC): collected_response = [] client = self.get_client() with client.messages.stream( - model="claude-4-sonnet-20250514", + model="claude-haiku-4-5-20251001", max_tokens=20000, thinking={"type": "enabled", "budget_tokens": 16000}, messages=[ diff --git a/tests/pass_through_tests/test_anthropic_passthrough.py b/tests/pass_through_tests/test_anthropic_passthrough.py index c4ae00768c2..d42e06937dc 100644 --- a/tests/pass_through_tests/test_anthropic_passthrough.py +++ b/tests/pass_through_tests/test_anthropic_passthrough.py @@ -333,7 +333,7 @@ async def test_anthropic_messages_streaming_cost_injection(): } payload = { - "model": "claude-4-sonnet-20250514", + "model": "claude-haiku-4-5-20251001", "max_tokens": 10, "stream": True, "messages": [{"role": "user", "content": "Say 'Hi'"}], diff --git a/tests/proxy_unit_tests/test_auth_checks.py b/tests/proxy_unit_tests/test_auth_checks.py index 5636a55c95a..d9f4a6e56b8 100644 --- a/tests/proxy_unit_tests/test_auth_checks.py +++ b/tests/proxy_unit_tests/test_auth_checks.py @@ -1153,3 +1153,320 @@ async def test_can_key_call_model_via_access_group_ids(): valid_token=user_api_key_object, llm_router=router, ) + + +# --------------------------------------------------------------------------- +# _key_access_group_grants_model (key access group overriding team restriction) +# --------------------------------------------------------------------------- + + +def _patch_proxy_server_globals(): + """Patch proxy_server's prisma_client and user_api_key_cache to non-None mocks + so the helper's None-guard doesn't short-circuit. The actual values don't + matter because get_access_object is patched separately to return fixtures.""" + from unittest.mock import MagicMock, patch + + return [ + patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), + patch("litellm.proxy.proxy_server.user_api_key_cache", MagicMock()), + patch("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock()), + ] + + +def _fake_access_group( + access_group_id: str, + access_model_names=None, + assigned_team_ids=None, + assigned_key_ids=None, +): + from litellm.proxy._types import LiteLLM_AccessGroupTable + + return LiteLLM_AccessGroupTable( + access_group_id=access_group_id, + access_group_name=access_group_id, + access_model_names=access_model_names or [], + assigned_team_ids=assigned_team_ids or [], + assigned_key_ids=assigned_key_ids or [], + ) + + +@pytest.mark.asyncio +async def test_key_access_group_grants_model_when_team_authorized(): + """Group's assigned_team_ids includes the key's team and grants the model → True. + + This is the happy path equivalent of Andres's report: admin creates an + access group with assigned_team_ids=[team-a], grants claude-haiku-4-5, + attaches it to a key on team-a. Override fires. + """ + from unittest.mock import AsyncMock, patch + + from litellm.proxy.auth.auth_checks import _key_access_group_grants_model + + valid_token = UserAPIKeyAuth( + token="test-token", + models=[], + access_group_ids=["premium-group"], + team_id="team-a", + ) + team_object = LiteLLM_TeamTable( + team_id="team-a", + models=["mock-success"], + access_group_ids=[], # deliberately not synced — the access group itself authorizes + ) + + fake_ag = _fake_access_group( + access_group_id="premium-group", + access_model_names=["claude-haiku-4-5"], + assigned_team_ids=["team-a"], + ) + + patches = _patch_proxy_server_globals() + [ + patch( + "litellm.proxy.auth.auth_checks.get_access_object", + new_callable=AsyncMock, + return_value=fake_ag, + ), + ] + for p in patches: + p.start() + try: + assert ( + await _key_access_group_grants_model( + model="claude-haiku-4-5", + valid_token=valid_token, + team_object=team_object, + llm_router=None, + ) + is True + ) + finally: + for p in patches: + p.stop() + + +@pytest.mark.asyncio +async def test_key_access_group_grants_model_when_key_directly_authorized(): + """Group's assigned_key_ids includes the key's token and grants the model → True. + + Per-key authorization path: an admin scopes a group directly to a key + (assigned_key_ids) without listing the team. + """ + from unittest.mock import AsyncMock, patch + + from litellm.proxy.auth.auth_checks import _key_access_group_grants_model + + valid_token = UserAPIKeyAuth( + token="test-token-hashed", + models=[], + access_group_ids=["per-key-group"], + team_id="team-a", + ) + team_object = LiteLLM_TeamTable( + team_id="team-a", + models=["mock-success"], + access_group_ids=[], + ) + + fake_ag = _fake_access_group( + access_group_id="per-key-group", + access_model_names=["claude-haiku-4-5"], + assigned_team_ids=[], + assigned_key_ids=["test-token-hashed"], + ) + + patches = _patch_proxy_server_globals() + [ + patch( + "litellm.proxy.auth.auth_checks.get_access_object", + new_callable=AsyncMock, + return_value=fake_ag, + ), + ] + for p in patches: + p.start() + try: + assert ( + await _key_access_group_grants_model( + model="claude-haiku-4-5", + valid_token=valid_token, + team_object=team_object, + llm_router=None, + ) + is True + ) + finally: + for p in patches: + p.stop() + + +@pytest.mark.asyncio +async def test_key_access_group_grants_model_when_key_has_no_groups(): + """Key with no access_group_ids → False (early return, no DB read).""" + from litellm.proxy.auth.auth_checks import _key_access_group_grants_model + + valid_token = UserAPIKeyAuth( + token="test-token", + models=[], + access_group_ids=[], + team_id="team-a", + ) + team_object = LiteLLM_TeamTable( + team_id="team-a", + models=["mock-success"], + access_group_ids=["any-group"], + ) + assert ( + await _key_access_group_grants_model( + model="claude-haiku-4-5", + valid_token=valid_token, + team_object=team_object, + llm_router=None, + ) + is False + ) + + +@pytest.mark.asyncio +async def test_key_access_group_grants_model_when_group_does_not_cover_model(): + """Group authorizes the team but does not grant the requested model → False.""" + from unittest.mock import AsyncMock, patch + + from litellm.proxy.auth.auth_checks import _key_access_group_grants_model + + valid_token = UserAPIKeyAuth( + token="test-token", + models=[], + access_group_ids=["basic-group"], + team_id="team-a", + ) + team_object = LiteLLM_TeamTable( + team_id="team-a", + models=["mock-success"], + access_group_ids=[], + ) + + fake_ag = _fake_access_group( + access_group_id="basic-group", + access_model_names=["gpt-4o-mini"], + assigned_team_ids=["team-a"], + ) + + patches = _patch_proxy_server_globals() + [ + patch( + "litellm.proxy.auth.auth_checks.get_access_object", + new_callable=AsyncMock, + return_value=fake_ag, + ), + ] + for p in patches: + p.start() + try: + assert ( + await _key_access_group_grants_model( + model="claude-haiku-4-5", + valid_token=valid_token, + team_object=team_object, + llm_router=None, + ) + is False + ) + finally: + for p in patches: + p.stop() + + +@pytest.mark.asyncio +async def test_key_access_group_grants_model_when_group_authorizes_neither(): + """ + Bypass regression test: a team member sets a foreign access group on their + key. The group grants the requested model but its assigned_team_ids / + assigned_key_ids do not include this caller's team or token. Override is + denied — the team's 401 propagates. + """ + from unittest.mock import AsyncMock, patch + + from litellm.proxy.auth.auth_checks import _key_access_group_grants_model + + valid_token = UserAPIKeyAuth( + token="team-a-token", + models=[], + access_group_ids=["team-b-premium"], + team_id="team-a", + ) + team_object = LiteLLM_TeamTable( + team_id="team-a", + models=["mock-success"], + access_group_ids=[], + ) + + fake_ag = _fake_access_group( + access_group_id="team-b-premium", + access_model_names=["claude-opus-4-5"], + assigned_team_ids=["team-b"], + assigned_key_ids=["team-b-token"], + ) + + patches = _patch_proxy_server_globals() + [ + patch( + "litellm.proxy.auth.auth_checks.get_access_object", + new_callable=AsyncMock, + return_value=fake_ag, + ), + ] + for p in patches: + p.start() + try: + assert ( + await _key_access_group_grants_model( + model="claude-opus-4-5", + valid_token=valid_token, + team_object=team_object, + llm_router=None, + ) + is False + ) + finally: + for p in patches: + p.stop() + + +@pytest.mark.asyncio +async def test_key_access_group_grants_model_when_get_access_object_raises(): + """Group lookup failure (404, network, etc.) is treated as no authorization.""" + from unittest.mock import AsyncMock, patch + + from litellm.proxy.auth.auth_checks import _key_access_group_grants_model + + valid_token = UserAPIKeyAuth( + token="test-token", + models=[], + access_group_ids=["missing-group"], + team_id="team-a", + ) + team_object = LiteLLM_TeamTable( + team_id="team-a", + models=["mock-success"], + access_group_ids=[], + ) + + patches = _patch_proxy_server_globals() + [ + patch( + "litellm.proxy.auth.auth_checks.get_access_object", + new_callable=AsyncMock, + side_effect=Exception("not found"), + ), + ] + for p in patches: + p.start() + try: + assert ( + await _key_access_group_grants_model( + model="claude-haiku-4-5", + valid_token=valid_token, + team_object=team_object, + llm_router=None, + ) + is False + ) + finally: + for p in patches: + p.stop() diff --git a/tests/test_litellm/integrations/arize/test_arize_utils.py b/tests/test_litellm/integrations/arize/test_arize_utils.py index a87a4167899..86c5448d468 100644 --- a/tests/test_litellm/integrations/arize/test_arize_utils.py +++ b/tests/test_litellm/integrations/arize/test_arize_utils.py @@ -273,6 +273,91 @@ def test_arize_set_attributes_responses_api(): ) +def test_set_usage_outputs_pydantic_completion_usage(): + """ + Regression test for https://github.com/BerriAI/litellm/issues/13672 + + `_set_usage_outputs` previously called `usage.get(...)` which crashes when + `usage` is a plain Pydantic model (e.g. openai.types.completion_usage.CompletionUsage) + that does not implement dict-style `.get()`. Same crash for nested + `output_tokens_details` / `completion_tokens_details`. + + The function must: + 1. Read total/prompt/completion tokens from a Pydantic usage without `.get`. + 2. Read reasoning_tokens from `completion_tokens_details` (chat completions API) + OR `output_tokens_details` (responses API), even when those nested objects + are Pydantic models without `.get`. + 3. Not raise AttributeError; not call span.record_exception. + """ + from unittest.mock import MagicMock + + from openai.types.completion_usage import ( + CompletionTokensDetails, + CompletionUsage, + ) + + from litellm.integrations.arize._utils import _set_usage_outputs + + span = MagicMock() + + # Plain OpenAI Pydantic model — has no `.get()` + usage = CompletionUsage( + completion_tokens=60, + prompt_tokens=40, + total_tokens=100, + completion_tokens_details=CompletionTokensDetails(reasoning_tokens=25), + ) + assert not hasattr(usage, "get"), "precondition: CompletionUsage must lack .get" + + response_obj = {"usage": usage} + + # Must not raise + _set_usage_outputs(span, response_obj, SpanAttributes) + + span.set_attribute.assert_any_call(SpanAttributes.LLM_TOKEN_COUNT_TOTAL, 100) + span.set_attribute.assert_any_call(SpanAttributes.LLM_TOKEN_COUNT_PROMPT, 40) + span.set_attribute.assert_any_call(SpanAttributes.LLM_TOKEN_COUNT_COMPLETION, 60) + # reasoning_tokens for chat completions live in completion_tokens_details + span.set_attribute.assert_any_call( + SpanAttributes.LLM_TOKEN_COUNT_COMPLETION_DETAILS_REASONING, 25 + ) + + +def test_set_usage_outputs_pydantic_response_api_usage(): + """ + Same crash also affects Responses API with `output_tokens_details` as a + Pydantic model that lacks `.get()`. Verifies the responses-API path. + """ + from unittest.mock import MagicMock + + from litellm.integrations.arize._utils import _set_usage_outputs + from litellm.types.llms.openai import OutputTokensDetails + + # Build an object that mimics openai ResponsesAPI usage but lacks `.get` + # (uses a plain class — not BaseLiteLLMOpenAIResponseObject) + class PlainResponsesUsage: + def __init__(self): + self.total_tokens = 370 + self.input_tokens = 120 + self.output_tokens = 250 + self.output_tokens_details = OutputTokensDetails(reasoning_tokens=180) + + usage = PlainResponsesUsage() + assert not hasattr(usage, "get") + + span = MagicMock() + response_obj = {"usage": usage} + + _set_usage_outputs(span, response_obj, SpanAttributes) + + span.set_attribute.assert_any_call(SpanAttributes.LLM_TOKEN_COUNT_TOTAL, 370) + span.set_attribute.assert_any_call(SpanAttributes.LLM_TOKEN_COUNT_PROMPT, 120) + span.set_attribute.assert_any_call(SpanAttributes.LLM_TOKEN_COUNT_COMPLETION, 250) + span.set_attribute.assert_any_call( + SpanAttributes.LLM_TOKEN_COUNT_COMPLETION_DETAILS_REASONING, 180 + ) + + class TestArizeLogger(CustomLogger): """ Custom logger implementation to capture standard_callback_dynamic_params. diff --git a/tests/test_litellm/integrations/test_langsmith_init.py b/tests/test_litellm/integrations/test_langsmith_init.py index 14b861355b7..5d6b7c74690 100644 --- a/tests/test_litellm/integrations/test_langsmith_init.py +++ b/tests/test_litellm/integrations/test_langsmith_init.py @@ -6,9 +6,18 @@ import pytest sys.path.insert(0, os.path.abspath("../..")) +import litellm from litellm.integrations.langsmith import LangsmithLogger +@pytest.fixture +def reset_redact_flag(): + """Reset redact_user_api_key_info between tests so global state doesn't leak.""" + original = litellm.redact_user_api_key_info + yield + litellm.redact_user_api_key_info = original + + class TestLangsmithLoggerInit: """Test cases for LangSmith logger initialization, particularly sampling rate handling. @@ -263,3 +272,78 @@ class TestLangsmithPrepareLogData: assert um["input_tokens"] == 100 assert um["output_tokens"] == 50 assert um["total_tokens"] == 150 + + +class TestLangsmithRedactUserApiKeyInfo: + """Verify litellm.redact_user_api_key_info is honored for LangSmith.""" + + def _logger(self): + return LangsmithLogger( + langsmith_api_key="test-key", + langsmith_project="test-project", + ) + + def _metadata_with_user_api_key_fields(self): + return { + "user_api_key_hash": "abc123", + "user_api_key_alias": "engineer-key", + "user_api_key_user_id": "default_user_id", + "user_api_key_team_id": "team-uuid", + "user_api_key_team_alias": "GNT", + "user_api_key_request_route": "/chat/completions", + "user_api_key_spend": 1.64, + "model": "gpt-4", + "requester_metadata": { + "user_api_key_team_id": "team-uuid", + "user_api_key_user_id": "default_user_id", + "session_id": "sess-1", + }, + } + + def test_redact_disabled_keeps_user_api_key_fields(self, reset_redact_flag): + """Flag off: user_api_key_* fields are preserved (no behavior change).""" + litellm.redact_user_api_key_info = False + logger = self._logger() + metadata = self._metadata_with_user_api_key_fields() + + extra = logger._build_extra_metadata(metadata) + + assert extra["user_api_key_hash"] == "abc123" + assert extra["user_api_key_team_id"] == "team-uuid" + assert extra["requester_metadata"]["user_api_key_user_id"] == "default_user_id" + + def test_redact_enabled_strips_top_level_user_api_key_fields( + self, reset_redact_flag + ): + """Flag on: top-level user_api_key_* keys removed; other keys preserved.""" + litellm.redact_user_api_key_info = True + logger = self._logger() + metadata = self._metadata_with_user_api_key_fields() + + extra = logger._build_extra_metadata(metadata) + + for key in ( + "user_api_key_hash", + "user_api_key_alias", + "user_api_key_user_id", + "user_api_key_team_id", + "user_api_key_team_alias", + "user_api_key_request_route", + "user_api_key_spend", + ): + assert key not in extra, f"{key} should be redacted at top level" + assert extra["model"] == "gpt-4" + + def test_redact_enabled_strips_nested_requester_metadata(self, reset_redact_flag): + """Flag on: nested requester_metadata.user_api_key_* removed; session_id still lifted.""" + litellm.redact_user_api_key_info = True + logger = self._logger() + metadata = self._metadata_with_user_api_key_fields() + + extra = logger._build_extra_metadata(metadata) + + nested = extra["requester_metadata"] + assert "user_api_key_team_id" not in nested + assert "user_api_key_user_id" not in nested + assert nested["session_id"] == "sess-1" + assert extra["session_id"] == "sess-1" diff --git a/tests/test_litellm/integrations/test_opentelemetry.py b/tests/test_litellm/integrations/test_opentelemetry.py index b31bbca8893..962806c5b52 100644 --- a/tests/test_litellm/integrations/test_opentelemetry.py +++ b/tests/test_litellm/integrations/test_opentelemetry.py @@ -259,6 +259,189 @@ class TestOpenTelemetryProviderInitialization(unittest.TestCase): ), "Existing LoggerProvider should be respected and not overridden" +class TestOpenTelemetryDualHandlerIsolation(unittest.TestCase): + """Two OpenTelemetry handlers coexisting via skip_set_global=True + must each get their own provider for every signal (tracer/meter/logger).""" + + @staticmethod + def _wire_span_processor(exporter): + """Context manager: while active, the next OpenTelemetry instance + wires its TracerProvider to `exporter`.""" + return patch.object( + OpenTelemetry, + "_get_span_processor", + lambda self, dynamic_headers=None: SimpleSpanProcessor(exporter), + ) + + def test_skip_set_global_creates_isolated_tracer_provider(self): + from opentelemetry.sdk.trace import TracerProvider as SDKTracerProvider + + fake_existing = SDKTracerProvider() + own_exporter = InMemorySpanExporter() + cfg = OpenTelemetryConfig( + exporter="console", service_name="iso-test", skip_set_global=True + ) + with ( + patch.object(trace, "get_tracer_provider", return_value=fake_existing), + patch.object(trace, "set_tracer_provider") as mock_set, + self._wire_span_processor(own_exporter), + ): + handler = OpenTelemetry(config=cfg) + + self.assertIsNot(handler._tracer_provider, fake_existing) + mock_set.assert_not_called() + + handler.tracer.start_span("isolation_check").end() + handler._tracer_provider.force_flush(2000) + self.assertEqual( + [s.name for s in own_exporter.get_finished_spans()], + ["isolation_check"], + ) + + def test_skip_set_global_via_callback_name_back_compat(self): + from opentelemetry.sdk.trace import TracerProvider as SDKTracerProvider + + fake_existing = SDKTracerProvider() + cfg = OpenTelemetryConfig(exporter="console", service_name="lf-back-compat") + with ( + patch.object(trace, "get_tracer_provider", return_value=fake_existing), + patch.object(trace, "set_tracer_provider"), + self._wire_span_processor(InMemorySpanExporter()), + ): + handler = OpenTelemetry(config=cfg, callback_name="langfuse_otel") + + self.assertIsNot(handler._tracer_provider, fake_existing) + + def test_default_behavior_reuses_existing_sdk_tracer_provider(self): + from opentelemetry.sdk.trace import TracerProvider as SDKTracerProvider + + fake_existing = SDKTracerProvider() + with patch.object(trace, "get_tracer_provider", return_value=fake_existing): + handler = OpenTelemetry(config=OpenTelemetryConfig(service_name="shared")) + self.assertIs(handler._tracer_provider, fake_existing) + + def test_skip_set_global_creates_isolated_meter_provider(self): + from opentelemetry import metrics + from opentelemetry.sdk.metrics import MeterProvider as SDKMeterProvider + + fake_existing = SDKMeterProvider() + cfg = OpenTelemetryConfig( + exporter="console", + service_name="meter-iso-test", + enable_metrics=True, + skip_set_global=True, + ) + with ( + patch.object(metrics, "get_meter_provider", return_value=fake_existing), + patch.object(metrics, "set_meter_provider") as mock_set, + self._wire_span_processor(InMemorySpanExporter()), + ): + handler = OpenTelemetry(config=cfg) + + self.assertIsNot(handler._meter_provider, fake_existing) + mock_set.assert_not_called() + + def test_skip_set_global_creates_isolated_logger_provider(self): + from opentelemetry import _logs + from opentelemetry.sdk._logs import LoggerProvider as SDKLoggerProvider + + fake_existing = SDKLoggerProvider() + cfg = OpenTelemetryConfig( + exporter="console", + service_name="logger-iso-test", + enable_events=True, + skip_set_global=True, + ) + with ( + patch.object(_logs, "get_logger_provider", return_value=fake_existing), + patch.object(_logs, "set_logger_provider") as mock_set, + self._wire_span_processor(InMemorySpanExporter()), + ): + handler = OpenTelemetry(config=cfg) + + self.assertIsNot(handler._logger_provider, fake_existing) + mock_set.assert_not_called() + + def test_emitted_logs_route_to_isolated_logger_provider(self): + # End-to-end: emitted logs land in the handler's private LoggerProvider, + # not the global one. Guards against get_logger() bypassing self._logger_provider. + from opentelemetry import _logs + from opentelemetry.sdk._logs import LoggerProvider as SDKLoggerProvider + + global_exporter = InMemoryLogExporter() + fake_existing = SDKLoggerProvider() + fake_existing.add_log_record_processor( + SimpleLogRecordProcessor(global_exporter) + ) + + private_exporter = InMemoryLogExporter() + cfg = OpenTelemetryConfig( + exporter="console", + service_name="logger-emit-test", + enable_events=True, + skip_set_global=True, + ) + with ( + patch.object(_logs, "get_logger_provider", return_value=fake_existing), + patch.object(_logs, "set_logger_provider"), + patch.object( + OpenTelemetry, "_get_log_exporter", return_value=private_exporter + ), + self._wire_span_processor(InMemorySpanExporter()), + ): + handler = OpenTelemetry(config=cfg) + + span = handler.tracer.start_span("emit-test") + handler._emit_semantic_logs( + kwargs={"messages": [{"role": "user", "content": "hi"}]}, + response_obj={"choices": []}, + span=span, + ) + span.end() + handler._logger_provider.force_flush(2000) + + self.assertGreater(len(private_exporter.get_finished_logs()), 0) + self.assertEqual(len(global_exporter.get_finished_logs()), 0) + + def test_two_handlers_each_receive_their_own_spans(self): + # Handler A gets explicit injection (production-ish: claims the global). + exporter_a = InMemorySpanExporter() + provider_a = TracerProvider() + provider_a.add_span_processor(SimpleSpanProcessor(exporter_a)) + handler_a = OpenTelemetry( + config=OpenTelemetryConfig(service_name="handler-a"), + tracer_provider=provider_a, + ) + + # Handler B comes along with the global appearing to be A's provider. + exporter_b = InMemorySpanExporter() + cfg_b = OpenTelemetryConfig( + exporter="console", service_name="handler-b", skip_set_global=True + ) + with ( + patch.object(trace, "get_tracer_provider", return_value=provider_a), + patch.object(trace, "set_tracer_provider"), + self._wire_span_processor(exporter_b), + ): + handler_b = OpenTelemetry(config=cfg_b) + + self.assertIsNot(handler_a._tracer_provider, handler_b._tracer_provider) + + handler_a.tracer.start_span("from_handler_a").end() + handler_b.tracer.start_span("from_handler_b").end() + provider_a.force_flush(2000) + handler_b._tracer_provider.force_flush(2000) + + self.assertEqual( + sorted(s.name for s in exporter_a.get_finished_spans()), + ["from_handler_a"], + ) + self.assertEqual( + sorted(s.name for s in exporter_b.get_finished_spans()), + ["from_handler_b"], + ) + + class TestOpenTelemetry(unittest.TestCase): POLL_INTERVAL = 0.05 POLL_TIMEOUT = 2.0 diff --git a/tests/test_litellm/integrations/test_prometheus_api_promql_escape.py b/tests/test_litellm/integrations/test_prometheus_api_promql_escape.py new file mode 100644 index 00000000000..262ca6b6922 --- /dev/null +++ b/tests/test_litellm/integrations/test_prometheus_api_promql_escape.py @@ -0,0 +1,150 @@ +""" +Tests for VERIA-53: PromQL string-literal quoting in +``get_daily_spend_from_prometheus``. + +PromQL string literals follow Go's escape rules +(https://prometheus.io/docs/prometheus/latest/querying/basics/). JSON's +quoting is a strict subset of Go's, so ``json.dumps`` produces a literal +Prometheus parses identically. +""" + +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest + + +def test_quote_safe_input_round_trips(): + from litellm.integrations.prometheus_helpers.prometheus_api import ( + _quote_promql_string_literal, + ) + + assert _quote_promql_string_literal("sk-abc123") == '"sk-abc123"' + assert _quote_promql_string_literal("hash:deadbeef") == '"hash:deadbeef"' + + +def test_quote_escapes_double_quote(): + from litellm.integrations.prometheus_helpers.prometheus_api import ( + _quote_promql_string_literal, + ) + + # A bare double quote would otherwise terminate the label matcher and + # let the attacker append `, foo="..."} or sum(...)`. + assert _quote_promql_string_literal('hello"injected') == '"hello\\"injected"' + + +def test_quote_escapes_backslash(): + from litellm.integrations.prometheus_helpers.prometheus_api import ( + _quote_promql_string_literal, + ) + + assert _quote_promql_string_literal('a\\"b') == '"a\\\\\\"b"' + + +def test_quote_escapes_newlines_and_control_chars(): + """Beyond the security minimum, the canonical Go/JSON escape also + handles control characters that would otherwise produce an invalid + PromQL string literal.""" + from litellm.integrations.prometheus_helpers.prometheus_api import ( + _quote_promql_string_literal, + ) + + assert _quote_promql_string_literal("a\nb") == '"a\\nb"' + assert _quote_promql_string_literal("a\tb") == '"a\\tb"' + assert _quote_promql_string_literal("a\rb") == '"a\\rb"' + + +@pytest.mark.asyncio +async def test_get_daily_spend_does_not_pass_raw_quote_into_query(): + from litellm.integrations.prometheus_helpers import prometheus_api + + captured = {} + + class _FakeResponse: + def json(self): + return {"data": {"result": []}} + + async def _capture(url, params): + captured["url"] = url + captured["params"] = params + return _FakeResponse() + + fake_client = MagicMock() + fake_client.get = AsyncMock(side_effect=_capture) + + with patch.object(prometheus_api, "PROMETHEUS_URL", "http://prom.example"): + with patch.object(prometheus_api, "async_http_handler", fake_client): + await prometheus_api.get_daily_spend_from_prometheus( + api_key='sk-victim"} or sum(other_metric{a="b' + ) + + rendered_query = captured["params"]["query"] + # The legitimate matcher framing must still be intact: one outer + # `delta()` window, one inner `hashed_api_key="..."` matcher. + assert rendered_query.startswith( + 'sum(delta(litellm_spend_metric_total{hashed_api_key="' + ) + assert rendered_query.endswith('"}[1d]))') + + # Every injected `"` from the attacker payload appears as `\"` so the + # PromQL parser treats them as literal characters inside the matcher + # value, never as the terminator that would let the rest parse as + # PromQL syntax. + inner = rendered_query[ + len('sum(delta(litellm_spend_metric_total{hashed_api_key="') : -len('"}[1d]))') + ] + assert '"' not in inner.replace('\\"', "") + + +@pytest.mark.asyncio +async def test_get_daily_spend_with_no_api_key_uses_unfiltered_query(): + from litellm.integrations.prometheus_helpers import prometheus_api + + captured = {} + + class _FakeResponse: + def json(self): + return {"data": {"result": []}} + + async def _capture(url, params): + captured["params"] = params + return _FakeResponse() + + fake_client = MagicMock() + fake_client.get = AsyncMock(side_effect=_capture) + + with patch.object(prometheus_api, "PROMETHEUS_URL", "http://prom.example"): + with patch.object(prometheus_api, "async_http_handler", fake_client): + await prometheus_api.get_daily_spend_from_prometheus(api_key=None) + + assert captured["params"]["query"] == "sum(delta(litellm_spend_metric_total[1d]))" + + +@pytest.mark.asyncio +async def test_get_daily_spend_legitimate_hashed_key_unchanged(): + """A normal hex hashed_api_key flows through `json.dumps` as itself + plus the surrounding quotes — no spurious escaping that would break + real lookups.""" + from litellm.integrations.prometheus_helpers import prometheus_api + + captured = {} + + class _FakeResponse: + def json(self): + return {"data": {"result": []}} + + async def _capture(url, params): + captured["params"] = params + return _FakeResponse() + + fake_client = MagicMock() + fake_client.get = AsyncMock(side_effect=_capture) + + legit_key = "a" * 64 # 64-char hex-ish hashed key + with patch.object(prometheus_api, "PROMETHEUS_URL", "http://prom.example"): + with patch.object(prometheus_api, "async_http_handler", fake_client): + await prometheus_api.get_daily_spend_from_prometheus(api_key=legit_key) + + assert ( + captured["params"]["query"] + == f'sum(delta(litellm_spend_metric_total{{hashed_api_key="{legit_key}"}}[1d]))' + ) diff --git a/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py b/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py index 77284a64cf7..a7a2b7720d7 100644 --- a/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py +++ b/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py @@ -411,6 +411,46 @@ def test_generic_cost_per_token_gpt55_pro(): ) +@pytest.mark.parametrize( + "model,expected_none,expected_xhigh,expected_minimal", + [ + # Verified against OpenAI's live API on 2026-04-24: + # gpt-5.5 -> supports: none, low, medium, high, xhigh + # gpt-5.5-pro -> supports: medium, high, xhigh + # Neither supports "minimal"; gpt-5.5-pro additionally does not support "none". + # The JSON must reflect this so LiteLLM rejects unsupported values locally + # (or drops them with drop_params=True) instead of round-tripping to OpenAI + # for a 400. + ("gpt-5.5", True, True, False), + ("gpt-5.5-2026-04-23", True, True, False), + ("gpt-5.5-pro", False, True, False), + ("gpt-5.5-pro-2026-04-23", False, True, False), + ], +) +def test_gpt55_reasoning_effort_flags_match_live_openai_api( + model, expected_none, expected_xhigh, expected_minimal +): + """Pin reasoning_effort capability flags to OpenAI's actual API contract. + + Observed via `POST /v1/chat/completions` with reasoning_effort=minimal: + ``Unsupported value: 'reasoning_effort' does not support 'minimal' with + this model``. gpt-5.5-pro additionally rejects 'none' and 'low'. + """ + os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True" + litellm.model_cost = litellm.get_model_cost_map(url="") + + m = litellm.model_cost[model] + assert ( + m.get("supports_none_reasoning_effort") is expected_none + ), f"{model}: supports_none_reasoning_effort expected {expected_none}" + assert ( + m.get("supports_xhigh_reasoning_effort") is expected_xhigh + ), f"{model}: supports_xhigh_reasoning_effort expected {expected_xhigh}" + assert ( + m.get("supports_minimal_reasoning_effort") is expected_minimal + ), f"{model}: supports_minimal_reasoning_effort expected {expected_minimal}" + + @pytest.mark.parametrize( "base_model,dated_model", [ diff --git a/tests/test_litellm/llms/anthropic/chat/test_anthropic_chat_transformation.py b/tests/test_litellm/llms/anthropic/chat/test_anthropic_chat_transformation.py index e8da50f0ec7..cbb9b21d620 100644 --- a/tests/test_litellm/llms/anthropic/chat/test_anthropic_chat_transformation.py +++ b/tests/test_litellm/llms/anthropic/chat/test_anthropic_chat_transformation.py @@ -1687,9 +1687,7 @@ def test_max_effort_rejected_for_opus_45(): messages = [{"role": "user", "content": "Test"}] - with pytest.raises( - ValueError, match="effort='max' is not supported by this model" - ): + with pytest.raises(ValueError, match="effort='max' is not supported by this model"): optional_params = {"output_config": {"effort": "max"}} config.transform_request( model="claude-opus-4-5-20251101", @@ -2251,9 +2249,7 @@ def test_max_effort_rejected_for_sonnet_46(): config = AnthropicConfig() messages = [{"role": "user", "content": "Test"}] - with pytest.raises( - ValueError, match="effort='max' is not supported by this model" - ): + with pytest.raises(ValueError, match="effort='max' is not supported by this model"): config.transform_request( model="claude-sonnet-4-6-20260219", messages=messages, @@ -2315,6 +2311,30 @@ def test_effort_beta_header_not_injected_for_46_models(): assert result is False, f"is_effort_used should return False for {model}" +@pytest.mark.parametrize( + "model", + [ + "claude-opus-4-5-20251101", + "claude-opus-4-6-20250514", + "claude-sonnet-4-6-20260219", + "claude-opus-4-7", + ], +) +def test_reasoning_effort_none_omits_thinking_and_output_config(model): + """reasoning_effort="none" must omit thinking and output_config from the request.""" + config = AnthropicConfig() + + result = config.map_openai_params( + non_default_params={"reasoning_effort": "none"}, + optional_params={}, + model=model, + drop_params=False, + ) + + assert "thinking" not in result + assert "output_config" not in result + + def test_effort_beta_header_still_injected_for_older_models(): """ Test that is_effort_used still returns True for pre-4.6 models diff --git a/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_handler_output_config_passthrough.py b/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_handler_output_config_passthrough.py new file mode 100644 index 00000000000..615dc5cfebc --- /dev/null +++ b/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_handler_output_config_passthrough.py @@ -0,0 +1,232 @@ +""" +Regression tests for output_config passthrough through the Anthropic +``/v1/messages`` → ``/chat/completions`` adapter. + +Background — what was broken: +* When a client sent ``output_config`` to ``/v1/messages`` and the request + was routed to a non-Anthropic backend (Azure OpenAI, Fireworks, Bedrock + Nova, etc.), the adapter forwarded the raw Anthropic-shaped ``output_config`` + field as-is into the OpenAI-format ``completion_kwargs``. The non-Anthropic + backend then rejected the request with 400 "Extra inputs are not permitted". +* The translator above the re-merge already extracts the meaningful parts of + ``output_config`` (``format`` → ``response_format``, ``effort`` → + ``reasoning_effort`` for non-Claude targets), so re-adding the raw key was + always either redundant (Anthropic-family) or harmful (non-Anthropic). + +Tests cover (consolidating PRs #23706 and #22727): +1. ``output_config`` is excluded from the post-translation re-merge. +2. ``ANTHROPIC_ONLY_REQUEST_KEYS`` constant is exported and contains + ``output_config`` so future maintainers know where to extend it. +3. The translator-extracted fields (``response_format`` / ``reasoning_effort``) + are still present after the strip — the strip removes only the raw + Anthropic-shaped duplicate. +4. Helper-level coverage for empty ``extra_kwargs`` (PR #22727 Greptile P2 — + the original ``or {}`` pattern silently substituted a default and prevented + the fallback inference path from being exercised). +""" + +import os +import sys +from unittest.mock import MagicMock, patch + +import pytest + +# Anchor sys.path to this file's location — not the working-directory-relative +# pattern Greptile flagged on PR #23706. Resolves correctly regardless of +# where pytest is invoked from. +sys.path.insert( + 0, os.path.abspath(os.path.join(os.path.dirname(__file__), "../../../../../..")) +) + +from litellm.llms.anthropic.experimental_pass_through.adapters.handler import ( + ANTHROPIC_ONLY_REQUEST_KEYS, + LiteLLMMessagesToCompletionTransformationHandler, +) + +MESSAGES = [{"role": "user", "content": "hello"}] + + +def _call_prepare(extra_kwargs, model="gpt-4o", output_format=None, **overrides): + """ + Drive ``_prepare_completion_kwargs`` with the minimum scaffolding needed. + + ``output_format`` is a top-level parameter on the function, so callers + pass it explicitly here rather than tucking it into ``extra_kwargs``. + + Uses an explicit-None check on ``extra_kwargs`` so callers can test the + falsy-empty-dict path. The fallback ``or {}`` pattern PR #22727 used here + masked the no-extra-kwargs case from ever exercising the test's intent. + """ + return LiteLLMMessagesToCompletionTransformationHandler._prepare_completion_kwargs( + max_tokens=overrides.get("max_tokens", 1024), + messages=overrides.get("messages", MESSAGES), + model=model, + metadata=None, + stop_sequences=None, + stream=False, + system=None, + temperature=None, + thinking=None, + tool_choice=None, + tools=None, + top_k=None, + top_p=None, + output_format=output_format, + extra_kwargs=extra_kwargs, + ) + + +class TestAnthropicOnlyRequestKeysExport: + """The exclusion list must be a public, named constant for maintainability — + Greptile P2 on PR #23706: ``excluded_keys`` was silently growing as a + point-fix pattern. A named module-level constant gives reviewers a single + grep target when extending Anthropic-only fields.""" + + def test_constant_exposed(self): + assert isinstance(ANTHROPIC_ONLY_REQUEST_KEYS, frozenset) + + def test_contains_output_config(self): + assert "output_config" in ANTHROPIC_ONLY_REQUEST_KEYS + + +class TestOutputConfigStrippedFromCompletionKwargs: + """``output_config`` must not survive the post-translation re-merge into + ``completion_kwargs`` regardless of the target provider — the translator + has already consumed its meaningful parts.""" + + def test_output_config_with_effort_is_stripped(self): + extra_kwargs = { + "custom_llm_provider": "azure", + "output_config": {"effort": "high"}, + } + + result = _call_prepare(extra_kwargs=extra_kwargs) + + # Returns (completion_kwargs, original_messages, ...) — first element + # is the dict we care about. + completion_kwargs = result[0] if isinstance(result, tuple) else result + assert "output_config" not in completion_kwargs, ( + "Raw output_config must not be forwarded — non-Anthropic backends " + "reject it with 400 'Extra inputs are not permitted'" + ) + + def test_output_config_format_translated_to_response_format(self): + """When ``output_config`` carries structured-output ``format``, the + translator now maps it to OpenAI's ``response_format`` so non-Anthropic + backends see the schema in their native shape. The raw + ``output_config`` key is still stripped from ``completion_kwargs`` — + only the translated ``response_format`` survives. + + Before this PR, only the legacy top-level ``output_format`` was + translated; ``output_config.format`` was silently dropped on the + adapter path even when the schema was correctly supplied (issue + flagged by Greptile review of the initial fix). + """ + schema = { + "type": "object", + "additionalProperties": False, + "properties": {"name": {"type": "string"}}, + } + extra_kwargs = { + "custom_llm_provider": "azure", + "output_config": {"format": {"type": "json_schema", "schema": schema}}, + } + + result = _call_prepare(extra_kwargs=extra_kwargs) + completion_kwargs = result[0] if isinstance(result, tuple) else result + + # Raw Anthropic-shaped key is gone (would 400 on non-Anthropic backends). + assert "output_config" not in completion_kwargs + # Translated OpenAI-shaped key is present so the schema actually + # reaches the downstream backend. + assert "response_format" in completion_kwargs, ( + "output_config.format must be translated to response_format — " + "without this, structured-output schemas are silently dropped on " + "the adapter path" + ) + + def test_output_format_top_level_still_translates(self): + """Regression guard: the legacy top-level ``output_format`` field must + continue to translate to ``response_format``. The new + ``output_config.format`` path must not break this existing behavior.""" + schema = {"type": "object", "properties": {"name": {"type": "string"}}} + result = _call_prepare( + extra_kwargs={"custom_llm_provider": "azure"}, + output_format={"type": "json_schema", "schema": schema}, + ) + completion_kwargs = result[0] if isinstance(result, tuple) else result + + assert "response_format" in completion_kwargs + + def test_output_format_takes_precedence_over_output_config_format(self): + """When both top-level ``output_format`` and ``output_config.format`` + are present, the legacy top-level ``output_format`` wins. Documents + which one the translator picks rather than leaving it implementation- + defined.""" + winning_schema = { + "type": "object", + "properties": {"top_level": {"type": "string"}}, + } + losing_schema = { + "type": "object", + "properties": {"nested": {"type": "string"}}, + } + result = _call_prepare( + extra_kwargs={ + "custom_llm_provider": "azure", + "output_config": { + "format": {"type": "json_schema", "schema": losing_schema} + }, + }, + output_format={"type": "json_schema", "schema": winning_schema}, + ) + completion_kwargs = result[0] if isinstance(result, tuple) else result + + assert "response_format" in completion_kwargs + # Verify the winning_schema (top-level output_format) was used, + # not the losing one nested under output_config. + rendered = str(completion_kwargs["response_format"]) + assert "top_level" in rendered + assert "nested" not in rendered + + def test_other_extra_kwargs_still_passed_through(self): + """Regression guard: the strip must be narrow. Unrelated fields like + ``api_key`` / ``timeout`` continue to flow through.""" + extra_kwargs = { + "custom_llm_provider": "azure", + "output_config": {"effort": "high"}, + "timeout": 30, + "user": "end-user-123", + } + + result = _call_prepare(extra_kwargs=extra_kwargs) + completion_kwargs = result[0] if isinstance(result, tuple) else result + + assert "output_config" not in completion_kwargs + assert completion_kwargs.get("timeout") == 30 + assert completion_kwargs.get("user") == "end-user-123" + + +class TestEmptyExtraKwargsPath: + """Greptile P2 on PR #22727: ``extra_kwargs or {default}`` substitutes a + default for an explicitly-passed empty dict, hiding the no-extra-kwargs + path. The new explicit-None pattern lets ``extra_kwargs={}`` reach the + code under test as written.""" + + def test_explicit_empty_dict_does_not_substitute_default(self): + # Explicit empty dict must be honored — not silently replaced with a + # default that adds back a custom_llm_provider this test wants absent. + result = _call_prepare(extra_kwargs={}) + completion_kwargs = result[0] if isinstance(result, tuple) else result + + # No output_config because nothing supplied it. + assert "output_config" not in completion_kwargs + + def test_none_extra_kwargs_handled_safely(self): + """The signature documents ``extra_kwargs: Optional[Dict] = None``; + passing None must not crash with KeyError or AttributeError.""" + result = _call_prepare(extra_kwargs=None) + # Just exercising the path; assert no exception and we get back a + # dict-like result. + completion_kwargs = result[0] if isinstance(result, tuple) else result + assert isinstance(completion_kwargs, dict) diff --git a/tests/test_litellm/llms/azure/test_azure_cost_calculation.py b/tests/test_litellm/llms/azure/test_azure_cost_calculation.py new file mode 100644 index 00000000000..53c91032b34 --- /dev/null +++ b/tests/test_litellm/llms/azure/test_azure_cost_calculation.py @@ -0,0 +1,75 @@ +""" +Test Azure OpenAI cost calculator — service_tier pricing. +""" + +import pytest + +import litellm +from litellm.llms.azure.cost_calculation import cost_per_token +from litellm.types.utils import Usage + + +# Register a test model with tier-specific pricing +TEST_MODEL = "test-azure-gpt-4.1" +TEST_MODEL_COST = { + TEST_MODEL: { + "input_cost_per_token": 0.001, + "output_cost_per_token": 0.002, + "input_cost_per_token_priority": 0.01, + "output_cost_per_token_priority": 0.02, + "input_cost_per_token_flex": 0.0005, + "output_cost_per_token_flex": 0.001, + "litellm_provider": "azure", + "max_tokens": 8192, + } +} + + +class TestAzureServiceTierCostCalculation: + """Test that service_tier is passed through Azure cost calculation.""" + + @pytest.fixture(autouse=True) + def register_test_model(self): + litellm.register_model(model_cost=TEST_MODEL_COST) + + def test_service_tier_priority_higher_cost(self): + """Priority tier should cost more than standard.""" + usage = Usage(prompt_tokens=1000, completion_tokens=500, total_tokens=1500) + + standard_prompt, standard_completion = cost_per_token( + model=TEST_MODEL, usage=usage + ) + priority_prompt, priority_completion = cost_per_token( + model=TEST_MODEL, usage=usage, service_tier="priority" + ) + + assert priority_prompt > standard_prompt + assert priority_completion > standard_completion + + def test_service_tier_flex_lower_cost(self): + """Flex tier should cost less than standard.""" + usage = Usage(prompt_tokens=1000, completion_tokens=500, total_tokens=1500) + + standard_prompt, standard_completion = cost_per_token( + model=TEST_MODEL, usage=usage + ) + flex_prompt, flex_completion = cost_per_token( + model=TEST_MODEL, usage=usage, service_tier="flex" + ) + + assert flex_prompt < standard_prompt + assert flex_completion < standard_completion + + def test_service_tier_none_returns_standard(self): + """service_tier=None should return standard pricing.""" + usage = Usage(prompt_tokens=1000, completion_tokens=500, total_tokens=1500) + + none_prompt, none_completion = cost_per_token( + model=TEST_MODEL, usage=usage, service_tier=None + ) + standard_prompt, standard_completion = cost_per_token( + model=TEST_MODEL, usage=usage, service_tier="standard" + ) + + assert abs(none_prompt - standard_prompt) < 1e-10 + assert abs(none_completion - standard_completion) < 1e-10 diff --git a/tests/test_litellm/llms/azure_ai/test_azure_ai_cost_calculator.py b/tests/test_litellm/llms/azure_ai/test_azure_ai_cost_calculator.py index 37add41b83f..20260c744f8 100644 --- a/tests/test_litellm/llms/azure_ai/test_azure_ai_cost_calculator.py +++ b/tests/test_litellm/llms/azure_ai/test_azure_ai_cost_calculator.py @@ -451,3 +451,51 @@ class TestAzureModelRouterCostBreakdown: assert logging_obj.cost_breakdown["additional_costs"][ "Azure Model Router Flat Cost" ] == pytest.approx(expected_flat_cost, rel=1e-9) + + +class TestAzureAIServiceTierCostCalculation: + """Test that service_tier is passed through Azure AI cost calculation.""" + + @pytest.fixture(autouse=True) + def register_test_model(self): + import litellm + litellm.register_model(model_cost={ + "test-azure-ai-model": { + "input_cost_per_token": 0.001, + "output_cost_per_token": 0.002, + "input_cost_per_token_priority": 0.01, + "output_cost_per_token_priority": 0.02, + "input_cost_per_token_flex": 0.0005, + "output_cost_per_token_flex": 0.001, + "litellm_provider": "azure_ai", + "max_tokens": 8192, + } + }) + + def test_service_tier_priority_higher_cost(self): + """Priority tier should cost more than standard for azure_ai.""" + usage = Usage(prompt_tokens=1000, completion_tokens=500, total_tokens=1500) + + standard_prompt, standard_completion = cost_per_token( + model="test-azure-ai-model", usage=usage + ) + priority_prompt, priority_completion = cost_per_token( + model="test-azure-ai-model", usage=usage, service_tier="priority" + ) + + assert priority_prompt > standard_prompt + assert priority_completion > standard_completion + + def test_service_tier_flex_lower_cost(self): + """Flex tier should cost less than standard for azure_ai.""" + usage = Usage(prompt_tokens=1000, completion_tokens=500, total_tokens=1500) + + standard_prompt, standard_completion = cost_per_token( + model="test-azure-ai-model", usage=usage + ) + flex_prompt, flex_completion = cost_per_token( + model="test-azure-ai-model", usage=usage, service_tier="flex" + ) + + assert flex_prompt < standard_prompt + assert flex_completion < standard_completion diff --git a/tests/test_litellm/llms/bedrock/chat/test_converse_transformation.py b/tests/test_litellm/llms/bedrock/chat/test_converse_transformation.py index 589bcb4fc0e..5bb51ac619c 100644 --- a/tests/test_litellm/llms/bedrock/chat/test_converse_transformation.py +++ b/tests/test_litellm/llms/bedrock/chat/test_converse_transformation.py @@ -288,6 +288,28 @@ def test_reasoning_with_forced_tool_choice_switches_to_auto(): assert optional_params["tool_choice"] == {"auto": {}} +@pytest.mark.parametrize( + "model", + [ + "bedrock/converse/us.anthropic.claude-opus-4-5-20251101-v1:0", + "bedrock/converse/us.anthropic.claude-opus-4-6-v1", + "bedrock/converse/us.anthropic.claude-opus-4-7", + ], +) +def test_reasoning_effort_none_omits_thinking_for_anthropic_converse(model): + """reasoning_effort="none" must omit thinking from the Bedrock Converse request.""" + config = AmazonConverseConfig() + + optional_params = config.map_openai_params( + non_default_params={"reasoning_effort": "none"}, + optional_params={}, + model=model, + drop_params=False, + ) + + assert "thinking" not in optional_params + + def test_get_supported_openai_params(): config = AmazonConverseConfig() supported_params = config.get_supported_openai_params( diff --git a/tests/test_litellm/llms/crusoe/__init__.py b/tests/test_litellm/llms/crusoe/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/test_litellm/llms/crusoe/test_crusoe.py b/tests/test_litellm/llms/crusoe/test_crusoe.py new file mode 100644 index 00000000000..0a05126919a --- /dev/null +++ b/tests/test_litellm/llms/crusoe/test_crusoe.py @@ -0,0 +1,135 @@ +import os +from unittest.mock import patch + +CRUSOE_API_BASE = "https://managed-inference-api-proxy.crusoecloud.com/v1" + + +def test_crusoe_json_registry(): + """Test Crusoe is registered in the JSON provider registry""" + from litellm.llms.openai_like.json_loader import JSONProviderRegistry + + assert JSONProviderRegistry.exists("crusoe") + config = JSONProviderRegistry.get("crusoe") + assert config is not None + assert config.base_url == CRUSOE_API_BASE + assert config.api_key_env == "CRUSOE_API_KEY" + assert config.api_base_env == "CRUSOE_API_BASE" + + +def test_crusoe_dynamic_config_defaults(): + """Test dynamic config returns correct default API base""" + from litellm.llms.openai_like.dynamic_config import create_config_class + from litellm.llms.openai_like.json_loader import JSONProviderRegistry + + config = create_config_class(JSONProviderRegistry.get("crusoe"))() + + with patch.dict(os.environ, {}, clear=True): + api_base, api_key = config._get_openai_compatible_provider_info(None, None) + + assert api_base == CRUSOE_API_BASE + assert api_key is None + + +def test_crusoe_dynamic_config_env_vars(): + """Test dynamic config reads CRUSOE_API_KEY and CRUSOE_API_BASE from env""" + from litellm.llms.openai_like.dynamic_config import create_config_class + from litellm.llms.openai_like.json_loader import JSONProviderRegistry + + config = create_config_class(JSONProviderRegistry.get("crusoe"))() + + with patch.dict( + os.environ, + {"CRUSOE_API_KEY": "test-key", "CRUSOE_API_BASE": "https://custom.crusoe.com/v1"}, + ): + api_base, api_key = config._get_openai_compatible_provider_info(None, None) + + assert api_base == "https://custom.crusoe.com/v1" + assert api_key == "test-key" + + +def test_crusoe_dynamic_config_explicit_params(): + """Test explicit params override env vars""" + from litellm.llms.openai_like.dynamic_config import create_config_class + from litellm.llms.openai_like.json_loader import JSONProviderRegistry + + config = create_config_class(JSONProviderRegistry.get("crusoe"))() + + with patch.dict(os.environ, {"CRUSOE_API_KEY": "env-key"}): + api_base, api_key = config._get_openai_compatible_provider_info( + "https://override.crusoe.com/v1", "override-key" + ) + + assert api_base == "https://override.crusoe.com/v1" + assert api_key == "override-key" + + +def test_crusoe_supported_params(): + """Test dynamic config returns standard OpenAI params""" + from litellm.llms.openai_like.dynamic_config import create_config_class + from litellm.llms.openai_like.json_loader import JSONProviderRegistry + + config = create_config_class(JSONProviderRegistry.get("crusoe"))() + params = config.get_supported_openai_params(model="meta-llama/Llama-3.3-70B-Instruct") + + assert isinstance(params, list) + assert len(params) > 0 + assert "temperature" in params + assert "max_tokens" in params + assert "stream" in params + + +def test_crusoe_param_mapping_max_completion_tokens(): + """Test max_completion_tokens is mapped to max_tokens for Crusoe""" + from litellm.llms.openai_like.dynamic_config import create_config_class + from litellm.llms.openai_like.json_loader import JSONProviderRegistry + + config = create_config_class(JSONProviderRegistry.get("crusoe"))() + optional_params = config.map_openai_params( + non_default_params={"max_completion_tokens": 1024}, + optional_params={}, + model="meta-llama/Llama-3.3-70B-Instruct", + drop_params=False, + ) + + assert "max_tokens" in optional_params, "max_completion_tokens should be mapped to max_tokens" + assert optional_params["max_tokens"] == 1024 + assert "max_completion_tokens" not in optional_params + + +def test_crusoe_provider_detection_by_prefix(): + """Test crusoe/model prefix is correctly routed""" + from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider + + model, provider, _, _ = get_llm_provider("crusoe/meta-llama/Llama-3.3-70B-Instruct") + assert provider == "crusoe" + assert model == "meta-llama/Llama-3.3-70B-Instruct" + + +def test_crusoe_model_list_populated(): + """Test Crusoe models are present in model_prices_and_context_window.json""" + import litellm + + original_model_cost = litellm.model_cost + original_env = os.environ.get("LITELLM_LOCAL_MODEL_COST_MAP") + try: + os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True" + litellm.model_cost = litellm.get_model_cost_map(url="") + + expected = [ + "crusoe/meta-llama/Llama-3.3-70B-Instruct", + "crusoe/deepseek-ai/DeepSeek-R1-0528", + "crusoe/deepseek-ai/DeepSeek-V3-0324", + "crusoe/Qwen/Qwen3-235B-A22B-Instruct-2507", + "crusoe/moonshotai/Kimi-K2-Thinking", + "crusoe/openai/gpt-oss-120b", + "crusoe/google/gemma-3-12b-it", + ] + for model in expected: + assert model in litellm.model_cost, f"{model} not found in model_cost" + assert litellm.model_cost[model].get("litellm_provider") == "crusoe" + finally: + litellm.model_cost = original_model_cost + if original_env is None: + os.environ.pop("LITELLM_LOCAL_MODEL_COST_MAP", None) + else: + os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = original_env diff --git a/tests/test_litellm/llms/openai/test_gpt5_transformation.py b/tests/test_litellm/llms/openai/test_gpt5_transformation.py index 6383bfc9e18..ebf7681f2f3 100644 --- a/tests/test_litellm/llms/openai/test_gpt5_transformation.py +++ b/tests/test_litellm/llms/openai/test_gpt5_transformation.py @@ -536,6 +536,81 @@ def test_gpt5_unknown_model_passes_through_minimal(config: OpenAIConfig): assert params["reasoning_effort"] == "minimal" +def test_gpt5_5_pro_rejects_reasoning_effort_low(config: OpenAIConfig): + """gpt-5.5-pro only accepts {medium, high, xhigh} — 'low' must raise. + + Verified against OpenAI's live API: /v1/chat/completions with + reasoning_effort='low' on gpt-5.5-pro returns HTTP 400. + """ + with pytest.raises(litellm.utils.UnsupportedParamsError): + config.map_openai_params( + non_default_params={"reasoning_effort": "low"}, + optional_params={}, + model="gpt-5.5-pro", + drop_params=False, + ) + + +def test_gpt5_5_pro_dated_rejects_reasoning_effort_low(config: OpenAIConfig): + """Dated snapshot must inherit the base alias's low-rejection behavior.""" + with pytest.raises(litellm.utils.UnsupportedParamsError): + config.map_openai_params( + non_default_params={"reasoning_effort": "low"}, + optional_params={}, + model="gpt-5.5-pro-2026-04-23", + drop_params=False, + ) + + +def test_gpt5_5_pro_drops_reasoning_effort_low_when_requested(config: OpenAIConfig): + """drop_params=True silently strips 'low' instead of round-tripping a 400.""" + params = config.map_openai_params( + non_default_params={"reasoning_effort": "low"}, + optional_params={}, + model="gpt-5.5-pro", + drop_params=True, + ) + assert "reasoning_effort" not in params + + +def test_gpt5_5_chat_allows_reasoning_effort_low(config: OpenAIConfig): + """gpt-5.5 (chat) supports 'low'; flag absent → opt-out check passes.""" + params = config.map_openai_params( + non_default_params={"reasoning_effort": "low"}, + optional_params={}, + model="gpt-5.5", + drop_params=False, + ) + assert params["reasoning_effort"] == "low" + + +def test_gpt5_unknown_model_passes_through_low(config: OpenAIConfig): + """Unknown gpt-5 models pass 'low' through (opt-out, not opt-in).""" + params = config.map_openai_params( + non_default_params={"reasoning_effort": "low"}, + optional_params={}, + model="gpt-5.4-turbo-preview", + drop_params=False, + ) + assert params["reasoning_effort"] == "low" + + +def test_gpt5_low_explicitly_disabled_check(gpt5_config: OpenAIGPT5Config): + """supports_low_reasoning_effort=false → disabled; missing/true → not disabled.""" + assert gpt5_config._is_reasoning_effort_level_explicitly_disabled( + "gpt-5.5-pro", "low" + ) + assert gpt5_config._is_reasoning_effort_level_explicitly_disabled( + "gpt-5.5-pro-2026-04-23", "low" + ) + assert not gpt5_config._is_reasoning_effort_level_explicitly_disabled( + "gpt-5.5", "low" + ) + assert not gpt5_config._is_reasoning_effort_level_explicitly_disabled( + "gpt-5.4", "low" + ) + + def test_gpt5_normalizes_reasoning_effort_dict_with_summary(config: OpenAIConfig): """Dict with summary/generate_summary is normalized for chat completions.""" params = config.map_openai_params( diff --git a/tests/test_litellm/llms/vertex_ai/files/test_vertex_ai_files_transformation.py b/tests/test_litellm/llms/vertex_ai/files/test_vertex_ai_files_transformation.py index a41fafa0e87..7c063c72607 100644 --- a/tests/test_litellm/llms/vertex_ai/files/test_vertex_ai_files_transformation.py +++ b/tests/test_litellm/llms/vertex_ai/files/test_vertex_ai_files_transformation.py @@ -1,7 +1,9 @@ """ Tests for VertexAIFilesConfig transformation methods (Issues 5-7). +Includes tests for Vertex AI batch output transformation to OpenAI format. """ +import json import urllib.parse from types import MappingProxyType from urllib.parse import parse_qs, urlparse @@ -10,7 +12,12 @@ import httpx import pytest from unittest.mock import MagicMock -from litellm.llms.vertex_ai.files.transformation import VertexAIFilesConfig +from litellm.llms.vertex_ai.files.transformation import ( + VertexAIFilesConfig, + VertexAIJsonlFilesTransformation, + _get_litellm_batch_custom_id_from_labels, + _sanitize_gcp_label_value, +) from litellm.types.llms.openai import OpenAIFileObject, HttpxBinaryResponseContent from openai.types.file_deleted import FileDeleted @@ -262,6 +269,108 @@ class TestTransformFileContent: assert isinstance(result, HttpxBinaryResponseContent) assert result.response.content == b'{"line": 1}\n{"line": 2}\n' + def test_should_not_mutate_caller_logging_obj_for_batch_output_transform( + self, config, monkeypatch + ): + original_model = "vertex_ai/original-model" + original_start_time = 123.456 + original_optional_params = {"temperature": 0.1} + raw_response = httpx.Response( + status_code=200, + content=json.dumps( + { + "status": "", + "processed_time": "2024-11-01T18:13:16.826+00:00", + "request": {"labels": {"litellm_custom_id": "request-1"}}, + "response": { + "candidates": [ + {"content": {"parts": [{"text": "ok"}], "role": "model"}} + ], + "modelVersion": "gemini-2.0-flash-001@default", + }, + } + ).encode("utf-8"), + headers={"content-type": "application/octet-stream"}, + request=httpx.Request("GET", "https://example.com"), + ) + logging_obj = MagicMock() + logging_obj.model = original_model + logging_obj.start_time = original_start_time + logging_obj.optional_params = original_optional_params + captured = {} + + def mock_transform_single( + vertex_output, + vertex_gemini_config, + logging_obj, + mock_httpx_response, + ): + captured["logging_obj"] = logging_obj + logging_obj.model = "gemini-2.0-flash-001" + logging_obj.start_time = 789.0 + return { + "custom_id": vertex_output["request"]["labels"]["litellm_custom_id"] + } + + monkeypatch.setattr( + config, + "_transform_single_vertex_batch_output_to_openai", + mock_transform_single, + ) + + result = config.transform_file_content_response( + raw_response=raw_response, + logging_obj=logging_obj, + litellm_params={}, + ) + + assert captured["logging_obj"] is not logging_obj + assert logging_obj.model == original_model + assert logging_obj.start_time == original_start_time + assert logging_obj.optional_params == original_optional_params + assert result.response is not raw_response + + def test_should_skip_batch_output_transformation_when_opt_out_flag_set( + self, config, monkeypatch + ): + """When `litellm.disable_vertex_batch_output_transformation` is True the + Vertex predictions.jsonl content must be returned untouched, so callers + that parse raw `candidates`/`modelVersion` keep working.""" + import litellm + + raw_jsonl = json.dumps( + { + "status": "", + "processed_time": "2024-11-01T18:13:16.826+00:00", + "request": {"labels": {"litellm_custom_id": "request-1"}}, + "response": { + "candidates": [ + {"content": {"parts": [{"text": "ok"}], "role": "model"}} + ], + "modelVersion": "gemini-2.0-flash-001@default", + }, + } + ).encode("utf-8") + raw_response = httpx.Response( + status_code=200, + content=raw_jsonl, + headers={"content-type": "application/octet-stream"}, + request=httpx.Request("GET", "https://example.com"), + ) + + monkeypatch.setattr( + litellm, "disable_vertex_batch_output_transformation", True, raising=False + ) + + result = config.transform_file_content_response( + raw_response=raw_response, + logging_obj=MagicMock(), + litellm_params={}, + ) + + assert isinstance(result, HttpxBinaryResponseContent) + assert result.response.content == raw_jsonl + class TestTransformDeleteFile: def test_should_build_correct_gcs_delete_url(self, config): @@ -360,3 +469,726 @@ class TestTransformDeleteFile: "gs://prod-bucket/litellm-vertex-files/publishers/google/" "models/gemini-2.0-flash-001/abc-123" ) + + +class TestVertexBatchOutputTransformation: + """Test transformation of Vertex AI batch outputs to OpenAI format""" + + def test_transform_successful_vertex_batch_output(self, config): + """Test transformation of a successful Vertex AI batch output""" + # Sample Vertex AI batch output (based on actual format) + vertex_output = { + "status": "", + "processed_time": "2024-11-01T18:13:16.826+00:00", + "request": { + "contents": [{"role": "user", "parts": [{"text": "Hello world!"}]}], + "labels": {"litellm_custom_id": "request-1"}, + }, + "response": { + "candidates": [ + { + "content": { + "parts": [{"text": "Hello! How can I help you today?"}], + "role": "model", + }, + "finishReason": "STOP", + } + ], + "modelVersion": "gemini-2.0-flash-001@default", + "usageMetadata": { + "promptTokenCount": 10, + "candidatesTokenCount": 20, + "totalTokenCount": 30, + }, + }, + } + + content = json.dumps(vertex_output).encode("utf-8") + transformed_content = config._try_transform_vertex_batch_output_to_openai( + content + ) + result = json.loads(transformed_content.decode("utf-8")) + + # Verify OpenAI format + assert "id" in result + assert "custom_id" in result + assert "response" in result + assert "error" in result + + # Verify custom_id was extracted from labels + assert result["custom_id"] == "request-1" + + # Verify response structure + assert result["response"]["status_code"] == 200 + assert "body" in result["response"] + + # Verify body has OpenAI format + body = result["response"]["body"] + assert "choices" in body + assert "usage" in body + assert "model" in body + + # Verify choices + assert len(body["choices"]) > 0 + choice = body["choices"][0] + assert "message" in choice + assert "content" in choice["message"] + assert "Hello! How can I help you today?" in choice["message"]["content"] + + def test_transform_error_vertex_batch_output(self, config): + """Test transformation of an error Vertex AI batch output""" + vertex_output = { + "status": "Error: Invalid request", + "processed_time": "2024-11-01T18:13:16.826+00:00", + "request": { + "contents": [{"role": "user", "parts": [{"text": "Hello world!"}]}], + "labels": {"litellm_custom_id": "request-error"}, + }, + "response": {}, + } + + content = json.dumps(vertex_output).encode("utf-8") + transformed_content = config._try_transform_vertex_batch_output_to_openai( + content + ) + result = json.loads(transformed_content.decode("utf-8")) + + # Per OpenAI Batch output spec, error entries set response to null + # and populate the top-level error object. + assert result["response"] is None + assert result["error"] is not None + assert "Invalid request" in result["error"]["message"] + assert result["error"]["code"] == "vertex_ai_error" + assert result["custom_id"] == "request-error" + + def test_transform_exception_path_sets_response_null(self, config): + """ + The except-Exception branch in _transform_single_vertex_batch_output_to_openai + must also emit response=null per the OpenAI Batch output spec. The outer + _try_transform path swallows exceptions and falls back to original content, + so this test invokes the single-line transformer directly with a vertex_gemini_config + stub that raises during transformation. + """ + from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import ( + VertexGeminiConfig, + ) + + vertex_output = { + "status": "", + "processed_time": "2024-11-01T18:13:16.826+00:00", + "request": { + "contents": [{"role": "user", "parts": [{"text": "Hello world!"}]}], + "labels": {"litellm_custom_id": "request-boom"}, + }, + "response": {"modelVersion": "gemini-2.0-flash-001@default"}, + } + + class _RaisingGeminiConfig(VertexGeminiConfig): + def _transform_google_generate_content_to_openai_model_response( + self, *args, **kwargs + ): + raise ValueError("simulated transform failure") + + mock_response = httpx.Response( + status_code=200, + headers={"content-type": "application/json"}, + request=httpx.Request(method="POST", url="https://example.com"), + ) + + result = config._transform_single_vertex_batch_output_to_openai( + vertex_output=vertex_output, + vertex_gemini_config=_RaisingGeminiConfig(), + logging_obj=MagicMock(), + mock_httpx_response=mock_response, + ) + + assert result["response"] is None + assert result["error"] is not None + assert result["error"]["code"] == "transformation_error" + assert "simulated transform failure" in result["error"]["message"] + assert result["custom_id"] == "request-boom" + + def test_transform_vertex_batch_output_legacy_labels_only_sanitized(self, config): + """Older LiteLLM batches only stored litellm_custom_id (sanitized); read path still works.""" + vertex_output = { + "status": "", + "processed_time": "2024-11-01T18:13:16.826+00:00", + "request": { + "contents": [{"role": "user", "parts": [{"text": "Hello world!"}]}], + "labels": {"litellm_custom_id": "myrequest-1"}, + }, + "response": { + "candidates": [ + { + "content": { + "parts": [{"text": "Hello!"}], + "role": "model", + }, + "finishReason": "STOP", + } + ], + "modelVersion": "gemini-2.0-flash-001@default", + "usageMetadata": { + "promptTokenCount": 10, + "candidatesTokenCount": 20, + "totalTokenCount": 30, + }, + }, + } + + content = json.dumps(vertex_output).encode("utf-8") + transformed_content = config._try_transform_vertex_batch_output_to_openai( + content + ) + result = json.loads(transformed_content.decode("utf-8")) + + assert result["custom_id"] == "myrequest-1" + + def test_transform_multiple_vertex_batch_outputs(self, config): + """Test transformation of multiple Vertex AI batch outputs (JSONL)""" + vertex_outputs = [ + { + "status": "", + "processed_time": "2024-11-01T18:13:16.826+00:00", + "request": { + "contents": [ + {"role": "user", "parts": [{"text": "First request"}]} + ], + "labels": {"litellm_custom_id": "request-1"}, + }, + "response": { + "candidates": [ + { + "content": { + "parts": [{"text": "First response"}], + "role": "model", + }, + "finishReason": "STOP", + } + ], + "modelVersion": "gemini-2.0-flash-001@default", + "usageMetadata": { + "promptTokenCount": 5, + "candidatesTokenCount": 10, + "totalTokenCount": 15, + }, + }, + }, + { + "status": "", + "processed_time": "2024-11-01T18:13:17.826+00:00", + "request": { + "contents": [ + {"role": "user", "parts": [{"text": "Second request"}]} + ], + "labels": {"litellm_custom_id": "request-2"}, + }, + "response": { + "candidates": [ + { + "content": { + "parts": [{"text": "Second response"}], + "role": "model", + }, + "finishReason": "STOP", + } + ], + "modelVersion": "gemini-2.0-flash-001@default", + "usageMetadata": { + "promptTokenCount": 6, + "candidatesTokenCount": 11, + "totalTokenCount": 17, + }, + }, + }, + ] + + content = "\n".join(json.dumps(output) for output in vertex_outputs).encode( + "utf-8" + ) + transformed_content = config._try_transform_vertex_batch_output_to_openai( + content + ) + lines = transformed_content.decode("utf-8").strip().split("\n") + + assert len(lines) == 2 + + for i, line in enumerate(lines): + result = json.loads(line) + assert "id" in result + assert "response" in result + assert result["response"]["status_code"] == 200 + assert result["custom_id"] == f"request-{i+1}" + body = result["response"]["body"] + assert "choices" in body + assert len(body["choices"]) > 0 + + def test_transform_vertex_batch_output_with_first_line_prompt_feedback( + self, config, monkeypatch + ): + """Test that promptFeedback-only first lines are detected as Vertex batch output.""" + vertex_outputs = [ + { + "status": "", + "processed_time": "2024-11-01T18:13:16.826+00:00", + "request": {"labels": {"litellm_custom_id": "blocked-request"}}, + "response": { + "promptFeedback": {"blockReason": "SAFETY"}, + "modelVersion": "gemini-2.0-flash-001@default", + }, + }, + { + "status": "", + "processed_time": "2024-11-01T18:13:17.826+00:00", + "request": {"labels": {"litellm_custom_id": "request-2"}}, + "response": {"candidates": [{"content": {"parts": [{"text": "ok"}]}}]}, + }, + ] + + def mock_transform_single( + vertex_output, + vertex_gemini_config, + logging_obj, + mock_httpx_response, + ): + return { + "custom_id": vertex_output["request"]["labels"]["litellm_custom_id"] + } + + monkeypatch.setattr( + config, + "_transform_single_vertex_batch_output_to_openai", + mock_transform_single, + ) + + content = "\n".join(json.dumps(output) for output in vertex_outputs).encode( + "utf-8" + ) + transformed_content = config._try_transform_vertex_batch_output_to_openai( + content + ) + results = [ + json.loads(line) for line in transformed_content.decode("utf-8").split("\n") + ] + + assert [result["custom_id"] for result in results] == [ + "blocked-request", + "request-2", + ] + + def test_batch_detection_requires_candidates_or_non_empty_status(self, config): + """Test that JSONL with a blank status but no candidates is returned as-is.""" + non_batch_output = { + "status": "", + "processed_time": "2024-11-01T18:13:16.826+00:00", + "request": {"metadata": "not a Vertex batch request"}, + "response": {"metadata": "not a Gemini response"}, + } + + content = json.dumps(non_batch_output).encode("utf-8") + transformed_content = config._try_transform_vertex_batch_output_to_openai( + content + ) + + assert transformed_content == content + + def test_reuses_batch_transform_helpers_per_jsonl_file(self, config, monkeypatch): + """Test that heavy helper objects are reused while transforming a JSONL file.""" + vertex_outputs = [ + { + "status": "", + "processed_time": "2024-11-01T18:13:16.826+00:00", + "request": {"labels": {"litellm_custom_id": f"request-{i}"}}, + "response": {"candidates": [{"content": {"parts": [{"text": "ok"}]}}]}, + } + for i in range(2) + ] + helper_ids = [] + + def mock_transform_single( + vertex_output, + vertex_gemini_config, + logging_obj, + mock_httpx_response, + ): + helper_ids.append( + ( + id(vertex_gemini_config), + id(logging_obj), + id(mock_httpx_response), + ) + ) + return { + "custom_id": vertex_output["request"]["labels"]["litellm_custom_id"] + } + + monkeypatch.setattr( + config, + "_transform_single_vertex_batch_output_to_openai", + mock_transform_single, + ) + + content = "\n".join(json.dumps(output) for output in vertex_outputs).encode( + "utf-8" + ) + transformed_content = config._try_transform_vertex_batch_output_to_openai( + content + ) + + assert len(transformed_content.decode("utf-8").strip().split("\n")) == 2 + assert len(set(helper_ids)) == 1 + + def test_non_batch_output_passthrough(self, config): + """Test that non-batch output is returned as-is""" + regular_content = b"This is just a regular file content" + transformed_content = config._try_transform_vertex_batch_output_to_openai( + regular_content + ) + assert transformed_content == regular_content + + def test_invalid_json_passthrough(self, config): + """Test that invalid JSON is returned as-is""" + invalid_content = b'{"invalid": json content}' + transformed_content = config._try_transform_vertex_batch_output_to_openai( + invalid_content + ) + assert transformed_content == invalid_content + + +class TestTryTransformDoesNotMutateCallerLoggingObj: + """Regression tests: _try_transform_vertex_batch_output_to_openai must not mutate + the caller's logging_obj (model, start_time, optional_params).""" + + def _make_vertex_batch_line(self) -> bytes: + return json.dumps( + { + "status": "", + "processed_time": "2024-11-01T18:13:16.826+00:00", + "request": { + "contents": [{"role": "user", "parts": [{"text": "Hello world!"}]}], + "labels": {"litellm_custom_id": "request-1"}, + }, + "response": { + "candidates": [ + { + "content": { + "parts": [{"text": "Hi!"}], + "role": "model", + }, + "finishReason": "STOP", + } + ], + "modelVersion": "gemini-2.0-flash-001@default", + "usageMetadata": { + "promptTokenCount": 5, + "candidatesTokenCount": 3, + "totalTokenCount": 8, + }, + }, + } + ).encode("utf-8") + + def test_should_not_overwrite_model_on_caller_logging_obj(self, config): + sentinel_model = "original-caller-model" + logging_obj = MagicMock() + logging_obj.model = sentinel_model + logging_obj.optional_params = {"temperature": 0.9} + + config._try_transform_vertex_batch_output_to_openai( + content=self._make_vertex_batch_line(), + logging_obj=logging_obj, + ) + + assert ( + logging_obj.model == sentinel_model + ), "logging_obj.model was mutated by _try_transform_vertex_batch_output_to_openai" + + def test_should_not_overwrite_start_time_on_caller_logging_obj(self, config): + sentinel_start = 1234567890.0 + logging_obj = MagicMock() + logging_obj.start_time = sentinel_start + logging_obj.optional_params = {} + + config._try_transform_vertex_batch_output_to_openai( + content=self._make_vertex_batch_line(), + logging_obj=logging_obj, + ) + + assert ( + logging_obj.start_time == sentinel_start + ), "logging_obj.start_time was mutated by _try_transform_vertex_batch_output_to_openai" + + def test_should_not_overwrite_optional_params_on_caller_logging_obj(self, config): + sentinel_params = {"temperature": 0.5, "top_p": 0.9} + logging_obj = MagicMock() + logging_obj.optional_params = sentinel_params + + config._try_transform_vertex_batch_output_to_openai( + content=self._make_vertex_batch_line(), + logging_obj=logging_obj, + ) + + assert ( + logging_obj.optional_params is sentinel_params + ), "logging_obj.optional_params was replaced by _try_transform_vertex_batch_output_to_openai" + assert logging_obj.optional_params == { + "temperature": 0.5, + "top_p": 0.9, + }, "logging_obj.optional_params contents were mutated" + + def test_should_still_transform_content_correctly(self, config): + logging_obj = MagicMock() + logging_obj.model = "original-model" + logging_obj.start_time = 9999.0 + logging_obj.optional_params = {"max_tokens": 100} + + result = config._try_transform_vertex_batch_output_to_openai( + content=self._make_vertex_batch_line(), + logging_obj=logging_obj, + ) + + # Transformation should still succeed + transformed = json.loads(result.decode("utf-8")) + assert transformed["custom_id"] == "request-1" + assert transformed["response"]["status_code"] == 200 + + +class TestVertexBatchCustomIdLabels: + """Test custom_id handling in batch transformations""" + + def test_custom_id_added_to_labels_in_vertex_request(self): + """Test that custom_id from OpenAI format is added as a label in Vertex AI format""" + transformation = VertexAIJsonlFilesTransformation() + + openai_jsonl_content = [ + { + "custom_id": "request-1", + "method": "POST", + "url": "/v1/chat/completions", + "body": { + "model": "gemini-1.5-flash-001", + "messages": [{"role": "user", "content": "What is 2+2?"}], + "max_tokens": 10, + }, + } + ] + + vertex_jsonl_content = ( + transformation._transform_openai_jsonl_content_to_vertex_ai_jsonl_content( + openai_jsonl_content + ) + ) + + assert len(vertex_jsonl_content) == 1 + vertex_request = vertex_jsonl_content[0] + + # Verify labels were added + assert "labels" in vertex_request["request"] + assert "litellm_custom_id" in vertex_request["request"]["labels"] + assert vertex_request["request"]["labels"]["litellm_custom_id"] == "request-1" + raw_label = vertex_request["request"]["labels"]["litellm_custom_id_raw"] + assert raw_label != "request-1" + assert _sanitize_gcp_label_value(raw_label) == raw_label + + def test_long_custom_id_round_trips_across_raw_label_chunks(self): + """Test that long custom_ids are not truncated in raw labels.""" + transformation = VertexAIJsonlFilesTransformation() + custom_id_a = "shared-prefix-that-is-longer-than-thirty-six-bytes-A" + custom_id_b = "shared-prefix-that-is-longer-than-thirty-six-bytes-B" + + openai_jsonl_content = [ + { + "custom_id": custom_id, + "method": "POST", + "url": "/v1/chat/completions", + "body": { + "model": "gemini-1.5-flash-001", + "messages": [{"role": "user", "content": "Question"}], + }, + } + for custom_id in (custom_id_a, custom_id_b) + ] + + vertex_jsonl_content = ( + transformation._transform_openai_jsonl_content_to_vertex_ai_jsonl_content( + openai_jsonl_content + ) + ) + labels_a = vertex_jsonl_content[0]["request"]["labels"] + labels_b = vertex_jsonl_content[1]["request"]["labels"] + + assert "litellm_custom_id_raw_1" in labels_a + assert "litellm_custom_id_raw_1" in labels_b + assert labels_a["litellm_custom_id_raw"] == labels_b["litellm_custom_id_raw"] + assert ( + labels_a["litellm_custom_id_raw_1"] != labels_b["litellm_custom_id_raw_1"] + ) + assert _get_litellm_batch_custom_id_from_labels(labels_a) == custom_id_a + assert _get_litellm_batch_custom_id_from_labels(labels_b) == custom_id_b + + def test_multiple_requests_each_get_their_own_label(self): + """Test that multiple requests each get their own custom_id label""" + transformation = VertexAIJsonlFilesTransformation() + + openai_jsonl_content = [ + { + "custom_id": f"request-{i+1}", + "method": "POST", + "url": "/v1/chat/completions", + "body": { + "model": "gemini-1.5-flash-001", + "messages": [{"role": "user", "content": f"Question {i+1}"}], + }, + } + for i in range(3) + ] + + vertex_jsonl_content = ( + transformation._transform_openai_jsonl_content_to_vertex_ai_jsonl_content( + openai_jsonl_content + ) + ) + + assert len(vertex_jsonl_content) == 3 + + for i, vertex_request in enumerate(vertex_jsonl_content): + expected_custom_id = f"request-{i+1}" + assert ( + vertex_request["request"]["labels"]["litellm_custom_id"] + == expected_custom_id + ) + raw_label = vertex_request["request"]["labels"]["litellm_custom_id_raw"] + assert raw_label != expected_custom_id + assert _sanitize_gcp_label_value(raw_label) == raw_label + + def test_request_without_custom_id_has_no_label(self): + """Test that requests without custom_id don't get a label""" + transformation = VertexAIJsonlFilesTransformation() + + openai_jsonl_content = [ + { + "method": "POST", + "url": "/v1/chat/completions", + "body": { + "model": "gemini-1.5-flash-001", + "messages": [{"role": "user", "content": "Question"}], + }, + } + ] + + vertex_jsonl_content = ( + transformation._transform_openai_jsonl_content_to_vertex_ai_jsonl_content( + openai_jsonl_content + ) + ) + + # Should not have labels if no custom_id was provided + assert "labels" not in vertex_jsonl_content[0]["request"] + + def test_end_to_end_custom_id_round_trip(self): + """ + Test the full round trip: OpenAI format -> Vertex AI format -> Vertex AI output -> OpenAI output + Verify that custom_id is preserved through the entire flow. + """ + transformation = VertexAIJsonlFilesTransformation() + config = VertexAIFilesConfig() + + # Step 1: Transform OpenAI input to Vertex AI format (mixed case exercises raw label) + openai_input = [ + { + "custom_id": "MyRequest-1", + "method": "POST", + "url": "/v1/chat/completions", + "body": { + "model": "gemini-1.5-flash-001", + "messages": [{"role": "user", "content": "Hello"}], + }, + } + ] + + vertex_input = ( + transformation._transform_openai_jsonl_content_to_vertex_ai_jsonl_content( + openai_input + ) + ) + + # Verify both labels are GCP-safe and encoded raw preserves round-trip. + assert ( + vertex_input[0]["request"]["labels"]["litellm_custom_id"] == "myrequest-1" + ) + raw_label = vertex_input[0]["request"]["labels"]["litellm_custom_id_raw"] + assert raw_label != "MyRequest-1" + assert _sanitize_gcp_label_value(raw_label) == raw_label + + # Step 2: Simulate Vertex AI batch output (with the label echoed back) + vertex_output = { + "status": "", + "processed_time": "2024-11-01T18:13:16.826+00:00", + "request": vertex_input[0]["request"], + "response": { + "candidates": [ + { + "content": {"parts": [{"text": "Hi there!"}], "role": "model"}, + "finishReason": "STOP", + } + ], + "modelVersion": "gemini-2.0-flash-001@default", + "usageMetadata": { + "promptTokenCount": 5, + "candidatesTokenCount": 10, + "totalTokenCount": 15, + }, + }, + } + + # Step 3: Transform Vertex AI output back to OpenAI format + content = json.dumps(vertex_output).encode("utf-8") + transformed_content = config._try_transform_vertex_batch_output_to_openai( + content + ) + openai_output = json.loads(transformed_content.decode("utf-8")) + + # Step 4: Verify custom_id was preserved (original casing, not sanitized label) + assert openai_output["custom_id"] == "MyRequest-1" + assert openai_output["response"]["status_code"] == 200 + + def test_custom_id_label_sanitization(self): + """Test that custom_id values are sanitized to meet GCP label constraints""" + transformation = VertexAIJsonlFilesTransformation() + + # Test sanitization function + assert _sanitize_gcp_label_value("MyRequest-1") == "myrequest-1" + assert _sanitize_gcp_label_value("Request.With.Dots") == "request_with_dots" + assert _sanitize_gcp_label_value("Request With Spaces") == "request_with_spaces" + assert _sanitize_gcp_label_value("Request@#$%Special") == "request____special" + + # Test max length (63 chars) + long_id = "a" * 100 + assert len(_sanitize_gcp_label_value(long_id)) == 63 + + # Test in actual transformation + openai_input = [ + { + "custom_id": "MyRequest-1", + "method": "POST", + "url": "/v1/chat/completions", + "body": { + "model": "gemini-1.5-flash-001", + "messages": [{"role": "user", "content": "Hello"}], + }, + } + ] + + vertex_input = ( + transformation._transform_openai_jsonl_content_to_vertex_ai_jsonl_content( + openai_input + ) + ) + + # Verify both labels are safe for GCP labels. + assert ( + vertex_input[0]["request"]["labels"]["litellm_custom_id"] == "myrequest-1" + ) + raw_label = vertex_input[0]["request"]["labels"]["litellm_custom_id_raw"] + assert raw_label != "MyRequest-1" + assert _sanitize_gcp_label_value(raw_label) == raw_label diff --git a/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_ai_gemini_transformation.py b/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_ai_gemini_transformation.py index 831d1ef464b..977c53280a9 100644 --- a/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_ai_gemini_transformation.py +++ b/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_ai_gemini_transformation.py @@ -86,6 +86,76 @@ def test_check_if_part_exists_in_parts_camel_case_snake_case(): assert check_if_part_exists_in_parts(parts_mixed, part_mixed_casing) +def test_cached_content_respects_modify_params_for_cache_incompatible_fields(): + """Regression: cachedContent drops system/tools/toolConfig only when modify_params=True.""" + import litellm + + cache_name = "projects/p/locations/us-central1/cachedContents/abc123" + messages = [ + {"role": "system", "content": "You are helpful"}, + {"role": "user", "content": "hi"}, + ] + optional_params = { + "tools": [ + { + "functionDeclarations": [ + {"name": "get_weather", "description": "Get weather"}, + ] + } + ], + "tool_choice": {"functionCallingConfig": {"mode": "AUTO"}}, + } + + original_modify_params = litellm.modify_params + try: + # With modify_params=False (default), keep fields even with cachedContent. + litellm.modify_params = False + result = _transform_request_body( + messages=list(messages), + model="gemini-2.5-pro", + optional_params=dict(optional_params), + custom_llm_provider="vertex_ai", + litellm_params={}, + cached_content=cache_name, + ) + assert result.get("cachedContent") == cache_name + assert "system_instruction" in result + assert "tools" in result + assert "toolConfig" in result + assert "contents" in result + + # With modify_params=True, drop cache-incompatible fields. + litellm.modify_params = True + result_modify_true = _transform_request_body( + messages=list(messages), + model="gemini-2.5-pro", + optional_params=dict(optional_params), + custom_llm_provider="vertex_ai", + litellm_params={}, + cached_content=cache_name, + ) + assert result_modify_true.get("cachedContent") == cache_name + assert "system_instruction" not in result_modify_true + assert "tools" not in result_modify_true + assert "toolConfig" not in result_modify_true + assert "contents" in result_modify_true + + # Without cache, fields are always included. + result_no_cache = _transform_request_body( + messages=list(messages), + model="gemini-2.5-pro", + optional_params=dict(optional_params), + custom_llm_provider="vertex_ai", + litellm_params={}, + cached_content=None, + ) + assert "system_instruction" in result_no_cache + assert "tools" in result_no_cache + assert "toolConfig" in result_no_cache + finally: + litellm.modify_params = original_modify_params + + # Tests for issue #14556: Labels field provider-aware filtering def test_google_genai_excludes_labels(): """Test that Google GenAI/AI Studio endpoints exclude labels when custom_llm_provider='gemini'""" @@ -1291,6 +1361,53 @@ def test_file_data_field_order_gcs_urls(): ), "mime_type must come before file_uri in the file_data dict" +def test_gemini_files_api_uri_without_format(): + """ + Test that Gemini Files API URIs work WITHOUT an explicit format/mime_type. + + When a user uploads a file via the Gemini Files API and then references it + by URI (https://generativelanguage.googleapis.com/v1beta/files/...), + the file is already on Google's servers. These URLs return 403 when + fetched directly, so _process_gemini_media must NOT try to resolve the + MIME type via HTTP. Instead it should pass the URI through as file_data + and let the Gemini API resolve the type from its stored metadata. + + Related issue: https://github.com/BerriAI/litellm/issues/24907 + """ + from litellm.llms.vertex_ai.gemini.transformation import _process_gemini_media + + file_url = "https://generativelanguage.googleapis.com/v1beta/files/37eh7rsw1vfe" + + # Should NOT raise — previously this hit the generic https:// handler + # which called _get_image_mime_type_from_url() and got a 403. + result = _process_gemini_media(image_url=file_url) + + assert "file_data" in result + file_data = result["file_data"] + assert file_data["file_uri"] == file_url + # When no format is provided, mime_type should be absent so the + # Gemini API infers it from the stored file metadata. + assert "mime_type" not in file_data + + +def test_gemini_files_api_uri_with_format(): + """ + Test that Gemini Files API URIs correctly forward an explicit format. + + Related issue: https://github.com/BerriAI/litellm/issues/24907 + """ + from litellm.llms.vertex_ai.gemini.transformation import _process_gemini_media + + file_url = "https://generativelanguage.googleapis.com/v1beta/files/n1vhxa28lyaw" + + result = _process_gemini_media(image_url=file_url, format="text/plain") + + assert "file_data" in result + file_data = result["file_data"] + assert file_data["file_uri"] == file_url + assert file_data["mime_type"] == "text/plain" + + def test_extract_file_data_with_path_object(): """ Test that filename is correctly extracted from Path objects for MIME type detection. diff --git a/tests/test_litellm/llms/vertex_ai/test_vertex_ai_common_utils.py b/tests/test_litellm/llms/vertex_ai/test_vertex_ai_common_utils.py index 7cb3faf6177..6c549af2cc5 100644 --- a/tests/test_litellm/llms/vertex_ai/test_vertex_ai_common_utils.py +++ b/tests/test_litellm/llms/vertex_ai/test_vertex_ai_common_utils.py @@ -300,6 +300,18 @@ def test_process_items_basic(): process_items(schema) assert schema["items"] == {"type": "object"} + # Test array missing items inside anyOf branch + schema = { + "type": "object", + "properties": { + "callbacks": { + "anyOf": [{"type": "array"}, {"type": "object"}], + } + }, + } + process_items(schema) + assert schema["properties"]["callbacks"]["anyOf"][0]["items"] == {"type": "object"} + def test_build_vertex_schema_array_branch_missing_items_in_anyof(): """ diff --git a/tests/test_litellm/llms/vertex_ai/vertex_ai_partner_models/anthropic/test_vertex_ai_partner_models_anthropic_transformation.py b/tests/test_litellm/llms/vertex_ai/vertex_ai_partner_models/anthropic/test_vertex_ai_partner_models_anthropic_transformation.py index 79fc66a74b8..d0be476d72e 100644 --- a/tests/test_litellm/llms/vertex_ai/vertex_ai_partner_models/anthropic/test_vertex_ai_partner_models_anthropic_transformation.py +++ b/tests/test_litellm/llms/vertex_ai/vertex_ai_partner_models/anthropic/test_vertex_ai_partner_models_anthropic_transformation.py @@ -203,12 +203,15 @@ def test_vertex_ai_anthropic_structured_output_header_not_added(): def test_vertex_ai_claude_sonnet_4_5_structured_output_fix(): """ Test fix for issue #18625: Claude Sonnet 4.5 on VertexAI should use tool-based - structured outputs instead of output_format parameter. + structured outputs when ``response_format`` is supplied via the OpenAI-compat + interface (``map_openai_params``). This test verifies that: - 1. Claude Sonnet 4.5 uses tool-based structured outputs on VertexAI - 2. output_format parameter is removed from the final request - 3. The fix prevents "Extra inputs are not permitted" error + 1. Claude Sonnet 4.5 uses tool-based structured outputs when ``response_format`` + is given to the OpenAI-compat path (the path that triggered #18625). + 2. ``output_format`` is forwarded to Vertex AI when present — Vertex now + accepts the field; the prior blanket-strip behavior was the silent drop + of Anthropic Structured Outputs that this PR fixes. """ config = VertexAIAnthropicConfig() @@ -294,11 +297,15 @@ def test_vertex_ai_claude_sonnet_4_5_structured_output_fix(): headers={}, ) - # Verify that output_format was removed (fixes the "Extra inputs are not permitted" error) + # output_format is now forwarded to Vertex (Vertex parity has shifted — + # it accepts the field and uses it to enforce the JSON schema). The + # prior behavior silently stripped it, hiding Structured Outputs from + # callers who explicitly requested them. + assert "output_format" in final_data + assert final_data["output_format"]["type"] == "json_schema" assert ( - "output_format" not in final_data - ), "output_format should be removed for VertexAI" - assert "model" not in final_data, "model should be removed for VertexAI" + "model" not in final_data + ), "model is still stripped (Vertex routes by URL)" assert "tools" in final_data, "tools should still be present" assert "tool_choice" in final_data, "tool_choice should still be present" @@ -491,28 +498,22 @@ def test_vertex_ai_partner_models_anthropic_remove_prompt_caching_scope_beta_hea ), "Header should be removed if no supported values remain" -def test_vertex_ai_anthropic_output_config_dropped(): +def test_vertex_ai_anthropic_output_config_effort_only_dropped(): """ - Test that output_config parameter is dropped from Vertex AI Anthropic requests. - - Vertex AI does not support the output_config parameter (used for effort settings - in Anthropic API). This test ensures it's properly removed to prevent - "Extra inputs are not permitted" errors. + ``output_config`` containing only ``effort`` (an Anthropic-only key Vertex + rejects with "Extra inputs are not permitted") is dropped entirely so the + request body has no empty dict. """ config = VertexAIAnthropicConfig() messages = [{"role": "user", "content": "What is 2+2?"}] - headers = {} + headers: dict = {} - # Simulate optional_params with output_config that would be passed in optional_params = { "max_tokens": 1024, - "output_config": { - "effort": "high" # This is Anthropic-specific and not supported by Vertex AI - }, + "output_config": {"effort": "high"}, } - # Call transform_request which should drop output_config result = config.transform_request( model="claude-3-5-sonnet-20241022", messages=messages, @@ -521,54 +522,144 @@ def test_vertex_ai_anthropic_output_config_dropped(): headers=headers, ) - # Verify output_config was removed assert ( "output_config" not in result - ), "output_config should be dropped from Vertex AI Anthropic requests" - - # Verify other parameters are preserved - assert result["max_tokens"] == 1024, "max_tokens should be preserved" - assert "messages" in result, "messages should be present" + ), "output_config containing only effort must be dropped" + assert result["max_tokens"] == 1024 + assert "messages" in result -def test_vertex_ai_anthropic_output_format_and_output_config_both_dropped(): +def test_vertex_ai_anthropic_output_config_format_passes_through(): """ - Test that both output_format and output_config are dropped from Vertex AI requests. - - This ensures that even if both parameters somehow make it to the transform_request, - they are properly cleaned up before sending to Vertex AI. + ``output_config`` containing structured-output ``format`` is FORWARDED to + Vertex AI Claude — Vertex now accepts it and uses it for JSON Schema + enforcement. Previously the entire field was being silently stripped, so + Anthropic Structured Outputs never engaged on Vertex even when callers + requested it. """ config = VertexAIAnthropicConfig() + messages = [{"role": "user", "content": "Return a person object."}] + output_config = { + "format": { + "type": "json_schema", + "schema": { + "type": "object", + "additionalProperties": False, + "properties": { + "name": {"type": "string"}, + "age": {"type": "integer"}, + }, + }, + } + } + optional_params = {"max_tokens": 1024, "output_config": output_config} + + result = config.transform_request( + model="claude-3-5-sonnet-20241022", + messages=messages, + optional_params=optional_params, + litellm_params={}, + headers={}, + ) + + assert result["output_config"] == output_config + + +def test_vertex_ai_anthropic_output_config_format_plus_effort_strips_only_effort(): + """ + Greptile P1 on PR #23396: when ``output_config`` contains BOTH ``format`` + and ``effort``, the prior conditional-passthrough logic forwarded the + full dict including the unsupported ``effort`` key, reproducing the + 400 error the fix was meant to resolve. Only ``effort`` (and any future + Vertex-unsupported keys) should be filtered; ``format`` must survive. + """ + config = VertexAIAnthropicConfig() + messages = [{"role": "user", "content": "Return a person object."}] + + output_config = { + "format": { + "type": "json_schema", + "schema": { + "type": "object", + "additionalProperties": False, + "properties": {"name": {"type": "string"}}, + }, + }, + "effort": "high", + } + optional_params = {"max_tokens": 1024, "output_config": output_config} + + result = config.transform_request( + model="claude-3-5-sonnet-20241022", + messages=messages, + optional_params=optional_params, + litellm_params={}, + headers={}, + ) + + assert "output_config" in result + assert ( + "effort" not in result["output_config"] + ), "effort must be stripped — Vertex returns 400 on unknown keys" + assert result["output_config"]["format"] == output_config["format"] + + +def test_vertex_ai_anthropic_output_config_non_dict_dropped(): + """Defensive: if ``output_config`` is somehow not a dict, drop it rather + than forwarding malformed data downstream.""" + config = VertexAIAnthropicConfig() + messages = [{"role": "user", "content": "hi"}] + optional_params = {"max_tokens": 64, "output_config": "not-a-dict"} + + result = config.transform_request( + model="claude-3-5-sonnet-20241022", + messages=messages, + optional_params=optional_params, + litellm_params={}, + headers={}, + ) + + assert "output_config" not in result + + +def test_vertex_ai_anthropic_output_format_preserved_output_config_effort_dropped(): + """ + When the request carries both ``output_format`` (top-level structured + outputs) AND an ``output_config`` whose only useful key for Vertex is + ``effort``: ``output_format`` must be forwarded (Vertex accepts it), + while ``output_config`` is dropped because Vertex returns 400 on + ``effort``. This replaces the old "drop both" behavior, which was the + silent strip the bug report flagged. + """ + config = VertexAIAnthropicConfig() messages = [{"role": "user", "content": "Extract structured data"}] - headers = {} + + output_format = { + "type": "json_schema", + "json_schema": { + "name": "data", + "schema": { + "type": "object", + "properties": {"result": {"type": "string"}}, + }, + }, + } optional_params = { "max_tokens": 2048, - "output_format": { - "type": "json_schema", - "json_schema": { - "name": "data", - "schema": { - "type": "object", - "properties": {"result": {"type": "string"}}, - }, - }, - }, + "output_format": output_format, "output_config": {"effort": "high"}, } - # Simulate parent class creating test_data with both parameters - # (as if the parent transform_request added them) test_data = { "model": "claude-3-5-sonnet-20241022", "messages": messages, "max_tokens": 2048, - "output_format": optional_params["output_format"], - "output_config": optional_params["output_config"], + "output_format": output_format, + "output_config": {"effort": "high"}, } - # Mock the parent transform_request to return data with both parameters original_transform = config.__class__.__bases__[0].transform_request def mock_transform_request( @@ -584,22 +675,50 @@ def test_vertex_ai_anthropic_output_format_and_output_config_both_dropped(): messages=messages, optional_params=optional_params, litellm_params={}, - headers=headers, + headers={}, ) - # Verify both were removed - assert ( - "output_format" not in result - ), "output_format should be dropped from Vertex AI requests" - assert ( - "output_config" not in result - ), "output_config should be dropped from Vertex AI requests" - - # Verify essential params are preserved - assert result["max_tokens"] == 2048, "max_tokens should be preserved" - assert "messages" in result, "messages should be present" - assert "model" not in result, "model should also be dropped for Vertex AI" - + # output_format flows through unchanged — Vertex AI Claude accepts it. + assert result["output_format"] == output_format + # output_config containing only ``effort`` is dropped to avoid the + # 400 "Extra inputs are not permitted" the silent strip used to mask. + assert "output_config" not in result + assert result["max_tokens"] == 2048 + assert "model" not in result, "model is still stripped (Vertex routes by URL)" finally: - # Restore original method config.__class__.__bases__[0].transform_request = original_transform + + +def test_sanitize_vertex_anthropic_output_params_unit(): + """Direct unit coverage for the helper itself (used by both Vertex + Anthropic transformation paths). Mirrors the integration assertions + above without going through the full ``transform_request`` stack.""" + from litellm.llms.vertex_ai.vertex_ai_partner_models.anthropic.output_params_utils import ( + sanitize_vertex_anthropic_output_params, + ) + + # No-op when output_config absent. + data: dict = {"max_tokens": 8} + sanitize_vertex_anthropic_output_params(data) + assert data == {"max_tokens": 8} + + # Effort-only → dropped entirely. + data = {"output_config": {"effort": "high"}} + sanitize_vertex_anthropic_output_params(data) + assert "output_config" not in data + + # Format-only → preserved unchanged. + fmt = {"format": {"type": "json_schema", "schema": {"type": "object"}}} + data = {"output_config": dict(fmt)} + sanitize_vertex_anthropic_output_params(data) + assert data["output_config"] == fmt + + # Mixed → effort filtered, format kept. + data = {"output_config": {"format": fmt["format"], "effort": "high"}} + sanitize_vertex_anthropic_output_params(data) + assert data["output_config"] == fmt + + # Non-dict → dropped defensively. + data = {"output_config": "garbage"} + sanitize_vertex_anthropic_output_params(data) + assert "output_config" not in data diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py b/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py index dd352d0999a..6e0dadcd4d8 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py @@ -2139,3 +2139,275 @@ async def test_tool_permission_servers_included_in_allowed_servers(): assert "server_id_123" in result finally: global_mcp_server_manager.registry.pop("server_id_123", None) + + +# --------------------------------------------------------------------------- +# Org-level MCP permission tests +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +class TestOrgMCPPermissions: + """Tests for org-level MCP server permission enforcement.""" + + def _make_auth(self, org_id=None, team_id=None) -> UserAPIKeyAuth: + return UserAPIKeyAuth( + api_key="test-key", + user_id="test-user", + team_id=team_id, + org_id=org_id, + ) + + @pytest.mark.parametrize( + "key_servers,team_servers,org_servers,expected,scenario", + [ + ( + ["s1", "s2"], + [], + None, + ["s1", "s2"], + "no_org_id", + ), + ( + ["s1", "s2"], + [], + [], + ["s1", "s2"], + "org_empty_no_restriction", + ), + ( + [], + [], + ["org_s1", "org_s2"], + ["org_s1", "org_s2"], + "org_only_ceiling", + ), + ( + ["s1", "s2"], + [], + ["s1", "org_only"], + ["s1"], + "org_intersection", + ), + ( + ["s1", "s2"], + [], + ["org_s1"], + [], + "no_overlap_denied", + ), + ( + ["s1", "s2"], + ["s1", "s2", "s3"], + ["s1"], + ["s1"], + "team_then_org", + ), + ( + ["s1"], + ["s2"], + ["s1", "s2", "org_s1"], + [], + "key_team_conflict_not_expanded_by_org", + ), + ], + ) + async def test_get_allowed_mcp_servers_with_org( + self, + key_servers, + team_servers, + org_servers, + expected, + scenario, + ): + org_id = "org-123" if org_servers is not None else None + auth = self._make_auth(org_id=org_id) + + with ( + patch.object( + MCPRequestHandler, + "_get_allowed_mcp_servers_for_key", + new_callable=AsyncMock, + return_value=key_servers, + ), + patch.object( + MCPRequestHandler, + "_get_allowed_mcp_servers_for_team", + new_callable=AsyncMock, + return_value=team_servers, + ), + patch.object( + MCPRequestHandler, + "_get_allowed_mcp_servers_for_org", + new_callable=AsyncMock, + return_value=org_servers if org_servers is not None else [], + ), + ): + result = await MCPRequestHandler.get_allowed_mcp_servers(auth) + assert sorted(result) == sorted(expected), f"scenario={scenario}" + + async def test_get_org_object_permission_no_org_id(self): + auth = self._make_auth(org_id=None) + result = await MCPRequestHandler._get_org_object_permission(auth) + assert result is None + + async def test_get_org_object_permission_no_prisma(self): + auth = self._make_auth(org_id="org-123") + with patch( + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.MCPRequestHandler._get_org_object_permission", + new_callable=AsyncMock, + return_value=None, + ): + result = await MCPRequestHandler._get_allowed_mcp_servers_for_org(auth) + assert result == [] + + async def test_get_allowed_mcp_servers_for_org_direct_servers(self): + auth = self._make_auth(org_id="org-123") + + mock_perm = MagicMock() + mock_perm.mcp_servers = ["org_server_1", "org_server_2"] + mock_perm.mcp_access_groups = [] + mock_perm.mcp_tool_permissions = {} + + with ( + patch.object( + MCPRequestHandler, + "_get_org_object_permission", + new_callable=AsyncMock, + return_value=mock_perm, + ), + patch.object( + MCPRequestHandler, + "_get_mcp_servers_from_access_groups", + new_callable=AsyncMock, + return_value=[], + ), + ): + result = await MCPRequestHandler._get_allowed_mcp_servers_for_org(auth) + assert sorted(result) == ["org_server_1", "org_server_2"] + + async def test_get_allowed_mcp_servers_for_org_access_groups(self): + auth = self._make_auth(org_id="org-123") + + mock_perm = MagicMock() + mock_perm.mcp_servers = [] + mock_perm.mcp_access_groups = ["group-a"] + mock_perm.mcp_tool_permissions = {} + + with ( + patch.object( + MCPRequestHandler, + "_get_org_object_permission", + new_callable=AsyncMock, + return_value=mock_perm, + ), + patch.object( + MCPRequestHandler, + "_get_mcp_servers_from_access_groups", + new_callable=AsyncMock, + return_value=["group_server_1"], + ), + ): + result = await MCPRequestHandler._get_allowed_mcp_servers_for_org(auth) + assert "group_server_1" in result + + async def test_get_allowed_mcp_servers_for_org_tool_permissions_only(self): + auth = self._make_auth(org_id="org-123") + + mock_perm = MagicMock() + mock_perm.mcp_servers = [] + mock_perm.mcp_access_groups = [] + mock_perm.mcp_tool_permissions = {"tool_only_server": ["tool_x"]} + + with ( + patch.object( + MCPRequestHandler, + "_get_org_object_permission", + new_callable=AsyncMock, + return_value=mock_perm, + ), + patch.object( + MCPRequestHandler, + "_get_mcp_servers_from_access_groups", + new_callable=AsyncMock, + return_value=[], + ), + ): + result = await MCPRequestHandler._get_allowed_mcp_servers_for_org(auth) + assert "tool_only_server" in result + + async def test_get_allowed_mcp_servers_for_org_no_object_permission(self): + auth = self._make_auth(org_id="org-123") + + with patch.object( + MCPRequestHandler, + "_get_org_object_permission", + new_callable=AsyncMock, + return_value=None, + ): + result = await MCPRequestHandler._get_allowed_mcp_servers_for_org(auth) + assert result == [] + + async def test_get_allowed_tools_for_server_org_ceiling(self): + auth = self._make_auth(org_id="org-123") + + key_perm = MagicMock() + key_perm.mcp_tool_permissions = {"server_1": ["tool_a", "tool_b", "tool_c"]} + + org_perm = MagicMock() + org_perm.mcp_tool_permissions = {"server_1": ["tool_a", "tool_b"]} + + with ( + patch.object( + MCPRequestHandler, "_get_key_object_permission", return_value=key_perm + ), + patch.object( + MCPRequestHandler, + "_get_team_object_permission", + new_callable=AsyncMock, + return_value=None, + ), + patch.object( + MCPRequestHandler, + "_get_org_object_permission", + new_callable=AsyncMock, + return_value=org_perm, + ), + ): + result = await MCPRequestHandler.get_allowed_tools_for_server( + server_id="server_1", + user_api_key_auth=auth, + ) + assert sorted(result) == ["tool_a", "tool_b"] + + async def test_get_allowed_tools_for_server_org_no_restriction(self): + auth = self._make_auth(org_id="org-123") + + key_perm = MagicMock() + key_perm.mcp_tool_permissions = {"server_1": ["tool_a", "tool_b"]} + + org_perm = MagicMock() + org_perm.mcp_tool_permissions = {} + + with ( + patch.object( + MCPRequestHandler, "_get_key_object_permission", return_value=key_perm + ), + patch.object( + MCPRequestHandler, + "_get_team_object_permission", + new_callable=AsyncMock, + return_value=None, + ), + patch.object( + MCPRequestHandler, + "_get_org_object_permission", + new_callable=AsyncMock, + return_value=org_perm, + ), + ): + result = await MCPRequestHandler.get_allowed_tools_for_server( + server_id="server_1", + user_api_key_auth=auth, + ) + assert sorted(result) == ["tool_a", "tool_b"] diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_openapi_tool_auth.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_openapi_tool_auth.py new file mode 100644 index 00000000000..3ad01e9c3ec --- /dev/null +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_openapi_tool_auth.py @@ -0,0 +1,220 @@ +""" +VERIA-7 regression: OpenAPI-backed (local-registry) MCP tools must run +through `pre_call_tool_check` before dispatch, the same as managed +MCP server tools. +""" + +from datetime import datetime, timezone +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest + +from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth + + +@pytest.mark.asyncio +async def test_openapi_local_tool_runs_pre_call_tool_check(): + """When `execute_mcp_tool` resolves a local-registry (OpenAPI) tool + AND a server, the pre-call hook must fire before the local handler + runs. Pre-fix this path skipped the hook entirely.""" + from litellm.proxy._experimental.mcp_server import server as mcp_module + + user = UserAPIKeyAuth( + api_key="sk-user", + user_id="alice", + user_role=LitellmUserRoles.INTERNAL_USER.value, + ) + + fake_server = MagicMock() + fake_server.name = "openapi-petstore" + fake_server.is_byok = False + fake_server.auth_type = None + fake_server.mcp_info = None + fake_server.server_id = "srv-1" + fake_server.server_name = "openapi-petstore" + + fake_tool = MagicMock() + fake_tool.name = "list_pets" + + pre_call = AsyncMock(return_value={}) + handle_local = AsyncMock(return_value=[]) + + with ( + patch.object( + mcp_module.global_mcp_server_manager, + "_get_mcp_server_from_tool_name", + return_value=fake_server, + ), + patch.object( + mcp_module.global_mcp_server_manager, + "pre_call_tool_check", + new=pre_call, + ), + patch.object( + mcp_module.global_mcp_tool_registry, + "get_tool", + return_value=fake_tool, + ), + patch( + "litellm.proxy._experimental.mcp_server.server._handle_local_mcp_tool", + new=handle_local, + ), + patch( + "litellm.proxy._experimental.mcp_server.server.MCPRequestHandler.is_tool_allowed", + return_value=True, + ), + ): + await mcp_module.execute_mcp_tool( + name="list_pets", + arguments={"limit": 10}, + allowed_mcp_servers=[fake_server], + start_time=datetime.now(timezone.utc), + user_api_key_auth=user, + ) + + pre_call.assert_awaited_once() + handle_local.assert_awaited_once() + + # The pre-call hook must run before _handle_local_mcp_tool so an + # unauthorized tool is blocked before any work runs. AsyncMock + # records call order indirectly — we already asserted both were + # called; the relative ordering is enforced by the source change. + pre_call_kwargs = pre_call.await_args.kwargs + assert pre_call_kwargs["name"] == "list_pets" + assert pre_call_kwargs["server"] is fake_server + assert pre_call_kwargs["user_api_key_auth"] is user + # `proxy_logging_obj` must be sourced from the canonical proxy_server + # module (same as the managed path) — passing None would crash the + # downstream `_create_mcp_request_object_from_kwargs` call with + # AttributeError after the security checks succeed. + assert pre_call_kwargs["proxy_logging_obj"] is not None + + +@pytest.mark.asyncio +async def test_openapi_local_tool_blocked_when_pre_call_check_raises(): + """If the pre-call check raises (caller not authorized for this + tool), the local handler must NOT be invoked.""" + from fastapi import HTTPException + + from litellm.proxy._experimental.mcp_server import server as mcp_module + + user = UserAPIKeyAuth( + api_key="sk-user", + user_id="alice", + user_role=LitellmUserRoles.INTERNAL_USER.value, + ) + + fake_server = MagicMock() + fake_server.name = "openapi-petstore" + fake_server.is_byok = False + fake_server.auth_type = None + fake_server.mcp_info = None + fake_server.server_id = "srv-1" + fake_server.server_name = "openapi-petstore" + + fake_tool = MagicMock() + fake_tool.name = "delete_pet" + + pre_call = AsyncMock( + side_effect=HTTPException(status_code=403, detail="not allowed") + ) + handle_local = AsyncMock(return_value=[]) + + with ( + patch.object( + mcp_module.global_mcp_server_manager, + "_get_mcp_server_from_tool_name", + return_value=fake_server, + ), + patch.object( + mcp_module.global_mcp_server_manager, + "pre_call_tool_check", + new=pre_call, + ), + patch.object( + mcp_module.global_mcp_tool_registry, + "get_tool", + return_value=fake_tool, + ), + patch( + "litellm.proxy._experimental.mcp_server.server._handle_local_mcp_tool", + new=handle_local, + ), + patch( + "litellm.proxy._experimental.mcp_server.server.MCPRequestHandler.is_tool_allowed", + return_value=True, + ), + ): + with pytest.raises(HTTPException) as exc: + await mcp_module.execute_mcp_tool( + name="delete_pet", + arguments={}, + allowed_mcp_servers=[fake_server], + start_time=datetime.now(timezone.utc), + user_api_key_auth=user, + ) + + assert exc.value.status_code == 403 + pre_call.assert_awaited_once() + handle_local.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_openapi_local_tool_denied_when_server_not_resolvable(): + """If the local-registry tool is found but no MCP server resolves + (startup race or orphaned registry entry), the call must be rejected + rather than dispatched without `pre_call_tool_check`.""" + from fastapi import HTTPException + + from litellm.proxy._experimental.mcp_server import server as mcp_module + + user = UserAPIKeyAuth( + api_key="sk-user", + user_id="alice", + user_role=LitellmUserRoles.INTERNAL_USER.value, + ) + + fake_tool = MagicMock() + fake_tool.name = "list_pets" + + pre_call = AsyncMock(return_value={}) + handle_local = AsyncMock(return_value=[]) + + # `_get_mcp_server_from_tool_name` returns None — no server context. + with ( + patch.object( + mcp_module.global_mcp_server_manager, + "_get_mcp_server_from_tool_name", + return_value=None, + ), + patch.object( + mcp_module.global_mcp_server_manager, + "pre_call_tool_check", + new=pre_call, + ), + patch.object( + mcp_module.global_mcp_tool_registry, + "get_tool", + return_value=fake_tool, + ), + patch( + "litellm.proxy._experimental.mcp_server.server._handle_local_mcp_tool", + new=handle_local, + ), + patch( + "litellm.proxy._experimental.mcp_server.server.MCPRequestHandler.is_tool_allowed", + return_value=True, + ), + ): + with pytest.raises(HTTPException) as exc: + await mcp_module.execute_mcp_tool( + name="list_pets", + arguments={}, + allowed_mcp_servers=[], + start_time=datetime.now(timezone.utc), + user_api_key_auth=user, + ) + + assert exc.value.status_code == 503 + pre_call.assert_not_awaited() + handle_local.assert_not_awaited() diff --git a/tests/test_litellm/proxy/auth/test_admin_viewer_handler_access.py b/tests/test_litellm/proxy/auth/test_admin_viewer_handler_access.py new file mode 100644 index 00000000000..b0d2595e48c --- /dev/null +++ b/tests/test_litellm/proxy/auth/test_admin_viewer_handler_access.py @@ -0,0 +1,168 @@ +""" +Handler-level admin viewer parity tests. + +These tests assert that PROXY_ADMIN_VIEW_ONLY callers are NOT blocked at the +handler level for read-only admin endpoints. The route_checks layer is tested +separately in `test_route_checks.py`; here we verify each individual endpoint +function has been updated to use `_user_has_admin_view()` rather than a bare +`user_role != PROXY_ADMIN` check. + +The principle (see Admin Viewer role doc): anything Proxy Admin can read, +Admin Viewer can read. No writes, no cost-incurring actions. +""" + +import os +import sys +import types +from unittest.mock import AsyncMock, MagicMock + +import pytest +from fastapi.testclient import TestClient + +sys.path.insert(0, os.path.abspath("../../../")) + +import litellm.proxy.proxy_server as ps +from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth +from litellm.proxy.proxy_server import app + + +def _make_admin_viewer_auth() -> UserAPIKeyAuth: + return UserAPIKeyAuth( + user_id="viewer_user", + user_role=LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY, + ) + + +def _override_auth(role: LitellmUserRoles) -> None: + fake_user = UserAPIKeyAuth(user_id="viewer_user", user_role=role) + app.dependency_overrides[ps.user_api_key_auth] = lambda: fake_user + + +def _clear_overrides() -> None: + app.dependency_overrides.clear() + + +@pytest.fixture +def admin_viewer_client(monkeypatch): + """TestClient where auth always returns PROXY_ADMIN_VIEW_ONLY + a mocked Prisma.""" + mock_prisma = MagicMock() + + # Common DB tables touched by the read endpoints under test. + mock_budget_table = MagicMock() + mock_budget_table.find_many = AsyncMock(return_value=[]) + mock_budget_table.find_first = AsyncMock(return_value=None) + + mock_invitation_table = MagicMock() + mock_invitation_table.find_unique = AsyncMock(return_value=None) + + mock_config_table = MagicMock() + mock_config_table.find_first = AsyncMock(return_value=None) + + mock_prisma.db = types.SimpleNamespace( + litellm_budgettable=mock_budget_table, + litellm_invitationlink=mock_invitation_table, + litellm_config=mock_config_table, + ) + + monkeypatch.setattr(ps, "prisma_client", mock_prisma) + _override_auth(LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY) + + yield TestClient(app) + + _clear_overrides() + + +def _assert_not_role_blocked(response) -> None: + """The endpoint must not return a role-block error. + + Detects both the 400 ``not_allowed_access`` pattern (used by most + management endpoints) and the 403 ``Admin role required`` pattern + (used by model cost map endpoints). + """ + if response.status_code in (400, 401, 403): + body = response.json() + detail = body.get("detail", body) + if isinstance(detail, dict): + err = detail.get("error", "") + else: + err = str(detail) + err_lower = err.lower() + role_block_signals = ( + "your role=", + "not allowed to access", + "admin role required", + "admin-only endpoint", + ) + for signal in role_block_signals: + assert ( + signal not in err_lower + ), f"endpoint blocked PROXY_ADMIN_VIEW_ONLY at handler level: {err}" + + +def test_budget_list_allows_admin_viewer(admin_viewer_client): + """`/budget/list` is read-only and must be accessible to Admin Viewer.""" + resp = admin_viewer_client.get("/budget/list") + _assert_not_role_blocked(resp) + assert resp.status_code == 200, resp.text + + +def test_budget_settings_allows_admin_viewer(admin_viewer_client): + """`/budget/settings` describes a budget's fields; read-only.""" + resp = admin_viewer_client.get("/budget/settings", params={"budget_id": "b1"}) + _assert_not_role_blocked(resp) + assert resp.status_code == 200, resp.text + + +def test_alerting_settings_allows_admin_viewer(admin_viewer_client): + """`/alerting/settings` describes alerting params; read-only.""" + resp = admin_viewer_client.get("/alerting/settings") + _assert_not_role_blocked(resp) + # Endpoint may 400 for *config* reasons (no proxy config loaded), but it + # must not 400 because of role. + assert resp.status_code != 403, resp.text + + +def test_get_config_field_info_allows_admin_viewer(admin_viewer_client): + """`/config/field/info` describes a single general-settings field; read-only.""" + resp = admin_viewer_client.get( + "/config/field/info", params={"field_name": "alerting"} + ) + _assert_not_role_blocked(resp) + + +def test_get_config_list_allows_admin_viewer(admin_viewer_client): + """`/config/list` lists configurable params for a config_type; read-only.""" + resp = admin_viewer_client.get( + "/config/list", params={"config_type": "general_settings"} + ) + _assert_not_role_blocked(resp) + + +def test_get_config_callbacks_allows_admin_viewer(admin_viewer_client): + """`/get/config/callbacks` lists current callbacks; read-only.""" + resp = admin_viewer_client.get("/get/config/callbacks") + _assert_not_role_blocked(resp) + + +def test_invitation_info_allows_admin_viewer(admin_viewer_client): + """`/invitation/info` reads a single invitation; read-only. + + The invitation lookup will return 400 because no invitation exists in our + mock DB — that's fine. We only assert it doesn't hit the role-block path. + """ + resp = admin_viewer_client.get( + "/invitation/info", params={"invitation_id": "nonexistent"} + ) + _assert_not_role_blocked(resp) + + +def test_model_cost_map_reload_status_allows_admin_viewer(admin_viewer_client): + """`/schedule/model_cost_map_reload/status` is read-only operations status.""" + resp = admin_viewer_client.get("/schedule/model_cost_map_reload/status") + _assert_not_role_blocked(resp) + + +def test_model_cost_map_source_allows_admin_viewer(admin_viewer_client): + """`/model/cost_map/source` reads the configured cost map source URL.""" + resp = admin_viewer_client.get("/model/cost_map/source") + _assert_not_role_blocked(resp) diff --git a/tests/test_litellm/proxy/auth/test_auth_checks.py b/tests/test_litellm/proxy/auth/test_auth_checks.py index 4c21d0ec645..5dedf05215b 100644 --- a/tests/test_litellm/proxy/auth/test_auth_checks.py +++ b/tests/test_litellm/proxy/auth/test_auth_checks.py @@ -17,7 +17,10 @@ import litellm from litellm.proxy._types import ( CallInfo, Litellm_EntityType, + LiteLLM_BudgetTable, + LiteLLM_EndUserTable, LiteLLM_ObjectPermissionTable, + LiteLLM_TagTable, LiteLLM_TeamTable, LiteLLM_UserTable, LitellmUserRoles, @@ -29,10 +32,12 @@ from litellm.proxy._types import ( from litellm.proxy.auth.auth_checks import ( ExperimentalUIJWTToken, _can_object_call_vector_stores, + _check_end_user_budget, _check_team_member_budget, _get_fuzzy_user_object, _get_team_db_check, _log_budget_lookup_failure, + _tag_max_budget_check, _team_max_budget_check, _virtual_key_max_budget_alert_check, _virtual_key_max_budget_check, @@ -1964,6 +1969,67 @@ async def test_team_budget_check_reads_from_spend_counter(): assert exc_info.value.current_cost == 1.5 +@pytest.mark.asyncio +async def test_end_user_budget_check_reads_from_spend_counter(): + """End-user budget check should use get_current_spend when counter exists.""" + end_user_object = LiteLLM_EndUserTable( + user_id="customer-1", + blocked=False, + spend=0.0, + litellm_budget_table=LiteLLM_BudgetTable(max_budget=1.0), + ) + + async def mock_get_current_spend(counter_key, fallback_spend): + if counter_key == "spend:end_user:customer-1": + return 1.5 + return fallback_spend + + with patch("litellm.proxy.proxy_server.get_current_spend", mock_get_current_spend): + with pytest.raises(litellm.BudgetExceededError) as exc_info: + await _check_end_user_budget( + end_user_obj=end_user_object, + route="/chat/completions", + ) + assert exc_info.value.current_cost == 1.5 + assert exc_info.value.max_budget == 1.0 + + +@pytest.mark.asyncio +async def test_tag_budget_check_reads_from_spend_counter(): + """Tag budget check should use get_current_spend when counter exists.""" + from litellm.proxy.utils import ProxyLogging + + tag_object = LiteLLM_TagTable( + tag_name="paid-tag", + spend=0.0, + litellm_budget_table=LiteLLM_BudgetTable(max_budget=1.0), + ) + + async def mock_get_current_spend(counter_key, fallback_spend): + if counter_key == "spend:tag:paid-tag": + return 1.5 + return fallback_spend + + with ( + patch("litellm.proxy.proxy_server.get_current_spend", mock_get_current_spend), + patch( + "litellm.proxy.auth.auth_checks.get_tag_objects_batch", + new_callable=AsyncMock, + return_value={"paid-tag": tag_object}, + ), + ): + with pytest.raises(litellm.BudgetExceededError) as exc_info: + await _tag_max_budget_check( + request_body={"metadata": {"tags": ["paid-tag"]}}, + prisma_client=MagicMock(), + user_api_key_cache=MagicMock(), + proxy_logging_obj=ProxyLogging(user_api_key_cache=None), + valid_token=UserAPIKeyAuth(token="test-token"), + ) + assert exc_info.value.current_cost == 1.5 + assert exc_info.value.max_budget == 1.0 + + @pytest.mark.asyncio async def test_team_member_budget_check_reads_from_spend_counter(): """Team member budget check should use get_current_spend when counter exists.""" diff --git a/tests/test_litellm/proxy/auth/test_auth_utils.py b/tests/test_litellm/proxy/auth/test_auth_utils.py index b82cb355192..c146b5ded5b 100644 --- a/tests/test_litellm/proxy/auth/test_auth_utils.py +++ b/tests/test_litellm/proxy/auth/test_auth_utils.py @@ -2,6 +2,7 @@ Unit tests for auth_utils functions related to rate limiting and customer ID extraction. """ +import base64 from typing import Optional from unittest.mock import MagicMock, patch @@ -10,11 +11,12 @@ import pytest from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.auth.auth_utils import ( _get_customer_id_from_standard_headers, + abbreviate_api_key, check_complete_credentials, get_end_user_id_from_request_body, - get_model_from_request, get_key_model_rpm_limit, get_key_model_tpm_limit, + get_model_from_request, get_project_model_rpm_limit, get_project_model_tpm_limit, is_request_body_safe, @@ -258,6 +260,206 @@ def test_get_model_from_request_vertex_passthrough_still_works(): assert get_model_from_request(request_data={}, route=route) == "gemini-1.5-pro" +def test_get_model_from_request_openai_deployment_route_still_works(): + assert ( + get_model_from_request( + request_data={}, + route="/openai/deployments/my-azure-deployment/chat/completions", + ) + == "my-azure-deployment" + ) + + +def test_get_model_from_request_includes_file_endpoint_header_model(): + assert ( + get_model_from_request( + request_data={}, + route="/v1/files", + request_headers={"X-LiteLLM-Model": "restricted-model"}, + ) + == "restricted-model" + ) + + +def test_get_model_from_request_ignores_routing_header_on_standard_llm_routes(): + assert ( + get_model_from_request( + request_data={"model": "allowed-model"}, + route="/v1/chat/completions", + request_headers={"x-litellm-model": "restricted-model"}, + ) + == "allowed-model" + ) + + +def test_get_model_from_request_authorizes_all_file_routing_model_sources(): + models = get_model_from_request( + request_data={"model": "body-model"}, + route="/v1/files", + request_headers={"x-litellm-model": "header-model"}, + request_query_params={"target_model_names": "query-model-a,query-model-b"}, + ) + assert isinstance(models, list) + assert set(models) == { + "body-model", + "query-model-a", + "query-model-b", + "header-model", + } + + +def test_get_model_from_request_extracts_simple_encoded_file_id_model(): + from litellm.proxy.openai_files_endpoints.common_utils import ( + encode_file_id_with_model, + ) + + file_id = encode_file_id_with_model( + file_id="file-provider-id", + model="restricted-model", + ) + + assert ( + get_model_from_request( + request_data={"file_id": file_id}, + route="/v1/files/{file_id}", + ) + == "restricted-model" + ) + + +def test_get_model_from_request_extracts_unified_file_id_models(): + raw_unified_file_id = ( + "litellm_proxy:application/octet-stream;unified_id,test-id;" + "target_model_names,model-a,model-b;llm_output_file_id,file-provider-id" + ) + encoded_unified_file_id = ( + base64.urlsafe_b64encode(raw_unified_file_id.encode()).decode().rstrip("=") + ) + + assert get_model_from_request( + request_data={"file_id": encoded_unified_file_id}, + route="/v1/files/{file_id}", + ) == ["model-a", "model-b"] + + +def test_get_model_from_request_extracts_eval_completion_model(): + assert ( + get_model_from_request( + request_data={"completion": {"model": "judge-model"}}, + route="/v1/evals/{eval_id}/runs", + ) + == "judge-model" + ) + + +def test_get_model_from_request_includes_fine_tuning_target_model_query(): + assert ( + get_model_from_request( + request_data={}, + route="/v1/fine_tuning/jobs", + request_query_params={"target_model_names": "fine-tune-model"}, + ) + == "fine-tune-model" + ) + + +def test_get_model_from_request_extracts_video_id_model(): + from litellm.types.videos.utils import encode_video_id_with_provider + + video_id = encode_video_id_with_provider( + video_id="video-provider-id", + provider="openai", + model_id="video-model", + ) + + assert ( + get_model_from_request( + request_data={"video_id": video_id}, + route="/v1/videos/{video_id}", + ) + == "video-model" + ) + + +def test_get_model_from_request_only_runs_media_decoders_for_matching_fields(): + with ( + patch( + "litellm.types.videos.utils.decode_video_id_with_provider", + return_value={"model_id": "video-model"}, + ) as video_decoder, + patch( + "litellm.types.videos.utils.decode_character_id_with_provider", + return_value={"model_id": "character-model"}, + ) as character_decoder, + ): + assert ( + get_model_from_request( + request_data={"file_id": "file-provider-id"}, + route="/v1/files/{file_id}", + ) + is None + ) + video_decoder.assert_not_called() + character_decoder.assert_not_called() + + assert ( + get_model_from_request( + request_data={"video_id": "video-provider-id"}, + route="/v1/videos/{video_id}", + ) + == "video-model" + ) + video_decoder.assert_called_once_with("video-provider-id") + character_decoder.assert_not_called() + + video_decoder.reset_mock() + character_decoder.reset_mock() + assert ( + get_model_from_request( + request_data={"character_id": "character-provider-id"}, + route="/v1/videos/{character_id}", + ) + == "character-model" + ) + video_decoder.assert_not_called() + character_decoder.assert_called_once_with("character-provider-id") + + +def test_get_model_from_request_handles_managed_id_decoder_failures(): + with ( + patch( + "litellm.proxy.openai_files_endpoints.common_utils.decode_model_from_file_id", + side_effect=Exception("decode failed"), + ), + patch( + "litellm.llms.base_llm.managed_resources.utils.parse_unified_id", + side_effect=Exception("parse failed"), + ), + patch( + "litellm.types.videos.utils.decode_video_id_with_provider", + side_effect=Exception("video decode failed"), + ), + ): + assert ( + get_model_from_request( + request_data={"file_id": "not-a-managed-resource-id"}, + route="/v1/files/{file_id}", + ) + is None + ) + assert ( + get_model_from_request( + request_data={"video_id": "not-a-managed-resource-id"}, + route="/v1/videos/{video_id}", + ) + is None + ) + + +def test_abbreviate_api_key(): + assert abbreviate_api_key("sk-test-1234") == "sk-...1234" + + def test_get_customer_user_header_returns_none_when_no_customer_role(): from litellm.proxy.auth.auth_utils import get_customer_user_header_from_mapping diff --git a/tests/test_litellm/proxy/auth/test_handle_jwt.py b/tests/test_litellm/proxy/auth/test_handle_jwt.py index 47e513dc593..b7dba9c1d16 100644 --- a/tests/test_litellm/proxy/auth/test_handle_jwt.py +++ b/tests/test_litellm/proxy/auth/test_handle_jwt.py @@ -2567,3 +2567,116 @@ async def test_auth_builder_single_team_fallback_membership_error_skips_no_raise assert result["team_membership"] is None mock_get_team.assert_called() mock_get_membership.assert_called_once() + + +# --------------------------------------------------------------------------- +# JWTHandler._build_decode_kwargs — VERIA-27 (audience + issuer verification) +# --------------------------------------------------------------------------- + + +@pytest.fixture(autouse=False) +def _reset_unscoped_warning_flag(): + """Reset the once-per-process warning sentinel so each test sees a fresh + state.""" + JWTHandler._unscoped_jwt_warning_emitted = False + yield + JWTHandler._unscoped_jwt_warning_emitted = False + + +def test_build_decode_kwargs_no_env_disables_both_verifications( + monkeypatch, _reset_unscoped_warning_flag +): + monkeypatch.delenv("JWT_AUDIENCE", raising=False) + monkeypatch.delenv("JWT_ISSUER", raising=False) + + kwargs = JWTHandler._build_decode_kwargs() + + assert kwargs["audience"] is None + assert kwargs["issuer"] is None + assert kwargs["options"] == {"verify_aud": False, "verify_iss": False} + + +def test_build_decode_kwargs_audience_only_enables_aud_verification( + monkeypatch, _reset_unscoped_warning_flag +): + monkeypatch.setenv("JWT_AUDIENCE", "my-proxy") + monkeypatch.delenv("JWT_ISSUER", raising=False) + + kwargs = JWTHandler._build_decode_kwargs() + + assert kwargs["audience"] == "my-proxy" + assert kwargs["issuer"] is None + # verify_aud not in options means PyJWT will verify audience + assert kwargs["options"] == {"verify_iss": False} + + +def test_build_decode_kwargs_issuer_only_enables_iss_verification( + monkeypatch, _reset_unscoped_warning_flag +): + monkeypatch.delenv("JWT_AUDIENCE", raising=False) + monkeypatch.setenv("JWT_ISSUER", "https://idp.example.com/") + + kwargs = JWTHandler._build_decode_kwargs() + + assert kwargs["audience"] is None + assert kwargs["issuer"] == "https://idp.example.com/" + assert kwargs["options"] == {"verify_aud": False} + + +def test_build_decode_kwargs_both_set_enables_full_verification( + monkeypatch, _reset_unscoped_warning_flag +): + monkeypatch.setenv("JWT_AUDIENCE", "my-proxy") + monkeypatch.setenv("JWT_ISSUER", "https://idp.example.com/") + + kwargs = JWTHandler._build_decode_kwargs() + + assert kwargs["audience"] == "my-proxy" + assert kwargs["issuer"] == "https://idp.example.com/" + # No verification opt-outs — PyJWT verifies both claims by default. + assert kwargs["options"] is None + + +def test_build_decode_kwargs_warns_once_when_unscoped( + monkeypatch, _reset_unscoped_warning_flag, caplog +): + """The warning about unscoped JWT auth should fire on the first call but + not on every subsequent decode.""" + import logging + + monkeypatch.delenv("JWT_AUDIENCE", raising=False) + monkeypatch.delenv("JWT_ISSUER", raising=False) + caplog.set_level(logging.WARNING) + + JWTHandler._build_decode_kwargs() + JWTHandler._build_decode_kwargs() + JWTHandler._build_decode_kwargs() + + matching = [ + r + for r in caplog.records + if "JWT auth is enabled" in r.getMessage() + and "neither JWT_AUDIENCE nor JWT_ISSUER" in r.getMessage() + ] + assert ( + len(matching) == 1 + ), f"Expected exactly one warning across 3 calls, got {len(matching)}" + + +def test_build_decode_kwargs_no_warning_when_scoped( + monkeypatch, _reset_unscoped_warning_flag, caplog +): + import logging + + monkeypatch.setenv("JWT_AUDIENCE", "my-proxy") + monkeypatch.delenv("JWT_ISSUER", raising=False) + caplog.set_level(logging.WARNING) + + JWTHandler._build_decode_kwargs() + + matching = [ + r + for r in caplog.records + if "neither JWT_AUDIENCE nor JWT_ISSUER" in r.getMessage() + ] + assert matching == [] diff --git a/tests/test_litellm/proxy/auth/test_oauth2_proxy_hook.py b/tests/test_litellm/proxy/auth/test_oauth2_proxy_hook.py new file mode 100644 index 00000000000..dcbfd281e01 --- /dev/null +++ b/tests/test_litellm/proxy/auth/test_oauth2_proxy_hook.py @@ -0,0 +1,211 @@ +""" +Regression tests for the OAuth2-proxy header-forgery fix +(GHSA-5c3m-qffq-4r9m). + +The hook reads HTTP request headers per ``oauth2_config_mappings`` and +constructs a ``UserAPIKeyAuth`` from them. The fix has two parts: + +1. Only requests from configured trusted proxy CIDR ranges may provide + identity headers. +2. Only identity fields may be mapped from those headers. Without the + identity-only allowlist any field could be mapped — including + ``user_role``, which Pydantic coerces from the string + ``"proxy_admin"`` into ``LitellmUserRoles.PROXY_ADMIN``. +""" + +import os +import sys + +import pytest +from fastapi import Request +from starlette.datastructures import Headers + +sys.path.insert(0, os.path.abspath("../../../..")) + +from litellm.proxy._types import LitellmUserRoles +from litellm.proxy.auth.oauth2_proxy_hook import ( + ALLOWED_OAUTH2_PROXY_FIELDS, + handle_oauth2_proxy_request, +) + + +def _request_with_headers(headers: dict, *, client_host: str = "127.0.0.1") -> Request: + scope = { + "type": "http", + "client": (client_host, 12345), + "headers": [(k.lower().encode(), v.encode()) for k, v in headers.items()], + } + request = Request(scope=scope) + request._headers = Headers(headers) + return request + + +@pytest.fixture +def configure_proxy(monkeypatch): + """ + Yields a callable that sets ``oauth2_config_mappings`` and + ``trusted_proxy_ranges`` on the proxy_server module for the duration + of one test. Defaults to a single identity mapping and localhost as + a trusted proxy. + """ + import litellm.proxy.proxy_server as proxy_server + + def _configure(*, mappings=None, trusted_proxy_ranges=("127.0.0.1/32",)): + if mappings is None: + mappings = {"user_id": "x-user-id"} + settings = { + "oauth2_config_mappings": mappings, + "trusted_proxy_ranges": trusted_proxy_ranges, + } + monkeypatch.setattr( + proxy_server, + "general_settings", + settings, + raising=False, + ) + + return _configure + + +@pytest.mark.asyncio +async def test_returns_auth_for_simple_user_id_mapping(configure_proxy): + configure_proxy() + request = _request_with_headers({"x-user-id": "alice"}) + + auth = await handle_oauth2_proxy_request(request) + + assert auth.user_id == "alice" + assert auth.user_role is None + + +@pytest.mark.asyncio +async def test_rejects_identity_headers_without_trusted_proxy_ranges(configure_proxy): + configure_proxy(trusted_proxy_ranges=None) + request = _request_with_headers({"x-user-id": "alice"}) + + with pytest.raises(ValueError, match="trusted_proxy_ranges"): + await handle_oauth2_proxy_request(request) + + +@pytest.mark.asyncio +async def test_rejects_identity_headers_from_untrusted_direct_client(configure_proxy): + configure_proxy(trusted_proxy_ranges=["10.0.0.0/24"]) + request = _request_with_headers({"x-user-id": "alice"}, client_host="203.0.113.10") + + with pytest.raises(ValueError, match="not trusted"): + await handle_oauth2_proxy_request(request) + + +@pytest.mark.parametrize( + "privileged_field", + [ + # The GHSA-5c3m-qffq-4r9m primary privesc field. + "user_role", + # Key-level enforcement bypass shapes. + "api_key", + "token", + "permissions", + "allowed_routes", + "max_budget", + "spend", + "tpm_limit", + "rpm_limit", + "model_max_budget", + "metadata", + # User-level enforcement bypass — flagged by Greptile as a denylist gap. + "user_max_budget", + "user_tpm_limit", + "user_rpm_limit", + "user_spend", + # Team / org / end-user / region — same class, all denied by the + # identity-only allowlist. + "team_max_budget", + "team_spend", + "team_member_tpm_limit", + "organization_max_budget", + "organization_tpm_limit", + "end_user_max_budget", + "allowed_model_region", + # Anything not on ALLOWED_OAUTH2_PROXY_FIELDS is blocked, even + # fabricated field names admins might try. + "definitely_not_a_real_field", + ], +) +@pytest.mark.asyncio +async def test_refuses_to_map_non_identity_fields(configure_proxy, privileged_field): + # GHSA-5c3m-qffq-4r9m attack shape: admin maps a privileged field + # to a header and a caller forges the value. The allowlist rejects + # any non-identity mapping at request time, regardless of whether + # the field ever appeared on a denylist — which is the whole reason + # we use an allowlist instead. + configure_proxy(mappings={privileged_field: f"x-{privileged_field}"}) + request = _request_with_headers({f"x-{privileged_field}": "proxy_admin"}) + + with pytest.raises(ValueError) as exc: + await handle_oauth2_proxy_request(request) + assert privileged_field in str(exc.value) + + +@pytest.mark.parametrize("identity_field", sorted(ALLOWED_OAUTH2_PROXY_FIELDS)) +def test_allowlist_is_identity_only(identity_field): + # Lock in the allowlist's intent: only identity-assertion fields are + # safe to populate from a header. If anyone proposes adding budget / + # spend / role / permission to ``ALLOWED_OAUTH2_PROXY_FIELDS``, this + # assertion forces them to update the test deliberately. + assert identity_field in { + "user_id", + "user_email", + "team_id", + "team_alias", + "org_id", + "models", + } + + +@pytest.mark.asyncio +async def test_user_role_header_forgery_attack_is_blocked(configure_proxy): + # End-to-end form of the privesc: with ``user_role`` mapped, the + # forged ``X-User-Role: proxy_admin`` header would have produced + # a ``UserAPIKeyAuth(user_role=PROXY_ADMIN)``. Now rejected before + # any auth object is constructed. + configure_proxy( + mappings={"user_id": "x-user-id", "user_role": "x-user-role"}, + ) + request = _request_with_headers( + { + "x-user-id": "attacker", + "x-user-role": LitellmUserRoles.PROXY_ADMIN.value, + } + ) + + with pytest.raises(ValueError, match="user_role"): + await handle_oauth2_proxy_request(request) + + +@pytest.mark.asyncio +async def test_safe_fields_still_pass_through(configure_proxy): + # The documented use case for OAuth2 proxy auth: identity assertion + # from a trusted upstream. Must remain unaffected by the denylist. + configure_proxy( + mappings={ + "user_id": "x-user-id", + "user_email": "x-user-email", + "team_id": "x-team-id", + "models": "x-models", + }, + ) + request = _request_with_headers( + { + "x-user-id": "alice", + "x-user-email": "alice@example.com", + "x-team-id": "team-corp", + "x-models": "gpt-4, gpt-3.5-turbo", + } + ) + + auth = await handle_oauth2_proxy_request(request) + + assert auth.user_id == "alice" + assert auth.user_email == "alice@example.com" + assert auth.team_id == "team-corp" + assert auth.models == ["gpt-4", "gpt-3.5-turbo"] diff --git a/tests/test_litellm/proxy/auth/test_route_checks.py b/tests/test_litellm/proxy/auth/test_route_checks.py index a5d405cfc2d..39f832256d0 100644 --- a/tests/test_litellm/proxy/auth/test_route_checks.py +++ b/tests/test_litellm/proxy/auth/test_route_checks.py @@ -1198,6 +1198,349 @@ def test_proxy_admin_viewer_can_access_audit_logs(route): ) +# ── Admin Viewer parity: Logs page endpoints ────────────────────────────────── +# +# The Admin Viewer (PROXY_ADMIN_VIEW_ONLY) role is documented as +# "view all keys, view all spend" and follows a read-parity-with-Proxy-Admin +# rule. The UI Logs page is the most user-visible failure mode: filtering and +# log details break entirely when these routes are blocked at the route_checks +# layer, even though the underlying handlers already gate on PROXY_ADMIN_VIEW_ONLY. +# +# Each route below corresponds to a network call made by the Logs page +# (ui/litellm-dashboard/src/components/view_logs/) — see the comment on each. +ADMIN_VIEWER_LOGS_PAGE_ROUTES = [ + # Main paginated log list — uiSpendLogsCall in log_filter_logic.tsx & index.tsx + "/spend/logs/ui", + # Single-log detail drawer — fetched on row click in LogDetailsDrawer + "/spend/logs/ui/abc-request-id", + # Multi-call session drawer — sessionSpendLogsCall in LogDetailsDrawer + "/spend/logs/session/ui", + # End User filter dropdown — allEndUsersCall in index.tsx + "/customer/list", + "/customer/info", + # Cost estimation — used by some log views + "/cost/estimate", + # Public spend logs / spend tracking routes that admin viewer should read + "/spend/logs", + "/spend/keys", + "/spend/users", + "/spend/tags", + "/spend/calculate", +] + + +@pytest.mark.parametrize("route", ADMIN_VIEWER_LOGS_PAGE_ROUTES) +def test_proxy_admin_viewer_can_access_logs_page_endpoints(route): + """ + PROXY_ADMIN_VIEW_ONLY must pass route_checks for every endpoint the UI + Logs page depends on. Without these, the page renders empty / errors. + """ + user_obj = LiteLLM_UserTable( + user_id="viewer_user", + user_email="viewer@example.com", + user_role=LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY.value, + ) + valid_token = UserAPIKeyAuth( + user_id="viewer_user", + user_role=LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY.value, + ) + request = MagicMock(spec=Request) + request.query_params = {} + + try: + RouteChecks.non_proxy_admin_allowed_routes_check( + user_obj=user_obj, + _user_role=LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY.value, + route=route, + request=request, + valid_token=valid_token, + request_data={}, + ) + except Exception as e: + pytest.fail( + f"proxy_admin_viewer should be able to access {route}. Got error: {str(e)}" + ) + + +@pytest.mark.parametrize("route", ADMIN_VIEWER_LOGS_PAGE_ROUTES) +def test_internal_user_blocked_from_admin_viewer_logs_routes(route): + """ + The Logs-page route opening above must NOT also widen access for + INTERNAL_USER. Plain internal users still see only their own logs and + must be blocked from proxy-wide spend tracking + customer routes. + """ + user_obj = LiteLLM_UserTable( + user_id="internal_user", + user_email="user@example.com", + user_role=LitellmUserRoles.INTERNAL_USER.value, + ) + valid_token = UserAPIKeyAuth( + user_id="internal_user", + user_role=LitellmUserRoles.INTERNAL_USER.value, + ) + request = MagicMock(spec=Request) + request.query_params = {} + + # Routes already in `spend_tracking_routes` (which is part of + # `internal_user_routes`) are intentionally accessible to internal users + # for their own scoped spend — those handlers enforce per-user filtering. + # /cost/estimate is similarly per-user. The /customer/* routes are + # admin-only. + INTERNAL_USER_BLOCKED_SUBSET = { + "/customer/list", + "/customer/info", + } + if route not in INTERNAL_USER_BLOCKED_SUBSET: + return + + with pytest.raises(Exception) as exc_info: + RouteChecks.non_proxy_admin_allowed_routes_check( + user_obj=user_obj, + _user_role=LitellmUserRoles.INTERNAL_USER.value, + route=route, + request=request, + valid_token=valid_token, + request_data={}, + ) + assert "Only proxy admin" in str(exc_info.value) + + +# ── Admin Viewer parity: Settings/observability read endpoints ──────────────── +# +# These are GET endpoints accessible to PROXY_ADMIN that the UI exposes to +# admin viewers via sidebar items gated by `all_admin_roles` (which includes +# proxy_admin_viewer). Without these, the Logging & Alerts, Caching, Budgets, +# and Admin Settings pages break for admin viewers. +ADMIN_VIEWER_SETTINGS_ROUTES = [ + # Logging & Alerts page + "/callbacks/list", + "/callbacks/configs", + "/get/config/callbacks", + "/alerting/settings", + # Admin Settings / Router Settings pages + "/config/list", + "/config/field/info", + # Budgets page + "/budget/list", + "/budget/settings", + # Invitation viewing (admin viewer cannot create/delete; can read) + "/invitation/info", + # Guardrails / Policies pages (read-only views) + "/guardrails/list", + "/v2/guardrails/list", + "/guardrails/submissions", + "/guardrails/submissions/some-guardrail-id", + "/guardrails/usage/overview", + "/policies/attachments/list", + # MCP semantic filter settings (read) + "/get/mcp_semantic_filter_settings", + # Model cost map (read-only status / source) + "/schedule/model_cost_map_reload/status", + "/model/cost_map/source", +] + + +@pytest.mark.parametrize("route", ADMIN_VIEWER_SETTINGS_ROUTES) +def test_proxy_admin_viewer_can_access_settings_read_endpoints(route): + """ + PROXY_ADMIN_VIEW_ONLY must pass route_checks for the read-only + settings/observability endpoints exposed in admin-only sidebar groups. + """ + user_obj = LiteLLM_UserTable( + user_id="viewer_user", + user_email="viewer@example.com", + user_role=LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY.value, + ) + valid_token = UserAPIKeyAuth( + user_id="viewer_user", + user_role=LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY.value, + ) + request = MagicMock(spec=Request) + request.query_params = {} + + try: + RouteChecks.non_proxy_admin_allowed_routes_check( + user_obj=user_obj, + _user_role=LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY.value, + route=route, + request=request, + valid_token=valid_token, + request_data={}, + ) + except Exception as e: + pytest.fail( + f"proxy_admin_viewer should be able to access {route}. Got error: {str(e)}" + ) + + +# ── Admin Viewer parity: default-allow GET semantics ───────────────────────── +# +# The route-check layer is structured to default-allow safe HTTP methods +# (GET / HEAD / OPTIONS) for PROXY_ADMIN_VIEW_ONLY. This eliminates the +# whack-a-mole where every newly-added GET endpoint silently 403'd until +# someone remembered to add it to admin_viewer_routes. +# +# These tests pin the new contract: +# - Any GET endpoint not on the LLM/inference path is readable. +# - Any unsafe method (POST/PUT/PATCH/DELETE) outside the explicit allow +# sets is still 403. + +# Routes the user reported as broken in production — they're in disparate +# corners of the codebase and represent the long tail of GETs we'd otherwise +# need to enumerate manually. Default-allow makes them all work. +ADMIN_VIEWER_REPORTED_GET_ROUTES = [ + "/in_product_nudges", + "/health/latest", + "/credentials", + "/v1/mcp/network/client-ip", + "/claude-code/plugins", + "/policy/templates", + # Routes we already had to enumerate manually (regression coverage). + "/spend/logs/ui", + "/customer/list", + "/guardrails/list", + "/policies/attachments/list", + # Hypothetical future GETs — must not require an allowlist entry. + "/some/future/read/endpoint", + "/another/admin-tool/status", +] + + +@pytest.mark.parametrize("route", ADMIN_VIEWER_REPORTED_GET_ROUTES) +def test_proxy_admin_viewer_default_allows_any_get(route): + """ + PROXY_ADMIN_VIEW_ONLY must be able to GET any non-inference endpoint. + + This is a structural guarantee: the route-check defaults to allow for + safe HTTP methods so we don't have to maintain an explicit allowlist. + """ + user_obj = LiteLLM_UserTable( + user_id="viewer_user", + user_email="viewer@example.com", + user_role=LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY.value, + ) + valid_token = UserAPIKeyAuth( + user_id="viewer_user", + user_role=LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY.value, + ) + request = MagicMock(spec=Request) + request.method = "GET" + request.query_params = {} + request.url = MagicMock() + request.url.path = route + + try: + RouteChecks.non_proxy_admin_allowed_routes_check( + user_obj=user_obj, + _user_role=LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY.value, + route=route, + request=request, + valid_token=valid_token, + request_data={}, + ) + except Exception as e: + pytest.fail(f"proxy_admin_viewer GET should default-allow {route!r}. Got: {e}") + + +@pytest.mark.parametrize( + "route", + [ + # Random path that isn't in any allowlist — POST must still 403. + "/some/future/write/endpoint", + # Hard-blocked write routes. + "/user/new", + "/team/new", + "/key/generate", + "/model/new", + ], +) +def test_proxy_admin_viewer_post_blocked_outside_allowlists(route): + """ + Default-allow only applies to safe HTTP methods. POST/PUT/PATCH/DELETE + on a route not in any allow set must still 403. + """ + user_obj = LiteLLM_UserTable( + user_id="viewer_user", + user_email="viewer@example.com", + user_role=LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY.value, + ) + valid_token = UserAPIKeyAuth( + user_id="viewer_user", + user_role=LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY.value, + ) + request = MagicMock(spec=Request) + request.method = "POST" + request.query_params = {} + + with pytest.raises(HTTPException) as exc_info: + RouteChecks.non_proxy_admin_allowed_routes_check( + user_obj=user_obj, + _user_role=LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY.value, + route=route, + request=request, + valid_token=valid_token, + request_data={}, + ) + assert exc_info.value.status_code == 403 + + +# ── Admin Viewer: management_routes write endpoints stay blocked ───────────── +# +# `management_routes` is a mix of reads (info/list, handled via the safe-method +# branch — GET) and writes. The route_checks layer must NOT blanket-allow the +# whole set on POST — that would let Admin Viewer mutate teams, JWT mappings, +# and bulk-update keys, violating the "no writes, ever" rule. +# +# These cases pin the gap closed (Greptile P1 review, 2026-04-30). +ADMIN_VIEWER_MANAGEMENT_ROUTE_WRITES = [ + # Team writes + "/team/block", + "/team/unblock", + "/team/permissions_update", + # JWT key mapping writes + "/jwt/key/mapping/new", + "/jwt/key/mapping/update", + "/jwt/key/mapping/delete", + # Key writes (existing _ADMIN_VIEWER_BLOCKED_WRITE_ROUTES doesn't list bulk + # update or per-key reset-spend, so the management_routes fallback was the + # only thing keeping them out — and it was permissive, not restrictive). + "/key/bulk_update", + "/key/some-key-id/reset_spend", +] + + +@pytest.mark.parametrize("route", ADMIN_VIEWER_MANAGEMENT_ROUTE_WRITES) +def test_proxy_admin_viewer_post_blocked_for_management_route_writes(route): + """ + Admin Viewer must be blocked on POST to write endpoints in + `management_routes`, even when the specific route is not in + `_ADMIN_VIEWER_BLOCKED_WRITE_ROUTES`. + """ + user_obj = LiteLLM_UserTable( + user_id="viewer_user", + user_email="viewer@example.com", + user_role=LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY.value, + ) + valid_token = UserAPIKeyAuth( + user_id="viewer_user", + user_role=LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY.value, + ) + request = MagicMock(spec=Request) + request.method = "POST" + request.query_params = {} + + with pytest.raises(HTTPException) as exc_info: + RouteChecks.non_proxy_admin_allowed_routes_check( + user_obj=user_obj, + _user_role=LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY.value, + route=route, + request=request, + valid_token=valid_token, + request_data={}, + ) + assert exc_info.value.status_code == 403 + + class TestModelsRouteExemptFromDisableLLMEndpoints: """ Test that /models and /v1/models are exempt from DISABLE_LLM_API_ENDPOINTS. diff --git a/tests/test_litellm/proxy/auth/test_router_override_fallback_auth.py b/tests/test_litellm/proxy/auth/test_router_override_fallback_auth.py new file mode 100644 index 00000000000..fc0e9aec501 --- /dev/null +++ b/tests/test_litellm/proxy/auth/test_router_override_fallback_auth.py @@ -0,0 +1,237 @@ +""" +VERIA-44: ``router_settings_override.fallbacks`` must be validated +against the API key's model allowlist at auth time. Without this, the +override is promoted to per-request kwargs after auth and lets a caller +execute requests against models their API key cannot call. +""" + +from typing import List +from unittest.mock import AsyncMock, patch + +import pytest + +from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth +from litellm.proxy.auth.user_api_key_auth import ( + _enforce_key_and_fallback_model_access, + iter_router_fallback_model_names, +) + + +def _key_with_models(models: List[str]) -> UserAPIKeyAuth: + return UserAPIKeyAuth( + api_key="hashed", + user_id="caller", + user_role=LitellmUserRoles.INTERNAL_USER, + models=models, + ) + + +# ── iter_router_fallback_model_names ───────────────────────────────────────── + + +def testiter_router_fallback_model_names_router_config_shape(): + """Router-config shape: ``[{primary: [fallback_list]}]``.""" + assert list( + iter_router_fallback_model_names( + [{"gpt-3.5-turbo": ["gpt-4", "claude-3"]}, {"gpt-4o": ["o1"]}] + ) + ) == ["gpt-4", "claude-3", "o1"] + + +def testiter_router_fallback_model_names_simple_string_shape(): + """Simple top-level shape: list of strings.""" + assert list(iter_router_fallback_model_names(["gpt-4", "claude-3"])) == [ + "gpt-4", + "claude-3", + ] + + +def testiter_router_fallback_model_names_client_side_shape(): + """ClientSideFallbackModel shape: ``[{"model": "..."}]``.""" + assert list( + iter_router_fallback_model_names([{"model": "gpt-4"}, {"model": "claude-3"}]) + ) == ["gpt-4", "claude-3"] + + +def testiter_router_fallback_model_names_empty_or_none(): + assert list(iter_router_fallback_model_names(None)) == [] + assert list(iter_router_fallback_model_names([])) == [] + assert list(iter_router_fallback_model_names("not a list")) == [] + + +# ── _enforce_key_and_fallback_model_access ──────────────────────────────────── + + +@pytest.mark.asyncio +async def test_router_override_fallbacks_validated_against_key_allowlist(): + """A fallback nested inside ``router_settings_override`` is validated + against the API key's allowed models — not just the top-level + ``fallbacks`` field.""" + valid_token = _key_with_models(["gpt-3.5-turbo"]) + request_data = { + "model": "gpt-3.5-turbo", + "router_settings_override": { + "fallbacks": [{"gpt-3.5-turbo": ["unauthorized-model"]}], + }, + } + + seen_models: List[str] = [] + + async def fake_can_key_call_model(model, llm_model_list, valid_token, llm_router): + seen_models.append(model) + + with ( + patch( + "litellm.proxy.auth.user_api_key_auth.can_key_call_model", + side_effect=fake_can_key_call_model, + ), + patch( + "litellm.proxy.auth.user_api_key_auth.is_valid_fallback_model", + new=AsyncMock(), + ), + ): + await _enforce_key_and_fallback_model_access( + valid_token=valid_token, + request_data=request_data, + route="/v1/chat/completions", + request=None, + llm_model_list=None, + llm_router=None, + ) + + # Both the primary model and the override-nested fallback must be + # checked against the API key's allowlist. + assert "gpt-3.5-turbo" in seen_models + assert "unauthorized-model" in seen_models + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "fallback_field", + [ + "fallbacks", + "context_window_fallbacks", + "content_policy_fallbacks", + ], +) +async def test_router_override_all_fallback_fields_validated(fallback_field): + """All three fallback fields the router accepts as per-request kwargs + are validated — context_window_fallbacks and content_policy_fallbacks + are promoted in route_llm_request.py too.""" + valid_token = _key_with_models(["gpt-3.5-turbo"]) + request_data = { + "model": "gpt-3.5-turbo", + "router_settings_override": { + fallback_field: [{"gpt-3.5-turbo": ["smuggled-model"]}], + }, + } + + seen: List[str] = [] + + async def fake_can_key_call_model(model, llm_model_list, valid_token, llm_router): + seen.append(model) + + with ( + patch( + "litellm.proxy.auth.user_api_key_auth.can_key_call_model", + side_effect=fake_can_key_call_model, + ), + patch( + "litellm.proxy.auth.user_api_key_auth.is_valid_fallback_model", + new=AsyncMock(), + ), + ): + await _enforce_key_and_fallback_model_access( + valid_token=valid_token, + request_data=request_data, + route="/v1/chat/completions", + request=None, + llm_model_list=None, + llm_router=None, + ) + + assert "smuggled-model" in seen + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "fallback_field", + [ + "fallbacks", + "context_window_fallbacks", + "content_policy_fallbacks", + ], +) +async def test_top_level_fallback_fields_validated(fallback_field): + """All three top-level fallback fields are forwarded to the router as + per-request kwargs, so all three must be validated against the API + key's allowlist. Greptile P1 follow-up: previously only the + ``fallbacks`` field was walked at the top level.""" + valid_token = _key_with_models(["gpt-3.5-turbo"]) + request_data = { + "model": "gpt-3.5-turbo", + fallback_field: [{"gpt-3.5-turbo": ["top-level-smuggled"]}], + } + + seen: List[str] = [] + + async def fake_can_key_call_model(model, llm_model_list, valid_token, llm_router): + seen.append(model) + + with ( + patch( + "litellm.proxy.auth.user_api_key_auth.can_key_call_model", + side_effect=fake_can_key_call_model, + ), + patch( + "litellm.proxy.auth.user_api_key_auth.is_valid_fallback_model", + new=AsyncMock(), + ), + ): + await _enforce_key_and_fallback_model_access( + valid_token=valid_token, + request_data=request_data, + route="/v1/chat/completions", + request=None, + llm_model_list=None, + llm_router=None, + ) + + assert "top-level-smuggled" in seen + + +@pytest.mark.asyncio +async def test_router_override_without_fallbacks_does_not_break_auth(): + """``router_settings_override`` set without any fallback fields is a + no-op for the auth check — only the primary model is validated.""" + valid_token = _key_with_models(["gpt-3.5-turbo"]) + request_data = { + "model": "gpt-3.5-turbo", + "router_settings_override": {"num_retries": 3, "timeout": 30}, + } + + seen: List[str] = [] + + async def fake_can_key_call_model(model, llm_model_list, valid_token, llm_router): + seen.append(model) + + with ( + patch( + "litellm.proxy.auth.user_api_key_auth.can_key_call_model", + side_effect=fake_can_key_call_model, + ), + patch( + "litellm.proxy.auth.user_api_key_auth.is_valid_fallback_model", + new=AsyncMock(), + ), + ): + await _enforce_key_and_fallback_model_access( + valid_token=valid_token, + request_data=request_data, + route="/v1/chat/completions", + request=None, + llm_model_list=None, + llm_router=None, + ) + + assert seen == ["gpt-3.5-turbo"] diff --git a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py index 08f4bd0ebff..dd24ac87495 100644 --- a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py +++ b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py @@ -1,8 +1,7 @@ -import asyncio import json import os import sys -from typing import Tuple +from types import SimpleNamespace from unittest.mock import ANY, AsyncMock, MagicMock, patch sys.path.insert( @@ -15,6 +14,8 @@ import litellm.proxy.proxy_server from litellm.caching.dual_cache import DualCache from litellm.proxy._types import ( LiteLLM_JWTAuth, + LiteLLM_BudgetTable, + LiteLLM_EndUserTable, LiteLLM_UserTable, LitellmUserRoles, ProxyErrorTypes, @@ -23,8 +24,10 @@ from litellm.proxy._types import ( JWTRoutingOverride, ) from litellm.proxy.auth.handle_jwt import JWTHandler +from litellm.proxy.auth.auth_checks import get_key_object, _cache_key_object from litellm.proxy.auth.route_checks import RouteChecks from litellm.proxy.auth.user_api_key_auth import ( + _reserve_budget_after_common_checks, _run_centralized_common_checks, _run_post_custom_auth_checks, get_api_key, @@ -32,6 +35,13 @@ from litellm.proxy.auth.user_api_key_auth import ( ) +class _RoutingRequest: + def __init__(self, headers=None, query_params=None): + self.headers = headers or {} + self.query_params = query_params or {} + self.state = SimpleNamespace() + + def test_get_api_key(): bearer_token = "Bearer sk-12345678" api_key = "sk-12345678" @@ -49,6 +59,74 @@ def test_get_api_key(): ) == (api_key, passed_in_key) +@pytest.mark.asyncio +async def test_should_clear_stale_budget_reservation_when_budget_checks_skip(): + user_api_key_auth_obj = UserAPIKeyAuth( + token="test_token", + budget_reservation={ + "reserved_cost": 0.5, + "entries": [{"counter_key": "spend:key:test_token"}], + }, + ) + + await _reserve_budget_after_common_checks( + user_api_key_auth_obj=user_api_key_auth_obj, + request_data={"model": "free-model"}, + route="/v1/chat/completions", + llm_router=None, + team_object=None, + user_object=None, + prisma_client=None, + user_api_key_cache=MagicMock(), + proxy_logging_obj=MagicMock(), + skip_budget_checks=True, + ) + + assert user_api_key_auth_obj.budget_reservation is None + + +@pytest.mark.asyncio +async def test_should_not_reuse_cached_key_object_for_request_state(): + key_cache = DualCache() + cached_key = UserAPIKeyAuth( + token="cached-token", + request_route="/old-route", + budget_reservation={ + "reserved_cost": 0.5, + "entries": [{"counter_key": "spend:key:cached-token"}], + }, + ) + + await _cache_key_object( + hashed_token="cached-token", + user_api_key_obj=cached_key, + user_api_key_cache=key_cache, + proxy_logging_obj=None, + ) + + first_request_key = await get_key_object( + hashed_token="cached-token", + prisma_client=MagicMock(), + user_api_key_cache=key_cache, + ) + first_request_key.budget_reservation = { + "reserved_cost": 0.9, + "entries": [{"counter_key": "spend:key:cached-token"}], + } + first_request_key.request_route = "/chat/completions" + + second_request_key = await get_key_object( + hashed_token="cached-token", + prisma_client=MagicMock(), + user_api_key_cache=key_cache, + ) + + assert first_request_key is not cached_key + assert second_request_key is not first_request_key + assert second_request_key.budget_reservation is None + assert second_request_key.request_route is None + + @pytest.mark.asyncio async def test_custom_auth_does_not_enforce_key_model_access_by_default(): valid_token = UserAPIKeyAuth(token="test_token", models=["gpt-4o-mini"]) @@ -107,6 +185,39 @@ async def test_custom_auth_honors_key_level_model_access_restriction_allowed_wit ) +@pytest.mark.asyncio +async def test_custom_auth_enforces_key_model_access_from_file_route_header_with_opt_in(): + valid_token = UserAPIKeyAuth(token="test_token", models=["allowed-model"]) + request = _RoutingRequest(headers={"x-litellm-model": "restricted-model"}) + + with ( + patch( + "litellm.proxy.auth.user_api_key_auth.can_key_call_model", + new_callable=AsyncMock, + ) as mock_can_key, + patch( + "litellm.proxy.auth.user_api_key_auth.common_checks", new_callable=AsyncMock + ), + patch( + "litellm.proxy.proxy_server.general_settings", + {"custom_auth_run_common_checks": True}, + ), + ): + await _run_post_custom_auth_checks( + valid_token=valid_token, + request=request, + request_data={}, + route="/v1/files", + parent_otel_span=None, + ) + mock_can_key.assert_awaited_once_with( + model="restricted-model", + llm_model_list=ANY, + valid_token=valid_token, + llm_router=ANY, + ) + + @pytest.mark.asyncio async def test_custom_auth_honors_key_level_model_access_restriction_denied_with_opt_in(): valid_token = UserAPIKeyAuth(token="test_token", models=["gpt-4o-mini"]) @@ -1752,7 +1863,11 @@ async def test_team_metadata_refreshed_from_team_object_during_auth(): from starlette.datastructures import URL from starlette.requests import Request - from litellm.proxy._types import LiteLLM_TeamTableCachedObj, LitellmUserRoles, UserAPIKeyAuth + from litellm.proxy._types import ( + LiteLLM_TeamTableCachedObj, + LitellmUserRoles, + UserAPIKeyAuth, + ) from litellm.proxy.auth.user_api_key_auth import _user_api_key_auth_builder api_key = "sk-test-team-metadata-refresh" @@ -1833,16 +1948,17 @@ async def test_team_metadata_refreshed_from_team_object_during_auth(): request_data={}, ) - assert result.team_metadata == {"guardrails": ["test-guardrail-333"]}, ( - f"team_metadata was not updated from fresh team object. Got: {result.team_metadata}" - ) + assert result.team_metadata == { + "guardrails": ["test-guardrail-333"] + }, f"team_metadata was not updated from fresh team object. Got: {result.team_metadata}" finally: for k, v in _originals.items(): setattr(_proxy_server_mod, k, v) - + + # --------------------------------------------------------------------------- - + # _run_centralized_common_checks — centralized authz gate # --------------------------------------------------------------------------- @@ -1859,7 +1975,7 @@ def _proxy_attrs_for_centralized_checks( """ return { "prisma_client": None, - "user_api_key_cache": MagicMock(), + "user_api_key_cache": DualCache(), "proxy_logging_obj": MagicMock(), "general_settings": ({"custom_auth_run_common_checks": True} if flag else {}), "llm_router": None, @@ -2120,6 +2236,81 @@ async def test_centralized_common_checks_propagates_end_user_budget_error(): setattr(_proxy_server_mod, k, v) +@pytest.mark.asyncio +async def test_centralized_common_checks_reserves_request_end_user_budget(): + """Regression: reservation runs before user_api_key_auth() copies the + request end-user onto the token, so centralized checks must pass the + locally extracted end_user_id/end_user_object into reservation.""" + import litellm.proxy.proxy_server as _proxy_server_mod + from fastapi import Request + from starlette.datastructures import URL + + token = UserAPIKeyAuth(api_key="sk-test", user_id="u") + request = Request(scope={"type": "http", "headers": []}) + request._url = URL(url="/chat/completions") + request_data = { + "model": "gpt-4o", + "messages": [{"role": "user", "content": "hello"}], + "user": "alice", + } + end_user_object = LiteLLM_EndUserTable( + user_id="alice", + blocked=False, + spend=0.0, + litellm_budget_table=LiteLLM_BudgetTable(max_budget=1.0), + ) + + attrs = _proxy_attrs_for_centralized_checks(user_custom_auth=None) + counter_cache = DualCache() + attrs["spend_counter_cache"] = counter_cache + originals = {a: getattr(_proxy_server_mod, a, None) for a in attrs} + try: + for k, v in attrs.items(): + setattr(_proxy_server_mod, k, v) + with ( + patch( + "litellm.proxy.auth.user_api_key_auth.get_end_user_object", + new_callable=AsyncMock, + return_value=end_user_object, + ), + patch( + "litellm.proxy.auth.user_api_key_auth.common_checks", + new_callable=AsyncMock, + ), + patch( + "litellm.proxy.spend_tracking.budget_reservation.estimate_request_max_cost", + return_value=0.6, + ), + ): + assert token.end_user_id is None + + await _run_centralized_common_checks( + user_api_key_auth_obj=token, + request=request, + request_data=request_data, + route="/chat/completions", + ) + + finally: + for k, v in originals.items(): + setattr(_proxy_server_mod, k, v) + + assert token.end_user_id is None + assert token.budget_reservation is not None + assert token.budget_reservation["entries"] == [ + { + "counter_key": "spend:end_user:alice", + "entity_type": "EndUser", + "entity_id": "alice", + "reserved_cost": 0.6, + "applied_adjustment": 0.0, + } + ] + assert counter_cache.in_memory_cache.get_cache( + key="spend:end_user:alice" + ) == pytest.approx(0.6) + + @pytest.mark.asyncio async def test_centralized_common_checks_short_circuits_when_master_key_unset(): """master_key=None is no-auth dev mode — admin-only routes and diff --git a/tests/test_litellm/proxy/db/test_db_spend_update_writer.py b/tests/test_litellm/proxy/db/test_db_spend_update_writer.py index 4d584349342..c6017752814 100644 --- a/tests/test_litellm/proxy/db/test_db_spend_update_writer.py +++ b/tests/test_litellm/proxy/db/test_db_spend_update_writer.py @@ -1133,8 +1133,10 @@ async def test_update_daily_spend_logs_detailed_error_on_batch_upsert_failure(): mock_proxy_logging = MagicMock() mock_proxy_logging.failure_handler = AsyncMock() - # Mock the logger to capture exception calls - with patch.object(verbose_proxy_logger, "exception") as mock_exception_logger: + # Capture the ERROR-level log emitted by the spend_log_error helper. + # We assert against the formatted message instead of patching a specific + # logger method so the test stays valid as the helper evolves. + with patch.object(verbose_proxy_logger, "error") as mock_error_logger: # Call the method and expect it to raise the exception with pytest.raises(Exception, match="Unique constraint violation"): await DBSpendUpdateWriter._update_daily_spend( @@ -1148,17 +1150,20 @@ async def test_update_daily_spend_logs_detailed_error_on_batch_upsert_failure(): unique_constraint_name="user_id_date_api_key_model_custom_llm_provider_mcp_namespaced_tool_name_endpoint", ) - # Verify that exception was logged with detailed information - assert mock_exception_logger.called - call_args = mock_exception_logger.call_args[0][0] - assert "Daily user spend batch upsert failed" in call_args - assert "Table: litellm_dailyuserspend" in call_args + # Verify that the error was logged with detailed information. + # spend_log_error formats the message via ``%`` interpolation, so + # render the call args before asserting on substrings. + assert mock_error_logger.called + call = mock_error_logger.call_args + formatted = call.args[0] % call.args[1:] + assert "Daily user spend batch upsert failed" in formatted + assert "Table: litellm_dailyuserspend" in formatted assert ( "Constraint: user_id_date_api_key_model_custom_llm_provider_mcp_namespaced_tool_name_endpoint" - in call_args + in formatted ) - assert "Batch size: 1" in call_args - assert "Unique constraint violation" in call_args + assert "Batch size: 1" in formatted + assert "Unique constraint violation" in formatted @pytest.mark.asyncio diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_presidio.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_presidio.py index 38ea42285c1..565bf83c6a2 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_presidio.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_presidio.py @@ -2119,6 +2119,103 @@ async def test_streaming_unmask_path_bytes_passthrough(): assert chunks[0] == byte_chunk +@pytest.mark.asyncio +async def test_apply_to_output_streaming_unknown_events_passthrough(): + """ + Regression test: /v1/responses-style event objects (neither bytes nor + ModelResponseStream) must be preserved in order and not dropped. + """ + guardrail = _OPTIONAL_PresidioPIIMasking( + mock_testing=True, + apply_to_output=True, + ) + + class FakeResponsesEvent: + def __init__(self, event_type: str): + self.type = event_type + + events = [ + FakeResponsesEvent("response.created"), + FakeResponsesEvent("response.output_text.delta"), + FakeResponsesEvent("response.completed"), + ] + + async def mock_stream(): + for event in events: + yield event + + mock_user_api_key = UserAPIKeyAuth(api_key="test-key") + received = [] + async for chunk in guardrail.async_post_call_streaming_iterator_hook( + user_api_key_dict=mock_user_api_key, + response=mock_stream(), + request_data={}, + ): + received.append(chunk) + + # Preserve exact objects and ordering so clients receive full event lifecycle. + assert received == events + assert [e.type for e in received] == [ + "response.created", + "response.output_text.delta", + "response.completed", + ] + + +@pytest.mark.asyncio +async def test_apply_to_output_streaming_mixed_chunks_flushes_and_warns(): + """ + Regression test for mixed stream shape: + a buffered ModelResponseStream chunk followed by unknown responses-style + events should be preserved, and masking skip should be visible via warnings. + """ + guardrail = _OPTIONAL_PresidioPIIMasking( + mock_testing=True, + apply_to_output=True, + ) + + class FakeResponsesEvent: + def __init__(self, event_type: str): + self.type = event_type + + model_chunk = ModelResponseStream( + id="chatcmpl-mixed-1", + choices=[], + created=1, + model="gpt-4", + object="chat.completion.chunk", + system_fingerprint=None, + ) + response_completed = FakeResponsesEvent("response.completed") + + async def mock_stream(): + yield model_chunk + yield response_completed + + mock_user_api_key = UserAPIKeyAuth(api_key="test-key") + received = [] + with patch( + "litellm.proxy.guardrails.guardrail_hooks.presidio.verbose_proxy_logger" + ) as mock_logger: + async for chunk in guardrail.async_post_call_streaming_iterator_hook( + user_api_key_dict=mock_user_api_key, + response=mock_stream(), + request_data={}, + ): + received.append(chunk) + + # Preserve original ordering across mixed stream types. + assert received == [model_chunk, response_completed] + + # Two warnings are expected: + # 1) mixed stream detected + unmasked flush + # 2) passthrough mode skipped output masking + assert mock_logger.warning.call_count == 2 + warning_messages = [call.args[0] for call in mock_logger.warning.call_args_list] + assert any("mixed stream detected" in msg for msg in warning_messages) + assert any("unknown event objects" in msg for msg in warning_messages) + + # --------------------------------------------------------------------------- # Fix 4: apply_guardrail unmask path for input_type="response" # --------------------------------------------------------------------------- diff --git a/tests/test_litellm/proxy/guardrails/test_deferred_guardrail_logging.py b/tests/test_litellm/proxy/guardrails/test_deferred_guardrail_logging.py index 59b6e24f430..e10258c0829 100644 --- a/tests/test_litellm/proxy/guardrails/test_deferred_guardrail_logging.py +++ b/tests/test_litellm/proxy/guardrails/test_deferred_guardrail_logging.py @@ -758,6 +758,60 @@ class TestDeferredStreamingClosure: apply_guardrail_called is False ), "apply_guardrail guardrails must be SKIPPED in deferred path" + @pytest.mark.asyncio + async def test_streaming_iterator_hook_skipped_in_deferred_path(self): + """regression test: guardrails that define async_post_call_streaming_iterator_hook must be SKIPPED in _run_deferred_stream_guardrails. + The iterator hook already scanned the assembled response in the streaming + pipeline""" + success_hook_called = False + + class IteratorHookGuardrail(CustomGuardrail): + def __init__(self): + super().__init__( + guardrail_name="iterator-hook", + default_on=True, + event_hook=GuardrailEventHooks.post_call, + ) + + async def async_post_call_streaming_iterator_hook( + self, user_api_key_dict, response, request_data + ): + async for chunk in response: + yield chunk + + async def async_post_call_success_hook( + self, data: dict, user_api_key_dict: UserAPIKeyAuth, response: Any + ) -> Any: + nonlocal success_hook_called + success_hook_called = True + return response + + mock_logging_obj = MagicMock() + mock_logging_obj.model_call_details = {"metadata": {}} + + async def track_async_success(*args, **kwargs): + pass + + mock_logging_obj.async_success_handler = track_async_success + + guardrail = IteratorHookGuardrail() + + with patch("litellm.callbacks", [guardrail]): + await ProxyBaseLLMRequestProcessing._run_deferred_stream_guardrails( + captured_data={"model": "gpt-4", "metadata": {}}, + captured_user_api_key_dict=UserAPIKeyAuth(api_key="test"), + captured_logging_obj=mock_logging_obj, + assembled_response=MagicMock(), + cache_hit=False, + ) + + await asyncio.sleep(0) + + assert success_hook_called is False, ( + "Guardrails that implement async_post_call_streaming_iterator_hook " + "must be SKIPPED in deferred path — the iterator hook already ran" + ) + @pytest.mark.asyncio async def test_hooks_receive_merged_guardrail_data(self): """Hooks must receive guardrail_data (the merged dict from diff --git a/tests/test_litellm/proxy/guardrails/test_qostodian_nexus_guardrail.py b/tests/test_litellm/proxy/guardrails/test_qostodian_nexus_guardrail.py new file mode 100644 index 00000000000..6daa3e1430d --- /dev/null +++ b/tests/test_litellm/proxy/guardrails/test_qostodian_nexus_guardrail.py @@ -0,0 +1,252 @@ +""" +Unit tests for Qostodian Nexus (by Qohash) integration. + +Tests verify: +1. QostodianNexus can be instantiated with default and custom values +2. Qostodian Nexus is registered in SupportedGuardrailIntegrations +3. Guardrail initializer and class registries contain Qostodian Nexus +4. Configuration parameters are properly passed through +5. QostodianNexusConfigModel works correctly +""" + +import os +import pytest +from unittest.mock import MagicMock + + +def test_qostodian_nexus_initialization_with_defaults(): + """Test QostodianNexus initializes with default values.""" + import os + from unittest.mock import patch + from litellm.proxy.guardrails.guardrail_hooks.qohash import QostodianNexus + + # Unset env var so the hardcoded default is used + env = {k: v for k, v in os.environ.items() if k != "QOSTODIAN_NEXUS_API_BASE"} + with patch.dict(os.environ, env, clear=True): + guardrail = QostodianNexus() + + # Should use default api_base + assert guardrail.api_base is not None + assert "nexus:8800" in guardrail.api_base + + +def test_qostodian_nexus_initialization_with_custom_api_base(): + """Test QostodianNexus initializes with custom api_base.""" + from litellm.proxy.guardrails.guardrail_hooks.qohash import QostodianNexus + + custom_api_base = "http://custom-nexus:9000" + guardrail = QostodianNexus(api_base=custom_api_base) + + assert custom_api_base in guardrail.api_base + + +def test_qostodian_nexus_in_supported_guardrail_integrations(): + """Test that Qostodian Nexus is registered in SupportedGuardrailIntegrations enum.""" + from litellm.types.guardrails import SupportedGuardrailIntegrations + + # Check enum contains QOSTODIAN_NEXUS + assert hasattr(SupportedGuardrailIntegrations, "QOSTODIAN_NEXUS") + assert SupportedGuardrailIntegrations.QOSTODIAN_NEXUS.value == "qostodian_nexus" + + # Check it's in the list of all values + all_values = [e.value for e in SupportedGuardrailIntegrations] + assert "qostodian_nexus" in all_values + + +def test_qostodian_nexus_in_guardrail_initializer_registry(): + """Test that Qostodian Nexus is registered in guardrail_initializer_registry.""" + from litellm.proxy.guardrails.guardrail_hooks.qohash import ( + guardrail_initializer_registry, + ) + + assert "qostodian_nexus" in guardrail_initializer_registry + assert callable(guardrail_initializer_registry["qostodian_nexus"]) + + +def test_qostodian_nexus_in_guardrail_class_registry(): + """Test that Qostodian Nexus is registered in guardrail_class_registry.""" + from litellm.proxy.guardrails.guardrail_hooks.qohash import ( + guardrail_class_registry, + QostodianNexus, + ) + + assert "qostodian_nexus" in guardrail_class_registry + assert guardrail_class_registry["qostodian_nexus"] == QostodianNexus + + +def test_qostodian_nexus_config_model_initialization(): + """Test QostodianNexusConfigModel can be instantiated.""" + from litellm.types.proxy.guardrails.guardrail_hooks.qohash import ( + QostodianNexusConfigModel, + ) + + config = QostodianNexusConfigModel( + api_base="http://test:8800", + ) + + assert config.api_base == "http://test:8800" + + +def test_qostodian_nexus_config_model_defaults(): + """Test QostodianNexusConfigModel uses correct defaults.""" + from litellm.types.proxy.guardrails.guardrail_hooks.qohash import ( + QostodianNexusConfigModel, + ) + + config = QostodianNexusConfigModel() + + assert config.api_base is None + + +def test_qostodian_nexus_config_model_ui_friendly_name(): + """Test QostodianNexusConfigModel returns correct UI friendly name.""" + from litellm.types.proxy.guardrails.guardrail_hooks.qohash import ( + QostodianNexusConfigModel, + ) + + ui_name = QostodianNexusConfigModel.ui_friendly_name() + assert ui_name == "Qostodian Nexus" + + +def test_qostodian_nexus_initializer_function(): + """Test the initialize_guardrail function.""" + from litellm.proxy.guardrails.guardrail_hooks.qohash import initialize_guardrail + from litellm.types.guardrails import LitellmParams, Guardrail + from unittest.mock import patch + + # Mock litellm.logging_callback_manager + with patch("litellm.logging_callback_manager") as mock_manager: + mock_manager.add_litellm_callback = MagicMock() + + # Create test params + litellm_params = LitellmParams( + guardrail="qostodian_nexus", + mode="pre_call", + api_base="http://test:8800", + default_on=True, + ) + + guardrail_config: Guardrail = {"guardrail_name": "test-qostodian-nexus"} + + # Call initializer + result = initialize_guardrail(litellm_params, guardrail_config) + + # Verify callback was added + mock_manager.add_litellm_callback.assert_called_once() + + # Verify returned instance has correct properties + assert result is not None + assert "test:8800" in result.api_base + + +def test_qostodian_nexus_inherits_from_generic_guardrail_api(): + """Test that QostodianNexus inherits from GenericGuardrailAPI.""" + from litellm.proxy.guardrails.guardrail_hooks.qohash import QostodianNexus + from litellm.proxy.guardrails.guardrail_hooks.generic_guardrail_api.generic_guardrail_api import ( + GenericGuardrailAPI, + ) + + assert issubclass(QostodianNexus, GenericGuardrailAPI) + + +def test_qostodian_nexus_guardrail_name_constant(): + """Test that GUARDRAIL_NAME constant is defined correctly.""" + from litellm.proxy.guardrails.guardrail_hooks.qohash.qohash import GUARDRAIL_NAME + + assert GUARDRAIL_NAME == "qostodian_nexus" + + +def test_qostodian_nexus_get_config_model(): + """Test that QostodianNexus returns the correct config model.""" + from litellm.proxy.guardrails.guardrail_hooks.qohash import QostodianNexus + from litellm.types.proxy.guardrails.guardrail_hooks.qohash import ( + QostodianNexusConfigModel, + ) + + config_model = QostodianNexus.get_config_model() + + assert config_model is not None + assert config_model == QostodianNexusConfigModel + + +def test_qostodian_nexus_env_vars(): + """Test that QOSTODIAN_NEXUS_API_BASE env var is picked up correctly.""" + import os + from unittest.mock import patch + from litellm.proxy.guardrails.guardrail_hooks.qohash import QostodianNexus + + with patch.dict(os.environ, {"QOSTODIAN_NEXUS_API_BASE": "http://new-api:8800"}): + guardrail = QostodianNexus() + assert "new-api:8800" in guardrail.api_base + + +def test_qostodian_nexus_config_model_field_descriptions(): + """Test that QostodianNexusConfigModel has correct field descriptions.""" + from litellm.types.proxy.guardrails.guardrail_hooks.qohash import ( + QostodianNexusConfigModel, + ) + + # Check that field descriptions mention the correct env vars + api_base_field = QostodianNexusConfigModel.model_fields["api_base"] + assert "QOSTODIAN_NEXUS_API_BASE" in api_base_field.description + + +def test_qostodian_nexus_unified_detection(): + """ + Test that QostodianNexus is properly detected by LiteLLM's unified guardrail system. + + This verifies the fix for the detection bug where QostodianNexus wasn't being + recognized because apply_guardrail was only inherited, not in the class's own __dict__. + """ + from litellm.proxy.guardrails.guardrail_hooks.qohash import QostodianNexus + + # Create an instance (this is how LiteLLM uses it) + instance = QostodianNexus(api_base="http://test:8800") + + # Test the exact detection logic used in litellm/proxy/utils.py:868 + # use_unified = "apply_guardrail" in type(callback).__dict__ + use_unified = "apply_guardrail" in type(instance).__dict__ + + # Should be detected as using unified guardrail system + assert use_unified is True, ( + "QostodianNexus should be detected by unified guardrail system. " + "The apply_guardrail method must be present in QostodianNexus.__dict__" + ) + + # Also verify the method is callable + assert hasattr(instance, "apply_guardrail") + assert callable(instance.apply_guardrail) + + +def test_qostodian_nexus_builtin_extra_headers(): + """Test that QostodianNexus includes built-in x-qostodian-nexus-identifiers-* headers.""" + from litellm.proxy.guardrails.guardrail_hooks.qohash import QostodianNexus + + instance = QostodianNexus() + + expected_headers = [ + "x-qostodian-nexus-identifiers-trace", + "x-qostodian-nexus-identifiers-source", + "x-qostodian-nexus-identifiers-container", + "x-qostodian-nexus-identifiers-identity", + ] + + for header in expected_headers: + assert header in instance.extra_headers, ( + f"Expected built-in header '{header}' to be in extra_headers" + ) + + +def test_qostodian_nexus_extra_headers_merged(): + """Test that caller-supplied extra_headers are merged with built-in headers.""" + from litellm.proxy.guardrails.guardrail_hooks.qohash import QostodianNexus + + custom_header = "x-custom-correlation-id" + instance = QostodianNexus(extra_headers=[custom_header]) + + # Built-in headers should be present + assert "x-qostodian-nexus-identifiers-trace" in instance.extra_headers + # Custom header should also be present + assert custom_header in instance.extra_headers + # No duplicates + assert len(instance.extra_headers) == len(set(instance.extra_headers)) diff --git a/tests/test_litellm/proxy/hooks/test_batch_file_validation.py b/tests/test_litellm/proxy/hooks/test_batch_file_validation.py new file mode 100644 index 00000000000..7f1006543bb --- /dev/null +++ b/tests/test_litellm/proxy/hooks/test_batch_file_validation.py @@ -0,0 +1,285 @@ +""" +VERIA-39 regression tests: + +- The batch input-file token counter must measure embeddings (`input`) + and text-completion (`prompt`) payloads, not only chat (`messages`). +- The batch rate-limiter pre-call hook must reject batch files that name + models the caller is not authorized to use. +""" + +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest +from fastapi import HTTPException + +from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth + + +# --------------------------------------------------------------------------- +# Token counter — covers all three batch payload shapes +# --------------------------------------------------------------------------- + + +def test_token_counter_counts_chat_messages(): + from litellm.batches.batch_utils import _get_batch_job_input_file_usage + + usage = _get_batch_job_input_file_usage( + file_content_dictionary=[ + { + "body": { + "model": "gpt-4o-mini", + "messages": [{"role": "user", "content": "hello"}], + } + } + ] + ) + assert usage.prompt_tokens > 0 + + +def test_token_counter_counts_text_completion_prompt(): + """Pre-fix this returned 0 tokens (the function only inspected + `messages`), letting `prompt`-style batches slip past TPM limits.""" + from litellm.batches.batch_utils import _get_batch_job_input_file_usage + + usage = _get_batch_job_input_file_usage( + file_content_dictionary=[ + {"body": {"model": "gpt-3.5-turbo-instruct", "prompt": "hello world"}} + ] + ) + assert usage.prompt_tokens > 0 + + +def test_token_counter_counts_embedding_input_string(): + from litellm.batches.batch_utils import _get_batch_job_input_file_usage + + usage = _get_batch_job_input_file_usage( + file_content_dictionary=[ + {"body": {"model": "text-embedding-3-small", "input": "hello world"}} + ] + ) + assert usage.prompt_tokens > 0 + + +def test_token_counter_counts_embedding_input_list(): + from litellm.batches.batch_utils import _get_batch_job_input_file_usage + + usage = _get_batch_job_input_file_usage( + file_content_dictionary=[ + { + "body": { + "model": "text-embedding-3-small", + "input": ["hello", "world"], + } + } + ] + ) + assert usage.prompt_tokens > 0 + + +def test_token_counter_counts_text_completion_prompt_list(): + from litellm.batches.batch_utils import _get_batch_job_input_file_usage + + usage = _get_batch_job_input_file_usage( + file_content_dictionary=[ + { + "body": { + "model": "gpt-3.5-turbo-instruct", + "prompt": ["alpha", "beta"], + } + } + ] + ) + assert usage.prompt_tokens > 0 + + +def test_token_counter_counts_pre_tokenized_prompt_int_list(): + """OpenAI's text-completion API accepts a single pre-tokenized prompt as + a list of ints. Each int is one token; pre-fix this shape was silently + counted as zero, leaving a TPM bypass.""" + from litellm.batches.batch_utils import _get_batch_job_input_file_usage + + usage = _get_batch_job_input_file_usage( + file_content_dictionary=[ + { + "body": { + "model": "gpt-3.5-turbo-instruct", + "prompt": [1, 2, 3, 4, 5], + } + } + ] + ) + assert usage.prompt_tokens == 5 + + +def test_token_counter_counts_pre_tokenized_prompt_list_of_int_lists(): + """Multiple pre-tokenized prompts (`list[list[int]]`) — the most + important bypass shape. A 1000-token batch must report 1000 tokens, + not zero.""" + from litellm.batches.batch_utils import _get_batch_job_input_file_usage + + usage = _get_batch_job_input_file_usage( + file_content_dictionary=[ + { + "body": { + "model": "gpt-3.5-turbo-instruct", + "prompt": [[1] * 250, [2] * 250, [3] * 500], + } + } + ] + ) + assert usage.prompt_tokens == 1000 + + +def test_token_counter_counts_pre_tokenized_input_for_embeddings(): + """Same shape applies to embeddings (`input`).""" + from litellm.batches.batch_utils import _get_batch_job_input_file_usage + + usage = _get_batch_job_input_file_usage( + file_content_dictionary=[ + { + "body": { + "model": "text-embedding-3-small", + "input": [[1, 2, 3], [4, 5, 6]], + } + } + ] + ) + assert usage.prompt_tokens == 6 + + +# --------------------------------------------------------------------------- +# Model extractor +# --------------------------------------------------------------------------- + + +def test_model_extractor_returns_distinct_models(): + from litellm.batches.batch_utils import _get_models_from_batch_input_file_content + + models = _get_models_from_batch_input_file_content( + [ + {"body": {"model": "gpt-4o", "messages": []}}, + {"body": {"model": "gpt-4o", "messages": []}}, # duplicate + {"body": {"model": "gpt-4o-mini", "messages": []}}, + {"body": {}}, # missing model + ] + ) + assert models == ["gpt-4o", "gpt-4o-mini"] + + +# --------------------------------------------------------------------------- +# Pre-call hook model validation +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_pre_call_rejects_unauthorized_model_in_batch_file(): + """Pre-fix the hook only validated the outer `model` parameter and + forwarded the file as-is. With this fix, a model named inside the + JSONL that the caller cannot use must trigger a 403.""" + from litellm.proxy.hooks.batch_rate_limiter import _PROXY_BatchRateLimiter + + rate_limiter = _PROXY_BatchRateLimiter( + internal_usage_cache=MagicMock(), + parallel_request_limiter=MagicMock(), + ) + + # Simulated decoded batch file: caller is restricted to gpt-3.5 + # but the JSONL points at gpt-4o. + file_dict = [ + {"body": {"model": "gpt-4o", "messages": [{"role": "user", "content": "x"}]}} + ] + + user = UserAPIKeyAuth( + api_key="sk-restricted", + user_id="alice", + models=["gpt-3.5-turbo"], + user_role=LitellmUserRoles.INTERNAL_USER.value, + ) + + # `can_key_call_model` raises a ProxyException for non-allowed models. + async def _raise_unauthorized(**kwargs): + raise Exception( + f"Key not allowed to access model. This key only has access to models={kwargs['valid_token'].models}" + ) + + with ( + patch( + "litellm.proxy.auth.auth_checks.can_key_call_model", + new=AsyncMock(side_effect=_raise_unauthorized), + ), + patch("litellm.proxy.proxy_server.llm_router", None), + ): + with pytest.raises(HTTPException) as exc: + await rate_limiter._enforce_batch_file_model_access( + user_api_key_dict=user, + file_content_as_dict=file_dict, + ) + + assert exc.value.status_code == 403 + assert "gpt-4o" in str(exc.value.detail) + + +@pytest.mark.asyncio +async def test_pre_call_allows_authorized_model_in_batch_file(): + """If every model in the JSONL is on the caller's allowlist, the hook + must not raise.""" + from litellm.proxy.hooks.batch_rate_limiter import _PROXY_BatchRateLimiter + + rate_limiter = _PROXY_BatchRateLimiter( + internal_usage_cache=MagicMock(), + parallel_request_limiter=MagicMock(), + ) + + file_dict = [ + { + "body": { + "model": "gpt-3.5-turbo", + "messages": [{"role": "user", "content": "x"}], + } + } + ] + + user = UserAPIKeyAuth( + api_key="sk-ok", + user_id="alice", + models=["gpt-3.5-turbo"], + user_role=LitellmUserRoles.INTERNAL_USER.value, + ) + + with ( + patch( + "litellm.proxy.auth.auth_checks.can_key_call_model", + new=AsyncMock(return_value=True), + ), + patch("litellm.proxy.proxy_server.llm_router", None), + ): + # Should not raise + await rate_limiter._enforce_batch_file_model_access( + user_api_key_dict=user, + file_content_as_dict=file_dict, + ) + + +@pytest.mark.asyncio +async def test_pre_call_skips_check_when_no_models_present(): + """Files without any `body.model` (corrupt or empty) must not 500; + the rate limiter logs a warning elsewhere and proceeds.""" + from litellm.proxy.hooks.batch_rate_limiter import _PROXY_BatchRateLimiter + + rate_limiter = _PROXY_BatchRateLimiter( + internal_usage_cache=MagicMock(), + parallel_request_limiter=MagicMock(), + ) + user = UserAPIKeyAuth(api_key="sk-ok", user_id="alice") + + # Should not raise even though `can_key_call_model` is the default + # (would fail). The early-return on empty models keeps the call out + # entirely. + await rate_limiter._enforce_batch_file_model_access( + user_api_key_dict=user, + file_content_as_dict=[], + ) + await rate_limiter._enforce_batch_file_model_access( + user_api_key_dict=user, + file_content_as_dict=[{"body": {}}], + ) diff --git a/tests/test_litellm/proxy/hooks/test_max_budget_limiter.py b/tests/test_litellm/proxy/hooks/test_max_budget_limiter.py new file mode 100644 index 00000000000..0074d7062b8 --- /dev/null +++ b/tests/test_litellm/proxy/hooks/test_max_budget_limiter.py @@ -0,0 +1,208 @@ +""" +Unit tests for the personal-budget pre-call hook. + +The reservation path (added in PR #26845) atomically pre-fills the same +`spend:user:{user_id}` counter this hook reads, admitting at a strict-`<` +boundary. Re-checking with `>=` after reservation would reject requests the +reservation already admitted when the reservation fills the counter to +exactly `max_budget` (e.g. requests with no `max_tokens` cap fall back to +reserving the smallest remaining headroom). + +These tests pin the skip-when-reserved behavior and guard against drift. +""" + +from unittest.mock import AsyncMock, patch + +import pytest +from fastapi import HTTPException + +from litellm.caching.caching import DualCache +from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.hooks.max_budget_limiter import _PROXY_MaxBudgetLimiter + + +def _make_user_api_key_auth( + user_id: str = "user-1", + user_max_budget: float = 10.0, + user_spend: float = 0.0, + team_id=None, + budget_reservation=None, +) -> UserAPIKeyAuth: + return UserAPIKeyAuth( + api_key="sk-test", + user_id=user_id, + user_max_budget=user_max_budget, + user_spend=user_spend, + team_id=team_id, + budget_reservation=budget_reservation, + ) + + +@pytest.mark.asyncio +async def test_under_budget_passes(): + handler = _PROXY_MaxBudgetLimiter() + user_api_key_dict = _make_user_api_key_auth(user_max_budget=10.0) + + with patch( + "litellm.proxy.proxy_server.get_current_spend", + new=AsyncMock(return_value=3.0), + ): + result = await handler.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=DualCache(), + data={}, + call_type="completion", + ) + + assert result is None + + +@pytest.mark.asyncio +async def test_over_budget_rejects_without_reservation(): + handler = _PROXY_MaxBudgetLimiter() + user_api_key_dict = _make_user_api_key_auth(user_max_budget=10.0) + + with patch( + "litellm.proxy.proxy_server.get_current_spend", + new=AsyncMock(return_value=10.0), + ): + with pytest.raises(HTTPException) as exc_info: + await handler.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=DualCache(), + data={}, + call_type="completion", + ) + + assert exc_info.value.status_code == 429 + assert "Max budget limit reached." in exc_info.value.detail + + +@pytest.mark.asyncio +async def test_skips_when_user_counter_is_reserved(): + """ + Reservation atomically pre-fills `spend:user:{user_id}` and admits the + request. The legacy `>=` check must not double-enforce on the same + counter — that's what produced the boundary regression where a fresh + user with no `max_tokens` cap got 429'd on their first request. + """ + handler = _PROXY_MaxBudgetLimiter() + user_api_key_dict = _make_user_api_key_auth( + user_id="user-1", + user_max_budget=10.0, + budget_reservation={ + "reserved_cost": 10.0, + "entries": [ + { + "counter_key": "spend:user:user-1", + "entity_type": "User", + "entity_id": "user-1", + "reserved_cost": 10.0, + "applied_adjustment": 0.0, + } + ], + "finalized": False, + }, + ) + + # `get_current_spend` would return 10.0 here (counter pre-filled by the + # reservation). The hook must skip without reading it. + with patch( + "litellm.proxy.proxy_server.get_current_spend", + new=AsyncMock(return_value=10.0), + ) as mock_get_spend: + result = await handler.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=DualCache(), + data={}, + call_type="completion", + ) + + assert result is None + mock_get_spend.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_does_not_skip_when_reservation_covers_a_different_counter(): + """ + A reservation that only covers e.g. `spend:team:{team_id}` (not the user + counter) must not exempt the user-budget check. + """ + handler = _PROXY_MaxBudgetLimiter() + user_api_key_dict = _make_user_api_key_auth( + user_id="user-1", + user_max_budget=10.0, + budget_reservation={ + "reserved_cost": 5.0, + "entries": [ + { + "counter_key": "spend:team:team-x", + "entity_type": "Team", + "entity_id": "team-x", + "reserved_cost": 5.0, + "applied_adjustment": 0.0, + } + ], + "finalized": False, + }, + ) + + with patch( + "litellm.proxy.proxy_server.get_current_spend", + new=AsyncMock(return_value=10.0), + ): + with pytest.raises(HTTPException) as exc_info: + await handler.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=DualCache(), + data={}, + call_type="completion", + ) + + assert exc_info.value.status_code == 429 + + +@pytest.mark.asyncio +async def test_team_keys_skip_personal_budget(): + handler = _PROXY_MaxBudgetLimiter() + user_api_key_dict = _make_user_api_key_auth( + user_max_budget=10.0, + team_id="team-1", + ) + + with patch( + "litellm.proxy.proxy_server.get_current_spend", + new=AsyncMock(return_value=999.0), + ) as mock_get_spend: + result = await handler.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=DualCache(), + data={}, + call_type="completion", + ) + + assert result is None + mock_get_spend.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_no_max_budget_passes(): + handler = _PROXY_MaxBudgetLimiter() + user_api_key_dict = UserAPIKeyAuth( + api_key="sk-test", + user_id="user-1", + ) + + with patch( + "litellm.proxy.proxy_server.get_current_spend", + new=AsyncMock(return_value=999.0), + ) as mock_get_spend: + result = await handler.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=DualCache(), + data={}, + call_type="completion", + ) + + assert result is None + mock_get_spend.assert_not_awaited() diff --git a/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py b/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py index d4fb5b72719..370477c3605 100644 --- a/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py +++ b/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py @@ -1029,6 +1029,83 @@ async def test_team_member_rate_limits_v3(): ), "Team member TPM limit should be set" +@pytest.mark.asyncio +async def test_team_member_rate_limits_v3_raises_429_when_over_limit(): + """ + When should_rate_limit reports OVER_LIMIT for the team_member descriptor, the + pre-call hook raises HTTP 429 with rate_limit headers — same contract as + test_rpm_api_key_rate_limits_v3 / test_tpm_api_key_rate_limits_v3. + """ + _api_key = hash_token("sk-12345") + _team_id = "team_123" + _user_id = "user_456" + + user_api_key_dict = UserAPIKeyAuth( + api_key=_api_key, + team_id=_team_id, + user_id=_user_id, + team_member_rpm_limit=10, + team_member_tpm_limit=1000, + ) + + local_cache = DualCache() + parallel_request_handler = _PROXY_MaxParallelRequestsHandler( + internal_usage_cache=InternalUsageCache(local_cache) + ) + + captured_descriptors = None + + async def mock_should_rate_limit(descriptors, **kwargs): + nonlocal captured_descriptors + captured_descriptors = descriptors + return { + "overall_code": "OVER_LIMIT", + "statuses": [ + { + "code": "OVER_LIMIT", + "current_limit": 10, + "limit_remaining": -1, + "rate_limit_type": "requests", + "descriptor_key": "team_member", + }, + { + "code": "OK", + "current_limit": 1000, + "limit_remaining": 500, + "rate_limit_type": "tokens", + "descriptor_key": "team_member", + }, + ], + } + + parallel_request_handler.should_rate_limit = mock_should_rate_limit + + error = None + try: + await parallel_request_handler.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=local_cache, + data={"model": "gpt-3.5-turbo"}, + call_type="", + ) + except HTTPException as e: + error = e + assert e.status_code == 429 + assert "rate_limit_type" in e.headers + assert e.headers.get("rate_limit_type") == "requests" + assert "retry-after" in e.headers + + assert error is not None, "An Exception must be thrown" + assert captured_descriptors is not None, "Rate limit descriptors should be captured" + team_member_descriptor = None + for descriptor in captured_descriptors: + if descriptor["key"] == "team_member": + team_member_descriptor = descriptor + break + assert team_member_descriptor is not None + assert team_member_descriptor["value"] == f"{_team_id}:{_user_id}" + + @pytest.mark.asyncio async def test_dynamic_rate_limiting_v3(): """ diff --git a/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py b/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py index 65e7f744c85..771e10a54a0 100644 --- a/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py +++ b/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py @@ -1,9 +1,7 @@ -import json import os import sys import pytest -from fastapi.testclient import TestClient sys.path.insert( 0, os.path.abspath("../../../..") @@ -13,8 +11,11 @@ from datetime import datetime from unittest.mock import AsyncMock, MagicMock, patch from litellm.proxy._types import UserAPIKeyAuth -from litellm.proxy.hooks.proxy_track_cost_callback import _ProxyDBLogger -from litellm.types.utils import StandardLoggingPayload +from litellm.proxy.hooks.proxy_track_cost_callback import ( + _ProxyDBLogger, + _get_budget_reservation_from_metadata, + _update_database_and_spend_counters, +) @pytest.mark.asyncio @@ -62,7 +63,6 @@ async def test_async_post_call_failure_hook(): # Check the arguments passed to update_database call_args = mock_update_database.call_args[1] - print("call_args", json.dumps(call_args, indent=4, default=str)) assert call_args["token"] == "test_api_key" assert call_args["response_cost"] == 0.0 assert call_args["user_id"] == "test_user_id" @@ -128,6 +128,440 @@ async def test_async_post_call_failure_hook_non_llm_route(): mock_update_database.assert_not_called() +@pytest.mark.asyncio +async def test_async_post_call_failure_hook_releases_budget_reservation_before_route_skip(): + logger = _ProxyDBLogger() + budget_reservation = {"reserved_cost": 0.5, "entries": []} + user_api_key_dict = UserAPIKeyAuth( + api_key="test_api_key", + request_route="/custom/route", + budget_reservation=budget_reservation, + ) + + with ( + patch( + "litellm.proxy.spend_tracking.budget_reservation.release_budget_reservation", + new_callable=AsyncMock, + ) as mock_release_budget_reservation, + patch( + "litellm.proxy.db.db_spend_update_writer.DBSpendUpdateWriter.update_database", + new_callable=AsyncMock, + ) as mock_update_database, + ): + await logger.async_post_call_failure_hook( + request_data={}, + original_exception=Exception("Test exception"), + user_api_key_dict=user_api_key_dict, + ) + + assert mock_release_budget_reservation.await_count == 1 + assert ( + mock_release_budget_reservation.await_args.kwargs["budget_reservation"] + is user_api_key_dict.budget_reservation + ) + mock_update_database.assert_not_called() + + +@pytest.mark.asyncio +async def test_should_continue_failure_tracking_when_budget_release_fails(): + logger = _ProxyDBLogger() + budget_reservation = {"reserved_cost": 0.5, "entries": []} + user_api_key_dict = UserAPIKeyAuth( + api_key="test_api_key", + user_id="test_user_id", + team_id="test_team_id", + request_route="/chat/completions", + budget_reservation=budget_reservation, + ) + + with ( + patch( + "litellm.proxy.spend_tracking.budget_reservation.release_budget_reservation", + new_callable=AsyncMock, + side_effect=RuntimeError("redis unavailable"), + ) as mock_release_budget_reservation, + patch( + "litellm.proxy.hooks.proxy_track_cost_callback._invalidate_budget_reservation_counters", + new_callable=AsyncMock, + ) as mock_invalidate_budget_reservation_counters, + patch( + "litellm.proxy.db.db_spend_update_writer.DBSpendUpdateWriter.update_database", + new_callable=AsyncMock, + ) as mock_update_database, + patch( + "litellm.proxy.hooks.proxy_track_cost_callback.verbose_proxy_logger.exception", + ) as mock_log_exception, + ): + await logger.async_post_call_failure_hook( + request_data={ + "model": "gpt-4", + "messages": [{"role": "user", "content": "Hello"}], + }, + original_exception=Exception("provider failed"), + user_api_key_dict=user_api_key_dict, + ) + + assert mock_release_budget_reservation.await_count == 1 + assert ( + mock_release_budget_reservation.await_args.kwargs["budget_reservation"] + is user_api_key_dict.budget_reservation + ) + assert mock_invalidate_budget_reservation_counters.await_count == 1 + assert ( + mock_invalidate_budget_reservation_counters.await_args.kwargs[ + "budget_reservation" + ] + is user_api_key_dict.budget_reservation + ) + assert user_api_key_dict.budget_reservation["finalized"] is True + mock_log_exception.assert_called_once() + mock_update_database.assert_called_once() + + +@pytest.mark.asyncio +async def test_track_cost_callback_releases_budget_reservation_when_spend_tracking_skips(): + logger = _ProxyDBLogger() + budget_reservation = {"reserved_cost": 0.5, "entries": []} + user_api_key_auth = UserAPIKeyAuth(budget_reservation=budget_reservation) + + kwargs = { + "model": "gpt-4", + "litellm_params": { + "metadata": { + "user_api_key_auth": user_api_key_auth, + }, + }, + "standard_logging_object": { + "response_cost": 0.1, + "request_tags": None, + }, + "stream": False, + } + + with patch( + "litellm.proxy.spend_tracking.budget_reservation.release_budget_reservation", + new_callable=AsyncMock, + ) as mock_release_budget_reservation: + await logger._PROXY_track_cost_callback( + kwargs=kwargs, + completion_response=None, + start_time=datetime.now(), + end_time=datetime.now(), + ) + + mock_release_budget_reservation.assert_awaited_once_with( + budget_reservation=budget_reservation, + ) + + +@pytest.mark.asyncio +async def test_track_cost_callback_releases_budget_reservation_when_response_cost_missing(): + logger = _ProxyDBLogger() + budget_reservation = {"reserved_cost": 0.5, "entries": []} + user_api_key_auth = UserAPIKeyAuth(budget_reservation=budget_reservation) + + kwargs = { + "model": "gpt-4", + "call_type": "acompletion", + "litellm_params": { + "metadata": { + "user_api_key_auth": user_api_key_auth, + }, + }, + "standard_logging_object": { + "response_cost": None, + "response_cost_failure_debug_info": "missing custom price", + "request_tags": None, + }, + "stream": False, + } + + with ( + patch( + "litellm.proxy.proxy_server.proxy_logging_obj", + ) as mock_proxy_logging, + patch( + "litellm.proxy.spend_tracking.budget_reservation.release_budget_reservation", + new_callable=AsyncMock, + ) as mock_release_budget_reservation, + ): + mock_proxy_logging.failed_tracking_alert = AsyncMock() + + await logger._PROXY_track_cost_callback( + kwargs=kwargs, + completion_response=None, + start_time=datetime.now(), + end_time=datetime.now(), + ) + + mock_release_budget_reservation.assert_awaited_once_with( + budget_reservation=budget_reservation, + ) + + +def test_get_budget_reservation_from_metadata_handles_dict_auth_object(): + budget_reservation = { + "reserved_cost": 0.5, + "entries": [{"counter_key": "spend:key:test_api_key"}], + } + + assert ( + _get_budget_reservation_from_metadata( + metadata={"user_api_key_auth": dict(UserAPIKeyAuth())} + ) + is None + ) + assert ( + _get_budget_reservation_from_metadata( + metadata={ + "user_api_key_auth": UserAPIKeyAuth( + budget_reservation=budget_reservation + ) + } + ) + == budget_reservation + ) + assert ( + _get_budget_reservation_from_metadata( + metadata={ + "user_api_key_auth": dict( + UserAPIKeyAuth(budget_reservation=budget_reservation) + ) + } + ) + == budget_reservation + ) + assert ( + _get_budget_reservation_from_metadata( + metadata={"user_api_key_budget_reservation": budget_reservation} + ) + is budget_reservation + ) + + +@pytest.mark.asyncio +async def test_update_database_and_spend_counters_releases_reservation_when_db_update_fails(): + proxy_logging_obj = MagicMock() + proxy_logging_obj.db_spend_update_writer.update_database = AsyncMock( + side_effect=Exception("db unavailable") + ) + increment_spend_counters = AsyncMock() + budget_reservation = {"reserved_cost": 0.5, "entries": []} + + with patch( + "litellm.proxy.spend_tracking.budget_reservation.release_budget_reservation", + new_callable=AsyncMock, + ) as mock_release_budget_reservation: + with pytest.raises(Exception, match="db unavailable"): + await _update_database_and_spend_counters( + proxy_logging_obj=proxy_logging_obj, + increment_spend_counters=increment_spend_counters, + user_api_key="test_api_key", + user_id="test_user_id", + end_user_id=None, + team_id="test_team_id", + org_id="test_org_id", + kwargs={}, + completion_response=None, + start_time=datetime.now(), + end_time=datetime.now(), + response_cost=0.2, + budget_reservation=budget_reservation, + ) + + mock_release_budget_reservation.assert_awaited_once_with( + budget_reservation=budget_reservation, + ) + + increment_spend_counters.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_update_database_and_spend_counters_preserves_db_exception_when_release_fails(): + proxy_logging_obj = MagicMock() + db_exception = RuntimeError("db unavailable") + proxy_logging_obj.db_spend_update_writer.update_database = AsyncMock( + side_effect=db_exception + ) + increment_spend_counters = AsyncMock() + budget_reservation = {"reserved_cost": 0.5, "entries": []} + + with ( + patch( + "litellm.proxy.spend_tracking.budget_reservation.release_budget_reservation", + new_callable=AsyncMock, + side_effect=RuntimeError("release unavailable"), + ) as mock_release_budget_reservation, + patch( + "litellm.proxy.hooks.proxy_track_cost_callback.verbose_proxy_logger.exception", + ) as mock_log_exception, + patch( + "litellm.proxy.hooks.proxy_track_cost_callback._invalidate_budget_reservation_counters", + new_callable=AsyncMock, + side_effect=RuntimeError("invalidate unavailable"), + ) as mock_invalidate_budget_reservation_counters, + ): + with pytest.raises(RuntimeError) as exc_info: + await _update_database_and_spend_counters( + proxy_logging_obj=proxy_logging_obj, + increment_spend_counters=increment_spend_counters, + user_api_key="test_api_key", + user_id="test_user_id", + end_user_id=None, + team_id="test_team_id", + org_id="test_org_id", + kwargs={}, + completion_response=None, + start_time=datetime.now(), + end_time=datetime.now(), + response_cost=0.2, + budget_reservation=budget_reservation, + ) + + assert exc_info.value is db_exception + mock_release_budget_reservation.assert_awaited_once_with( + budget_reservation=budget_reservation, + ) + mock_invalidate_budget_reservation_counters.assert_awaited_once_with( + budget_reservation=budget_reservation, + ) + assert mock_log_exception.call_count == 2 + mock_log_exception.assert_any_call( + "Failed to release budget reservation after database update failed" + ) + mock_log_exception.assert_any_call( + "Failed to invalidate budget reservation counters after release failed" + ) + + increment_spend_counters.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_update_database_and_spend_counters_updates_counters_after_db_update(): + proxy_logging_obj = MagicMock() + proxy_logging_obj.db_spend_update_writer.update_database = AsyncMock() + increment_spend_counters = AsyncMock() + budget_reservation = {"reserved_cost": 0.5, "entries": []} + + await _update_database_and_spend_counters( + proxy_logging_obj=proxy_logging_obj, + increment_spend_counters=increment_spend_counters, + user_api_key="test_api_key", + user_id="test_user_id", + end_user_id="test_end_user_id", + team_id="test_team_id", + org_id="test_org_id", + kwargs={}, + completion_response=None, + start_time=datetime.now(), + end_time=datetime.now(), + response_cost=0.2, + budget_reservation=budget_reservation, + request_tags=["tag-a"], + ) + + proxy_logging_obj.db_spend_update_writer.update_database.assert_awaited_once() + increment_spend_counters.assert_awaited_once_with( + token="test_api_key", + team_id="test_team_id", + user_id="test_user_id", + response_cost=0.2, + org_id="test_org_id", + budget_reservation=budget_reservation, + end_user_id="test_end_user_id", + tags=["tag-a"], + ) + + +@pytest.mark.asyncio +async def test_update_database_and_spend_counters_invalidates_reservation_when_counter_update_fails(): + proxy_logging_obj = MagicMock() + proxy_logging_obj.db_spend_update_writer.update_database = AsyncMock() + increment_spend_counters = AsyncMock(side_effect=Exception("counter unavailable")) + budget_reservation = { + "reserved_cost": 0.5, + "entries": [{"counter_key": "spend:key:test_api_key"}], + } + + with patch( + "litellm.proxy.spend_tracking.budget_reservation.invalidate_budget_reservation_counters", + new_callable=AsyncMock, + ) as mock_invalidate_budget_reservation_counters: + with pytest.raises(Exception, match="counter unavailable"): + await _update_database_and_spend_counters( + proxy_logging_obj=proxy_logging_obj, + increment_spend_counters=increment_spend_counters, + user_api_key="test_api_key", + user_id="test_user_id", + end_user_id=None, + team_id="test_team_id", + org_id="test_org_id", + kwargs={}, + completion_response=None, + start_time=datetime.now(), + end_time=datetime.now(), + response_cost=0.2, + budget_reservation=budget_reservation, + ) + + mock_invalidate_budget_reservation_counters.assert_awaited_once_with( + budget_reservation=budget_reservation, + ) + assert budget_reservation["finalized"] is True + + proxy_logging_obj.db_spend_update_writer.update_database.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_update_database_and_spend_counters_preserves_counter_exception_when_invalidation_fails(): + proxy_logging_obj = MagicMock() + proxy_logging_obj.db_spend_update_writer.update_database = AsyncMock() + counter_exception = RuntimeError("counter unavailable") + increment_spend_counters = AsyncMock(side_effect=counter_exception) + budget_reservation = { + "reserved_cost": 0.5, + "entries": [{"counter_key": "spend:key:test_api_key"}], + } + + with ( + patch( + "litellm.proxy.spend_tracking.budget_reservation.invalidate_budget_reservation_counters", + new_callable=AsyncMock, + side_effect=RuntimeError("invalidate unavailable"), + ) as mock_invalidate_budget_reservation_counters, + patch( + "litellm.proxy.hooks.proxy_track_cost_callback.verbose_proxy_logger.exception", + ) as mock_log_exception, + ): + with pytest.raises(RuntimeError) as exc_info: + await _update_database_and_spend_counters( + proxy_logging_obj=proxy_logging_obj, + increment_spend_counters=increment_spend_counters, + user_api_key="test_api_key", + user_id="test_user_id", + end_user_id=None, + team_id="test_team_id", + org_id="test_org_id", + kwargs={}, + completion_response=None, + start_time=datetime.now(), + end_time=datetime.now(), + response_cost=0.2, + budget_reservation=budget_reservation, + ) + + assert exc_info.value is counter_exception + mock_invalidate_budget_reservation_counters.assert_awaited_once_with( + budget_reservation=budget_reservation, + ) + mock_log_exception.assert_called_once_with( + "Failed to invalidate budget reservation counters after spend counter update failed" + ) + assert budget_reservation["finalized"] is True + + proxy_logging_obj.db_spend_update_writer.update_database.assert_awaited_once() + + @pytest.mark.asyncio async def test_track_cost_callback_skips_when_no_standard_logging_object(): """ @@ -344,7 +778,7 @@ async def test_enrich_failure_metadata_skips_when_no_api_key(): "user_api_key_team_id": None, "user_api_key_team_alias": None, } - result = await _ProxyDBLogger._enrich_failure_metadata_with_key_info(metadata) + await _ProxyDBLogger._enrich_failure_metadata_with_key_info(metadata) mock_get_key.assert_not_called() @@ -556,3 +990,80 @@ async def test_async_post_call_failure_hook_uses_actual_start_time(): # Duration should be approximately 60 seconds, not 0 duration = (call_args["end_time"] - call_args["start_time"]).total_seconds() assert duration >= 55, f"Duration should be ~60s, got {duration}s" + + +async def _invoke_failure_hook_with_raised_exception(): + """Run the failure hook with an exception that has a real ``__traceback__``. + + Returns the metadata dict that was forwarded to ``update_database`` so the + caller can assert on its ``error_information`` payload. + """ + logger = _ProxyDBLogger() + user_api_key_dict = UserAPIKeyAuth( + api_key="test_api_key", + user_id="u", + team_id="t", + ) + request_data = { + "model": "gpt-4", + "messages": [{"role": "user", "content": "hi"}], + "metadata": {}, + "proxy_server_request": {}, + } + + try: + raise RuntimeError("boom-with-traceback") + except RuntimeError as exc: + original_exception = exc + + with patch( + "litellm.proxy.db.db_spend_update_writer.DBSpendUpdateWriter.update_database", + new_callable=AsyncMock, + ) as mock_update_database: + await logger.async_post_call_failure_hook( + request_data=request_data, + original_exception=original_exception, + user_api_key_dict=user_api_key_dict, + ) + call_args = mock_update_database.call_args[1] + return call_args["kwargs"]["litellm_params"]["metadata"] + + +@pytest.mark.asyncio +async def test_failure_hook_keeps_error_information_traceback_by_default(monkeypatch): + """Without the opt-in env var, the SpendLogs row carries the full traceback.""" + monkeypatch.delenv("LITELLM_SUPPRESS_SPEND_LOG_TRACEBACKS", raising=False) + + metadata = await _invoke_failure_hook_with_raised_exception() + + error_information = metadata["error_information"] + assert error_information["error_class"] == "RuntimeError" + assert error_information["error_message"] == "boom-with-traceback" + assert error_information["traceback"], "expected a non-empty traceback by default" + + +@pytest.mark.asyncio +async def test_failure_hook_drops_error_information_traceback_when_env_set( + monkeypatch, +): + """With the opt-in env var, the traceback key is omitted from the + SpendLogs row entirely so the per-row Metadata pane in the UI (which + renders ``error_information`` as a JSON blob) doesn't show a noisy empty + ``"traceback": ""`` line. The other fields (error_class / error_message / + error_code) are preserved.""" + import logging + + from litellm._logging import verbose_proxy_logger + + monkeypatch.setenv("LITELLM_SUPPRESS_SPEND_LOG_TRACEBACKS", "true") + original_level = verbose_proxy_logger.level + verbose_proxy_logger.setLevel(logging.INFO) + try: + metadata = await _invoke_failure_hook_with_raised_exception() + finally: + verbose_proxy_logger.setLevel(original_level) + + error_information = metadata["error_information"] + assert "traceback" not in error_information + assert error_information["error_class"] == "RuntimeError" + assert error_information["error_message"] == "boom-with-traceback" diff --git a/tests/test_litellm/proxy/management_endpoints/test_common_utils.py b/tests/test_litellm/proxy/management_endpoints/test_common_utils.py index a1e7fe59cab..f898763d2cb 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_common_utils.py +++ b/tests/test_litellm/proxy/management_endpoints/test_common_utils.py @@ -481,3 +481,39 @@ class TestSetObjectMetadataField: ): _set_object_metadata_field(team, "model_rpm_limit", {"x": 1}) assert team.metadata == {"model_rpm_limit": {"x": 1}} + + +class TestRequireCallerUserIdForNonAdmin: + """ + Security regression: service-account keys (user_id=None) must not bypass + the non-admin scoping branch on analytics endpoints. + """ + + def test_returns_user_id_when_present(self): + from litellm.proxy.management_endpoints.common_utils import ( + require_caller_user_id_for_non_admin, + ) + + key_dict = UserAPIKeyAuth( + user_id="user-abc", + user_role=LitellmUserRoles.INTERNAL_USER, + ) + assert require_caller_user_id_for_non_admin(key_dict) == "user-abc" + + def test_raises_403_when_user_id_is_none(self): + from fastapi import HTTPException + + from litellm.proxy.management_endpoints.common_utils import ( + require_caller_user_id_for_non_admin, + ) + + # Simulates a service-account key (user_id forced to None at key creation) + service_account_key = UserAPIKeyAuth( + user_id=None, + user_role=LitellmUserRoles.INTERNAL_USER, + ) + with pytest.raises(HTTPException) as exc_info: + require_caller_user_id_for_non_admin(service_account_key) + + assert exc_info.value.status_code == 403 + assert "Service-account keys" in str(exc_info.value.detail) diff --git a/tests/test_litellm/proxy/management_endpoints/test_internal_user_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_internal_user_endpoints.py index a1ba7ecd677..f4dc85dad99 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_internal_user_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_internal_user_endpoints.py @@ -1732,6 +1732,107 @@ async def test_get_user_daily_activity_non_admin_cannot_view_other_users(monkeyp assert call_kwargs.kwargs["entity_id"] == "regular-user-123" +@pytest.mark.asyncio +async def test_get_user_daily_activity_rejects_service_account_caller(monkeypatch): + """ + Security regression: a non-admin caller with user_id=None (a service-account + key, where user_id is forced to None at key creation) must not be able to + bypass the entity filter and read every tenant's daily spend. + + Before the fix, the endpoint silently defaulted user_id to + user_api_key_dict.user_id, which is itself None for service-account keys. + None != None is False, the same-user check passed, and entity_id=None + flowed into get_daily_activity, where the SQL builder treats None as + "no filter". + """ + from unittest.mock import AsyncMock, MagicMock + + from fastapi import HTTPException + + from litellm.proxy.management_endpoints.internal_user_endpoints import ( + get_user_daily_activity, + ) + + mock_prisma_client = MagicMock() + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) + + # Tripwire: ensure get_daily_activity is never reached + mock_get_daily = AsyncMock() + monkeypatch.setattr( + "litellm.proxy.management_endpoints.internal_user_endpoints.get_daily_activity", + mock_get_daily, + ) + + service_account_key = UserAPIKeyAuth( + user_id=None, # service-account keys have user_id forced to None + user_role=LitellmUserRoles.INTERNAL_USER, + ) + + with pytest.raises(HTTPException) as exc_info: + await get_user_daily_activity( + start_date="2025-01-01", + end_date="2025-01-31", + model=None, + api_key=None, + user_id=None, + page=1, + page_size=50, + timezone=None, + user_api_key_dict=service_account_key, + ) + + assert exc_info.value.status_code == 403 + assert "Service-account keys" in str(exc_info.value.detail) + mock_get_daily.assert_not_called() + + +@pytest.mark.asyncio +async def test_get_user_daily_activity_aggregated_rejects_service_account_caller( + monkeypatch, +): + """ + Same security regression as + test_get_user_daily_activity_rejects_service_account_caller, on the + aggregated route. Same shape, raw-SQL builder, same fix. + """ + from unittest.mock import AsyncMock, MagicMock + + from fastapi import HTTPException + + from litellm.proxy.management_endpoints.internal_user_endpoints import ( + get_user_daily_activity_aggregated, + ) + + mock_prisma_client = MagicMock() + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) + + mock_get_daily_agg = AsyncMock() + monkeypatch.setattr( + "litellm.proxy.management_endpoints.internal_user_endpoints.get_daily_activity_aggregated", + mock_get_daily_agg, + ) + + service_account_key = UserAPIKeyAuth( + user_id=None, + user_role=LitellmUserRoles.INTERNAL_USER, + ) + + with pytest.raises(HTTPException) as exc_info: + await get_user_daily_activity_aggregated( + start_date="2025-01-01", + end_date="2025-01-31", + model=None, + api_key=None, + user_id=None, + timezone=None, + user_api_key_dict=service_account_key, + ) + + assert exc_info.value.status_code == 403 + assert "Service-account keys" in str(exc_info.value.detail) + mock_get_daily_agg.assert_not_called() + + @pytest.mark.asyncio async def test_get_user_daily_activity_aggregated_admin_global_view(monkeypatch): """ @@ -1949,11 +2050,13 @@ async def test_delete_user_rejects_org_admin_deleting_outside_scope(mocker): 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 + 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 @@ -2631,3 +2734,107 @@ class TestGetUserIdFromRequestValidation: request = self._make_request(f"user_id={exact_id}") result = get_user_id_from_request(request) assert result == exact_id + + +# --------------------------------------------------------------------------- +# VERIA-60: /user/info post-decode re-authorization +# --------------------------------------------------------------------------- + + +def test_enforce_user_info_access_admin_bypass(): + """Proxy admins must always be allowed past the re-check.""" + from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth + from litellm.proxy.management_endpoints.internal_user_endpoints import ( + _enforce_user_info_access, + ) + + admin = UserAPIKeyAuth( + user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN.value + ) + # Should not raise even when querying a different user + _enforce_user_info_access(user_id="someone_else", user_api_key_dict=admin) + + +def test_enforce_user_info_access_view_only_admin_blocked_from_other_users(): + """PROXY_ADMIN_VIEW_ONLY is not a true admin for /user/info — the upstream + route check applies the same `user_id == valid_token.user_id` rule, so the + re-check here must mirror that and deny cross-user lookups.""" + import pytest + from fastapi import HTTPException + + from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth + from litellm.proxy.management_endpoints.internal_user_endpoints import ( + _enforce_user_info_access, + ) + + viewer = UserAPIKeyAuth( + user_id="viewer", + user_role=LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY.value, + ) + with pytest.raises(HTTPException) as exc_info: + _enforce_user_info_access(user_id="someone_else", user_api_key_dict=viewer) + assert exc_info.value.status_code == 403 + + +def test_enforce_user_info_access_view_only_admin_can_read_own(): + from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth + from litellm.proxy.management_endpoints.internal_user_endpoints import ( + _enforce_user_info_access, + ) + + viewer = UserAPIKeyAuth( + user_id="viewer", + user_role=LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY.value, + ) + _enforce_user_info_access(user_id="viewer", user_api_key_dict=viewer) + + +def test_enforce_user_info_access_owner_allowed(): + from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth + from litellm.proxy.management_endpoints.internal_user_endpoints import ( + _enforce_user_info_access, + ) + + user = UserAPIKeyAuth( + user_id="alice", user_role=LitellmUserRoles.INTERNAL_USER.value + ) + _enforce_user_info_access(user_id="alice", user_api_key_dict=user) + + +def test_enforce_user_info_access_no_user_id_allowed(): + """No user_id in query → handler resolves to caller's own id later, so + this branch must not raise.""" + from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth + from litellm.proxy.management_endpoints.internal_user_endpoints import ( + _enforce_user_info_access, + ) + + user = UserAPIKeyAuth( + user_id="alice", user_role=LitellmUserRoles.INTERNAL_USER.value + ) + _enforce_user_info_access(user_id=None, user_api_key_dict=user) + + +def test_enforce_user_info_access_blocks_cross_user_lookup(): + """A non-admin caller may not query another user's row, even if URL + re-parsing produced a user_id that differs from the one the route check + saw (the VERIA-60 bypass).""" + import pytest + from fastapi import HTTPException + + from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth + from litellm.proxy.management_endpoints.internal_user_endpoints import ( + _enforce_user_info_access, + ) + + attacker = UserAPIKeyAuth( + user_id="attacker space", # original (URL-decoded) id seen by route check + user_role=LitellmUserRoles.INTERNAL_USER.value, + ) + + with pytest.raises(HTTPException) as exc_info: + # Re-parsed id (with literal '+') belongs to the victim + _enforce_user_info_access(user_id="victim+target", user_api_key_dict=attacker) + + assert exc_info.value.status_code == 403 + assert "key not allowed to access this user's info" in str(exc_info.value.detail) diff --git a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py index f256ffd8661..b292e8d0cae 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py @@ -5021,6 +5021,299 @@ async def test_list_keys_non_admin_user_id_auto_set(): ) +def _make_member_team_table( + team_id: str, + member_user_id: str, + member_role: str = "user", + team_member_permissions=None, +): + """Build a LiteLLM_TeamTable with a single member, suitable for list_keys tests.""" + from litellm.proxy._types import LiteLLM_TeamTable, Member + + return LiteLLM_TeamTable( + team_id=team_id, + members_with_roles=[Member(user_id=member_user_id, role=member_role)], + team_member_permissions=team_member_permissions, + ) + + +async def _invoke_list_keys_and_capture_helper_kwargs( + user_api_key_dict, + team_objects, + *, + include_team_keys: bool = True, +): + """ + Invoke list_keys with mocked dependencies and return the kwargs that + list_keys passes to _list_key_helper (so tests can assert on + admin_team_ids / member_team_ids classification). + """ + from unittest.mock import Mock, patch + + from litellm.proxy._types import LiteLLM_UserTable + + mock_prisma_client = AsyncMock() + mock_user_info = LiteLLM_UserTable( + user_id=user_api_key_dict.user_id, + user_email="member@example.com", + teams=[t.team_id for t in team_objects], + organization_memberships=[], + ) + mock_list_key_helper = AsyncMock( + return_value={ + "keys": [], + "total_count": 0, + "current_page": 1, + "total_pages": 0, + } + ) + with ( + patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client), + patch( + "litellm.proxy.management_endpoints.key_management_endpoints.validate_key_list_check", + return_value=mock_user_info, + ), + patch( + "litellm.proxy.management_endpoints.key_management_endpoints._fetch_user_team_objects", + AsyncMock(return_value=team_objects), + ), + patch( + "litellm.proxy.management_endpoints.key_management_endpoints._list_key_helper", + mock_list_key_helper, + ), + ): + await list_keys( + request=Mock(), + user_api_key_dict=user_api_key_dict, + include_team_keys=include_team_keys, + status=None, + ) + mock_list_key_helper.assert_called_once() + return mock_list_key_helper.call_args.kwargs + + +@pytest.mark.asyncio +async def test_list_keys_team_member_with_key_list_permission_sees_all_team_keys(): + """ + Bug fix: when a team has /key/list in team_member_permissions, regular + team members must get full key visibility for that team — same as a team + admin would. This means other members' keys AND service account keys + (user_id=NULL) must be returned, not only the caller's own keys. + + This test pins down list_keys' classification: the team must be passed + to _list_key_helper as a full-visibility team (admin_team_ids), not as + a service-account-only team (member_team_ids). + """ + member_user_id = "member-user-1" + team_id = "team-with-permission" + user_api_key_dict = UserAPIKeyAuth( + user_role=LitellmUserRoles.INTERNAL_USER, + user_id=member_user_id, + ) + team_objects = [ + _make_member_team_table( + team_id=team_id, + member_user_id=member_user_id, + member_role="user", + team_member_permissions=["/key/list"], + ) + ] + + helper_kwargs = await _invoke_list_keys_and_capture_helper_kwargs( + user_api_key_dict=user_api_key_dict, + team_objects=team_objects, + ) + + admin_team_ids = helper_kwargs.get("admin_team_ids") or [] + member_team_ids = helper_kwargs.get("member_team_ids") or [] + assert team_id in admin_team_ids, ( + "team granting /key/list permission must be classified as full-visibility " + "(admin_team_ids), but got " + f"admin_team_ids={admin_team_ids}, member_team_ids={member_team_ids}" + ) + + +@pytest.mark.asyncio +async def test_list_keys_team_member_without_key_list_permission_only_service_accounts(): + """ + Without the /key/list permission, the existing scoping must hold: the + team is classified as member-only, so only service account keys + (user_id=NULL) for that team are visible to the caller. + """ + member_user_id = "member-user-2" + team_id = "team-no-permission" + user_api_key_dict = UserAPIKeyAuth( + user_role=LitellmUserRoles.INTERNAL_USER, + user_id=member_user_id, + ) + team_objects = [ + _make_member_team_table( + team_id=team_id, + member_user_id=member_user_id, + member_role="user", + team_member_permissions=None, + ) + ] + + helper_kwargs = await _invoke_list_keys_and_capture_helper_kwargs( + user_api_key_dict=user_api_key_dict, + team_objects=team_objects, + ) + + admin_team_ids = helper_kwargs.get("admin_team_ids") or [] + member_team_ids = helper_kwargs.get("member_team_ids") or [] + assert team_id not in admin_team_ids + assert team_id in member_team_ids + + +@pytest.mark.asyncio +async def test_list_keys_team_member_with_permission_in_one_team_only(): + """ + Granular: a user is a member of two teams. Only one team grants + /key/list — the other does not. The classification must respect the + per-team permission, not leak full visibility across both teams. + """ + member_user_id = "member-user-3" + team_with_permission = "team-A" + team_without_permission = "team-B" + user_api_key_dict = UserAPIKeyAuth( + user_role=LitellmUserRoles.INTERNAL_USER, + user_id=member_user_id, + ) + team_objects = [ + _make_member_team_table( + team_id=team_with_permission, + member_user_id=member_user_id, + member_role="user", + team_member_permissions=["/key/list"], + ), + _make_member_team_table( + team_id=team_without_permission, + member_user_id=member_user_id, + member_role="user", + team_member_permissions=[], + ), + ] + + helper_kwargs = await _invoke_list_keys_and_capture_helper_kwargs( + user_api_key_dict=user_api_key_dict, + team_objects=team_objects, + ) + + admin_team_ids = helper_kwargs.get("admin_team_ids") or [] + member_team_ids = helper_kwargs.get("member_team_ids") or [] + assert team_with_permission in admin_team_ids + assert team_without_permission not in admin_team_ids + assert team_without_permission in member_team_ids + + +@pytest.mark.asyncio +async def test_list_keys_team_admin_unaffected_by_member_permission_logic(): + """ + Sanity: a team admin's classification is unchanged by the new + permission-aware path. They still appear in admin_team_ids (full + visibility) regardless of the team_member_permissions value. + """ + admin_user_id = "team-admin-user" + team_id = "team-with-admin" + user_api_key_dict = UserAPIKeyAuth( + user_role=LitellmUserRoles.INTERNAL_USER, + user_id=admin_user_id, + ) + team_objects = [ + _make_member_team_table( + team_id=team_id, + member_user_id=admin_user_id, + member_role="admin", + team_member_permissions=None, + ) + ] + + helper_kwargs = await _invoke_list_keys_and_capture_helper_kwargs( + user_api_key_dict=user_api_key_dict, + team_objects=team_objects, + ) + + admin_team_ids = helper_kwargs.get("admin_team_ids") or [] + assert team_id in admin_team_ids + + +def test_build_key_filter_conditions_full_visibility_team_includes_service_accounts(): + """ + Direct check on the SQL filter: when a team is in the full-visibility + set (admin_team_ids), the filter clause for that team is + {"team_id": {"in": [...]}} with NO user_id constraint — so service + account keys (user_id=NULL) AND other members' keys are returned. + + This is the SQL-level proof that a member with /key/list permission + will see service account keys for the team. + """ + from litellm.proxy.management_endpoints.key_management_endpoints import ( + _build_key_filter_conditions, + ) + + full_visibility_team = "team-full-vis" + where = _build_key_filter_conditions( + user_id="member-user-x", + team_id=None, + organization_id=None, + key_alias=None, + key_hash=None, + exclude_team_id=None, + admin_team_ids=[full_visibility_team], + member_team_ids=[full_visibility_team], + include_created_by_keys=False, + ) + + serialized = json.dumps(where) + # Full-visibility team filter: no user_id restriction + assert ( + json.dumps({"team_id": {"in": [full_visibility_team]}}) in serialized + ), f"expected unrestricted team_id IN clause, got: {serialized}" + # No service-account-only AND clause for this team (it would be redundant + # and would erroneously narrow the visibility back to user_id=NULL). + sa_only_clause = json.dumps( + {"AND": [{"team_id": {"in": [full_visibility_team]}}, {"user_id": None}]} + ) + assert ( + sa_only_clause not in serialized + ), f"team in admin_team_ids must not also be filtered to user_id=NULL: {serialized}" + + +def test_build_key_filter_conditions_member_only_team_restricts_to_service_accounts(): + """ + Existing-behavior pin: when a team is ONLY in member_team_ids (not + admin_team_ids), the filter for that team must be + {"AND": [{"team_id": {"in": [...]}}, {"user_id": None}]} — i.e. only + service accounts visible. This is the "no permission" baseline that + must keep working. + """ + from litellm.proxy.management_endpoints.key_management_endpoints import ( + _build_key_filter_conditions, + ) + + member_only_team = "team-member-only" + where = _build_key_filter_conditions( + user_id="member-user-y", + team_id=None, + organization_id=None, + key_alias=None, + key_hash=None, + exclude_team_id=None, + admin_team_ids=[], + member_team_ids=[member_only_team], + include_created_by_keys=False, + ) + + serialized = json.dumps(where) + expected = json.dumps( + {"AND": [{"team_id": {"in": [member_only_team]}}, {"user_id": None}]} + ) + assert ( + expected in serialized + ), f"member-only team must be restricted to user_id=NULL keys, got: {serialized}" + + @pytest.mark.asyncio async def test_generate_key_negative_max_budget(): """ diff --git a/tests/test_litellm/proxy/management_endpoints/test_project_org_authz.py b/tests/test_litellm/proxy/management_endpoints/test_project_org_authz.py new file mode 100644 index 00000000000..bd982480d60 --- /dev/null +++ b/tests/test_litellm/proxy/management_endpoints/test_project_org_authz.py @@ -0,0 +1,195 @@ +""" +Unit tests for the VERIA-55 fixes: + +- Project update permission must be evaluated against the project's *current* + team, not a team supplied in the request body. +- Key update may not assign a key to an organization the caller is not a + member of. +""" + +from unittest.mock import AsyncMock, MagicMock + +import pytest +from fastapi import HTTPException + +from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth + + +# --------------------------------------------------------------------------- +# /project/update — _check_user_permission_for_project +# --------------------------------------------------------------------------- + + +def _make_prisma_with_team(team_id: str, admins: list): + prisma = MagicMock() + team_row = MagicMock() + team_row.team_id = team_id + team_row.admins = admins + prisma.db.litellm_teamtable.find_unique = AsyncMock(return_value=team_row) + return prisma + + +@pytest.mark.asyncio +async def test_project_perm_check_uses_current_team_not_caller_supplied(): + """The permission check must look at the project's existing team. Even + if the caller is admin of an unrelated team, they must not pass when no + explicit team_object is forced through.""" + from enterprise.litellm_enterprise.proxy.management_endpoints.project_endpoints import ( + _check_user_permission_for_project, + ) + + # Project lives on team-A, caller is admin only of team-B. + prisma = _make_prisma_with_team(team_id="team-A", admins=["alice"]) + caller = UserAPIKeyAuth( + user_id="bob", + user_role=LitellmUserRoles.INTERNAL_USER.value, + ) + + has_perm = await _check_user_permission_for_project( + user_api_key_dict=caller, + team_id="team-A", + prisma_client=prisma, + ) + assert has_perm is False + prisma.db.litellm_teamtable.find_unique.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_project_perm_check_allows_team_admin_of_existing_team(): + from enterprise.litellm_enterprise.proxy.management_endpoints.project_endpoints import ( + _check_user_permission_for_project, + ) + + prisma = _make_prisma_with_team(team_id="team-A", admins=["alice"]) + alice = UserAPIKeyAuth( + user_id="alice", + user_role=LitellmUserRoles.INTERNAL_USER.value, + ) + + has_perm = await _check_user_permission_for_project( + user_api_key_dict=alice, + team_id="team-A", + prisma_client=prisma, + ) + assert has_perm is True + + +@pytest.mark.asyncio +async def test_project_perm_check_proxy_admin_always_allowed(): + from enterprise.litellm_enterprise.proxy.management_endpoints.project_endpoints import ( + _check_user_permission_for_project, + ) + + prisma = MagicMock() + admin = UserAPIKeyAuth( + user_id="root", + user_role=LitellmUserRoles.PROXY_ADMIN.value, + ) + + has_perm = await _check_user_permission_for_project( + user_api_key_dict=admin, + team_id="team-A", + prisma_client=prisma, + ) + assert has_perm is True + # Admin shortcut should not even hit the DB. + prisma.db.litellm_teamtable.find_unique.assert_not_called() + + +# --------------------------------------------------------------------------- +# /key/update — _validate_caller_can_assign_key_org +# --------------------------------------------------------------------------- + + +def _make_prisma_with_user_orgs(user_id: str, org_ids: list): + prisma = MagicMock() + user_row = MagicMock() + user_row.organization_memberships = [ + MagicMock(organization_id=org_id) for org_id in org_ids + ] + prisma.db.litellm_usertable.find_unique = AsyncMock(return_value=user_row) + return prisma + + +@pytest.mark.asyncio +async def test_assign_key_org_allows_member(): + from litellm.proxy.management_endpoints.key_management_endpoints import ( + _validate_caller_can_assign_key_org, + ) + + prisma = _make_prisma_with_user_orgs("alice", ["org-1", "org-2"]) + caller = UserAPIKeyAuth( + user_id="alice", + user_role=LitellmUserRoles.INTERNAL_USER.value, + ) + # Should not raise. + await _validate_caller_can_assign_key_org( + user_api_key_dict=caller, + organization_id="org-2", + prisma_client=prisma, + ) + + +@pytest.mark.asyncio +async def test_assign_key_org_blocks_non_member(): + """The IDOR: caller asks to point a key at an org they don't belong to.""" + from litellm.proxy.management_endpoints.key_management_endpoints import ( + _validate_caller_can_assign_key_org, + ) + + prisma = _make_prisma_with_user_orgs("alice", ["org-1"]) + caller = UserAPIKeyAuth( + user_id="alice", + user_role=LitellmUserRoles.INTERNAL_USER.value, + ) + with pytest.raises(HTTPException) as exc_info: + await _validate_caller_can_assign_key_org( + user_api_key_dict=caller, + organization_id="someone-elses-org", + prisma_client=prisma, + ) + assert exc_info.value.status_code == 403 + assert "someone-elses-org" in str(exc_info.value.detail) + + +@pytest.mark.asyncio +async def test_assign_key_org_blocks_caller_without_user_id(): + from litellm.proxy.management_endpoints.key_management_endpoints import ( + _validate_caller_can_assign_key_org, + ) + + prisma = MagicMock() + caller = UserAPIKeyAuth( + user_role=LitellmUserRoles.INTERNAL_USER.value, + ) + with pytest.raises(HTTPException) as exc_info: + await _validate_caller_can_assign_key_org( + user_api_key_dict=caller, + organization_id="org-1", + prisma_client=prisma, + ) + assert exc_info.value.status_code == 403 + + +@pytest.mark.asyncio +async def test_assign_key_org_blocks_caller_with_no_memberships(): + from litellm.proxy.management_endpoints.key_management_endpoints import ( + _validate_caller_can_assign_key_org, + ) + + prisma = MagicMock() + user_row = MagicMock() + user_row.organization_memberships = None + prisma.db.litellm_usertable.find_unique = AsyncMock(return_value=user_row) + + caller = UserAPIKeyAuth( + user_id="alice", + user_role=LitellmUserRoles.INTERNAL_USER.value, + ) + with pytest.raises(HTTPException) as exc_info: + await _validate_caller_can_assign_key_org( + user_api_key_dict=caller, + organization_id="org-1", + prisma_client=prisma, + ) + assert exc_info.value.status_code == 403 diff --git a/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py b/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py index a0ae95df589..69798744f7f 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py +++ b/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py @@ -1841,6 +1841,7 @@ class TestCustomUISSO: "x-forwarded-for": "192.168.1.1", } mock_request.base_url = "https://test.litellm.ai/" + mock_request.client.host = "10.0.0.10" # Mock the custom handler mock_custom_handler = MagicMock(spec=CustomSSOLoginHandler) @@ -1866,36 +1867,73 @@ class TestCustomUISSO: "litellm.proxy.proxy_server.user_custom_ui_sso_sign_in_handler", mock_custom_handler, ): - with patch.object( - SSOAuthenticationHandler, - "get_redirect_response_from_openid", - return_value=mock_redirect_response, - ) as mock_get_redirect: - # Act - result = ( + with patch( + "litellm.proxy.proxy_server.general_settings", + {"trusted_proxy_ranges": ["10.0.0.0/24"]}, + ): + with patch.object( + SSOAuthenticationHandler, + "get_redirect_response_from_openid", + return_value=mock_redirect_response, + ) as mock_get_redirect: + # Act + result = await EnterpriseCustomSSOHandler.handle_custom_ui_sso_sign_in( + request=mock_request + ) + + # Assert + # Verify the custom handler was called with the request + mock_custom_handler.handle_custom_ui_sso_sign_in.assert_called_once_with( + request=mock_request + ) + + # Verify the redirect response was generated with correct OpenID + mock_get_redirect.assert_called_once_with( + result=expected_openid, + request=mock_request, + received_response=None, + generic_client_id=None, + ui_access_mode=None, + ) + + # Verify the result is the redirect response + assert result == mock_redirect_response + assert result.status_code == 303 + + @pytest.mark.asyncio + async def test_handle_custom_ui_sso_sign_in_rejects_untrusted_proxy(self): + """Custom UI SSO rejects spoofed identity headers from direct clients.""" + from enterprise.litellm_enterprise.proxy.auth.custom_sso_handler import ( + EnterpriseCustomSSOHandler, + ) + from litellm.integrations.custom_sso_handler import CustomSSOLoginHandler + + mock_request = MagicMock(spec=Request) + mock_request.headers = { + "x-litellm-user-id": "admin", + "x-litellm-user-email": "admin@example.com", + } + mock_request.base_url = "https://test.litellm.ai/" + mock_request.client.host = "203.0.113.10" + + mock_custom_handler = MagicMock(spec=CustomSSOLoginHandler) + mock_custom_handler.handle_custom_ui_sso_sign_in = AsyncMock() + + with patch("litellm.proxy.proxy_server.premium_user", True): + with patch( + "litellm.proxy.proxy_server.user_custom_ui_sso_sign_in_handler", + mock_custom_handler, + ): + with patch( + "litellm.proxy.proxy_server.general_settings", + {"trusted_proxy_ranges": ["10.0.0.0/24"]}, + ): + with pytest.raises(ValueError, match="not trusted"): await EnterpriseCustomSSOHandler.handle_custom_ui_sso_sign_in( request=mock_request ) - ) - # Assert - # Verify the custom handler was called with the request - mock_custom_handler.handle_custom_ui_sso_sign_in.assert_called_once_with( - request=mock_request - ) - - # Verify the redirect response was generated with correct OpenID - mock_get_redirect.assert_called_once_with( - result=expected_openid, - request=mock_request, - received_response=None, - generic_client_id=None, - ui_access_mode=None, - ) - - # Verify the result is the redirect response - assert result == mock_redirect_response - assert result.status_code == 303 + mock_custom_handler.handle_custom_ui_sso_sign_in.assert_not_called() @pytest.mark.asyncio async def test_custom_ui_sso_handler_execution_with_real_class(self): @@ -1946,6 +1984,7 @@ class TestCustomUISSO: "x-forwarded-for": "10.0.0.1", } mock_request.base_url = "https://custom.litellm.ai/" + mock_request.client.host = "10.0.0.20" # Mock the redirect response method mock_redirect_response = MagicMock() @@ -1956,34 +1995,36 @@ class TestCustomUISSO: "litellm.proxy.proxy_server.user_custom_ui_sso_sign_in_handler", test_handler_instance, ): - with patch.object( - SSOAuthenticationHandler, - "get_redirect_response_from_openid", - return_value=mock_redirect_response, - ) as mock_get_redirect: - # Act - result = ( - await EnterpriseCustomSSOHandler.handle_custom_ui_sso_sign_in( + with patch( + "litellm.proxy.proxy_server.general_settings", + {"trusted_proxy_ranges": ["10.0.0.0/24"]}, + ): + with patch.object( + SSOAuthenticationHandler, + "get_redirect_response_from_openid", + return_value=mock_redirect_response, + ) as mock_get_redirect: + # Act + result = await EnterpriseCustomSSOHandler.handle_custom_ui_sso_sign_in( request=mock_request ) - ) - # Assert that our custom handler was executed - assert test_handler_instance.method_called is True - assert test_handler_instance.received_request == mock_request + # Assert that our custom handler was executed + assert test_handler_instance.method_called is True + assert test_handler_instance.received_request == mock_request - # Verify the redirect response was called with the OpenID from our custom handler - mock_get_redirect.assert_called_once() - call_args = mock_get_redirect.call_args.kwargs + # Verify the redirect response was called with the OpenID from our custom handler + mock_get_redirect.assert_called_once() + call_args = mock_get_redirect.call_args.kwargs - # Verify the OpenID object has the expected values from our custom handler - openid_result = call_args["result"] - assert openid_result.id == "custom_test_user_456" - assert openid_result.email == "custom@example.com" - assert openid_result.first_name == "Custom" - assert openid_result.last_name == "Handler" - assert openid_result.display_name == "Custom Handler Test" - assert openid_result.provider == "custom" + # Verify the OpenID object has the expected values from our custom handler + openid_result = call_args["result"] + assert openid_result.id == "custom_test_user_456" + assert openid_result.email == "custom@example.com" + assert openid_result.first_name == "Custom" + assert openid_result.last_name == "Handler" + assert openid_result.display_name == "Custom Handler Test" + assert openid_result.provider == "custom" # Verify the request and other parameters were passed correctly assert call_args["request"] == mock_request @@ -5767,3 +5808,324 @@ class TestSyncUserRoleFromJwtRoleMap: ) prisma.db.litellm_usertable.update.assert_not_called() + + +# ── VERIA-34 regression: PKCE state-to-session-cookie binding ─────────────── + + +class TestPKCEStateCookieBinding: + """The Generic SSO PKCE flow used the URL ``state`` parameter as a + cache-key for the PKCE ``code_verifier`` without binding the state to + the caller's browser. An attacker who pre-mints a state + cached + verifier could hand the link to a victim and capture the resulting + access token. Fix: set ``litellm_oauth_state`` HttpOnly cookie on + the redirect; verify the URL state matches the cookie before doing + the PKCE token exchange.""" + + @pytest.mark.asyncio + async def test_redirect_response_sets_oauth_state_cookie_when_pkce_enabled(self): + """``get_generic_sso_redirect_response`` must set + ``litellm_oauth_state`` on the redirect response when PKCE is on so + the callback can verify it later. The cookie must carry HttpOnly, + SameSite=Lax, and (because no http request was supplied to the + helper) the production-safe ``Secure`` default.""" + from fastapi.responses import RedirectResponse + + from litellm.proxy.management_endpoints.ui_sso import ( + SSOAuthenticationHandler, + ) + + mock_redirect = RedirectResponse( + url="https://idp.example.com/authorize?state=test-state-xyz" + ) + mock_generic_sso = MagicMock() + mock_generic_sso.__enter__ = MagicMock(return_value=mock_generic_sso) + mock_generic_sso.__exit__ = MagicMock(return_value=None) + mock_generic_sso.get_login_redirect = AsyncMock(return_value=mock_redirect) + + with patch.dict( + os.environ, + { + "GENERIC_CLIENT_STATE": "test-state-xyz", + "GENERIC_CLIENT_USE_PKCE": "true", + }, + ): + response = await SSOAuthenticationHandler.get_generic_sso_redirect_response( + generic_sso=mock_generic_sso, + state=None, + generic_authorization_endpoint="https://idp.example.com/authorize", + ) + + assert response is not None + cookie_headers = response.headers.getlist("set-cookie") + cookie_str = next( + (c for c in cookie_headers if "litellm_oauth_state=" in c), None + ) + assert ( + cookie_str is not None + ), f"litellm_oauth_state cookie not set; got: {cookie_headers}" + assert "test-state-xyz" in cookie_str + assert "HttpOnly" in cookie_str + assert "SameSite=lax" in cookie_str + # No incoming Request supplied → ``Secure`` defaults to True so a + # network observer on plain HTTP cannot read the state value. + assert "Secure" in cookie_str + + @pytest.mark.asyncio + async def test_redirect_response_omits_oauth_state_cookie_when_pkce_disabled( + self, + ): + """Non-PKCE flows delegate to fastapi-sso's own session-cookie + binding; we do not set our cookie there because it would never be + validated (and could collide with a concurrent PKCE session in + the same browser).""" + from fastapi.responses import RedirectResponse + + from litellm.proxy.management_endpoints.ui_sso import ( + SSOAuthenticationHandler, + ) + + mock_redirect = RedirectResponse( + url="https://idp.example.com/authorize?state=test-state-xyz" + ) + mock_generic_sso = MagicMock() + mock_generic_sso.__enter__ = MagicMock(return_value=mock_generic_sso) + mock_generic_sso.__exit__ = MagicMock(return_value=None) + mock_generic_sso.get_login_redirect = AsyncMock(return_value=mock_redirect) + + with patch.dict( + os.environ, + { + "GENERIC_CLIENT_STATE": "test-state-xyz", + "GENERIC_CLIENT_USE_PKCE": "false", + }, + ): + response = await SSOAuthenticationHandler.get_generic_sso_redirect_response( + generic_sso=mock_generic_sso, + state=None, + generic_authorization_endpoint="https://idp.example.com/authorize", + ) + + assert response is not None + cookie_headers = response.headers.getlist("set-cookie") + assert not any( + "litellm_oauth_state=" in c for c in cookie_headers + ), f"litellm_oauth_state cookie set on non-PKCE flow; got: {cookie_headers}" + + @pytest.mark.asyncio + async def test_redirect_response_drops_secure_flag_for_http_dev(self): + """When the incoming request is plain HTTP (local dev), ``Secure`` + must be dropped so the browser will actually attach the cookie on + the callback hop.""" + from fastapi.responses import RedirectResponse + + from litellm.proxy.management_endpoints.ui_sso import ( + SSOAuthenticationHandler, + ) + + mock_redirect = RedirectResponse( + url="http://idp.local/authorize?state=local-dev-state" + ) + mock_generic_sso = MagicMock() + mock_generic_sso.__enter__ = MagicMock(return_value=mock_generic_sso) + mock_generic_sso.__exit__ = MagicMock(return_value=None) + mock_generic_sso.get_login_redirect = AsyncMock(return_value=mock_redirect) + + http_request = MagicMock(spec=Request) + http_request.url.scheme = "http" + + with patch.dict( + os.environ, + { + "GENERIC_CLIENT_STATE": "local-dev-state", + "GENERIC_CLIENT_USE_PKCE": "true", + }, + ): + response = await SSOAuthenticationHandler.get_generic_sso_redirect_response( + generic_sso=mock_generic_sso, + state=None, + generic_authorization_endpoint="http://idp.local/authorize", + request=http_request, + ) + + cookie_headers = response.headers.getlist("set-cookie") + cookie_str = next( + (c for c in cookie_headers if "litellm_oauth_state=" in c), None + ) + assert cookie_str is not None + assert "Secure" not in cookie_str + + @pytest.mark.asyncio + async def test_pkce_callback_rejects_missing_cookie(self): + """When PKCE is enabled and a code_verifier is in the cache, the + callback must reject a request that has no ``litellm_oauth_state`` + cookie (browser-to-server binding missing).""" + from litellm.proxy._types import ProxyException + from litellm.proxy.management_endpoints.ui_sso import ( + SSOAuthenticationHandler, + get_generic_sso_response, + ) + + mock_request = MagicMock(spec=Request) + mock_request.query_params = { + "state": "attacker-minted-state", + "code": "auth-code", + } + # No oauth_state cookie set → request.cookies.get returns None. + mock_request.cookies = {} + + with ( + patch.object( + SSOAuthenticationHandler, + "prepare_token_exchange_parameters", + AsyncMock( + return_value={ + "code_verifier": "attacker-cached-verifier", + "_pkce_cache_key": "pkce_verifier:attacker-minted-state", + } + ), + ), + patch("fastapi_sso.sso.base.DiscoveryDocument"), + patch("fastapi_sso.sso.generic.create_provider", return_value=MagicMock()), + patch.dict( + os.environ, + { + "GENERIC_CLIENT_SECRET": "x", + "GENERIC_AUTHORIZATION_ENDPOINT": "https://idp.example.com/auth", + "GENERIC_TOKEN_ENDPOINT": "https://idp.example.com/token", + "GENERIC_USERINFO_ENDPOINT": "https://idp.example.com/userinfo", + "GENERIC_CLIENT_USE_PKCE": "true", + }, + ), + pytest.raises(ProxyException) as exc_info, + ): + await get_generic_sso_response( + request=mock_request, + jwt_handler=MagicMock(spec=JWTHandler), + generic_client_id="cid", + redirect_url="https://proxy.example.com/sso/callback", + sso_jwt_handler=None, + ) + + assert "state" in str(exc_info.value.message).lower() + + @pytest.mark.asyncio + async def test_pkce_callback_rejects_state_cookie_mismatch(self): + """The Login-CSRF shape: attacker mints state ``A``, victim's browser + carries cookie state ``B``. The callback must reject.""" + from litellm.proxy._types import ProxyException + from litellm.proxy.management_endpoints.ui_sso import ( + SSOAuthenticationHandler, + get_generic_sso_response, + ) + + mock_request = MagicMock(spec=Request) + mock_request.query_params = { + "state": "attacker-minted-state", + "code": "auth-code", + } + mock_request.cookies = {"litellm_oauth_state": "victim-browser-state"} + + with ( + patch.object( + SSOAuthenticationHandler, + "prepare_token_exchange_parameters", + AsyncMock( + return_value={ + "code_verifier": "verifier", + "_pkce_cache_key": "pkce_verifier:attacker-minted-state", + } + ), + ), + patch("fastapi_sso.sso.base.DiscoveryDocument"), + patch("fastapi_sso.sso.generic.create_provider", return_value=MagicMock()), + patch.dict( + os.environ, + { + "GENERIC_CLIENT_SECRET": "x", + "GENERIC_AUTHORIZATION_ENDPOINT": "https://idp.example.com/auth", + "GENERIC_TOKEN_ENDPOINT": "https://idp.example.com/token", + "GENERIC_USERINFO_ENDPOINT": "https://idp.example.com/userinfo", + "GENERIC_CLIENT_USE_PKCE": "true", + }, + ), + pytest.raises(ProxyException) as exc_info, + ): + await get_generic_sso_response( + request=mock_request, + jwt_handler=MagicMock(spec=JWTHandler), + generic_client_id="cid", + redirect_url="https://proxy.example.com/sso/callback", + sso_jwt_handler=None, + ) + + assert "state" in str(exc_info.value.message).lower() + + @pytest.mark.asyncio + async def test_pkce_callback_accepts_matching_state_cookie(self): + """Happy path: URL state and cookie state match (the legitimate + flow where the same browser that started the redirect lands on + the callback) → the PKCE token exchange proceeds.""" + from litellm.proxy.management_endpoints.ui_sso import ( + SSOAuthenticationHandler, + get_generic_sso_response, + ) + + mock_request = MagicMock(spec=Request) + mock_request.query_params = {"state": "matched-state", "code": "auth-code"} + mock_request.cookies = {"litellm_oauth_state": "matched-state"} + + with ( + patch.object( + SSOAuthenticationHandler, + "prepare_token_exchange_parameters", + AsyncMock( + return_value={ + "code_verifier": "verifier", + "_pkce_cache_key": "pkce_verifier:matched-state", + } + ), + ), + patch.object( + SSOAuthenticationHandler, + "_pkce_token_exchange", + AsyncMock( + return_value={ + "access_token": "tok", + "id_token": "id", + "sub": "user@example.com", + "email": "user@example.com", + } + ), + ), + patch.object( + SSOAuthenticationHandler, + "_delete_pkce_verifier", + AsyncMock(), + ), + patch("fastapi_sso.sso.base.DiscoveryDocument"), + patch("fastapi_sso.sso.generic.create_provider", return_value=MagicMock()), + patch.dict( + os.environ, + { + "GENERIC_CLIENT_SECRET": "x", + "GENERIC_AUTHORIZATION_ENDPOINT": "https://idp.example.com/auth", + "GENERIC_TOKEN_ENDPOINT": "https://idp.example.com/token", + "GENERIC_USERINFO_ENDPOINT": "https://idp.example.com/userinfo", + "GENERIC_CLIENT_USE_PKCE": "true", + }, + ), + ): + jwt_handler = MagicMock(spec=JWTHandler) + jwt_handler.get_team_ids_from_jwt.return_value = [] + result, _, _ = await get_generic_sso_response( + request=mock_request, + jwt_handler=jwt_handler, + generic_client_id="cid", + redirect_url="https://proxy.example.com/sso/callback", + sso_jwt_handler=None, + ) + + # State-cookie check passed, so the function got past the early + # ProxyException raise and produced an SSO result object. + assert result is not None diff --git a/tests/test_litellm/proxy/management_endpoints/usage_endpoints/test_ai_usage_chat.py b/tests/test_litellm/proxy/management_endpoints/usage_endpoints/test_ai_usage_chat.py index 66a18e2edb4..e8a74e41dae 100644 --- a/tests/test_litellm/proxy/management_endpoints/usage_endpoints/test_ai_usage_chat.py +++ b/tests/test_litellm/proxy/management_endpoints/usage_endpoints/test_ai_usage_chat.py @@ -409,3 +409,60 @@ class TestStreamUsageAiChat: end_date="2025-01-31", user_id="my-user-id", ) + + +class TestUsageAiChatServiceAccountGuard: + """ + Security regression: a non-admin caller with user_id=None (service-account + key) must be rejected at the endpoint boundary, before any tool dispatch. + """ + + @pytest.mark.asyncio + async def test_non_admin_with_user_id_none_is_rejected(self): + from fastapi import HTTPException + + from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth + from litellm.proxy.management_endpoints.usage_endpoints.endpoints import ( + ChatMessage, + UsageAIChatRequest, + usage_ai_chat, + ) + + service_account_key = UserAPIKeyAuth( + user_id=None, + user_role=LitellmUserRoles.INTERNAL_USER, + ) + request = MagicMock() + body = UsageAIChatRequest( + messages=[ChatMessage(role="user", content="hi")], + model="gpt-4o-mini", + ) + + with pytest.raises(HTTPException) as exc_info: + await usage_ai_chat( + data=body, + request=request, + user_api_key_dict=service_account_key, + ) + + assert exc_info.value.status_code == 403 + assert "Service-account keys" in str(exc_info.value.detail) + + def test_resolve_fetch_kwargs_tripwire_fires_on_none_user_id(self): + """ + Defense-in-depth: if a future endpoint forgets the entry guard and + a non-admin caller with user_id=None reaches _resolve_fetch_kwargs, + the tripwire must fire rather than issuing an unscoped query. + """ + from litellm.proxy.management_endpoints.usage_endpoints.ai_usage_chat import ( + _resolve_fetch_kwargs, + ) + + with pytest.raises(ValueError) as exc_info: + _resolve_fetch_kwargs( + fn_name="get_usage_data", + fn_args={"start_date": "2025-01-01", "end_date": "2025-01-31"}, + user_id=None, + is_admin=False, + ) + assert "Endpoint-level guard missing" in str(exc_info.value) diff --git a/tests/test_litellm/proxy/spend_tracking/test_spend_log_error_logger.py b/tests/test_litellm/proxy/spend_tracking/test_spend_log_error_logger.py new file mode 100644 index 00000000000..50d58c5849c --- /dev/null +++ b/tests/test_litellm/proxy/spend_tracking/test_spend_log_error_logger.py @@ -0,0 +1,139 @@ +""" +Unit tests for ``litellm.proxy.spend_tracking.spend_log_error_logger``. + +The helper exists to let proxy operators silence the multi-line stack traces +that the spend-tracking machinery normally emits on 4xx/5xx and DB errors. +These tests cover: + + * the env-var gating behavior (opt-in, off by default), + * the interaction between the env var and the proxy log level (DEBUG always + keeps the traceback, INFO/WARNING honors the opt-in), and + * the fact that ``spend_log_error`` always emits an ERROR-level record so + operators can still see the failure summary. +""" + +import logging + +import pytest + +from litellm._logging import verbose_proxy_logger +from litellm.proxy.spend_tracking.spend_log_error_logger import ( + SUPPRESS_SPEND_LOG_TRACEBACKS_ENV, + should_suppress_spend_log_tracebacks, + spend_log_error, +) + + +@pytest.fixture +def reset_env_and_level(monkeypatch): + """Restore both the env var and proxy logger level after each test.""" + monkeypatch.delenv(SUPPRESS_SPEND_LOG_TRACEBACKS_ENV, raising=False) + original_level = verbose_proxy_logger.level + yield monkeypatch + verbose_proxy_logger.setLevel(original_level) + + +def test_should_suppress_default_is_false(reset_env_and_level): + """With no env var set, suppression is off so existing operators see no change.""" + verbose_proxy_logger.setLevel(logging.INFO) + assert should_suppress_spend_log_tracebacks() is False + + +@pytest.mark.parametrize("value", ["true", "True", "TRUE"]) +def test_should_suppress_when_env_true_at_info(reset_env_and_level, value): + reset_env_and_level.setenv(SUPPRESS_SPEND_LOG_TRACEBACKS_ENV, value) + verbose_proxy_logger.setLevel(logging.INFO) + assert should_suppress_spend_log_tracebacks() is True + + +@pytest.mark.parametrize("value", ["false", "False", "no", "0", "", "garbage"]) +def test_should_not_suppress_when_env_falsy(reset_env_and_level, value): + if value == "": + # ``""`` would be ambiguous; ensure the var is genuinely unset. + reset_env_and_level.delenv(SUPPRESS_SPEND_LOG_TRACEBACKS_ENV, raising=False) + else: + reset_env_and_level.setenv(SUPPRESS_SPEND_LOG_TRACEBACKS_ENV, value) + verbose_proxy_logger.setLevel(logging.INFO) + assert should_suppress_spend_log_tracebacks() is False + + +def test_debug_level_overrides_suppression(reset_env_and_level): + """DEBUG always shows the traceback even when the env var is set.""" + reset_env_and_level.setenv(SUPPRESS_SPEND_LOG_TRACEBACKS_ENV, "true") + verbose_proxy_logger.setLevel(logging.DEBUG) + assert should_suppress_spend_log_tracebacks() is False + + +def test_spend_log_error_includes_traceback_by_default(reset_env_and_level, caplog): + """Default behavior: ERROR record carries exc_info so the formatter renders it.""" + verbose_proxy_logger.setLevel(logging.INFO) + caplog.set_level(logging.ERROR, logger=verbose_proxy_logger.name) + + try: + raise ValueError("boom") + except ValueError as e: + spend_log_error("update failed: %s", str(e), exc=e) + + assert len(caplog.records) == 1 + record = caplog.records[0] + assert record.levelno == logging.ERROR + assert "update failed: boom" in record.getMessage() + assert record.exc_info is not None + assert record.exc_info[0] is ValueError + + +def test_spend_log_error_drops_traceback_when_env_set(reset_env_and_level, caplog): + """Opt-in path: ERROR record still emitted, but exc_info is stripped.""" + reset_env_and_level.setenv(SUPPRESS_SPEND_LOG_TRACEBACKS_ENV, "true") + verbose_proxy_logger.setLevel(logging.INFO) + caplog.set_level(logging.ERROR, logger=verbose_proxy_logger.name) + + try: + raise ValueError("boom") + except ValueError as e: + spend_log_error("update failed: %s", str(e), exc=e) + + assert len(caplog.records) == 1 + record = caplog.records[0] + assert record.levelno == logging.ERROR + assert "update failed: boom" in record.getMessage() + assert record.exc_info is None + + +def test_spend_log_error_keeps_traceback_at_debug_even_with_env( + reset_env_and_level, caplog +): + """DEBUG operators always get tracebacks; the env var doesn't apply.""" + reset_env_and_level.setenv(SUPPRESS_SPEND_LOG_TRACEBACKS_ENV, "true") + verbose_proxy_logger.setLevel(logging.DEBUG) + caplog.set_level(logging.DEBUG, logger=verbose_proxy_logger.name) + + try: + raise RuntimeError("boom-at-debug") + except RuntimeError as e: + spend_log_error("update failed: %s", str(e), exc=e) + + error_records = [r for r in caplog.records if r.levelno == logging.ERROR] + assert len(error_records) == 1 + record = error_records[0] + assert record.exc_info is not None + assert record.exc_info[0] is RuntimeError + + +def test_spend_log_error_uses_active_exception_when_exc_omitted( + reset_env_and_level, caplog +): + """When called inside an ``except`` block without ``exc=``, the active + exception's traceback should still be attached.""" + verbose_proxy_logger.setLevel(logging.INFO) + caplog.set_level(logging.ERROR, logger=verbose_proxy_logger.name) + + try: + raise KeyError("missing") + except KeyError: + spend_log_error("update failed without exc kwarg") + + assert len(caplog.records) == 1 + record = caplog.records[0] + assert record.exc_info is not None + assert record.exc_info[0] is KeyError diff --git a/tests/test_litellm/proxy/test_budget_reservation.py b/tests/test_litellm/proxy/test_budget_reservation.py new file mode 100644 index 00000000000..070b232066a --- /dev/null +++ b/tests/test_litellm/proxy/test_budget_reservation.py @@ -0,0 +1,1495 @@ +from datetime import datetime, timedelta, timezone +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest + +import litellm +from litellm.caching.dual_cache import DualCache +from litellm.proxy._types import ( + LiteLLM_BudgetTable, + LiteLLM_EndUserTable, + LiteLLM_OrganizationTable, + LiteLLM_TagTable, + LiteLLM_TeamMembership, + LiteLLM_TeamTable, + LiteLLM_UserTable, + UserAPIKeyAuth, +) +from litellm.proxy.spend_tracking.budget_reservation import ( + estimate_request_max_cost, + get_budget_window_start, + invalidate_budget_reservation_counters, + release_budget_reservation, + reserve_budget_for_request, +) +from litellm.proxy.utils import ProxyLogging + + +@pytest.fixture() +def spend_counter_state(): + import litellm.proxy.proxy_server as ps + + original_counter_cache = ps.spend_counter_cache + original_key_cache = ps.user_api_key_cache + original_prisma_client = ps.prisma_client + + counter_cache = DualCache() + key_cache = DualCache() + ps.spend_counter_cache = counter_cache + ps.user_api_key_cache = key_cache + ps.prisma_client = None + + try: + yield counter_cache, key_cache + finally: + ps.spend_counter_cache = original_counter_cache + ps.user_api_key_cache = original_key_cache + ps.prisma_client = original_prisma_client + + +def _request_body() -> dict: + return { + "model": "gpt-4o-mini", + "messages": [{"role": "user", "content": "hello"}], + "max_tokens": 10, + } + + +def test_should_not_serialize_budget_reservation_on_user_api_key_auth(): + auth = UserAPIKeyAuth( + token="key-budget-runtime-state", + budget_reservation={ + "reserved_cost": 0.5, + "entries": [{"counter_key": "spend:key:key-budget-runtime-state"}], + }, + ) + + assert "budget_reservation" not in auth.model_dump() + assert "budget_reservation" not in auth.model_dump(exclude_none=True) + assert "budget_reservation" not in auth.model_dump_json() + + +@pytest.mark.asyncio +async def test_should_shrink_second_key_reservation_to_remaining_budget( + spend_counter_state, +): + counter_cache, key_cache = spend_counter_state + proxy_logging_obj = ProxyLogging(user_api_key_cache=key_cache) + valid_token = UserAPIKeyAuth( + token="key-budget-race", + spend=0.0, + max_budget=1.0, + ) + + with patch( + "litellm.proxy.spend_tracking.budget_reservation.estimate_request_max_cost", + return_value=0.6, + ): + reservation = await reserve_budget_for_request( + request_body=_request_body(), + route="/chat/completions", + llm_router=None, + valid_token=valid_token, + team_object=None, + user_object=None, + prisma_client=None, + user_api_key_cache=key_cache, + proxy_logging_obj=proxy_logging_obj, + ) + assert reservation is not None + assert ( + counter_cache.in_memory_cache.get_cache(key="spend:key:key-budget-race") + == 0.6 + ) + + second_reservation = await reserve_budget_for_request( + request_body=_request_body(), + route="/chat/completions", + llm_router=None, + valid_token=valid_token, + team_object=None, + user_object=None, + prisma_client=None, + user_api_key_cache=key_cache, + proxy_logging_obj=proxy_logging_obj, + ) + assert second_reservation is not None + assert second_reservation["reserved_cost"] == pytest.approx(0.4) + assert counter_cache.in_memory_cache.get_cache( + key="spend:key:key-budget-race" + ) == pytest.approx(1.0) + + with pytest.raises(litellm.BudgetExceededError): + await reserve_budget_for_request( + request_body=_request_body(), + route="/chat/completions", + llm_router=None, + valid_token=valid_token, + team_object=None, + user_object=None, + prisma_client=None, + user_api_key_cache=key_cache, + proxy_logging_obj=proxy_logging_obj, + ) + + assert counter_cache.in_memory_cache.get_cache( + key="spend:key:key-budget-race" + ) == pytest.approx(1.0) + + await release_budget_reservation(second_reservation) + assert counter_cache.in_memory_cache.get_cache( + key="spend:key:key-budget-race" + ) == pytest.approx(0.6) + await release_budget_reservation(reservation) + + +@pytest.mark.asyncio +async def test_should_shrink_second_end_user_reservation_to_remaining_budget( + spend_counter_state, +): + counter_cache, key_cache = spend_counter_state + proxy_logging_obj = ProxyLogging(user_api_key_cache=key_cache) + valid_token = UserAPIKeyAuth( + token="key-budget-end-user", + end_user_id="end-user-budget-race", + ) + end_user_object = LiteLLM_EndUserTable( + user_id="end-user-budget-race", + blocked=False, + spend=0.0, + litellm_budget_table=LiteLLM_BudgetTable(max_budget=1.0), + ) + + with patch( + "litellm.proxy.spend_tracking.budget_reservation.estimate_request_max_cost", + return_value=0.6, + ): + reservation = await reserve_budget_for_request( + request_body=_request_body(), + route="/chat/completions", + llm_router=None, + valid_token=valid_token, + team_object=None, + user_object=None, + prisma_client=None, + user_api_key_cache=key_cache, + proxy_logging_obj=proxy_logging_obj, + end_user_object=end_user_object, + ) + assert reservation is not None + assert counter_cache.in_memory_cache.get_cache( + key="spend:end_user:end-user-budget-race" + ) == pytest.approx(0.6) + + second_reservation = await reserve_budget_for_request( + request_body=_request_body(), + route="/chat/completions", + llm_router=None, + valid_token=valid_token, + team_object=None, + user_object=None, + prisma_client=None, + user_api_key_cache=key_cache, + proxy_logging_obj=proxy_logging_obj, + end_user_object=end_user_object, + ) + assert second_reservation is not None + assert second_reservation["reserved_cost"] == pytest.approx(0.4) + assert counter_cache.in_memory_cache.get_cache( + key="spend:end_user:end-user-budget-race" + ) == pytest.approx(1.0) + + with pytest.raises(litellm.BudgetExceededError): + await reserve_budget_for_request( + request_body=_request_body(), + route="/chat/completions", + llm_router=None, + valid_token=valid_token, + team_object=None, + user_object=None, + prisma_client=None, + user_api_key_cache=key_cache, + proxy_logging_obj=proxy_logging_obj, + end_user_object=end_user_object, + ) + + assert counter_cache.in_memory_cache.get_cache( + key="spend:end_user:end-user-budget-race" + ) == pytest.approx(1.0) + + await release_budget_reservation(second_reservation) + assert counter_cache.in_memory_cache.get_cache( + key="spend:end_user:end-user-budget-race" + ) == pytest.approx(0.6) + + from litellm.proxy.proxy_server import increment_spend_counters + + await increment_spend_counters( + token=None, + team_id=None, + user_id=None, + response_cost=0.2, + budget_reservation=reservation, + end_user_id="end-user-budget-race", + ) + + assert counter_cache.in_memory_cache.get_cache( + key="spend:end_user:end-user-budget-race" + ) == pytest.approx(0.2) + + +@pytest.mark.asyncio +async def test_should_shrink_second_tag_reservation_to_remaining_budget( + spend_counter_state, +): + counter_cache, key_cache = spend_counter_state + proxy_logging_obj = ProxyLogging(user_api_key_cache=key_cache) + valid_token = UserAPIKeyAuth(token="key-budget-tag") + request_body = _request_body() + request_body["metadata"] = { + "tags": ["tag-budget-race", "tag-without-budget", "tag-budget-race"] + } + await key_cache.async_set_cache( + key="tag:tag-budget-race", + value=LiteLLM_TagTable( + tag_name="tag-budget-race", + spend=0.0, + budget_id="tag-budget-id", + litellm_budget_table=LiteLLM_BudgetTable(max_budget=1.0), + ).model_dump(), + ) + await key_cache.async_set_cache( + key="tag:tag-without-budget", + value=LiteLLM_TagTable( + tag_name="tag-without-budget", + spend=0.0, + ).model_dump(), + ) + prisma_client = MagicMock() + prisma_client.db.litellm_tagtable.find_many = AsyncMock(return_value=[]) + + with patch( + "litellm.proxy.spend_tracking.budget_reservation.estimate_request_max_cost", + return_value=0.6, + ): + reservation = await reserve_budget_for_request( + request_body=request_body, + route="/chat/completions", + llm_router=None, + valid_token=valid_token, + team_object=None, + user_object=None, + prisma_client=prisma_client, + user_api_key_cache=key_cache, + proxy_logging_obj=proxy_logging_obj, + ) + assert reservation is not None + assert reservation["entries"] == [ + { + "counter_key": "spend:tag:tag-budget-race", + "entity_type": "Tag", + "entity_id": "tag-budget-race", + "reserved_cost": 0.6, + "applied_adjustment": 0.0, + } + ] + assert counter_cache.in_memory_cache.get_cache( + key="spend:tag:tag-budget-race" + ) == pytest.approx(0.6) + assert ( + counter_cache.in_memory_cache.get_cache(key="spend:tag:tag-without-budget") + is None + ) + + second_reservation = await reserve_budget_for_request( + request_body=request_body, + route="/chat/completions", + llm_router=None, + valid_token=valid_token, + team_object=None, + user_object=None, + prisma_client=prisma_client, + user_api_key_cache=key_cache, + proxy_logging_obj=proxy_logging_obj, + ) + assert second_reservation is not None + assert second_reservation["reserved_cost"] == pytest.approx(0.4) + assert counter_cache.in_memory_cache.get_cache( + key="spend:tag:tag-budget-race" + ) == pytest.approx(1.0) + + with pytest.raises(litellm.BudgetExceededError): + await reserve_budget_for_request( + request_body=request_body, + route="/chat/completions", + llm_router=None, + valid_token=valid_token, + team_object=None, + user_object=None, + prisma_client=prisma_client, + user_api_key_cache=key_cache, + proxy_logging_obj=proxy_logging_obj, + ) + + assert counter_cache.in_memory_cache.get_cache( + key="spend:tag:tag-budget-race" + ) == pytest.approx(1.0) + + await release_budget_reservation(second_reservation) + assert counter_cache.in_memory_cache.get_cache( + key="spend:tag:tag-budget-race" + ) == pytest.approx(0.6) + + from litellm.proxy.proxy_server import increment_spend_counters + + await increment_spend_counters( + token=None, + team_id=None, + user_id=None, + response_cost=0.2, + budget_reservation=reservation, + tags=["tag-budget-race"], + ) + + assert counter_cache.in_memory_cache.get_cache( + key="spend:tag:tag-budget-race" + ) == pytest.approx(0.2) + + +@pytest.mark.asyncio +async def test_should_seed_and_update_end_user_and_tag_counters_without_reservation( + spend_counter_state, +): + counter_cache, key_cache = spend_counter_state + await key_cache.async_set_cache( + key="end_user_id:customer-1", + value=LiteLLM_EndUserTable( + user_id="customer-1", + blocked=False, + spend=4.0, + litellm_budget_table=LiteLLM_BudgetTable(max_budget=10.0), + ).model_dump(), + ) + await key_cache.async_set_cache( + key="tag:paid-tag", + value=LiteLLM_TagTable( + tag_name="paid-tag", + spend=7.0, + ).model_dump(), + ) + await key_cache.async_set_cache( + key="tag:other-tag", + value=LiteLLM_TagTable( + tag_name="other-tag", + spend=2.0, + ).model_dump(), + ) + + from litellm.proxy.proxy_server import increment_spend_counters + + await increment_spend_counters( + token=None, + team_id=None, + user_id=None, + response_cost=0.50, + end_user_id="customer-1", + tags=["paid-tag", "paid-tag", "other-tag", ""], + ) + + assert counter_cache.in_memory_cache.get_cache( + key="spend:end_user:customer-1" + ) == pytest.approx(4.50) + assert counter_cache.in_memory_cache.get_cache( + key="spend:tag:paid-tag" + ) == pytest.approx(7.50) + assert counter_cache.in_memory_cache.get_cache( + key="spend:tag:other-tag" + ) == pytest.approx(2.50) + + +@pytest.mark.asyncio +async def test_should_reserve_team_member_and_org_budget_counters(spend_counter_state): + counter_cache, key_cache = spend_counter_state + proxy_logging_obj = ProxyLogging(user_api_key_cache=key_cache) + valid_token = UserAPIKeyAuth( + token="key-budget-shared", + spend=0.0, + max_budget=1.0, + user_id="user-budget-shared", + team_id="team-budget-shared", + org_id="org-budget-shared", + ) + team_object = LiteLLM_TeamTable( + team_id="team-budget-shared", + spend=0.0, + max_budget=1.0, + ) + user_object = LiteLLM_UserTable( + user_id="user-budget-shared", + spend=0.0, + ) + await key_cache.async_set_cache( + key="team_membership:user-budget-shared:team-budget-shared", + value=LiteLLM_TeamMembership( + user_id="user-budget-shared", + team_id="team-budget-shared", + spend=0.1, + litellm_budget_table=LiteLLM_BudgetTable(max_budget=1.0), + ).model_dump(), + ) + await key_cache.async_set_cache( + key="org_id:org-budget-shared:with_budget", + value=LiteLLM_OrganizationTable( + organization_id="org-budget-shared", + organization_alias="shared-org", + budget_id="org-budget-id", + spend=0.1, + models=[], + created_by="test", + updated_by="test", + litellm_budget_table=LiteLLM_BudgetTable(max_budget=1.0), + ).model_dump(), + ) + + with patch( + "litellm.proxy.spend_tracking.budget_reservation.estimate_request_max_cost", + return_value=0.3, + ): + reservation = await reserve_budget_for_request( + request_body=_request_body(), + route="/chat/completions", + llm_router=None, + valid_token=valid_token, + team_object=team_object, + user_object=user_object, + prisma_client=None, + user_api_key_cache=key_cache, + proxy_logging_obj=proxy_logging_obj, + ) + + assert counter_cache.in_memory_cache.get_cache( + key="spend:team_member:user-budget-shared:team-budget-shared" + ) == pytest.approx(0.4) + assert counter_cache.in_memory_cache.get_cache( + key="spend:org:org-budget-shared" + ) == pytest.approx(0.4) + + await release_budget_reservation(reservation) + + +@pytest.mark.asyncio +async def test_should_seed_org_counter_from_with_budget_cache(spend_counter_state): + counter_cache, key_cache = spend_counter_state + await key_cache.async_set_cache( + key="org_id:org-counter-with-budget:with_budget", + value=LiteLLM_OrganizationTable( + organization_id="org-counter-with-budget", + organization_alias="shared-org", + budget_id="org-budget-id", + spend=2.0, + models=[], + created_by="test", + updated_by="test", + litellm_budget_table=LiteLLM_BudgetTable(max_budget=10.0), + ).model_dump(), + ) + + from litellm.proxy.proxy_server import increment_spend_counters + + await increment_spend_counters( + token=None, + team_id=None, + user_id=None, + org_id="org-counter-with-budget", + response_cost=0.25, + ) + + assert counter_cache.in_memory_cache.get_cache( + key="spend:org:org-counter-with-budget" + ) == pytest.approx(2.25) + + +@pytest.mark.asyncio +async def test_should_seed_org_counter_from_plain_org_cache(spend_counter_state): + counter_cache, key_cache = spend_counter_state + await key_cache.async_set_cache( + key="org_id:org-counter-plain", + value=LiteLLM_OrganizationTable( + organization_id="org-counter-plain", + organization_alias="shared-org", + budget_id="org-budget-id", + spend=2.0, + models=[], + created_by="test", + updated_by="test", + ).model_dump(), + ) + + from litellm.proxy.proxy_server import increment_spend_counters + + await increment_spend_counters( + token=None, + team_id=None, + user_id=None, + org_id="org-counter-plain", + response_cost=0.25, + ) + + assert counter_cache.in_memory_cache.get_cache( + key="spend:org:org-counter-plain" + ) == pytest.approx(2.25) + + +@pytest.mark.asyncio +async def test_should_cap_known_estimate_to_remaining_budget( + spend_counter_state, +): + counter_cache, key_cache = spend_counter_state + proxy_logging_obj = ProxyLogging(user_api_key_cache=key_cache) + valid_token = UserAPIKeyAuth( + token="key-budget-known-estimate-cap", + spend=0.9, + max_budget=1.0, + ) + counter_cache.in_memory_cache.set_cache( + key="spend:key:key-budget-known-estimate-cap", + value=0.9, + ) + + with patch( + "litellm.proxy.spend_tracking.budget_reservation.estimate_request_max_cost", + return_value=0.6, + ): + reservation = await reserve_budget_for_request( + request_body=_request_body(), + route="/chat/completions", + llm_router=None, + valid_token=valid_token, + team_object=None, + user_object=None, + prisma_client=None, + user_api_key_cache=key_cache, + proxy_logging_obj=proxy_logging_obj, + ) + + assert reservation is not None + assert reservation["reserved_cost"] == pytest.approx(0.1) + assert counter_cache.in_memory_cache.get_cache( + key="spend:key:key-budget-known-estimate-cap" + ) == pytest.approx(1.0) + + await release_budget_reservation(reservation) + assert counter_cache.in_memory_cache.get_cache( + key="spend:key:key-budget-known-estimate-cap" + ) == pytest.approx(0.9) + + +@pytest.mark.asyncio +async def test_should_reserve_remaining_budget_when_output_cap_missing( + spend_counter_state, +): + counter_cache, key_cache = spend_counter_state + proxy_logging_obj = ProxyLogging(user_api_key_cache=key_cache) + valid_token = UserAPIKeyAuth( + token="key-budget-uncapped", + spend=0.2, + max_budget=1.0, + ) + await key_cache.async_set_cache( + key="key-budget-uncapped", + value=valid_token, + ) + request_body = _request_body() + request_body.pop("max_tokens") + + with patch( + "litellm.proxy.spend_tracking.budget_reservation._get_model_cost_info", + return_value={ + "input_cost_per_token": 0.0, + "output_cost_per_token": 100.0, + "max_output_tokens": 200000, + }, + ): + assert ( + estimate_request_max_cost( + request_body=request_body, + route="/chat/completions", + llm_router=None, + ) + is None + ) + reservation = await reserve_budget_for_request( + request_body=request_body, + route="/chat/completions", + llm_router=None, + valid_token=valid_token, + team_object=None, + user_object=None, + prisma_client=None, + user_api_key_cache=key_cache, + proxy_logging_obj=proxy_logging_obj, + ) + + assert reservation is not None + assert reservation["reserved_cost"] == pytest.approx(0.8) + assert counter_cache.in_memory_cache.get_cache( + key="spend:key:key-budget-uncapped" + ) == pytest.approx(1.0) + + await release_budget_reservation(reservation) + + +@pytest.mark.asyncio +async def test_should_shrink_uncapped_reservation_when_counter_advances( + spend_counter_state, + monkeypatch, +): + counter_cache, key_cache = spend_counter_state + proxy_logging_obj = ProxyLogging(user_api_key_cache=key_cache) + valid_token = UserAPIKeyAuth( + token="key-budget-uncapped-race", + spend=0.2, + max_budget=1.0, + ) + request_body = _request_body() + request_body.pop("max_tokens") + + from litellm.proxy.spend_tracking import budget_reservation + + async def stale_counter_read(counter): + await counter_cache.async_increment_cache( + key=counter.counter_key, + value=0.3, + ) + return 0.2 + + monkeypatch.setattr( + budget_reservation, + "_get_current_counter_value", + stale_counter_read, + ) + + with patch( + "litellm.proxy.spend_tracking.budget_reservation.estimate_request_max_cost", + return_value=None, + ): + reservation = await reserve_budget_for_request( + request_body=request_body, + route="/chat/completions", + llm_router=None, + valid_token=valid_token, + team_object=None, + user_object=None, + prisma_client=None, + user_api_key_cache=key_cache, + proxy_logging_obj=proxy_logging_obj, + ) + + assert reservation is not None + assert reservation["reserved_cost"] == pytest.approx(0.7) + assert counter_cache.in_memory_cache.get_cache( + key="spend:key:key-budget-uncapped-race" + ) == pytest.approx(1.0) + + await release_budget_reservation(reservation) + + assert counter_cache.in_memory_cache.get_cache( + key="spend:key:key-budget-uncapped-race" + ) == pytest.approx(0.3) + + +@pytest.mark.asyncio +async def test_should_shrink_uncapped_reservation_multiple_times( + spend_counter_state, + monkeypatch, +): + counter_cache, key_cache = spend_counter_state + proxy_logging_obj = ProxyLogging(user_api_key_cache=key_cache) + valid_token = UserAPIKeyAuth( + token="key-budget-double-resize", + spend=0.2, + max_budget=1.0, + team_id="team-budget-double-resize", + ) + team_object = LiteLLM_TeamTable( + team_id="team-budget-double-resize", + spend=0.2, + max_budget=1.0, + ) + request_body = _request_body() + request_body.pop("max_tokens") + + from litellm.proxy.spend_tracking import budget_reservation + + stale_spend_by_counter_key = { + "spend:key:key-budget-double-resize": 0.3, + "spend:team:team-budget-double-resize": 0.4, + } + + async def stale_counter_read(counter): + await counter_cache.async_increment_cache( + key=counter.counter_key, + value=stale_spend_by_counter_key[counter.counter_key], + ) + return 0.2 + + monkeypatch.setattr( + budget_reservation, + "_get_current_counter_value", + stale_counter_read, + ) + + with patch( + "litellm.proxy.spend_tracking.budget_reservation.estimate_request_max_cost", + return_value=None, + ): + reservation = await reserve_budget_for_request( + request_body=request_body, + route="/chat/completions", + llm_router=None, + valid_token=valid_token, + team_object=team_object, + user_object=None, + prisma_client=None, + user_api_key_cache=key_cache, + proxy_logging_obj=proxy_logging_obj, + ) + + assert reservation is not None + assert reservation["reserved_cost"] == pytest.approx(0.6) + assert [entry["reserved_cost"] for entry in reservation["entries"]] == [ + pytest.approx(0.6), + pytest.approx(0.6), + ] + assert [entry["applied_adjustment"] for entry in reservation["entries"]] == [ + pytest.approx(0.0), + pytest.approx(0.0), + ] + assert counter_cache.in_memory_cache.get_cache( + key="spend:key:key-budget-double-resize" + ) == pytest.approx(0.9) + assert counter_cache.in_memory_cache.get_cache( + key="spend:team:team-budget-double-resize" + ) == pytest.approx(1.0) + + await release_budget_reservation(reservation) + + assert counter_cache.in_memory_cache.get_cache( + key="spend:key:key-budget-double-resize" + ) == pytest.approx(0.3) + assert counter_cache.in_memory_cache.get_cache( + key="spend:team:team-budget-double-resize" + ) == pytest.approx(0.4) + + +def test_should_start_window_without_reset_at_at_duration_boundary(): + before = datetime.now(timezone.utc) - timedelta(hours=1) + + window_start = get_budget_window_start({"budget_duration": "1h"}) + + after = datetime.now(timezone.utc) - timedelta(hours=1) + assert window_start is not None + assert before <= window_start <= after + + +@pytest.mark.asyncio +async def test_should_skip_budget_window_with_unparseable_duration( + spend_counter_state, +): + counter_cache, key_cache = spend_counter_state + proxy_logging_obj = ProxyLogging(user_api_key_cache=key_cache) + valid_token = UserAPIKeyAuth( + token="key-budget-malformed-window", + spend=0.9, + max_budget=10.0, + budget_limits=[ + { + "budget_duration": "not-a-duration", + "max_budget": 1.0, + } + ], + ) + counter_cache.in_memory_cache.set_cache( + key="spend:key:key-budget-malformed-window", + value=0.9, + ) + + with patch( + "litellm.proxy.spend_tracking.budget_reservation.estimate_request_max_cost", + return_value=0.2, + ): + reservation = await reserve_budget_for_request( + request_body=_request_body(), + route="/chat/completions", + llm_router=None, + valid_token=valid_token, + team_object=None, + user_object=None, + prisma_client=None, + user_api_key_cache=key_cache, + proxy_logging_obj=proxy_logging_obj, + ) + + assert reservation is not None + assert [entry["counter_key"] for entry in reservation["entries"]] == [ + "spend:key:key-budget-malformed-window" + ] + assert counter_cache.in_memory_cache.get_cache( + key="spend:key:key-budget-malformed-window" + ) == pytest.approx(1.1) + assert ( + counter_cache.in_memory_cache.get_cache( + key="spend:key:key-budget-malformed-window:window:not-a-duration" + ) + is None + ) + + await release_budget_reservation(reservation) + assert counter_cache.in_memory_cache.get_cache( + key="spend:key:key-budget-malformed-window" + ) == pytest.approx(0.9) + + +@pytest.mark.asyncio +async def test_should_skip_window_reservation_when_db_baseline_unavailable( + spend_counter_state, +): + counter_cache, key_cache = spend_counter_state + proxy_logging_obj = ProxyLogging(user_api_key_cache=key_cache) + valid_token = UserAPIKeyAuth( + token="key-budget-window-db-unavailable", + budget_limits=[ + { + "budget_duration": "1h", + "max_budget": 1.0, + } + ], + ) + + with patch( + "litellm.proxy.spend_tracking.budget_reservation.estimate_request_max_cost", + return_value=0.5, + ): + reservation = await reserve_budget_for_request( + request_body=_request_body(), + route="/chat/completions", + llm_router=None, + valid_token=valid_token, + team_object=None, + user_object=None, + prisma_client=None, + user_api_key_cache=key_cache, + proxy_logging_obj=proxy_logging_obj, + ) + + assert reservation is None + assert ( + counter_cache.in_memory_cache.get_cache( + key="spend:key:key-budget-window-db-unavailable:window:1h" + ) + is None + ) + + +@pytest.mark.asyncio +async def test_should_skip_reservation_when_counter_increment_fails( + spend_counter_state, + monkeypatch, +): + counter_cache, key_cache = spend_counter_state + proxy_logging_obj = ProxyLogging(user_api_key_cache=key_cache) + valid_token = UserAPIKeyAuth( + token="key-budget-reserve-unavailable", + spend=0.0, + max_budget=1.0, + ) + + async def fail_increment_cache(*args, **kwargs): + raise RuntimeError("counter unavailable") + + monkeypatch.setattr(counter_cache, "async_increment_cache", fail_increment_cache) + + with ( + patch( + "litellm.proxy.spend_tracking.budget_reservation.estimate_request_max_cost", + return_value=0.5, + ), + patch( + "litellm.proxy.spend_tracking.budget_reservation.verbose_proxy_logger.warning" + ) as mock_warning, + ): + reservation = await reserve_budget_for_request( + request_body=_request_body(), + route="/chat/completions", + llm_router=None, + valid_token=valid_token, + team_object=None, + user_object=None, + prisma_client=None, + user_api_key_cache=key_cache, + proxy_logging_obj=proxy_logging_obj, + ) + + assert reservation is None + assert mock_warning.call_count >= 1 + assert ( + counter_cache.in_memory_cache.get_cache( + key="spend:key:key-budget-reserve-unavailable" + ) + is None + ) + + +@pytest.mark.asyncio +async def test_should_skip_reservation_when_counter_initialization_fails( + spend_counter_state, +): + counter_cache, key_cache = spend_counter_state + proxy_logging_obj = ProxyLogging(user_api_key_cache=key_cache) + valid_token = UserAPIKeyAuth( + token="key-budget-reserve-init-unavailable", + spend=0.0, + max_budget=1.0, + ) + + with ( + patch( + "litellm.proxy.spend_tracking.budget_reservation.estimate_request_max_cost", + return_value=0.5, + ), + patch( + "litellm.proxy.proxy_server._ensure_spend_counter_initialized", + side_effect=RuntimeError("redis unavailable"), + ), + patch( + "litellm.proxy.spend_tracking.budget_reservation.verbose_proxy_logger.warning" + ) as mock_warning, + ): + reservation = await reserve_budget_for_request( + request_body=_request_body(), + route="/chat/completions", + llm_router=None, + valid_token=valid_token, + team_object=None, + user_object=None, + prisma_client=None, + user_api_key_cache=key_cache, + proxy_logging_obj=proxy_logging_obj, + ) + + assert reservation is None + assert mock_warning.call_count >= 1 + assert ( + counter_cache.in_memory_cache.get_cache( + key="spend:key:key-budget-reserve-init-unavailable" + ) + is None + ) + + +@pytest.mark.asyncio +async def test_should_release_tracked_entry_when_reservation_fails_after_increment( + spend_counter_state, +): + counter_cache, key_cache = spend_counter_state + proxy_logging_obj = ProxyLogging(user_api_key_cache=key_cache) + valid_token = UserAPIKeyAuth( + token="key-budget-reserve-after-increment-failure", + spend=0.0, + max_budget=1.0, + ) + + import litellm.proxy.proxy_server as ps + + original_increment_counter = ps._increment_spend_counter_cache + first_increment = True + + async def fail_after_increment(counter_key: str, increment: float): + nonlocal first_increment + if first_increment: + first_increment = False + await counter_cache.async_increment_cache(key=counter_key, value=increment) + raise RuntimeError("lost increment response") + return await original_increment_counter( + counter_key=counter_key, + increment=increment, + ) + + with ( + patch( + "litellm.proxy.spend_tracking.budget_reservation.estimate_request_max_cost", + return_value=0.5, + ), + patch( + "litellm.proxy.proxy_server._increment_spend_counter_cache", + side_effect=fail_after_increment, + ), + patch( + "litellm.proxy.proxy_server._invalidate_spend_counter", + side_effect=RuntimeError("invalidate unavailable"), + ), + ): + reservation = await reserve_budget_for_request( + request_body=_request_body(), + route="/chat/completions", + llm_router=None, + valid_token=valid_token, + team_object=None, + user_object=None, + prisma_client=None, + user_api_key_cache=key_cache, + proxy_logging_obj=proxy_logging_obj, + ) + + assert reservation is None + assert counter_cache.in_memory_cache.get_cache( + key="spend:key:key-budget-reserve-after-increment-failure" + ) == pytest.approx(0.0) + + +@pytest.mark.asyncio +async def test_should_not_re_read_uncapped_budget_after_reservation_fallback( + spend_counter_state, + monkeypatch, +): + _, key_cache = spend_counter_state + proxy_logging_obj = ProxyLogging(user_api_key_cache=key_cache) + valid_token = UserAPIKeyAuth( + token="key-budget-uncapped-read-once", + spend=0.2, + max_budget=1.0, + ) + + from litellm.proxy.spend_tracking import budget_reservation + + current_counter_reads = [] + + async def mock_get_current_counter_value(counter): + current_counter_reads.append(counter.counter_key) + return counter.fallback_spend + + async def mock_reserve_counter(counter, reservation_cost): + return None + + monkeypatch.setattr( + budget_reservation, + "_get_current_counter_value", + mock_get_current_counter_value, + ) + monkeypatch.setattr( + budget_reservation, + "_reserve_counter", + mock_reserve_counter, + ) + + with patch( + "litellm.proxy.spend_tracking.budget_reservation.estimate_request_max_cost", + return_value=None, + ): + reservation = await reserve_budget_for_request( + request_body=_request_body(), + route="/chat/completions", + llm_router=None, + valid_token=valid_token, + team_object=None, + user_object=None, + prisma_client=None, + user_api_key_cache=key_cache, + proxy_logging_obj=proxy_logging_obj, + ) + + assert reservation is not None + assert reservation["reserved_cost"] == pytest.approx(0.8) + assert current_counter_reads == ["spend:key:key-budget-uncapped-read-once"] + + +@pytest.mark.asyncio +async def test_should_reconcile_reserved_counter_to_actual_spend( + spend_counter_state, +): + counter_cache, key_cache = spend_counter_state + proxy_logging_obj = ProxyLogging(user_api_key_cache=key_cache) + valid_token = UserAPIKeyAuth( + token="key-budget-reconcile", + spend=0.0, + max_budget=1.0, + ) + + with patch( + "litellm.proxy.spend_tracking.budget_reservation.estimate_request_max_cost", + return_value=0.6, + ): + reservation = await reserve_budget_for_request( + request_body=_request_body(), + route="/chat/completions", + llm_router=None, + valid_token=valid_token, + team_object=None, + user_object=None, + prisma_client=None, + user_api_key_cache=key_cache, + proxy_logging_obj=proxy_logging_obj, + ) + + from litellm.proxy.proxy_server import increment_spend_counters + + await increment_spend_counters( + token="key-budget-reconcile", + team_id="team-without-budget", + user_id=None, + response_cost=0.2, + budget_reservation=reservation, + ) + + assert counter_cache.in_memory_cache.get_cache( + key="spend:key:key-budget-reconcile" + ) == pytest.approx(0.2) + assert counter_cache.in_memory_cache.get_cache( + key="spend:team:team-without-budget" + ) == pytest.approx(0.2) + + +@pytest.mark.asyncio +async def test_should_release_reservation_on_failure(spend_counter_state): + counter_cache, key_cache = spend_counter_state + proxy_logging_obj = ProxyLogging(user_api_key_cache=key_cache) + valid_token = UserAPIKeyAuth( + token="key-budget-release", + spend=0.0, + max_budget=1.0, + ) + + with patch( + "litellm.proxy.spend_tracking.budget_reservation.estimate_request_max_cost", + return_value=0.4, + ): + reservation = await reserve_budget_for_request( + request_body=_request_body(), + route="/chat/completions", + llm_router=None, + valid_token=valid_token, + team_object=None, + user_object=None, + prisma_client=None, + user_api_key_cache=key_cache, + proxy_logging_obj=proxy_logging_obj, + ) + + await release_budget_reservation(reservation) + await release_budget_reservation(reservation) + + assert counter_cache.in_memory_cache.get_cache( + key="spend:key:key-budget-release" + ) == pytest.approx(0.0) + + +@pytest.mark.asyncio +async def test_should_retry_partial_release_without_double_decrement( + spend_counter_state, + monkeypatch, +): + counter_cache, key_cache = spend_counter_state + proxy_logging_obj = ProxyLogging(user_api_key_cache=key_cache) + valid_token = UserAPIKeyAuth( + token="key-budget-partial-release", + spend=0.0, + max_budget=1.0, + team_id="team-budget-partial-release", + ) + team_object = LiteLLM_TeamTable( + team_id="team-budget-partial-release", + spend=0.0, + max_budget=1.0, + ) + + with patch( + "litellm.proxy.spend_tracking.budget_reservation.estimate_request_max_cost", + return_value=0.4, + ): + reservation = await reserve_budget_for_request( + request_body=_request_body(), + route="/chat/completions", + llm_router=None, + valid_token=valid_token, + team_object=team_object, + user_object=None, + prisma_client=None, + user_api_key_cache=key_cache, + proxy_logging_obj=proxy_logging_obj, + ) + + original_increment_cache = counter_cache.async_increment_cache + fail_next_team_release = True + + async def flaky_increment_cache(key, value, *args, **kwargs): + nonlocal fail_next_team_release + if ( + key == "spend:team:team-budget-partial-release" + and value < 0 + and fail_next_team_release + ): + fail_next_team_release = False + raise RuntimeError("simulated counter failure") + return await original_increment_cache(key=key, value=value, *args, **kwargs) + + monkeypatch.setattr(counter_cache, "async_increment_cache", flaky_increment_cache) + + with pytest.raises(RuntimeError): + await release_budget_reservation(reservation) + + assert counter_cache.in_memory_cache.get_cache( + key="spend:key:key-budget-partial-release" + ) == pytest.approx(0.0) + assert counter_cache.in_memory_cache.get_cache( + key="spend:team:team-budget-partial-release" + ) == pytest.approx(0.4) + + await release_budget_reservation(reservation) + + assert counter_cache.in_memory_cache.get_cache( + key="spend:key:key-budget-partial-release" + ) == pytest.approx(0.0) + assert counter_cache.in_memory_cache.get_cache( + key="spend:team:team-budget-partial-release" + ) == pytest.approx(0.0) + + +@pytest.mark.asyncio +async def test_should_preserve_budget_error_and_continue_partial_cleanup( + spend_counter_state, + monkeypatch, +): + counter_cache, key_cache = spend_counter_state + proxy_logging_obj = ProxyLogging(user_api_key_cache=key_cache) + valid_token = UserAPIKeyAuth( + token="key-budget-cleanup-failure", + spend=0.0, + max_budget=1.0, + team_id="team-budget-cleanup-failure", + ) + team_object = LiteLLM_TeamTable( + team_id="team-budget-cleanup-failure", + spend=0.3, + max_budget=0.3, + ) + await key_cache.async_set_cache( + key="team_id:team-budget-cleanup-failure", + value=team_object, + ) + + original_increment_cache = counter_cache.async_increment_cache + fail_key_cleanup = True + + async def flaky_increment_cache(key, value, *args, **kwargs): + nonlocal fail_key_cleanup + if key == "spend:key:key-budget-cleanup-failure" and value < 0: + if fail_key_cleanup: + fail_key_cleanup = False + raise RuntimeError("simulated cleanup failure") + return await original_increment_cache(key=key, value=value, *args, **kwargs) + + monkeypatch.setattr(counter_cache, "async_increment_cache", flaky_increment_cache) + + with ( + patch( + "litellm.proxy.spend_tracking.budget_reservation.estimate_request_max_cost", + return_value=0.4, + ), + patch( + "litellm.proxy.spend_tracking.budget_reservation.verbose_proxy_logger.exception" + ) as mock_log_exception, + ): + with pytest.raises(litellm.BudgetExceededError): + await reserve_budget_for_request( + request_body=_request_body(), + route="/chat/completions", + llm_router=None, + valid_token=valid_token, + team_object=team_object, + user_object=None, + prisma_client=None, + user_api_key_cache=key_cache, + proxy_logging_obj=proxy_logging_obj, + ) + + assert ( + counter_cache.in_memory_cache.get_cache( + key="spend:key:key-budget-cleanup-failure" + ) + is None + ) + assert counter_cache.in_memory_cache.get_cache( + key="spend:team:team-budget-cleanup-failure" + ) == pytest.approx(0.3) + mock_log_exception.assert_called() + + +@pytest.mark.asyncio +async def test_should_not_create_negative_counter_when_release_counter_is_missing( + spend_counter_state, +): + counter_cache, _ = spend_counter_state + reservation = { + "reserved_cost": 0.4, + "entries": [ + { + "counter_key": "spend:key:key-budget-missing-release", + "reserved_cost": 0.4, + "applied_adjustment": 0.0, + } + ], + "finalized": False, + } + + with pytest.raises(RuntimeError, match="missing counter"): + await release_budget_reservation(reservation) + + assert ( + counter_cache.in_memory_cache.get_cache( + key="spend:key:key-budget-missing-release" + ) + is None + ) + assert reservation["finalized"] is False + + +@pytest.mark.asyncio +async def test_should_invalidate_counter_when_release_would_underflow( + spend_counter_state, +): + counter_cache, _ = spend_counter_state + await counter_cache.async_increment_cache( + key="spend:key:key-budget-underflow-release", + value=0.1, + ) + reservation = { + "reserved_cost": 0.4, + "entries": [ + { + "counter_key": "spend:key:key-budget-underflow-release", + "reserved_cost": 0.4, + "applied_adjustment": 0.0, + } + ], + "finalized": False, + } + + with pytest.raises(RuntimeError, match="negative"): + await release_budget_reservation(reservation) + + assert ( + counter_cache.in_memory_cache.get_cache( + key="spend:key:key-budget-underflow-release" + ) + is None + ) + assert reservation["finalized"] is False + + +@pytest.mark.asyncio +async def test_should_invalidate_non_numeric_counter_during_release( + spend_counter_state, +): + counter_cache, _ = spend_counter_state + counter_cache.in_memory_cache.set_cache( + key="spend:key:key-budget-nonnumeric-release", + value="stale", + ) + reservation = { + "reserved_cost": 0.4, + "entries": [ + { + "counter_key": "spend:key:key-budget-nonnumeric-release", + "reserved_cost": 0.4, + "applied_adjustment": 0.0, + } + ], + "finalized": False, + } + + with pytest.raises(RuntimeError, match="non-numeric"): + await release_budget_reservation(reservation) + + assert ( + counter_cache.in_memory_cache.get_cache( + key="spend:key:key-budget-nonnumeric-release" + ) + is None + ) + assert reservation["finalized"] is False + + +@pytest.mark.asyncio +async def test_should_invalidate_reserved_counters_after_persisted_spend_failure( + spend_counter_state, +): + counter_cache, _ = spend_counter_state + await counter_cache.async_increment_cache( + key="spend:key:key-budget-invalidate", + value=0.4, + ) + await counter_cache.async_increment_cache( + key="spend:team:team-budget-invalidate", + value=0.4, + ) + + await invalidate_budget_reservation_counters( + { + "reserved_cost": 0.4, + "entries": [ + {"counter_key": "spend:key:key-budget-invalidate"}, + {"counter_key": "spend:team:team-budget-invalidate"}, + ], + } + ) + + assert ( + counter_cache.in_memory_cache.get_cache(key="spend:key:key-budget-invalidate") + is None + ) + assert ( + counter_cache.in_memory_cache.get_cache(key="spend:team:team-budget-invalidate") + is None + ) + + +@pytest.mark.asyncio +async def test_should_reserve_all_budgeted_counters(spend_counter_state): + counter_cache, key_cache = spend_counter_state + proxy_logging_obj = ProxyLogging(user_api_key_cache=key_cache) + valid_token = UserAPIKeyAuth( + token="key-budget-all", + spend=0.0, + max_budget=1.0, + team_id="team-budget-all", + ) + team_object = LiteLLM_TeamTable( + team_id="team-budget-all", + spend=0.0, + max_budget=1.0, + ) + + with patch( + "litellm.proxy.spend_tracking.budget_reservation.estimate_request_max_cost", + return_value=0.3, + ): + reservation = await reserve_budget_for_request( + request_body=_request_body(), + route="/chat/completions", + llm_router=None, + valid_token=valid_token, + team_object=team_object, + user_object=None, + prisma_client=None, + user_api_key_cache=key_cache, + proxy_logging_obj=proxy_logging_obj, + ) + + assert ( + counter_cache.in_memory_cache.get_cache(key="spend:key:key-budget-all") == 0.3 + ) + assert ( + counter_cache.in_memory_cache.get_cache(key="spend:team:team-budget-all") == 0.3 + ) + + await release_budget_reservation(reservation) diff --git a/tests/test_litellm/proxy/test_langfuse_passthrough_security.py b/tests/test_litellm/proxy/test_langfuse_passthrough_security.py new file mode 100644 index 00000000000..5ef3c38c09d --- /dev/null +++ b/tests/test_litellm/proxy/test_langfuse_passthrough_security.py @@ -0,0 +1,102 @@ +import socket + +import pytest +from fastapi import HTTPException + +import litellm +from litellm.proxy.vertex_ai_endpoints.langfuse_endpoints import ( + _build_langfuse_proxy_target, + _get_langfuse_proxy_credentials, +) + + +def test_dynamic_langfuse_host_requires_dynamic_credentials(monkeypatch): + monkeypatch.setattr(litellm, "user_url_validation", True, raising=False) + monkeypatch.setenv("LANGFUSE_PUBLIC_KEY", "global-public") + monkeypatch.setenv("LANGFUSE_SECRET_KEY", "global-secret") + + with pytest.raises(HTTPException) as exc: + _get_langfuse_proxy_credentials( + dynamic_host_supplied=True, + dynamic_langfuse_public_key=None, + dynamic_langfuse_secret_key=None, + ) + + assert exc.value.status_code == 400 + + +def test_global_langfuse_host_can_use_env_credentials(monkeypatch): + monkeypatch.setenv("LANGFUSE_PUBLIC_KEY", "global-public") + monkeypatch.setenv("LANGFUSE_SECRET_KEY", "global-secret") + + public_key, secret_key = _get_langfuse_proxy_credentials( + dynamic_host_supplied=False, + dynamic_langfuse_public_key=None, + dynamic_langfuse_secret_key=None, + ) + + assert public_key == "global-public" + assert secret_key == "global-secret" + + +@pytest.mark.parametrize( + "endpoint", + [ + "../api/public/projects", + "%2e%2e/api/public/projects", + "%252e%252e%252fapi/public/projects", + "api\\public\\projects", + "%2f%2fattacker.example/api", + ], +) +def test_langfuse_proxy_target_rejects_traversal_paths(endpoint): + with pytest.raises(HTTPException) as exc: + _build_langfuse_proxy_target( + endpoint=endpoint, + base_target_url="https://cloud.langfuse.com", + dynamic_host_supplied=False, + ) + + assert exc.value.status_code == 400 + + +def test_dynamic_langfuse_proxy_target_rejects_internal_host(monkeypatch): + monkeypatch.setattr(litellm, "user_url_validation", True, raising=False) + + with pytest.raises(HTTPException) as exc: + _build_langfuse_proxy_target( + endpoint="api/public/projects", + base_target_url="http://127.0.0.1:3000", + dynamic_host_supplied=True, + ) + + assert exc.value.status_code == 400 + + +def test_dynamic_langfuse_proxy_target_preserves_host_header_for_http(monkeypatch): + monkeypatch.setattr(litellm, "user_url_validation", True, raising=False) + + def fake_getaddrinfo(host, port, proto): + assert host == "langfuse.example" + assert port == 80 + assert proto == socket.IPPROTO_TCP + return [ + ( + socket.AF_INET, + socket.SOCK_STREAM, + socket.IPPROTO_TCP, + "", + ("8.8.8.8", 80), + ) + ] + + monkeypatch.setattr(socket, "getaddrinfo", fake_getaddrinfo) + + target_url, headers = _build_langfuse_proxy_target( + endpoint="api/public/projects", + base_target_url="http://langfuse.example", + dynamic_host_supplied=True, + ) + + assert target_url == "http://8.8.8.8/api/public/projects" + assert headers["Host"] == "langfuse.example" diff --git a/tests/test_litellm/proxy/test_lazy_openapi_snapshot.py b/tests/test_litellm/proxy/test_lazy_openapi_snapshot.py index 8bc39c93eeb..64cb931888b 100644 --- a/tests/test_litellm/proxy/test_lazy_openapi_snapshot.py +++ b/tests/test_litellm/proxy/test_lazy_openapi_snapshot.py @@ -1,6 +1,89 @@ +import sys +from types import ModuleType, SimpleNamespace + from litellm.proxy._lazy_openapi_snapshot import _normalize_operation_ids +def test_generate_snapshot_uses_shared_operation_id_reservations(monkeypatch): + from litellm.proxy import _lazy_openapi_snapshot + + route_a = SimpleNamespace(path="/feature-a/items") + route_b = SimpleNamespace(path="/feature-b/items") + fake_app = SimpleNamespace( + title="LiteLLM test", + version="0.0.0", + routes=[route_a, route_b], + ) + + fake_feature_a_module = ModuleType("fake_feature_a") + fake_feature_b_module = ModuleType("fake_feature_b") + monkeypatch.setitem(sys.modules, "fake_feature_a", fake_feature_a_module) + monkeypatch.setitem(sys.modules, "fake_feature_b", fake_feature_b_module) + + fake_lazy_features_module = ModuleType("litellm.proxy._lazy_features") + fake_lazy_features_module.LAZY_FEATURES = [ + SimpleNamespace( + name="feature-a", + module_path="fake_feature_a", + path_prefixes=("/feature-a",), + register_fn=lambda app, module: None, + ), + SimpleNamespace( + name="feature-b", + module_path="fake_feature_b", + path_prefixes=("/feature-b",), + register_fn=lambda app, module: None, + ), + ] + monkeypatch.setitem( + sys.modules, "litellm.proxy._lazy_features", fake_lazy_features_module + ) + + def fake_get_openapi(title, version, routes): + path = routes[0].path + return { + "paths": {path: {"get": {"operationId": "shared_operation_id_get"}}}, + "components": {"schemas": {"Example": {"type": "object"}}}, + } + + def fake_ensure_unique_openapi_operation_ids(schema, reserved_operation_ids): + for path_item in schema["paths"].values(): + operation = path_item["get"] + operation_id = operation["operationId"] + if operation_id in reserved_operation_ids: + operation_id = f"{operation_id}_2" + operation["operationId"] = operation_id + reserved_operation_ids.add(operation_id) + return schema + + fake_proxy_server_module = ModuleType("litellm.proxy.proxy_server") + fake_proxy_server_module.app = fake_app + fake_proxy_server_module.ensure_unique_openapi_operation_ids = ( + fake_ensure_unique_openapi_operation_ids + ) + monkeypatch.setitem( + sys.modules, "litellm.proxy.proxy_server", fake_proxy_server_module + ) + monkeypatch.setattr("fastapi.openapi.utils.get_openapi", fake_get_openapi) + + fragments = _lazy_openapi_snapshot.generate_snapshot() + + assert ( + fragments["feature-a"]["paths"]["/feature-a/items"]["get"]["operationId"] + == "shared_operation_id_get" + ) + assert ( + fragments["feature-b"]["paths"]["/feature-b/items"]["get"]["operationId"] + == "shared_operation_id_get_2" + ) + assert fragments["feature-a"]["paths"]["/feature-a/items"]["get"]["tags"] == [ + "feature-a" + ] + assert fragments["feature-b"]["paths"]["/feature-b/items"]["get"]["tags"] == [ + "feature-b" + ] + + def test_normalize_operation_ids_uses_each_http_method(): paths = { "/proxy/{endpoint}": { diff --git a/tests/test_litellm/proxy/test_litellm_pre_call_utils.py b/tests/test_litellm/proxy/test_litellm_pre_call_utils.py index 6be1a9ecef0..4e24d8af653 100644 --- a/tests/test_litellm/proxy/test_litellm_pre_call_utils.py +++ b/tests/test_litellm/proxy/test_litellm_pre_call_utils.py @@ -2846,6 +2846,61 @@ async def test_add_guardrails_from_policy_engine_accepts_dynamic_policies_and_po assert "metadata" in data +@pytest.mark.asyncio +async def test_api_created_global_policy_applies_to_new_key_without_restart(): + """ + Regression: policies created at runtime via policy builder must apply + immediately when attached globally, even if the server started with no + initialized policy config. + """ + from litellm.proxy.policy_engine.attachment_registry import get_attachment_registry + from litellm.proxy.policy_engine.policy_registry import get_policy_registry + from litellm.types.proxy.policy_engine import ( + Policy, + PolicyAttachment, + PolicyGuardrails, + ) + + data = { + "model": "gpt-4", + "messages": [{"role": "user", "content": "Hello"}], + "metadata": {}, + } + user_api_key_dict = UserAPIKeyAuth(api_key="test-key") + + policy_registry = get_policy_registry() + attachment_registry = get_attachment_registry() + policy_registry._policies = {} + policy_registry._policies_by_id = {} + policy_registry._initialized = False + attachment_registry._attachments = [] + attachment_registry._initialized = False + + try: + policy_registry.add_policy( + "runtime-global-policy", + Policy(guardrails=PolicyGuardrails(add=["runtime-guardrail"])), + ) + attachment_registry.add_attachment( + PolicyAttachment(policy="runtime-global-policy", scope="*") + ) + + await add_guardrails_from_policy_engine( + data=data, + metadata_variable_name="metadata", + user_api_key_dict=user_api_key_dict, + ) + + assert "runtime-guardrail" in data["metadata"]["guardrails"] + assert "runtime-global-policy" in data["metadata"]["applied_policies"] + finally: + policy_registry._policies = {} + policy_registry._policies_by_id = {} + policy_registry._initialized = False + attachment_registry._attachments = [] + attachment_registry._initialized = False + + @pytest.mark.asyncio async def test_add_guardrails_from_policy_engine_policy_version_by_id(): """ diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py index 465ce579e09..3f19db36c3f 100644 --- a/tests/test_litellm/proxy/test_proxy_server.py +++ b/tests/test_litellm/proxy/test_proxy_server.py @@ -5,7 +5,7 @@ import os import socket import subprocess import sys -from datetime import datetime, timezone +from datetime import datetime, timedelta, timezone from pathlib import Path from unittest import mock from unittest.mock import AsyncMock, MagicMock, mock_open, patch @@ -5084,8 +5084,12 @@ async def test_init_and_increment_spend_counter_reseeds_from_db_on_counter_miss( @pytest.mark.asyncio async def test_reseed_spend_from_db_user_and_org_prefixes(): - """User and org counters must reseed from their own DB tables, not - fall through to 0.0 like the other counters do today.""" + """User and org counters reseed from their own DB tables. + + End-user and tag counters use the already fetched auth objects passed as + fallback_spend, so this reseed helper must not add extra per-request DB + reads for them. + """ from litellm.proxy.db.spend_counter_reseed import SpendCounterReseed user_row = MagicMock() @@ -5095,6 +5099,8 @@ async def test_reseed_spend_from_db_user_and_org_prefixes(): fake_prisma = MagicMock() fake_prisma.db.litellm_usertable.find_unique = AsyncMock(return_value=user_row) + fake_prisma.db.litellm_endusertable.find_unique = AsyncMock() + fake_prisma.db.litellm_tagtable.find_unique = AsyncMock() fake_prisma.db.litellm_organizationtable.find_unique = AsyncMock( return_value=org_row ) @@ -5104,6 +5110,18 @@ async def test_reseed_spend_from_db_user_and_org_prefixes(): where={"user_id": "alice"} ) + assert ( + await SpendCounterReseed.from_db( + fake_prisma, + "spend:end_user:customer-1", + ) + is None + ) + fake_prisma.db.litellm_endusertable.find_unique.assert_not_awaited() + + assert await SpendCounterReseed.from_db(fake_prisma, "spend:tag:paid-tag") is None + fake_prisma.db.litellm_tagtable.find_unique.assert_not_awaited() + assert await SpendCounterReseed.from_db(fake_prisma, "spend:org:acme") == 305.0 fake_prisma.db.litellm_organizationtable.find_unique.assert_awaited_once_with( where={"organization_id": "acme"} @@ -5133,6 +5151,468 @@ async def test_reseed_spend_from_db_skips_window_variant_keys(): fake_prisma.db.litellm_teamtable.find_unique.assert_not_awaited() +@pytest.mark.asyncio +async def test_window_spend_counter_reseeds_from_spend_logs_on_counter_miss(): + from litellm.caching.dual_cache import DualCache + from litellm.proxy.proxy_server import _init_and_increment_window_spend_counter + + counter_cache = DualCache() + window_start = datetime.now(timezone.utc) - timedelta(hours=1) + fake_prisma = MagicMock() + fake_prisma.db.litellm_spendlogs.group_by = AsyncMock( + return_value=[{"api_key": "key-window", "_sum": {"spend": 2.25}}] + ) + + import litellm.proxy.proxy_server as ps + + orig_counter, orig_prisma = ps.spend_counter_cache, ps.prisma_client + ps.spend_counter_cache = counter_cache + ps.prisma_client = fake_prisma + try: + await _init_and_increment_window_spend_counter( + counter_key="spend:key:key-window:window:1h", + entity_type="Key", + entity_id="key-window", + window_start=window_start, + increment=0.5, + ) + + fake_prisma.db.litellm_spendlogs.group_by.assert_awaited_once_with( + by=["api_key"], + where={"api_key": "key-window", "startTime": {"gte": window_start}}, + sum={"spend": True}, + ) + assert counter_cache.in_memory_cache.get_cache( + key="spend:key:key-window:window:1h" + ) == pytest.approx(2.75) + finally: + ps.spend_counter_cache = orig_counter + ps.prisma_client = orig_prisma + + +@pytest.mark.asyncio +async def test_init_spend_counter_redis_clean_miss_skips_stale_in_memory(): + from litellm.caching.dual_cache import DualCache + from litellm.proxy.proxy_server import _init_and_increment_spend_counter + + counter_cache = DualCache() + counter_key = "spend:team:team-stale-local" + counter_cache.in_memory_cache.set_cache(key=counter_key, value=10.0) + + redis_store: dict = {} + + async def redis_increment(key, value, **_): + redis_store[key] = (redis_store.get(key) or 0.0) + value + return redis_store[key] + + fake_redis = AsyncMock() + fake_redis.async_get_cache = AsyncMock(return_value=None) + fake_redis.async_increment = AsyncMock(side_effect=redis_increment) + counter_cache.redis_cache = fake_redis + + db_row = MagicMock() + db_row.spend = 42.0 + fake_prisma = MagicMock() + fake_prisma.db.litellm_teamtable.find_unique = AsyncMock(return_value=db_row) + + import litellm.proxy.proxy_server as ps + + orig_counter, orig_prisma, orig_user = ( + ps.spend_counter_cache, + ps.prisma_client, + ps.user_api_key_cache, + ) + ps.spend_counter_cache = counter_cache + ps.prisma_client = fake_prisma + ps.user_api_key_cache = DualCache() + try: + await _init_and_increment_spend_counter( + counter_key=counter_key, + source_cache_key="team_id:team-stale-local", + increment=1.5, + ) + + fake_prisma.db.litellm_teamtable.find_unique.assert_awaited_once_with( + where={"team_id": "team-stale-local"} + ) + assert redis_store[counter_key] == pytest.approx(43.5) + assert counter_cache.in_memory_cache.get_cache( + key=counter_key + ) == pytest.approx(43.5) + finally: + ps.spend_counter_cache = orig_counter + ps.prisma_client = orig_prisma + ps.user_api_key_cache = orig_user + + +@pytest.mark.asyncio +async def test_window_spend_counter_redis_clean_miss_skips_stale_in_memory(): + from litellm.caching.dual_cache import DualCache + from litellm.proxy.proxy_server import _init_and_increment_window_spend_counter + + counter_cache = DualCache() + counter_key = "spend:key:key-window-stale-local:window:1h" + counter_cache.in_memory_cache.set_cache(key=counter_key, value=100.0) + window_start = datetime.now(timezone.utc) - timedelta(hours=1) + + redis_store: dict = {} + + async def redis_increment(key, value, **_): + redis_store[key] = (redis_store.get(key) or 0.0) + value + return redis_store[key] + + async def redis_set_cache(key, value, **_): + if key in redis_store: + return False + redis_store[key] = value + return True + + fake_redis = AsyncMock() + fake_redis.async_get_cache = AsyncMock(return_value=None) + fake_redis.async_set_cache = AsyncMock(side_effect=redis_set_cache) + fake_redis.async_increment = AsyncMock(side_effect=redis_increment) + counter_cache.redis_cache = fake_redis + + fake_prisma = MagicMock() + fake_prisma.db.litellm_spendlogs.group_by = AsyncMock( + return_value=[{"api_key": "key-window-stale-local", "_sum": {"spend": 2.25}}] + ) + + import litellm.proxy.proxy_server as ps + + orig_counter, orig_prisma = ps.spend_counter_cache, ps.prisma_client + ps.spend_counter_cache = counter_cache + ps.prisma_client = fake_prisma + try: + await _init_and_increment_window_spend_counter( + counter_key=counter_key, + entity_type="Key", + entity_id="key-window-stale-local", + window_start=window_start, + increment=0.5, + ) + + fake_prisma.db.litellm_spendlogs.group_by.assert_awaited_once_with( + by=["api_key"], + where={ + "api_key": "key-window-stale-local", + "startTime": {"gte": window_start}, + }, + sum={"spend": True}, + ) + assert redis_store[counter_key] == pytest.approx(2.75) + assert counter_cache.in_memory_cache.get_cache( + key=counter_key + ) == pytest.approx(2.75) + finally: + ps.spend_counter_cache = orig_counter + ps.prisma_client = orig_prisma + + +@pytest.mark.asyncio +async def test_window_spend_counter_redis_concurrent_seed_does_not_double_seed(): + from litellm.caching.dual_cache import DualCache + from litellm.proxy.proxy_server import _init_and_increment_window_spend_counter + + counter_cache = DualCache() + counter_key = "spend:key:key-window-concurrent-seed:window:1h" + window_start = datetime.now(timezone.utc) - timedelta(hours=1) + redis_store = {counter_key: 2.75} + redis_reads = 0 + + async def redis_get_cache(key): + nonlocal redis_reads + redis_reads += 1 + if redis_reads <= 2: + return None + return redis_store.get(key) + + async def redis_increment(key, value, **_): + redis_store[key] = (redis_store.get(key) or 0.0) + value + return redis_store[key] + + fake_redis = AsyncMock() + fake_redis.async_get_cache = AsyncMock(side_effect=redis_get_cache) + fake_redis.async_set_cache = AsyncMock(return_value=False) + fake_redis.async_increment = AsyncMock(side_effect=redis_increment) + counter_cache.redis_cache = fake_redis + + fake_prisma = MagicMock() + fake_prisma.db.litellm_spendlogs.group_by = AsyncMock( + return_value=[ + {"api_key": "key-window-concurrent-seed", "_sum": {"spend": 2.25}} + ] + ) + + import litellm.proxy.proxy_server as ps + + orig_counter, orig_prisma = ps.spend_counter_cache, ps.prisma_client + ps.spend_counter_cache = counter_cache + ps.prisma_client = fake_prisma + try: + await _init_and_increment_window_spend_counter( + counter_key=counter_key, + entity_type="Key", + entity_id="key-window-concurrent-seed", + window_start=window_start, + increment=0.5, + ) + + fake_redis.async_set_cache.assert_awaited_once_with( + key=counter_key, + value=2.25, + nx=True, + ) + assert redis_store[counter_key] == pytest.approx(3.25) + assert counter_cache.in_memory_cache.get_cache( + key=counter_key + ) == pytest.approx(3.25) + finally: + ps.spend_counter_cache = orig_counter + ps.prisma_client = orig_prisma + + +@pytest.mark.asyncio +async def test_window_spend_counter_skips_invalid_window_start(): + from litellm.caching.dual_cache import DualCache + from litellm.proxy.proxy_server import _init_and_increment_window_spend_counter + + counter_cache = DualCache() + + import litellm.proxy.proxy_server as ps + + orig_counter = ps.spend_counter_cache + ps.spend_counter_cache = counter_cache + try: + await _init_and_increment_window_spend_counter( + counter_key="spend:key:key-invalid-window:window:not-a-duration", + entity_type="Key", + entity_id="key-invalid-window", + window_start=None, + increment=0.5, + ) + + assert ( + counter_cache.in_memory_cache.get_cache( + key="spend:key:key-invalid-window:window:not-a-duration" + ) + is None + ) + finally: + ps.spend_counter_cache = orig_counter + + +@pytest.mark.asyncio +async def test_window_spend_counter_does_not_seed_zero_when_db_unavailable(): + from litellm.caching.dual_cache import DualCache + from litellm.proxy.proxy_server import _ensure_window_spend_counter_initialized + + counter_cache = DualCache() + counter_key = "spend:key:key-window-db-unavailable:window:1h" + + import litellm.proxy.proxy_server as ps + + orig_counter, orig_prisma = ps.spend_counter_cache, ps.prisma_client + ps.spend_counter_cache = counter_cache + ps.prisma_client = None + try: + initialized = await _ensure_window_spend_counter_initialized( + counter_key=counter_key, + entity_type="Key", + entity_id="key-window-db-unavailable", + window_start=datetime.now(timezone.utc) - timedelta(hours=1), + ) + + assert initialized is False + assert counter_cache.in_memory_cache.get_cache(key=counter_key) is None + finally: + ps.spend_counter_cache = orig_counter + ps.prisma_client = orig_prisma + + +@pytest.mark.asyncio +async def test_increment_spend_counters_finalizes_after_unreserved_increments(): + from litellm.caching.dual_cache import DualCache + from litellm.proxy.proxy_server import increment_spend_counters + + counter_cache = DualCache() + counter_cache.in_memory_cache.set_cache( + key="spend:key:key-finalize-after-increments", + value=0.5, + ) + budget_reservation = { + "reserved_cost": 0.5, + "entries": [ + { + "counter_key": "spend:key:key-finalize-after-increments", + "entity_type": "Key", + "entity_id": "key-finalize-after-increments", + "reserved_cost": 0.5, + "applied_adjustment": 0.0, + } + ], + "finalized": False, + } + incremented_counters = [] + + async def assert_reservation_not_finalized_yet(**kwargs): + assert budget_reservation["finalized"] is False + incremented_counters.append(kwargs["counter_key"]) + + import litellm.proxy.proxy_server as ps + + orig_counter, orig_user = ps.spend_counter_cache, ps.user_api_key_cache + ps.spend_counter_cache = counter_cache + ps.user_api_key_cache = DualCache() + try: + with patch( + "litellm.proxy.proxy_server._init_and_increment_spend_counter", + new=AsyncMock(side_effect=assert_reservation_not_finalized_yet), + ): + await increment_spend_counters( + token="key-finalize-after-increments", + team_id="team-finalize-after-increments", + user_id=None, + response_cost=0.25, + budget_reservation=budget_reservation, + ) + + assert incremented_counters == ["spend:team:team-finalize-after-increments"] + assert budget_reservation["finalized"] is True + assert counter_cache.in_memory_cache.get_cache( + key="spend:key:key-finalize-after-increments" + ) == pytest.approx(0.25) + finally: + ps.spend_counter_cache = orig_counter + ps.user_api_key_cache = orig_user + + +@pytest.mark.asyncio +async def test_increment_spend_counters_finalizes_none_cost_reservation(): + from litellm.caching.dual_cache import DualCache + from litellm.proxy.proxy_server import increment_spend_counters + + counter_cache = DualCache() + counter_cache.in_memory_cache.set_cache( + key="spend:key:key-finalize-none-cost", + value=0.5, + ) + budget_reservation = { + "reserved_cost": 0.5, + "entries": [ + { + "counter_key": "spend:key:key-finalize-none-cost", + "entity_type": "Key", + "entity_id": "key-finalize-none-cost", + "reserved_cost": 0.5, + "applied_adjustment": 0.0, + } + ], + "finalized": False, + } + + import litellm.proxy.proxy_server as ps + + orig_counter = ps.spend_counter_cache + ps.spend_counter_cache = counter_cache + try: + await increment_spend_counters( + token="key-finalize-none-cost", + team_id=None, + user_id=None, + response_cost=None, + budget_reservation=budget_reservation, + ) + + assert budget_reservation["finalized"] is True + assert counter_cache.in_memory_cache.get_cache( + key="spend:key:key-finalize-none-cost" + ) == pytest.approx(0.0) + finally: + ps.spend_counter_cache = orig_counter + + +@pytest.mark.asyncio +async def test_increment_spend_counters_invalidates_bad_reserved_counter_without_failing(): + from litellm.caching.dual_cache import DualCache + from litellm.proxy.proxy_server import increment_spend_counters + + counter_cache = DualCache() + budget_reservation = { + "reserved_cost": 0.5, + "entries": [ + { + "counter_key": "spend:key:key-bad-reserved-counter", + "entity_type": "Key", + "entity_id": "key-bad-reserved-counter", + "reserved_cost": 0.5, + "applied_adjustment": 0.0, + } + ], + "finalized": False, + } + + import litellm.proxy.proxy_server as ps + + orig_counter = ps.spend_counter_cache + ps.spend_counter_cache = counter_cache + try: + with patch( + "litellm.proxy.proxy_server.verbose_proxy_logger.warning" + ) as mock_warning: + await increment_spend_counters( + token="key-bad-reserved-counter", + team_id=None, + user_id=None, + response_cost=0.25, + budget_reservation=budget_reservation, + ) + + mock_warning.assert_called_once() + assert budget_reservation["finalized"] is True + assert ( + counter_cache.in_memory_cache.get_cache( + key="spend:key:key-bad-reserved-counter" + ) + is None + ) + finally: + ps.spend_counter_cache = orig_counter + + +@pytest.mark.asyncio +async def test_increment_spend_counter_invalidates_stale_cache_on_redis_failure(): + from litellm.caching.dual_cache import DualCache + from litellm.proxy.proxy_server import _increment_spend_counter_cache + + counter_cache = DualCache() + counter_cache.in_memory_cache.set_cache(key="spend:team:redis-fail", value=4.0) + fake_redis = AsyncMock() + fake_redis.async_increment = AsyncMock(side_effect=RuntimeError("redis down")) + fake_redis.async_delete_cache = AsyncMock() + counter_cache.redis_cache = fake_redis + + import litellm.proxy.proxy_server as ps + + orig_counter = ps.spend_counter_cache + ps.spend_counter_cache = counter_cache + try: + with pytest.raises(RuntimeError): + await _increment_spend_counter_cache( + counter_key="spend:team:redis-fail", + increment=0.5, + ) + + assert ( + counter_cache.in_memory_cache.get_cache(key="spend:team:redis-fail") is None + ) + fake_redis.async_delete_cache.assert_awaited_once_with( + key="spend:team:redis-fail" + ) + finally: + ps.spend_counter_cache = orig_counter + + @pytest.mark.asyncio async def test_get_current_spend_reseeds_from_db_when_counter_missing(): """ @@ -5181,6 +5661,9 @@ async def test_get_current_spend_reseeds_from_db_when_counter_missing(): assert ("spend:team_member:user-1:team-1", 362.0) in [ (w["key"], w["value"]) for w in recorded_warms ] + assert counter_cache.in_memory_cache.get_cache( + key="spend:team_member:user-1:team-1" + ) == pytest.approx(362.0) finally: ps.spend_counter_cache = orig_counter ps.prisma_client = orig_prisma diff --git a/tests/test_litellm/proxy/test_route_llm_request.py b/tests/test_litellm/proxy/test_route_llm_request.py index bfea21e705e..98b0b6be025 100644 --- a/tests/test_litellm/proxy/test_route_llm_request.py +++ b/tests/test_litellm/proxy/test_route_llm_request.py @@ -241,6 +241,53 @@ async def test_route_request_with_router_settings_override_preserves_existing(): assert call_kwargs["timeout"] == 30 +def test_mock_testing_kwarg_names_matches_dataclass(): + """``_MOCK_TESTING_KWARG_NAMES`` is hardcoded to avoid a cyclic import + against ``litellm.types.router``. This test guards against drift — + if a new ``mock_testing_*`` field is added to ``MockRouterTestingParams`` + the strip list must be updated to keep covering it.""" + from dataclasses import fields + + from litellm.proxy.route_llm_request import _MOCK_TESTING_KWARG_NAMES + from litellm.types.router import MockRouterTestingParams + + assert set(_MOCK_TESTING_KWARG_NAMES) == { + f.name for f in fields(MockRouterTestingParams) + } + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "mock_flag", + [ + "mock_testing_fallbacks", + "mock_testing_context_fallbacks", + "mock_testing_content_policy_fallbacks", + ], +) +async def test_route_request_strips_mock_testing_flags(mock_flag): + """VERIA-44: router-internal testing flags must not survive a + user-supplied request body. Without this strip, an attacker can + combine ``mock_testing_fallbacks=true`` with an unauthorized fallback + in ``router_settings_override`` to deterministically execute requests + against restricted models.""" + data = { + "model": "gpt-3.5-turbo", + "messages": [{"role": "user", "content": "Hello"}], + mock_flag: True, + } + llm_router = MagicMock() + llm_router.acompletion.return_value = "ok" + + await route_request(data, llm_router, None, "acompletion") + + call_kwargs = llm_router.acompletion.call_args[1] + assert mock_flag not in call_kwargs + # The flag is also gone from the original data dict so any subsequent + # processing (e.g. logging) doesn't see it either. + assert mock_flag not in data + + @pytest.mark.parametrize( "route_type", ["agenerate_content", "agenerate_content_stream"] ) diff --git a/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_endpoints.py b/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_endpoints.py index 6dd0e0e68a6..e67a04c749a 100644 --- a/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_endpoints.py +++ b/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_endpoints.py @@ -156,8 +156,11 @@ async def test_update_request_data_with_litellm_managed_vector_store_registry(): vector_store_id="test_store_id" ) - # Test with no vector store registry - with patch.object(litellm, "vector_store_registry", None): + # Test with no vector store registry or DB fallback + with ( + patch.object(litellm, "vector_store_registry", None), + patch("litellm.proxy.proxy_server.prisma_client", None), + ): original_data = {"existing_key": "existing_value"} result = await _update_request_data_with_litellm_managed_vector_store_registry( data=original_data, vector_store_id=vector_store_id diff --git a/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_tenant_guard.py b/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_tenant_guard.py new file mode 100644 index 00000000000..48262afd363 --- /dev/null +++ b/tests/test_litellm/proxy/vector_store_endpoints/test_vector_store_tenant_guard.py @@ -0,0 +1,540 @@ +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest +from fastapi import HTTPException, Request, Response + +import litellm +from litellm.proxy._types import LiteLLM_ManagedVectorStoresTable, UserAPIKeyAuth + + +def _mock_request() -> MagicMock: + request = MagicMock(spec=Request) + request.headers = {} + request.method = "POST" + request.query_params = {} + request.url.path = "/v1/vector_stores/vs_path/search" + return request + + +@pytest.mark.asyncio +async def test_vector_store_search_forces_path_id_over_body_id(): + from litellm.proxy.vector_store_endpoints.endpoints import vector_store_search + + captured_data = {} + + async def fake_base_process(self, **kwargs): + captured_data.update(self.data) + return {"ok": True} + + request = _mock_request() + with ( + patch( + "litellm.proxy.proxy_server._read_request_body", + new=AsyncMock( + return_value={ + "vector_store_id": "vs_body_victim", + "query": "test", + } + ), + ), + patch.object(litellm, "vector_store_registry", None), + patch("litellm.proxy.proxy_server.prisma_client", None), + patch( + "litellm.proxy.vector_store_endpoints.endpoints.ProxyBaseLLMRequestProcessing.base_process_llm_request", + new=fake_base_process, + ), + ): + response = await vector_store_search( + request=request, + vector_store_id="vs_path_allowed", + fastapi_response=Response(), + user_api_key_dict=UserAPIKeyAuth(team_id="team-a"), + ) + + assert response == {"ok": True} + assert captured_data["vector_store_id"] == "vs_path_allowed" + + +@pytest.mark.asyncio +async def test_vector_store_file_create_forces_path_id_over_body_id(): + from litellm.proxy.vector_store_files_endpoints.endpoints import ( + vector_store_file_create, + ) + + captured_data = {} + + async def fake_base_process(self, **kwargs): + captured_data.update(self.data) + return {"ok": True} + + mock_registry = MagicMock() + mock_registry.get_litellm_managed_vector_store_from_registry.return_value = { + "vector_store_id": "vs_path_allowed", + "custom_llm_provider": "openai", + "team_id": "team-a", + } + + request = _mock_request() + with ( + patch( + "litellm.proxy.proxy_server._read_request_body", + new=AsyncMock( + return_value={ + "vector_store_id": "vs_body_victim", + "file_id": "file_123", + } + ), + ), + patch.object(litellm, "vector_store_registry", mock_registry), + patch( + "litellm.proxy.vector_store_files_endpoints.endpoints.ProxyBaseLLMRequestProcessing.base_process_llm_request", + new=fake_base_process, + ), + ): + response = await vector_store_file_create( + vector_store_id="vs_path_allowed", + request=request, + fastapi_response=Response(), + user_api_key_dict=UserAPIKeyAuth(team_id="team-a"), + ) + + assert response == {"ok": True} + assert captured_data["vector_store_id"] == "vs_path_allowed" + assert captured_data["custom_llm_provider"] == "openai" + mock_registry.get_litellm_managed_vector_store_from_registry.assert_called_once_with( + vector_store_id="vs_path_allowed" + ) + + +@pytest.mark.asyncio +async def test_vector_store_file_create_denies_other_team_path_store(): + from litellm.proxy.vector_store_files_endpoints.endpoints import ( + vector_store_file_create, + ) + + mock_registry = MagicMock() + mock_registry.get_litellm_managed_vector_store_from_registry.return_value = { + "vector_store_id": "vs_other_team", + "custom_llm_provider": "openai", + "team_id": "team-b", + } + + request = _mock_request() + with ( + patch( + "litellm.proxy.proxy_server._read_request_body", + new=AsyncMock(return_value={"file_id": "file_123"}), + ), + patch.object(litellm, "vector_store_registry", mock_registry), + patch( + "litellm.proxy.vector_store_files_endpoints.endpoints.ProxyBaseLLMRequestProcessing.base_process_llm_request", + new=AsyncMock(), + ) as mock_base_process, + ): + with pytest.raises(HTTPException) as exc_info: + await vector_store_file_create( + vector_store_id="vs_other_team", + request=request, + fastapi_response=Response(), + user_api_key_dict=UserAPIKeyAuth(team_id="team-a"), + ) + + assert exc_info.value.status_code == 403 + mock_base_process.assert_not_called() + + +@pytest.mark.asyncio +async def test_rag_query_denies_nested_other_team_vector_store(): + from litellm.proxy.rag_endpoints.endpoints import rag_query + + mock_registry = MagicMock() + mock_registry.get_litellm_managed_vector_store_from_registry.return_value = { + "vector_store_id": "vs_other_team", + "custom_llm_provider": "openai", + "team_id": "team-b", + } + + request = _mock_request() + with ( + patch( + "litellm.proxy.rag_endpoints.endpoints._read_request_body", + new=AsyncMock( + return_value={ + "model": "gpt-4o-mini", + "messages": [{"role": "user", "content": "hello"}], + "retrieval_config": {"vector_store_id": "vs_other_team"}, + } + ), + ), + patch.object(litellm, "vector_store_registry", mock_registry), + patch( + "litellm.proxy.rag_endpoints.endpoints.litellm.aquery", + new=AsyncMock(), + ) as mock_aquery, + ): + with pytest.raises(HTTPException) as exc_info: + await rag_query( + request=request, + fastapi_response=Response(), + user_api_key_dict=UserAPIKeyAuth(team_id="team-a"), + ) + + assert exc_info.value.status_code == 403 + mock_aquery.assert_not_called() + + +@pytest.mark.asyncio +async def test_rag_ingest_denies_nested_other_team_vector_store(): + from litellm.proxy.rag_endpoints.endpoints import rag_ingest + + mock_registry = MagicMock() + mock_registry.get_litellm_managed_vector_store_from_registry.return_value = { + "vector_store_id": "vs_other_team", + "custom_llm_provider": "openai", + "team_id": "team-b", + } + + request = _mock_request() + with ( + patch( + "litellm.proxy.rag_endpoints.endpoints.parse_rag_ingest_request", + new=AsyncMock( + return_value=( + { + "vector_store": { + "custom_llm_provider": "openai", + "vector_store_id": "vs_other_team", + } + }, + None, + "https://example.com/file.txt", + None, + ) + ), + ), + patch.object(litellm, "vector_store_registry", mock_registry), + patch( + "litellm.proxy.rag_endpoints.endpoints.litellm.aingest", + new=AsyncMock(), + ) as mock_aingest, + ): + with pytest.raises(HTTPException) as exc_info: + await rag_ingest( + request=request, + fastapi_response=Response(), + user_api_key_dict=UserAPIKeyAuth(team_id="team-a"), + ) + + assert exc_info.value.status_code == 403 + mock_aingest.assert_not_called() + + +def test_rag_payload_scan_rejects_excessive_nesting(): + from litellm.constants import DEFAULT_MAX_RECURSE_DEPTH + from litellm.proxy.rag_endpoints.endpoints import ( + _collect_vector_store_ids_from_payload, + ) + + payload = {} + current = payload + for _ in range(DEFAULT_MAX_RECURSE_DEPTH + 1): + current["nested"] = {} + current = current["nested"] + current["vector_store_id"] = "vs_too_deep" + + with pytest.raises(HTTPException) as exc_info: + _collect_vector_store_ids_from_payload(payload) + + assert exc_info.value.status_code == 400 + + +def test_rag_payload_scan_accepts_vector_store_id_at_depth_limit(): + from litellm.constants import DEFAULT_MAX_RECURSE_DEPTH + from litellm.proxy.rag_endpoints.endpoints import ( + _collect_vector_store_ids_from_payload, + ) + + payload = {} + current = payload + for _ in range(DEFAULT_MAX_RECURSE_DEPTH): + current["nested"] = {} + current = current["nested"] + current["vector_store_id"] = "vs_at_limit" + + assert _collect_vector_store_ids_from_payload(payload) == {"vs_at_limit"} + + +def test_rag_payload_scan_ignores_primitive_list_beyond_depth_limit(): + from litellm.constants import DEFAULT_MAX_RECURSE_DEPTH + from litellm.proxy.rag_endpoints.endpoints import ( + _collect_vector_store_ids_from_payload, + ) + + payload = {} + current = payload + for _ in range(DEFAULT_MAX_RECURSE_DEPTH): + current["nested"] = {} + current = current["nested"] + current["labels"] = ["alpha", "beta"] + + assert _collect_vector_store_ids_from_payload(payload) == set() + + +@pytest.mark.asyncio +async def test_responses_file_search_denies_other_team_vector_store(): + from litellm.proxy.common_request_processing import ( + _authorize_response_file_search_vector_stores, + ) + + mock_registry = MagicMock() + mock_registry.get_litellm_managed_vector_store_from_registry.return_value = { + "vector_store_id": "vs_other_team", + "custom_llm_provider": "openai", + "team_id": "team-b", + } + + with patch.object(litellm, "vector_store_registry", mock_registry): + with pytest.raises(HTTPException) as exc_info: + await _authorize_response_file_search_vector_stores( + data={ + "tools": [ + { + "type": "file_search", + "vector_store_ids": ["vs_other_team"], + } + ] + }, + user_api_key_dict=UserAPIKeyAuth(team_id="team-a"), + ) + + assert exc_info.value.status_code == 403 + + +@pytest.mark.asyncio +async def test_vertex_discovery_denies_other_team_vector_store_credentials(): + from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import ( + _base_vertex_proxy_route, + ) + + request = _mock_request() + request.method = "GET" + vector_store_credentials = { + "vector_store_id": "vs_other_team", + "custom_llm_provider": "vertex_ai", + "team_id": "team-b", + } + + with patch( + "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.user_api_key_auth", + new=AsyncMock(return_value=UserAPIKeyAuth(team_id="team-a")), + ): + with pytest.raises(HTTPException) as exc_info: + await _base_vertex_proxy_route( + endpoint="projects/p/locations/us-central1/dataStores/vs_other_team", + request=request, + fastapi_response=Response(), + get_vertex_pass_through_handler=MagicMock(), + router_credentials=vector_store_credentials, + ) + + assert exc_info.value.status_code == 403 + + +@pytest.mark.asyncio +async def test_get_managed_vector_store_uses_shared_cache_helper_for_db_fallback(): + from litellm.proxy.vector_store_endpoints.utils import ( + get_litellm_managed_vector_store, + ) + + mock_registry = MagicMock() + mock_registry.get_litellm_managed_vector_store_from_registry.return_value = None + cache_helper = AsyncMock( + return_value=[ + LiteLLM_ManagedVectorStoresTable( + vector_store_id="vs_cached", + custom_llm_provider="openai", + vector_store_name=None, + vector_store_description=None, + vector_store_metadata=None, + created_at=None, + updated_at=None, + litellm_credential_name=None, + litellm_params={"api_base": "https://example.com"}, + team_id="team-a", + user_id=None, + ) + ] + ) + + with ( + patch.object(litellm, "vector_store_registry", mock_registry), + patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), + patch("litellm.proxy.proxy_server.user_api_key_cache", MagicMock()), + patch("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock()), + patch( + "litellm.proxy.auth.auth_checks.get_managed_vector_store_rows_by_uuids", + new=cache_helper, + ), + ): + vector_store = await get_litellm_managed_vector_store( + vector_store_id="vs_cached" + ) + + assert vector_store is not None + assert vector_store["vector_store_id"] == "vs_cached" + assert vector_store["team_id"] == "team-a" + cache_helper.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_get_managed_vector_store_fails_closed_on_lookup_error(): + from litellm.proxy.vector_store_endpoints.utils import ( + get_litellm_managed_vector_store, + ) + + mock_registry = MagicMock() + mock_registry.get_litellm_managed_vector_store_from_registry.side_effect = ( + RuntimeError("registry unavailable") + ) + + with patch.object(litellm, "vector_store_registry", mock_registry): + with pytest.raises(HTTPException) as exc_info: + await get_litellm_managed_vector_store(vector_store_id="vs_registry_only") + + assert exc_info.value.status_code == 500 + + +@pytest.mark.asyncio +async def test_vertex_discovery_allows_unregistered_provider_native_datastore_id(): + from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import ( + vertex_discovery_proxy_route, + ) + + request = _mock_request() + request.method = "GET" + + with ( + patch( + "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.get_litellm_managed_vector_store", + new=AsyncMock(return_value=None), + ) as mock_lookup, + patch( + "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints._base_vertex_proxy_route", + new=AsyncMock(return_value={"ok": True}), + ) as mock_base_route, + ): + response = await vertex_discovery_proxy_route( + endpoint="projects/p/locations/us-central1/dataStores/vs_unknown", + request=request, + fastapi_response=Response(), + ) + + assert response == {"ok": True} + mock_lookup.assert_awaited_once_with(vector_store_id="vs_unknown") + assert mock_base_route.call_args.kwargs["router_credentials"] is None + + +@pytest.mark.asyncio +async def test_milvus_passthrough_denies_other_team_vector_store_index(): + from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import ( + milvus_proxy_route, + ) + + request = _mock_request() + request.url.path = "/milvus/v2/vectordb/entities/search" + + index_object = MagicMock() + index_object.litellm_params.vector_store_name = "tenant-b-store" + index_object.litellm_params.vector_store_index = "tenant_b_collection" + + mock_index_registry = MagicMock() + mock_index_registry.is_vector_store_index.return_value = True + mock_index_registry.get_vector_store_index_by_name.return_value = index_object + + mock_vector_registry = MagicMock() + mock_vector_registry.get_litellm_managed_vector_store_from_registry_by_name.return_value = { + "vector_store_id": "vs_other_team", + "custom_llm_provider": "milvus", + "team_id": "team-b", + "litellm_params": {"api_base": "https://milvus.example.com"}, + } + + with ( + patch( + "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.ProviderConfigManager.get_provider_vector_stores_config", + return_value=MagicMock(), + ), + patch( + "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.get_request_body", + new=AsyncMock(return_value={"collectionName": "managed_index"}), + ), + patch( + "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.is_allowed_to_call_vector_store_endpoint", + return_value=True, + ), + patch.object(litellm, "vector_store_index_registry", mock_index_registry), + patch.object(litellm, "vector_store_registry", mock_vector_registry), + ): + with pytest.raises(HTTPException) as exc_info: + await milvus_proxy_route( + endpoint="v2/vectordb/entities/search", + request=request, + fastapi_response=Response(), + user_api_key_dict=UserAPIKeyAuth(team_id="team-a"), + ) + + assert exc_info.value.status_code == 403 + + +@pytest.mark.asyncio +async def test_azure_passthrough_denies_other_team_vector_store_index(): + from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import ( + azure_proxy_route, + ) + + request = _mock_request() + request.url.path = "/azure/indexes/managed_index/docs/search" + + index_object = MagicMock() + index_object.litellm_params.vector_store_name = "tenant-b-store" + + mock_index_registry = MagicMock() + mock_index_registry.is_vector_store_index.side_effect = ( + lambda vector_store_index_name: vector_store_index_name == "managed_index" + ) + mock_index_registry.get_vector_store_index_by_name.return_value = index_object + + mock_vector_registry = MagicMock() + mock_vector_registry.get_litellm_managed_vector_store_from_registry_by_name.return_value = { + "vector_store_id": "vs_other_team", + "custom_llm_provider": "azure_ai", + "team_id": "team-b", + "litellm_params": {"api_base": "https://azure.example.com"}, + } + + with ( + patch("litellm.proxy.proxy_server.llm_router", MagicMock()), + patch( + "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.is_passthrough_request_using_router_model", + return_value=False, + ), + patch( + "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.ProviderConfigManager.get_provider_vector_stores_config", + return_value=MagicMock(), + ), + patch( + "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.is_allowed_to_call_vector_store_endpoint", + return_value=True, + ), + patch.object(litellm, "vector_store_index_registry", mock_index_registry), + patch.object(litellm, "vector_store_registry", mock_vector_registry), + ): + with pytest.raises(HTTPException) as exc_info: + await azure_proxy_route( + endpoint="indexes/managed_index/docs/search", + request=request, + fastapi_response=Response(), + user_api_key_dict=UserAPIKeyAuth(team_id="team-a"), + ) + + assert exc_info.value.status_code == 403 diff --git a/tests/test_litellm/router_utils/test_router_utils_common_utils.py b/tests/test_litellm/router_utils/test_router_utils_common_utils.py index 465c6669ceb..efa5f2382dc 100644 --- a/tests/test_litellm/router_utils/test_router_utils_common_utils.py +++ b/tests/test_litellm/router_utils/test_router_utils_common_utils.py @@ -4,6 +4,7 @@ from unittest.mock import Mock import pytest from litellm import Router +from litellm.proxy._types import UserAPIKeyAuth from litellm.router_utils.common_utils import ( _deployment_supports_web_search, add_model_file_id_mappings, @@ -365,6 +366,49 @@ def test_invalidate_model_group_info_cache(): assert router._cached_get_model_group_info.cache_info().currsize == 0 +def test_filter_deployments_by_model_access_groups_access_group_only_key(): + """ + Access-group-only keys should only route to deployments in allowed groups, + even when multiple deployments share the same public model name. + """ + router = Router( + model_list=[ + { + "model_name": "gpt-5", + "litellm_params": {"model": "openai/gpt-5.1", "api_key": "key-1"}, + "model_info": {"access_groups": ["AG1"]}, + }, + { + "model_name": "gpt-5", + "litellm_params": {"model": "openai/gpt-4o", "api_key": "key-2"}, + "model_info": {"access_groups": ["AG2"]}, + }, + ] + ) + + scoped_key = UserAPIKeyAuth( + api_key="hashed-key", + team_id="team-2", + models=["AG2"], + team_models=["AG2"], + ) + + filtered = router._filter_deployments_by_model_access_groups( + model="gpt-5", + healthy_deployments=router._get_all_deployments(model_name="gpt-5"), + request_kwargs={ + "metadata": { + "user_api_key_team_id": "team-2", + "user_api_key_auth": scoped_key, + } + }, + request_team_id="team-2", + ) + + assert len(filtered) == 1 + assert filtered[0].get("model_info", {}).get("access_groups") == ["AG2"] + + class TestAddModelFileIdMappings: """Test cases for add_model_file_id_mappings. diff --git a/tests/test_litellm/test_openai_embedding_encoding_format_default.py b/tests/test_litellm/test_openai_embedding_encoding_format_default.py new file mode 100644 index 00000000000..94e4e3c81e5 --- /dev/null +++ b/tests/test_litellm/test_openai_embedding_encoding_format_default.py @@ -0,0 +1,124 @@ +from unittest.mock import MagicMock, patch + +import pytest + +from litellm import embedding + + +@pytest.mark.parametrize( + "set_env, env_value, expected", + [ + (False, None, "float"), + (True, "base64", "base64"), + ], +) +def test_openai_embedding_encoding_format_default( + monkeypatch, set_env, env_value, expected +): + monkeypatch.delenv("LITELLM_DEFAULT_EMBEDDING_ENCODING_FORMAT", raising=False) + if set_env: + monkeypatch.setenv("LITELLM_DEFAULT_EMBEDDING_ENCODING_FORMAT", env_value) + + mock_response = MagicMock() + mock_response.parse.return_value = MagicMock( + model_dump=lambda: { + "data": [{"embedding": [0.1, 0.2, 0.3], "index": 0}], + "model": "text-embedding-ada-002", + "object": "list", + "usage": {"prompt_tokens": 1, "total_tokens": 1}, + } + ) + mock_response.headers = {} + + with patch( + "litellm.llms.openai.openai.OpenAIChatCompletion._get_openai_client" + ) as mock_get_client: + mock_client_instance = MagicMock() + mock_get_client.return_value = mock_client_instance + mock_client_instance.embeddings.with_raw_response.create.return_value = ( + mock_response + ) + + embedding( + model="text-embedding-ada-002", + input="Hello world", + ) + + call_kwargs = ( + mock_client_instance.embeddings.with_raw_response.create.call_args[1] + ) + assert call_kwargs["encoding_format"] == expected + + +@pytest.mark.parametrize("env_none", ["none", "NONE", " none "]) +def test_openai_embedding_encoding_format_env_none_omits_param( + monkeypatch, env_none +): + """LITELLM_DEFAULT_EMBEDDING_ENCODING_FORMAT=none omits encoding_format (provider default).""" + monkeypatch.setenv("LITELLM_DEFAULT_EMBEDDING_ENCODING_FORMAT", env_none) + + mock_response = MagicMock() + mock_response.parse.return_value = MagicMock( + model_dump=lambda: { + "data": [{"embedding": [0.1, 0.2, 0.3], "index": 0}], + "model": "text-embedding-ada-002", + "object": "list", + "usage": {"prompt_tokens": 1, "total_tokens": 1}, + } + ) + mock_response.headers = {} + + with patch( + "litellm.llms.openai.openai.OpenAIChatCompletion._get_openai_client" + ) as mock_get_client: + mock_client_instance = MagicMock() + mock_get_client.return_value = mock_client_instance + mock_client_instance.embeddings.with_raw_response.create.return_value = ( + mock_response + ) + + embedding( + model="text-embedding-ada-002", + input="Hello world", + ) + + call_kwargs = ( + mock_client_instance.embeddings.with_raw_response.create.call_args[1] + ) + assert "encoding_format" not in call_kwargs + + +def test_openai_embedding_encoding_format_explicit_overrides_env(monkeypatch): + """Request `encoding_format` wins over LITELLM_DEFAULT_EMBEDDING_ENCODING_FORMAT.""" + monkeypatch.setenv("LITELLM_DEFAULT_EMBEDDING_ENCODING_FORMAT", "float") + + mock_response = MagicMock() + mock_response.parse.return_value = MagicMock( + model_dump=lambda: { + "data": [{"embedding": [0.1, 0.2, 0.3], "index": 0}], + "model": "text-embedding-ada-002", + "object": "list", + "usage": {"prompt_tokens": 1, "total_tokens": 1}, + } + ) + mock_response.headers = {} + + with patch( + "litellm.llms.openai.openai.OpenAIChatCompletion._get_openai_client" + ) as mock_get_client: + mock_client_instance = MagicMock() + mock_get_client.return_value = mock_client_instance + mock_client_instance.embeddings.with_raw_response.create.return_value = ( + mock_response + ) + + embedding( + model="text-embedding-ada-002", + input="Hello world", + encoding_format="base64", + ) + + call_kwargs = ( + mock_client_instance.embeddings.with_raw_response.create.call_args[1] + ) + assert call_kwargs["encoding_format"] == "base64" diff --git a/tests/test_litellm/test_router.py b/tests/test_litellm/test_router.py index 4df8003338c..48facace528 100644 --- a/tests/test_litellm/test_router.py +++ b/tests/test_litellm/test_router.py @@ -3267,3 +3267,433 @@ async def test_multiregion_team_failover_between_regions(): "response from us-east-1", "response from us-west-2", ] + + +def test_access_group_scoped_key_filters_deployments_with_same_public_model(): + """ + If a key can access a model only via access group membership, + router candidate deployments for that public model should be constrained + to deployments in the allowed access group. + """ + from litellm.proxy._types import UserAPIKeyAuth + + router = litellm.Router( + model_list=[ + { + "model_name": "gpt-5", + "litellm_params": { + "model": "openai/gpt-5.1", + "api_key": "key1", + "mock_response": "response-via-AG1", + }, + "model_info": {"access_groups": ["AG1"]}, + }, + { + "model_name": "gpt-5", + "litellm_params": { + "model": "openai/gpt-4o", + "api_key": "key2", + "mock_response": "response-via-AG2", + }, + "model_info": {"access_groups": ["AG2"]}, + }, + ] + ) + + scoped_key = UserAPIKeyAuth( + api_key="hashed-key", + team_id="team2", + models=["AG2"], + team_models=["AG2"], + ) + + _model, deployments = router._common_checks_available_deployment( + model="gpt-5", + request_kwargs={ + "metadata": { + "user_api_key_team_id": "team2", + "user_api_key_auth": scoped_key, + } + }, + ) + + assert len(deployments) == 1 + assert deployments[0].get("model_info", {}).get("access_groups") == ["AG2"] + + seen = set() + for _ in range(20): + response = router.completion( + model="gpt-5", + messages=[{"role": "user", "content": "hello"}], + metadata={"user_api_key_team_id": "team2", "user_api_key_auth": scoped_key}, + ) + seen.add(response.choices[0].message.content) + + assert seen == {"response-via-AG2"} + + +def test_explicit_model_access_does_not_force_access_group_filtering(): + """ + If a key has explicit model access in addition to access group entries, + do not force access-group-only filtering for deployment selection. + """ + from litellm.proxy._types import UserAPIKeyAuth + + router = litellm.Router( + model_list=[ + { + "model_name": "gpt-5", + "litellm_params": { + "model": "openai/gpt-5.1", + "api_key": "key1", + "mock_response": "response-via-AG1", + }, + "model_info": {"access_groups": ["AG1"]}, + }, + { + "model_name": "gpt-5", + "litellm_params": { + "model": "openai/gpt-4o", + "api_key": "key2", + "mock_response": "response-via-AG2", + }, + "model_info": {"access_groups": ["AG2"]}, + }, + ] + ) + + explicit_key = UserAPIKeyAuth( + api_key="hashed-key", + team_id="team2", + models=["AG2", "gpt-5"], + team_models=["AG2", "gpt-5"], + ) + + _model, deployments = router._common_checks_available_deployment( + model="gpt-5", + request_kwargs={ + "metadata": { + "user_api_key_team_id": "team2", + "user_api_key_auth": explicit_key, + } + }, + ) + + deployment_groups = [ + d.get("model_info", {}).get("access_groups") for d in deployments + ] + assert ["AG1"] in deployment_groups + assert ["AG2"] in deployment_groups + + +def test_access_group_filter_empty_does_not_bypass_via_litellm_model_fallback( + monkeypatch: pytest.MonkeyPatch, +): + """ + When access-group filtering removes all candidates, _get_deployment_by_litellm_model + must not run: it does not re-apply access groups and could return blocked deployments + that share the same litellm_params.model as the request model string. + + ``get_model_access_groups`` is patched to expose AG1 for the public model (so the + access-group filter runs with a non-empty allowed set) while every deployment + returned for that name is AG2-only — filtered to empty. Without the guard, the + litellm-model fallback would return both rows because ``litellm_params.model`` matches. + """ + from litellm.proxy._types import UserAPIKeyAuth + + router = litellm.Router( + model_list=[ + { + "model_name": "gpt-5", + "litellm_params": { + "model": "gpt-5", + "api_key": "key1", + "mock_response": "blocked-dep-1", + }, + "model_info": {"access_groups": ["AG2"]}, + }, + { + "model_name": "gpt-5", + "litellm_params": { + "model": "gpt-5", + "api_key": "key2", + "mock_response": "blocked-dep-2", + }, + "model_info": {"access_groups": ["AG2"]}, + }, + ] + ) + + orig_groups = router.get_model_access_groups + + def fake_get_model_access_groups( + model_name=None, model_access_group=None, team_id=None + ): + if model_name == "gpt-5" and model_access_group is None: + return {"AG1": ["gpt-5"], "AG2": ["gpt-5"]} + return orig_groups( + model_name=model_name, + model_access_group=model_access_group, + team_id=team_id, + ) + + monkeypatch.setattr(router, "get_model_access_groups", fake_get_model_access_groups) + + scoped_key = UserAPIKeyAuth( + api_key="hashed-key", + team_id="team2", + models=["AG1"], + team_models=["AG1"], + ) + + with pytest.raises(litellm.BadRequestError): + router._common_checks_available_deployment( + model="gpt-5", + request_kwargs={ + "metadata": { + "user_api_key_team_id": "team2", + "user_api_key_auth": scoped_key, + } + }, + ) + + +def test_access_group_block_does_not_silently_use_default_fallback_model( + monkeypatch: pytest.MonkeyPatch, +): + """ + When access-group filtering empties candidates for model X, the router must not use + ``fallbacks`` default ``*`` routing to model Y: Y may have no ``access_groups``, so + ``_filter_deployments_by_model_access_groups`` would not constrain Y and the caller + would be served despite being blocked from X. + """ + from litellm.proxy._types import UserAPIKeyAuth + + router = litellm.Router( + model_list=[ + { + "model_name": "gpt-5", + "litellm_params": { + "model": "gpt-5", + "api_key": "key1", + "mock_response": "blocked-dep-1", + }, + "model_info": {"access_groups": ["AG2"]}, + }, + { + "model_name": "gpt-5", + "litellm_params": { + "model": "gpt-5", + "api_key": "key2", + "mock_response": "blocked-dep-2", + }, + "model_info": {"access_groups": ["AG2"]}, + }, + { + "model_name": "gpt-4-fallback", + "litellm_params": { + "model": "gpt-4", + "api_key": "fallback-key", + "mock_response": "should-not-reach", + }, + }, + ], + fallbacks=[{"*": ["gpt-4-fallback"]}], + ) + + orig_groups = router.get_model_access_groups + + def fake_get_model_access_groups( + model_name=None, model_access_group=None, team_id=None + ): + if model_name == "gpt-5" and model_access_group is None: + return {"AG1": ["gpt-5"], "AG2": ["gpt-5"]} + return orig_groups( + model_name=model_name, + model_access_group=model_access_group, + team_id=team_id, + ) + + monkeypatch.setattr(router, "get_model_access_groups", fake_get_model_access_groups) + + scoped_key = UserAPIKeyAuth( + api_key="hashed-key", + team_id="team2", + models=["AG1"], + team_models=["AG1"], + ) + + with pytest.raises(litellm.BadRequestError): + router._common_checks_available_deployment( + model="gpt-5", + request_kwargs={ + "metadata": { + "user_api_key_team_id": "team2", + "user_api_key_auth": scoped_key, + } + }, + ) + + +def test_access_group_block_via_litellm_model_branch_does_not_use_default_fallback( + monkeypatch: pytest.MonkeyPatch, +): + """ + When the by-name lookup returns no deployments and the litellm-model fallback + branch finds candidates that access-group filtering then empties, the router + must not fall through to default ``fallbacks`` routing — the default fallback + model may have no ``access_groups`` and would short-circuit the filter, + silently serving a caller blocked by access-group restrictions. + """ + from litellm.proxy._types import UserAPIKeyAuth + + router = litellm.Router( + model_list=[ + { + "model_name": "gpt-5-alias", + "litellm_params": { + "model": "gpt-5", + "api_key": "key1", + "mock_response": "blocked-dep-1", + }, + "model_info": {"access_groups": ["AG2"]}, + }, + { + "model_name": "gpt-4-fallback", + "litellm_params": { + "model": "gpt-4", + "api_key": "fallback-key", + "mock_response": "should-not-reach", + }, + }, + ], + fallbacks=[{"*": ["gpt-4-fallback"]}], + ) + + orig_groups = router.get_model_access_groups + + def fake_get_model_access_groups( + model_name=None, model_access_group=None, team_id=None + ): + if model_name == "gpt-5" and model_access_group is None: + return {"AG1": ["gpt-5"], "AG2": ["gpt-5"]} + return orig_groups( + model_name=model_name, + model_access_group=model_access_group, + team_id=team_id, + ) + + monkeypatch.setattr(router, "get_model_access_groups", fake_get_model_access_groups) + + scoped_key = UserAPIKeyAuth( + api_key="hashed-key", + team_id="team2", + models=["AG1"], + team_models=["AG1"], + ) + + with pytest.raises(litellm.BadRequestError): + router._common_checks_available_deployment( + model="gpt-5", + request_kwargs={ + "metadata": { + "user_api_key_team_id": "team2", + "user_api_key_auth": scoped_key, + } + }, + ) + + +def test_try_early_resolve_deployments_for_model_not_in_names(): + """ + Direct coverage for ``_try_early_resolve_deployments_for_model_not_in_names``: + + - Returns ``None`` when the requested model is already in ``self.model_names`` + (the by-name lookup path will handle it). + - Returns ``None`` when there are no team deployments, no pattern matches, and + no default deployment to fall back to. + - Returns the pattern-router match when the model matches a wildcard route. + - Returns the default deployment with the request model substituted in when one + is configured, without mutating the stored default. + """ + router_in_names = litellm.Router( + model_list=[ + { + "model_name": "gpt-5", + "litellm_params": { + "model": "openai/gpt-5", + "api_key": "key1", + }, + }, + ] + ) + + assert ( + router_in_names._try_early_resolve_deployments_for_model_not_in_names( + model="gpt-5", request_team_id=None + ) + is None + ) + assert ( + router_in_names._try_early_resolve_deployments_for_model_not_in_names( + model="some-unknown-model", request_team_id=None + ) + is None + ) + + pattern_router = litellm.Router( + model_list=[ + { + "model_name": "openai/*", + "litellm_params": { + "model": "openai/*", + "api_key": "key-pattern", + }, + }, + ] + ) + + pattern_result = ( + pattern_router._try_early_resolve_deployments_for_model_not_in_names( + model="openai/gpt-4o-mini", request_team_id=None + ) + ) + assert pattern_result is not None + resolved_model, pattern_deployments = pattern_result + assert resolved_model == "openai/gpt-4o-mini" + assert isinstance(pattern_deployments, list) and len(pattern_deployments) == 1 + + default_router = litellm.Router( + model_list=[ + { + "model_name": "named-model", + "litellm_params": { + "model": "openai/gpt-4o", + "api_key": "key-named", + }, + }, + ] + ) + default_router.default_deployment = { + "model_name": "default", + "litellm_params": { + "model": "openai/will-be-overridden", + "api_key": "key-default", + }, + } + + default_result = ( + default_router._try_early_resolve_deployments_for_model_not_in_names( + model="brand-new-model", request_team_id=None + ) + ) + assert default_result is not None + resolved_model, default_deployment = default_result + assert resolved_model == "brand-new-model" + assert isinstance(default_deployment, dict) + assert default_deployment["litellm_params"]["model"] == "brand-new-model" + # The original default_deployment must not be mutated. + assert ( + default_router.default_deployment["litellm_params"]["model"] + == "openai/will-be-overridden" + ) diff --git a/tests/test_litellm/test_secret_redaction.py b/tests/test_litellm/test_secret_redaction.py index 6cbb2fd7b8f..8a0a2221c11 100644 --- a/tests/test_litellm/test_secret_redaction.py +++ b/tests/test_litellm/test_secret_redaction.py @@ -9,11 +9,11 @@ from litellm._logging import ( JsonFormatter, _redact_string, _secret_filter, - _setup_json_exception_handlers, verbose_logger, verbose_proxy_logger, verbose_router_logger, ) +from litellm.litellm_core_utils.secret_redaction import redact_string SECRET = "sk-proj-abc123def456ghi789jklmnopqrst" @@ -57,12 +57,12 @@ def test_redact_string_catches_secret_patterns(): SECRET, ] for secret in cases: - result = _redact_string("msg: " + secret) + result = redact_string("msg: " + secret) assert secret not in result, f"{secret!r} was not redacted" assert "REDACTED" in result normal = "Loaded model gpt-4 with 3 replicas on us-east-1" - assert _redact_string(normal) == normal + assert redact_string(normal) == normal def test_filter_redacts_secrets_in_logger_output(): @@ -155,7 +155,7 @@ def test_x_api_key_regex_does_not_consume_json_delimiters(): """x-api-key pattern must stop before closing quotes/braces so JSON stays valid.""" # Simulates a JSON log line containing an x-api-key header value json_line = '{"headers": {"x-api-key": "secret123"}, "status": 200}' - result = _redact_string(json_line) + result = redact_string(json_line) # The secret value should be redacted assert "secret123" not in result assert "REDACTED" in result @@ -234,12 +234,12 @@ def test_key_name_redaction_catches_secrets_in_dict_repr(): "'slack_webhook_url': 'https://hooks.slack.com/services/T00/B00/xxx'", ] for secret_line in cases: - result = _redact_string(secret_line) + result = redact_string(secret_line) assert "REDACTED" in result, f"Key-name redaction missed: {secret_line!r}" # Non-sensitive keys should NOT be redacted safe = "'enable_jwt_auth': True, 'store_model_in_db': True" - assert _redact_string(safe) == safe + assert redact_string(safe) == safe def test_key_name_redaction_in_general_settings_dict(): @@ -277,7 +277,7 @@ _SAMPLE_SA_JSON = ( def test_pem_private_key_redacted_in_json(): - result = _redact_string(_SAMPLE_SA_JSON) + result = redact_string(_SAMPLE_SA_JSON) assert "MIIEvQIBADA" not in result assert "-----BEGIN" not in result @@ -286,12 +286,12 @@ def test_pem_private_key_redacted_in_dict_repr(): import json sa = json.loads(_SAMPLE_SA_JSON) - result = _redact_string(str(sa)) + result = redact_string(str(sa)) assert "MIIEvQIBADA" not in result def test_service_account_blob_fully_redacted(): - result = _redact_string(f"Got={_SAMPLE_SA_JSON}") + result = redact_string(f"Got={_SAMPLE_SA_JSON}") assert "my-proj-123" not in result assert "sa@my-proj.iam.gserviceaccount.com" not in result assert "abc123def" not in result @@ -320,22 +320,22 @@ def test_vertex_traceback_redacts_pem(): "Unable to load vertex credentials from environment. " f"Got={_SAMPLE_SA_JSON}" ) - result = _redact_string(traceback_text) + result = redact_string(traceback_text) assert "MIIEvQIBADA" not in result assert "-----BEGIN" not in result def test_gcp_oauth_token_redacted(): - result = _redact_string("access token ya29.c.c0ASRK0GZvXlongtokenhere") + result = redact_string("access token ya29.c.c0ASRK0GZvXlongtokenhere") assert "ya29." not in result assert "REDACTED" in result def test_non_pem_private_key_value_redacted(): - result = _redact_string("'private_key': 'some-non-pem-secret-value'") + result = redact_string("'private_key': 'some-non-pem-secret-value'") assert "some-non-pem-secret" not in result def test_normal_vertex_log_not_redacted(): msg = "Vertex: Loading vertex credentials, is_file_path=True, current dir /app" - assert _redact_string(msg) == msg + assert redact_string(msg) == msg diff --git a/ui/litellm-dashboard/public/assets/logos/qohash.jpg b/ui/litellm-dashboard/public/assets/logos/qohash.jpg new file mode 100644 index 00000000000..50227ab3910 Binary files /dev/null and b/ui/litellm-dashboard/public/assets/logos/qohash.jpg differ diff --git a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/ModelsAndEndpointsView.tsx b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/ModelsAndEndpointsView.tsx index 7c162e2056b..944c56833e5 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/ModelsAndEndpointsView.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/models-and-endpoints/ModelsAndEndpointsView.tsx @@ -17,7 +17,7 @@ import { RefreshIcon } from "@heroicons/react/outline"; import { useQueryClient } from "@tanstack/react-query"; import { Col, Grid, Icon, Tab, TabGroup, TabList, TabPanel, TabPanels } from "@tremor/react"; import type { UploadProps } from "antd"; -import { Form, Typography } from "antd"; +import { Form } from "antd"; import { PlusCircleOutlined } from "@ant-design/icons"; import React, { useEffect, useMemo, useState } from "react"; import AddModelTab from "../../../components/add_model/add_model_tab"; @@ -251,15 +251,9 @@ const ModelsAndEndpointsView: React.FC = ({ premiumUser, te const isLoading = isLoadingModels || isLoadingModelCostMap || isLoadingCredentials || isLoadingUISettings; - if (userRole && userRole == "Admin Viewer") { - const { Title, Paragraph } = Typography; - return ( -
- Access Denied - Ask your proxy admin for access to view all models -
- ); - } + // Admin Viewer can view all models read-only — page render proceeds; the + // individual write-action tabs (Add Model, LLM Credentials, etc.) are + // gated separately below. const handleOk = async () => { try { @@ -395,107 +389,154 @@ const ModelsAndEndpointsView: React.FC = ({ premiumUser, te modelAccessGroups={availableModelAccessGroups} /> ) : ( - - -
- {all_admin_roles.includes(userRole) ? All Models : Your Models} - {!shouldHideAddModelTab && Add Model} - {all_admin_roles.includes(userRole) && LLM Credentials} - {all_admin_roles.includes(userRole) && Pass-Through Endpoints} - {all_admin_roles.includes(userRole) && Health Status} - {all_admin_roles.includes(userRole) && Model Retry Settings} - {all_admin_roles.includes(userRole) && Model Group Alias} - {all_admin_roles.includes(userRole) && Price Data Reload} -
- -
- {lastRefreshed && Last Refreshed: {lastRefreshed}} - -
-
- - - {!shouldHideAddModelTab && ( - - { + // Build a single source-of-truth list of {tab, panel} pairs. + // Conditionally-hidden tabs (e.g. "Add Model" for non-admin) get + // filtered out as a unit so tab indices and panel indices can + // never drift apart — Tremor's TabList and TabPanels filter + // falsy children inconsistently, which previously caused + // "click LLM Credentials, see nothing" for Admin Viewer. + const isAdmin = all_admin_roles.includes(userRole); + const visibleTabs: Array<{ tab: React.ReactElement; panel: React.ReactElement }> = [ + { + tab: {isAdmin ? "All Models" : "Your Models"}, + panel: ( + - - )} - - - - - - - - - - - - - - - -
+ ), + }, + ]; + if (!shouldHideAddModelTab) { + visibleTabs.push({ + tab: Add Model, + panel: ( + + + + ), + }); + } + if (isAdmin) { + visibleTabs.push( + { + tab: LLM Credentials, + panel: ( + + + + ), + }, + { + tab: Pass-Through Endpoints, + panel: ( + + + + ), + }, + { + tab: Health Status, + panel: ( + + + + ), + }, + { + tab: Model Retry Settings, + panel: ( + + ), + }, + { + tab: Model Group Alias, + panel: ( + + + + ), + }, + { + tab: Price Data Reload, + panel: , + }, + ); + } + return ( + + +
{visibleTabs.map((t) => t.tab)}
+ +
+ {lastRefreshed && Last Refreshed: {lastRefreshed}} + +
+
+ {visibleTabs.map((t) => t.panel)} +
+ ); + })() )} diff --git a/ui/litellm-dashboard/src/app/login/LoginPage.test.tsx b/ui/litellm-dashboard/src/app/login/LoginPage.test.tsx index 866b7d0f172..cd58c51a862 100644 --- a/ui/litellm-dashboard/src/app/login/LoginPage.test.tsx +++ b/ui/litellm-dashboard/src/app/login/LoginPage.test.tsx @@ -289,4 +289,85 @@ describe("LoginPage", () => { expect(ssoButton).toBeInTheDocument(); expect(ssoButton).toBeDisabled(); }); + + describe("URL ?token= legacy path is rejected (security regression test)", () => { + const originalLocation = window.location; + + beforeEach(() => { + Object.defineProperty(window, "location", { + value: { + ...originalLocation, + href: "http://localhost:3000/ui/login?token=attacker.jwt.value", + pathname: "/ui/login", + search: "?token=attacker.jwt.value", + }, + writable: true, + }); + document.cookie = + "token=; expires=Thu, 01 Jan 1970 00:00:00 GMT; path=/; SameSite=Lax"; + }); + + afterEach(() => { + Object.defineProperty(window, "location", { + value: originalLocation, + writable: true, + }); + }); + + it("must not set a token cookie or redirect to /ui/?login=success when ?token= is in the URL", async () => { + (useUIConfig as ReturnType).mockReturnValue({ + data: { + auto_redirect_to_sso: false, + server_root_path: "/", + proxy_base_url: null, + sso_configured: false, + }, + isLoading: false, + }); + (getCookie as ReturnType).mockReturnValue(null); + (isJwtExpired as ReturnType).mockReturnValue(false); + + const queryClient = createQueryClient(); + render( + + + , + ); + + await waitFor(() => { + expect(screen.getByRole("heading", { name: "Login" })).toBeInTheDocument(); + }); + + expect(document.cookie).not.toContain("token=attacker.jwt.value"); + expect(mockReplace).not.toHaveBeenCalledWith("/ui/?login=success"); + }); + + it("must not overwrite an existing valid session cookie when ?token= is in the URL", async () => { + (useUIConfig as ReturnType).mockReturnValue({ + data: { + auto_redirect_to_sso: false, + server_root_path: "/", + proxy_base_url: null, + sso_configured: false, + }, + isLoading: false, + }); + (getCookie as ReturnType).mockReturnValue("legitimate-session-jwt"); + (isJwtExpired as ReturnType).mockReturnValue(false); + + const queryClient = createQueryClient(); + render( + + + , + ); + + await waitFor(() => { + expect(mockReplace).toHaveBeenCalledWith("/ui"); + }); + + expect(document.cookie).not.toContain("token=attacker.jwt.value"); + expect(mockReplace).not.toHaveBeenCalledWith("/ui/?login=success"); + }); + }); }); diff --git a/ui/litellm-dashboard/src/app/login/LoginPage.tsx b/ui/litellm-dashboard/src/app/login/LoginPage.tsx index 7ad3e32ef5c..74ee9f9de59 100644 --- a/ui/litellm-dashboard/src/app/login/LoginPage.tsx +++ b/ui/litellm-dashboard/src/app/login/LoginPage.tsx @@ -66,21 +66,6 @@ function LoginPageContent() { return; } - // Backwards compat: handle direct token in URL (legacy flow) - const urlToken = params.get("token"); - if (urlToken && !isJwtExpired(urlToken)) { - document.cookie = `token=${urlToken}; path=/; SameSite=Lax`; - params.delete("token"); - const cleanSearch = params.toString(); - window.history.replaceState( - null, - "", - window.location.pathname + (cleanSearch ? `?${cleanSearch}` : ""), - ); - router.replace("/ui/?login=success"); - return; - } - // If switching workers on a control plane, clear the old token and show login const switchingWorker = params.has("worker"); if (switchingWorker && uiConfig?.is_control_plane) { diff --git a/ui/litellm-dashboard/src/app/page.tsx b/ui/litellm-dashboard/src/app/page.tsx index a8553d5405b..06bf3b68d05 100644 --- a/ui/litellm-dashboard/src/app/page.tsx +++ b/ui/litellm-dashboard/src/app/page.tsx @@ -325,9 +325,6 @@ function CreateKeyPageContent() { if (decoded.user_role) { const formattedUserRole = formatUserRole(decoded.user_role); setUserRole(formattedUserRole); - if (formattedUserRole == "Admin Viewer") { - setPage("usage"); - } } if (decoded.user_email) { diff --git a/ui/litellm-dashboard/src/components/AIHub/ModelHubTable.tsx b/ui/litellm-dashboard/src/components/AIHub/ModelHubTable.tsx index 39695c1348b..75058157a65 100644 --- a/ui/litellm-dashboard/src/components/AIHub/ModelHubTable.tsx +++ b/ui/litellm-dashboard/src/components/AIHub/ModelHubTable.tsx @@ -22,7 +22,7 @@ import { modelHubPublicModelsCall, } from "@/components/networking"; import PublicModelHub from "@/components/public_model_hub"; -import { isAdminRole } from "@/utils/roles"; +import { isAdminRole, isProxyAdminRole } from "@/utils/roles"; import { CopyOutlined } from "@ant-design/icons"; import { Badge, Button, Card, Tab, TabGroup, TabList, TabPanel, TabPanels, Text, Title } from "@tremor/react"; import { Modal } from "antd"; @@ -61,6 +61,10 @@ interface ModelGroupInfo { } const ModelHubTable: React.FC = ({ accessToken, publicPage, premiumUser, userRole }) => { + // Admin Viewer follows the read-parity rule: see the AI Hub catalog, but + // cannot toggle public visibility (write). + const canModify = isProxyAdminRole(userRole || ""); + const [publicPageAllowed, setPublicPageAllowed] = useState(false); const [modelHubData, setModelHubData] = useState(null); const [loading, setLoading] = useState(true); @@ -420,7 +424,7 @@ const ModelHubTable: React.FC = ({ accessToken, publicPage, {/* Useful Links Management Section for Admins */} - {isAdminRole(userRole || "") && ( + {canModify && (
@@ -441,7 +445,7 @@ const ModelHubTable: React.FC = ({ accessToken, publicPage, {/* Model Filters and Table */} {/* Header with Make Public Button */} - {publicPage == false && isAdminRole(userRole || "") && ( + {publicPage == false && canModify && (
@@ -470,7 +474,7 @@ const ModelHubTable: React.FC = ({ accessToken, publicPage, {/* Header with Make Public Button */} - {publicPage == false && isAdminRole(userRole || "") && ( + {publicPage == false && canModify && (
@@ -496,7 +500,7 @@ const ModelHubTable: React.FC = ({ accessToken, publicPage, {/* Header with Make Public Button */} - {publicPage == false && isAdminRole(userRole || "") && ( + {publicPage == false && canModify && (
@@ -520,7 +524,7 @@ const ModelHubTable: React.FC = ({ accessToken, publicPage, {/* Skill Hub Tab */} - {publicPage == false && isAdminRole(userRole || "") && ( + {publicPage == false && canModify && (
+ {canModify && ( + + )} diff --git a/ui/litellm-dashboard/src/components/Settings/RouterSettings/Fallbacks/Fallbacks.tsx b/ui/litellm-dashboard/src/components/Settings/RouterSettings/Fallbacks/Fallbacks.tsx index 9c5933aba3a..b4ab7ddddba 100644 --- a/ui/litellm-dashboard/src/components/Settings/RouterSettings/Fallbacks/Fallbacks.tsx +++ b/ui/litellm-dashboard/src/components/Settings/RouterSettings/Fallbacks/Fallbacks.tsx @@ -8,6 +8,7 @@ import DeleteResourceModal from "../../../common_components/DeleteResourceModal" import { ProviderLogo } from "../../../molecules/models/ProviderLogo"; import NotificationsManager from "../../../molecules/notifications_manager"; import { getCallbacksCall, setCallbacksCall } from "../../../networking"; +import { isProxyAdminRole } from "@/utils/roles"; import AddFallbacks from "./AddFallbacks"; type FallbackEntry = { [modelName: string]: string[] }; @@ -243,15 +244,19 @@ const Fallbacks: React.FC = ({ accessToken, userRole, userID, mo }; const hasFallbacks = Array.isArray(routerSettings.fallbacks) && routerSettings.fallbacks.length > 0; + // Admin Viewer follows the read-parity rule: see fallbacks, no writes. + const canModify = isProxyAdminRole(userRole ?? ""); return ( <> - data.model_name) : []} - accessToken={accessToken || ""} - value={routerSettings.fallbacks || []} - onChange={handleFallbacksChange} - /> + {canModify && ( + data.model_name) : []} + accessToken={accessToken || ""} + value={routerSettings.fallbacks || []} + onChange={handleFallbacksChange} + /> + )} {!hasFallbacks ? (
@@ -280,30 +285,34 @@ const Fallbacks: React.FC = ({ accessToken, userRole, userID, mo {renderFallbacksChain(key, Array.isArray(value) ? value : [], getProviderFromModel)} - - testFallbackModelResponse(Object.keys(item)[0], accessToken || "")} - className="cursor-pointer hover:text-blue-600" - /> - - - handleDeleteClick(item)} - onKeyDown={(e) => e.key === "Enter" && handleDeleteClick(item)} - className="cursor-pointer inline-flex" - > - - - + {canModify && ( + <> + + testFallbackModelResponse(Object.keys(item)[0], accessToken || "")} + className="cursor-pointer hover:text-blue-600" + /> + + + handleDeleteClick(item)} + onKeyDown={(e) => e.key === "Enter" && handleDeleteClick(item)} + className="cursor-pointer inline-flex" + > + + + + + )} )), diff --git a/ui/litellm-dashboard/src/components/budgets/budget_panel.tsx b/ui/litellm-dashboard/src/components/budgets/budget_panel.tsx index e42d0569652..d90737b130e 100644 --- a/ui/litellm-dashboard/src/components/budgets/budget_panel.tsx +++ b/ui/litellm-dashboard/src/components/budgets/budget_panel.tsx @@ -28,6 +28,8 @@ import { useBudgets, useDeleteBudget } from "@/app/(dashboard)/hooks/budgets/use import BudgetModal from "./budget_modal"; import EditBudgetModal from "./edit_budget_modal"; import { CREATE_END_USER_CURL_COMMAND, CHAT_COMPLETIONS_CURL_COMMAND, OPENAI_SDK_PYTHON_CODE } from "./constants"; +import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized"; +import { isProxyAdminRole } from "@/utils/roles"; interface BudgetSettingsPageProps { accessToken: string | null; @@ -47,6 +49,10 @@ const BudgetPanel: React.FC = ({ accessToken }) => { const [selectedBudget, setSelectedBudget] = useState(null); const [isDeleteModalVisible, setIsDeleteModalVisible] = useState(false); + const { userRole } = useAuthorized(); + // Admin Viewer follows the read-parity rule: see budgets, no writes. + const canModify = isProxyAdminRole(userRole ?? ""); + const { data: budgetList = [] } = useBudgets(); const deleteBudget = useDeleteBudget(); @@ -89,9 +95,11 @@ const BudgetPanel: React.FC = ({ accessToken }) => { return (
- + {canModify && ( + + )} Budgets @@ -133,18 +141,22 @@ const BudgetPanel: React.FC = ({ accessToken }) => { {value.max_budget ? value.max_budget : "n/a"} {value.tpm_limit ? value.tpm_limit : "n/a"} {value.rpm_limit ? value.rpm_limit : "n/a"} - handleEditCall(value)} - dataTestId="edit-budget-button" - /> - handleDeleteClick(value)} - dataTestId="delete-budget-button" - /> + {canModify && ( + <> + handleEditCall(value)} + dataTestId="edit-budget-button" + /> + handleDeleteClick(value)} + dataTestId="delete-budget-button" + /> + + )} ))} diff --git a/ui/litellm-dashboard/src/components/guardrails/guardrail_info_helpers.tsx b/ui/litellm-dashboard/src/components/guardrails/guardrail_info_helpers.tsx index 2286eba7768..ac4b787e96a 100644 --- a/ui/litellm-dashboard/src/components/guardrails/guardrail_info_helpers.tsx +++ b/ui/litellm-dashboard/src/components/guardrails/guardrail_info_helpers.tsx @@ -52,6 +52,7 @@ export const guardrail_provider_map: Record = { Promptguard: "promptguard", LlmAsAJudge: "llm_as_a_judge", Xecguard: "xecguard", + QostodianNexus: "qostodian_nexus", }; // Function to populate provider map from API response - updates the original map @@ -138,6 +139,7 @@ export const guardrailLogoMap: Record = { "LiteLLM Content Filter": `${asset_logos_folder}litellm_logo.jpg`, "LiteLLM LLM as a Judge": `${asset_logos_folder}litellm_logo.jpg`, "Akto": `${asset_logos_folder}akto.svg`, + "Qostodian Nexus": `${asset_logos_folder}qohash.jpg`, }; export const getGuardrailLogoAndName = (guardrailValue: string): { logo: string; displayName: string } => { diff --git a/ui/litellm-dashboard/src/components/leftnav.test.tsx b/ui/litellm-dashboard/src/components/leftnav.test.tsx index 08b66c382b8..0e88a1a3603 100644 --- a/ui/litellm-dashboard/src/components/leftnav.test.tsx +++ b/ui/litellm-dashboard/src/components/leftnav.test.tsx @@ -5,10 +5,11 @@ import Sidebar from "./leftnav"; vi.mock("../utils/roles", () => { return { - all_admin_roles: ["admin"], + all_admin_roles: ["admin", "admin_viewer"], internalUserRoles: ["internal"], rolesWithWriteAccess: ["admin", "internal"], - isAdminRole: (role: string) => role === "admin", + rolesAllowedToViewWriteScopedPages: ["admin", "internal", "admin_viewer"], + isAdminRole: (role: string) => role === "admin" || role === "admin_viewer", isUserTeamAdminForAnyTeam: () => false, }; }); @@ -135,6 +136,53 @@ describe("Sidebar (leftnav)", () => { expect(duplicates).toHaveLength(0); }); + describe("Admin Viewer parity", () => { + // Admin Viewer follows a "read parity with Proxy Admin, no writes, no + // cost-incurring actions" rule. Playground stays hidden (incurs LLM + // cost); Models + Endpoints and Agents must be visible read-only. + const adminViewerAuth = { + userId: "admin-viewer-user-id", + accessToken: "test-access-token", + userRole: "admin_viewer", + token: "test-token", + userEmail: "viewer@example.com", + premiumUser: false, + disabledPersonalKeyCreation: false, + showSSOBanner: false, + }; + + it("hides Playground from Admin Viewer (cost-incurring action)", () => { + mockUseAuthorized.mockReturnValueOnce(adminViewerAuth); + renderWithProviders(); + expect(screen.queryByText("Playground")).not.toBeInTheDocument(); + }); + + it("shows Models + Endpoints to Admin Viewer (read-only)", () => { + mockUseAuthorized.mockReturnValueOnce(adminViewerAuth); + renderWithProviders(); + expect(screen.getByText("Models + Endpoints")).toBeInTheDocument(); + }); + + it("shows Agents (under Agentic) to Admin Viewer (read-only)", async () => { + mockUseAuthorized.mockReturnValueOnce(adminViewerAuth); + renderWithProviders(); + // Agents is now nested under the "Agentic" submenu — expand parent + // first to render the children, then assert Agents is visible. + act(() => { + fireEvent.click(screen.getByText("Agentic")); + }); + await waitFor(() => { + expect(screen.getByText("Agents")).toBeInTheDocument(); + }); + }); + + it("shows Logs to Admin Viewer", () => { + mockUseAuthorized.mockReturnValueOnce(adminViewerAuth); + renderWithProviders(); + expect(screen.getByText("Logs")).toBeInTheDocument(); + }); + }); + it("should show Organizations tab for organization admins", () => { mockUseAuthorized.mockReturnValueOnce({ userId: "org-admin-user-id", diff --git a/ui/litellm-dashboard/src/components/leftnav.tsx b/ui/litellm-dashboard/src/components/leftnav.tsx index cecc99739b1..f2fe1fe96ec 100644 --- a/ui/litellm-dashboard/src/components/leftnav.tsx +++ b/ui/litellm-dashboard/src/components/leftnav.tsx @@ -32,7 +32,14 @@ import { import type { MenuProps } from "antd"; import { ConfigProvider, Layout, Menu } from "antd"; import { useMemo } from "react"; -import { all_admin_roles, internalUserRoles, isAdminRole, isUserTeamAdminForAnyTeam, rolesWithWriteAccess } from "../utils/roles"; +import { + all_admin_roles, + internalUserRoles, + isAdminRole, + isUserTeamAdminForAnyTeam, + rolesAllowedToViewWriteScopedPages, + rolesWithWriteAccess, +} from "../utils/roles"; import NewBadge from "./common_components/NewBadge"; import type { Organization } from "./networking"; import UsageIndicator from "./UsageIndicator"; @@ -118,7 +125,9 @@ const menuGroups: MenuGroup[] = [ page: "models", label: "Models + Endpoints", icon: , - roles: rolesWithWriteAccess, + // Admin Viewer can view models read-only (write actions are + // hidden inside the page); Playground above stays write-only. + roles: rolesAllowedToViewWriteScopedPages, }, { key: "agentic", @@ -131,7 +140,9 @@ const menuGroups: MenuGroup[] = [ page: "agents", label: "Agents", icon: , - roles: rolesWithWriteAccess, + // Admin Viewer can view agents read-only (write actions are + // hidden inside the page); Playground above stays write-only. + roles: rolesAllowedToViewWriteScopedPages, }, { key: "workflows", diff --git a/ui/litellm-dashboard/src/components/model_add/credentials.test.tsx b/ui/litellm-dashboard/src/components/model_add/credentials.test.tsx index 2504b2a9789..6d555621ca8 100644 --- a/ui/litellm-dashboard/src/components/model_add/credentials.test.tsx +++ b/ui/litellm-dashboard/src/components/model_add/credentials.test.tsx @@ -30,7 +30,7 @@ const createQueryClient = () => describe("CredentialsPanel", () => { it("should render", () => { - mockUseAuthorized.mockReturnValue({ accessToken: "test-token" }); + mockUseAuthorized.mockReturnValue({ accessToken: "test-token", userRole: "Admin" }); mockUseCredentials.mockReturnValue({ data: { credentials: [] }, refetch: vi.fn(), @@ -54,7 +54,7 @@ describe("CredentialsPanel", () => { }, ]; - mockUseAuthorized.mockReturnValue({ accessToken: "test-token" }); + mockUseAuthorized.mockReturnValue({ accessToken: "test-token", userRole: "Admin" }); mockUseCredentials.mockReturnValue({ data: { credentials }, refetch: vi.fn(), @@ -70,7 +70,7 @@ describe("CredentialsPanel", () => { }); it("should display empty state when no credentials are provided", () => { - mockUseAuthorized.mockReturnValue({ accessToken: "test-token" }); + mockUseAuthorized.mockReturnValue({ accessToken: "test-token", userRole: "Admin" }); mockUseCredentials.mockReturnValue({ data: { credentials: [] }, refetch: vi.fn(), @@ -86,7 +86,7 @@ describe("CredentialsPanel", () => { }); it("should open add modal when add button is clicked", async () => { - mockUseAuthorized.mockReturnValue({ accessToken: "test-token" }); + mockUseAuthorized.mockReturnValue({ accessToken: "test-token", userRole: "Admin" }); mockUseCredentials.mockReturnValue({ data: { credentials: [] }, refetch: vi.fn(), @@ -108,4 +108,63 @@ describe("CredentialsPanel", () => { expect(screen.getByText("Add New Credential")).toBeInTheDocument(); }); }); + + describe("Admin Viewer write-action gating", () => { + // Admin Viewer can VIEW credentials but must not be able to add / edit / + // delete them. The page shows the credential list read-only. + const credentials: CredentialItem[] = [ + { + credential_name: "openai-key", + credential_values: {}, + credential_info: { custom_llm_provider: "openai" }, + }, + ]; + + it("hides the Add Credential button for Admin Viewer", () => { + mockUseAuthorized.mockReturnValue({ + accessToken: "test-token", + userRole: "Admin Viewer", + }); + mockUseCredentials.mockReturnValue({ + data: { credentials }, + refetch: vi.fn(), + }); + + render( + + + , + ); + + // Credential row still renders (read parity). + expect(screen.getByText("openai-key")).toBeInTheDocument(); + // But no Add Credential button (write blocked). + expect( + screen.queryByRole("button", { name: /add credential/i }), + ).not.toBeInTheDocument(); + }); + + it("hides Edit / Delete buttons on existing credentials for Admin Viewer", () => { + mockUseAuthorized.mockReturnValue({ + accessToken: "test-token", + userRole: "Admin Viewer", + }); + mockUseCredentials.mockReturnValue({ + data: { credentials }, + refetch: vi.fn(), + }); + + const { container } = render( + + + , + ); + + // The Actions cell should be empty (no edit/delete buttons rendered). + // We rely on the row being visible but containing no ` + {canModifyCredentials && ( + + )}
Configured credentials for different AI providers. Add and manage your API credentials.
@@ -166,22 +171,26 @@ const CredentialsPanel: React.FC = ({ uploadProps }) => { {renderProviderBadge((credential.credential_info?.custom_llm_provider as string) || "-")} - - + {canModify && ( + <> + + + + )}