mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
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:
commit
48421a4723
78 changed files with 4505 additions and 513 deletions
|
|
@ -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: |
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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 ####
|
||||
|
||||
|
|
|
|||
|
|
@ -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"))
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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":
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
171
litellm/proxy/_experimental/mcp_server/client_allowlist.py
Normal file
171
litellm/proxy/_experimental/mcp_server/client_allowlist.py
Normal 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
|
||||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
165
litellm/proxy/common_utils/responses_stream_errors.py
Normal file
165
litellm/proxy/common_utils/responses_stream_errors.py
Normal 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"
|
||||
|
|
@ -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__"
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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":
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -152,4 +152,7 @@ class PrismaBatch(Protocol):
|
|||
@property
|
||||
def litellm_modelaccessgroupbudgettable(self) -> BatchTable: ...
|
||||
|
||||
@property
|
||||
def litellm_projecttable(self) -> BatchTable: ...
|
||||
|
||||
async def commit(self) -> None: ...
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"])
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
],
|
||||
|
|
|
|||
82
tests/integration/pricing/test_off_peak_pricing.py
Normal file
82
tests/integration/pricing/test_off_peak_pricing.py
Normal 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)
|
||||
|
|
@ -13,5 +13,4 @@ litellm_settings:
|
|||
host: os.environ/REDIS_HOST
|
||||
port: os.environ/REDIS_PORT
|
||||
router_settings:
|
||||
num_retries: 0
|
||||
disable_cooldowns: true
|
||||
|
|
|
|||
|
|
@ -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"]})
|
||||
|
|
|
|||
|
|
@ -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"""
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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")]
|
||||
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
[
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"),
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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", '
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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 ────────────────────────────────────────────────────
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
# ---------------------------------------------------------------------------
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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 (
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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 (
|
||||
|
|
|
|||
|
|
@ -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}),
|
||||
]
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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"));
|
||||
});
|
||||
});
|
||||
|
|
|
|||
|
|
@ -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>
|
||||
);
|
||||
};
|
||||
|
|
|
|||
26
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
26
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue