chore: merge main into litellm_mcp_ui_prompts_resources

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
joshua 2026-09-19 00:51:37 +00:00
commit 48421a4723
78 changed files with 4505 additions and 513 deletions

View file

@ -1785,6 +1785,12 @@ jobs:
- wait_for_service:
url: http://localhost:4000
timeout: "300"
- run:
name: Seed the routing strategy through /config/update
command: |
curl --noproxy '*' -sSf -X POST http://localhost:4000/config/update \
-H 'Authorization: Bearer sk-1234' -H 'Content-Type: application/json' \
-d '{"router_settings": {"routing_strategy": "usage-based-routing-v2"}}'
- run:
name: Run tests
command: |

View file

@ -125,6 +125,9 @@ start_proxy() {
start_proxy 4000 proxy.log
proxy_pid="$launched_pid"
.venv/bin/python .circleci/scripts/wait_integration_services.py
curl --noproxy '*' -sSf -X POST "$INTEGRATION_PROXY_URL/config/update" \
-H "Authorization: Bearer $LITELLM_MASTER_KEY" -H 'Content-Type: application/json' \
-d '{"router_settings": {"num_retries": 0}}' > "$results/seed-router-settings.json"
if [ "$suite" = management ]; then
export INTEGRATION_PEER_URL=http://127.0.0.1:4001
start_proxy 4001 peer.log

View file

@ -17,6 +17,7 @@ from fastapi import HTTPException
import litellm
from litellm._logging import verbose_proxy_logger
from litellm.constants import DEFAULT_OPENAI_MODERATIONS_MODEL
from litellm.integrations.custom_logger import CustomLogger
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.guardrails._content_utils import iter_message_text
@ -24,11 +25,9 @@ from litellm.types.utils import CallTypesLiteral
class _ENTERPRISE_OpenAI_Moderation(CustomLogger):
def __init__(self):
self.model_name = (
litellm.openai_moderations_model_name or "text-moderation-latest"
) # pass the model_name you initialized on litellm.Router()
pass
@property
def model_name(self) -> str:
return litellm.openai_moderations_model_name or DEFAULT_OPENAI_MODERATIONS_MODEL
#### CALL HOOKS - proxy only ####

View file

@ -158,6 +158,8 @@ DEFAULT_SEMANTIC_GUARD_EMBEDDING_MODEL: Final = str(
)
DEFAULT_SEMANTIC_GUARD_SIMILARITY_THRESHOLD = float(os.getenv("DEFAULT_SEMANTIC_GUARD_SIMILARITY_THRESHOLD", 0.75))
DEFAULT_OPENAI_MODERATIONS_MODEL: Final = "omni-moderation-latest"
# MCP OAuth2 Client Credentials Defaults
MCP_OAUTH2_TOKEN_EXPIRY_BUFFER_SECONDS: Final = int(os.getenv("MCP_OAUTH2_TOKEN_EXPIRY_BUFFER_SECONDS", "60"))
MCP_OAUTH2_TOKEN_CACHE_MAX_SIZE: Final = int(os.getenv("MCP_OAUTH2_TOKEN_CACHE_MAX_SIZE", "200"))

View file

@ -464,6 +464,7 @@ class RateLimitError(openai.RateLimitError):
rate_limit_type: str | RateLimitType | None = None,
headers: dict[str, str] | None = None,
detail: Any = None,
body: object | None = None,
):
self.status_code = 429
self.message = f"litellm.RateLimitError: {message}"
@ -507,7 +508,7 @@ class RateLimitError(openai.RateLimitError):
),
)
super().__init__(
self.message, response=self.response, body=None
self.message, response=self.response, body=body
) # Call the base class constructor with the parameters it needs
self.code = "429"
self.type = "throttling_error"
@ -765,6 +766,7 @@ class InternalServerError(openai.InternalServerError):
litellm_debug_info: str | None = None,
max_retries: int | None = None,
num_retries: int | None = None,
body: object | None = None,
):
self.status_code = 500
self.message = f"litellm.InternalServerError: {message}"
@ -783,7 +785,7 @@ class InternalServerError(openai.InternalServerError):
),
)
super().__init__(
self.message, response=self.response, body=None
self.message, response=self.response, body=body
) # Call the base class constructor with the parameters it needs
def __str__(self):
@ -815,6 +817,7 @@ class APIError(openai.APIError):
litellm_debug_info: str | None = None,
max_retries: int | None = None,
num_retries: int | None = None,
body: object | None = None,
):
self.status_code = status_code
self.message = f"litellm.APIError: {message}"
@ -825,7 +828,7 @@ class APIError(openai.APIError):
self.num_retries = num_retries
if request is None:
request = httpx.Request(method="POST", url="https://api.openai.com/v1")
super().__init__(self.message, request=request, body=None)
super().__init__(self.message, request=request, body=body)
def __str__(self):
_message = self.message

View file

@ -224,6 +224,12 @@ _FINISH_REASON_MAP: Final[dict[str, OpenAIChatCompletionFinishReason]] = {
"IMAGE_PROHIBITED_CONTENT": "content_filter",
"TOO_MANY_TOOL_CALLS": "stop",
"MALFORMED_RESPONSE": "stop",
"NO_IMAGE": "content_filter",
"IMAGE_RECITATION": "content_filter",
"IMAGE_OTHER": "content_filter",
"ESCALATION": "content_filter",
"UNEXPECTED_TOOL_CALL": "stop",
"MISSING_THOUGHT_SIGNATURE": "stop",
# Zhipu GLM
"network_error": "stop",
"sensitive": "content_filter",

View file

@ -307,6 +307,7 @@ def _map_openai_exception(
model=model,
llm_provider=custom_llm_provider,
response=response,
body=getattr(original_exception, "body", None),
)
elif ExceptionCheckers.is_error_str_context_window_exceeded(error_str):
raise ContextWindowExceededError(
@ -381,6 +382,7 @@ def _map_openai_exception(
message=f"{exception_provider} - {message}",
model=model,
llm_provider=custom_llm_provider,
body=getattr(original_exception, "body", None),
)
elif "Request too large" in error_str:
raise RateLimitError(
@ -389,6 +391,7 @@ def _map_openai_exception(
llm_provider=custom_llm_provider,
response=response,
litellm_debug_info=extra_information,
body=getattr(original_exception, "body", None),
)
elif (
"The api_key client option must be set either by passing api_key to the client or by setting the OPENAI_API_KEY environment variable"
@ -460,6 +463,7 @@ def _map_openai_exception(
llm_provider=custom_llm_provider,
response=response,
litellm_debug_info=extra_information,
body=getattr(original_exception, "body", None),
)
elif original_exception.status_code == 500:
raise InternalServerError(
@ -468,6 +472,7 @@ def _map_openai_exception(
llm_provider=custom_llm_provider,
response=response,
litellm_debug_info=extra_information,
body=getattr(original_exception, "body", None),
)
elif original_exception.status_code == 502:
raise BadGatewayError(

View file

@ -1410,6 +1410,8 @@ class LiteLLMAnthropicMessagesAdapter:
return "max_tokens"
elif openai_finish_reason == "tool_calls":
return "tool_use"
elif openai_finish_reason in ["content_filter", "refusal"]:
return "refusal"
return "end_turn"
@staticmethod

View file

@ -108,6 +108,7 @@ BEDROCK_COMPUTER_USE_TOOLS: Final = [
"bash_",
"text_editor_",
]
BEDROCK_OPENAI_COMPAT_MIN_MAX_TOKENS: Final = 16
# Beta header patterns that are not supported by Bedrock Converse API
# These will be filtered out to prevent errors
@ -378,6 +379,10 @@ class AmazonConverseConfig(BaseConfig):
def _is_openai_gpt_reasoning_model(model: str) -> bool:
return re.search(r"openai\.gpt-\d", model) is not None
@staticmethod
def _requires_min_max_tokens(model: str) -> bool:
return re.search(r"openai\.gpt-\d|xai\.grok-", model) is not None
def _is_nova_2_model(self, model: str) -> bool:
"""
Check if the model is a Nova 2 model that supports reasoningConfig.
@ -1000,7 +1005,11 @@ class AmazonConverseConfig(BaseConfig):
is_thinking_enabled=is_thinking_enabled,
)
if param == "max_tokens" or param == "max_completion_tokens":
optional_params["maxTokens"] = value
optional_params["maxTokens"] = (
max(value, BEDROCK_OPENAI_COMPAT_MIN_MAX_TOKENS)
if isinstance(value, int) and self._requires_min_max_tokens(model)
else value
)
if param == "stream":
optional_params["stream"] = value
if param == "stop":

View file

@ -6,7 +6,7 @@ import time
from collections.abc import Callable, Mapping, Sequence
from copy import deepcopy
from functools import partial
from typing import TYPE_CHECKING, Any, Final, Literal, Optional, Union, cast
from typing import TYPE_CHECKING, Any, Final, Literal, Optional, Union, cast, get_args
import httpx
@ -57,6 +57,7 @@ from litellm.types.llms.vertex_ai import (
ContentType,
FunctionCallingConfig,
FunctionDeclaration,
GeminiFinishReason,
GeminiThinkingConfig,
GenerateContentResponseBody,
HttpxPartType,
@ -1330,25 +1331,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
"IMAGE_PROHIBITED_CONTENT": "The token generation was stopped as the response was flagged for prohibited image content.",
}
_GEMINI_FINISH_REASON_KEYS = frozenset(
{
"STOP",
"MAX_TOKENS",
"SAFETY",
"RECITATION",
"FINISH_REASON_UNSPECIFIED",
"MALFORMED_FUNCTION_CALL",
"LANGUAGE",
"OTHER",
"BLOCKLIST",
"PROHIBITED_CONTENT",
"SPII",
"IMAGE_SAFETY",
"IMAGE_PROHIBITED_CONTENT",
"TOO_MANY_TOOL_CALLS",
"MALFORMED_RESPONSE",
}
)
_GEMINI_FINISH_REASON_KEYS: Final[frozenset[str]] = frozenset(get_args(GeminiFinishReason))
@staticmethod
def get_finish_reason_mapping() -> dict[str, OpenAIChatCompletionFinishReason]:
@ -2232,22 +2215,23 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
grounding_metadata: Final[list[dict]] = []
url_context_metadata: Final[list[dict]] = []
image_response: list[ImageURLListItem] | None = None
safety_ratings: Final[list] = []
citation_metadata: Final[list] = []
chat_completion_message: Final[ChatCompletionResponseMessage] = {"role": "assistant"}
chat_completion_logprobs: ChoiceLogprobs | None = None
tools: list[ChatCompletionToolCallChunk] | None = []
functions: ChatCompletionToolCallFunctionChunk | None = None
thinking_blocks: list[ChatCompletionThinkingBlock] | None = None
reasoning_content: str | None = None
thought_signatures: Sequence[str] | None = None
server_side_tool_invocations: list[dict[str, object]] | None = None
for idx, candidate in enumerate(_candidates):
if "content" not in candidate:
if "content" not in candidate and "finishReason" not in candidate:
continue
image_response: list[ImageURLListItem] | None = None
chat_completion_message: ChatCompletionResponseMessage = {"role": "assistant"}
chat_completion_logprobs: ChoiceLogprobs | None = None
tools: list[ChatCompletionToolCallChunk] | None = None
functions: ChatCompletionToolCallFunctionChunk | None = None
thinking_blocks: list[ChatCompletionThinkingBlock] | None = None
reasoning_content: str | None = None
thought_signatures: Sequence[str] | None = None
server_side_tool_invocations: list[dict[str, object]] | None = None
# Extract metadata using helper function
(
candidate_grounding_metadata,
@ -2261,7 +2245,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
safety_ratings.extend(candidate_safety_ratings)
citation_metadata.extend(candidate_citation_metadata)
if "parts" in candidate["content"]:
if "content" in candidate and candidate["content"] and "parts" in candidate["content"]:
(
content,
reasoning_content,
@ -2368,14 +2352,18 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
)
model_response.choices.append(choice)
elif isinstance(model_response, ModelResponse):
native_finish_reason = candidate.get("finishReason")
choice = litellm.Choices(
finish_reason=VertexGeminiConfig._check_finish_reason(
chat_completion_message, candidate.get("finishReason")
chat_completion_message, native_finish_reason
),
index=candidate.get("index", idx),
message=chat_completion_message,
logprobs=chat_completion_logprobs,
enhancements=None,
provider_specific_fields=(
{"native_finish_reason": native_finish_reason} if native_finish_reason is not None else None
),
)
model_response.choices.append(choice)
@ -3173,12 +3161,10 @@ class ModelResponseIterator:
self.has_seen_tool_calls = True
break
# _process_candidates skips candidates without a "content" part, so a
# content-less chunk leaves choices empty and the downstream streaming
# handler hits IndexError on choices[0]. This covers the final chunk
# (finishReason, no content) and mid-stream metadata-only chunks
# (grounding/web-search/thought, no content and no finishReason — seen
# with web_search + reasoning) by emitting an empty-delta choice.
# _process_candidates skips candidates with neither "content" nor
# "finishReason", so a metadata-only chunk (grounding/web-search/thought,
# seen with web_search + reasoning) leaves choices empty and the downstream
# streaming handler hits IndexError on choices[0]. Emit an empty-delta choice.
if not model_response.choices and _candidates:
from litellm.types.utils import Delta, StreamingChoices

View file

@ -19244,7 +19244,20 @@
"supports_function_calling": true,
"supports_anthropic_thinking_payload": true,
"supports_prompt_caching": true,
"supports_tool_choice": true
"supports_tool_choice": true,
"supports_reasoning": true,
"supported_modalities": [
"text",
"image",
"audio",
"video"
],
"supported_output_modalities": [
"text"
],
"supports_vision": true,
"supports_audio_input": true,
"supports_video_input": true
},
"databricks/databricks-gemini-2-5-pro": {
"cache_creation_input_token_cost": 1.24999e-06,
@ -19265,7 +19278,21 @@
"supports_function_calling": true,
"supports_anthropic_thinking_payload": true,
"supports_prompt_caching": true,
"supports_tool_choice": true
"supports_tool_choice": true,
"deprecation_date": "2026-10-02",
"supports_reasoning": true,
"supported_modalities": [
"text",
"image",
"audio",
"video"
],
"supported_output_modalities": [
"text"
],
"supports_vision": true,
"supports_audio_input": true,
"supports_video_input": true
},
"databricks/databricks-gemini-3-1-flash-lite": {
"cache_creation_input_token_cost": 3.1248e-07,
@ -19285,7 +19312,19 @@
"source": "https://www.databricks.com/product/pricing/proprietary-foundation-model-serving",
"supports_function_calling": true,
"supports_prompt_caching": true,
"supports_tool_choice": true
"supports_tool_choice": true,
"supported_modalities": [
"text",
"image",
"audio",
"video"
],
"supported_output_modalities": [
"text"
],
"supports_vision": true,
"supports_audio_input": true,
"supports_video_input": true
},
"databricks/databricks-gemini-3-1-flash-image": {
"litellm_provider": "databricks",
@ -19347,7 +19386,20 @@
"source": "https://www.databricks.com/product/pricing/proprietary-foundation-model-serving",
"supports_function_calling": true,
"supports_prompt_caching": true,
"supports_tool_choice": true
"supports_tool_choice": true,
"supports_reasoning": true,
"supported_modalities": [
"text",
"image",
"audio",
"video"
],
"supported_output_modalities": [
"text"
],
"supports_vision": true,
"supports_audio_input": true,
"supports_video_input": true
},
"databricks/databricks-gemini-3-flash": {
"cache_creation_input_token_cost": 6.2503e-07,
@ -19367,7 +19419,19 @@
"source": "https://www.databricks.com/product/pricing/proprietary-foundation-model-serving",
"supports_function_calling": true,
"supports_prompt_caching": true,
"supports_tool_choice": true
"supports_tool_choice": true,
"supported_modalities": [
"text",
"image",
"audio",
"video"
],
"supported_output_modalities": [
"text"
],
"supports_vision": true,
"supports_audio_input": true,
"supports_video_input": true
},
"databricks/databricks-gemini-3-pro": {
"cache_creation_input_token_cost": 2.49998e-06,
@ -21433,7 +21497,11 @@
"mode": "chat",
"supports_tool_choice": true,
"supports_function_calling": true,
"supports_image_size": false
"supports_image_size": false,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_vision": true,
"supports_audio_input": true
},
"deepinfra/google/gemini-2.5-pro": {
"max_tokens": 1000000,
@ -21444,7 +21512,11 @@
"litellm_provider": "deepinfra",
"mode": "chat",
"supports_tool_choice": true,
"supports_function_calling": true
"supports_function_calling": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_vision": true,
"supports_audio_input": true
},
"deepinfra/google/gemma-3-12b-it": {
"max_tokens": 131072,
@ -29340,26 +29412,6 @@
"supports_parallel_function_calling": true,
"supports_vision": true
},
"github_copilot/gemini-2.5-pro": {
"litellm_provider": "github_copilot",
"max_input_tokens": 128000,
"max_output_tokens": 64000,
"max_tokens": 64000,
"mode": "chat",
"supports_function_calling": true,
"supports_parallel_function_calling": true,
"supports_vision": true
},
"github_copilot/gemini-3-pro-preview": {
"litellm_provider": "github_copilot",
"max_input_tokens": 128000,
"max_output_tokens": 64000,
"max_tokens": 64000,
"mode": "chat",
"supports_function_calling": true,
"supports_parallel_function_calling": true,
"supports_vision": true
},
"github_copilot/gpt-3.5-turbo": {
"litellm_provider": "github_copilot",
"max_input_tokens": 16384,
@ -30014,17 +30066,6 @@
"output_cost_per_token": 8.8e-07,
"supports_function_calling": true
},
"gmi/google/gemini-3-pro-preview": {
"input_cost_per_token": 2e-06,
"litellm_provider": "gmi",
"max_input_tokens": 1048576,
"max_output_tokens": 65536,
"max_tokens": 65536,
"mode": "chat",
"output_cost_per_token": 1.2e-05,
"supports_function_calling": true,
"supports_vision": true
},
"gmi/google/gemini-3-flash-preview": {
"input_cost_per_token": 5e-07,
"litellm_provider": "gmi",
@ -30034,7 +30075,8 @@
"mode": "chat",
"output_cost_per_token": 3e-06,
"supports_function_calling": true,
"supports_vision": true
"supports_vision": true,
"supports_system_messages": true
},
"gmi/moonshotai/Kimi-K2-Thinking": {
"input_cost_per_token": 8e-07,
@ -40163,7 +40205,12 @@
"supports_response_schema": true,
"supports_vision": true,
"supports_native_streaming": true,
"supports_image_size": false
"supports_image_size": false,
"supports_reasoning": true,
"supports_system_messages": true,
"supports_pdf_input": true,
"supports_audio_input": true,
"supports_video_input": true
},
"oci/google.gemini-2.5-pro": {
"input_cost_per_token": 1.25e-06,
@ -40177,7 +40224,12 @@
"supports_function_calling": true,
"supports_response_schema": true,
"supports_vision": true,
"supports_native_streaming": true
"supports_native_streaming": true,
"supports_reasoning": true,
"supports_system_messages": true,
"supports_pdf_input": true,
"supports_audio_input": true,
"supports_video_input": true
},
"oci/google.gemini-2.5-flash-lite": {
"input_cost_per_token": 7.5e-08,
@ -40192,7 +40244,12 @@
"supports_response_schema": true,
"supports_vision": true,
"supports_native_streaming": true,
"supports_image_size": false
"supports_image_size": false,
"supports_reasoning": false,
"supports_system_messages": true,
"supports_pdf_input": true,
"supports_audio_input": true,
"supports_video_input": true
},
"oci/cohere.command-a-vision": {
"input_cost_per_token": 1.56e-06,
@ -41435,7 +41492,7 @@
"max_tokens": 65535,
"mode": "chat",
"output_cost_per_token": 2.5e-06,
"supports_audio_output": true,
"supports_audio_output": false,
"supports_function_calling": true,
"supports_response_schema": true,
"supports_system_messages": true,
@ -41449,7 +41506,8 @@
"supports_audio_input": true,
"supports_pdf_input": true,
"supports_reasoning": true,
"supports_web_search": false
"supports_web_search": false,
"supports_video_input": true
},
"openrouter/google/gemini-2.5-pro": {
"cache_creation_input_token_cost": 3.75e-07,
@ -41462,7 +41520,7 @@
"max_tokens": 65536,
"mode": "chat",
"output_cost_per_token": 1e-05,
"supports_audio_output": true,
"supports_audio_output": false,
"supports_function_calling": true,
"supports_response_schema": true,
"supports_system_messages": true,
@ -41478,7 +41536,8 @@
"supports_audio_input": true,
"supports_pdf_input": true,
"supports_reasoning": true,
"supports_web_search": false
"supports_web_search": false,
"supports_video_input": true
},
"openrouter/google/gemini-3-pro-preview": {
"cache_read_input_token_cost": 2e-07,
@ -41563,7 +41622,8 @@
"supports_url_context": true,
"supports_vision": true,
"supports_web_search": false,
"tpm": 800000
"tpm": 800000,
"supports_video_input": true
},
"openrouter/google/gemini-3.1-flash-lite-preview": {
"cache_creation_input_token_cost": 8.33333333333333e-08,
@ -41690,7 +41750,8 @@
"supports_system_messages": true,
"supports_tool_choice": true,
"supports_vision": true,
"supports_web_search": false
"supports_web_search": false,
"supports_video_input": true
},
"openrouter/gryphe/mythomax-l2-13b": {
"input_cost_per_token": 8e-08,
@ -44413,12 +44474,16 @@
"output_cost_per_token": 1.2e-05,
"litellm_provider": "replicate",
"mode": "chat",
"supports_function_calling": true,
"supports_parallel_function_calling": true,
"supports_function_calling": false,
"supports_parallel_function_calling": false,
"supports_vision": true,
"supports_system_messages": true,
"supports_tool_choice": true,
"supports_response_schema": true
"supports_tool_choice": false,
"supports_response_schema": false,
"input_cost_per_token_above_200k_tokens": 4e-06,
"output_cost_per_token_above_200k_tokens": 1.8e-05,
"supports_audio_input": true,
"supports_video_input": true
},
"replicate/anthropic/claude-4.5-sonnet": {
"input_cost_per_token": 3e-06,
@ -44487,17 +44552,19 @@
"supports_response_schema": true
},
"replicate/google/gemini-2.5-flash": {
"input_cost_per_token": 2.5e-06,
"input_cost_per_token": 3e-07,
"output_cost_per_token": 2.5e-06,
"litellm_provider": "replicate",
"mode": "chat",
"supports_function_calling": true,
"supports_parallel_function_calling": true,
"supports_function_calling": false,
"supports_parallel_function_calling": false,
"supports_vision": true,
"supports_system_messages": true,
"supports_tool_choice": true,
"supports_response_schema": true,
"supports_image_size": false
"supports_tool_choice": false,
"supports_response_schema": false,
"supports_image_size": false,
"supports_reasoning": true,
"supports_video_input": true
},
"replicate/openai/gpt-oss-120b": {
"input_cost_per_token": 1.8e-07,
@ -48076,10 +48143,15 @@
"supports_function_calling": true,
"supports_tool_choice": true,
"supports_response_schema": true,
"supports_image_size": false
"supports_image_size": false,
"cache_read_input_token_cost": 3e-08,
"supports_reasoning": true,
"supports_pdf_input": true,
"supports_web_search": true,
"supports_prompt_caching": true
},
"vercel_ai_gateway/google/gemini-2.5-pro": {
"input_cost_per_token": 2.5e-06,
"input_cost_per_token": 1.25e-06,
"litellm_provider": "vercel_ai_gateway",
"max_input_tokens": 1048576,
"max_output_tokens": 65536,
@ -48089,7 +48161,15 @@
"supports_vision": true,
"supports_function_calling": true,
"supports_tool_choice": true,
"supports_response_schema": true
"supports_response_schema": true,
"input_cost_per_token_above_200k_tokens": 2.5e-06,
"output_cost_per_token_above_200k_tokens": 1.5e-05,
"cache_read_input_token_cost": 1.25e-07,
"cache_read_input_token_cost_above_200k_tokens": 2.5e-07,
"supports_reasoning": true,
"supports_pdf_input": true,
"supports_web_search": true,
"supports_prompt_caching": true
},
"vercel_ai_gateway/google/gemini-embedding-001": {
"input_cost_per_token": 1.5e-07,
@ -62224,7 +62304,8 @@
"supports_response_schema": true,
"supports_reasoning": true,
"supports_vision": true,
"source": "https://deepinfra.com/pricing"
"source": "https://deepinfra.com/pricing",
"supports_audio_input": true
},
"deepinfra/XiaomiMiMo/MiMo-V2.5": {
"max_tokens": 262144,
@ -62524,7 +62605,8 @@
"supports_response_schema": true,
"supports_reasoning": true,
"supports_vision": true,
"source": "https://deepinfra.com/pricing"
"source": "https://deepinfra.com/pricing",
"supports_audio_input": true
},
"deepinfra/google/gemini-3.7-flash": {
"max_tokens": 1000000,
@ -62538,7 +62620,8 @@
"supports_response_schema": true,
"supports_reasoning": true,
"supports_vision": true,
"source": "https://deepinfra.com/pricing"
"source": "https://deepinfra.com/pricing",
"supports_audio_input": true
},
"deepinfra/inclusionAI/Ling-3.0-flash": {
"max_tokens": 131072,
@ -62970,7 +63053,8 @@
"supports_response_schema": true,
"supports_reasoning": true,
"supports_vision": true,
"source": "https://deepinfra.com/pricing"
"source": "https://deepinfra.com/pricing",
"supports_audio_input": true
},
"deepinfra/XiaomiMiMo/MiMo-V2.5-Pro": {
"max_tokens": 1048576,
@ -65365,7 +65449,8 @@
"deprecation_date": "2026-10-20",
"input_cost_per_audio_token": 3e-07,
"supports_prompt_caching": true,
"supports_web_search": false
"supports_web_search": false,
"supports_video_input": true
},
"openrouter/google/gemini-3.5-flash": {
"cache_creation_input_token_cost": 8.33333333333333e-08,
@ -65388,7 +65473,8 @@
"cache_read_input_token_cost": 1.5e-07,
"input_cost_per_audio_token": 3e-06,
"supports_prompt_caching": true,
"supports_web_search": false
"supports_web_search": false,
"supports_video_input": true
},
"openrouter/google/gemini-3.5-flash-lite": {
"cache_creation_input_token_cost": 8.33333333333333e-08,
@ -65411,7 +65497,8 @@
"cache_read_input_token_cost": 3e-08,
"input_cost_per_audio_token": 3e-07,
"supports_prompt_caching": true,
"supports_web_search": false
"supports_web_search": false,
"supports_video_input": true
},
"openrouter/google/gemini-3.6-flash": {
"cache_creation_input_token_cost": 4.16666666666667e-08,
@ -65434,7 +65521,8 @@
"cache_read_input_token_cost": 7.5e-08,
"input_cost_per_audio_token": 7.5e-07,
"supports_prompt_caching": true,
"supports_web_search": false
"supports_web_search": false,
"supports_video_input": true
},
"openrouter/google/gemini-3.7-flash": {
"cache_creation_input_token_cost": 4.16666666666667e-08,
@ -65457,7 +65545,8 @@
"cache_read_input_token_cost": 7.5e-08,
"input_cost_per_audio_token": 7.5e-07,
"supports_prompt_caching": true,
"supports_web_search": false
"supports_web_search": false,
"supports_video_input": true
},
"openrouter/google/gemini-3.8-flash": {
"cache_creation_input_token_cost": 4.16666666666667e-08,
@ -65480,7 +65569,8 @@
"cache_read_input_token_cost": 7.5e-08,
"input_cost_per_audio_token": 7.5e-07,
"supports_prompt_caching": true,
"supports_web_search": false
"supports_web_search": false,
"supports_video_input": true
},
"openrouter/openai/gpt-4o-mini": {
"input_cost_per_token": 1.5e-07,
@ -67108,7 +67198,8 @@
"supports_pdf_input": true,
"supports_audio_input": true,
"supports_prompt_caching": true,
"supports_web_search": false
"supports_web_search": false,
"supports_video_input": true
},
"openrouter/qwen/qwen3-max-thinking": {
"input_cost_per_token": 7.8e-07,
@ -71974,7 +72065,8 @@
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
"supports_web_search": false
"supports_web_search": false,
"supports_video_input": true
},
"openrouter/google/gemini-2.5-flash:batch": {
"cache_read_input_audio_token_cost": 1e-07,
@ -71997,7 +72089,8 @@
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
"supports_web_search": false
"supports_web_search": false,
"supports_video_input": true
},
"openrouter/google/gemini-2.5-pro:batch": {
"cache_read_input_audio_token_cost": 1.25e-07,
@ -72023,7 +72116,8 @@
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
"supports_web_search": false
"supports_web_search": false,
"supports_video_input": true
},
"openrouter/google/gemini-3-flash-preview:batch": {
"input_cost_per_audio_token": 5e-07,
@ -72043,7 +72137,8 @@
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
"supports_web_search": false
"supports_web_search": false,
"supports_video_input": true
},
"openrouter/google/gemini-3.1-flash-lite:batch": {
"cache_read_input_audio_token_cost": 2.5e-08,
@ -72065,7 +72160,8 @@
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
"supports_web_search": false
"supports_web_search": false,
"supports_video_input": true
},
"openrouter/google/gemini-3.1-pro-preview:batch": {
"input_cost_per_audio_token": 1e-06,
@ -72087,7 +72183,8 @@
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
"supports_web_search": false
"supports_web_search": false,
"supports_video_input": true
},
"openrouter/google/gemini-3.5-flash-lite:batch": {
"cache_read_input_audio_token_cost": 1.5e-08,
@ -72109,7 +72206,8 @@
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
"supports_web_search": false
"supports_web_search": false,
"supports_video_input": true
},
"openrouter/google/gemini-3.5-flash:batch": {
"cache_read_input_audio_token_cost": 1.5e-07,
@ -72131,7 +72229,8 @@
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
"supports_web_search": false
"supports_web_search": false,
"supports_video_input": true
},
"openrouter/google/gemini-3.6-flash:batch": {
"cache_creation_input_token_cost": 4.16666666666667e-08,
@ -72154,7 +72253,8 @@
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
"supports_web_search": false
"supports_web_search": false,
"supports_video_input": true
},
"openrouter/google/gemini-3.7-flash:batch": {
"cache_creation_input_token_cost": 4.16666666666667e-08,
@ -72177,7 +72277,8 @@
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
"supports_web_search": false
"supports_web_search": false,
"supports_video_input": true
},
"openrouter/google/gemini-3.8-flash:batch": {
"cache_creation_input_token_cost": 4.16666666666667e-08,
@ -72200,7 +72301,8 @@
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
"supports_web_search": false
"supports_web_search": false,
"supports_video_input": true
},
"openrouter/ibm-granite/granite-4.0-h-micro": {
"input_cost_per_token": 1.7e-08,

View file

@ -21,6 +21,7 @@ litellm/proxy/_experimental/mcp_server/
user_api_key_auth_mcp.py # LiteLLM admission auth and MCP request headers
token_exchange.py # OAuth token exchange handling [unchanged; V1TokenExchangeAdapter delegates here]
litellm_auth_handler.py # authenticated-user adapter for MCP sessions
client_allowlist.py # gateway-level client application allowlist (mcp_allowed_clients); leaf module, no litellm.proxy imports
outbound_credentials/ # NEW — typed upstream-credential resolution (resolve_credentials + arms)
__init__.py # public surface: resolve_credentials, the configs, CredError
result.py # Ok | Error union (pure stdlib)

View file

@ -0,0 +1,171 @@
"""
Gateway-level allowlist of MCP client applications (``general_settings.mcp_allowed_clients``).
Each entry pairs an admin-chosen ``alias`` (shown in the dashboard and logs) with the ``value`` that
identifies the client. Only the value is compared, exactly and case-sensitively.
A caller that authenticated with a JWT is identified by the claim named in
``litellm_jwtauth.mcp_client_id_jwt_field``, a value asserted by the identity provider.
Every other caller is identified by the header named in ``general_settings.mcp_client_id_header``,
which the client picks itself, so that source is a policy control rather than a security boundary.
While the allowlist is set, a caller with no usable identity source is rejected.
"""
from collections.abc import Mapping
from dataclasses import dataclass
from types import MappingProxyType
from typing import Final, Literal
from pydantic import TypeAdapter, ValidationError
from typing_extensions import ReadOnly, TypedDict
from litellm._logging import verbose_logger
from litellm.litellm_core_utils.dot_notation_indexing import get_nested_value
from litellm.types.mcp import MCPAllowedClient
MCP_ALLOWED_CLIENTS_SETTING: Final = "mcp_allowed_clients"
MCP_CLIENT_ID_HEADER_SETTING: Final = "mcp_client_id_header"
MCP_CLIENT_ID_JWT_FIELD_SETTING: Final = "mcp_client_id_jwt_field"
_JWT_AUTH_SETTING: Final = "litellm_jwtauth"
_ALLOWED_CLIENTS_ADAPTER: Final[TypeAdapter[list[MCPAllowedClient]]] = TypeAdapter(list[MCPAllowedClient])
_OPTIONAL_NAME_ADAPTER: Final[TypeAdapter[str | None]] = TypeAdapter(str | None)
_OPTIONAL_MAPPING_ADAPTER: Final[TypeAdapter[dict[str, object] | None]] = TypeAdapter(dict[str, object] | None)
_NOBODY: Final[Mapping[str, str]] = MappingProxyType({})
class MCPClientForbiddenBody(TypedDict):
error: ReadOnly[Literal["Forbidden"]]
details: ReadOnly[str]
@dataclass(frozen=True, slots=True)
class MCPClientAllowlist:
"""``aliases_by_value`` maps each admitted identity value to the alias the admin gave it."""
aliases_by_value: Mapping[str, str]
jwt_field: str | None
header: str | None
@dataclass(frozen=True, slots=True)
class MCPClientIdentity:
client_id: str
source: Literal["jwt", "header"]
source_name: str
@property
def description(self) -> str:
return f"'{self.client_id}' (from {'JWT claim' if self.source == 'jwt' else 'header'} '{self.source_name}')"
@dataclass(frozen=True, slots=True)
class MCPClientRejection:
details: str
@property
def response_body(self) -> MCPClientForbiddenBody:
body: Final[MCPClientForbiddenBody] = {"error": "Forbidden", "details": self.details}
return body
def _unidentified_rejection(reason: str) -> MCPClientRejection:
return MCPClientRejection(
details=f"{reason} This gateway only admits client applications listed in {MCP_ALLOWED_CLIENTS_SETTING}."
)
def parse_allowed_mcp_clients(raw_setting: object) -> Mapping[str, str] | None:
"""Value-to-alias mapping; None when the setting is absent (not enforced). A malformed setting admits nobody."""
if raw_setting is None:
return None
try:
clients: Final = _ALLOWED_CLIENTS_ADAPTER.validate_python(raw_setting)
except ValidationError:
verbose_logger.warning(
"%s is not a list of {alias, value} entries (%r); rejecting every MCP client until it is fixed",
MCP_ALLOWED_CLIENTS_SETTING,
raw_setting,
)
return _NOBODY
return MappingProxyType({client.value: client.alias for client in clients})
def _parse_optional_name(setting_name: str, raw_setting: object) -> str | None:
try:
name: Final = _OPTIONAL_NAME_ADAPTER.validate_python(raw_setting)
except ValidationError:
verbose_logger.warning("%s is not a string (%r); ignoring it", setting_name, raw_setting)
return None
return name or None
def _jwt_field_from_general_settings(general_settings: Mapping[str, object]) -> str | None:
try:
jwt_auth: Final = _OPTIONAL_MAPPING_ADAPTER.validate_python(general_settings.get(_JWT_AUTH_SETTING))
except ValidationError:
return None
if jwt_auth is None:
return None
return _parse_optional_name(
f"{_JWT_AUTH_SETTING}.{MCP_CLIENT_ID_JWT_FIELD_SETTING}", jwt_auth.get(MCP_CLIENT_ID_JWT_FIELD_SETTING)
)
def load_mcp_client_allowlist(general_settings: Mapping[str, object]) -> MCPClientAllowlist | None:
"""None when ``mcp_allowed_clients`` is unset, which admits every client."""
allowed_clients: Final = parse_allowed_mcp_clients(general_settings.get(MCP_ALLOWED_CLIENTS_SETTING))
if allowed_clients is None:
return None
header: Final = _parse_optional_name(
MCP_CLIENT_ID_HEADER_SETTING, general_settings.get(MCP_CLIENT_ID_HEADER_SETTING)
)
return MCPClientAllowlist(
aliases_by_value=allowed_clients,
jwt_field=_jwt_field_from_general_settings(general_settings),
header=header.lower() if header is not None else None,
)
def resolve_mcp_client_identity(
allowlist: MCPClientAllowlist,
jwt_claims: Mapping[str, object] | None,
headers: Mapping[str, str],
) -> MCPClientIdentity | MCPClientRejection:
"""A JWT caller is identified by its configured claim alone, so a header can never override the IdP."""
if jwt_claims is not None and allowlist.jwt_field is not None:
claim: Final[object] = get_nested_value(data=jwt_claims, key_path=allowlist.jwt_field)
if isinstance(claim, str) and claim:
return MCPClientIdentity(client_id=claim, source="jwt", source_name=allowlist.jwt_field)
return _unidentified_rejection(
f"The JWT presented has no '{allowlist.jwt_field}' claim naming the client application."
)
if allowlist.header is None:
configured: Final = (
f"litellm_jwtauth.{MCP_CLIENT_ID_JWT_FIELD_SETTING} for JWT callers or {MCP_CLIENT_ID_HEADER_SETTING}"
)
return _unidentified_rejection(
f"No client identity source is configured for this request; set {configured} in general_settings."
)
header_value: Final = headers.get(allowlist.header)
if header_value:
return MCPClientIdentity(client_id=header_value, source="header", source_name=allowlist.header)
return _unidentified_rejection(f"The request has no '{allowlist.header}' header naming the client application.")
def check_mcp_client_allowed(
allowlist: MCPClientAllowlist | None,
jwt_claims: Mapping[str, object] | None,
headers: Mapping[str, str],
) -> MCPClientRejection | None:
if allowlist is None:
return None
identity: Final = resolve_mcp_client_identity(allowlist, jwt_claims, headers)
if isinstance(identity, MCPClientRejection):
return identity
alias: Final = allowlist.aliases_by_value.get(identity.client_id)
if alias is None:
return MCPClientRejection(
details=f"MCP client {identity.description} is not listed in this gateway's {MCP_ALLOWED_CLIENTS_SETTING}."
)
verbose_logger.debug("Admitted MCP client '%s' identified as %s", alias, identity.description)
return None

View file

@ -197,6 +197,7 @@ if MCP_AVAILABLE:
filter_tools_by_allowed_tools,
filter_tools_by_key_team_permissions,
fire_mcp_tool_call_failure_logging,
reject_disallowed_mcp_client,
)
class MCPCatalogPrompt(Prompt):
@ -907,6 +908,7 @@ if MCP_AVAILABLE:
MCPRequestHandler,
)
reject_disallowed_mcp_client(request.headers, user_api_key_dict)
try:
mcp_server_name = _as_query_str(mcp_server_name)
toolset_name = _as_query_str(toolset_name)
@ -1211,6 +1213,7 @@ if MCP_AVAILABLE:
proxy_logging_obj,
)
reject_disallowed_mcp_client(request.headers, user_api_key_dict)
try:
user_api_key_dict = await acting_user_auth(user_api_key_dict)
data = await request.json()

View file

@ -47,6 +47,11 @@ from litellm.proxy._experimental.mcp_server.byok_credential_cache import (
cache_byok_credential,
get_cached_byok_credential,
)
from litellm.proxy._experimental.mcp_server.client_allowlist import (
MCPClientAllowlist,
check_mcp_client_allowed,
load_mcp_client_allowlist,
)
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
get_request_base_url,
)
@ -72,6 +77,7 @@ from litellm.proxy._experimental.mcp_server.oauth_utils import (
get_route_relative_request_path,
well_known_root_suffix,
)
from litellm.proxy._experimental.mcp_server.ui_session_utils import is_ui_session_credential
from litellm.proxy._experimental.mcp_server.utils import (
LITELLM_MCP_SERVER_DESCRIPTION,
LITELLM_MCP_SERVER_NAME,
@ -3701,6 +3707,25 @@ if MCP_AVAILABLE:
mcp_servers_from_path = [servers_and_path]
return mcp_servers_from_path
def _load_mcp_client_allowlist() -> MCPClientAllowlist | None:
from litellm.proxy.proxy_server import general_settings
return load_mcp_client_allowlist(general_settings)
def reject_disallowed_mcp_client(headers: Mapping[str, str], user_api_key_auth: UserAPIKeyAuth | None) -> None:
"""Gate every MCP tool surface on ``mcp_allowed_clients``; the dashboard's own session is not a client app."""
if user_api_key_auth is not None and is_ui_session_credential(user_api_key_auth):
return
rejection: Final = check_mcp_client_allowed(
allowlist=_load_mcp_client_allowlist(),
jwt_claims=user_api_key_auth.jwt_claims if user_api_key_auth is not None else None,
headers=headers,
)
if rejection is None:
return
verbose_logger.warning("Rejected MCP request from a disallowed client application: %s", rejection.details)
raise HTTPException(status_code=403, detail=rejection.response_body)
async def extract_mcp_auth_context(scope, path):
"""
Extracts mcp_servers from the path and processes the MCP request for auth context.
@ -4537,6 +4562,7 @@ if MCP_AVAILABLE:
oauth2_headers,
raw_headers,
) = await extract_mcp_auth_context(scope, path)
reject_disallowed_mcp_client(StarletteRequest(scope).headers, user_api_key_auth)
scoped_server_endpoint: Final = len(_get_mcp_servers_in_path(path) or []) == 1
# Extract client IP for MCP access control
@ -4865,6 +4891,7 @@ if MCP_AVAILABLE:
oauth2_headers,
raw_headers,
) = await extract_mcp_auth_context(scope, path)
reject_disallowed_mcp_client(StarletteRequest(scope).headers, user_api_key_auth)
scoped_server_endpoint: Final = len(_get_mcp_servers_in_path(path) or []) == 1
# Extract client IP for MCP access control

View file

@ -34,6 +34,7 @@ from litellm.types.llms.openai import (
ResponsesAPIResponse,
)
from litellm.types.mcp import (
MCPAllowedClient,
MCPAuth,
MCPAuthType,
MCPCredentials,
@ -2904,6 +2905,14 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase):
None,
description="Custom CIDR ranges that define internal/private networks for MCP access control. When set, only these ranges are treated as internal. Defaults to RFC 1918 private ranges (10.0.0.0/8, 172.16.0.0/12, 192.168.0.0/16, 127.0.0.0/8).",
)
mcp_allowed_clients: list[MCPAllowedClient] | None = Field(
None,
description="MCP client applications admitted by the gateway, each an {alias, value} pair where alias is the name shown in the dashboard and logs and value is the identity that must match exactly. When set, every MCP request must carry a client identity equal to one of the values: a JWT caller is identified by the claim named in litellm_jwtauth.mcp_client_id_jwt_field, any other caller by the header named in mcp_client_id_header. A request with no resolvable identity, or an unlisted one, is rejected with 403. Unset means every client is admitted.",
)
mcp_client_id_header: str | None = Field(
None,
description="Request header whose value names the calling MCP client application (for example 'x-mcp-client') for callers that did not authenticate with a JWT, used only while mcp_allowed_clients is set. The client picks this value itself, so it is a policy control rather than a security boundary; prefer litellm_jwtauth.mcp_client_id_jwt_field where callers use JWTs.",
)
mcp_trusted_proxy_ranges: list[str] | None = Field(
None,
description="CIDR ranges of trusted reverse proxies. When set, X-Forwarded-For and X-Forwarded-* origin headers are only trusted from these IPs.",
@ -5120,6 +5129,15 @@ class LiteLLM_JWTAuth(LiteLLMPydanticObjectBase):
"then agent_name, and the request is rejected when it matches neither."
),
)
mcp_client_id_jwt_field: str | None = Field(
default=None,
description=(
"The field in the JWT token that identifies the MCP client application (harness) making the request, "
"e.g. 'azp' or 'client_id'. Supports dot notation. Only consulted while general_settings.mcp_allowed_clients "
"is set: the claim value must be listed there or the MCP request is rejected with 403. Distinct from "
"agent_id_jwt_field, which identifies an AI agent rather than the client software."
),
)
public_key_ttl: float = 600
public_key_stale_ttl: float = Field(
default=DEFAULT_JWKS_STALE_TTL,
@ -5405,6 +5423,7 @@ class DBSpendUpdateTransactions(TypedDict):
team_member_list_transactions: dict[str, float] | None
org_list_transactions: dict[str, float] | None
org_member_list_transactions: ReadOnly[dict[str, float] | None]
project_list_transactions: ReadOnly[dict[str, float] | None]
tag_list_transactions: dict[str, float] | None
agent_list_transactions: dict[str, float] | None
model_access_group_list_transactions: ReadOnly[dict[str, float] | None]

View file

@ -94,6 +94,8 @@ from litellm.proxy.common_utils.user_api_key_cache import (
model_access_group_registry_cache_key,
model_access_group_spend_counter_key,
object_permission_cache_key,
project_cache_key,
project_spend_counter_key,
tag_cache_key,
tag_registry_cache_key,
team_membership_auth_cache_key,
@ -5680,16 +5682,22 @@ async def _project_max_budget_check(
if project_object.litellm_budget_table is not None:
max_budget = project_object.litellm_budget_table.max_budget
if (
max_budget is not None
and project_object.spend is not None
and math.isfinite(max_budget)
and project_object.spend > max_budget
):
if max_budget is None or max_budget <= 0 or not math.isfinite(max_budget):
return
from litellm.proxy.proxy_server import get_current_spend
project_spend: Final = await get_current_spend(
counter_key=project_spend_counter_key(project_object.project_id),
fallback_spend=project_object.spend or 0.0,
max_budget=max_budget,
)
if project_spend >= max_budget:
if valid_token:
call_info: Final = CallInfo(
token=valid_token.token,
spend=project_object.spend,
spend=project_spend,
max_budget=max_budget,
user_id=valid_token.user_id,
team_id=valid_token.team_id,
@ -5705,9 +5713,9 @@ async def _project_max_budget_check(
)
raise litellm.BudgetExceededError(
current_cost=project_object.spend,
current_cost=project_spend,
max_budget=max_budget,
message=f"Budget has been exceeded! Project={project_object.project_id} Current cost: {project_object.spend}, Max budget: {max_budget}",
message=f"Budget has been exceeded! Project={project_object.project_id} Current cost: {project_spend}, Max budget: {max_budget}",
entity_type=Litellm_EntityType.PROJECT.value,
entity_id=project_object.project_id,
)
@ -5757,10 +5765,6 @@ async def _project_soft_budget_check(
)
def _project_cache_key(project_id: str) -> str:
return f"project_id:{project_id}"
async def get_project_object(
project_id: str,
prisma_client: PrismaClient | None,
@ -5778,7 +5782,7 @@ async def get_project_object(
return None
# Check cache first
cache_key: Final = _project_cache_key(project_id)
cache_key: Final = project_cache_key(project_id)
deserialized_project: Final = await user_api_key_cache.async_get_cache(
key=cache_key,
model_type=LiteLLM_ProjectTableCachedObj,
@ -5820,7 +5824,7 @@ async def delete_cached_project_object(
from litellm.proxy.common_utils.auth_cache_invalidation_pubsub import evict_and_broadcast
await evict_and_broadcast(
cache_keys=(_project_cache_key(project_id),),
cache_keys=(project_cache_key(project_id),),
user_api_key_cache=user_api_key_cache,
)

View file

@ -40,6 +40,8 @@ from litellm.proxy.common_utils.user_api_key_cache import (
end_user_cache_key,
model_access_group_cache_key,
model_access_group_spend_counter_key,
project_cache_key,
project_spend_counter_key,
tag_cache_key,
)
from litellm.proxy.db.budget_window_spend_writer import roll_window_spend_row
@ -48,6 +50,7 @@ from litellm.proxy.db.exception_handler import call_with_db_reconnect_retry
from litellm.proxy.utils import PrismaClient, ProxyLogging
from litellm.repositories.organization_repository import OrganizationRepository
from litellm.repositories.prisma_protocols import SpendLinkedTable
from litellm.repositories.project_repository import ProjectRepository
from litellm.repositories.table_repositories import (
EndUserRepository,
ModelAccessGroupBudgetRepository,
@ -114,6 +117,11 @@ class _ModelAccessGroupRow(_BudgetLinkedRow, Protocol):
def access_group_name(self) -> str: ...
class _ProjectRow(_BudgetLinkedRow, Protocol):
@property
def project_id(self) -> str: ...
class _EndUserRow(_BudgetLinkedRow, Protocol):
@property
def user_id(self) -> str: ...
@ -184,6 +192,14 @@ def _model_access_group_cache_keys(row: _ModelAccessGroupRow) -> tuple[str, ...]
return (model_access_group_cache_key(row.access_group_name),)
def _project_counter_key(row: _ProjectRow) -> str:
return project_spend_counter_key(row.project_id)
def _project_cache_keys(row: _ProjectRow) -> tuple[str, ...]:
return (project_cache_key(row.project_id),)
def _enduser_counter_key(row: _EndUserRow) -> str:
return f"spend:end_user:{row.user_id}"
@ -754,6 +770,11 @@ class ResetBudgetJob:
where=_budget_link_where(budget_ids, _SPENT_ROWS_WHERE),
log_subject="model access groups",
)
projects: Final[tuple[_ProjectRow, ...]] = await self._fetch_linked_rows(
table=ProjectRepository(self.prisma_client).table,
where=_budget_link_where(budget_ids, _SPENT_ROWS_WHERE),
log_subject="projects",
)
rollover_caps: Final[Mapping[str, float]] = MappingProxyType(
{ # mutable-ok: MappingProxyType wraps a one-shot dict comprehension
b.budget_id: cap
@ -786,6 +807,7 @@ class ResetBudgetJob:
(_model_access_group_counter_key(row), _row_carried_spend(row, rollover_caps))
for row in model_access_groups
),
*((_project_counter_key(row), _row_carried_spend(row, rollover_caps)) for row in projects),
),
rollover_caps=rollover_caps,
cache_keys=(
@ -794,6 +816,7 @@ class ResetBudgetJob:
*(key for row in orgs for key in _org_cache_keys(row)),
*(key for row in tags for key in _tag_cache_keys(row)),
*(key for row in model_access_groups for key in _model_access_group_cache_keys(row)),
*(key for row in projects for key in _project_cache_keys(row)),
),
)
@ -820,6 +843,7 @@ class ResetBudgetJob:
_queue_budget_linked_resets(uow.organizations, cascade, extra=_SPENT_ROWS_WHERE)
_queue_budget_linked_resets(uow.tags, cascade, extra=_SPENT_ROWS_WHERE)
_queue_budget_linked_resets(uow.model_access_groups, cascade, extra=_SPENT_ROWS_WHERE)
_queue_budget_linked_resets(uow.projects, cascade, extra=_SPENT_ROWS_WHERE)
_queue_enduser_resets(uow.endusers, cascade)
for budget_id, budget_reset_at in cascade.budget_resets:
uow.budgets.queue_window_advance(budget_id=budget_id, budget_reset_at=budget_reset_at)

View file

@ -0,0 +1,165 @@
import time
from collections.abc import Mapping
from http import HTTPStatus
from types import MappingProxyType
from typing import Final
from pydantic import BaseModel, ConfigDict, field_validator
from litellm._logging import redact_internal_details_from_client_message
from litellm._uuid import uuid
from litellm.exceptions import MidStreamFallbackError
from litellm.types.llms.openai import ResponseFailedEvent, ResponsesAPIResponse, ResponsesAPIStreamEvents
class _ResponseIdentity(BaseModel):
model_config = ConfigDict(frozen=True, from_attributes=True)
id: str | None = None
model: str | None = None
created_at: int | None = None
class _StreamEvent(BaseModel):
model_config = ConfigDict(frozen=True, from_attributes=True)
type: str | None = None
sequence_number: int | None = None
response: _ResponseIdentity | None = None
class _FailureDetails(BaseModel):
model_config = ConfigDict(frozen=True, from_attributes=True)
message: str | None = None
code: str | int | None = None
type: str | None = None
status_code: int | None = None
@field_validator("message", mode="before")
@classmethod
def normalize_message(cls, value: object) -> str | None:
return value if isinstance(value, str) else None
@field_validator("code", mode="before")
@classmethod
def normalize_code(cls, value: object) -> str | int | None:
return value if isinstance(value, (str, int)) and not isinstance(value, bool) else None
@field_validator("type", mode="before")
@classmethod
def normalize_type(cls, value: object) -> str | None:
return value if isinstance(value, str) else None
def _original_failure(exception: Exception) -> Exception:
current = exception # rebind-ok: the recursion gate requires iterative wrapper traversal
while isinstance(current, MidStreamFallbackError) and current.original_exception is not None:
current = current.original_exception
return current
def _failure_details(original: Exception) -> _FailureDetails:
mapped: Final = _FailureDetails.model_validate(original)
body: Final = getattr(original, "body", None)
if not isinstance(body, Mapping):
return mapped
upstream: Final = _FailureDetails.model_validate(body)
return _FailureDetails(
message=upstream.message or mapped.message,
code=upstream.code if upstream.code is not None else mapped.code,
type=upstream.type or mapped.type,
status_code=mapped.status_code,
)
_CLIENT_ERROR_CODES: Final = MappingProxyType(
{
int(HTTPStatus.UNAUTHORIZED): "authentication_error",
int(HTTPStatus.FORBIDDEN): "permission_error",
int(HTTPStatus.NOT_FOUND): "not_found_error",
int(HTTPStatus.REQUEST_TIMEOUT): "request_timeout",
int(HTTPStatus.TOO_MANY_REQUESTS): "rate_limit_exceeded",
}
)
def _status_error_code(status_code: int | None) -> str:
if status_code is None or not HTTPStatus.BAD_REQUEST <= status_code < HTTPStatus.INTERNAL_SERVER_ERROR:
return "server_error"
return _CLIENT_ERROR_CODES.get(status_code, "invalid_request_error")
def _response_error_code(details: _FailureDetails) -> str:
for value in (details.code, details.type):
if value == "insufficient_quota":
return "insufficient_quota"
if value in (429, "429") or isinstance(value, str) and value.startswith("rate_limit"):
return "rate_limit_exceeded"
if isinstance(details.code, str) and details.code and not details.code.isdecimal():
return details.code
return _status_error_code(details.status_code)
class ResponsesStreamErrorState:
def __init__(self) -> None:
self.response_id: str | None = None
self.model: str | None = None
self.created_at: int | None = None
self.sequence_number = -1
self.terminal_emitted = False
self._pending_event: _StreamEvent | None = None
def observe_chunk(self, chunk: object) -> None:
self._pending_event = _StreamEvent.model_validate(chunk) if isinstance(chunk, (BaseModel, Mapping)) else None
def mark_emitted(self, frame: str | bytes) -> str | bytes:
event: Final = self._pending_event
if event is None:
return frame
if event.sequence_number is not None:
self.sequence_number = max(self.sequence_number, event.sequence_number)
if event.response is not None:
self.response_id = event.response.id or self.response_id
self.model = event.response.model or self.model
if event.response.created_at is not None:
self.created_at = event.response.created_at
if event.type in ("response.completed", "response.failed", "response.incomplete"):
self.terminal_emitted = True
return frame
def format_failure(self, exception: Exception) -> str | None:
if self.terminal_emitted:
return None
original: Final = _original_failure(exception)
details: Final = _failure_details(original)
response: Final = ResponsesAPIResponse.model_validate(
MappingProxyType(
{
"id": self.response_id or f"resp_{uuid.uuid4().hex}",
"object": "response",
"created_at": self.created_at if self.created_at is not None else int(time.time()),
"model": self.model,
"status": "failed",
"output": (),
"error": MappingProxyType(
{
"code": _response_error_code(details),
"message": redact_internal_details_from_client_message(details.message or str(original)),
}
),
}
)
)
event: Final = ResponseFailedEvent.model_validate(
MappingProxyType(
{
"type": ResponsesAPIStreamEvents.RESPONSE_FAILED,
"response": response,
"sequence_number": self.sequence_number + 1,
}
)
)
payload: Final = event.model_dump_json(exclude_none=True)
self.terminal_emitted = True
return f"event: response.failed\ndata: {payload}\n\n"

View file

@ -325,6 +325,14 @@ def model_access_group_spend_counter_key(access_group_name: str) -> str:
return f"spend:model_access_group:{access_group_name}"
def project_cache_key(project_id: str) -> str:
return f"project_id:{project_id}"
def project_spend_counter_key(project_id: str) -> str:
return f"spend:project:{project_id}"
#: Cached under ``end_user_restricted_registry_cache_key`` when the restricted set exceeds
#: ``END_USER_RESTRICTED_REGISTRY_MAX_SIZE``: registry unusable, fall back to the per-id fetch.
END_USER_RESTRICTED_REGISTRY_OVERFLOW_SENTINEL: Final = "__end_user_restricted_registry_overflow__"

View file

@ -15,7 +15,7 @@ import traceback
from collections.abc import Mapping, Sequence
from datetime import datetime, timedelta, timezone
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, Literal, Protocol, cast, overload
from typing import TYPE_CHECKING, Any, Final, Literal, Protocol, TypeVar, cast, overload
from urllib.parse import quote, unquote
from typing_extensions import LiteralString, ReadOnly, TypedDict
@ -31,6 +31,7 @@ from litellm.constants import (
from litellm.litellm_core_utils.litellm_logging import coerce_model_access_groups
from litellm.litellm_core_utils.safe_json_loads import safe_json_loads
from litellm.proxy._types import (
DB_CONNECTION_ERROR_TYPES,
DB_RETRY_SAFE_ERROR_TYPES,
BaseDailySpendTransaction,
DailyAgentSpendTransaction,
@ -46,6 +47,7 @@ from litellm.proxy._types import (
SpendUpdateQueueItem,
ToolDiscoveryQueueItem,
)
from litellm.proxy.common_utils.user_api_key_cache import project_cache_key
from litellm.proxy.db.daily_spend_bulk_upsert import (
DAILY_SPEND_TABLES,
build_bulk_upsert,
@ -64,6 +66,7 @@ from litellm.proxy.db.db_transaction_queue.window_spend_update_queue import (
WindowSpendTransaction,
WindowSpendUpdateQueue,
)
from litellm.proxy.db.exception_handler import PrismaDBExceptionHandler
from litellm.proxy.route_llm_request import ROUTE_ENDPOINT_MAPPING
from litellm.proxy.spend_tracking.compression_savings import (
extract_compression_saved_tokens,
@ -122,6 +125,7 @@ class _SpendBatch(Protocol):
litellm_teammembership: BatchTable
litellm_organizationtable: BatchTable
litellm_organizationmembership: BatchTable
litellm_projecttable: BatchTable
litellm_tagtable: BatchTable
litellm_agentstable: BatchTable
litellm_modelaccessgroupbudgettable: BatchTable
@ -145,6 +149,30 @@ class _SpendTransactionManager(Protocol):
async def __aexit__(self, exc_type: object, exc_value: object, traceback: object) -> bool | None: ...
_DailySpendTransactionT = TypeVar("_DailySpendTransactionT", bound=BaseDailySpendTransaction)
class _DailySpendCommit(Protocol[_DailySpendTransactionT]):
async def __call__(
self,
*,
n_retry_times: int,
prisma_client: PrismaClient,
proxy_logging_obj: ProxyLogging,
daily_spend_transactions: dict[str, _DailySpendTransactionT],
) -> None: ...
_DATA_REJECTED_SQLSTATE_CLASSES: Final = frozenset({"22", "23"})
def _daily_spend_commit_failure_is_requeue_safe(e: Exception) -> bool:
if isinstance(e, DB_CONNECTION_ERROR_TYPES):
return isinstance(e, DB_RETRY_SAFE_ERROR_TYPES)
sqlstate: Final = PrismaDBExceptionHandler.postgres_sqlstate(e)
return sqlstate is None or sqlstate[:2] not in _DATA_REJECTED_SQLSTATE_CLASSES
def _timed_request_duration_ms(
payload: dict | SpendLogsPayload,
request_status: Literal["success", "failure"],
@ -300,6 +328,7 @@ class DBSpendUpdateWriter:
start_time: datetime,
end_time: datetime,
response_cost: float | None,
project_id: str | None = None,
) -> bool:
"""Record the request's spend, answering whether its cost still needs charging.
@ -382,6 +411,7 @@ class DBSpendUpdateWriter:
hashed_token=hashed_token,
team_id=team_id,
org_id=org_id,
project_id=project_id,
end_user_id=end_user_id,
prisma_client=prisma_client,
litellm_proxy_budget_name=litellm_proxy_budget_name,
@ -678,6 +708,7 @@ class DBSpendUpdateWriter:
litellm_proxy_budget_name: str | None,
payload: SpendLogsPayload,
request_model_access_groups: Sequence[str] = (),
project_id: str | None = None,
):
"""
Runs all 13 spend-update helpers sequentially inside a single asyncio task.
@ -741,6 +772,18 @@ class DBSpendUpdateWriter:
traceback.format_exc(),
)
try:
await self._update_project_db(
response_cost=response_cost,
project_id=project_id,
prisma_client=prisma_client,
)
except Exception: # noqa: BLE001 # a project enqueue failure must not skip the sibling spend writes
verbose_proxy_logger.debug(
"_batch_database_updates: _update_project_db failed: %s",
traceback.format_exc(),
)
try:
await self._update_tag_db(
response_cost=response_cost,
@ -1003,6 +1046,32 @@ class DBSpendUpdateWriter:
)
raise e
async def _update_project_db(
self,
response_cost: float | None,
project_id: str | None,
prisma_client: PrismaClient | None,
) -> None:
if project_id is None or prisma_client is None:
return
try:
await self.spend_update_queue.add_update(
update=SpendUpdateQueueItem(
entity_type=Litellm_EntityType.PROJECT,
entity_id=project_id,
response_cost=response_cost,
)
)
except Exception as e:
spend_log_error(
"Spend tracking - failed to enqueue project spend update. project_id=%s, response_cost=%s - %s",
project_id,
response_cost,
str(e),
exc=e,
)
raise e
async def _update_agent_db(
self,
response_cost: float | None,
@ -1240,18 +1309,19 @@ class DBSpendUpdateWriter:
if db_spend_update_transactions is not None:
verbose_proxy_logger.info(
"Spend tracking - committing spend updates from Redis to DB: "
"keys=%d, users=%d, teams=%d, orgs=%d, end_users=%d, team_members=%d, org_members=%d, tags=%d, "
"agents=%d, model_access_groups=%d",
len(db_spend_update_transactions.get("key_list_transactions") or {}),
len(db_spend_update_transactions.get("user_list_transactions") or {}),
len(db_spend_update_transactions.get("team_list_transactions") or {}),
len(db_spend_update_transactions.get("org_list_transactions") or {}),
len(db_spend_update_transactions.get("end_user_list_transactions") or {}),
len(db_spend_update_transactions.get("team_member_list_transactions") or {}),
len(db_spend_update_transactions.get("org_member_list_transactions") or {}),
len(db_spend_update_transactions.get("tag_list_transactions") or {}),
len(db_spend_update_transactions.get("agent_list_transactions") or {}),
len(db_spend_update_transactions.get("model_access_group_list_transactions") or {}),
"keys=%d, users=%d, teams=%d, orgs=%d, end_users=%d, team_members=%d, org_members=%d, "
"projects=%d, tags=%d, agents=%d, model_access_groups=%d",
len(db_spend_update_transactions.get("key_list_transactions") or ()),
len(db_spend_update_transactions.get("user_list_transactions") or ()),
len(db_spend_update_transactions.get("team_list_transactions") or ()),
len(db_spend_update_transactions.get("org_list_transactions") or ()),
len(db_spend_update_transactions.get("end_user_list_transactions") or ()),
len(db_spend_update_transactions.get("team_member_list_transactions") or ()),
len(db_spend_update_transactions.get("org_member_list_transactions") or ()),
len(db_spend_update_transactions.get("project_list_transactions") or ()),
len(db_spend_update_transactions.get("tag_list_transactions") or ()),
len(db_spend_update_transactions.get("agent_list_transactions") or ()),
len(db_spend_update_transactions.get("model_access_group_list_transactions") or ()),
)
await self._commit_spend_updates_to_db(
prisma_client=prisma_client,
@ -1328,6 +1398,36 @@ class DBSpendUpdateWriter:
cronjob_id=DB_SPEND_UPDATE_JOB_NAME,
)
async def _flush_daily_spend_queue(
self,
queue: DailySpendUpdateQueue,
entity_type: Literal["user", "team", "org", "tag", "end_user", "agent"],
commit: _DailySpendCommit[_DailySpendTransactionT],
n_retry_times: int,
prisma_client: PrismaClient,
proxy_logging_obj: ProxyLogging,
) -> None:
transactions: Final = await queue.flush_and_get_aggregated_daily_spend_update_transactions()
try:
await commit(
n_retry_times=n_retry_times,
prisma_client=prisma_client,
proxy_logging_obj=proxy_logging_obj,
daily_spend_transactions=cast(dict[str, _DailySpendTransactionT], transactions),
)
except Exception as e: # noqa: BLE001 # whatever failed here, the other tables must still flush
if not transactions:
return
spend_log_error(
"Spend tracking - failed to commit daily %s spend updates. "
"Re-queued %d rows for retry on next tick. Error: %s",
entity_type,
len(transactions),
str(e),
exc=e,
)
await queue.add_update(transactions)
async def _commit_spend_updates_to_db_without_redis_buffer(
self,
prisma_client: PrismaClient,
@ -1356,74 +1456,59 @@ class DBSpendUpdateWriter:
################## Daily Spend Update Transactions ##################
# Aggregate all in memory daily spend transactions and commit to db
daily_spend_update_transactions: Final = cast(
dict[str, DailyUserSpendTransaction],
await self.daily_spend_update_queue.flush_and_get_aggregated_daily_spend_update_transactions(),
)
await DBSpendUpdateWriter.update_daily_user_spend(
await self._flush_daily_spend_queue(
queue=self.daily_spend_update_queue,
entity_type="user",
commit=DBSpendUpdateWriter.update_daily_user_spend,
n_retry_times=n_retry_times,
prisma_client=prisma_client,
proxy_logging_obj=proxy_logging_obj,
daily_spend_transactions=daily_spend_update_transactions,
)
################## Daily Team Spend Update Transactions ##################
# Aggregate all in memory daily team spend transactions and commit to db
daily_team_spend_update_transactions: Final = cast(
dict[str, DailyTeamSpendTransaction],
await self.daily_team_spend_update_queue.flush_and_get_aggregated_daily_spend_update_transactions(),
)
await DBSpendUpdateWriter.update_daily_team_spend(
await self._flush_daily_spend_queue(
queue=self.daily_team_spend_update_queue,
entity_type="team",
commit=DBSpendUpdateWriter.update_daily_team_spend,
n_retry_times=n_retry_times,
prisma_client=prisma_client,
proxy_logging_obj=proxy_logging_obj,
daily_spend_transactions=daily_team_spend_update_transactions,
)
################## Daily Organization Spend Update Transactions ##################
# Aggregate all in memory daily org spend transactions and commit to db
daily_org_spend_update_transactions: Final = cast(
dict[str, DailyOrganizationSpendTransaction],
await self.daily_org_spend_update_queue.flush_and_get_aggregated_daily_spend_update_transactions(),
)
await DBSpendUpdateWriter.update_daily_org_spend(
await self._flush_daily_spend_queue(
queue=self.daily_org_spend_update_queue,
entity_type="org",
commit=DBSpendUpdateWriter.update_daily_org_spend,
n_retry_times=n_retry_times,
prisma_client=prisma_client,
proxy_logging_obj=proxy_logging_obj,
daily_spend_transactions=daily_org_spend_update_transactions,
)
# NOTE: Daily tag spend is committed by a separate scheduler job.
################## Daily End-User Spend Update Transactions ##################
# Aggregate all in memory daily end-user spend transactions and commit to db
daily_end_user_spend_update_transactions: Final = cast(
dict[str, DailyEndUserSpendTransaction],
await self.daily_end_user_spend_update_queue.flush_and_get_aggregated_daily_spend_update_transactions(),
)
await DBSpendUpdateWriter.update_daily_end_user_spend(
await self._flush_daily_spend_queue(
queue=self.daily_end_user_spend_update_queue,
entity_type="end_user",
commit=DBSpendUpdateWriter.update_daily_end_user_spend,
n_retry_times=n_retry_times,
prisma_client=prisma_client,
proxy_logging_obj=proxy_logging_obj,
daily_spend_transactions=daily_end_user_spend_update_transactions,
)
################## Daily Agent Spend Update Transactions ##################
# Aggregate all in memory daily agent spend transactions and commit to db
daily_agent_spend_update_transactions: Final = cast(
dict[str, DailyAgentSpendTransaction],
await self.daily_agent_spend_update_queue.flush_and_get_aggregated_daily_spend_update_transactions(),
)
await DBSpendUpdateWriter.update_daily_agent_spend(
await self._flush_daily_spend_queue(
queue=self.daily_agent_spend_update_queue,
entity_type="agent",
commit=DBSpendUpdateWriter.update_daily_agent_spend,
n_retry_times=n_retry_times,
prisma_client=prisma_client,
proxy_logging_obj=proxy_logging_obj,
daily_spend_transactions=daily_agent_spend_update_transactions,
)
################## Budget Window Spend Update Transactions ##################
@ -1460,19 +1545,15 @@ class DBSpendUpdateWriter:
Commit only tag spend updates to database.
This is called by a separate scheduler job at a longer interval.
"""
daily_tag_spend_update_transactions: Final = cast(
dict[str, DailyTagSpendTransaction],
await self.daily_tag_spend_update_queue.flush_and_get_aggregated_daily_spend_update_transactions(),
await self._flush_daily_spend_queue(
queue=self.daily_tag_spend_update_queue,
entity_type="tag",
commit=DBSpendUpdateWriter.update_daily_tag_spend,
n_retry_times=n_retry_times,
prisma_client=prisma_client,
proxy_logging_obj=proxy_logging_obj,
)
if daily_tag_spend_update_transactions:
await DBSpendUpdateWriter.update_daily_tag_spend(
n_retry_times=n_retry_times,
prisma_client=prisma_client,
proxy_logging_obj=proxy_logging_obj,
daily_spend_transactions=daily_tag_spend_update_transactions,
)
async def _commit_daily_tag_spend_to_db_with_redis(
self,
prisma_client: PrismaClient,
@ -1797,6 +1878,22 @@ class DBSpendUpdateWriter:
proxy_logging_obj=proxy_logging_obj,
)
### UPDATE PROJECT TABLE ###
project_list_transactions: Final = db_spend_update_transactions.get("project_list_transactions")
await DBSpendUpdateWriter._update_entity_spend_in_db(
entity_name="Project",
transactions=project_list_transactions,
table_accessor="litellm_projecttable",
where_field="project_id",
n_retry_times=n_retry_times,
prisma_client=prisma_client,
proxy_logging_obj=proxy_logging_obj,
)
await DBSpendUpdateWriter._invalidate_project_caches(
project_ids=tuple(project_list_transactions or ()),
proxy_logging_obj=proxy_logging_obj,
)
### UPDATE TAG TABLE ###
tag_list_transactions: Final = db_spend_update_transactions["tag_list_transactions"]
await DBSpendUpdateWriter._update_entity_spend_in_db(
@ -1835,11 +1932,23 @@ class DBSpendUpdateWriter:
proxy_logging_obj=proxy_logging_obj,
)
@staticmethod
async def _invalidate_project_caches(project_ids: Sequence[str], proxy_logging_obj: ProxyLogging | None) -> None:
if not project_ids or proxy_logging_obj is None:
return
user_api_key_cache: Final = proxy_logging_obj.call_details.get("user_api_key_cache")
if user_api_key_cache is None:
return
for project_id in project_ids:
await user_api_key_cache.async_delete_cache(key=project_cache_key(project_id))
@staticmethod
async def _update_entity_spend_in_db(
entity_name: str,
transactions: dict[str, float] | None,
table_accessor: Literal["litellm_tagtable", "litellm_agentstable", "litellm_modelaccessgroupbudgettable"],
table_accessor: Literal[
"litellm_tagtable", "litellm_agentstable", "litellm_modelaccessgroupbudgettable", "litellm_projecttable"
],
where_field: str,
n_retry_times: int,
prisma_client: PrismaClient,
@ -2031,13 +2140,25 @@ class DBSpendUpdateWriter:
sql, params = build_bulk_upsert(table=table, batch=merged_batch)
await prisma_client.db.execute_raw(sql, *params)
except Exception as batch_error:
# Log detailed error information for debugging batch upsert failures
# This helps diagnose issues like unique constraint violations
if _daily_spend_commit_failure_is_requeue_safe(batch_error):
spend_log_error(
"Daily %s spend batch upsert failed. Table: %s, Rows: %d, Error: %s",
entity_type,
table.name,
len(transactions_to_process),
str(batch_error),
exc=batch_error,
)
raise
for key in transactions_to_process:
daily_spend_transactions.pop(key, None)
spend_log_error(
"Daily %s spend batch upsert failed. Table: %s, Rows: %d, Error: %s",
"Spend tracking - dropped %d daily %s spend rows: the failed statement may have "
"applied or the database refused the data, so re-sending it is not safe. "
"Table: %s, Error: %s",
len(transactions_to_process),
entity_type,
table.name,
len(transactions_to_process),
str(batch_error),
exc=batch_error,
)

View file

@ -70,6 +70,7 @@ _SpendTransactionField: TypeAlias = Literal[
"team_member_list_transactions",
"org_list_transactions",
"org_member_list_transactions",
"project_list_transactions",
"tag_list_transactions",
"agent_list_transactions",
"model_access_group_list_transactions",
@ -83,6 +84,7 @@ _SPEND_TRANSACTION_FIELDS: Final[tuple[_SpendTransactionField, ...]] = (
"team_member_list_transactions",
"org_list_transactions",
"org_member_list_transactions",
"project_list_transactions",
"tag_list_transactions",
"agent_list_transactions",
"model_access_group_list_transactions",
@ -418,6 +420,10 @@ class RedisUpdateBuffer:
Litellm_EntityType.ORGANIZATION_MEMBER,
db_spend_update_transactions.get("org_member_list_transactions"),
),
(
Litellm_EntityType.PROJECT,
db_spend_update_transactions.get("project_list_transactions"),
),
(
Litellm_EntityType.TAG,
db_spend_update_transactions.get("tag_list_transactions"),
@ -885,6 +891,7 @@ class RedisUpdateBuffer:
org_member_list_transactions=_merged_entity_transactions(
list_of_transactions, "org_member_list_transactions"
),
project_list_transactions=_merged_entity_transactions(list_of_transactions, "project_list_transactions"),
tag_list_transactions=_merged_entity_transactions(list_of_transactions, "tag_list_transactions"),
agent_list_transactions=_merged_entity_transactions(list_of_transactions, "agent_list_transactions"),
model_access_group_list_transactions=_merged_entity_transactions(

View file

@ -138,6 +138,7 @@ class SpendUpdateQueue(BaseUpdateQueue):
team_member_list_transactions={},
org_list_transactions={},
org_member_list_transactions={},
project_list_transactions={},
tag_list_transactions={},
agent_list_transactions={},
model_access_group_list_transactions={},
@ -152,6 +153,7 @@ class SpendUpdateQueue(BaseUpdateQueue):
Litellm_EntityType.TEAM_MEMBER: "team_member_list_transactions",
Litellm_EntityType.ORGANIZATION: "org_list_transactions",
Litellm_EntityType.ORGANIZATION_MEMBER: "org_member_list_transactions",
Litellm_EntityType.PROJECT: "project_list_transactions",
Litellm_EntityType.TAG: "tag_list_transactions",
Litellm_EntityType.AGENT: "agent_list_transactions",
Litellm_EntityType.MODEL_ACCESS_GROUP: "model_access_group_list_transactions",
@ -192,6 +194,8 @@ class SpendUpdateQueue(BaseUpdateQueue):
transactions_dict = db_spend_update_transactions["org_list_transactions"]
elif dict_key == "org_member_list_transactions":
transactions_dict = db_spend_update_transactions["org_member_list_transactions"]
elif dict_key == "project_list_transactions":
transactions_dict = db_spend_update_transactions["project_list_transactions"]
elif dict_key == "tag_list_transactions":
transactions_dict = db_spend_update_transactions["tag_list_transactions"]
elif dict_key == "agent_list_transactions":

View file

@ -1,6 +1,8 @@
from collections.abc import Awaitable, Callable, Iterator
from typing import Any, Final, TypeVar
from pydantic import TypeAdapter, ValidationError
from litellm._logging import verbose_proxy_logger
from litellm.proxy._types import (
DB_CONNECTION_ERROR_TYPES,
@ -17,6 +19,8 @@ _TRANSIENT_DB_UNAVAILABLE_MESSAGE: Final = (
"Service Unavailable, the authentication database is temporarily unreachable. Please retry shortly."
)
_DATABASE_ERROR_META: Final = TypeAdapter(dict[str, object])
def _exception_chain(e: BaseException) -> Iterator[BaseException]:
current = e # rebind-ok: advances one link per iteration of the bounded walk
@ -221,6 +225,20 @@ class PrismaDBExceptionHandler:
or "write conflict or a deadlock" in error_message
)
@staticmethod
def postgres_sqlstate(e: Exception) -> str | None:
"""The SQLSTATE Postgres attached to a failed statement, as prisma surfaces it, or None."""
import prisma
if not isinstance(e, _exception_types(prisma.errors.DataError)):
return None
try:
meta: Final = _DATABASE_ERROR_META.validate_python(getattr(e, "meta", None))
except ValidationError:
return None
code: Final = meta.get("code")
return code if isinstance(code, str) else None
@staticmethod
def is_read_only_transaction_error(e: Exception) -> bool:
"""True iff ``e`` is Postgres SQLSTATE 25006 surfaced through prisma: the

View file

@ -26,6 +26,7 @@ from litellm.proxy._types import Litellm_EntityType
from litellm.proxy.db.db_lookup_gate import db_lookup_gate
from litellm.proxy.spend_tracking.spend_counter_batch import read_batched_spend_counter, record_spend_counter_value
from litellm.repositories.organization_repository import OrganizationRepository
from litellm.repositories.project_repository import ProjectRepository
from litellm.repositories.table_repositories import (
BudgetWindowSpendRepository,
EndUserRepository,
@ -77,6 +78,7 @@ class SpendCounterReseed:
spend:team_member:{uid}:{tid} -> LiteLLM_TeamMembership.spend
spend:user:{user_id} -> LiteLLM_UserTable.spend
spend:org:{org_id} -> LiteLLM_OrganizationTable.spend
spend:project:{project_id} -> LiteLLM_ProjectTable.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()
@ -157,6 +159,9 @@ class SpendCounterReseed:
row = await OrganizationRepository(prisma_client).table.find_unique(
where={"organization_id": org_id}
)
elif counter_key.startswith("spend:project:"):
project_id: Final = counter_key[len("spend:project:") :]
row = await ProjectRepository(prisma_client).table.find_unique(where={"project_id": project_id})
else:
return None
except Exception:

View file

@ -267,6 +267,7 @@ class _ProxyDBLogger(CustomLogger):
start_time=actual_start_time,
end_time=datetime.now(),
org_id=user_api_key_dict.org_id,
project_id=user_api_key_dict.project_id,
)
@log_db_metrics
@ -318,6 +319,11 @@ class _ProxyDBLogger(CustomLogger):
user_id: Final = cast(str | None, metadata.get("user_api_key_user_id", None))
team_id: Final = cast(str | None, metadata.get("user_api_key_team_id", None))
org_id: Final = cast(str | None, metadata.get("user_api_key_org_id", None))
project_id: Final = (
project_id_value
if isinstance(project_id_value := metadata.get("user_api_key_project_id"), str)
else None
)
key_alias: Final = cast(str | None, metadata.get("user_api_key_alias", None))
end_user_max_budget: Final = metadata.get("user_api_end_user_max_budget", None)
sl_object: Final[StandardLoggingPayload | None] = kwargs.get("standard_logging_object", None)
@ -368,6 +374,7 @@ class _ProxyDBLogger(CustomLogger):
budget_reservation=budget_reservation,
request_tags=tags,
model_access_groups=model_access_groups,
project_id=project_id,
)
if not charged:
return
@ -501,6 +508,8 @@ class _ProxyDBLogger(CustomLogger):
metadata["user_api_key_team_id"] = key_obj.team_id
if metadata.get("user_api_key_org_id") is None:
metadata["user_api_key_org_id"] = key_obj.org_id
if metadata.get("user_api_key_project_id") is None:
metadata["user_api_key_project_id"] = key_obj.project_id
except Exception:
verbose_proxy_logger.debug(
"Failed to enrich failure metadata with key info for api_key=%s",
@ -651,6 +660,7 @@ async def _update_database_and_spend_counters(
budget_reservation: dict | None,
request_tags: list[str] | None = None,
model_access_groups: Sequence[str] | None = None,
project_id: str | None = None,
) -> bool:
if budget_reservation is not None:
await _reconcile_budget_reservation_before_db_update(
@ -668,6 +678,7 @@ async def _update_database_and_spend_counters(
start_time=start_time,
end_time=end_time,
org_id=org_id,
project_id=project_id,
)
except Exception:
if budget_reservation is not None:
@ -698,6 +709,7 @@ async def _update_database_and_spend_counters(
tags=request_tags,
request_started_at=start_time,
model_access_groups=model_access_groups,
project_id=project_id,
)
except Exception:
if budget_reservation is not None:

View file

@ -293,9 +293,9 @@ async def _clone_team_default_budget_for_member(
member budget. Returns the new budget_id, or None if the default budget
no longer exists in the DB.
Used when adding a new team member without an explicit per-member budget,
so the member starts with the team default's values but gets their own
private budget row (which can be edited independently).
Used when adding a new team member with a per-member ``budget_duration``
but no other per-member limit, so the member keeps the team default's
values in their own private budget row while the reset window differs.
``budget_duration_override`` replaces the default's reset window for this
member while keeping the default's other limits, so an admin can set a
@ -346,14 +346,21 @@ async def _resolve_member_budget_id(
"""
Resolve the budget a new team member should be linked to.
Explicit per-member limits create a fresh budget. Otherwise the team's
default member budget is cloned (with ``budget_duration`` overriding its
reset window while keeping its other limits). A lone ``budget_duration``
with no team default creates a window-only budget. With nothing set the
member gets no budget, though ``add_new_member`` still writes its membership row.
Explicit per-member limits create a fresh budget. Otherwise the member is
linked to the team's shared default member budget, so later ``/team/update``
changes reach them; ``/team/member_update`` clones that row on first write.
A lone ``budget_duration`` clones the default with the reset window
overridden, or creates a window-only budget when there is no team default.
With nothing set the member gets no budget, though ``add_new_member`` still writes its membership row.
"""
has_explicit_limit: Final = max_budget_in_team is not None or allowed_models is not None
if not has_explicit_limit and default_team_budget_id is not None and budget_duration is None:
default_budget: Final = await _budget_table(prisma_client, tx).find_unique(
where={"budget_id": default_team_budget_id}
)
return default_team_budget_id if default_budget is not None else None
if not has_explicit_limit and default_team_budget_id is not None:
return await _clone_team_default_budget_for_member(
prisma_client=prisma_client,

View file

@ -421,6 +421,7 @@ from litellm.proxy.common_utils.periodic_reload_schedule import (
)
from litellm.proxy.common_utils.proxy_state import ProxyState
from litellm.proxy.common_utils.reset_budget_job import ResetBudgetJob
from litellm.proxy.common_utils.responses_stream_errors import ResponsesStreamErrorState
from litellm.proxy.common_utils.scheduled_job_stagger import (
apply_scheduled_job_stagger,
attach_job_timing_logger,
@ -438,6 +439,8 @@ from litellm.proxy.common_utils.user_api_key_cache import (
get_management_object_ttl,
model_access_group_cache_key,
model_access_group_spend_counter_key,
project_cache_key,
project_spend_counter_key,
tag_cache_key,
)
from litellm.proxy.config_resolvers import SettingsStore, resolve_fields
@ -2836,6 +2839,7 @@ async def increment_spend_counters(
tags: list[str] | None = None,
request_started_at: datetime | None = None,
model_access_groups: Sequence[str] | None = None,
project_id: str | None = None,
):
"""
Atomically increment spend counters for budget enforcement.
@ -2857,6 +2861,7 @@ async def increment_spend_counters(
end_user_id=end_user_id,
tags=tags,
model_access_groups=model_access_groups,
project_id=project_id,
),
):
await _increment_spend_counters_batched(
@ -2870,6 +2875,7 @@ async def increment_spend_counters(
tags=tags,
request_started_at=request_started_at,
model_access_groups=model_access_groups,
project_id=project_id,
)
@ -2884,6 +2890,7 @@ async def _increment_spend_counters_batched(
tags: list[str] | None,
request_started_at: datetime | None,
model_access_groups: Sequence[str] | None,
project_id: str | None = None,
):
"""Runs inside one spend counter batch: the reservation reconcile and the warm checks share a single MGET."""
reserved_counter_keys: Final = await _reconcile_budget_reservation_for_counter_update(
@ -3084,6 +3091,13 @@ async def _increment_spend_counters_batched(
)
if org_id is not None
else None,
_prepare_project_spend_increment(
project_id=project_id,
response_cost=cost,
reserved_counter_keys=reserved_counter_keys,
)
if project_id is not None
else None,
)
if coro is not None
)
@ -3236,6 +3250,23 @@ async def _prepare_org_spend_increment(
return (pending,) if pending is not None else ()
async def _prepare_project_spend_increment(
project_id: str | None,
response_cost: float,
reserved_counter_keys: set[str],
) -> tuple[PendingSpendIncrement, ...]:
if project_id is None:
return ()
pending: Final = await _prepare_unreserved_spend_counter_increment(
counter_key=project_spend_counter_key(project_id),
source_cache_key=project_cache_key(project_id),
increment=response_cost,
reserved_counter_keys=reserved_counter_keys,
)
return (pending,) if pending is not None else ()
async def _prepare_unreserved_spend_counter_increment(
counter_key: str,
source_cache_key: str | list[str],
@ -6889,9 +6920,7 @@ class ProxyConfig:
self._add_callbacks_from_db_config(config_data)
# router settings
await self._add_router_settings_from_db_config(
config_data=config_data, llm_router=llm_router, prisma_client=prisma_client
)
await self._add_router_settings_from_db_config(llm_router=llm_router, prisma_client=prisma_client)
return still_desired_ids
@ -7099,13 +7128,11 @@ class ProxyConfig:
async def _add_router_settings_from_db_config(
self,
config_data: Mapping[str, object],
llm_router: Router | None,
prisma_client: PrismaClient | None,
) -> None:
if llm_router is None or prisma_client is None:
return
self.router_settings.load_yaml(_as_settings_mapping(config_data.get("router_settings")))
db_router_settings: Final[_ConfigParamRow | None] = await _config_param_table(prisma_client).find_first(
where={"param_name": "router_settings"}
)
@ -8868,6 +8895,7 @@ def _format_streaming_sse_chunk(chunk: str | bytes) -> str | bytes:
_SSE_FRAME_DELIMITERS: Final = ("\r\n\r\n", "\n\n", "\r\r")
_OPENAI_STREAM_DONE_FRAME: Final = "data: [DONE]\n\n"
_MAX_RAW_SSE_BUFFER_CHARS: Final = 8 * 1024 * 1024
@ -9092,10 +9120,13 @@ async def async_data_generator(
user_api_key_dict: UserAPIKeyAuth,
request_data: dict,
request: Request | None = None,
*,
responses_stream_errors: bool = False,
):
verbose_proxy_logger.debug("inside generator")
stream_completed = False
client_disconnected = False
error_state: Final = ResponsesStreamErrorState() if responses_stream_errors else None
try:
error_message: str | None = None
requested_model_from_client: Final = _get_client_requested_model_for_streaming(request_data=request_data)
@ -9206,6 +9237,8 @@ async def async_data_generator(
fallback_metadata_event_sent = True
continue
if error_state is not None:
error_state.observe_chunk(cast(object, chunk)) # cast-ok: the helper validates legacy untyped chunks
raw_passthrough = False
if isinstance(chunk, BaseModel):
chunk = _serialize_streaming_chunk(chunk)
@ -9240,8 +9273,13 @@ async def async_data_generator(
if not raw_passthrough:
try:
yield _format_streaming_sse_chunk(chunk=chunk)
if error_state is not None:
yield error_state.mark_emitted(_format_streaming_sse_chunk(chunk=chunk))
else:
yield _format_streaming_sse_chunk(chunk=chunk)
except Exception as e:
if error_state is not None:
raise
yield f"data: {e}\n\n"
if pending_fallback_event:
@ -9265,8 +9303,7 @@ async def async_data_generator(
yield error_message
# OpenAI-compatible streams terminate with data: [DONE]; Google GenAI (?alt=sse) does not.
if not request_data.get("_litellm_skip_openai_stream_done"):
done_message: Final = "[DONE]"
yield f"data: {done_message}\n\n"
yield _OPENAI_STREAM_DONE_FRAME
except (asyncio.CancelledError, GeneratorExit):
# Client disconnected mid-stream. CancelledError / GeneratorExit are
# BaseException, so they bypass the success/failure logging callbacks
@ -9291,6 +9328,14 @@ async def async_data_generator(
e,
)
if error_state is not None:
stream_completed = True
error_frame: Final = error_state.format_failure(e)
if error_frame is not None:
yield error_frame
if not request_data.get("_litellm_skip_openai_stream_done"):
yield _OPENAI_STREAM_DONE_FRAME
return
if isinstance(e, HTTPException):
raise e
elif isinstance(e, StreamingCallbackError):
@ -9327,12 +9372,15 @@ def select_data_generator(
user_api_key_dict: UserAPIKeyAuth,
request_data: dict,
request: Request | None = None,
*,
responses_stream_errors: bool = False,
):
return async_data_generator(
response=response,
user_api_key_dict=user_api_key_dict,
request_data=request_data,
request=request,
responses_stream_errors=responses_stream_errors,
)
@ -17000,6 +17048,48 @@ async def update_config(
if prisma_client is None:
raise Exception("No DB Connected")
requested_general_settings: Final[Mapping[str, JsonValue]] = (
config_info.general_settings.model_dump(exclude_none=True, exclude_unset=True)
if config_info.general_settings is not None
else {}
)
raw_litellm_settings: Final[Mapping[str, JsonValue]] = _CONFIG_SECTION_VALUES.validate_python(
config_info.litellm_settings if config_info.litellm_settings is not None else {}
)
incoming_success_callback: Final = raw_litellm_settings.get("success_callback")
updated_litellm_settings: Final[Mapping[str, JsonValue]] = _CONFIG_SECTION_VALUES.validate_python(
{
**raw_litellm_settings,
**(
{"success_callback": normalize_callback_names(incoming_success_callback)}
if isinstance(incoming_success_callback, list)
else {}
),
}
)
typed_router_settings: Final[Mapping[str, JsonValue]] = (
config_info.router_settings.model_dump(exclude_none=True, exclude_unset=True)
if config_info.router_settings is not None
else {}
)
router_settings_updates: Final[Mapping[str, JsonValue]] = {
**typed_router_settings,
**(
{
key: value
for key, value in raw_router_settings.items()
if key not in typed_router_settings and value is not None
}
if isinstance(raw_router_settings, dict)
else {}
),
}
proxy_config.reject_config_owned_writes(
section_name="general_settings", changed_keys=requested_general_settings
)
proxy_config.reject_config_owned_writes(section_name="litellm_settings", changed_keys=raw_litellm_settings)
proxy_config.reject_config_owned_writes(section_name="router_settings", changed_keys=router_settings_updates)
async def _read_section(param_name: str) -> dict:
row: Final[_ConfigParamRow | None] = await _config_param_table(prisma_client).find_first(
where={"param_name": param_name}
@ -17026,8 +17116,7 @@ async def update_config(
if config_info.general_settings is not None:
existing = await _read_section("general_settings")
before_general_settings: Final = copy.deepcopy(existing)
updates: Mapping[str, JsonValue] = config_info.general_settings.dict(exclude_none=True)
for k, v in updates.items():
for k, v in requested_general_settings.items():
if k == "alert_to_webhook_url":
if "alerting" not in existing:
existing["alerting"] = ["slack"]
@ -17070,15 +17159,9 @@ async def update_config(
if config_info.litellm_settings is not None:
existing = await _read_section("litellm_settings")
before_litellm_settings: Final = copy.deepcopy(existing)
updated_litellm_settings: Final = dict(config_info.litellm_settings)
incoming_cb = updated_litellm_settings.get("success_callback")
if isinstance(incoming_cb, list):
updated_litellm_settings["success_callback"] = normalize_callback_names(incoming_cb)
merged: Final = {**existing, **updated_litellm_settings}
incoming_cb = updated_litellm_settings.get("success_callback")
incoming_cb: Final = updated_litellm_settings.get("success_callback")
existing_cb: Final = existing.get("success_callback")
if isinstance(incoming_cb, list):
if isinstance(existing_cb, list):
@ -17101,15 +17184,6 @@ async def update_config(
if isinstance(raw_router_settings, dict):
existing = await _read_section("router_settings")
before_router_settings: Final = copy.deepcopy(existing)
typed_router_settings: Final = (
config_info.router_settings.dict(exclude_none=True) if config_info.router_settings is not None else {}
)
raw_router_settings_without_none: Final = {
key: value
for key, value in raw_router_settings.items()
if key not in typed_router_settings and value is not None
}
router_settings_updates: Final = {**typed_router_settings, **raw_router_settings_without_none}
new_router_settings: Final = {**existing, **router_settings_updates}
await _upsert_section("router_settings", new_router_settings)
asyncio.create_task(
@ -17175,6 +17249,8 @@ _GENERAL_SETTINGS_CONFIG_LIST_FIELD_TYPES: Final[Mapping[str, str]] = MappingPro
"maximum_spend_logs_cleanup_run_budget": "String",
"maximum_spend_logs_cleanup_batch_timeout": "String",
"mcp_internal_ip_ranges": "List",
"mcp_allowed_clients": "TypedDictionary",
"mcp_client_id_header": "String",
"mcp_trusted_proxy_ranges": "List",
"mcp_xff_num_trusted_hops": "Integer",
"always_include_stream_usage": "Boolean",

View file

@ -3,6 +3,7 @@ import json
import time
from collections.abc import AsyncIterator, Awaitable, Mapping
from enum import Enum
from functools import partial
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, NamedTuple, Protocol, cast, get_args
from uuid import uuid4
@ -243,6 +244,7 @@ async def responses_api(
version,
)
native_data_generator: Final = partial(select_data_generator, responses_stream_errors=True)
data = await _read_request_body(request=request)
# Check if polling via cache should be used for this request
@ -329,7 +331,7 @@ async def responses_api(
llm_router=llm_router,
proxy_config=proxy_config,
proxy_logging_obj=proxy_logging_obj,
select_data_generator=select_data_generator,
select_data_generator=native_data_generator,
user_model=user_model,
user_temperature=user_temperature,
user_request_timeout=user_request_timeout,
@ -355,7 +357,7 @@ async def responses_api(
llm_router=llm_router,
general_settings=general_settings,
proxy_config=proxy_config,
select_data_generator=select_data_generator,
select_data_generator=native_data_generator,
model=None,
user_model=user_model,
user_temperature=user_temperature,

View file

@ -75,6 +75,10 @@ class _StreamEventParser:
parse: Callable[[str], _StreamEvent] = staticmethod(json.loads)
def _sse_frame_data(frame: str) -> str | None:
return next((line[6:].strip() for line in frame.splitlines() if line.startswith("data: ")), None)
async def _never_receive() -> Message:
await asyncio.Event().wait()
raise AssertionError("unreachable")
@ -224,8 +228,7 @@ async def background_streaming_task(
if isinstance(chunk, bytes):
chunk = chunk.decode("utf-8")
if isinstance(chunk, str) and chunk.startswith("data: "):
chunk_data = chunk[6:].strip()
if isinstance(chunk, str) and (chunk_data := _sse_frame_data(chunk)) is not None:
if chunk_data == "[DONE]":
break

View file

@ -2,6 +2,7 @@ from __future__ import annotations
import asyncio
import json
import math
from collections.abc import Mapping, Sequence
from dataclasses import dataclass
from datetime import datetime, timedelta, timezone
@ -29,6 +30,8 @@ from litellm.proxy.common_utils.user_api_key_cache import (
end_user_cache_key,
model_access_group_cache_key,
model_access_group_spend_counter_key,
project_cache_key,
project_spend_counter_key,
tag_cache_key,
team_membership_reservation_cache_key,
)
@ -62,6 +65,7 @@ _COUNTER_ENTITY_TYPES: Final[Mapping[str, str]] = {
"Tag": Litellm_EntityType.TAG.value,
"Model access group": Litellm_EntityType.MODEL_ACCESS_GROUP.value,
"Organization": Litellm_EntityType.ORGANIZATION.value,
"Project": Litellm_EntityType.PROJECT.value,
}
@ -542,6 +546,13 @@ async def _get_budget_counters(
if org_counter is not None:
counters.append(org_counter)
project_counter: Final = await _get_project_budget_counter(
valid_token=valid_token,
user_api_key_cache=user_api_key_cache,
)
if project_counter is not None:
counters.append(project_counter)
return counters
@ -757,6 +768,36 @@ async def _get_org_budget_counter(
)
async def _get_project_budget_counter(
valid_token: UserAPIKeyAuth,
user_api_key_cache: UserApiKeyCache,
) -> _BudgetCounter | None:
if valid_token.project_id is None:
return None
source_cache_key: Final = project_cache_key(valid_token.project_id)
project_object: Final = await user_api_key_cache.async_get_cache(key=source_cache_key)
if project_object is None:
return None
project_budget_table: Final = _get_value(project_object, "litellm_budget_table")
if project_budget_table is None:
return None
project_max_budget: Final = _to_float(_get_value(project_budget_table, "max_budget"))
if project_max_budget is None or project_max_budget <= 0 or not math.isfinite(project_max_budget):
return None
return _BudgetCounter(
counter_key=project_spend_counter_key(valid_token.project_id),
source_cache_key=source_cache_key,
max_budget=project_max_budget,
fallback_spend=_to_float(_get_value(project_object, "spend")) or 0.0,
entity_type="Project",
entity_id=valid_token.project_id,
)
def _get_budget_limit_counters(
entity_prefix: str,
entity_type: str,

View file

@ -12,7 +12,10 @@ from pydantic import TypeAdapter
from litellm._logging import verbose_proxy_logger
from litellm.caching.redis_cache import RedisCache
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.common_utils.user_api_key_cache import model_access_group_spend_counter_key
from litellm.proxy.common_utils.user_api_key_cache import (
model_access_group_spend_counter_key,
project_spend_counter_key,
)
_CounterValues: Final = TypeAdapter(dict[str, float | None])
_NO_VALUES: Final[Mapping[str, float | None]] = MappingProxyType({})
@ -154,6 +157,8 @@ def _iter_admission_counter_keys(token: UserAPIKeyAuth, end_user_id: str | None)
yield f"spend:end_user:{end_user_id}"
if token.org_id is not None:
yield f"spend:org:{token.org_id}"
if token.project_id is not None:
yield project_spend_counter_key(token.project_id)
def admission_counter_keys(token: UserAPIKeyAuth, end_user_id: str | None) -> frozenset[str]:
@ -168,10 +173,12 @@ def post_call_counter_keys(
end_user_id: str | None,
tags: Sequence[object] | None,
model_access_groups: Sequence[object] | None,
project_id: str | None = None,
) -> frozenset[str]:
"""Every counter ``increment_spend_counters`` warm-checks, except budget windows which bind on read."""
entity_keys: Final = admission_counter_keys(
UserAPIKeyAuth(token=token, team_id=team_id, user_id=user_id, org_id=org_id), end_user_id
UserAPIKeyAuth(token=token, team_id=team_id, user_id=user_id, org_id=org_id, project_id=project_id),
end_user_id,
)
tag_keys: Final = frozenset(f"spend:tag:{tag}" for tag in tags or () if tag and isinstance(tag, str))
group_keys: Final = frozenset(

View file

@ -152,4 +152,7 @@ class PrismaBatch(Protocol):
@property
def litellm_modelaccessgroupbudgettable(self) -> BatchTable: ...
@property
def litellm_projecttable(self) -> BatchTable: ...
async def commit(self) -> None: ...

View file

@ -109,6 +109,7 @@ class BudgetCascadeUnitOfWork:
organizations: LinkedSpendResetWrites
tags: LinkedSpendResetWrites
model_access_groups: LinkedSpendResetWrites
projects: LinkedSpendResetWrites
endusers: LinkedSpendResetWrites
budgets: BudgetWindowWrites
@ -135,6 +136,7 @@ async def budget_cascade_unit_of_work(
organizations=LinkedSpendResetWrites(table=batch.litellm_organizationtable),
tags=LinkedSpendResetWrites(table=batch.litellm_tagtable),
model_access_groups=LinkedSpendResetWrites(table=batch.litellm_modelaccessgroupbudgettable),
projects=LinkedSpendResetWrites(table=batch.litellm_projecttable),
endusers=LinkedSpendResetWrites(table=batch.litellm_endusertable),
budgets=BudgetWindowWrites(table=batch.litellm_budgettable),
)

View file

@ -61,6 +61,7 @@ from litellm.types.llms.openai import (
ChatCompletionToolParamFunctionChunk,
ChatCompletionUserMessage,
GenericChatCompletionMessage,
IncompleteDetails,
InputTokensDetails,
OpenAIChatCompletionTextObject,
OpenAIMcpServerTool,
@ -111,6 +112,9 @@ ResponseTools: TypeAlias = Sequence[Mapping[str, object]] | None
ChatToolParam: TypeAlias = ChatCompletionToolParam | OpenAIMcpServerTool
NAMESPACE_DESCRIPTION_SEPARATOR: Final = "\n\n"
NAMESPACE_MEMBER_TYPES_WITH_CHAT_TOOLS: Final = frozenset({"function", "custom"})
_INCOMPLETE_REASON_BY_FINISH_REASON: Final[Mapping[str, Literal["max_output_tokens", "content_filter"]]] = (
MappingProxyType({"length": "max_output_tokens", "content_filter": "content_filter", "refusal": "content_filter"})
)
@dataclass(frozen=True, slots=True)
@ -2299,6 +2303,18 @@ class LiteLLMCompletionResponsesConfig:
# Default to completed for unknown finish reasons
return "completed"
@staticmethod
def _incomplete_details_for_finish_reason(
finish_reason: str | None,
existing: IncompleteDetails | None,
) -> IncompleteDetails | None:
if existing is not None:
return existing
if finish_reason is None:
return None
reason: Final = _INCOMPLETE_REASON_BY_FINISH_REASON.get(finish_reason)
return IncompleteDetails(reason=reason) if reason is not None else None
@staticmethod
def _tool_call_id_from_responses_item(item_id: str | None, call_id: str | None) -> str:
"""Bedrock Mantle returns a non-unique, index-based ``call_id`` (``call_0``,
@ -2415,13 +2431,18 @@ class LiteLLMCompletionResponsesConfig:
if choices and len(choices) > 0:
finish_reason = choices[0].finish_reason
incomplete_details: Final = LiteLLMCompletionResponsesConfig._incomplete_details_for_finish_reason(
finish_reason=finish_reason,
existing=getattr(chat_completion_response, "incomplete_details", None),
)
responses_api_response: Final[ResponsesAPIResponse] = ResponsesAPIResponse(
id=chat_completion_response.id,
created_at=chat_completion_response.created,
model=chat_completion_response.model,
object="response",
error=getattr(chat_completion_response, "error", None),
incomplete_details=getattr(chat_completion_response, "incomplete_details", None),
incomplete_details=incomplete_details,
instructions=getattr(chat_completion_response, "instructions", None),
metadata=getattr(chat_completion_response, "metadata", {}),
output=LiteLLMCompletionResponsesConfig._transform_chat_completion_choices_to_responses_output(

View file

@ -212,18 +212,21 @@ def _error_event_fields(error_obj: object) -> tuple[str, str | None, str | None]
raw_code = None
message: Final = str(raw_message) if raw_message is not None else "Response API in-stream error"
error_type: Final = raw_type if isinstance(raw_type, str) else None
code: Final = raw_code if isinstance(raw_code, str) else None
code: Final = str(raw_code) if isinstance(raw_code, (str, int)) and not isinstance(raw_code, bool) else None
return message, error_type, code
def _status_code_for_error_field(field: str) -> int | None:
if field.isdecimal() and 400 <= int(field) <= 599:
return int(field)
return _ERROR_CODE_HTTP_STATUS.get(field)
def _status_code_for_error_fields(error_type: str | None, error_code: str | None) -> int:
fields: Final = tuple(field for field in (error_code, error_type) if field is not None)
if any(field.startswith("rate_limit") or field == "insufficient_quota" for field in fields):
return 429
return next(
(_ERROR_CODE_HTTP_STATUS[field] for field in fields if field in _ERROR_CODE_HTTP_STATUS),
500,
)
return next((status for status in map(_status_code_for_error_field, fields) if status is not None), 500)
def _mid_stream_fallback_eligible(mapped_exception: Exception) -> bool:

View file

@ -425,22 +425,35 @@ class UrlContextMetadata(TypedDict, total=False):
urlMetadata: list[UrlMetadata]
GeminiFinishReason = Literal[
"FINISH_REASON_UNSPECIFIED",
"STOP",
"MAX_TOKENS",
"SAFETY",
"RECITATION",
"LANGUAGE",
"OTHER",
"BLOCKLIST",
"PROHIBITED_CONTENT",
"SPII",
"MALFORMED_FUNCTION_CALL",
"IMAGE_SAFETY",
"IMAGE_PROHIBITED_CONTENT",
"TOO_MANY_TOOL_CALLS",
"MALFORMED_RESPONSE",
"NO_IMAGE",
"IMAGE_RECITATION",
"IMAGE_OTHER",
"ESCALATION",
"UNEXPECTED_TOOL_CALL",
"MISSING_THOUGHT_SIGNATURE",
]
class Candidates(TypedDict, total=False):
index: int
content: HttpxContentType
finishReason: Literal[
"FINISH_REASON_UNSPECIFIED",
"STOP",
"MAX_TOKENS",
"SAFETY",
"RECITATION",
"OTHER",
"BLOCKLIST",
"PROHIBITED_CONTENT",
"SPII",
"MALFORMED_FUNCTION_CALL",
"IMAGE_SAFETY",
]
finishReason: GeminiFinishReason
safetyRatings: list[SafetyRatings]
citationMetadata: CitationMetadata
groundingMetadata: GroundingMetadata

View file

@ -91,6 +91,22 @@ class MCPPublicServer(BaseModel):
mcp_info: dict[str, Any] | None = None
class MCPAllowedClient(BaseModel):
"""One entry of `general_settings.mcp_allowed_clients`."""
model_config = ConfigDict(frozen=True)
alias: str = Field(
min_length=1,
description="Human-readable name for this client application, shown in the dashboard and in gateway logs.",
)
value: str = Field(
min_length=1,
description="Exact value of the JWT claim named in litellm_jwtauth.mcp_client_id_jwt_field, or of the "
"mcp_client_id_header header, that identifies this client application. Matched case-sensitively.",
)
class MCPToolSearchSettings(BaseModel):
"""`litellm_settings.mcp_tool_search`: how the native `mcp_tool_search` virtual tool ranks the caller's tools."""

View file

@ -19244,7 +19244,20 @@
"supports_function_calling": true,
"supports_anthropic_thinking_payload": true,
"supports_prompt_caching": true,
"supports_tool_choice": true
"supports_tool_choice": true,
"supports_reasoning": true,
"supported_modalities": [
"text",
"image",
"audio",
"video"
],
"supported_output_modalities": [
"text"
],
"supports_vision": true,
"supports_audio_input": true,
"supports_video_input": true
},
"databricks/databricks-gemini-2-5-pro": {
"cache_creation_input_token_cost": 1.24999e-06,
@ -19265,7 +19278,21 @@
"supports_function_calling": true,
"supports_anthropic_thinking_payload": true,
"supports_prompt_caching": true,
"supports_tool_choice": true
"supports_tool_choice": true,
"deprecation_date": "2026-10-02",
"supports_reasoning": true,
"supported_modalities": [
"text",
"image",
"audio",
"video"
],
"supported_output_modalities": [
"text"
],
"supports_vision": true,
"supports_audio_input": true,
"supports_video_input": true
},
"databricks/databricks-gemini-3-1-flash-lite": {
"cache_creation_input_token_cost": 3.1248e-07,
@ -19285,7 +19312,19 @@
"source": "https://www.databricks.com/product/pricing/proprietary-foundation-model-serving",
"supports_function_calling": true,
"supports_prompt_caching": true,
"supports_tool_choice": true
"supports_tool_choice": true,
"supported_modalities": [
"text",
"image",
"audio",
"video"
],
"supported_output_modalities": [
"text"
],
"supports_vision": true,
"supports_audio_input": true,
"supports_video_input": true
},
"databricks/databricks-gemini-3-1-flash-image": {
"litellm_provider": "databricks",
@ -19347,7 +19386,20 @@
"source": "https://www.databricks.com/product/pricing/proprietary-foundation-model-serving",
"supports_function_calling": true,
"supports_prompt_caching": true,
"supports_tool_choice": true
"supports_tool_choice": true,
"supports_reasoning": true,
"supported_modalities": [
"text",
"image",
"audio",
"video"
],
"supported_output_modalities": [
"text"
],
"supports_vision": true,
"supports_audio_input": true,
"supports_video_input": true
},
"databricks/databricks-gemini-3-flash": {
"cache_creation_input_token_cost": 6.2503e-07,
@ -19367,7 +19419,19 @@
"source": "https://www.databricks.com/product/pricing/proprietary-foundation-model-serving",
"supports_function_calling": true,
"supports_prompt_caching": true,
"supports_tool_choice": true
"supports_tool_choice": true,
"supported_modalities": [
"text",
"image",
"audio",
"video"
],
"supported_output_modalities": [
"text"
],
"supports_vision": true,
"supports_audio_input": true,
"supports_video_input": true
},
"databricks/databricks-gemini-3-pro": {
"cache_creation_input_token_cost": 2.49998e-06,
@ -21433,7 +21497,11 @@
"mode": "chat",
"supports_tool_choice": true,
"supports_function_calling": true,
"supports_image_size": false
"supports_image_size": false,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_vision": true,
"supports_audio_input": true
},
"deepinfra/google/gemini-2.5-pro": {
"max_tokens": 1000000,
@ -21444,7 +21512,11 @@
"litellm_provider": "deepinfra",
"mode": "chat",
"supports_tool_choice": true,
"supports_function_calling": true
"supports_function_calling": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_vision": true,
"supports_audio_input": true
},
"deepinfra/google/gemma-3-12b-it": {
"max_tokens": 131072,
@ -29340,26 +29412,6 @@
"supports_parallel_function_calling": true,
"supports_vision": true
},
"github_copilot/gemini-2.5-pro": {
"litellm_provider": "github_copilot",
"max_input_tokens": 128000,
"max_output_tokens": 64000,
"max_tokens": 64000,
"mode": "chat",
"supports_function_calling": true,
"supports_parallel_function_calling": true,
"supports_vision": true
},
"github_copilot/gemini-3-pro-preview": {
"litellm_provider": "github_copilot",
"max_input_tokens": 128000,
"max_output_tokens": 64000,
"max_tokens": 64000,
"mode": "chat",
"supports_function_calling": true,
"supports_parallel_function_calling": true,
"supports_vision": true
},
"github_copilot/gpt-3.5-turbo": {
"litellm_provider": "github_copilot",
"max_input_tokens": 16384,
@ -30014,17 +30066,6 @@
"output_cost_per_token": 8.8e-07,
"supports_function_calling": true
},
"gmi/google/gemini-3-pro-preview": {
"input_cost_per_token": 2e-06,
"litellm_provider": "gmi",
"max_input_tokens": 1048576,
"max_output_tokens": 65536,
"max_tokens": 65536,
"mode": "chat",
"output_cost_per_token": 1.2e-05,
"supports_function_calling": true,
"supports_vision": true
},
"gmi/google/gemini-3-flash-preview": {
"input_cost_per_token": 5e-07,
"litellm_provider": "gmi",
@ -30034,7 +30075,8 @@
"mode": "chat",
"output_cost_per_token": 3e-06,
"supports_function_calling": true,
"supports_vision": true
"supports_vision": true,
"supports_system_messages": true
},
"gmi/moonshotai/Kimi-K2-Thinking": {
"input_cost_per_token": 8e-07,
@ -40163,7 +40205,12 @@
"supports_response_schema": true,
"supports_vision": true,
"supports_native_streaming": true,
"supports_image_size": false
"supports_image_size": false,
"supports_reasoning": true,
"supports_system_messages": true,
"supports_pdf_input": true,
"supports_audio_input": true,
"supports_video_input": true
},
"oci/google.gemini-2.5-pro": {
"input_cost_per_token": 1.25e-06,
@ -40177,7 +40224,12 @@
"supports_function_calling": true,
"supports_response_schema": true,
"supports_vision": true,
"supports_native_streaming": true
"supports_native_streaming": true,
"supports_reasoning": true,
"supports_system_messages": true,
"supports_pdf_input": true,
"supports_audio_input": true,
"supports_video_input": true
},
"oci/google.gemini-2.5-flash-lite": {
"input_cost_per_token": 7.5e-08,
@ -40192,7 +40244,12 @@
"supports_response_schema": true,
"supports_vision": true,
"supports_native_streaming": true,
"supports_image_size": false
"supports_image_size": false,
"supports_reasoning": false,
"supports_system_messages": true,
"supports_pdf_input": true,
"supports_audio_input": true,
"supports_video_input": true
},
"oci/cohere.command-a-vision": {
"input_cost_per_token": 1.56e-06,
@ -41435,7 +41492,7 @@
"max_tokens": 65535,
"mode": "chat",
"output_cost_per_token": 2.5e-06,
"supports_audio_output": true,
"supports_audio_output": false,
"supports_function_calling": true,
"supports_response_schema": true,
"supports_system_messages": true,
@ -41449,7 +41506,8 @@
"supports_audio_input": true,
"supports_pdf_input": true,
"supports_reasoning": true,
"supports_web_search": false
"supports_web_search": false,
"supports_video_input": true
},
"openrouter/google/gemini-2.5-pro": {
"cache_creation_input_token_cost": 3.75e-07,
@ -41462,7 +41520,7 @@
"max_tokens": 65536,
"mode": "chat",
"output_cost_per_token": 1e-05,
"supports_audio_output": true,
"supports_audio_output": false,
"supports_function_calling": true,
"supports_response_schema": true,
"supports_system_messages": true,
@ -41478,7 +41536,8 @@
"supports_audio_input": true,
"supports_pdf_input": true,
"supports_reasoning": true,
"supports_web_search": false
"supports_web_search": false,
"supports_video_input": true
},
"openrouter/google/gemini-3-pro-preview": {
"cache_read_input_token_cost": 2e-07,
@ -41563,7 +41622,8 @@
"supports_url_context": true,
"supports_vision": true,
"supports_web_search": false,
"tpm": 800000
"tpm": 800000,
"supports_video_input": true
},
"openrouter/google/gemini-3.1-flash-lite-preview": {
"cache_creation_input_token_cost": 8.33333333333333e-08,
@ -41690,7 +41750,8 @@
"supports_system_messages": true,
"supports_tool_choice": true,
"supports_vision": true,
"supports_web_search": false
"supports_web_search": false,
"supports_video_input": true
},
"openrouter/gryphe/mythomax-l2-13b": {
"input_cost_per_token": 8e-08,
@ -44413,12 +44474,16 @@
"output_cost_per_token": 1.2e-05,
"litellm_provider": "replicate",
"mode": "chat",
"supports_function_calling": true,
"supports_parallel_function_calling": true,
"supports_function_calling": false,
"supports_parallel_function_calling": false,
"supports_vision": true,
"supports_system_messages": true,
"supports_tool_choice": true,
"supports_response_schema": true
"supports_tool_choice": false,
"supports_response_schema": false,
"input_cost_per_token_above_200k_tokens": 4e-06,
"output_cost_per_token_above_200k_tokens": 1.8e-05,
"supports_audio_input": true,
"supports_video_input": true
},
"replicate/anthropic/claude-4.5-sonnet": {
"input_cost_per_token": 3e-06,
@ -44487,17 +44552,19 @@
"supports_response_schema": true
},
"replicate/google/gemini-2.5-flash": {
"input_cost_per_token": 2.5e-06,
"input_cost_per_token": 3e-07,
"output_cost_per_token": 2.5e-06,
"litellm_provider": "replicate",
"mode": "chat",
"supports_function_calling": true,
"supports_parallel_function_calling": true,
"supports_function_calling": false,
"supports_parallel_function_calling": false,
"supports_vision": true,
"supports_system_messages": true,
"supports_tool_choice": true,
"supports_response_schema": true,
"supports_image_size": false
"supports_tool_choice": false,
"supports_response_schema": false,
"supports_image_size": false,
"supports_reasoning": true,
"supports_video_input": true
},
"replicate/openai/gpt-oss-120b": {
"input_cost_per_token": 1.8e-07,
@ -48076,10 +48143,15 @@
"supports_function_calling": true,
"supports_tool_choice": true,
"supports_response_schema": true,
"supports_image_size": false
"supports_image_size": false,
"cache_read_input_token_cost": 3e-08,
"supports_reasoning": true,
"supports_pdf_input": true,
"supports_web_search": true,
"supports_prompt_caching": true
},
"vercel_ai_gateway/google/gemini-2.5-pro": {
"input_cost_per_token": 2.5e-06,
"input_cost_per_token": 1.25e-06,
"litellm_provider": "vercel_ai_gateway",
"max_input_tokens": 1048576,
"max_output_tokens": 65536,
@ -48089,7 +48161,15 @@
"supports_vision": true,
"supports_function_calling": true,
"supports_tool_choice": true,
"supports_response_schema": true
"supports_response_schema": true,
"input_cost_per_token_above_200k_tokens": 2.5e-06,
"output_cost_per_token_above_200k_tokens": 1.5e-05,
"cache_read_input_token_cost": 1.25e-07,
"cache_read_input_token_cost_above_200k_tokens": 2.5e-07,
"supports_reasoning": true,
"supports_pdf_input": true,
"supports_web_search": true,
"supports_prompt_caching": true
},
"vercel_ai_gateway/google/gemini-embedding-001": {
"input_cost_per_token": 1.5e-07,
@ -62224,7 +62304,8 @@
"supports_response_schema": true,
"supports_reasoning": true,
"supports_vision": true,
"source": "https://deepinfra.com/pricing"
"source": "https://deepinfra.com/pricing",
"supports_audio_input": true
},
"deepinfra/XiaomiMiMo/MiMo-V2.5": {
"max_tokens": 262144,
@ -62524,7 +62605,8 @@
"supports_response_schema": true,
"supports_reasoning": true,
"supports_vision": true,
"source": "https://deepinfra.com/pricing"
"source": "https://deepinfra.com/pricing",
"supports_audio_input": true
},
"deepinfra/google/gemini-3.7-flash": {
"max_tokens": 1000000,
@ -62538,7 +62620,8 @@
"supports_response_schema": true,
"supports_reasoning": true,
"supports_vision": true,
"source": "https://deepinfra.com/pricing"
"source": "https://deepinfra.com/pricing",
"supports_audio_input": true
},
"deepinfra/inclusionAI/Ling-3.0-flash": {
"max_tokens": 131072,
@ -62970,7 +63053,8 @@
"supports_response_schema": true,
"supports_reasoning": true,
"supports_vision": true,
"source": "https://deepinfra.com/pricing"
"source": "https://deepinfra.com/pricing",
"supports_audio_input": true
},
"deepinfra/XiaomiMiMo/MiMo-V2.5-Pro": {
"max_tokens": 1048576,
@ -65365,7 +65449,8 @@
"deprecation_date": "2026-10-20",
"input_cost_per_audio_token": 3e-07,
"supports_prompt_caching": true,
"supports_web_search": false
"supports_web_search": false,
"supports_video_input": true
},
"openrouter/google/gemini-3.5-flash": {
"cache_creation_input_token_cost": 8.33333333333333e-08,
@ -65388,7 +65473,8 @@
"cache_read_input_token_cost": 1.5e-07,
"input_cost_per_audio_token": 3e-06,
"supports_prompt_caching": true,
"supports_web_search": false
"supports_web_search": false,
"supports_video_input": true
},
"openrouter/google/gemini-3.5-flash-lite": {
"cache_creation_input_token_cost": 8.33333333333333e-08,
@ -65411,7 +65497,8 @@
"cache_read_input_token_cost": 3e-08,
"input_cost_per_audio_token": 3e-07,
"supports_prompt_caching": true,
"supports_web_search": false
"supports_web_search": false,
"supports_video_input": true
},
"openrouter/google/gemini-3.6-flash": {
"cache_creation_input_token_cost": 4.16666666666667e-08,
@ -65434,7 +65521,8 @@
"cache_read_input_token_cost": 7.5e-08,
"input_cost_per_audio_token": 7.5e-07,
"supports_prompt_caching": true,
"supports_web_search": false
"supports_web_search": false,
"supports_video_input": true
},
"openrouter/google/gemini-3.7-flash": {
"cache_creation_input_token_cost": 4.16666666666667e-08,
@ -65457,7 +65545,8 @@
"cache_read_input_token_cost": 7.5e-08,
"input_cost_per_audio_token": 7.5e-07,
"supports_prompt_caching": true,
"supports_web_search": false
"supports_web_search": false,
"supports_video_input": true
},
"openrouter/google/gemini-3.8-flash": {
"cache_creation_input_token_cost": 4.16666666666667e-08,
@ -65480,7 +65569,8 @@
"cache_read_input_token_cost": 7.5e-08,
"input_cost_per_audio_token": 7.5e-07,
"supports_prompt_caching": true,
"supports_web_search": false
"supports_web_search": false,
"supports_video_input": true
},
"openrouter/openai/gpt-4o-mini": {
"input_cost_per_token": 1.5e-07,
@ -67108,7 +67198,8 @@
"supports_pdf_input": true,
"supports_audio_input": true,
"supports_prompt_caching": true,
"supports_web_search": false
"supports_web_search": false,
"supports_video_input": true
},
"openrouter/qwen/qwen3-max-thinking": {
"input_cost_per_token": 7.8e-07,
@ -71974,7 +72065,8 @@
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
"supports_web_search": false
"supports_web_search": false,
"supports_video_input": true
},
"openrouter/google/gemini-2.5-flash:batch": {
"cache_read_input_audio_token_cost": 1e-07,
@ -71997,7 +72089,8 @@
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
"supports_web_search": false
"supports_web_search": false,
"supports_video_input": true
},
"openrouter/google/gemini-2.5-pro:batch": {
"cache_read_input_audio_token_cost": 1.25e-07,
@ -72023,7 +72116,8 @@
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
"supports_web_search": false
"supports_web_search": false,
"supports_video_input": true
},
"openrouter/google/gemini-3-flash-preview:batch": {
"input_cost_per_audio_token": 5e-07,
@ -72043,7 +72137,8 @@
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
"supports_web_search": false
"supports_web_search": false,
"supports_video_input": true
},
"openrouter/google/gemini-3.1-flash-lite:batch": {
"cache_read_input_audio_token_cost": 2.5e-08,
@ -72065,7 +72160,8 @@
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
"supports_web_search": false
"supports_web_search": false,
"supports_video_input": true
},
"openrouter/google/gemini-3.1-pro-preview:batch": {
"input_cost_per_audio_token": 1e-06,
@ -72087,7 +72183,8 @@
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
"supports_web_search": false
"supports_web_search": false,
"supports_video_input": true
},
"openrouter/google/gemini-3.5-flash-lite:batch": {
"cache_read_input_audio_token_cost": 1.5e-08,
@ -72109,7 +72206,8 @@
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
"supports_web_search": false
"supports_web_search": false,
"supports_video_input": true
},
"openrouter/google/gemini-3.5-flash:batch": {
"cache_read_input_audio_token_cost": 1.5e-07,
@ -72131,7 +72229,8 @@
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
"supports_web_search": false
"supports_web_search": false,
"supports_video_input": true
},
"openrouter/google/gemini-3.6-flash:batch": {
"cache_creation_input_token_cost": 4.16666666666667e-08,
@ -72154,7 +72253,8 @@
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
"supports_web_search": false
"supports_web_search": false,
"supports_video_input": true
},
"openrouter/google/gemini-3.7-flash:batch": {
"cache_creation_input_token_cost": 4.16666666666667e-08,
@ -72177,7 +72277,8 @@
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
"supports_web_search": false
"supports_web_search": false,
"supports_video_input": true
},
"openrouter/google/gemini-3.8-flash:batch": {
"cache_creation_input_token_cost": 4.16666666666667e-08,
@ -72200,7 +72301,8 @@
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
"supports_web_search": false
"supports_web_search": false,
"supports_video_input": true
},
"openrouter/ibm-granite/granite-4.0-h-micro": {
"input_cost_per_token": 1.7e-08,

View file

@ -213,7 +213,6 @@ files_settings:
api_key: os.environ/OPENAI_API_KEY
router_settings:
routing_strategy: usage-based-routing-v2
redis_host: os.environ/REDIS_HOST
redis_password: os.environ/REDIS_PASSWORD
redis_port: os.environ/REDIS_PORT

View file

@ -161,7 +161,7 @@ class Scenario:
assert all(object_value(object_value(entry)["model_info"])["id"] != identity for entry in entries)
assert read_rows('SELECT model_id FROM "LiteLLM_ProxyModelTable" WHERE model_id = %s', (identity,)) == []
def model(self, **parameters: JsonValue) -> str:
def model(self, *, model_info: Mapping[str, JsonValue] | None = None, **parameters: JsonValue) -> str:
name: Final = f"integration-{uuid.uuid4().hex}"
created: Final = self.gateway.post(
"/model/new",
@ -173,7 +173,7 @@ class Scenario:
"api_base": f"{self.gateway.upstream_url}/v1",
**parameters,
},
"model_info": {},
"model_info": dict(model_info) if model_info is not None else {},
},
)
identity: Final = string_value(object_value(created["model_info"])["id"])

View file

@ -92,6 +92,12 @@
"tests/integration/pricing/test_price_precedence.py::test_same_upstream_aliases_keep_distinct_prices_after_reload": [
"quota_management.spend_tracking.alias_prices.remain_independent_on_reload"
],
"tests/integration/pricing/test_off_peak_pricing.py::test_open_off_peak_window_bills_off_peak_rates": [
"quota_management.spend_tracking.off_peak_pricing.open_window_bills_off_peak_rates"
],
"tests/integration/pricing/test_off_peak_pricing.py::test_closed_off_peak_window_bills_standard_rates": [
"quota_management.spend_tracking.off_peak_pricing.closed_window_bills_standard_rates"
],
"tests/integration/spend/test_cache_and_quota.py::test_generated_cache_sequences_preserve_content_usage_and_zero_hit_cost": [
"quota_management.response_cache.generated_sequences_preserve_content_and_accounting"
],

View file

@ -0,0 +1,82 @@
import json
from collections.abc import Mapping
from datetime import datetime, timedelta, timezone
from typing import Final
import pytest
from pydantic import JsonValue
from tests.integration._support.client import Gateway, Scenario, eventually, object_value, string_value
from tests.integration._support.database import read_rows
STANDARD_INPUT_RATE: Final = 0.001
STANDARD_OUTPUT_RATE: Final = 0.002
OFF_PEAK_INPUT_RATE: Final = 0.0001
OFF_PEAK_OUTPUT_RATE: Final = 0.0002
def off_peak_window(start_offset_hours: int, end_offset_hours: int) -> Mapping[str, JsonValue]:
now: Final = datetime.now(timezone.utc)
start: Final = now + timedelta(hours=start_offset_hours)
end: Final = now + timedelta(hours=end_offset_hours)
return {
"hours_utc": f"{start:%H:%M}-{end:%H:%M}",
"input_cost_per_token": OFF_PEAK_INPUT_RATE,
"output_cost_per_token": OFF_PEAK_OUTPUT_RATE,
}
def billed_model(scenario: Scenario, off_peak: Mapping[str, JsonValue]) -> str:
return scenario.model(
input_cost_per_token=STANDARD_INPUT_RATE,
output_cost_per_token=STANDARD_OUTPUT_RATE,
model_info={"off_peak_pricing": dict(off_peak)},
)
def assert_chat_bills_rates(gateway: Gateway, model: str, input_rate: float, output_rate: float) -> None:
response: Final = gateway.request(
"POST", "/v1/chat/completions", {"model": model, "messages": [{"role": "user", "content": "off peak control"}]}
)
assert response.status_code == 200, response.text
expected: Final = 20 * input_rate + 20 * output_rate
assert float(response.headers["x-litellm-response-cost"]) == pytest.approx(expected, rel=1e-6)
request_id: Final = string_value(object_value(response.json())["id"])
rows: Final = eventually(
lambda: read_rows(
'SELECT spend, metadata, prompt_tokens, completion_tokens FROM "LiteLLM_SpendLogs" WHERE request_id = %s',
(request_id,),
),
lambda values: len(values) == 1,
seconds=70,
)
assert rows[0]["prompt_tokens"] == 20
assert rows[0]["completion_tokens"] == 20
assert float(rows[0]["spend"]) == pytest.approx(expected, rel=1e-6)
metadata: Final = rows[0]["metadata"]
parsed: Final = json.loads(metadata) if isinstance(metadata, str) else object_value(metadata)
breakdown: Final = object_value(parsed["cost_breakdown"])
assert float(breakdown["input_cost"]) == pytest.approx(20 * input_rate, rel=1e-6)
assert float(breakdown["output_cost"]) == pytest.approx(20 * output_rate, rel=1e-6)
@pytest.mark.covers("quota_management.spend_tracking.off_peak_pricing.open_window_bills_off_peak_rates")
def test_open_off_peak_window_bills_off_peak_rates(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
model: Final = billed_model(scenario, off_peak_window(-1, 1))
entries: Final = gateway.get("/model/info")["data"]
assert isinstance(entries, list)
matching: Final = tuple(object_value(entry) for entry in entries if object_value(entry)["model_name"] == model)
assert len(matching) == 1
info: Final = object_value(matching[0]["model_info"])
off_peak: Final = object_value(info["off_peak_pricing"])
assert off_peak["input_cost_per_token"] == OFF_PEAK_INPUT_RATE
assert off_peak["output_cost_per_token"] == OFF_PEAK_OUTPUT_RATE
assert_chat_bills_rates(gateway, model, OFF_PEAK_INPUT_RATE, OFF_PEAK_OUTPUT_RATE)
@pytest.mark.covers("quota_management.spend_tracking.off_peak_pricing.closed_window_bills_standard_rates")
def test_closed_off_peak_window_bills_standard_rates(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
model: Final = billed_model(scenario, off_peak_window(2, 3))
assert_chat_bills_rates(gateway, model, STANDARD_INPUT_RATE, STANDARD_OUTPUT_RATE)

View file

@ -13,5 +13,4 @@ litellm_settings:
host: os.environ/REDIS_HOST
port: os.environ/REDIS_PORT
router_settings:
num_retries: 0
disable_cooldowns: true

View file

@ -3079,6 +3079,9 @@ async def test_update_config_success_callback_normalization():
async def add_deployment(self, prisma_client=None, proxy_logging_obj=None): # noqa: F811 # pytest fixture, not a redefinition
return None
def reject_config_owned_writes(self, *, section_name, changed_keys):
return None
setattr(proxy_server, "proxy_config", MockProxyConfig())
config_update = ConfigYAML(litellm_settings={"success_callback": ["SQS", "sQs"]})

View file

@ -1482,6 +1482,44 @@ class TestBackgroundStreamingTerminalEvents:
assert final_call.kwargs["status"] == "failed"
assert final_call.kwargs["error"] == error_payload
@pytest.mark.asyncio
async def test_named_event_failed_frame_sets_failed_status_and_error(self):
from litellm.proxy.response_polling.background_streaming import (
background_streaming_task,
)
error_payload = {
"code": "cyber_policy",
"message": "Your request was flagged for possible cybersecurity risk and was not completed",
}
failed_event = {
"type": "response.failed",
"sequence_number": 5,
"response": {"id": "resp_123", "status": "failed", "error": error_payload, "output": []},
}
async def _body_iterator():
yield b'data: {"type": "response.in_progress"}\n\n'
yield f"event: response.failed\ndata: {json.dumps(failed_event)}\n\n".encode()
yield b"data: [DONE]\n\n"
mock_response = Mock()
mock_response.body_iterator = _body_iterator()
handler = AsyncMock(spec=ResponsePollingHandler)
kwargs = _make_background_streaming_kwargs("poll_named_event", handler)
with patch( # test-quality-ok: the processor is built inside the task, same idiom as the sibling tests
"litellm.proxy.response_polling.background_streaming.ProxyBaseLLMRequestProcessing"
) as MockProcessor:
MockProcessor.return_value.base_process_llm_request = AsyncMock(
return_value=mock_response
)
await background_streaming_task(**kwargs)
final_call = handler.update_state.call_args_list[-1]
assert final_call.kwargs["status"] == "failed"
assert final_call.kwargs["error"] == error_payload
@pytest.mark.asyncio
async def test_response_incomplete_sets_incomplete_status_and_details(self):
"""Test that a response.incomplete stream event results in incomplete status"""

View file

@ -130,7 +130,6 @@ class TestPollingEndpointPreCallGuard:
"litellm.proxy.proxy_server.proxy_config": MagicMock(),
"litellm.proxy.proxy_server.proxy_logging_obj": AsyncMock(),
"litellm.proxy.proxy_server.redis_usage_cache": AsyncMock(),
"litellm.proxy.proxy_server.select_data_generator": None,
"litellm.proxy.proxy_server.user_api_base": None,
"litellm.proxy.proxy_server.user_max_tokens": None,
"litellm.proxy.proxy_server.user_model": None,

View file

@ -151,6 +151,12 @@ class TestMapFinishReasonGemini:
("IMAGE_PROHIBITED_CONTENT", "content_filter"),
("TOO_MANY_TOOL_CALLS", "stop"),
("MALFORMED_RESPONSE", "stop"),
("NO_IMAGE", "content_filter"),
("IMAGE_RECITATION", "content_filter"),
("IMAGE_OTHER", "content_filter"),
("ESCALATION", "content_filter"),
("UNEXPECTED_TOOL_CALL", "stop"),
("MISSING_THOUGHT_SIGNATURE", "stop"),
],
)
def test_gemini_finish_reasons(self, gemini_reason, expected):

View file

@ -1437,6 +1437,29 @@ def test_openai_compatible_vendor_400_keeps_body_but_not_headers():
assert not exc_info.value.response.headers
@pytest.mark.parametrize(
("status_code", "mapped_class"), [(429, litellm.RateLimitError), (500, litellm.InternalServerError)]
)
def test_openai_429_and_500_keep_body(status_code: int, mapped_class: type[openai.APIError]):
with pytest.raises(mapped_class) as exc_info:
exception_type(
model="gpt-5.4-mini",
original_exception=_openai_handler_error(
"server_error", {}, status_code=status_code, message="upstream cannot complete this response"
),
custom_llm_provider="openai",
completion_kwargs={},
extra_kwargs={},
)
assert exc_info.value.body == {
**_GUARDRAIL_BLOCK_ERROR,
"type": "server_error",
"code": str(status_code),
"message": "upstream cannot complete this response",
}
def test_litellm_proxy_repeated_response_header_keeps_each_value():
repeated = [("x-litellm-call-id", "call-guardrail"), ("set-cookie", "a=1"), ("set-cookie", "b=2")]

View file

@ -105,6 +105,46 @@ def test_translate_chat_length_takes_precedence_over_refusal():
assert result.get("stop_details") is None
def test_translate_chat_content_filter_to_anthropic_response():
response = ModelResponse(
id="chatcmpl-content-filter",
model="openai-model",
choices=[
Choices(
index=0,
finish_reason="content_filter",
message=Message(content=None, role="assistant"),
)
],
usage=Usage(prompt_tokens=1, completion_tokens=0, total_tokens=1),
)
result = LiteLLMAnthropicMessagesAdapter().translate_openai_response_to_anthropic(response)
assert result["content"] == []
assert result["stop_reason"] == "refusal"
def test_translate_chat_refusal_finish_reason_to_anthropic_response():
response = ModelResponse(
id="chatcmpl-refusal-reason",
model="openai-model",
choices=[
Choices(
index=0,
finish_reason="refusal",
message=Message(content=None, role="assistant"),
)
],
usage=Usage(prompt_tokens=1, completion_tokens=0, total_tokens=1),
)
result = LiteLLMAnthropicMessagesAdapter().translate_openai_response_to_anthropic(response)
assert result["content"] == []
assert result["stop_reason"] == "refusal"
def test_translate_streaming_openai_chunk_to_anthropic_content_block():
choices = [
StreamingChoices(

View file

@ -458,6 +458,34 @@ def test_reasoning_with_forced_tool_choice_switches_to_auto():
assert optional_params["tool_choice"] == {"auto": {}}
@pytest.mark.parametrize(
"model, param, value, expected_max_tokens",
[
("us.openai.gpt-6-astra", "max_tokens", 1, 16),
("us.openai.gpt-6-astra", "max_completion_tokens", 1, 16),
("us.openai.gpt-6-astra", "max_tokens", 64, 64),
("us.xai.grok-4.6", "max_tokens", 1, 16),
("global.xai.grok-4.6", "max_completion_tokens", 1, 16),
("us.xai.grok-4.6", "max_tokens", 32, 32),
("anthropic.claude-sonnet-4-5-20250929-v1:0", "max_tokens", 1, 1),
("arn:aws:bedrock:us-east-1:123456789012:inference-profile/us.openai.gpt-6-astra", "max_tokens", 1, 16),
("arn:aws:bedrock:us-east-1:123456789012:inference-profile/global.xai.grok-4.6", "max_tokens", 1, 16),
("arn:aws:bedrock:us-east-1:123456789012:application-inference-profile/abc123xyz", "max_tokens", 1, 1),
],
)
def test_map_openai_params_enforces_minimum_max_tokens_for_openai_compat_models(
model: str, param: str, value: int, expected_max_tokens: int
):
optional_params = AmazonConverseConfig().map_openai_params(
non_default_params={param: value},
optional_params={},
model=model,
drop_params=False,
)
assert optional_params["maxTokens"] == expected_max_tokens
@pytest.mark.parametrize(
"model",
[

View file

@ -2,7 +2,7 @@ import asyncio
import json
import re
from copy import deepcopy
from typing import Final, List, cast
from typing import Final, List, cast, get_args
from unittest.mock import MagicMock, patch
import httpx
@ -18,7 +18,7 @@ from litellm.llms.vertex_ai.common_utils import VertexAIError
from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import (
VertexGeminiConfig,
)
from litellm.types.llms.vertex_ai import UsageMetadata
from litellm.types.llms.vertex_ai import GeminiFinishReason, UsageMetadata
from litellm.types.utils import ChoiceLogprobs, Usage
from litellm.utils import CustomStreamWrapper
@ -940,6 +940,11 @@ def test_check_finish_reason():
)
def test_every_documented_gemini_finish_reason_has_an_explicit_mapping():
documented: Final = frozenset(get_args(GeminiFinishReason))
assert set(VertexGeminiConfig.get_finish_reason_mapping()) == documented
def test_finish_reason_unspecified_and_malformed_function_call():
"""
Test that FINISH_REASON_UNSPECIFIED and MALFORMED_FUNCTION_CALL
@ -968,6 +973,12 @@ def test_finish_reason_unspecified_and_malformed_function_call():
# Test new Gemini finish reasons
assert finish_reason_mappings["TOO_MANY_TOOL_CALLS"] == "stop"
assert finish_reason_mappings["MALFORMED_RESPONSE"] == "stop"
assert finish_reason_mappings["NO_IMAGE"] == "content_filter"
assert finish_reason_mappings["IMAGE_RECITATION"] == "content_filter"
assert finish_reason_mappings["IMAGE_OTHER"] == "content_filter"
assert finish_reason_mappings["ESCALATION"] == "content_filter"
assert finish_reason_mappings["UNEXPECTED_TOOL_CALL"] == "stop"
assert finish_reason_mappings["MISSING_THOUGHT_SIGNATURE"] == "stop"
def test_vertex_ai_usage_metadata_response_token_count():
@ -6074,3 +6085,210 @@ def test_prompt_blocked_chunk_keeps_served_model_version():
assert streaming_chunk.model == "gemini-3.8-flash-001"
assert streaming_chunk.choices[0].finish_reason == "content_filter"
def test_gemini_candidate_with_finish_reason_no_content_chat_completion():
config = VertexGeminiConfig()
completion_response = {
"candidates": [
{
"finishReason": "NO_IMAGE",
"index": 0,
}
],
"usageMetadata": {
"promptTokenCount": 19,
"candidatesTokenCount": 0,
"totalTokenCount": 19,
},
}
model_response = ModelResponse()
logging_obj = MagicMock()
raw_response = MagicMock()
raw_response.headers = {}
resp = config._transform_google_generate_content_to_openai_model_response(
completion_response=completion_response,
model_response=model_response,
model="gemini-2.5-flash-image",
logging_obj=logging_obj,
raw_response=raw_response,
)
assert len(resp.choices) == 1
assert resp.choices[0].finish_reason == "content_filter"
assert resp.choices[0].message.content is None
assert resp.choices[0].provider_specific_fields["native_finish_reason"] == "NO_IMAGE"
def test_gemini_candidate_with_finish_reason_no_content_anthropic_messages():
from litellm.llms.anthropic.experimental_pass_through.adapters.transformation import (
LiteLLMAnthropicMessagesAdapter,
)
config = VertexGeminiConfig()
completion_response = {
"candidates": [
{
"finishReason": "NO_IMAGE",
"index": 0,
}
],
"usageMetadata": {
"promptTokenCount": 19,
"candidatesTokenCount": 0,
"totalTokenCount": 19,
},
}
resp = config._transform_google_generate_content_to_openai_model_response(
completion_response=completion_response,
model_response=ModelResponse(),
model="gemini-2.5-flash-image",
logging_obj=MagicMock(),
raw_response=MagicMock(headers={}),
)
adapter = LiteLLMAnthropicMessagesAdapter()
anthropic_resp = adapter.translate_openai_response_to_anthropic(
response=resp,
tool_name_mapping={},
)
assert anthropic_resp["stop_reason"] == "refusal"
assert anthropic_resp["content"] == []
def test_gemini_candidate_with_finish_reason_no_content_responses_api():
from litellm.responses.litellm_completion_transformation.transformation import (
LiteLLMCompletionResponsesConfig,
)
config = VertexGeminiConfig()
completion_response = {
"candidates": [
{
"finishReason": "NO_IMAGE",
"index": 0,
}
],
"usageMetadata": {
"promptTokenCount": 19,
"candidatesTokenCount": 0,
"totalTokenCount": 19,
},
}
resp = config._transform_google_generate_content_to_openai_model_response(
completion_response=completion_response,
model_response=ModelResponse(),
model="gemini-2.5-flash-image",
logging_obj=MagicMock(),
raw_response=MagicMock(headers={}),
)
responses_resp = LiteLLMCompletionResponsesConfig.transform_chat_completion_response_to_responses_api_response(
request_input="Generate picture",
responses_api_request={},
chat_completion_response=resp,
)
assert responses_resp.status == "incomplete"
assert responses_resp.incomplete_details is not None
assert responses_resp.incomplete_details.reason == "content_filter"
def test_gemini_candidate_other_finish_reasons_no_content():
from litellm.llms.anthropic.experimental_pass_through.adapters.transformation import (
LiteLLMAnthropicMessagesAdapter,
)
from litellm.responses.litellm_completion_transformation.transformation import (
LiteLLMCompletionResponsesConfig,
)
config = VertexGeminiConfig()
max_tokens_response = {
"candidates": [{"finishReason": "MAX_TOKENS", "index": 0}],
"usageMetadata": {"promptTokenCount": 10, "candidatesTokenCount": 50, "totalTokenCount": 60},
}
resp_length = config._transform_google_generate_content_to_openai_model_response(
completion_response=max_tokens_response,
model_response=ModelResponse(),
model="gemini-2.5-flash",
logging_obj=MagicMock(),
raw_response=MagicMock(headers={}),
)
assert len(resp_length.choices) == 1
assert resp_length.choices[0].finish_reason == "length"
assert resp_length.choices[0].provider_specific_fields["native_finish_reason"] == "MAX_TOKENS"
anthropic_length = LiteLLMAnthropicMessagesAdapter().translate_openai_response_to_anthropic(
response=resp_length,
tool_name_mapping={},
)
assert anthropic_length["stop_reason"] == "max_tokens"
responses_length = LiteLLMCompletionResponsesConfig.transform_chat_completion_response_to_responses_api_response(
request_input="thinking request",
responses_api_request={},
chat_completion_response=resp_length,
)
assert responses_length.status == "incomplete"
assert responses_length.incomplete_details.reason == "max_output_tokens"
def test_gemini_candidate_with_finish_reason_no_content_streaming_chunk():
from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import (
ModelResponseIterator,
)
chunk: Final = {
"candidates": [{"finishReason": "NO_IMAGE", "index": 0}],
"usageMetadata": {"promptTokenCount": 19, "candidatesTokenCount": 0, "totalTokenCount": 19},
}
iterator: Final = ModelResponseIterator(streaming_response=[], sync_stream=True, logging_obj=MagicMock())
streaming_chunk: Final = iterator.chunk_parser(chunk)
assert len(streaming_chunk.choices) == 1
assert streaming_chunk.choices[0].finish_reason == "content_filter"
assert streaming_chunk.choices[0].delta.content is None
assert streaming_chunk.choices[0].delta.tool_calls is None
def test_gemini_multi_candidate_messages_do_not_share_state():
config: Final = VertexGeminiConfig()
completion_response: Final = {
"candidates": [
{
"content": {
"role": "model",
"parts": [
{"text": "Let me check the weather.", "thought": True},
{"functionCall": {"name": "get_weather", "args": {"city": "Paris"}}},
],
},
"finishReason": "STOP",
"index": 0,
},
{
"content": {"role": "model", "parts": [{"text": "It is sunny in Paris."}]},
"finishReason": "STOP",
"index": 1,
},
],
"usageMetadata": {"promptTokenCount": 10, "candidatesTokenCount": 20, "totalTokenCount": 30},
}
resp: Final = config._transform_google_generate_content_to_openai_model_response(
completion_response=completion_response,
model_response=ModelResponse(),
model="gemini-2.5-flash",
logging_obj=MagicMock(),
raw_response=MagicMock(headers={}),
)
assert len(resp.choices) == 2
assert resp.choices[0].finish_reason == "tool_calls"
assert resp.choices[0].message.tool_calls[0].function.name == "get_weather"
assert resp.choices[0].message.reasoning_content == "Let me check the weather."
assert resp.choices[1].finish_reason == "stop"
assert resp.choices[1].message.content == "It is sunny in Paris."
assert resp.choices[1].message.tool_calls is None
assert getattr(resp.choices[1].message, "reasoning_content", None) is None
assert resp.choices[1].provider_specific_fields["native_finish_reason"] == "STOP"

View file

@ -0,0 +1,217 @@
from collections.abc import Mapping
from typing import Final
import pytest
from litellm.proxy._experimental.mcp_server.client_allowlist import (
MCP_ALLOWED_CLIENTS_SETTING,
MCP_CLIENT_ID_HEADER_SETTING,
MCP_CLIENT_ID_JWT_FIELD_SETTING,
MCPClientAllowlist,
MCPClientIdentity,
MCPClientRejection,
check_mcp_client_allowed,
load_mcp_client_allowlist,
parse_allowed_mcp_clients,
resolve_mcp_client_identity,
)
ANTIGRAVITY: Final = {"alias": "Antigravity CLI", "value": "antigravity-cli"}
CODEX: Final = {"alias": "Codex", "value": "codex-mcp-client"}
ANTIGRAVITY_ONLY: Final[Mapping[str, str]] = {"antigravity-cli": "Antigravity CLI"}
JWT_ONLY: Final = MCPClientAllowlist(aliases_by_value=ANTIGRAVITY_ONLY, jwt_field="azp", header=None)
HEADER_ONLY: Final = MCPClientAllowlist(aliases_by_value=ANTIGRAVITY_ONLY, jwt_field=None, header="x-mcp-client")
JWT_AND_HEADER: Final = MCPClientAllowlist(aliases_by_value=ANTIGRAVITY_ONLY, jwt_field="azp", header="x-mcp-client")
NO_SOURCE: Final = MCPClientAllowlist(aliases_by_value=ANTIGRAVITY_ONLY, jwt_field=None, header=None)
NO_HEADERS: Final[Mapping[str, str]] = {}
_ALLOWLIST_SETTING_CASES: Final[tuple[tuple[object, Mapping[str, str] | None], ...]] = (
(None, None),
([], {}),
([ANTIGRAVITY], ANTIGRAVITY_ONLY),
([ANTIGRAVITY, CODEX], {"antigravity-cli": "Antigravity CLI", "codex-mcp-client": "Codex"}),
(
[ANTIGRAVITY, {"alias": "Antigravity (prod)", "value": "antigravity-cli"}],
{"antigravity-cli": "Antigravity (prod)"},
),
(
[ANTIGRAVITY, {"alias": "Antigravity CLI", "value": "antigravity-prod"}],
{**ANTIGRAVITY_ONLY, "antigravity-prod": "Antigravity CLI"},
),
(["antigravity-cli"], {}),
("antigravity-cli", {}),
([ANTIGRAVITY, 1], {}),
([{"alias": "Antigravity CLI"}], {}),
([{"value": "antigravity-cli"}], {}),
([{"alias": "", "value": "antigravity-cli"}], {}),
([{"alias": "Antigravity CLI", "value": ""}], {}),
([{"alias": "Antigravity CLI", "value": ["antigravity-cli"]}], {}),
(ANTIGRAVITY, {}),
)
@pytest.mark.parametrize(("raw_setting", "expected"), _ALLOWLIST_SETTING_CASES)
def test_parse_allowed_mcp_clients(raw_setting: object, expected: Mapping[str, str] | None) -> None:
assert parse_allowed_mcp_clients(raw_setting) == expected
def test_load_returns_none_when_the_allowlist_setting_is_absent_even_if_identity_sources_are_set() -> None:
settings: Final = {"litellm_jwtauth": {"mcp_client_id_jwt_field": "azp"}, "mcp_client_id_header": "x-mcp-client"}
assert load_mcp_client_allowlist(settings) is None
def test_load_reads_the_jwt_field_from_litellm_jwtauth_and_lowercases_the_header_name() -> None:
settings: Final = {
"mcp_allowed_clients": [ANTIGRAVITY, CODEX],
"litellm_jwtauth": {"user_id_jwt_field": "sub", "mcp_client_id_jwt_field": "resource_access.mcp.client"},
"mcp_client_id_header": "X-MCP-Client",
}
assert load_mcp_client_allowlist(settings) == MCPClientAllowlist(
aliases_by_value={"antigravity-cli": "Antigravity CLI", "codex-mcp-client": "Codex"},
jwt_field="resource_access.mcp.client",
header="x-mcp-client",
)
@pytest.mark.parametrize(
"settings",
(
{"mcp_allowed_clients": [ANTIGRAVITY]},
{"mcp_allowed_clients": [ANTIGRAVITY], "litellm_jwtauth": {}, "mcp_client_id_header": ""},
{"mcp_allowed_clients": [ANTIGRAVITY], "litellm_jwtauth": {"mcp_client_id_jwt_field": ""}},
{"mcp_allowed_clients": [ANTIGRAVITY], "litellm_jwtauth": "azp", "mcp_client_id_header": ["x"]},
),
)
def test_load_without_a_usable_identity_source_keeps_the_allowlist_but_no_source(
settings: Mapping[str, object],
) -> None:
assert load_mcp_client_allowlist(settings) == NO_SOURCE
@pytest.mark.parametrize("raw_setting", ("antigravity-cli", ["antigravity-cli"], [{"alias": "Antigravity CLI"}]))
def test_load_malformed_allowlist_admits_nobody(raw_setting: object) -> None:
loaded: Final = load_mcp_client_allowlist({"mcp_allowed_clients": raw_setting})
assert loaded is not None
assert loaded.aliases_by_value == {}
assert check_mcp_client_allowed(loaded, {"azp": "antigravity-cli"}, {"x-mcp-client": "antigravity-cli"}) is not None
def test_only_the_value_identifies_a_client_never_its_alias() -> None:
assert check_mcp_client_allowed(JWT_ONLY, {"azp": "Antigravity CLI"}, NO_HEADERS) is not None
assert check_mcp_client_allowed(HEADER_ONLY, None, {"x-mcp-client": "Antigravity CLI"}) is not None
def test_two_clients_may_share_an_alias_and_both_are_admitted() -> None:
settings: Final = {
"mcp_allowed_clients": [
{"alias": "Coding CLI", "value": "cli-dev"},
{"alias": "Coding CLI", "value": "cli-prod"},
],
"litellm_jwtauth": {"mcp_client_id_jwt_field": "azp"},
}
loaded: Final = load_mcp_client_allowlist(settings)
assert check_mcp_client_allowed(loaded, {"azp": "cli-dev"}, NO_HEADERS) is None
assert check_mcp_client_allowed(loaded, {"azp": "cli-prod"}, NO_HEADERS) is None
assert check_mcp_client_allowed(loaded, {"azp": "Coding CLI"}, NO_HEADERS) is not None
def test_unconfigured_allowlist_admits_callers_with_no_identity_at_all() -> None:
assert check_mcp_client_allowed(None, None, NO_HEADERS) is None
assert check_mcp_client_allowed(None, {"azp": "claude-code"}, {"x-mcp-client": "claude-code"}) is None
def test_jwt_claim_identifies_the_client() -> None:
assert resolve_mcp_client_identity(JWT_ONLY, {"azp": "antigravity-cli"}, NO_HEADERS) == MCPClientIdentity(
client_id="antigravity-cli", source="jwt", source_name="azp"
)
assert check_mcp_client_allowed(JWT_ONLY, {"azp": "antigravity-cli"}, NO_HEADERS) is None
def test_nested_jwt_claim_path_is_resolved_with_dot_notation() -> None:
nested: Final = MCPClientAllowlist(
aliases_by_value=ANTIGRAVITY_ONLY, jwt_field="resource_access.mcp.client", header=None
)
claims: Final = {"resource_access": {"mcp": {"client": "antigravity-cli"}}}
assert check_mcp_client_allowed(nested, claims, NO_HEADERS) is None
def test_unlisted_jwt_client_is_rejected_and_the_rejection_names_it() -> None:
rejection: Final = check_mcp_client_allowed(JWT_ONLY, {"azp": "claude-code"}, NO_HEADERS)
assert isinstance(rejection, MCPClientRejection)
assert "'claude-code'" in rejection.details
assert "azp" in rejection.details
assert MCP_ALLOWED_CLIENTS_SETTING in rejection.details
assert rejection.response_body == {"error": "Forbidden", "details": rejection.details}
@pytest.mark.parametrize("claims", ({"sub": "user-1"}, {"azp": ""}, {"azp": 42}, {"azp": ["antigravity-cli"]}))
def test_jwt_without_a_usable_client_claim_is_rejected(claims: Mapping[str, object]) -> None:
rejection: Final = check_mcp_client_allowed(JWT_ONLY, claims, NO_HEADERS)
assert isinstance(rejection, MCPClientRejection)
assert "azp" in rejection.details
def test_matching_is_exact_not_prefix_or_case_insensitive() -> None:
for spoof in ("Antigravity-Cli", "antigravity-cli-sdk", " antigravity-cli"):
assert check_mcp_client_allowed(JWT_ONLY, {"azp": spoof}, NO_HEADERS) is not None
assert check_mcp_client_allowed(HEADER_ONLY, None, {"x-mcp-client": spoof}) is not None
def test_configured_header_identifies_callers_without_a_jwt() -> None:
headers: Final = {"x-mcp-client": "antigravity-cli"}
assert resolve_mcp_client_identity(HEADER_ONLY, None, headers) == MCPClientIdentity(
client_id="antigravity-cli", source="header", source_name="x-mcp-client"
)
assert check_mcp_client_allowed(HEADER_ONLY, None, headers) is None
assert check_mcp_client_allowed(HEADER_ONLY, {}, headers) is None
def test_unlisted_or_missing_header_is_rejected() -> None:
unlisted: Final = check_mcp_client_allowed(HEADER_ONLY, None, {"x-mcp-client": "claude-code"})
assert isinstance(unlisted, MCPClientRejection)
assert "'claude-code'" in unlisted.details
for headers in (NO_HEADERS, {"x-mcp-client": ""}, {"x-other": "antigravity-cli"}):
missing = check_mcp_client_allowed(HEADER_ONLY, None, headers)
assert isinstance(missing, MCPClientRejection)
assert "x-mcp-client" in missing.details
def test_header_is_not_consulted_when_it_is_not_configured() -> None:
rejection: Final = check_mcp_client_allowed(JWT_ONLY, None, {"x-mcp-client": "antigravity-cli"})
assert isinstance(rejection, MCPClientRejection)
assert MCP_CLIENT_ID_HEADER_SETTING in rejection.details
assert MCP_CLIENT_ID_JWT_FIELD_SETTING in rejection.details
def test_jwt_caller_is_judged_by_its_claim_even_when_the_header_would_pass() -> None:
spoofed_header: Final = {"x-mcp-client": "antigravity-cli"}
assert check_mcp_client_allowed(JWT_AND_HEADER, {"azp": "claude-code"}, spoofed_header) is not None
assert check_mcp_client_allowed(JWT_AND_HEADER, {"sub": "user-1"}, spoofed_header) is not None
assert check_mcp_client_allowed(JWT_AND_HEADER, {"azp": "antigravity-cli"}, {"x-mcp-client": "claude-code"}) is None
def test_jwt_caller_with_an_empty_claim_set_cannot_fall_back_to_the_header() -> None:
rejection: Final = check_mcp_client_allowed(JWT_AND_HEADER, {}, {"x-mcp-client": "antigravity-cli"})
assert isinstance(rejection, MCPClientRejection)
assert "azp" in rejection.details
def test_non_jwt_caller_falls_back_to_the_header_when_both_sources_are_configured() -> None:
assert check_mcp_client_allowed(JWT_AND_HEADER, None, {"x-mcp-client": "antigravity-cli"}) is None
assert check_mcp_client_allowed(JWT_AND_HEADER, None, {"x-mcp-client": "claude-code"}) is not None
def test_allowlist_with_no_identity_source_rejects_everyone_and_says_what_to_configure() -> None:
rejection: Final = check_mcp_client_allowed(
NO_SOURCE, {"azp": "antigravity-cli"}, {"x-mcp-client": "antigravity-cli"}
)
assert isinstance(rejection, MCPClientRejection)
assert MCP_CLIENT_ID_JWT_FIELD_SETTING in rejection.details
assert MCP_CLIENT_ID_HEADER_SETTING in rejection.details
def test_empty_allowlist_rejects_an_identified_client() -> None:
empty: Final = MCPClientAllowlist(aliases_by_value={}, jwt_field="azp", header="x-mcp-client")
assert check_mcp_client_allowed(empty, {"azp": "antigravity-cli"}, NO_HEADERS) is not None
assert check_mcp_client_allowed(empty, None, {"x-mcp-client": "antigravity-cli"}) is not None

View file

@ -1,4 +1,5 @@
import asyncio
import contextlib
import contextvars
import os
from datetime import datetime, timedelta
@ -17,6 +18,8 @@ from mcp.types import (
TextContent,
TextResourceContents,
)
from pydantic import TypeAdapter
from starlette.types import Receive, Scope, Send
from litellm.proxy._types import (
LiteLLM_MCPServerTable,
@ -2035,6 +2038,220 @@ async def test_mcp_routing_initialize_to_stateful_no_session_to_stateless(
assert not any(name.startswith(b"x-mcp-debug") for name in headers)
_FORBIDDEN_BODY_ADAPTER: Final[TypeAdapter[dict[str, str]]] = TypeAdapter(dict[str, str])
_BODY_CHUNK_ADAPTER: Final[TypeAdapter[bytes]] = TypeAdapter(bytes)
_INITIALIZE: Final = (
b'{"jsonrpc":"2.0","id":0,"method":"initialize","params":{"protocolVersion":"2025-06-18",'
b'"capabilities":{},"clientInfo":{"name":"antigravity-cli","version":"1.0.0"}}}'
)
_TOOLS_LIST: Final = b'{"jsonrpc":"2.0","id":2,"method":"tools/list","params":{}}'
_ALLOWLIST_SETTINGS: Final[dict[str, object]] = {
"mcp_allowed_clients": [{"alias": "Antigravity CLI", "value": "antigravity-cli"}],
"litellm_jwtauth": {"mcp_client_id_jwt_field": "azp"},
"mcp_client_id_header": "x-mcp-client",
}
_LISTED_JWT: Final[dict[str, object]] = {"azp": "antigravity-cli", "sub": "user-1"}
_UNLISTED_JWT: Final[dict[str, object]] = {"azp": "claude-code", "sub": "user-1"}
_LISTED_HEADER: Final[list[tuple[bytes, bytes]]] = [(b"x-mcp-client", b"antigravity-cli")]
_UNLISTED_HEADER: Final[list[tuple[bytes, bytes]]] = [(b"x-mcp-client", b"claude-code")]
async def _drain_body(receive: Receive) -> bytes:
first: Final = await receive()
body: Final = _BODY_CHUNK_ADAPTER.validate_python(first.get("body", b""))
if not first.get("more_body", False):
return body
return body + await _drain_body(receive)
def _client_allowlist_patches(
settings: dict[str, object], jwt_claims: dict[str, object] | None
) -> contextlib.ExitStack:
stack: Final = contextlib.ExitStack()
stack.enter_context(
patch( # test-quality-ok: the ASGI handler resolves auth through a module-level function; no injection seam
"litellm.proxy._experimental.mcp_server.server.extract_mcp_auth_context",
new_callable=AsyncMock,
return_value=(UserAPIKeyAuth(user_id="allowlist-user", jwt_claims=jwt_claims), None, None, None, None, {}),
)
)
stack.enter_context(
patch( # test-quality-ok: module flag guarding lazy session-manager startup; no injection seam
"litellm.proxy._experimental.mcp_server.server._SESSION_MANAGERS_INITIALIZED", True
)
)
stack.enter_context(
patch( # test-quality-ok: the allowlist is read off this module global; no injection seam
"litellm.proxy.proxy_server.general_settings", settings
)
)
return stack
def _forbidden_body(denied: HTTPException) -> dict[str, str]:
return _FORBIDDEN_BODY_ADAPTER.validate_python(denied.detail)
@pytest.mark.asyncio
@pytest.mark.parametrize(
("jwt_claims", "headers", "request_body", "expected_fragment"),
(
(_UNLISTED_JWT, [], _INITIALIZE, "MCP client 'claude-code' (from JWT claim 'azp')"),
(_UNLISTED_JWT, _LISTED_HEADER, _INITIALIZE, "MCP client 'claude-code' (from JWT claim 'azp')"),
({"sub": "user-1"}, _LISTED_HEADER, _INITIALIZE, "no 'azp' claim"),
(None, _UNLISTED_HEADER, _INITIALIZE, "MCP client 'claude-code' (from header 'x-mcp-client')"),
(None, [], _INITIALIZE, "no 'x-mcp-client' header"),
(_UNLISTED_JWT, [(b"mcp-session-id", b"session-1")], _TOOLS_LIST, "MCP client 'claude-code'"),
),
)
async def test_streamable_http_rejects_unlisted_client_before_any_session_work(
jwt_claims: dict[str, object] | None,
headers: list[tuple[bytes, bytes]],
request_body: bytes,
expected_fragment: str,
) -> None:
from litellm.proxy._experimental.mcp_server import server as mcp_module
scope: Final[Scope] = {"type": "http", "method": "POST", "path": "/mcp", "headers": headers}
receive: Final = AsyncMock(return_value={"type": "http.request", "body": request_body, "more_body": False})
send: Final = AsyncMock()
stateful_handle: Final = AsyncMock()
stateless_handle: Final = AsyncMock()
session_cap: Final = AsyncMock(return_value=True)
with (
_client_allowlist_patches(_ALLOWLIST_SETTINGS, jwt_claims),
patch( # test-quality-ok: session managers are module singletons; the downstream call is the observable
"litellm.proxy._experimental.mcp_server.server.session_manager_stateful",
SimpleNamespace(handle_request=stateful_handle),
),
patch( # test-quality-ok: session managers are module singletons; the downstream call is the observable
"litellm.proxy._experimental.mcp_server.server.session_manager_stateless",
SimpleNamespace(handle_request=stateless_handle),
),
patch( # test-quality-ok: module-level cap check; asserting it is never reached is the point
"litellm.proxy._experimental.mcp_server.server._enforce_stateful_session_cap_for_owner", session_cap
),
pytest.raises(HTTPException) as denied,
):
await mcp_module.handle_streamable_http_mcp(scope, receive, send)
assert denied.value.status_code == 403
body: Final = _forbidden_body(denied.value)
assert body["error"] == "Forbidden"
assert expected_fragment in body["details"]
assert "mcp_allowed_clients" in body["details"]
receive.assert_not_awaited()
send.assert_not_awaited()
stateful_handle.assert_not_awaited()
stateless_handle.assert_not_awaited()
session_cap.assert_not_awaited()
@pytest.mark.asyncio
@pytest.mark.parametrize(
("settings", "jwt_claims", "headers"),
(
(_ALLOWLIST_SETTINGS, _LISTED_JWT, []),
(_ALLOWLIST_SETTINGS, _LISTED_JWT, _UNLISTED_HEADER),
(_ALLOWLIST_SETTINGS, None, _LISTED_HEADER),
({}, _UNLISTED_JWT, _UNLISTED_HEADER),
({}, None, []),
),
)
async def test_streamable_http_admits_listed_or_unrestricted_clients_and_hands_the_body_downstream(
settings: dict[str, object], jwt_claims: dict[str, object] | None, headers: list[tuple[bytes, bytes]]
) -> None:
from litellm.proxy._experimental.mcp_server import server as mcp_module
scope: Final[Scope] = {"type": "http", "method": "POST", "path": "/mcp", "headers": headers}
receive: Final = AsyncMock(
side_effect=[
{"type": "http.request", "body": _INITIALIZE[:20], "more_body": True},
{"type": "http.request", "body": _INITIALIZE[20:], "more_body": False},
]
)
send: Final = AsyncMock()
downstream_bodies: Final[list[bytes]] = []
async def handle_request(_: Scope, downstream_receive: Receive, __: Send) -> None:
downstream_bodies.append(await _drain_body(downstream_receive))
stateful_handle: Final = AsyncMock(side_effect=handle_request)
stateless_handle: Final = AsyncMock()
with (
_client_allowlist_patches(settings, jwt_claims),
patch( # test-quality-ok: session managers are module singletons; the downstream call is the observable
"litellm.proxy._experimental.mcp_server.server.session_manager_stateful",
SimpleNamespace(handle_request=stateful_handle),
),
patch( # test-quality-ok: session managers are module singletons; the downstream call is the observable
"litellm.proxy._experimental.mcp_server.server.session_manager_stateless",
SimpleNamespace(handle_request=stateless_handle),
),
):
await mcp_module.handle_streamable_http_mcp(scope, receive, send)
assert downstream_bodies == [_INITIALIZE]
stateless_handle.assert_not_awaited()
send.assert_not_awaited()
@pytest.mark.asyncio
@pytest.mark.parametrize(
("jwt_claims", "headers", "admitted"),
(
(_LISTED_JWT, [], True),
(None, _LISTED_HEADER, True),
(_UNLISTED_JWT, _LISTED_HEADER, False),
(None, _UNLISTED_HEADER, False),
(None, [], False),
),
)
async def test_sse_endpoint_applies_the_same_client_allowlist(
jwt_claims: dict[str, object] | None, headers: list[tuple[bytes, bytes]], admitted: bool
) -> None:
from litellm.proxy._experimental.mcp_server import server as mcp_module
scope: Final[Scope] = {"type": "http", "method": "POST", "path": "/mcp/sse", "headers": headers}
receive: Final = AsyncMock(return_value={"type": "http.request", "body": _INITIALIZE, "more_body": False})
send: Final = AsyncMock()
downstream_bodies: Final[list[bytes]] = []
async def handle_request(_: Scope, downstream_receive: Receive, __: Send) -> None:
downstream_bodies.append(await _drain_body(downstream_receive))
with (
_client_allowlist_patches(_ALLOWLIST_SETTINGS, jwt_claims),
patch( # test-quality-ok: module-level pre-auth probe unrelated to the allowlist under test; no injection seam
"litellm.proxy._experimental.mcp_server.server._raise_preemptive_401_for_unauthenticated_servers",
new_callable=AsyncMock,
),
patch( # test-quality-ok: module-level upstream auth probe unrelated to the allowlist under test; no injection seam
"litellm.proxy._experimental.mcp_server.server._check_passthrough_upstream_auth",
new_callable=AsyncMock,
),
patch.object( # test-quality-ok: SSE manager is a module singleton; the downstream call is the observable
mcp_module.sse_session_manager, "handle_request", side_effect=handle_request
),
):
if admitted:
await mcp_module.handle_sse_mcp(scope, receive, send)
assert downstream_bodies == [_INITIALIZE]
send.assert_not_awaited()
return
with pytest.raises(HTTPException) as denied:
await mcp_module.handle_sse_mcp(scope, receive, send)
assert denied.value.status_code == 403
body: Final = _forbidden_body(denied.value)
assert body["error"] == "Forbidden"
assert "mcp_allowed_clients" in body["details"]
assert downstream_bodies == []
send.assert_not_awaited()
@pytest.mark.asyncio
async def test_mcp_routing_chunked_initialize_to_stateful():
"""
@ -5422,11 +5639,12 @@ async def test_get_tools_from_mcp_servers_logs_list_tools_to_spendlogs_when_enab
Ensure list-tools logging path calls `async_success_handler` when enabled.
"""
try:
from mcp.types import Tool as MCPTool
from litellm.proxy._experimental.mcp_server.server import (
_get_tools_from_mcp_servers,
)
from litellm.proxy._types import UserAPIKeyAuth
from mcp.types import Tool as MCPTool
except ImportError:
pytest.skip("MCP server not available")
@ -8376,10 +8594,10 @@ async def test_fire_mcp_tool_call_logging_iserror_logs_failure():
"""Regression test: a CallToolResult with isError=True must go
down the failure logging path (async_failure_handler + post_call_failure_hook),
never async_success_handler."""
from litellm.proxy._experimental.mcp_server.exceptions import MCPToolResultError
from litellm.proxy._experimental.mcp_server.server import (
_fire_mcp_tool_call_logging,
)
from litellm.proxy._experimental.mcp_server.exceptions import MCPToolResultError
logging_obj = _mock_mcp_logging_obj()
proxy_logging_mock = _mock_mcp_proxy_logging()
@ -8729,11 +8947,11 @@ async def test_call_mcp_tool_skips_failure_hook_for_upstream_auth_error():
caller-must-reauth signal, not a failed call, so call_mcp_tool must re-raise it WITHOUT firing
post_call_failure_hook (which records a failure and can trip LLM exception alerts). The
streamable handler downgrades it to an informational isError result afterward."""
from litellm.proxy._experimental.mcp_server.exceptions import MCPUpstreamAuthError
from litellm.proxy._experimental.mcp_server.server import (
call_mcp_tool,
global_mcp_server_manager,
)
from litellm.proxy._experimental.mcp_server.exceptions import MCPUpstreamAuthError
from litellm.proxy._types import MCPTransport, UserAPIKeyAuth
from litellm.types.mcp_server.mcp_server_manager import MCPServer

View file

@ -4525,3 +4525,131 @@ class TestV1ResolvedOauth2Gate:
assert rest_endpoints._v1_resolved_oauth2_server_ids(["oauth2-srv"]) == set()
assert rest_endpoints._v1_resolved_oauth2_server_ids(["oauth2-srv", "delegate-srv"]) == {"delegate-srv"}
_CLIENT_ALLOWLIST_SETTINGS: Final[dict[str, object]] = {
"mcp_allowed_clients": [{"alias": "Antigravity CLI", "value": "antigravity-cli"}],
"litellm_jwtauth": {"mcp_client_id_jwt_field": "azp"},
"mcp_client_id_header": "x-mcp-client",
}
class TestClientAllowlistOnRestRoutes:
"""``mcp_allowed_clients`` must gate the REST tool facade exactly like the /mcp transports,
otherwise an unlisted harness can list and call tools by switching to /mcp-rest."""
pytestmark = pytest.mark.asyncio
@staticmethod
def _stub_listing(monkeypatch: pytest.MonkeyPatch) -> list[UserAPIKeyAuth]:
listed_for: list[UserAPIKeyAuth] = []
async def fake_contexts(user_api_key_auth: UserAPIKeyAuth) -> list[UserAPIKeyAuth]:
listed_for.append(user_api_key_auth)
return [user_api_key_auth]
async def fake_get_allowed_mcp_servers(
user_api_key_auth: UserAPIKeyAuth | None = None,
*,
keyless_source: bool = False,
) -> list[str]:
return []
monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", _CLIENT_ALLOWLIST_SETTINGS, raising=False)
monkeypatch.setattr(rest_endpoints, "build_effective_auth_contexts", fake_contexts, raising=False)
monkeypatch.setattr(
rest_endpoints.global_mcp_server_manager,
"get_allowed_mcp_servers",
fake_get_allowed_mcp_servers,
raising=False,
)
return listed_for
@pytest.mark.parametrize(
("caller", "headers", "expected_fragment"),
(
(UserAPIKeyAuth(jwt_claims={"azp": "claude-code"}), {"x-mcp-client": "antigravity-cli"}, "'claude-code'"),
(UserAPIKeyAuth(jwt_claims={}), {"x-mcp-client": "antigravity-cli"}, "no 'azp' claim"),
(UserAPIKeyAuth(), {"x-mcp-client": "claude-code"}, "'claude-code'"),
(UserAPIKeyAuth(), {}, "no 'x-mcp-client' header"),
),
)
async def test_tools_list_rejects_unlisted_clients_before_resolving_servers(
self,
monkeypatch: pytest.MonkeyPatch,
caller: UserAPIKeyAuth,
headers: dict[str, str],
expected_fragment: str,
) -> None:
listed_for: Final = self._stub_listing(monkeypatch)
request: Final = _build_request(headers, path="/mcp-rest/tools/list", method="GET")
with pytest.raises(HTTPException) as denied:
await rest_endpoints.list_tool_rest_api(
request, server_id=None, mcp_server_name=None, toolset_name=None, user_api_key_dict=caller
)
assert denied.value.status_code == 403
assert denied.value.detail["error"] == "Forbidden"
assert expected_fragment in denied.value.detail["details"]
assert "mcp_allowed_clients" in denied.value.detail["details"]
assert listed_for == []
@pytest.mark.parametrize(
("caller", "headers"),
(
(UserAPIKeyAuth(jwt_claims={"azp": "antigravity-cli"}), {"x-mcp-client": "claude-code"}),
(UserAPIKeyAuth(), {"x-mcp-client": "antigravity-cli"}),
),
)
async def test_tools_list_admits_listed_clients(
self, monkeypatch: pytest.MonkeyPatch, caller: UserAPIKeyAuth, headers: dict[str, str]
) -> None:
listed_for: Final = self._stub_listing(monkeypatch)
request: Final = _build_request(headers, path="/mcp-rest/tools/list", method="GET")
result: Final = await rest_endpoints.list_tool_rest_api(
request, server_id=None, mcp_server_name=None, toolset_name=None, user_api_key_dict=caller
)
assert result["tools"] == []
assert listed_for == [caller]
async def test_dashboard_session_is_not_treated_as_a_client_application(
self, monkeypatch: pytest.MonkeyPatch
) -> None:
from litellm.constants import UI_SESSION_TOKEN_TEAM_ID
listed_for: Final = self._stub_listing(monkeypatch)
session: Final = UserAPIKeyAuth(team_id=UI_SESSION_TOKEN_TEAM_ID, user_id="admin-user", user_role="proxy_admin")
request: Final = _build_request(path="/mcp-rest/tools/list", method="GET")
result: Final = await rest_endpoints.list_tool_rest_api(
request, server_id=None, mcp_server_name=None, toolset_name=None, user_api_key_dict=session
)
assert result["tools"] == []
assert listed_for == [session]
async def test_tools_call_rejects_unlisted_clients_before_reading_the_body(
self, monkeypatch: pytest.MonkeyPatch
) -> None:
monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", _CLIENT_ALLOWLIST_SETTINGS, raising=False)
acting: Final = AsyncMock()
monkeypatch.setattr(rest_endpoints, "acting_user_auth", acting, raising=False)
request: Final = _build_request(
{"x-mcp-client": "antigravity-cli"},
path="/mcp-rest/tools/call",
method="POST",
json_body={"server_id": "server-1", "name": "demo-tool", "arguments": {}},
)
with pytest.raises(HTTPException) as denied:
await rest_endpoints.call_tool_rest_api(
request, user_api_key_dict=UserAPIKeyAuth(jwt_claims={"azp": "claude-code"})
)
assert denied.value.status_code == 403
assert denied.value.detail["error"] == "Forbidden"
assert "'claude-code'" in denied.value.detail["details"]
acting.assert_not_awaited()

View file

@ -7545,6 +7545,70 @@ async def test_project_allowlist_enforced_when_key_models_empty():
assert exc_info.value.code == "403"
def _project_with_budget(spend: float, max_budget: float):
from litellm.proxy._types import LiteLLM_BudgetTable, LiteLLM_ProjectTableCachedObj
return LiteLLM_ProjectTableCachedObj(
project_id="p-budget",
team_id="t-1",
budget_id="b-1",
spend=spend,
litellm_budget_table=LiteLLM_BudgetTable(budget_id="b-1", max_budget=max_budget),
)
@pytest.mark.asyncio
@pytest.mark.parametrize(
"counter_spend, db_spend, max_budget, blocks",
[
pytest.param(5.0, 0.0, 5.0, True, id="counter-at-budget-blocks-despite-stale-db-row"),
pytest.param(4.99, 0.0, 5.0, False, id="counter-under-budget-admits"),
pytest.param(None, 5.0, 5.0, True, id="no-counter-falls-back-to-persisted-spend"),
pytest.param(None, 0.0, 5.0, False, id="no-counter-and-no-persisted-spend-admits"),
pytest.param(12.5, 12.5, 0.0, False, id="zero-budget-is-unbudgeted"),
pytest.param(12.5, 12.5, -1.0, False, id="negative-budget-is-unbudgeted"),
],
)
async def test_project_max_budget_check_blocks_only_when_live_spend_reaches_a_positive_budget(
counter_spend, db_spend, max_budget, blocks
):
from litellm.caching.dual_cache import DualCache
from litellm.proxy.auth.auth_checks import _project_max_budget_check
real_spend_counter_cache = DualCache()
if counter_spend is not None:
real_spend_counter_cache.in_memory_cache.set_cache(key="spend:project:p-budget", value=counter_spend)
valid_token = UserAPIKeyAuth(api_key="hashed-key", project_id="p-budget", team_id="t-1", user_id="u-1")
proxy_logging_obj = MagicMock()
proxy_logging_obj.budget_alerts = AsyncMock()
with patch( # test-quality-ok: injects a real DualCache for the module global, not a behavior mock
"litellm.proxy.proxy_server.spend_counter_cache", real_spend_counter_cache
):
if not blocks:
await _project_max_budget_check(
project_object=_project_with_budget(spend=db_spend, max_budget=max_budget),
valid_token=valid_token,
proxy_logging_obj=proxy_logging_obj,
)
await asyncio.sleep(0)
proxy_logging_obj.budget_alerts.assert_not_awaited()
return
with pytest.raises(litellm.BudgetExceededError) as exc_info:
await _project_max_budget_check(
project_object=_project_with_budget(spend=db_spend, max_budget=max_budget),
valid_token=valid_token,
proxy_logging_obj=proxy_logging_obj,
)
await asyncio.sleep(0)
assert exc_info.value.entity_type == Litellm_EntityType.PROJECT.value
assert exc_info.value.entity_id == "p-budget"
assert exc_info.value.current_cost == 5.0
proxy_logging_obj.budget_alerts.assert_awaited_once()
assert proxy_logging_obj.budget_alerts.await_args.kwargs["type"] == "project_budget"
def test_is_user_proxy_admin_rejects_view_only_admin():
"""This predicate skips `non_proxy_admin_allowed_routes_check` entirely, so an
Admin Viewer answering True here would gain every write route. Read parity for

View file

@ -102,6 +102,7 @@ class MockBatcher:
self.litellm_organizationtable = _Table("org", self)
self.litellm_tagtable = _Table("tag", self)
self.litellm_modelaccessgroupbudgettable = _Table("model_access_group", self)
self.litellm_projecttable = _Table("project", self)
self.litellm_endusertable = _Table("enduser", self)
async def commit(self):
@ -117,6 +118,7 @@ class MockDB:
self.litellm_organizationtable = MockTable()
self.litellm_tagtable = MockTable()
self.litellm_modelaccessgroupbudgettable = MockTable()
self.litellm_projecttable = MockTable()
self.batch_calls: List[Dict[str, Any]] = []
self.batchers: List[MockBatcher] = []
@ -1575,13 +1577,19 @@ _INVALIDATION_CASES = [
"spend:model_access_group:gpt-4-group",
{"model_access_group:gpt-4-group"},
),
(
"litellm_projecttable",
type("Project", (), {"project_id": "proj-1"}),
"spend:project:proj-1",
{"project_id:proj-1"},
),
]
@pytest.mark.parametrize(
"table_attr, linked_row, counter_key, cache_keys",
_INVALIDATION_CASES,
ids=["team_membership", "key", "org", "tag", "model_access_group"],
ids=["team_membership", "key", "org", "tag", "model_access_group", "project"],
)
def test_budget_table_reset_invalidates_counters_and_management_cache(
reset_budget_job, mock_prisma_client, monkeypatch, table_attr, linked_row, counter_key, cache_keys
@ -1826,6 +1834,24 @@ def test_budget_table_reset_invalidates_every_access_group_not_just_the_first(
counter_cache.in_memory_cache.delete_cache.assert_any_call(key=f"spend:model_access_group:{name}")
def test_project_reset_zeroes_spend_on_due_tiers(reset_budget_job, mock_prisma_client, monkeypatch):
_make_counter_invalidation_job(monkeypatch)
mock_prisma_client.data["budget"] = [_budget_row(budget_id="budget-due", budget_duration="7d")]
mock_prisma_client.db.litellm_projecttable.set_find_many_results(
[type("Project", (), {"project_id": "proj-1", "spend": 12.0, "budget_id": "budget-due"})]
)
asyncio.run(reset_budget_job.reset_budget_for_litellm_budget_table())
expected_where = {"budget_id": {"in": ["budget-due"]}, "spend": {"gt": 0}}
assert mock_prisma_client.db.litellm_projecttable.find_many_calls == [{"where": expected_where}]
writes = _batch_writes(mock_prisma_client, "project", op="update_many")
assert len(writes) == 1
assert writes[0]["where"] == expected_where
assert writes[0]["data"] == {"spend": 0}
assert mock_prisma_client.db.batchers[0].committed is True
def test_budget_cascade_carries_access_group_overage_when_rollover_enabled(
rollover_enabled, reset_budget_job, mock_prisma_client, monkeypatch
):
@ -1971,6 +1997,7 @@ def test_budget_cascade_writes_land_in_a_single_transaction(reset_budget_job, mo
("org", "update_many"),
("tag", "update_many"),
("model_access_group", "update_many"),
("project", "update_many"),
("enduser", "update_many"),
("budget", "update_many"),
}

View file

@ -1,6 +1,7 @@
import asyncio
import copy
import json
import logging
import re
@ -11,16 +12,20 @@ from types import SimpleNamespace
from typing import Final
from unittest.mock import AsyncMock, MagicMock, call, patch
import httpx
import pytest
from prisma.errors import RawQueryError
from redis.exceptions import DataError
import litellm
from litellm.proxy._types import Litellm_EntityType
from litellm._logging import verbose_proxy_logger
from litellm.proxy._types import Litellm_EntityType, SpendUpdateQueueItem
from litellm.proxy.db.db_spend_update_writer import (
_TEAM_ADVISORY_LOCK_SQL,
_TEAM_MEMBER_SPEND_SQL,
DBSpendUpdateWriter,
)
from litellm.proxy.db.db_transaction_queue.spend_update_queue import SpendUpdateQueue
from litellm.proxy.db.db_transaction_queue.window_spend_update_queue import (
build_window_spend_transaction,
)
@ -1146,6 +1151,114 @@ async def test_batch_database_updates_queues_org_member_spend_for_the_request_us
assert transactions["org_member_list_transactions"] == {"organization_id::org1::user_id::u1": 0.1}
@pytest.mark.asyncio
async def test_project_spend_is_persisted_to_project_table_and_project_cache_is_evicted():
db_writer: Final = DBSpendUpdateWriter()
await db_writer._batch_database_updates(
response_cost=0.25,
user_id="u1",
hashed_token="t1",
team_id="team-1",
org_id=None,
end_user_id=None,
prisma_client=MagicMock(),
litellm_proxy_budget_name=None,
payload={"request_id": "req-1", "model": "gpt-4o-mini", "spend": 0.25},
project_id="proj-1",
)
await db_writer._batch_database_updates(
response_cost=0.5,
user_id="u1",
hashed_token="t1",
team_id="team-1",
org_id=None,
end_user_id=None,
prisma_client=MagicMock(),
litellm_proxy_budget_name=None,
payload={"request_id": "req-2", "model": "gpt-4o-mini", "spend": 0.5},
project_id="proj-1",
)
transactions: Final = await db_writer.spend_update_queue.flush_and_get_aggregated_db_spend_update_transactions()
assert transactions["project_list_transactions"] == {"proj-1": 0.75}
assert transactions["team_member_list_transactions"] == {"team_id::team-1::user_id::u1": 0.75}
mock_batcher: Final = MagicMock()
mock_prisma_client: Final = MagicMock()
mock_prisma_client.db.tx = MagicMock(return_value=_good_tx(mock_batcher))
user_api_key_cache: Final = MagicMock()
user_api_key_cache.async_delete_cache = AsyncMock()
proxy_logging: Final = MagicMock()
proxy_logging.call_details = {"user_api_key_cache": user_api_key_cache}
await db_writer._commit_spend_updates_to_db(
prisma_client=mock_prisma_client,
n_retry_times=0,
proxy_logging_obj=proxy_logging,
db_spend_update_transactions=transactions,
)
mock_batcher.litellm_projecttable.update_many.assert_called_once_with(
where={"project_id": "proj-1"},
data={"spend": {"increment": 0.75}},
)
user_api_key_cache.async_delete_cache.assert_any_await(key="project_id:proj-1")
@pytest.mark.asyncio
async def test_batch_database_updates_without_project_id_touches_no_project_row():
db_writer: Final = DBSpendUpdateWriter()
await db_writer._batch_database_updates(
response_cost=0.1,
user_id="u1",
hashed_token="t1",
team_id=None,
org_id=None,
end_user_id=None,
prisma_client=MagicMock(),
litellm_proxy_budget_name=None,
payload={"request_id": "req-1", "model": "gpt-4o-mini", "spend": 0.1},
)
transactions: Final = await db_writer.spend_update_queue.flush_and_get_aggregated_db_spend_update_transactions()
assert transactions["project_list_transactions"] == {}
@pytest.mark.asyncio
async def test_failed_project_enqueue_is_reported_and_does_not_drop_the_rest_of_the_batch(
caplog: pytest.LogCaptureFixture,
):
class _ProjectRejectingQueue(SpendUpdateQueue):
async def add_update(self, update: SpendUpdateQueueItem):
if update.get("entity_type") is Litellm_EntityType.PROJECT:
raise RuntimeError("project enqueue boom")
await super().add_update(update)
db_writer: Final = DBSpendUpdateWriter()
db_writer.spend_update_queue = _ProjectRejectingQueue()
with caplog.at_level(logging.ERROR, logger=verbose_proxy_logger.name):
await db_writer._batch_database_updates(
response_cost=0.25,
user_id="u1",
hashed_token="t1",
team_id="team-1",
org_id="org-1",
end_user_id=None,
prisma_client=MagicMock(),
litellm_proxy_budget_name=None,
payload={"request_id": "req-1", "model": "gpt-4o-mini", "spend": 0.25, "request_tags": ["tag-1"]},
project_id="proj-1",
)
assert any("proj-1" in record.getMessage() for record in caplog.records if record.levelno >= logging.ERROR)
transactions: Final = await db_writer.spend_update_queue.flush_and_get_aggregated_db_spend_update_transactions()
assert transactions["project_list_transactions"] == {}
assert transactions["tag_list_transactions"] == {"tag-1": 0.25}
assert transactions["key_list_transactions"] == {"t1": 0.25}
assert transactions["team_list_transactions"] == {"team-1": 0.25}
@pytest.mark.asyncio
async def test_add_spend_log_transaction_to_daily_tag_transaction_with_request_id():
"""
@ -1665,6 +1778,33 @@ async def test_update_daily_spend_keeps_failed_transactions_for_retry():
assert daily_spend_transactions == expected
@pytest.mark.asyncio
async def test_update_daily_spend_drops_the_batch_whose_failure_cannot_be_resent():
"""A reply lost after the statement was sent may already have applied, so the batch is
taken out of the caller's dict before the error propagates: whichever requeue the caller
runs afterwards, the Redis restore included, cannot send it a second time."""
def lose_the_reply() -> int:
raise httpx.ReadTimeout("no reply")
prisma_client = _RecordingPrisma(execute_raw=lose_the_reply)
daily_spend_transactions = {"user-key": _daily_txn(user_id="user-1")}
proxy_logging_obj = MagicMock()
proxy_logging_obj.failure_handler = AsyncMock()
with pytest.raises(httpx.ReadTimeout):
await DBSpendUpdateWriter._update_daily_spend(
n_retry_times=0,
prisma_client=prisma_client,
proxy_logging_obj=proxy_logging_obj,
daily_spend_transactions=daily_spend_transactions,
entity_type="user",
entity_id_field="user_id",
)
assert daily_spend_transactions == {}
@pytest.mark.asyncio
async def test_commit_key_spend_updates_includes_last_active():
"""
@ -2841,6 +2981,162 @@ async def test_failed_window_spend_commit_requeues_the_increments_and_continues_
assert requeued == (transaction,)
class _DailySpendFakeDB(_WindowSpendFakeDB):
"""Records the daily rollup upserts it is handed and fails the ones aimed at one table."""
def __init__(self, failing_table: str | None, failure: Exception | None = None) -> None:
super().__init__()
self.failing_table = failing_table
self.failure = failure
self.execute_raw_calls: list[Statement] = []
async def execute_raw(self, query: str, *args: object) -> int:
if self.failing_table is not None and self.failing_table in query:
raise self.failure if self.failure is not None else Exception("connection reset")
self.execute_raw_calls.append((query, args))
return len(args)
def _daily_upserts(db: _DailySpendFakeDB, table: str) -> list[Statement]:
return [statement for statement in db.execute_raw_calls if table in statement[0]]
def _postgres_rejection(sqlstate: str) -> RawQueryError:
return RawQueryError(
data={"user_facing_error": {"error_code": "P2010", "meta": {"code": sqlstate, "message": "db error"}}}
)
@pytest.mark.parametrize(
("failure", "lands_on_the_next_tick"),
[
pytest.param(httpx.ReadTimeout("no reply"), False, id="reply lost after the statement was sent"),
pytest.param(httpx.ConnectError("refused"), True, id="statement never reached the database"),
pytest.param(_postgres_rejection("22021"), False, id="postgres refused the data itself"),
pytest.param(_postgres_rejection("23502"), False, id="postgres refused a constraint violation"),
pytest.param(_postgres_rejection("42P01"), True, id="table missing"),
pytest.param(_postgres_rejection("57014"), True, id="statement cancelled"),
],
)
@pytest.mark.asyncio
async def test_failed_daily_spend_commit_is_requeued_only_when_the_rows_are_provably_uncommitted(
failure: Exception, lands_on_the_next_tick: bool
):
"""A lost reply means the statement may already have applied, and re-sending it stacks a
second increment into the same transaction (LIT-4823); a row Postgres refuses would fail
every tick forever. Both are dropped loudly. Every other failure left nothing committed,
so its rows go back on the queue and land on the next tick."""
db_writer = DBSpendUpdateWriter()
await db_writer.daily_spend_update_queue.add_update({"user-key": _daily_txn(user_id="user-1")})
db = _DailySpendFakeDB(failing_table="LiteLLM_DailyUserSpend", failure=failure)
db_writer._flush_tool_discovery_queue = AsyncMock()
proxy_logging_obj = MagicMock()
proxy_logging_obj.failure_handler = AsyncMock()
await db_writer._commit_spend_updates_to_db_without_redis_buffer(
prisma_client=_WindowSpendFakePrisma(db), n_retry_times=0, proxy_logging_obj=proxy_logging_obj
)
db.failing_table = None
await db_writer._commit_spend_updates_to_db_without_redis_buffer(
prisma_client=_WindowSpendFakePrisma(db), n_retry_times=0, proxy_logging_obj=proxy_logging_obj
)
assert len(_daily_upserts(db, "LiteLLM_DailyUserSpend")) == (1 if lands_on_the_next_tick else 0)
assert db_writer.daily_spend_update_queue.update_queue.empty()
@pytest.mark.asyncio
async def test_failed_daily_spend_commit_drops_only_the_batch_that_was_sent():
"""A tick holding more than one batch of 100 rows sends them one statement at a time, and
a reply lost on one statement says nothing about the batches after it: only the batch that
was on the wire is dropped, the ones never sent go back on the queue and land next tick."""
db_writer = DBSpendUpdateWriter()
await db_writer.daily_spend_update_queue.add_update(
{f"user-{i:03d}": _daily_txn(user_id=f"user-{i:03d}") for i in range(150)}
)
db = _DailySpendFakeDB(failing_table="LiteLLM_DailyUserSpend", failure=httpx.ReadTimeout("no reply"))
db_writer._flush_tool_discovery_queue = AsyncMock()
proxy_logging_obj = MagicMock()
proxy_logging_obj.failure_handler = AsyncMock()
await db_writer._commit_spend_updates_to_db_without_redis_buffer(
prisma_client=_WindowSpendFakePrisma(db), n_retry_times=0, proxy_logging_obj=proxy_logging_obj
)
db.failing_table = None
await db_writer._commit_spend_updates_to_db_without_redis_buffer(
prisma_client=_WindowSpendFakePrisma(db), n_retry_times=0, proxy_logging_obj=proxy_logging_obj
)
(upsert,) = _daily_upserts(db, "LiteLLM_DailyUserSpend")
assert _row_values(upsert, "user_id") == [f"user-{i:03d}" for i in range(100, 150)]
assert db_writer.daily_spend_update_queue.update_queue.empty()
@pytest.mark.asyncio
async def test_failed_daily_spend_commit_requeues_the_rows_and_flushes_the_other_tables():
"""With the Redis buffer off, a daily batch that failed to commit was discarded along
with the tick's exception, so the Usage page stayed short of LiteLLM_SpendLogs for good.
The uncommitted rows must go back on their queue and land on the next tick, and the
other daily tables must still be flushed on the failing tick."""
db_writer = DBSpendUpdateWriter()
await db_writer.daily_spend_update_queue.add_update({"user-key": _daily_txn(user_id="user-1")})
team_txn = {key: value for key, value in _daily_txn().items() if key != "user_id"} | {"team_id": "team-1"}
await db_writer.daily_team_spend_update_queue.add_update({"team-key": team_txn})
db = _DailySpendFakeDB(failing_table="LiteLLM_DailyUserSpend")
db_writer._flush_tool_discovery_queue = AsyncMock()
proxy_logging_obj = MagicMock()
proxy_logging_obj.failure_handler = AsyncMock()
await db_writer._commit_spend_updates_to_db_without_redis_buffer(
prisma_client=_WindowSpendFakePrisma(db), n_retry_times=0, proxy_logging_obj=proxy_logging_obj
)
assert _daily_upserts(db, "LiteLLM_DailyUserSpend") == []
(team_upsert,) = _daily_upserts(db, "LiteLLM_DailyTeamSpend")
assert _row_values(team_upsert, "team_id") == ["team-1"]
db_writer._flush_tool_discovery_queue.assert_called_once()
db.failing_table = None
await db_writer._commit_spend_updates_to_db_without_redis_buffer(
prisma_client=_WindowSpendFakePrisma(db), n_retry_times=0, proxy_logging_obj=proxy_logging_obj
)
(user_upsert,) = _daily_upserts(db, "LiteLLM_DailyUserSpend")
assert _row_values(user_upsert, "user_id") == ["user-1"]
assert _row_values(user_upsert, "spend") == [0.1]
assert len(_daily_upserts(db, "LiteLLM_DailyTeamSpend")) == 1
assert db_writer.daily_spend_update_queue.update_queue.empty()
@pytest.mark.asyncio
async def test_failed_daily_tag_spend_commit_requeues_the_rows():
"""The tag rollup drains on its own scheduler job with the same no-Redis drop:
a failed LiteLLM_DailyTagSpend commit has to put the rows back for the next tick."""
db_writer = DBSpendUpdateWriter()
tag_txn = {key: value for key, value in _daily_txn().items() if key != "user_id"} | {"tag": "tag-1"}
await db_writer.daily_tag_spend_update_queue.add_update({"tag-key": tag_txn})
db = _DailySpendFakeDB(failing_table="LiteLLM_DailyTagSpend")
proxy_logging_obj = MagicMock()
proxy_logging_obj.failure_handler = AsyncMock()
await db_writer._commit_daily_tag_spend_to_db(
prisma_client=_WindowSpendFakePrisma(db), n_retry_times=0, proxy_logging_obj=proxy_logging_obj
)
assert _daily_upserts(db, "LiteLLM_DailyTagSpend") == []
assert not db_writer.daily_tag_spend_update_queue.update_queue.empty()
db.failing_table = None
await db_writer._commit_daily_tag_spend_to_db(
prisma_client=_WindowSpendFakePrisma(db), n_retry_times=0, proxy_logging_obj=proxy_logging_obj
)
(tag_upsert,) = _daily_upserts(db, "LiteLLM_DailyTagSpend")
assert _row_values(tag_upsert, "tag") == ["tag-1"]
assert _row_values(tag_upsert, "spend") == [0.1]
assert db_writer.daily_tag_spend_update_queue.update_queue.empty()
@pytest.mark.asyncio
async def test_failed_window_spend_commit_from_redis_is_restored_to_redis():
"""The Redis drain is destructive, so a failed window commit has to push

View file

@ -665,6 +665,28 @@ def test_is_deadlock_error_excludes_non_deadlocks(error):
assert PrismaDBExceptionHandler.is_deadlock_error(error) is False
@pytest.mark.parametrize(
("error", "sqlstate"),
[
(
RawQueryError(
data={"user_facing_error": {"error_code": "P2010", "meta": {"code": "22021", "message": "m"}}}
),
"22021",
),
(RawQueryError(data={"user_facing_error": {"error_code": "P2010", "meta": {"message": "m"}}}), None),
(RawQueryError(data={"user_facing_error": {"error_code": "P2010", "meta": {"code": 42, "message": "m"}}}), None),
(prisma_errors.DataError(data={"user_facing_error": {"meta": None}}), None),
(PrismaError("db error"), None),
(httpx.ReadTimeout("no reply"), None),
],
)
def test_postgres_sqlstate_reads_the_code_prisma_attached_to_the_failed_statement(error: Exception, sqlstate: str | None):
"""Only a prisma data error carrying Postgres's own error code yields a SQLSTATE; a
codeless or malformed payload, an engine-level error, and a transport error yield None."""
assert PrismaDBExceptionHandler.postgres_sqlstate(error) == sqlstate
READ_ONLY_CONNECTOR_ERROR: Final = (
"Error occurred during query execution:\nConnectorError(ConnectorError { user_facing_error: None, "
'kind: QueryError(PostgresError { code: "25006", message: "cannot execute UPDATE in a read-only transaction", '

View file

@ -67,12 +67,14 @@ class _FakePrismaClient:
error: Exception | None = None,
end_user_row: SimpleNamespace | None = None,
end_user_error: Exception | None = None,
project_row: SimpleNamespace | None = None,
) -> None:
self.db = SimpleNamespace(
litellm_budgetwindowspend=_FakeFindUniqueTable(row=row, error=error),
litellm_spendlogs=_FakeSpendLogsTable(total=spend_logs_total),
litellm_endusertable=_FakeFindUniqueTable(row=end_user_row, error=end_user_error),
litellm_verificationtoken=_InFlightCountingTable(),
litellm_projecttable=_FakeFindUniqueTable(row=project_row),
)
@ -428,6 +430,21 @@ async def test_from_db_bounds_in_flight_prisma_requests_across_counter_keys():
assert prisma.db.litellm_verificationtoken.max_in_flight == PROXY_DB_LOOKUP_MAX_CONCURRENCY
@pytest.mark.asyncio
async def test_from_db_reseeds_project_counter_from_the_project_row():
prisma: Final = _FakePrismaClient(project_row=SimpleNamespace(project_id="proj-1", spend=7.25))
assert await SpendCounterReseed.from_db(prisma_client=prisma, counter_key="spend:project:proj-1") == 7.25
assert prisma.db.litellm_projecttable.where_clauses == [{"project_id": "proj-1"}]
@pytest.mark.asyncio
async def test_from_db_returns_none_for_a_missing_project_row():
prisma: Final = _FakePrismaClient(project_row=None)
assert await SpendCounterReseed.from_db(prisma_client=prisma, counter_key="spend:project:proj-1") is None
@pytest.mark.asyncio
async def test_from_db_still_never_reads_the_end_user_row():
"""A cold end-user counter keeps seeding from the cached end-user object the auth

View file

@ -18,7 +18,9 @@ from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from httpx import Request, Response
import litellm
from litellm import DualCache
from litellm.constants import DEFAULT_OPENAI_MODERATIONS_MODEL
from litellm.proxy._types import UserAPIKeyAuth
from litellm.types.utils import Choices, Message, ModelResponse
@ -764,6 +766,40 @@ async def test_openai_moderation_inspects_multimodal_content(monkeypatch, user_a
assert seen_inputs == ["alpha beta"]
@pytest.mark.asyncio
@pytest.mark.parametrize(
("configured_after_init", "expected_model"),
[("omni-moderation-2024-09-26", "omni-moderation-2024-09-26"), (None, DEFAULT_OPENAI_MODERATIONS_MODEL)],
)
async def test_openai_moderation_reads_model_name_at_call_time(
monkeypatch, user_api_key, configured_after_init, expected_model
):
"""``litellm_settings`` applies ``callbacks`` and ``openai_moderations_model_name`` in YAML
order, so the hook must resolve the model when it runs, not when it is constructed."""
from enterprise.enterprise_hooks.openai_moderation import (
_ENTERPRISE_OpenAI_Moderation,
)
monkeypatch.setattr(litellm, "openai_moderations_model_name", None)
guard = _ENTERPRISE_OpenAI_Moderation()
monkeypatch.setattr(litellm, "openai_moderations_model_name", configured_after_init)
class FakeModeration:
results = [type("R", (), {"flagged": False})()]
fake_router = MagicMock()
fake_router.amoderation = AsyncMock(return_value=FakeModeration())
monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", fake_router, raising=False)
await guard.async_moderation_hook(
data={"messages": [{"role": "user", "content": "hello"}]},
user_api_key_dict=user_api_key,
call_type="acompletion",
)
fake_router.amoderation.assert_awaited_once_with(model=expected_model, input="hello")
# ── Google Text Moderation ────────────────────────────────────────────────────

View file

@ -586,6 +586,7 @@ async def test_update_database_and_spend_counters_updates_counters_after_db_upda
tags=["tag-a"],
request_started_at=start_time,
model_access_groups=("premium",),
project_id=None,
)
@ -1371,6 +1372,7 @@ async def test_enrich_failure_metadata_with_full_key_lookup():
mock_key_obj.user_id = "fetched-user-id"
mock_key_obj.team_id = "fetched-team-id"
mock_key_obj.org_id = "fetched-org-id"
mock_key_obj.project_id = "fetched-project-id"
mock_team_obj = MagicMock()
mock_team_obj.team_alias = "fetched-team-alias"
@ -1394,12 +1396,14 @@ async def test_enrich_failure_metadata_with_full_key_lookup():
"user_api_key_team_id": None,
"user_api_key_team_alias": None,
"user_api_key_org_id": None,
"user_api_key_project_id": None,
}
result = await _ProxyDBLogger._enrich_failure_metadata_with_key_info(metadata)
assert result["user_api_key_alias"] == "fetched-key-alias"
assert result["user_api_key_user_id"] == "fetched-user-id"
assert result["user_api_key_team_id"] == "fetched-team-id"
assert result["user_api_key_org_id"] == "fetched-org-id"
assert result["user_api_key_project_id"] == "fetched-project-id"
assert result["user_api_key_team_alias"] == "fetched-team-alias"

View file

@ -39,11 +39,10 @@ from litellm.proxy._types import (
from litellm.models.object_permission import LiteLLM_ObjectPermissionTable
from litellm.proxy.auth.auth_checks import (
_delete_cache_key_object,
_project_cache_key,
jwt_key_mapping_cache_key,
)
from litellm.proxy.auth.user_api_key_auth import UserAPIKeyAuth
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache, project_cache_key
from litellm.litellm_core_utils.duration_parser import duration_in_seconds
from litellm.proxy.management_endpoints.key_management_endpoints import (
_check_org_key_limits,
@ -19465,7 +19464,7 @@ async def test_regenerate_key_repoints_live_membership_not_the_key_row_it_read(
async def _cache_with_project(project_id: str, project_models: list[str]) -> UserApiKeyCache:
user_api_key_cache = UserApiKeyCache()
await user_api_key_cache.async_set_cache(
key=_project_cache_key(project_id),
key=project_cache_key(project_id),
value=LiteLLM_ProjectTableCachedObj(project_id=project_id, team_id="team-lit-5823", models=project_models),
model_type=LiteLLM_ProjectTableCachedObj,
)

View file

@ -1,13 +1,17 @@
import json
from collections.abc import Mapping
from datetime import datetime, timezone
from litellm._uuid import uuid
from unittest.mock import AsyncMock, MagicMock
from typing import Final
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
import litellm
from litellm._uuid import uuid
from litellm.proxy._types import (
LiteLLM_BudgetTable,
LiteLLM_TeamMembership,
LiteLLM_TeamTable,
LiteLLM_UserTable,
Member,
UserAPIKeyAuth,
@ -164,21 +168,12 @@ async def test_management_otel_span_redacts_nested_submission_env_var_secrets(
@pytest.mark.asyncio
async def test_add_new_member_clones_default_team_budget_id():
"""
Test that add_new_member CLONES the team's default member budget when
max_budget_in_team is None and a default_team_budget_id is provided.
Cloning (rather than sharing the same budget row) is what lets admins later
edit one member's budget without mutating every other member's budget.
"""
async def test_add_new_member_links_default_team_budget_id():
from litellm.proxy._types import LitellmUserRoles
# Setup test data
test_user_id = "test_user_123"
test_team_id = "test_team_456"
test_default_budget_id = "default_budget_789"
test_cloned_budget_id = "cloned_budget_xyz"
test_admin_name = "test_admin"
new_member = Member(user_id=test_user_id, role="user")
@ -202,36 +197,19 @@ async def test_add_new_member_clones_default_team_budget_id():
return_value=mock_user_response
)
# Mock the default budget row fetched for cloning.
mock_default_budget_row = MagicMock()
mock_default_budget_row.model_dump.return_value = {
"budget_id": test_default_budget_id,
"max_budget": 100.0,
"soft_budget": None,
"max_parallel_requests": None,
"tpm_limit": 1000,
"rpm_limit": None,
"model_max_budget": None,
"budget_duration": "1d",
"allowed_models": [],
}
mock_default_budget_row.budget_id = test_default_budget_id
mock_prisma_client.db.litellm_budgettable.find_unique = AsyncMock(
return_value=mock_default_budget_row
)
# Mock the cloned budget row that .create() returns.
mock_cloned_budget_row = MagicMock()
mock_cloned_budget_row.budget_id = test_cloned_budget_id
mock_prisma_client.db.litellm_budgettable.create = AsyncMock(
return_value=mock_cloned_budget_row
)
mock_prisma_client.db.litellm_budgettable.create = AsyncMock()
# Mock the team membership creation
mock_team_membership_response = MagicMock()
mock_team_membership_response.model_dump.return_value = {
"team_id": test_team_id,
"user_id": test_user_id,
"budget_id": test_cloned_budget_id,
"budget_id": test_default_budget_id,
"litellm_budget_table": None,
}
mock_prisma_client.db.litellm_teammembership.upsert = AsyncMock(
@ -251,33 +229,71 @@ async def test_add_new_member_clones_default_team_budget_id():
assert result_user is not None
assert result_user.user_id == test_user_id
# Membership should be linked to the new cloned budget, not the shared default.
assert result_team_membership is not None
assert result_team_membership.budget_id == test_cloned_budget_id
assert result_team_membership.budget_id != test_default_budget_id
assert result_team_membership.budget_id == test_default_budget_id
mock_prisma_client.db.litellm_usertable.upsert.assert_called_once()
mock_prisma_client.db.litellm_teammembership.upsert.assert_called_once()
# The clone must have happened: find_unique on the default, create for the clone.
mock_prisma_client.db.litellm_budgettable.find_unique.assert_called_once_with(
where={"budget_id": test_default_budget_id}
)
mock_prisma_client.db.litellm_budgettable.create.assert_called_once()
cloned_create_data = (
mock_prisma_client.db.litellm_budgettable.create.call_args.kwargs["data"]
)
# Cloned values from the default budget row
assert cloned_create_data["max_budget"] == 100.0
assert cloned_create_data["tpm_limit"] == 1000
assert cloned_create_data["budget_duration"] == "1d"
assert cloned_create_data["created_by"] == user_api_key_dict.user_id
mock_prisma_client.db.litellm_budgettable.create.assert_not_called()
team_membership_call_args = (
mock_prisma_client.db.litellm_teammembership.upsert.call_args
)
create_data = team_membership_call_args.kwargs["data"]["create"]
assert create_data["budget_id"] == test_cloned_budget_id
assert create_data["budget_id"] == test_default_budget_id
@pytest.mark.asyncio
async def test_add_new_member_no_budget_when_default_budget_row_is_missing():
from litellm.proxy._types import LitellmUserRoles
new_member = Member(user_id="missing-default-user", role="user")
user_api_key_dict = UserAPIKeyAuth(
user_id="admin_user", user_role=LitellmUserRoles.PROXY_ADMIN
)
mock_prisma_client = AsyncMock()
mock_user_response = MagicMock()
mock_user_response.model_dump.return_value = {
"user_id": "missing-default-user",
"user_email": None,
"teams": ["team-md"],
"user_role": "internal_user",
}
mock_prisma_client.db.litellm_usertable.upsert = AsyncMock(return_value=mock_user_response)
mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(
return_value=mock_user_response
)
mock_prisma_client.db.litellm_budgettable.find_unique = AsyncMock(return_value=None)
mock_prisma_client.db.litellm_budgettable.create = AsyncMock()
mock_membership = MagicMock()
mock_membership.model_dump.return_value = {
"team_id": "team-md",
"user_id": "missing-default-user",
"budget_id": None,
"litellm_budget_table": None,
}
mock_prisma_client.db.litellm_teammembership.upsert = AsyncMock(return_value=mock_membership)
_, result_team_membership = await add_new_member(
new_member=new_member,
max_budget_in_team=None,
prisma_client=mock_prisma_client,
team_id="team-md",
user_api_key_dict=user_api_key_dict,
litellm_proxy_admin_name="test_admin",
default_team_budget_id="deleted-default",
)
assert result_team_membership is not None
assert result_team_membership.budget_id is None
mock_prisma_client.db.litellm_budgettable.create.assert_not_called()
upsert_kwargs = mock_prisma_client.db.litellm_teammembership.upsert.call_args.kwargs
assert upsert_kwargs["data"]["create"] == {"user_id": "missing-default-user", "team_id": "team-md"}
@pytest.mark.asyncio
@ -661,18 +677,12 @@ async def test_add_new_member_persists_budget_duration_without_max_budget():
@pytest.mark.asyncio
async def test_add_new_member_with_user_email_clones_default_budget():
"""
Test add_new_member with user_email instead of user_id and a team default
budget. The default budget should be CLONED into a new private row for
this user, not shared with other members of the team.
"""
async def test_add_new_member_with_user_email_links_default_budget():
from litellm.proxy._types import LitellmUserRoles
test_user_email = "test@example.com"
test_team_id = "test_team_456"
test_default_budget_id = "default_budget_789"
test_cloned_budget_id = "cloned_budget_for_email_user"
test_admin_name = "test_admin"
new_member = Member(user_email=test_user_email, role="user")
@ -694,35 +704,18 @@ async def test_add_new_member_with_user_email_clones_default_budget():
}
mock_prisma_client.insert_data = AsyncMock(return_value=mock_user_response)
# Default budget that will be cloned
mock_default_budget_row = MagicMock()
mock_default_budget_row.model_dump.return_value = {
"budget_id": test_default_budget_id,
"max_budget": 25.0,
"soft_budget": None,
"max_parallel_requests": None,
"tpm_limit": None,
"rpm_limit": None,
"model_max_budget": None,
"budget_duration": None,
"allowed_models": [],
}
mock_default_budget_row.budget_id = test_default_budget_id
mock_prisma_client.db.litellm_budgettable.find_unique = AsyncMock(
return_value=mock_default_budget_row
)
# Cloned budget result
mock_cloned_budget_row = MagicMock()
mock_cloned_budget_row.budget_id = test_cloned_budget_id
mock_prisma_client.db.litellm_budgettable.create = AsyncMock(
return_value=mock_cloned_budget_row
)
mock_prisma_client.db.litellm_budgettable.create = AsyncMock()
mock_team_membership_response = MagicMock()
mock_team_membership_response.model_dump.return_value = {
"team_id": test_team_id,
"user_id": "generated_user_id",
"budget_id": test_cloned_budget_id,
"budget_id": test_default_budget_id,
"litellm_budget_table": None,
}
mock_prisma_client.db.litellm_teammembership.upsert = AsyncMock(
@ -742,9 +735,8 @@ async def test_add_new_member_with_user_email_clones_default_budget():
assert result_user is not None
assert result_user.user_email == test_user_email
# Membership should point at the cloned (private) budget, not the shared default.
assert result_team_membership is not None
assert result_team_membership.budget_id == test_cloned_budget_id
assert result_team_membership.budget_id == test_default_budget_id
mock_prisma_client.get_data.assert_called_once_with(
key_val={"user_email": test_user_email},
@ -758,11 +750,166 @@ async def test_add_new_member_with_user_email_clones_default_budget():
assert insert_data["user_email"] == test_user_email
assert insert_data["teams"] == [test_team_id]
# Confirm the clone path ran
mock_prisma_client.db.litellm_budgettable.find_unique.assert_called_once_with(
where={"budget_id": test_default_budget_id}
)
mock_prisma_client.db.litellm_budgettable.create.assert_called_once()
mock_prisma_client.db.litellm_budgettable.create.assert_not_called()
class _FakeBudgetTable:
def __init__(self) -> None:
self.rows: dict[str, dict[str, object]] = {}
def _record(self, budget_id: str) -> LiteLLM_BudgetTable:
row: Final = self.rows[budget_id]
return LiteLLM_BudgetTable(**{k: v for k, v in row.items() if k in LiteLLM_BudgetTable.model_fields})
async def create(
self, *, data: Mapping[str, object], include: Mapping[str, bool] | None = None
) -> LiteLLM_BudgetTable:
budget_id: Final = str(data.get("budget_id") or uuid.uuid4())
self.rows[budget_id] = {**data, "budget_id": budget_id}
return self._record(budget_id)
async def find_unique(self, *, where: Mapping[str, str]) -> LiteLLM_BudgetTable | None:
return self._record(where["budget_id"]) if where["budget_id"] in self.rows else None
async def update(self, *, where: Mapping[str, str], data: Mapping[str, object]) -> LiteLLM_BudgetTable:
self.rows[where["budget_id"]] = {**self.rows[where["budget_id"]], **data}
return self._record(where["budget_id"])
class _FakeMembershipTable:
def __init__(self, budgets: _FakeBudgetTable) -> None:
self.budgets: Final = budgets
self.budget_ids: dict[tuple[str, str], str | None] = {}
def membership(self, team_id: str, user_id: str) -> LiteLLM_TeamMembership:
budget_id: Final = self.budget_ids[(team_id, user_id)]
return LiteLLM_TeamMembership(
user_id=user_id,
team_id=team_id,
budget_id=budget_id,
litellm_budget_table=self.budgets._record(budget_id) if budget_id is not None else None,
)
@staticmethod
def _linked_budget_id(row: Mapping[str, object]) -> str | None:
budget_id: Final = row.get("budget_id")
if isinstance(budget_id, str):
return budget_id
connect: Final = row.get("litellm_budget_table")
if isinstance(connect, dict):
return connect["connect"]["budget_id"]
return None
async def upsert(
self,
*,
where: Mapping[str, Mapping[str, str]],
data: Mapping[str, Mapping[str, object]],
include: Mapping[str, bool] | None = None,
) -> LiteLLM_TeamMembership:
key: Final = where["user_id_team_id"]
membership_key: Final = (key["team_id"], key["user_id"])
if membership_key not in self.budget_ids:
self.budget_ids[membership_key] = self._linked_budget_id(data["create"])
elif "litellm_budget_table" in data["update"]:
self.budget_ids[membership_key] = self._linked_budget_id(data["update"])
return self.membership(*membership_key)
class _FakeUserTable:
async def upsert(self, *, where: Mapping[str, str], data: Mapping[str, Mapping[str, object]]) -> LiteLLM_UserTable:
return LiteLLM_UserTable(user_id=where["user_id"], teams=list(data["create"].get("teams", [])))
async def update_many(self, *, where: Mapping[str, object], data: Mapping[str, object]) -> int:
return 1
class _FakeDb:
def __init__(self) -> None:
self.litellm_budgettable: Final = _FakeBudgetTable()
self.litellm_teammembership: Final = _FakeMembershipTable(self.litellm_budgettable)
self.litellm_usertable: Final = _FakeUserTable()
@pytest.mark.asyncio
async def test_team_update_reaches_inherited_members_but_not_overridden_ones():
from litellm.proxy._types import LitellmUserRoles
from litellm.proxy.auth.auth_checks import _check_team_member_budget
from litellm.proxy.management_endpoints.common_utils import _upsert_budget_and_membership
from litellm.proxy.management_endpoints.team_endpoints import TeamMemberBudgetHandler
from litellm.proxy.utils import ProxyLogging
db: Final = _FakeDb()
prisma_client: Final = MagicMock()
prisma_client.db = db
admin: Final = UserAPIKeyAuth(user_id="admin_user", user_role=LitellmUserRoles.PROXY_ADMIN)
team_id: Final = "team-shared-default"
default_budget: Final = await db.litellm_budgettable.create(data={"budget_id": "team-default", "max_budget": 100.0})
team: Final = LiteLLM_TeamTable(team_id=team_id, metadata={"team_member_budget_id": default_budget.budget_id})
for user_id in ("inherits", "overridden"):
await add_new_member(
new_member=Member(user_id=user_id, role="user"),
max_budget_in_team=None,
prisma_client=prisma_client,
team_id=team_id,
user_api_key_dict=admin,
litellm_proxy_admin_name="admin",
default_team_budget_id=default_budget.budget_id,
)
await _upsert_budget_and_membership(
db,
team_id=team_id,
user_id="overridden",
existing_budget_id=default_budget.budget_id,
user_api_key_dict=admin,
budget_patch={"max_budget": 50.0},
team_default_budget_id=default_budget.budget_id,
)
assert db.litellm_teammembership.membership(team_id, "inherits").budget_id == default_budget.budget_id
assert db.litellm_teammembership.membership(team_id, "overridden").budget_id != default_budget.budget_id
assert db.litellm_budgettable.rows[default_budget.budget_id]["max_budget"] == 100.0
with patch( # test-quality-ok: update_budget reads this module global; no dependency injection seam exists
"litellm.proxy.proxy_server.prisma_client", prisma_client
):
await TeamMemberBudgetHandler.upsert_team_member_budget_table(
team_table=team,
user_api_key_dict=admin,
updated_kv={},
team_member_budget=1.0,
)
async def spend_from_membership(counter_key: str, fallback_spend: float, max_budget: float | None = None) -> float:
return fallback_spend
async def check(user_id: str, spend: float) -> None:
membership: Final = db.litellm_teammembership.membership(team_id, user_id).model_copy(update={"spend": spend})
with patch( # test-quality-ok: production auth reads this module global; no dependency injection seam exists
"litellm.proxy.proxy_server.get_current_spend", spend_from_membership
):
await _check_team_member_budget(
team_object=team,
user_object=LiteLLM_UserTable(user_id=user_id),
valid_token=UserAPIKeyAuth(token="tok", user_id=user_id, team_id=team_id),
prisma_client=prisma_client,
user_api_key_cache=MagicMock(),
proxy_logging_obj=ProxyLogging(user_api_key_cache=None),
team_membership=membership,
team_membership_loaded=True,
)
with pytest.raises(litellm.BudgetExceededError) as exc_info:
await check("inherits", spend=2.0)
assert exc_info.value.max_budget == 1.0
await check("overridden", spend=2.0)
with pytest.raises(litellm.BudgetExceededError) as exc_info:
await check("overridden", spend=60.0)
assert exc_info.value.max_budget == 50.0
@pytest.mark.asyncio

View file

@ -3399,9 +3399,8 @@ async def test_ProxyConfig__add_router_settings_from_db_config_updates_router():
fake_prisma.db.litellm_config.find_first = AsyncMock(
return_value=SimpleNamespace(param_value={"timeout": 30, "retries": 2, "fallbacks": []})
)
config_data = {"router_settings": {"timeout": 10}}
pc.router_settings.load_yaml({"timeout": 10})
await pc._add_router_settings_from_db_config(
config_data=config_data,
llm_router=fake_router,
prisma_client=fake_prisma,
)
@ -3421,7 +3420,7 @@ async def test_ProxyConfig__add_router_settings_from_db_config_updates_router():
async def test_ProxyConfig__add_router_settings_from_db_config_none_router_noop():
pc = ProxyConfig()
# No router and no prisma — should silently return.
await pc._add_router_settings_from_db_config(config_data={}, llm_router=None, prisma_client=None)
await pc._add_router_settings_from_db_config(llm_router=None, prisma_client=None)
# Error-style: bad call signature raises.
with pytest.raises(TypeError):
await pc._add_router_settings_from_db_config() # type: ignore[call-arg]

View file

@ -139,6 +139,108 @@ def test_config_update_persists_disable_cooldowns(client, auth_as, mock_prisma,
assert persisted["disable_cooldowns"] is True
@pytest.mark.parametrize(
("section", "store_attr", "yaml_values", "changed_values"),
[
("general_settings", "settings", {"alerting": ["slack"]}, {"alerting": ["email"]}),
("litellm_settings", "litellm_settings", {"success_callback": ["langfuse"]}, {"success_callback": ["otel"]}),
("router_settings", "router_settings", {"num_retries": 0}, {"num_retries": 2}),
],
)
def test_config_update_rejects_config_owned_keys_and_accepts_the_same_value(
client, auth_as, mock_prisma, monkeypatch, section, store_attr, yaml_values, changed_values
):
from litellm.proxy import proxy_server as ps
from litellm.proxy._types import LitellmUserRoles
table = _install_litellm_config(mock_prisma)
monkeypatch.setattr(ps, "prisma_client", mock_prisma)
monkeypatch.setattr(ps.proxy_config, "add_deployment", AsyncMock())
store = getattr(ps.proxy_config, store_attr)
store.load_yaml(yaml_values)
try:
with auth_as(LitellmUserRoles.PROXY_ADMIN):
rejected = client.post("/config/update", json={section: changed_values})
rejected_message = rejected.json()["error"]["message"]
table.upsert.assert_not_called()
accepted = client.post("/config/update", json={section: yaml_values})
finally:
store.load_yaml({})
assert rejected.status_code == 400
assert f"{section} key '{next(iter(yaml_values))}' is set in the config file and cannot be changed here" in (
rejected_message
)
assert accepted.status_code == 200
persisted = json.loads(table.upsert.call_args.kwargs["data"]["create"]["param_value"])
assert persisted[next(iter(yaml_values))] == yaml_values[next(iter(yaml_values))]
def test_config_update_persists_only_the_general_settings_keys_the_request_set(
client, auth_as, mock_prisma, monkeypatch
):
from litellm.proxy import proxy_server as ps
from litellm.proxy._types import LitellmUserRoles
table = _install_litellm_config(mock_prisma)
monkeypatch.setattr(ps, "prisma_client", mock_prisma)
monkeypatch.setattr(ps.proxy_config, "add_deployment", AsyncMock())
ps.proxy_config.settings.load_yaml({"health_check_interval": 60})
try:
with auth_as(LitellmUserRoles.PROXY_ADMIN):
response = client.post("/config/update", json={"general_settings": {"alerting_threshold": 600}})
finally:
ps.proxy_config.settings.load_yaml({})
assert response.status_code == 200
persisted = json.loads(table.upsert.call_args.kwargs["data"]["create"]["param_value"])
assert persisted == {"alerting_threshold": 600}
def test_config_update_persists_only_the_router_settings_keys_the_request_set(
client, auth_as, mock_prisma, monkeypatch
):
from litellm.proxy import proxy_server as ps
from litellm.proxy._types import LitellmUserRoles
table = _install_litellm_config(mock_prisma)
monkeypatch.setattr(ps, "prisma_client", mock_prisma)
monkeypatch.setattr(ps.proxy_config, "add_deployment", AsyncMock())
ps.proxy_config.router_settings.load_yaml({"model_group_alias": {"opus": "claude-opus-5"}})
try:
with auth_as(LitellmUserRoles.PROXY_ADMIN):
response = client.post(
"/config/update", json={"router_settings": {"retry_policy": {"TimeoutErrorRetries": 3}}}
)
finally:
ps.proxy_config.router_settings.load_yaml({})
assert response.status_code == 200, response.text
persisted = json.loads(table.upsert.call_args.kwargs["data"]["create"]["param_value"])
assert persisted == {"retry_policy": {"TimeoutErrorRetries": 3}}
def test_config_update_accepts_a_config_owned_success_callback_the_file_spells_in_mixed_case(
client, auth_as, mock_prisma, monkeypatch
):
from litellm.proxy import proxy_server as ps
from litellm.proxy._types import LitellmUserRoles
table = _install_litellm_config(mock_prisma)
monkeypatch.setattr(ps, "prisma_client", mock_prisma)
monkeypatch.setattr(ps.proxy_config, "add_deployment", AsyncMock())
ps.proxy_config.litellm_settings.load_yaml({"success_callback": ["Langfuse"]})
try:
with auth_as(LitellmUserRoles.PROXY_ADMIN):
response = client.post("/config/update", json={"litellm_settings": {"success_callback": ["Langfuse"]}})
finally:
ps.proxy_config.litellm_settings.load_yaml({})
assert response.status_code == 200
persisted = json.loads(table.upsert.call_args.kwargs["data"]["create"]["param_value"])
assert persisted["success_callback"] == ["langfuse"]
def test_config_update_rejects_assistants_config(client, auth_as, mock_prisma, monkeypatch):
from litellm.proxy import proxy_server as ps
from litellm.proxy._types import LitellmUserRoles

View file

@ -17,15 +17,20 @@ from __future__ import annotations
import asyncio
import json
from collections.abc import AsyncIterator
from typing import Final, Literal
from unittest.mock import AsyncMock, MagicMock
import httpx
import pytest
from fastapi import Response
from fastapi import HTTPException, Response
from fastapi.responses import StreamingResponse
from openai import APIError as OpenAIAPIError
from pydantic import BaseModel
import litellm
from litellm.constants import RETURN_RAW_MODEL_NAME_METADATA_KEY
import litellm.proxy.proxy_server as ps
from litellm.constants import RETURN_RAW_MODEL_NAME_METADATA_KEY
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.proxy_server import (
_apply_streaming_chunk_hooks,
@ -42,6 +47,12 @@ from litellm.proxy.proxy_server import (
data_generator,
select_data_generator,
)
from litellm.types.llms.openai import (
ResponseCompletedEvent,
ResponseCreatedEvent,
ResponseFailedEvent,
ResponsesAPIResponse,
)
from litellm.types.utils import Delta, ModelResponseStream, StreamingChoices, Usage
from .conftest import normalize
@ -872,6 +883,145 @@ async def test_async_data_generator_mid_stream_exception_yields_error_payload(
assert any(isinstance(item, str) and item.startswith('data: {"error":') for item in out)
_UPSTREAM_BODY: Final = {
"code": "cyber_policy",
"message": "Upstream rejected request: flagged for possible cybersecurity risk",
"type": None,
}
@pytest.mark.asyncio
@pytest.mark.parametrize(
"terminal,upstream_error,expected_code",
[
("completed", None, None),
("serialization_failure", None, "server_error"),
("failure_after_completed", None, None),
pytest.param(
"upstream_failure",
litellm.AuthenticationError(
message="Upstream rejected request", llm_provider="openai", model="gpt-6-astra"
),
"authentication_error", id="authentication_error",
),
pytest.param(
"upstream_failure",
OpenAIAPIError(
message="Upstream rejected request",
request=httpx.Request("POST", "https://streaming.example/v1/responses"),
body={"code": {"reason": "overloaded"}, "type": {"unexpected": "object"}},
),
"server_error", id="structured_provider_error_fields",
),
pytest.param(
"upstream_failure",
litellm.InternalServerError(
message="Upstream rejected request", llm_provider="openai", model="gpt-6-astra", body=_UPSTREAM_BODY
),
"cyber_policy", id="upstream_body_code_and_message",
),
*(
pytest.param(
"upstream_failure", HTTPException(status_code=status, detail="Upstream rejected request"),
code, id=f"http_{status}",
)
for status, code in (
(400, "invalid_request_error"), (403, "permission_error"), (404, "not_found_error"),
(408, "request_timeout"), (422, "invalid_request_error"), (500, "server_error"), (503, "server_error"),
)
),
],
)
async def test_responses_stream_keeps_tool_deltas_and_only_emits_a_valid_terminal(
terminal: Literal["completed", "serialization_failure", "failure_after_completed", "upstream_failure"],
upstream_error: HTTPException | OpenAIAPIError | None,
expected_code: str | None,
) -> None:
class ToolDelta(BaseModel):
type: Literal["response.function_call_arguments.delta"]
sequence_number: int
item_id: str
output_index: int
delta: str
class UnserializableTerminal(BaseModel):
type: Literal["response.completed"]
sequence_number: int
response: ResponsesAPIResponse
invalid: object
response: Final = ResponsesAPIResponse(id="resp_visible", created_at=1, model="gpt-6-astra", output=[])
created: Final = ResponseCreatedEvent.model_validate(
{"type": "response.created", "sequence_number": 0, "response": response}
)
completed: Final = ResponseCompletedEvent.model_validate(
{"type": "response.completed", "sequence_number": 2, "response": response}
)
tool_delta: Final = ToolDelta(
type="response.function_call_arguments.delta", sequence_number=1, item_id="fc_stream_error",
output_index=0, delta='{"path":"partial',
)
original_status: Final = (
upstream_error.status_code if isinstance(upstream_error, (HTTPException, litellm.AuthenticationError)) else None
)
async def upstream() -> AsyncIterator[BaseModel]:
yield created
yield tool_delta
if upstream_error is not None:
raise upstream_error
yield (
UnserializableTerminal(type="response.completed", sequence_number=2, response=response, invalid=object())
if terminal == "serialization_failure" else completed
)
if terminal == "failure_after_completed":
raise litellm.APIError(
status_code=500, message="Stream close failed", llm_provider="openai", model="gpt-6-astra"
)
frames: Final = [
frame
async for frame in select_data_generator(
response=upstream(),
user_api_key_dict=_user_auth(),
request_data={},
responses_stream_errors=True,
)
]
decoded: Final = tuple(frame.decode() if isinstance(frame, bytes) else frame for frame in frames)
event_frames: Final = tuple(frame for frame in decoded if frame != "data: [DONE]\n\n")
payloads: Final = tuple(
json.loads(next(line[6:] for line in frame.splitlines() if line.startswith("data: ")))
for frame in event_frames
)
assert decoded[-1] == "data: [DONE]\n\n"
assert len(decoded) == len(event_frames) + 1
assert payloads[0]["response"]["id"] == "resp_visible"
assert payloads[1] == tool_delta.model_dump()
assert len(payloads) == 3
if terminal in ("serialization_failure", "upstream_failure"):
failure: Final = ResponseFailedEvent.model_validate(payloads[-1])
assert event_frames[-1].startswith("event: response.failed\n")
assert failure.response.id == "resp_visible"
assert failure.response.status == "failed"
assert failure.response.error is not None
assert failure.response.error["code"] == expected_code
if upstream_error is None:
assert "serialize" in failure.response.error["message"].lower()
else:
assert "Upstream rejected request" in failure.response.error["message"]
if isinstance(upstream_error, litellm.InternalServerError):
assert failure.response.error["message"] == _UPSTREAM_BODY["message"]
if isinstance(upstream_error, (HTTPException, litellm.AuthenticationError)):
assert upstream_error.status_code == original_status
assert payloads[-1]["sequence_number"] > payloads[1]["sequence_number"]
else:
assert payloads[-1]["type"] == "response.completed"
assert payloads[-1]["sequence_number"] == 2
assert "error" not in payloads[-1]
# ---------------------------------------------------------------------------
# select_data_generator
# ---------------------------------------------------------------------------

View file

@ -3,10 +3,12 @@ Test for response_api_endpoints/endpoints.py
"""
import unittest
from typing import Any
from typing import Any, Final, Literal
from unittest.mock import AsyncMock, MagicMock, patch
import httpx
import pytest
import respx
from fastapi.testclient import TestClient
from httpx import Response
@ -14,6 +16,111 @@ import litellm
from litellm.proxy.proxy_server import app
@pytest.mark.asyncio
@pytest.mark.parametrize(
"path,error_kind",
[
("/v1/responses", "rate_limit"),
("/v1/responses", "numeric_rate_limit"),
("/v1/responses", "server_error"),
("/v1/responses", "response_failed"),
("/v1/responses", "cyber_policy"),
("/cursor/chat/completions", "server_error"),
("/v1/chat/completions", "server_error"),
],
)
async def test_streaming_upstream_errors_keep_the_client_protocol(
monkeypatch: pytest.MonkeyPatch,
path: str,
error_kind: Literal["rate_limit", "numeric_rate_limit", "server_error", "response_failed", "cyber_policy"],
) -> None:
import litellm.proxy.proxy_server as ps
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
model: Final = "gpt-6-astra"
message: Final = "Upstream cannot complete this response"
code: Final = {
"rate_limit": "rate_limit_exceeded", "numeric_rate_limit": "429",
"server_error": "server_error", "response_failed": "server_error", "cyber_policy": "cyber_policy",
}[error_kind]
error: Final = {"message": message, "code": code, "type": None, "param": "input"}
response: Final = {"id": "resp_upstream", "object": "response", "created_at": 1,
"status": "in_progress", "model": model, "output": [],
"parallel_tool_calls": True, "tool_choice": "auto", "tools": []}
created: Final = {"type": "response.created", "sequence_number": 0, "response": response}
tool_added: Final = {"type": "response.output_item.added", "sequence_number": 1, "output_index": 0,
"item": {"type": "function_call", "id": "fc_partial", "call_id": "call_partial",
"name": "read_file", "arguments": "", "status": "in_progress"}}
tool_delta: Final = {"type": "response.function_call_arguments.delta", "sequence_number": 2,
"item_id": "fc_partial", "output_index": 0, "delta": '{"path":"partial'}
failed: Final = (
{"type": "response.failed", "sequence_number": 9,
"response": {**response, "status": "failed", "error": error}}
if error_kind in ("response_failed", "cyber_policy") else {"type": "error", "error": error}
)
chat: Final = {"id": "chatcmpl_partial", "object": "chat.completion.chunk", "created": 1,
"model": model, "choices": [{"index": 0, "delta": {"content": "partial"},
"finish_reason": None}]}
is_chat: Final = path == "/v1/chat/completions"
partial: Final = path != "/v1/responses" or error_kind in ("numeric_rate_limit", "response_failed", "cyber_policy")
response_events: Final = (created, tool_added, tool_delta, failed) if partial else (failed,)
upstream_events: Final = (chat, {"error": error}) if is_chat else response_events
wire: Final = "".join("data: " + json.dumps(event) + "\n\n" for event in upstream_events)
upstream_url: Final = "https://streaming.example/v1"
router: Final = litellm.Router(
model_list=[{"model_name": model, "litellm_params": {
"model": "openai/" + model, "api_base": upstream_url, "api_key": "fixture-key"}}],
num_retries=0,
)
monkeypatch.setattr(ps, "llm_router", router)
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
monkeypatch.setitem(app.dependency_overrides, user_api_key_auth, _auth_override)
with respx.mock as transport:
transport.post(upstream_url + ("/chat/completions" if is_chat else "/responses")).respond(
200, content=wire, headers={"Content-Type": "text/event-stream"}
)
async with httpx.AsyncClient(transport=httpx.ASGITransport(app), base_url="http://testserver") as client:
result: Final = await client.post(
path, json={
"model": model, "stream": True,
**({"messages": [{"role": "user", "content": "hello"}]} if is_chat else {"input": "hello"}),
},
)
frames: Final = tuple(frame for frame in result.text.split("\n\n") if "data: " in frame)
events: Final = tuple(
json.loads(next(line[6:] for line in frame.splitlines() if line.startswith("data: ")))
for frame in frames if "data: [DONE]" not in frame
)
assert result.status_code == 200, result.text
assert message in result.text
if path == "/v1/responses":
assert frames[-1] == "data: [DONE]", result.text
assert frames[-2].startswith("event: response.failed\n"), result.text
if partial:
assert [event["type"] for event in events] == [
"response.created", "response.output_item.added",
"response.function_call_arguments.delta", "response.failed",
]
assert events[2]["delta"] == tool_delta["delta"]
assert events[-1]["sequence_number"] == events[-2]["sequence_number"] + 1
assert events[-1]["response"]["id"] == events[0]["response"]["id"]
else:
assert [event["type"] for event in events] == ["response.failed"]
assert events[0]["sequence_number"] == 0
assert events[0]["response"]["id"].startswith("resp_")
assert events[-1]["response"]["status"] == "failed"
assert events[-1]["response"]["error"]["code"] == {
"rate_limit": "rate_limit_exceeded", "numeric_rate_limit": "rate_limit_exceeded",
"server_error": "server_error", "response_failed": "server_error", "cyber_policy": "cyber_policy",
}[error_kind]
assert events[-1]["response"]["error"]["message"] == message
else:
assert events[0]["object"] == "chat.completion.chunk", result.text
assert "response.failed" not in result.text
assert "error" in events[-1]
class TestResponsesAPIEndpoints(unittest.TestCase):
@pytest.mark.asyncio
@patch("litellm.proxy.proxy_server.llm_router")

View file

@ -1209,6 +1209,7 @@ async def test_api_key_preserved_through_failure_hook_to_database():
start_time,
end_time,
org_id,
project_id=None,
):
"""Mock update_database and capture the payload it creates"""
from litellm.proxy.spend_tracking.spend_tracking_utils import (

View file

@ -22,6 +22,7 @@ from litellm.proxy._types import (
LiteLLM_EndUserTable,
Litellm_EntityType,
LiteLLM_OrganizationTable,
LiteLLM_ProjectTableCachedObj,
LiteLLM_TagTable,
LiteLLM_TeamMembership,
LiteLLM_TeamTable,
@ -631,6 +632,154 @@ async def test_should_reserve_team_member_and_org_budget_counters(spend_counter_
await release_budget_reservation(reservation)
def _project_scoped_token() -> UserAPIKeyAuth:
return UserAPIKeyAuth(
token="key-project-scoped",
spend=0.0,
user_id="user-proj",
team_id="team-proj",
project_id="proj-1",
)
async def _seed_project_scoped_budgets(
key_cache: DualCache,
team_member_spend: float,
team_member_max_budget: float,
project_spend: float,
project_max_budget: float,
) -> None:
await key_cache.async_set_cache(
key="team_membership:user-proj:team-proj",
value=LiteLLM_TeamMembership(
user_id="user-proj",
team_id="team-proj",
spend=team_member_spend,
litellm_budget_table=LiteLLM_BudgetTable(max_budget=team_member_max_budget),
).model_dump(),
)
await key_cache.async_set_cache(
key="project_id:proj-1",
value=LiteLLM_ProjectTableCachedObj(
project_id="proj-1",
team_id="team-proj",
budget_id="project-budget-id",
spend=project_spend,
litellm_budget_table=LiteLLM_BudgetTable(max_budget=project_max_budget),
).model_dump(),
)
@pytest.mark.asyncio
async def test_should_reserve_project_and_team_member_counters_for_project_scoped_key(spend_counter_state):
counter_cache, key_cache = spend_counter_state
proxy_logging_obj = ProxyLogging(user_api_key_cache=key_cache)
await _seed_project_scoped_budgets(
key_cache,
team_member_spend=0.1,
team_member_max_budget=1.0,
project_spend=0.2,
project_max_budget=1.0,
)
estimated = estimate_request_max_cost(request_body=_request_body(), route="/chat/completions", llm_router=None)
assert estimated is not None and estimated > 0
reservation = await reserve_budget_for_request(
request_body=_request_body(),
route="/chat/completions",
llm_router=None,
valid_token=_project_scoped_token(),
team_object=LiteLLM_TeamTable(team_id="team-proj", spend=0.0, max_budget=None),
user_object=LiteLLM_UserTable(user_id="user-proj", spend=0.0),
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:team_member:user-proj:team-proj") == pytest.approx(
0.1 + estimated
)
assert counter_cache.in_memory_cache.get_cache(key="spend:project:proj-1") == pytest.approx(0.2 + estimated)
from litellm.proxy.proxy_server import increment_spend_counters
await increment_spend_counters(
token="key-project-scoped",
team_id="team-proj",
user_id="user-proj",
response_cost=0.05,
budget_reservation=reservation,
project_id="proj-1",
)
assert counter_cache.in_memory_cache.get_cache(key="spend:project:proj-1") == pytest.approx(0.25)
assert counter_cache.in_memory_cache.get_cache(key="spend:team_member:user-proj:team-proj") == pytest.approx(0.15)
@pytest.mark.asyncio
async def test_exhausted_team_member_budget_still_blocks_project_scoped_key(spend_counter_state):
counter_cache, key_cache = spend_counter_state
proxy_logging_obj = ProxyLogging(user_api_key_cache=key_cache)
await _seed_project_scoped_budgets(
key_cache,
team_member_spend=1.0,
team_member_max_budget=1.0,
project_spend=0.0,
project_max_budget=100.0,
)
with pytest.raises(litellm.BudgetExceededError) as exc_info:
await reserve_budget_for_request(
request_body=_request_body(),
route="/chat/completions",
llm_router=None,
valid_token=_project_scoped_token(),
team_object=LiteLLM_TeamTable(team_id="team-proj", spend=0.0, max_budget=None),
user_object=LiteLLM_UserTable(user_id="user-proj", spend=0.0),
prisma_client=None,
user_api_key_cache=key_cache,
proxy_logging_obj=proxy_logging_obj,
)
assert "TeamMember=user-proj:team-proj" in str(exc_info.value)
assert counter_cache.in_memory_cache.get_cache(key="spend:project:proj-1") in (None, pytest.approx(0.0))
@pytest.mark.asyncio
async def test_exhausted_project_budget_blocks_project_scoped_key(spend_counter_state):
counter_cache, key_cache = spend_counter_state
proxy_logging_obj = ProxyLogging(user_api_key_cache=key_cache)
await _seed_project_scoped_budgets(
key_cache,
team_member_spend=0.0,
team_member_max_budget=100.0,
project_spend=5.0,
project_max_budget=5.0,
)
with pytest.raises(litellm.BudgetExceededError) as exc_info:
await reserve_budget_for_request(
request_body=_request_body(),
route="/chat/completions",
llm_router=None,
valid_token=_project_scoped_token(),
team_object=LiteLLM_TeamTable(team_id="team-proj", spend=0.0, max_budget=None),
user_object=LiteLLM_UserTable(user_id="user-proj", spend=0.0),
prisma_client=None,
user_api_key_cache=key_cache,
proxy_logging_obj=proxy_logging_obj,
)
assert "Project=proj-1" in str(exc_info.value)
assert exc_info.value.entity_type == Litellm_EntityType.PROJECT.value
assert counter_cache.in_memory_cache.get_cache(key="spend:team_member:user-proj:team-proj") in (
None,
pytest.approx(0.0),
)
@pytest.mark.asyncio
async def test_should_not_reserve_user_budget_counter_for_team_key(spend_counter_state):
"""The reservation path mirrors the read path: no personal user counter for a team key.

View file

@ -4974,8 +4974,8 @@ async def test_add_router_settings_from_db_config_merge_logic():
mock_prisma_client.db.litellm_config.find_first = AsyncMock(return_value=mock_db_config)
# Call the method under test
proxy_config.router_settings.load_yaml(config_data["router_settings"])
await proxy_config._add_router_settings_from_db_config(
config_data=config_data,
llm_router=mock_router,
prisma_client=mock_prisma_client,
)
@ -5029,9 +5029,7 @@ async def test_invalid_db_routing_groups_do_not_abort_other_router_settings():
mock_prisma_client = MagicMock()
mock_prisma_client.db.litellm_config.find_first = AsyncMock(return_value=mock_db_config)
await ProxyConfig()._add_router_settings_from_db_config(
config_data={}, llm_router=router, prisma_client=mock_prisma_client
)
await ProxyConfig()._add_router_settings_from_db_config(llm_router=router, prisma_client=mock_prisma_client)
assert router.num_retries == 7
assert router._model_to_group == {"m1": "g1"}
@ -5053,9 +5051,7 @@ async def test_valid_db_routing_groups_still_replace_router_groups():
mock_prisma_client = MagicMock()
mock_prisma_client.db.litellm_config.find_first = AsyncMock(return_value=mock_db_config)
await ProxyConfig()._add_router_settings_from_db_config(
config_data={}, llm_router=router, prisma_client=mock_prisma_client
)
await ProxyConfig()._add_router_settings_from_db_config(llm_router=router, prisma_client=mock_prisma_client)
assert router.num_retries == 7
assert router._model_to_group == {"m2": "g2"}
@ -5098,8 +5094,8 @@ async def test_add_router_settings_from_db_config_empty_db_lists_do_not_clobber_
mock_prisma_client = MagicMock()
mock_prisma_client.db.litellm_config.find_first = AsyncMock(return_value=mock_db_config)
proxy_config.router_settings.load_yaml(config_data["router_settings"])
await proxy_config._add_router_settings_from_db_config(
config_data=config_data,
llm_router=mock_router,
prisma_client=mock_prisma_client,
)
@ -5135,8 +5131,8 @@ async def test_add_router_settings_from_db_config_empty_db_list_still_clears_unc
mock_prisma_client = MagicMock()
mock_prisma_client.db.litellm_config.find_first = AsyncMock(return_value=mock_db_config)
proxy_config.router_settings.load_yaml(config_data["router_settings"])
await proxy_config._add_router_settings_from_db_config(
config_data=config_data,
llm_router=mock_router,
prisma_client=mock_prisma_client,
)
@ -5160,8 +5156,8 @@ async def test_add_router_settings_from_db_config_edge_cases():
mock_router.update_settings = MagicMock()
# Test Case 1: No router provided
proxy_config.router_settings.load_yaml({"test": "value"})
await proxy_config._add_router_settings_from_db_config(
config_data={"router_settings": {"test": "value"}},
llm_router=None,
prisma_client=MagicMock(),
)
@ -5169,8 +5165,8 @@ async def test_add_router_settings_from_db_config_edge_cases():
mock_router.update_settings.assert_not_called()
# Test Case 2: No prisma client provided
proxy_config.router_settings.load_yaml({"test": "value"})
await proxy_config._add_router_settings_from_db_config(
config_data={"router_settings": {"test": "value"}},
llm_router=mock_router,
prisma_client=None,
)
@ -5183,8 +5179,8 @@ async def test_add_router_settings_from_db_config_edge_cases():
config_data = {"router_settings": {"routing_strategy": "usage-based"}}
proxy_config.router_settings.load_yaml(config_data["router_settings"])
await proxy_config._add_router_settings_from_db_config(
config_data=config_data,
llm_router=mock_router,
prisma_client=mock_prisma_client,
)
@ -5198,8 +5194,8 @@ async def test_add_router_settings_from_db_config_edge_cases():
mock_db_config.param_value = {"db_setting": "db_value"}
mock_prisma_client.db.litellm_config.find_first = AsyncMock(return_value=mock_db_config)
proxy_config.router_settings.load_yaml({})
await proxy_config._add_router_settings_from_db_config(
config_data={}, # No router_settings in config
llm_router=mock_router,
prisma_client=mock_prisma_client,
)
@ -5211,9 +5207,8 @@ async def test_add_router_settings_from_db_config_edge_cases():
# Test Case 5: Both config and DB router_settings are None/empty
mock_prisma_client.db.litellm_config.find_first = AsyncMock(return_value=None)
await proxy_config._add_router_settings_from_db_config(
config_data={}, llm_router=mock_router, prisma_client=mock_prisma_client
)
proxy_config.router_settings.load_yaml({})
await proxy_config._add_router_settings_from_db_config(llm_router=mock_router, prisma_client=mock_prisma_client)
# Should not call update_settings when no settings exist
mock_router.update_settings.assert_not_called()
@ -5225,8 +5220,8 @@ async def test_add_router_settings_from_db_config_edge_cases():
config_data = {"router_settings": {"config_setting": "config_value"}}
proxy_config.router_settings.load_yaml(config_data["router_settings"])
await proxy_config._add_router_settings_from_db_config(
config_data=config_data,
llm_router=mock_router,
prisma_client=mock_prisma_client,
)
@ -5275,8 +5270,8 @@ async def test_add_router_settings_shallow_merge_behavior():
mock_prisma_client = MagicMock()
mock_prisma_client.db.litellm_config.find_first = AsyncMock(return_value=mock_db_config)
proxy_config.router_settings.load_yaml(config_data["router_settings"])
await proxy_config._add_router_settings_from_db_config(
config_data=config_data,
llm_router=mock_router,
prisma_client=mock_prisma_client,
)
@ -5298,6 +5293,36 @@ async def test_add_router_settings_shallow_merge_behavior():
assert merged_settings["top_level"] == "config_top"
@pytest.mark.asyncio
async def test_router_settings_reload_keeps_db_values_writable(tmp_path, monkeypatch):
from litellm.proxy.proxy_server import ProxyConfig
config_path: Final = tmp_path / "config.yaml"
config_path.write_text(yaml.safe_dump({"model_list": [], "router_settings": {"disable_cooldowns": True}}))
db_row: Final = types.SimpleNamespace(param_value={"num_retries": 0})
async def read_config_row(_prisma_client, param_name):
return db_row if param_name == "router_settings" else None
mock_prisma_client: Final = MagicMock()
mock_prisma_client.db.litellm_config.find_first = AsyncMock(return_value=db_row)
mock_router: Final = MagicMock()
monkeypatch.setattr(proxy_server_module, "get_config_param", read_config_row)
monkeypatch.setattr(proxy_server_module, "prisma_client", mock_prisma_client)
monkeypatch.setattr(proxy_server_module, "store_model_in_db", True)
monkeypatch.setattr(proxy_server_module, "user_config_file_path", None)
proxy_config: Final = ProxyConfig()
for _ in range(2):
await proxy_config.get_config(config_file_path=str(config_path))
await proxy_config._add_router_settings_from_db_config(llm_router=mock_router, prisma_client=mock_prisma_client)
assert mock_router.update_settings.call_args.kwargs == {"disable_cooldowns": True, "num_retries": 0}
assert proxy_config.router_settings.source("num_retries") == "db"
assert proxy_config.router_settings.rejected_writes({"num_retries": 3}) == ()
assert proxy_config.router_settings.rejected_writes({"disable_cooldowns": False}) == ("disable_cooldowns",)
@pytest.mark.asyncio
async def test_model_info_v1_oci_secrets_not_leaked():
"""
@ -7585,7 +7610,15 @@ async def test_update_general_settings_db_pass_through_endpoint_cannot_override_
assert still_open.api_key is None
@pytest.fixture
def app_routes_restored():
routes_before: Final = tuple(app.router.routes)
yield
app.router.routes[:] = routes_before
@pytest.mark.asyncio
@pytest.mark.usefixtures("app_routes_restored")
async def test_deleting_the_stored_pass_through_row_takes_the_route_out_of_service():
"""A pass-through route the database declared has to stop serving when that row is
deleted. The proxy's own registry of live pass-through routes is what decides whether
@ -7624,6 +7657,7 @@ async def test_deleting_the_stored_pass_through_row_takes_the_route_out_of_servi
@pytest.mark.asyncio
@pytest.mark.usefixtures("app_routes_restored")
async def test_a_stored_pass_through_row_never_disturbs_the_config_declared_routes():
"""``pass_through_endpoints`` is config-owned once the file declares it, so writing and then
deleting a stored row resolves to the same list both times and the config file's routes keep
@ -14028,32 +14062,29 @@ async def test_authoritative_floor_spend_keeps_a_reset_marker_written_during_the
@pytest.mark.asyncio
async def test_login_throttle_settings_are_not_hot_applied_from_the_database():
"""LIT-5285: a stored sign-in limit does not take effect on a live worker.
_update_general_settings copies an allowlist of keys out of the DB row on every config
poll. Adding these to it would let a stored value outrank config.yaml without a restart,
so an operator locked out by a bad value could not fix it by editing YAML and restarting.
"""
async def test_login_throttle_limits_from_the_config_file_outrank_the_database(monkeypatch):
import litellm.proxy.proxy_server as ps
from litellm.proxy.proxy_server import ProxyConfig
original = dict(ps.general_settings)
try:
ps.general_settings.clear()
await ProxyConfig()._update_general_settings(
db_general_settings={
"max_failed_login_attempts_per_source": 999,
"failed_login_window_seconds": 1,
"failed_login_block_seconds": 1,
}
)
assert "max_failed_login_attempts_per_source" not in ps.general_settings
assert "failed_login_window_seconds" not in ps.general_settings
assert "failed_login_block_seconds" not in ps.general_settings
finally:
ps.general_settings.clear()
ps.general_settings.update(original)
monkeypatch.setattr(
ps,
"general_settings",
{
"max_failed_login_attempts_per_source": 10,
"failed_login_window_seconds": 60,
"failed_login_block_seconds": 300,
},
)
await ProxyConfig()._update_general_settings(
db_general_settings={
"max_failed_login_attempts_per_source": 999,
"failed_login_window_seconds": 1,
"failed_login_block_seconds": 1,
}
)
assert ps.general_settings.get("max_failed_login_attempts_per_source") == 10
assert ps.general_settings.get("failed_login_window_seconds") == 60
assert ps.general_settings.get("failed_login_block_seconds") == 300
@pytest.mark.asyncio
@ -14283,6 +14314,29 @@ async def test_update_general_settings_keeps_yaml_openai_websocket_passthrough()
assert ps.general_settings["enable_openai_websocket_passthrough"] is False
def test_settings_store_exposes_dashboard_saved_mcp_client_allowlist_to_the_mcp_gateway() -> None:
from litellm.proxy._experimental.mcp_server.client_allowlist import MCPClientAllowlist, load_mcp_client_allowlist
from litellm.proxy.proxy_server import ProxyConfig
settings: Final = ProxyConfig().settings
settings.load_yaml({"litellm_jwtauth": {"mcp_client_id_jwt_field": "azp"}})
assert load_mcp_client_allowlist(settings) is None
settings.apply_db_row(
"general_settings",
{
"mcp_allowed_clients": [{"alias": "Antigravity CLI", "value": "antigravity-cli"}],
"mcp_client_id_header": "X-MCP-Client",
},
)
assert load_mcp_client_allowlist(settings) == MCPClientAllowlist(
aliases_by_value={"antigravity-cli": "Antigravity CLI"}, jwt_field="azp", header="x-mcp-client"
)
settings.apply_db_row("general_settings", {"mcp_client_id_header": "X-MCP-Client"})
assert load_mcp_client_allowlist(settings) is None
async def test_token_counter_keeps_the_event_loop_free_during_a_huggingface_count(monkeypatch):
from tests.large_text import text
from tests.test_litellm.litellm_core_utils.event_loop_lag import (

View file

@ -35,6 +35,7 @@ class FakeBatch:
self.litellm_organizationtable = FakeBatchTable("litellm_organizationtable", self.calls)
self.litellm_tagtable = FakeBatchTable("litellm_tagtable", self.calls)
self.litellm_modelaccessgroupbudgettable = FakeBatchTable("litellm_modelaccessgroupbudgettable", self.calls)
self.litellm_projecttable = FakeBatchTable("litellm_projecttable", self.calls)
self.litellm_endusertable = FakeBatchTable("litellm_endusertable", self.calls)
async def commit(self) -> None:
@ -94,6 +95,7 @@ async def test_budget_cascade_dependents_and_window_advance_share_one_batch():
uow.organizations.queue_spend_zero(where=linked)
uow.tags.queue_spend_zero(where=linked)
uow.model_access_groups.queue_spend_zero(where=linked)
uow.projects.queue_spend_zero(where=linked)
uow.endusers.queue_spend_zero(where={"user_id": {"in": ["enduser-1"]}})
uow.budgets.queue_window_advance(budget_id="budget-1", budget_reset_at=reset_at)
assert batch.commit_count == 0
@ -105,6 +107,7 @@ async def test_budget_cascade_dependents_and_window_advance_share_one_batch():
("litellm_organizationtable.update_many", linked, {"spend": 0}),
("litellm_tagtable.update_many", linked, {"spend": 0}),
("litellm_modelaccessgroupbudgettable.update_many", linked, {"spend": 0}),
("litellm_projecttable.update_many", linked, {"spend": 0}),
("litellm_endusertable.update_many", {"user_id": {"in": ["enduser-1"]}}, {"spend": 0}),
("litellm_budgettable.update_many", {"budget_id": "budget-1"}, {"budget_reset_at": reset_at}),
]

View file

@ -4938,3 +4938,65 @@ class TestStreamingSnapshotItemIds:
reasoning_items = _bridged_output_items(completed_event.response, "reasoning")
assert len(reasoning_items) == 1
assert reasoning_items[0].id == streamed_event.item_id
def test_transform_chat_completion_response_incomplete_details():
from litellm.types.llms.openai import IncompleteDetails
resp_length = ModelResponse(
id="resp-length",
choices=[Choices(index=0, finish_reason="length", message=Message(content="cutoff", role="assistant"))],
model="gpt-4o",
)
result_length = LiteLLMCompletionResponsesConfig.transform_chat_completion_response_to_responses_api_response(
request_input="test prompt",
responses_api_request={},
chat_completion_response=resp_length,
)
assert result_length.status == "incomplete"
assert result_length.incomplete_details is not None
assert result_length.incomplete_details.reason == "max_output_tokens"
resp_filter = ModelResponse(
id="resp-filter",
choices=[Choices(index=0, finish_reason="content_filter", message=Message(content=None, role="assistant"))],
model="gpt-4o",
)
result_filter = LiteLLMCompletionResponsesConfig.transform_chat_completion_response_to_responses_api_response(
request_input="test prompt",
responses_api_request={},
chat_completion_response=resp_filter,
)
assert result_filter.status == "incomplete"
assert result_filter.incomplete_details is not None
assert result_filter.incomplete_details.reason == "content_filter"
resp_refusal = ModelResponse(
id="resp-refusal",
choices=[Choices(index=0, finish_reason="refusal", message=Message(content=None, role="assistant"))],
model="gpt-4o",
)
result_refusal = LiteLLMCompletionResponsesConfig.transform_chat_completion_response_to_responses_api_response(
request_input="test prompt",
responses_api_request={},
chat_completion_response=resp_refusal,
)
assert result_refusal.status == "incomplete"
assert result_refusal.incomplete_details is not None
assert result_refusal.incomplete_details.reason == "content_filter"
existing_details = IncompleteDetails(reason="content_filter")
resp_existing = ModelResponse(
id="resp-existing",
choices=[Choices(index=0, finish_reason="length", message=Message(content="cutoff", role="assistant"))],
model="gpt-4o",
)
resp_existing.incomplete_details = existing_details
result_existing = LiteLLMCompletionResponsesConfig.transform_chat_completion_response_to_responses_api_response(
request_input="test prompt",
responses_api_request={},
chat_completion_response=resp_existing,
)
assert result_existing.status == "incomplete"
assert result_existing.incomplete_details == existing_details

View file

@ -340,6 +340,36 @@ def test_maybe_raise_for_response_failed_event_with_dict_error():
assert exc_info.value.status_code == 429
@pytest.mark.parametrize("code", [429, "429"])
def test_response_failed_numeric_code_maps_to_its_http_status(code: int | str):
iterator = _make_iterator()
mock_response_obj = Mock()
mock_response_obj.error = {"code": code, "message": "throttled"}
chunk = Mock()
chunk.type = "response.failed"
chunk.response = mock_response_obj
with pytest.raises(MidStreamFallbackError) as exc_info:
iterator._maybe_raise_for_error_event(chunk)
assert exc_info.value.status_code == 429
assert isinstance(exc_info.value.original_exception, litellm.RateLimitError)
def test_response_failed_unknown_code_keeps_upstream_code_and_message_on_mapped_exception():
iterator = _make_iterator()
upstream_message = "This content was flagged for possible cybersecurity risk."
mock_response_obj = Mock()
mock_response_obj.error = {"code": "cyber_policy", "message": upstream_message}
chunk = Mock()
chunk.type = "response.failed"
chunk.response = mock_response_obj
with pytest.raises(MidStreamFallbackError) as exc_info:
iterator._maybe_raise_for_error_event(chunk)
mapped = exc_info.value.original_exception
assert isinstance(mapped, litellm.InternalServerError)
assert mapped.code == "cyber_policy"
assert mapped.body == {"message": upstream_message, "type": None, "code": "cyber_policy"}
def test_maybe_raise_for_error_event_null_error_obj():
"""error chunk with no error field: message and code default; wrapped as 500."""
iterator = _make_iterator()
@ -523,6 +553,9 @@ def test_every_openai_sdk_response_error_code_has_explicit_status_mapping():
("failed_to_download_image", 400),
("image_file_not_found", 400),
("totally_unknown_future_code", 500),
("429", 429),
("503", 503),
("200", 500),
],
)
def test_status_code_for_documented_response_error_codes(code: str, expected_status: int):

View file

@ -360,7 +360,7 @@ async def test_config_update_persists_and_reads_back_retry_policy(monkeypatch):
async def _apply_router_settings(*args, **kwargs):
await proxy_server.proxy_config._add_router_settings_from_db_config(
config_data={}, llm_router=router, prisma_client=prisma_client
llm_router=router, prisma_client=prisma_client
)
monkeypatch.setattr(proxy_server, "prisma_client", prisma_client)

View file

@ -1,7 +1,8 @@
import { render, screen, waitFor } from "@testing-library/react";
import { fireEvent, render, screen, waitFor, within } from "@testing-library/react";
import userEvent from "@testing-library/user-event";
import { describe, it, expect, vi, beforeEach } from "vitest";
import MCPNetworkSettings from "./MCPNetworkSettings";
import { toast } from "@/lib/toast";
import {
getGeneralSettingsCall,
updateConfigFieldSetting,
@ -16,8 +17,31 @@ vi.mock("@/components/networking", () => ({
fetchMCPClientIp: vi.fn(),
}));
vi.mock("@/lib/toast", () => ({
toast: { success: vi.fn(), fromError: vi.fn() },
}));
const renderSettings = () => render(<MCPNetworkSettings accessToken="tok" />);
const ANTIGRAVITY = { alias: "Antigravity CLI", value: "antigravity-cli" };
const CODEX = { alias: "Codex", value: "codex-mcp-client" };
const clientCard = (alias: string) => screen.getByRole("button", { name: new RegExp(`^${alias}`) });
const fillClientDialog = async (alias: string, value: string) => {
const dialog = await screen.findByRole("dialog");
fireEvent.change(within(dialog).getByRole("textbox", { name: "Alias" }), { target: { value: alias } });
fireEvent.change(within(dialog).getByRole("textbox", { name: "Value" }), { target: { value } });
return dialog;
};
const addClient = async (alias: string, value: string) => {
await userEvent.click(screen.getByRole("button", { name: "Add client" }));
const dialog = await fillClientDialog(alias, value);
await userEvent.click(within(dialog).getByRole("button", { name: "Add" }));
await waitFor(() => expect(screen.queryByRole("dialog")).not.toBeInTheDocument());
};
describe("MCPNetworkSettings", () => {
beforeEach(() => {
vi.clearAllMocks();
@ -82,25 +106,399 @@ describe("MCPNetworkSettings", () => {
expect(screen.getByText("203.0.113.0/24")).toBeInTheDocument();
});
it("saves the configured ranges", async () => {
it("saves the configured ranges once they change", async () => {
vi.mocked(fetchMCPClientIp).mockResolvedValue("203.0.113.45");
vi.mocked(getGeneralSettingsCall).mockResolvedValue([
{ field_name: "mcp_internal_ip_ranges", field_value: ["10.0.0.0/8"] },
]);
renderSettings();
await userEvent.click(await screen.findByText("203.0.113.0/24"));
await userEvent.click(await screen.findByRole("button", { name: /Save/ }));
await waitFor(() =>
expect(updateConfigFieldSetting).toHaveBeenCalledWith("tok", "mcp_internal_ip_ranges", ["10.0.0.0/8"]),
expect(updateConfigFieldSetting).toHaveBeenCalledWith("tok", "mcp_internal_ip_ranges", [
"10.0.0.0/8",
"203.0.113.0/24",
]),
);
expect(deleteConfigFieldSetting).not.toHaveBeenCalled();
expect(deleteConfigFieldSetting).not.toHaveBeenCalledWith("tok", "mcp_internal_ip_ranges");
});
it("clears the setting instead of saving an empty list", async () => {
it("clears a stored range setting instead of saving an empty list", async () => {
vi.mocked(getGeneralSettingsCall).mockResolvedValue([
{ field_name: "mcp_internal_ip_ranges", field_value: ["10.0.0.0/8"] },
]);
renderSettings();
await userEvent.click(await screen.findByRole("button", { name: /Save/ }));
await userEvent.click(await screen.findByRole("button", { name: "Remove 10.0.0.0/8" }));
await userEvent.click(screen.getByRole("button", { name: /Save/ }));
await waitFor(() => expect(deleteConfigFieldSetting).toHaveBeenCalledWith("tok", "mcp_internal_ip_ranges"));
expect(updateConfigFieldSetting).not.toHaveBeenCalled();
});
it("does not write settings that were never stored and are still empty", async () => {
renderSettings();
await userEvent.click(await screen.findByRole("button", { name: /Save/ }));
await waitFor(() => expect(toast.success).toHaveBeenCalledWith("MCP network settings saved"));
expect(deleteConfigFieldSetting).not.toHaveBeenCalled();
expect(updateConfigFieldSetting).not.toHaveBeenCalled();
});
it("labels the section Allowed Clients and renders each stored client as a card showing alias and value", async () => {
vi.mocked(getGeneralSettingsCall).mockResolvedValue([
{ field_name: "mcp_allowed_clients", field_value: [ANTIGRAVITY, CODEX] },
]);
renderSettings();
expect(await screen.findByText("Allowed Clients")).toBeVisible();
expect(screen.queryByText(/Allowed Client IDs/)).not.toBeInTheDocument();
expect(screen.queryByText(/Allowed Client Applications/)).not.toBeInTheDocument();
expect(clientCard("Antigravity CLI")).toHaveTextContent("antigravity-cli");
expect(clientCard("Codex")).toHaveTextContent("codex-mcp-client");
expect(screen.queryByRole("dialog")).not.toBeInTheDocument();
});
it("opens an edit dialog when a client card is clicked, prefilled with that client's alias and value", async () => {
vi.mocked(getGeneralSettingsCall).mockResolvedValue([
{ field_name: "mcp_allowed_clients", field_value: [ANTIGRAVITY, CODEX] },
]);
renderSettings();
await screen.findByText("Allowed Clients");
await userEvent.click(clientCard("Codex"));
const dialog = await screen.findByRole("dialog", { name: "Edit client" });
expect(within(dialog).getByRole("textbox", { name: "Alias" })).toHaveValue("Codex");
expect(within(dialog).getByRole("textbox", { name: "Value" })).toHaveValue("codex-mcp-client");
});
it("warns that a stored allowlist in the old plain-string shape denies every client and lets Save remove it", async () => {
vi.mocked(getGeneralSettingsCall).mockResolvedValue([
{ field_name: "mcp_allowed_clients", field_value: ["antigravity-cli"] },
]);
renderSettings();
expect(await screen.findByText(/stored allowlist is not a list of alias and value pairs/)).toBeVisible();
expect(screen.queryByText("antigravity-cli")).not.toBeInTheDocument();
await userEvent.click(screen.getByRole("button", { name: /Save/ }));
await waitFor(() => expect(deleteConfigFieldSetting).toHaveBeenCalledWith("tok", "mcp_allowed_clients"));
expect(updateConfigFieldSetting).not.toHaveBeenCalled();
await waitFor(() => expect(screen.queryByText(/stored allowlist is not a list/)).not.toBeInTheDocument());
});
it("replaces a stored allowlist in the old plain-string shape with the clients the admin adds", async () => {
vi.mocked(getGeneralSettingsCall).mockResolvedValue([
{ field_name: "mcp_allowed_clients", field_value: ["antigravity-cli"] },
]);
renderSettings();
await screen.findByText(/stored allowlist is not a list of alias and value pairs/);
await addClient(ANTIGRAVITY.alias, ANTIGRAVITY.value);
await userEvent.click(screen.getByRole("button", { name: /Save/ }));
await waitFor(() =>
expect(updateConfigFieldSetting).toHaveBeenCalledWith("tok", "mcp_allowed_clients", [ANTIGRAVITY]),
);
expect(deleteConfigFieldSetting).not.toHaveBeenCalledWith("tok", "mcp_allowed_clients");
await waitFor(() => expect(screen.queryByText(/stored allowlist is not a list/)).not.toBeInTheDocument());
});
it("treats a stored entry with an empty alias or value as denying every client, like the gateway does", async () => {
vi.mocked(getGeneralSettingsCall).mockResolvedValue([
{ field_name: "mcp_allowed_clients", field_value: [ANTIGRAVITY, { alias: "", value: "claude-code" }] },
]);
renderSettings();
expect(await screen.findByText(/stored allowlist is not a list of alias and value pairs/)).toBeVisible();
expect(screen.queryByRole("button", { name: /^Antigravity CLI/ })).not.toBeInTheDocument();
});
it("adds clients as alias and value pairs and saves them under mcp_allowed_clients", async () => {
renderSettings();
await screen.findByText("Allowed Clients");
await addClient(" Antigravity CLI ", " antigravity-cli ");
await addClient("Codex", "codex-mcp-client");
await userEvent.click(screen.getByRole("button", { name: /Save/ }));
await waitFor(() =>
expect(updateConfigFieldSetting).toHaveBeenCalledWith("tok", "mcp_allowed_clients", [ANTIGRAVITY, CODEX]),
);
expect(deleteConfigFieldSetting).not.toHaveBeenCalledWith("tok", "mcp_allowed_clients");
});
it("edits a stored client's value through its dialog and saves the new value", async () => {
vi.mocked(getGeneralSettingsCall).mockResolvedValue([
{ field_name: "mcp_allowed_clients", field_value: [ANTIGRAVITY] },
]);
renderSettings();
await screen.findByText("Allowed Clients");
await userEvent.click(clientCard("Antigravity CLI"));
const dialog = await screen.findByRole("dialog");
fireEvent.change(within(dialog).getByRole("textbox", { name: "Value" }), {
target: { value: " 0oa1b2c3d4e5f6g7h8i9 " },
});
await userEvent.click(within(dialog).getByRole("button", { name: "Done" }));
await waitFor(() => expect(screen.queryByRole("dialog")).not.toBeInTheDocument());
expect(clientCard("Antigravity CLI")).toHaveTextContent("0oa1b2c3d4e5f6g7h8i9");
await userEvent.click(screen.getByRole("button", { name: /Save/ }));
await waitFor(() =>
expect(updateConfigFieldSetting).toHaveBeenCalledWith("tok", "mcp_allowed_clients", [
{ alias: "Antigravity CLI", value: "0oa1b2c3d4e5f6g7h8i9" },
]),
);
});
it("keeps a stored client untouched when its dialog is cancelled", async () => {
vi.mocked(getGeneralSettingsCall).mockResolvedValue([
{ field_name: "mcp_allowed_clients", field_value: [ANTIGRAVITY] },
]);
renderSettings();
await screen.findByText("Allowed Clients");
await userEvent.click(clientCard("Antigravity CLI"));
const dialog = await screen.findByRole("dialog");
fireEvent.change(within(dialog).getByRole("textbox", { name: "Value" }), { target: { value: "changed" } });
await userEvent.click(within(dialog).getByRole("button", { name: "Cancel" }));
await waitFor(() => expect(screen.queryByRole("dialog")).not.toBeInTheDocument());
expect(clientCard("Antigravity CLI")).toHaveTextContent("antigravity-cli");
await userEvent.click(screen.getByRole("button", { name: /Save/ }));
await waitFor(() => expect(toast.success).toHaveBeenCalledWith("MCP network settings saved"));
expect(updateConfigFieldSetting).not.toHaveBeenCalled();
});
it("will not add a client that has an alias but no value", async () => {
renderSettings();
await screen.findByText("Allowed Clients");
await userEvent.click(screen.getByRole("button", { name: "Add client" }));
const dialog = await fillClientDialog("Antigravity CLI", " ");
expect(within(dialog).getByRole("button", { name: "Add" })).toBeDisabled();
});
it("adds nothing when the add dialog is cancelled", async () => {
vi.mocked(getGeneralSettingsCall).mockResolvedValue([
{ field_name: "mcp_allowed_clients", field_value: [ANTIGRAVITY] },
]);
renderSettings();
await screen.findByText("Allowed Clients");
await userEvent.click(screen.getByRole("button", { name: "Add client" }));
const dialog = await fillClientDialog("Codex", "codex-mcp-client");
await userEvent.click(within(dialog).getByRole("button", { name: "Cancel" }));
await waitFor(() => expect(screen.queryByRole("dialog")).not.toBeInTheDocument());
expect(screen.queryByText("Codex")).not.toBeInTheDocument();
await userEvent.click(screen.getByRole("button", { name: /Save/ }));
await waitFor(() => expect(toast.success).toHaveBeenCalledWith("MCP network settings saved"));
expect(updateConfigFieldSetting).not.toHaveBeenCalled();
expect(deleteConfigFieldSetting).not.toHaveBeenCalled();
});
it("removes the right client from the middle of the list", async () => {
vi.mocked(getGeneralSettingsCall).mockResolvedValue([
{
field_name: "mcp_allowed_clients",
field_value: [ANTIGRAVITY, { alias: "Claude Code", value: "claude-code" }, CODEX],
},
]);
renderSettings();
await screen.findByText("Allowed Clients");
await userEvent.click(clientCard("Claude Code"));
await userEvent.click(within(await screen.findByRole("dialog")).getByRole("button", { name: "Remove client" }));
await waitFor(() => expect(screen.queryByRole("dialog")).not.toBeInTheDocument());
expect(screen.queryByText("claude-code")).not.toBeInTheDocument();
await userEvent.click(screen.getByRole("button", { name: /Save/ }));
await waitFor(() =>
expect(updateConfigFieldSetting).toHaveBeenCalledWith("tok", "mcp_allowed_clients", [ANTIGRAVITY, CODEX]),
);
});
it("removes a client and clears the setting when the list becomes empty", async () => {
vi.mocked(getGeneralSettingsCall).mockResolvedValue([
{ field_name: "mcp_allowed_clients", field_value: [{ alias: "Claude Code", value: "claude-code" }] },
]);
renderSettings();
await screen.findByText("Allowed Clients");
await userEvent.click(clientCard("Claude Code"));
await userEvent.click(within(await screen.findByRole("dialog")).getByRole("button", { name: "Remove client" }));
await waitFor(() => expect(screen.queryByRole("dialog")).not.toBeInTheDocument());
expect(screen.queryByText("claude-code")).not.toBeInTheDocument();
await userEvent.click(screen.getByRole("button", { name: /Save/ }));
await waitFor(() => expect(deleteConfigFieldSetting).toHaveBeenCalledWith("tok", "mcp_allowed_clients"));
expect(updateConfigFieldSetting).not.toHaveBeenCalledWith("tok", "mcp_allowed_clients", expect.anything());
});
it("warns that a stored empty allowlist denies every client and lets Save remove it", async () => {
vi.mocked(getGeneralSettingsCall).mockResolvedValue([{ field_name: "mcp_allowed_clients", field_value: [] }]);
renderSettings();
expect(await screen.findByText(/An empty allowlist is currently stored, so every client is denied/)).toBeVisible();
await userEvent.click(screen.getByRole("button", { name: /Save/ }));
await waitFor(() => expect(deleteConfigFieldSetting).toHaveBeenCalledWith("tok", "mcp_allowed_clients"));
expect(updateConfigFieldSetting).not.toHaveBeenCalled();
await waitFor(() => expect(screen.queryByText(/An empty allowlist is currently stored/)).not.toBeInTheDocument());
});
it("does not show the deny-all warning when no allowlist is stored", async () => {
vi.mocked(getGeneralSettingsCall).mockResolvedValue([{ field_name: "mcp_allowed_clients", field_value: null }]);
renderSettings();
await screen.findByText("Allowed Clients");
expect(screen.queryByText(/every client is denied/)).not.toBeInTheDocument();
});
it("explains that JWT callers are identified by the configured claim and others by the opt-in header", async () => {
renderSettings();
expect(await screen.findByText(/litellm_jwtauth\.mcp_client_id_jwt_field/)).toBeVisible();
expect(screen.getByText(/Clients pick this value themselves, so it is a policy control/)).toBeVisible();
expect(screen.queryByText(/clientInfo/)).not.toBeInTheDocument();
});
it("renders the stored client identity header once settings load", async () => {
vi.mocked(getGeneralSettingsCall).mockResolvedValue([
{ field_name: "mcp_client_id_header", field_value: "x-mcp-client" },
]);
renderSettings();
expect(await screen.findByRole("textbox", { name: "Client identity header" })).toHaveValue("x-mcp-client");
});
it("saves a newly typed client identity header under mcp_client_id_header", async () => {
renderSettings();
fireEvent.change(await screen.findByRole("textbox", { name: "Client identity header" }), {
target: { value: " x-mcp-client " },
});
await userEvent.click(screen.getByRole("button", { name: /Save/ }));
await waitFor(() =>
expect(updateConfigFieldSetting).toHaveBeenCalledWith("tok", "mcp_client_id_header", "x-mcp-client"),
);
expect(updateConfigFieldSetting).not.toHaveBeenCalledWith("tok", "mcp_allowed_clients", expect.anything());
expect(deleteConfigFieldSetting).not.toHaveBeenCalled();
});
it("clears a stored client identity header when the field is emptied, so only JWT identity is trusted", async () => {
vi.mocked(getGeneralSettingsCall).mockResolvedValue([
{ field_name: "mcp_client_id_header", field_value: "x-mcp-client" },
]);
renderSettings();
fireEvent.change(await screen.findByRole("textbox", { name: "Client identity header" }), {
target: { value: "" },
});
await userEvent.click(screen.getByRole("button", { name: /Save/ }));
await waitFor(() => expect(deleteConfigFieldSetting).toHaveBeenCalledWith("tok", "mcp_client_id_header"));
expect(updateConfigFieldSetting).not.toHaveBeenCalled();
});
it("does not rewrite an unchanged client identity header on save", async () => {
vi.mocked(getGeneralSettingsCall).mockResolvedValue([
{ field_name: "mcp_client_id_header", field_value: "x-mcp-client" },
]);
renderSettings();
await screen.findByRole("textbox", { name: "Client identity header" });
await userEvent.click(screen.getByRole("button", { name: /Save/ }));
await waitFor(() => expect(toast.success).toHaveBeenCalled());
expect(updateConfigFieldSetting).not.toHaveBeenCalled();
expect(deleteConfigFieldSetting).not.toHaveBeenCalled();
});
it("keeps the private ranges and the allowed clients as independent settings on save", async () => {
vi.mocked(getGeneralSettingsCall).mockResolvedValue([
{ field_name: "mcp_internal_ip_ranges", field_value: ["10.0.0.0/8"] },
{ field_name: "mcp_allowed_clients", field_value: [ANTIGRAVITY] },
]);
renderSettings();
await screen.findByText("Allowed Clients");
await addClient("Codex", "codex-mcp-client");
await userEvent.click(screen.getByRole("button", { name: /Save/ }));
await waitFor(() =>
expect(updateConfigFieldSetting).toHaveBeenCalledWith("tok", "mcp_allowed_clients", [ANTIGRAVITY, CODEX]),
);
expect(updateConfigFieldSetting).not.toHaveBeenCalledWith("tok", "mcp_internal_ip_ranges", expect.anything());
expect(deleteConfigFieldSetting).not.toHaveBeenCalled();
});
it("still saves the allowed clients when the private range write fails, and reports the failure", async () => {
vi.mocked(getGeneralSettingsCall).mockResolvedValue([
{ field_name: "mcp_internal_ip_ranges", field_value: ["10.0.0.0/8"] },
]);
const rangeFailure = new Error("Field name=mcp_internal_ip_ranges not in config");
vi.mocked(deleteConfigFieldSetting).mockRejectedValue(rangeFailure);
renderSettings();
await userEvent.click(await screen.findByRole("button", { name: "Remove 10.0.0.0/8" }));
await addClient("Codex", "codex-mcp-client");
await userEvent.click(screen.getByRole("button", { name: /Save/ }));
await waitFor(() => expect(updateConfigFieldSetting).toHaveBeenCalledWith("tok", "mcp_allowed_clients", [CODEX]));
await waitFor(() => expect(toast.fromError).toHaveBeenCalledWith(rangeFailure));
expect(toast.success).not.toHaveBeenCalled();
});
it("writes the private ranges and the allowed clients one after the other, never concurrently", async () => {
vi.mocked(fetchMCPClientIp).mockResolvedValue("203.0.113.45");
let finishRangeWrite: (() => void) | undefined;
vi.mocked(updateConfigFieldSetting).mockImplementation(
(_token, fieldName) =>
new Promise<void>((resolve) => {
if (fieldName === "mcp_internal_ip_ranges") {
finishRangeWrite = resolve;
} else {
resolve();
}
}),
);
renderSettings();
await userEvent.click(await screen.findByText("203.0.113.0/24"));
await addClient("Codex", "codex-mcp-client");
await userEvent.click(screen.getByRole("button", { name: /Save/ }));
await waitFor(() =>
expect(updateConfigFieldSetting).toHaveBeenCalledWith("tok", "mcp_internal_ip_ranges", ["203.0.113.0/24"]),
);
expect(updateConfigFieldSetting).not.toHaveBeenCalledWith("tok", "mcp_allowed_clients", expect.anything());
finishRangeWrite?.();
await waitFor(() => expect(updateConfigFieldSetting).toHaveBeenCalledWith("tok", "mcp_allowed_clients", [CODEX]));
await waitFor(() => expect(toast.success).toHaveBeenCalledWith("MCP network settings saved"));
});
});

View file

@ -1,11 +1,21 @@
import React, { useState, useEffect } from "react";
import React, { useState, useEffect, useId } from "react";
import { Save, Plus, X } from "lucide-react";
import { Badge } from "@/components/ui/badge";
import { Button } from "@/components/ui/button";
import { Card } from "@/components/ui/card";
import {
Dialog,
DialogContent,
DialogDescription,
DialogFooter,
DialogHeader,
DialogTitle,
} from "@/components/ui/dialog";
import { Input } from "@/components/ui/input";
import { Label } from "@/components/ui/label";
import { UiLoadingSpinner } from "@/components/ui/ui-loading-spinner";
import { DeprecationBanner } from "@/components/DeprecationBanner";
import { toast } from "@/lib/toast";
import {
getGeneralSettingsCall,
updateConfigFieldSetting,
@ -26,12 +36,142 @@ function ipToSlash24(ip: string): string {
return `${parts[0]}.${parts[1]}.${parts[2]}.0/24`;
}
export interface AllowedClient {
readonly alias: string;
readonly value: string;
}
interface AllowedClientRow extends AllowedClient {
readonly key: string;
}
interface ClientDraft extends AllowedClient {
readonly key: string | null;
}
const isAllowedClient = (entry: unknown): entry is AllowedClient => {
if (typeof entry !== "object" || entry === null) return false;
const { alias, value } = entry as Partial<Record<keyof AllowedClient, unknown>>;
return typeof alias === "string" && typeof value === "string" && !isIncomplete({ alias, value });
};
type StoredAllowlist =
| { readonly kind: "absent" }
| { readonly kind: "clients"; readonly clients: AllowedClient[] }
| { readonly kind: "malformed" };
const ABSENT: StoredAllowlist = { kind: "absent" };
const parseStoredClients = (fieldValue: unknown): StoredAllowlist => {
if (fieldValue === null || fieldValue === undefined) return ABSENT;
if (Array.isArray(fieldValue) && fieldValue.every(isAllowedClient)) {
return { kind: "clients", clients: fieldValue.map(({ alias, value }) => ({ alias, value })) };
}
return { kind: "malformed" };
};
let nextRowKey = 0;
const newRow = (client: AllowedClient): AllowedClientRow => ({ ...client, key: `client-${nextRowKey++}` });
const trimClient = ({ alias, value }: AllowedClient): AllowedClient => ({ alias: alias.trim(), value: value.trim() });
const isIncomplete = ({ alias, value }: AllowedClient) => alias === "" || value === "";
const sameList = (a: string[], b: string[]) => a.length === b.length && a.every((value, i) => value === b[i]);
const sameClients = (a: AllowedClient[], b: AllowedClient[]) =>
a.length === b.length && a.every((client, i) => client.alias === b[i].alias && client.value === b[i].value);
const unchangedSinceLoad = (value: string[], stored: string[] | null) =>
stored === null ? value.length === 0 : value.length > 0 && sameList(value, stored);
const clientsUnchangedSinceLoad = (value: AllowedClient[], stored: StoredAllowlist) => {
switch (stored.kind) {
case "absent":
return value.length === 0;
case "clients":
return value.length > 0 && sameClients(value, stored.clients);
case "malformed":
return false;
}
};
const headerUnchangedSinceLoad = (value: string, stored: string | null) =>
stored === null ? value === "" : value !== "" && value === stored;
interface AllowedClientDialogProps {
readonly draft: ClientDraft | null;
readonly onChange: (draft: ClientDraft) => void;
readonly onCommit: () => void;
readonly onRemove: () => void;
readonly onClose: () => void;
}
const AllowedClientDialog: React.FC<AllowedClientDialogProps> = ({ draft, onChange, onCommit, onRemove, onClose }) => {
const aliasId = useId();
const valueId = useId();
if (draft === null) return null;
return (
<Dialog open onOpenChange={(open) => !open && onClose()}>
<DialogContent>
<DialogHeader>
<DialogTitle>{draft.key === null ? "Add client" : "Edit client"}</DialogTitle>
<DialogDescription>
The alias is the name shown in the dashboard and gateway logs. The value is the exact JWT claim or header
value that identifies the client, such as the OAuth client ID your identity provider issues.
</DialogDescription>
</DialogHeader>
<div className="grid gap-4">
<div className="grid gap-2">
<Label htmlFor={aliasId}>Alias</Label>
<Input
id={aliasId}
value={draft.alias}
placeholder="e.g. Coding CLI"
onChange={(e) => onChange({ ...draft, alias: e.target.value })}
/>
</div>
<div className="grid gap-2">
<Label htmlFor={valueId}>Value</Label>
<Input
id={valueId}
value={draft.value}
placeholder="e.g. 0oa1b2c3d4e5f6g7h8i9"
className="font-mono"
onChange={(e) => onChange({ ...draft, value: e.target.value })}
/>
</div>
</div>
<DialogFooter>
{draft.key !== null && (
<Button type="button" variant="destructive" className="sm:mr-auto" onClick={onRemove}>
Remove client
</Button>
)}
<Button type="button" variant="outline" onClick={onClose}>
Cancel
</Button>
<Button type="button" disabled={isIncomplete(trimClient(draft))} onClick={onCommit}>
{draft.key === null ? "Add" : "Done"}
</Button>
</DialogFooter>
</DialogContent>
</Dialog>
);
};
const MCPNetworkSettings: React.FC<MCPNetworkSettingsProps> = ({ accessToken }) => {
const [loading, setLoading] = useState(true);
const [saving, setSaving] = useState(false);
const [privateRanges, setPrivateRanges] = useState<string[]>([]);
const [allowedClients, setAllowedClients] = useState<AllowedClientRow[]>([]);
const [clientIdHeader, setClientIdHeader] = useState("");
const [storedRanges, setStoredRanges] = useState<string[] | null>(null);
const [storedClients, setStoredClients] = useState<StoredAllowlist>(ABSENT);
const [storedClientIdHeader, setStoredClientIdHeader] = useState<string | null>(null);
const [currentIp, setCurrentIp] = useState<string | null>(null);
const [rangeDraft, setRangeDraft] = useState("");
const [clientDraft, setClientDraft] = useState<ClientDraft | null>(null);
useEffect(() => {
loadSettings();
@ -44,8 +184,18 @@ const MCPNetworkSettings: React.FC<MCPNetworkSettingsProps> = ({ accessToken })
try {
const settings = await getGeneralSettingsCall(accessToken);
for (const field of settings) {
if (field.field_name === "mcp_internal_ip_ranges" && field.field_value) {
if (field.field_name === "mcp_internal_ip_ranges" && Array.isArray(field.field_value)) {
setPrivateRanges(field.field_value);
setStoredRanges(field.field_value);
}
if (field.field_name === "mcp_allowed_clients") {
const stored = parseStoredClients(field.field_value);
setAllowedClients(stored.kind === "clients" ? stored.clients.map(newRow) : []);
setStoredClients(stored);
}
if (field.field_name === "mcp_client_id_header" && typeof field.field_value === "string") {
setClientIdHeader(field.field_value);
setStoredClientIdHeader(field.field_value);
}
}
} catch (error) {
@ -63,20 +213,56 @@ const MCPNetworkSettings: React.FC<MCPNetworkSettingsProps> = ({ accessToken })
}
};
const persistRanges = async (token: string) => {
if (unchangedSinceLoad(privateRanges, storedRanges)) return;
if (privateRanges.length > 0) {
await updateConfigFieldSetting(token, "mcp_internal_ip_ranges", privateRanges);
setStoredRanges(privateRanges);
return;
}
await deleteConfigFieldSetting(token, "mcp_internal_ip_ranges");
setStoredRanges(null);
};
const persistAllowedClients = async (token: string) => {
const clients = allowedClients.map(({ alias, value }) => ({ alias, value }));
if (clientsUnchangedSinceLoad(clients, storedClients)) return;
if (clients.length > 0) {
await updateConfigFieldSetting(token, "mcp_allowed_clients", clients);
setStoredClients({ kind: "clients", clients });
return;
}
await deleteConfigFieldSetting(token, "mcp_allowed_clients");
setStoredClients(ABSENT);
};
const persistClientIdHeader = async (token: string) => {
const value = clientIdHeader.trim();
if (headerUnchangedSinceLoad(value, storedClientIdHeader)) return;
if (value !== "") {
await updateConfigFieldSetting(token, "mcp_client_id_header", value);
setStoredClientIdHeader(value);
return;
}
await deleteConfigFieldSetting(token, "mcp_client_id_header");
setStoredClientIdHeader(null);
};
const handleSave = async () => {
if (!accessToken) return;
setSaving(true);
try {
if (privateRanges.length > 0) {
await updateConfigFieldSetting(accessToken, "mcp_internal_ip_ranges", privateRanges);
} else {
await deleteConfigFieldSetting(accessToken, "mcp_internal_ip_ranges");
}
} catch (error) {
console.error("Failed to save MCP network settings:", error);
} finally {
setSaving(false);
const [rangeResult] = await Promise.allSettled([persistRanges(accessToken)]);
const [clientResult] = await Promise.allSettled([persistAllowedClients(accessToken)]);
const [headerResult] = await Promise.allSettled([persistClientIdHeader(accessToken)]);
setSaving(false);
const failures = [rangeResult, clientResult, headerResult].filter(
(result): result is PromiseRejectedResult => result.status === "rejected",
);
if (failures.length === 0) {
toast.success("MCP network settings saved");
return;
}
failures.forEach((failure) => toast.fromError(failure.reason));
};
const addSuggestedRange = (range: string) => {
@ -86,17 +272,37 @@ const MCPNetworkSettings: React.FC<MCPNetworkSettingsProps> = ({ accessToken })
};
// Commas separate entries, matching the old tokenised input.
const commitDraft = () => {
const added = rangeDraft
const splitDraft = (draft: string, existing: string[]) =>
draft
.split(",")
.map((r) => r.trim())
.filter((r) => r !== "" && !privateRanges.includes(r));
.filter((r) => r !== "" && !existing.includes(r));
const commitDraft = () => {
const added = splitDraft(rangeDraft, privateRanges);
if (added.length > 0) {
setPrivateRanges([...privateRanges, ...added]);
}
setRangeDraft("");
};
const commitClientDraft = () => {
if (clientDraft === null) return;
const client = trimClient(clientDraft);
setAllowedClients(
clientDraft.key === null
? [...allowedClients, newRow(client)]
: allowedClients.map((row) => (row.key === clientDraft.key ? { ...row, ...client } : row)),
);
setClientDraft(null);
};
const removeDraftedClient = () => {
if (clientDraft === null) return;
setAllowedClients(allowedClients.filter((row) => row.key !== clientDraft.key));
setClientDraft(null);
};
if (loading) {
return (
<div className="flex justify-center py-12">
@ -106,6 +312,8 @@ const MCPNetworkSettings: React.FC<MCPNetworkSettingsProps> = ({ accessToken })
}
const suggestedRange = currentIp ? ipToSlash24(currentIp) : null;
const storedAllowlistIsMalformed = storedClients.kind === "malformed";
const storedAllowlistIsEmpty = storedClients.kind === "clients" && storedClients.clients.length === 0;
return (
<div className="space-y-6 p-4">
@ -178,12 +386,88 @@ const MCPNetworkSettings: React.FC<MCPNetworkSettingsProps> = ({ accessToken })
</p>
</Card>
<div>
<p className="text-lg font-semibold">Allowed Clients</p>
<p className="mt-1 text-sm text-muted-foreground">
Only the MCP client applications listed here can use the gateway. Leave empty to allow every client. A client
that authenticates with a JWT is identified by the claim named in litellm_jwtauth.mcp_client_id_jwt_field in
your proxy config (for example azp or client_id), which your identity provider asserts and the client cannot
change. Any other client is identified by the request header configured below, if you enable one.
</p>
</div>
<Card className="p-6">
{storedAllowlistIsMalformed && (
<p className="mb-2 text-sm text-destructive">
The stored allowlist is not a list of alias and value pairs, so every client is denied. Add the clients you
want and save to replace it, or save with the list empty to remove it and allow every client again.
</p>
)}
{storedAllowlistIsEmpty && (
<p className="mb-2 text-sm text-destructive">
An empty allowlist is currently stored, so every client is denied. Save with the list empty to remove it and
allow every client again.
</p>
)}
{allowedClients.length > 0 && (
<div className="mb-3 grid gap-2 sm:grid-cols-2 lg:grid-cols-3">
{allowedClients.map((row) => (
<button
key={row.key}
type="button"
className="flex min-w-0 flex-col items-start gap-1 rounded-lg border border-border bg-background p-3 text-left hover:bg-muted focus-visible:ring-2 focus-visible:ring-ring focus-visible:outline-none"
onClick={() => setClientDraft(row)}
>
<span className="w-full truncate text-sm font-medium">{row.alias}</span>
<span className="w-full truncate font-mono text-xs text-muted-foreground">{row.value}</span>
</button>
))}
</div>
)}
<Button
type="button"
variant="outline"
size="sm"
onClick={() => setClientDraft({ key: null, alias: "", value: "" })}
>
<Plus />
Add client
</Button>
<p className="mt-2 text-xs text-muted-foreground">
Click a client to edit or remove it. Leave the list empty to allow every client. Every MCP request from an
unlisted client, or from one with no resolvable identity, gets a 403.
</p>
<div className="mt-6 mb-2 flex items-center">
<p className="text-sm font-medium">Client Identity Header (less secure)</p>
</div>
<Input
aria-label="Client identity header"
value={clientIdHeader}
placeholder="Leave empty to identify clients by JWT only, e.g. x-mcp-client"
onChange={(e) => setClientIdHeader(e.target.value)}
/>
<p className="mt-2 text-xs text-muted-foreground">
Optional header whose value names the client for callers without a JWT identity. Clients pick this value
themselves, so it is a policy control rather than a security boundary. Without it, callers that do not carry
the JWT claim are rejected while the allowlist is set.
</p>
</Card>
<div className="flex justify-end">
<Button onClick={handleSave} disabled={saving}>
<Save />
Save
</Button>
</div>
<AllowedClientDialog
draft={clientDraft}
onChange={setClientDraft}
onCommit={commitClientDraft}
onRemove={removeDraftedClient}
onClose={() => setClientDraft(null)}
/>
</div>
);
};

View file

@ -26978,6 +26978,16 @@ export interface components {
* @description Maximum retention period for spend logs (e.g., '7d' for 7 days). Logs older than this will be deleted.
*/
maximum_spend_logs_retention_period?: string | null;
/**
* Mcp Allowed Clients
* @description MCP client applications admitted by the gateway, each an {alias, value} pair where alias is the name shown in the dashboard and logs and value is the identity that must match exactly. When set, every MCP request must carry a client identity equal to one of the values: a JWT caller is identified by the claim named in litellm_jwtauth.mcp_client_id_jwt_field, any other caller by the header named in mcp_client_id_header. A request with no resolvable identity, or an unlisted one, is rejected with 403. Unset means every client is admitted.
*/
mcp_allowed_clients?: components["schemas"]["MCPAllowedClient"][] | null;
/**
* Mcp Client Id Header
* @description Request header whose value names the calling MCP client application (for example 'x-mcp-client') for callers that did not authenticate with a JWT, used only while mcp_allowed_clients is set. The client picks this value itself, so it is a policy control rather than a security boundary; prefer litellm_jwtauth.mcp_client_id_jwt_field where callers use JWTs.
*/
mcp_client_id_header?: string | null;
/**
* Mcp Internal Ip Ranges
* @description Custom CIDR ranges that define internal/private networks for MCP access control. When set, only these ranges are treated as internal. Defaults to RFC 1918 private ranges (10.0.0.0/8, 172.16.0.0/12, 192.168.0.0/16, 127.0.0.0/8).
@ -32640,6 +32650,22 @@ export interface components {
*/
status?: "healthy" | "unhealthy";
};
/**
* MCPAllowedClient
* @description One entry of `general_settings.mcp_allowed_clients`.
*/
MCPAllowedClient: {
/**
* Alias
* @description Human-readable name for this client application, shown in the dashboard and in gateway logs.
*/
alias: string;
/**
* Value
* @description Exact value of the JWT claim named in litellm_jwtauth.mcp_client_id_jwt_field, or of the mcp_client_id_header header, that identifies this client application. Matched case-sensitively.
*/
value: string;
};
/**
* MCPCatalogPrompt
* @description An MCP server's prompt as the upstream reports it. Subclassed only so the OpenAPI