Merge branch 'BerriAI:litellm_internal_staging' into fix/realtime-usage-detail-keys

This commit is contained in:
lmcdonald-godaddy 2026-06-03 11:23:55 -07:00 • committed by GitHub
commit 8dd404244d
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
247 changed files with 25784 additions and 1795 deletions

View file

@ -36,11 +36,13 @@ jobs:
tests/test_litellm/proxy/a2a
tests/test_litellm/proxy/discovery_endpoints
tests/test_litellm/proxy/health_endpoints
tests/test_litellm/proxy/shutdown
tests/test_litellm/proxy/public_endpoints
tests/test_litellm/proxy/prompts
tests/test_litellm/proxy/rag_endpoints
tests/test_litellm/proxy/realtime_endpoints
tests/test_litellm/proxy/ui_crud_endpoints
tests/test_litellm/proxy/utils
workers: 2
reruns: 2
artifact-name: proxy-endpoints

View file

@ -146,7 +146,7 @@ test-unit-proxy-core: install-test-deps
$(UV_RUN) pytest tests/test_litellm/proxy/auth tests/test_litellm/proxy/client tests/test_litellm/proxy/db tests/test_litellm/proxy/hooks tests/test_litellm/proxy/policy_engine --tb=short -vv -n 4 --durations=20
test-unit-proxy-misc: install-test-deps
$(UV_RUN) pytest tests/test_litellm/proxy/_experimental tests/test_litellm/proxy/agent_endpoints tests/test_litellm/proxy/anthropic_endpoints tests/test_litellm/proxy/common_utils tests/test_litellm/proxy/discovery_endpoints tests/test_litellm/proxy/experimental tests/test_litellm/proxy/google_endpoints tests/test_litellm/proxy/health_endpoints tests/test_litellm/proxy/image_endpoints tests/test_litellm/proxy/middleware tests/test_litellm/proxy/openai_files_endpoint tests/test_litellm/proxy/pass_through_endpoints tests/test_litellm/proxy/prompts tests/test_litellm/proxy/public_endpoints tests/test_litellm/proxy/response_api_endpoints tests/test_litellm/proxy/spend_tracking tests/test_litellm/proxy/ui_crud_endpoints tests/test_litellm/proxy/vector_store_endpoints tests/test_litellm/proxy/test_*.py --tb=short -vv -n 4 --durations=20
$(UV_RUN) pytest tests/test_litellm/proxy/_experimental tests/test_litellm/proxy/agent_endpoints tests/test_litellm/proxy/anthropic_endpoints tests/test_litellm/proxy/common_utils tests/test_litellm/proxy/discovery_endpoints tests/test_litellm/proxy/experimental tests/test_litellm/proxy/google_endpoints tests/test_litellm/proxy/health_endpoints tests/test_litellm/proxy/image_endpoints tests/test_litellm/proxy/middleware tests/test_litellm/proxy/openai_files_endpoint tests/test_litellm/proxy/pass_through_endpoints tests/test_litellm/proxy/prompts tests/test_litellm/proxy/public_endpoints tests/test_litellm/proxy/response_api_endpoints tests/test_litellm/proxy/shutdown tests/test_litellm/proxy/spend_tracking tests/test_litellm/proxy/ui_crud_endpoints tests/test_litellm/proxy/vector_store_endpoints tests/test_litellm/proxy/test_*.py --tb=short -vv -n 4 --durations=20
test-unit-integrations: install-test-deps
$(UV_RUN) pytest tests/test_litellm/integrations --tb=short -vv -n 4 --durations=20

View file

@ -285,11 +285,31 @@ db:
deployStandalone: true
# Lifecycle hooks for the LiteLLM container
#
# Prefer the native /health/drain preStop hook over a fixed `sleep`: it marks
# the pod NotReady and blocks only until in-flight requests actually finish
# (bounded by GRACEFUL_SHUTDOWN_TIMEOUT, default 30s), instead of always
# waiting the worst-case duration. The drain runs once (the preStop hook and
# the SIGTERM handler share it), so set terminationGracePeriodSeconds a few
# seconds above GRACEFUL_SHUTDOWN_TIMEOUT to leave room for teardown before
# SIGKILL.
#
# /health/drain is off by default; enable it with
# general_settings.enable_drain_endpoint: true. The kubelet calls preStop
# hooks without proxy credentials, so when the health port is reachable from
# other pods (the common case) also set
# general_settings.drain_endpoint_token (or the DRAIN_ENDPOINT_TOKEN env
# var) and send the same value on the X-Drain-Token header from the hook.
# Calls missing/wrong the token get a 401 and have no side effect.
# Example:
# lifecycle:
# preStop:
# exec:
# command: ["/bin/sh", "-c", "sleep 10"]
# httpGet:
# path: /health/drain
# port: 4000
# httpHeaders:
# - name: X-Drain-Token
# value: <same value as drain_endpoint_token>
lifecycle: {}
# Settings for Bitnami postgresql chart (if db.deployStandalone is true, ignored

View file

@ -0,0 +1,2 @@
-- AlterTable
ALTER TABLE "LiteLLM_MCPServerTable" ADD COLUMN IF NOT EXISTS "oauth_passthrough" BOOLEAN NOT NULL DEFAULT false;

View file

@ -325,6 +325,7 @@ model LiteLLM_MCPServerTable {
allow_all_keys Boolean @default(false)
available_on_public_internet Boolean @default(true)
delegate_auth_to_upstream Boolean @default(false)
oauth_passthrough Boolean @default(false)
is_byok Boolean @default(false)
byok_description String[] @default([])
byok_api_key_help_url String?

View file

@ -278,6 +278,7 @@ ovhcloud_key: Optional[str] = None
lemonade_key: Optional[str] = None
sap_service_key: Optional[str] = None
amazon_nova_api_key: Optional[str] = None
inception_key: Optional[str] = None
common_cloud_provider_auth_params: dict = {
"params": ["project", "region_name", "token"],
"providers": ["vertex_ai", "bedrock", "watsonx", "azure", "vertex_ai_beta"],
@ -551,6 +552,7 @@ cohere_models: Set = set()
cohere_chat_models: Set = set()
mistral_chat_models: Set = set()
text_completion_codestral_models: Set = set()
text_completion_inception_models: Set = set()
anthropic_models: Set = set()
openrouter_models: Set = set()
datarobot_models: Set = set()
@ -628,6 +630,7 @@ publicai_models: Set = set()
v0_models: Set = set()
morph_models: Set = set()
lambda_ai_models: Set = set()
inception_models: Set = set()
hyperbolic_models: Set = set()
black_forest_labs_models: Set = set()
recraft_models: Set = set()
@ -792,6 +795,8 @@ def add_known_models(model_cost_map: Optional[Dict] = None):
fireworks_ai_embedding_models.add(key)
elif value.get("litellm_provider") == "text-completion-codestral":
text_completion_codestral_models.add(key)
elif value.get("litellm_provider") == "text-completion-inception":
text_completion_inception_models.add(key)
elif value.get("litellm_provider") == "xai":
xai_models.add(key)
elif value.get("litellm_provider") == "zai":
@ -878,6 +883,8 @@ def add_known_models(model_cost_map: Optional[Dict] = None):
morph_models.add(key)
elif value.get("litellm_provider") == "lambda_ai":
lambda_ai_models.add(key)
elif value.get("litellm_provider") == "inception":
inception_models.add(key)
elif value.get("litellm_provider") == "hyperbolic":
hyperbolic_models.add(key)
elif value.get("litellm_provider") == "black_forest_labs":
@ -980,6 +987,7 @@ model_list = list(
| watsonx_models
| gemini_models
| text_completion_codestral_models
| text_completion_inception_models
| xai_models
| zai_models
| fal_ai_models
@ -1018,6 +1026,7 @@ model_list = list(
| v0_models
| morph_models
| lambda_ai_models
| inception_models
| black_forest_labs_models
| recraft_models
| cometapi_models
@ -1074,6 +1083,7 @@ models_by_provider: dict = {
"fireworks_ai": fireworks_ai_models | fireworks_ai_embedding_models,
"aleph_alpha": aleph_alpha_models,
"text-completion-codestral": text_completion_codestral_models,
"text-completion-inception": text_completion_inception_models,
"xai": xai_models,
"zai": zai_models,
"fal_ai": fal_ai_models,
@ -1118,6 +1128,7 @@ models_by_provider: dict = {
"v0": v0_models,
"morph": morph_models,
"lambda_ai": lambda_ai_models,
"inception": inception_models,
"hyperbolic": hyperbolic_models,
"black_forest_labs": black_forest_labs_models,
"recraft": recraft_models,
@ -1869,6 +1880,9 @@ if TYPE_CHECKING:
from .llms.codestral.completion.transformation import (
CodestralTextCompletionConfig as CodestralTextCompletionConfig,
)
from .llms.inception.completion.transformation import (
InceptionTextCompletionConfig as InceptionTextCompletionConfig,
)
from .llms.azure.azure import (
AzureOpenAIAssistantsAPIConfig as AzureOpenAIAssistantsAPIConfig,
)
@ -1937,6 +1951,9 @@ if TYPE_CHECKING:
from .llms.lambda_ai.chat.transformation import (
LambdaAIChatConfig as LambdaAIChatConfig,
)
from .llms.inception.chat.transformation import (
InceptionChatConfig as InceptionChatConfig,
)
from .llms.hyperbolic.chat.transformation import (
HyperbolicChatConfig as HyperbolicChatConfig,
)

View file

@ -267,6 +267,7 @@ LLM_CONFIG_NAMES = (
"AIMLChatConfig",
"VolcEngineChatConfig",
"CodestralTextCompletionConfig",
"InceptionTextCompletionConfig",
"AzureOpenAIAssistantsAPIConfig",
"HerokuChatConfig",
"CometAPIConfig",
@ -310,6 +311,7 @@ LLM_CONFIG_NAMES = (
"MorphChatConfig",
"RAGFlowConfig",
"LambdaAIChatConfig",
"InceptionChatConfig",
"HyperbolicChatConfig",
"VercelAIGatewayConfig",
"OVHCloudChatConfig",
@ -1040,6 +1042,10 @@ _LLM_CONFIGS_IMPORT_MAP = {
".llms.codestral.completion.transformation",
"CodestralTextCompletionConfig",
),
"InceptionTextCompletionConfig": (
".llms.inception.completion.transformation",
"InceptionTextCompletionConfig",
),
"AzureOpenAIAssistantsAPIConfig": (
".llms.azure.azure",
"AzureOpenAIAssistantsAPIConfig",
@ -1154,6 +1160,10 @@ _LLM_CONFIGS_IMPORT_MAP = {
"MorphChatConfig": (".llms.morph.chat.transformation", "MorphChatConfig"),
"RAGFlowConfig": (".llms.ragflow.chat.transformation", "RAGFlowConfig"),
"LambdaAIChatConfig": (".llms.lambda_ai.chat.transformation", "LambdaAIChatConfig"),
"InceptionChatConfig": (
".llms.inception.chat.transformation",
"InceptionChatConfig",
),
"HyperbolicChatConfig": (
".llms.hyperbolic.chat.transformation",
"HyperbolicChatConfig",

View file

@ -371,6 +371,8 @@ class ServiceLogging(CustomLogger):
service=ServiceTypes.LITELLM,
duration=_duration,
call_type=kwargs.get("call_type", "unknown"),
start_time=start_time,
end_time=end_time,
)
except Exception as e:
raise e

View file

@ -20,9 +20,20 @@ from litellm.a2a_protocol.litellm_completion_bridge.transformation import (
)
from litellm.a2a_protocol.providers.config_manager import A2AProviderConfigManager
# litellm_params key carrying the authenticated principal (hashed virtual key) so
# A2A provider configs can scope provider-side state (e.g. LangFlow session memory)
# per key instead of trusting the client-supplied A2A contextId.
A2A_USER_API_KEY_HASH_PARAM = "litellm_a2a_user_api_key_hash"
# Agent metadata fields stored in litellm_params that are not valid litellm.acompletion() kwargs
_AGENT_ONLY_PARAMS = frozenset(
{"is_public", "agent_name", "agent_id", "agent_card_params"}
{
"is_public",
"agent_name",
"agent_id",
"agent_card_params",
A2A_USER_API_KEY_HASH_PARAM,
}
)
@ -37,6 +48,8 @@ class A2ACompletionBridgeHandler:
params: Dict[str, Any],
litellm_params: Dict[str, Any],
api_base: Optional[str] = None,
*,
_skip_a2a_provider_routing: bool = False,
) -> Dict[str, Any]:
"""
Handle non-streaming A2A request via litellm.acompletion.
@ -50,25 +63,24 @@ class A2ACompletionBridgeHandler:
Returns:
A2A SendMessageResponse dict
"""
# Get provider config for custom_llm_provider
custom_llm_provider = litellm_params.get("custom_llm_provider")
a2a_provider_config = A2AProviderConfigManager.get_provider_config(
custom_llm_provider=custom_llm_provider,
model=litellm_params.get("model"),
)
# If provider config exists, use it
if a2a_provider_config is not None:
verbose_logger.info(f"A2A: Using provider config for {custom_llm_provider}")
response_data = await a2a_provider_config.handle_non_streaming(
request_id=request_id,
params=params,
api_base=api_base,
litellm_params=litellm_params,
if not _skip_a2a_provider_routing:
a2a_provider_config = A2AProviderConfigManager.get_provider_config(
custom_llm_provider=custom_llm_provider,
model=litellm_params.get("model"),
)
return response_data
if a2a_provider_config is not None:
verbose_logger.info(
f"A2A: Using provider config for {custom_llm_provider}"
)
return await a2a_provider_config.handle_non_streaming(
request_id=request_id,
params=params,
api_base=api_base,
litellm_params=litellm_params,
)
# Extract message from params
message = params.get("message", {})
@ -137,6 +149,8 @@ class A2ACompletionBridgeHandler:
params: Dict[str, Any],
litellm_params: Dict[str, Any],
api_base: Optional[str] = None,
*,
_skip_a2a_provider_routing: bool = False,
) -> AsyncIterator[Dict[str, Any]]:
"""
Handle streaming A2A request via litellm.acompletion with stream=True.
@ -156,28 +170,27 @@ class A2ACompletionBridgeHandler:
Yields:
A2A streaming response events
"""
# Get provider config for custom_llm_provider
custom_llm_provider = litellm_params.get("custom_llm_provider")
a2a_provider_config = A2AProviderConfigManager.get_provider_config(
custom_llm_provider=custom_llm_provider,
model=litellm_params.get("model"),
)
# If provider config exists, use it
if a2a_provider_config is not None:
verbose_logger.info(
f"A2A: Using provider config for {custom_llm_provider} (streaming)"
if not _skip_a2a_provider_routing:
a2a_provider_config = A2AProviderConfigManager.get_provider_config(
custom_llm_provider=custom_llm_provider,
model=litellm_params.get("model"),
)
async for chunk in a2a_provider_config.handle_streaming(
request_id=request_id,
params=params,
api_base=api_base,
litellm_params=litellm_params,
):
yield chunk
if a2a_provider_config is not None:
verbose_logger.info(
f"A2A: Using provider config for {custom_llm_provider} (streaming)"
)
return
async for chunk in a2a_provider_config.handle_streaming(
request_id=request_id,
params=params,
api_base=api_base,
litellm_params=litellm_params,
):
yield chunk
return
# Extract message from params
message = params.get("message", {})

View file

@ -159,7 +159,9 @@ async def _send_message_via_completion_bridge(
api_base=api_base,
)
return LiteLLMSendMessageResponse.from_dict(response_dict)
return LiteLLMSendMessageResponse.from_dict(
response_dict, request_id=str(request.id)
)
async def _execute_a2a_send_with_retry(
@ -317,15 +319,6 @@ async def asend_message(
)
card_url = getattr(agent_card, "url", None) if agent_card else None
context_id = trace_id or str(uuid.uuid4())
message = request.params.message
if isinstance(message, dict):
if message.get("context_id") is None:
message["context_id"] = context_id
else:
if getattr(message, "context_id", None) is None:
message.context_id = context_id
a2a_response = await _execute_a2a_send_with_retry(
a2a_client=a2a_client,
request=request,
@ -338,7 +331,9 @@ async def asend_message(
verbose_logger.info(f"A2A send_message completed, request_id={request.id}")
# Wrap in LiteLLM response type for _hidden_params support
response = LiteLLMSendMessageResponse.from_a2a_response(a2a_response)
response = LiteLLMSendMessageResponse.from_a2a_response(
a2a_response, request_id=str(request.id)
)
# Calculate token usage from request and response
response_dict = a2a_response.model_dump(mode="json", exclude_none=True)

View file

@ -48,6 +48,11 @@ class A2AProviderConfigManager:
return BedrockAgentCoreA2AConfig()
if custom_llm_provider == "langflow":
from litellm.a2a_protocol.providers.langflow.config import LangFlowA2AConfig
return LangFlowA2AConfig()
if custom_llm_provider == "watsonx_orchestrate":
from litellm.a2a_protocol.providers.watsonx_orchestrate.config import (
WatsonxOrchestrateA2AConfig,

View file

@ -0,0 +1,62 @@
from typing import Any, AsyncIterator, Dict, Optional
from litellm.a2a_protocol.litellm_completion_bridge.handler import (
A2A_USER_API_KEY_HASH_PARAM,
A2ACompletionBridgeHandler,
)
from litellm.a2a_protocol.providers.base import BaseA2AProviderConfig
from litellm.llms.langflow.a2a import merge_a2a_session_into_litellm_params
class LangFlowA2AConfig(BaseA2AProviderConfig):
"""A2A bridge for LangFlow: scopes contextId to the authenticated key as the
LangFlow session_id, then uses completion."""
async def handle_non_streaming(
self,
request_id: str,
params: Dict[str, Any],
api_base: Optional[str] = None,
**kwargs,
) -> Dict[str, Any]:
litellm_params = kwargs.get("litellm_params")
if not litellm_params:
raise ValueError(
"litellm_params is required for LangFlowA2AConfig "
"(must contain custom_llm_provider and model)"
)
litellm_params = merge_a2a_session_into_litellm_params(
litellm_params, params, litellm_params.get(A2A_USER_API_KEY_HASH_PARAM)
)
return await A2ACompletionBridgeHandler.handle_non_streaming(
request_id=request_id,
params=params,
litellm_params=litellm_params,
api_base=api_base,
_skip_a2a_provider_routing=True,
)
async def handle_streaming(
self,
request_id: str,
params: Dict[str, Any],
api_base: Optional[str] = None,
**kwargs,
) -> AsyncIterator[Dict[str, Any]]:
litellm_params = kwargs.get("litellm_params")
if not litellm_params:
raise ValueError(
"litellm_params is required for LangFlowA2AConfig "
"(must contain custom_llm_provider and model)"
)
litellm_params = merge_a2a_session_into_litellm_params(
litellm_params, params, litellm_params.get(A2A_USER_API_KEY_HASH_PARAM)
)
async for chunk in A2ACompletionBridgeHandler.handle_streaming(
request_id=request_id,
params=params,
litellm_params=litellm_params,
api_base=api_base,
_skip_a2a_provider_routing=True,
):
yield chunk

View file

@ -182,6 +182,14 @@ class ResponsesToCompletionBridgeHandler:
client=kwargs.get("client"),
)
# Pin the resolved provider so `responses()` doesn't re-run
# `get_llm_provider()` on the model string and strip a second
# provider prefix (see GitHub issue #28505). request_data already
# carries `custom_llm_provider` via the spread of
# `sanitized_litellm_params`; overwriting it on the dict (rather
# than adding an explicit kwarg) avoids the duplicate-keyword
# TypeError that would otherwise fire on the real bridge path.
request_data["custom_llm_provider"] = custom_llm_provider
result = responses(
**request_data,
)
@ -268,6 +276,13 @@ class ResponsesToCompletionBridgeHandler:
except Exception as e:
raise e
# Pin the resolved provider so `aresponses()` doesn't re-run
# `get_llm_provider()` on the model string and strip a second
# provider prefix (see GitHub issue #28505). Set on request_data
# rather than passed as a separate kwarg to avoid the duplicate-
# keyword TypeError when `sanitized_litellm_params` already
# carries `custom_llm_provider`.
request_data["custom_llm_provider"] = custom_llm_provider
result = await aresponses(
**request_data,
aresponses=True,

View file

@ -585,6 +585,7 @@ LITELLM_CHAT_PROVIDERS = [
"volcengine",
"codestral",
"text-completion-codestral",
"text-completion-inception",
"deepseek",
"sambanova",
"maritalk",
@ -620,6 +621,7 @@ LITELLM_CHAT_PROVIDERS = [
"oci",
"morph",
"lambda_ai",
"inception",
"vercel_ai_gateway",
"wandb",
"ovhcloud",
@ -779,6 +781,7 @@ openai_compatible_endpoints: List = [
"https://api.v0.dev/v1",
"https://api.morphllm.com/v1",
"https://api.lambda.ai/v1",
"https://api.inceptionlabs.ai/v1",
"https://api.hyperbolic.xyz/v1",
"https://ai-gateway.helicone.ai/",
"https://ai-gateway.vercel.sh/v1",
@ -835,6 +838,7 @@ openai_compatible_providers: List = [
"helicone",
"morph",
"lambda_ai",
"inception",
"hyperbolic",
"vercel_ai_gateway",
"aiml",

View file

@ -421,8 +421,16 @@ class MCPClient:
return factory
async def list_tools(self) -> List[MCPTool]:
"""List available tools from the server."""
async def list_tools(self, raise_on_error: bool = False) -> List[MCPTool]:
"""List available tools from the server.
Args:
raise_on_error: When True, re-raise exceptions instead of returning
an empty list. Used by the proxy's pass-through MCP flow so it
can surface upstream HTTP 401 responses as a proper 401 to the
MCP client (triggering the upstream OAuth flow) rather than
masking them as "connected, no tools".
"""
verbose_logger.debug(
f"MCP client listing tools from {self.server_url or 'stdio'}"
)
@ -458,6 +466,8 @@ class MCPClient:
"the MCP server may have crashed, disconnected, or timed out"
)
if raise_on_error:
raise
# Return empty list instead of raising to allow graceful degradation
return []

View file

@ -662,6 +662,16 @@ class CustomGuardrail(CustomLogger):
request_data["metadata"] = {}
_append_guardrail_info(request_data["metadata"])
# Emit the otel guardrail span here, where every guardrail execution lands,
# rather than relying on a post-call hook that does not fire on every path
# (e.g. a pass-through request that passes its guardrails).
try:
from litellm.integrations.otel.logger import emit_guardrail_span
emit_guardrail_span(slg)
except Exception:
pass
async def apply_guardrail(
self,
inputs: GenericGuardrailAPIInputs,

View file

@ -1012,6 +1012,13 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
litellm_params = kwargs.get("litellm_params", {}) or {}
_metadata = litellm_params.get("metadata", {}) or {}
proxy_span = _metadata.get("litellm_parent_otel_span", None)
# Fallback: check litellm_metadata (used by /v1/messages and other
# LITELLM_METADATA_ROUTES).
if proxy_span is None:
_litellm_metadata = litellm_params.get("litellm_metadata", {}) or {}
proxy_span = _litellm_metadata.get("litellm_parent_otel_span", None)
if (
proxy_span is not None
and getattr(proxy_span, "name", None) == LITELLM_PROXY_REQUEST_SPAN_NAME
@ -2668,6 +2675,10 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
)
def _to_ns(self, dt):
if dt is None:
return int(datetime.now().timestamp() * 1e9)
if isinstance(dt, (int, float)):
return int(dt * 1e9)
return int(dt.timestamp() * 1e9)
def _get_span_name(self, kwargs):
@ -2714,6 +2725,13 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
_metadata = litellm_params.get("metadata", {}) or {}
parent_otel_span = _metadata.get("litellm_parent_otel_span", None)
# Fallback: check litellm_metadata (used by /v1/messages and other
# LITELLM_METADATA_ROUTES that store proxy-internal metadata
# separately from the provider's native "metadata" field).
if parent_otel_span is None:
_litellm_metadata = litellm_params.get("litellm_metadata", {}) or {}
parent_otel_span = _litellm_metadata.get("litellm_parent_otel_span", None)
# Priority 1: Explicit parent span from metadata
if parent_otel_span is not None:
verbose_logger.debug(
@ -3287,6 +3305,32 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger):
value=int(status_code),
)
def record_error_attributes_on_span(
self,
span: Optional[Span],
exception: Optional[Exception],
status_code: int,
) -> None:
"""Stamp structured ``error.*`` attributes on the SERVER span from the
exception returned to the client, with ``error.code`` pinned to the real
response status. Idempotent (overwrites); emits no exception event."""
if span is None or exception is None:
return
from litellm.litellm_core_utils.litellm_logging import (
StandardLoggingPayloadSetup,
)
error_information = StandardLoggingPayloadSetup.get_error_information(
original_exception=exception
)
error_information["error_code"] = str(status_code)
self._record_exception_on_span(
span=span,
kwargs={
"standard_logging_object": {"error_information": error_information}
},
)
def set_preprocessing_duration_attribute(
self, span: Optional[Span], container: Any
) -> None:

View file

@ -39,20 +39,32 @@ def extract_opik_metadata(
standard_logging_metadata: Dict[str, Any],
) -> Dict[str, Any]:
"""
Extract and merge Opik metadata from request and requester.
Merge Opik metadata from three sources in increasing priority order:
1. user_api_key_auth_metadata– lowest priority (operator-level defaults)
2. litellm_metadata (request)– overrides auth-key defaults
3. requester_metadata – highest priority (e.g. proxy header overrides)
Args:
litellm_metadata: Metadata from litellm_params
standard_logging_metadata: Metadata from standard_logging_object
litellm_metadata: Metadata from litellm_params.mak
standard_logging_metadata: Metadata from standard_logging_object.
Returns:
Merged Opik metadata dictionary
Merged Opik metadata dictionary.
"""
opik_meta = litellm_metadata.get("opik", {}).copy()
# Start with auth-key defaults (lowest priority).
auth_meta = standard_logging_metadata.get("user_api_key_auth_metadata") or {}
opik_meta = (auth_meta.get("opik") or {}).copy()
# Request-level values override auth-key defaults.
request_opik = litellm_metadata.get("opik") or {}
opik_meta.update(request_opik)
# Requester-level values win over everything else.
requester_metadata = standard_logging_metadata.get("requester_metadata", {}) or {}
requester_opik = requester_metadata.get("opik", {}) or {}
opik_meta.update(requester_opik)
if requester_opik:
opik_meta.update(requester_opik)
_logging.verbose_logger.debug(
f"litellm_opik_metadata - {json.dumps(opik_meta, default=str)}"

View file

@ -32,20 +32,28 @@ from litellm.integrations.otel.model.payloads import (
LLMCallSpanData,
LLMRequestParams,
LLMUsage,
MCPToolCallSpanData,
ProxyRequestSpanData,
ServerInfo,
ServiceSpanData,
SpanError,
is_mcp_tool_call,
)
from litellm.integrations.otel.model.semconv import (
DB,
HTTP,
MCP,
Client,
Error,
GenAI,
GenAIOperation,
GenAIProvider,
HTTP,
JsonRpc,
LiteLLM,
MCPMethod,
Metric,
Network,
NetworkTransport,
Server,
resolve_operation,
resolve_provider,
@ -69,13 +77,19 @@ __all__ = [
"BAGGAGE_PROMOTED_KEYS",
"DB",
"DEFAULT_BAGGAGE_METADATA_KEYS",
"Client",
"Error",
"GenAI",
"GenAIOperation",
"GenAIProvider",
"HTTP",
"JsonRpc",
"LiteLLM",
"MCP",
"MCPMethod",
"Metric",
"Network",
"NetworkTransport",
"Server",
"resolve_operation",
"resolve_provider",
@ -92,11 +106,13 @@ __all__ = [
"LLMCallSpanData",
"LLMRequestParams",
"LLMUsage",
"MCPToolCallSpanData",
"ProxyRequestSpanData",
"RequestContext",
"RequestIdentity",
"ServerInfo",
"ServiceSpanData",
"SpanError",
"is_mcp_tool_call",
"promoted_baggage",
]

View file

@ -13,6 +13,7 @@ from litellm.integrations.otel.mappers.base import AttributeMapper, SpanData
from litellm.integrations.otel.model.payloads import (
GuardrailSpanData,
LLMCallSpanData,
MCPToolCallSpanData,
ServiceSpanData,
)
from litellm.integrations.otel.plumbing.providers import to_otel_span_kind
@ -22,6 +23,7 @@ from litellm.integrations.otel.model.spans import (
SpanRole,
guardrail_span_name,
llm_call_span_name,
mcp_tool_call_span_name,
service_span_name,
)
@ -30,6 +32,7 @@ from litellm.integrations.otel.model.spans import (
# have no builder here.
_NAME_BUILDERS: dict[SpanRole, Callable[..., str]] = {
SpanRole.LLM_CALL: llm_call_span_name,
SpanRole.MCP_TOOL_CALL: mcp_tool_call_span_name,
SpanRole.GUARDRAIL: guardrail_span_name,
# DB_CALL and SERVICE are both built from ServiceSpanData; they differ only in
# span kind (CLIENT vs INTERNAL) and attribute vocabulary, not in naming.
@ -121,10 +124,14 @@ class SpanEmitter:
Return the span, or ``None`` if it was deduplicated away. ``tracer``
overrides the bound tracer for this span, used for per-request routing.
"""
# Only LLM-call spans carry a dedup key; LLM-call and service spans
# carry an ``error`` field. ``isinstance`` narrows the type for mypy and
# keeps the engine free of duck-typed attribute reads.
dedup_key = data.identity.call_id if isinstance(data, LLMCallSpanData) else None
# LLM-call and MCP tool-call spans carry a dedup key (their request's
# call id), so a sync+async double-firing coalesces. ``isinstance`` narrows
# the type for mypy and keeps the engine free of duck-typed attribute reads.
dedup_key = (
data.identity.call_id
if isinstance(data, (LLMCallSpanData, MCPToolCallSpanData))
else None
)
if self._seen(dedup_key, role):
return None
span = self.start_span(
@ -160,7 +167,15 @@ class SpanEmitter:
span.set_attribute(key, value)
error = (
data.error
if isinstance(data, (LLMCallSpanData, ServiceSpanData, GuardrailSpanData))
if isinstance(
data,
(
LLMCallSpanData,
MCPToolCallSpanData,
ServiceSpanData,
GuardrailSpanData,
),
)
else None
)
if error and (error.error_type or error.message):

View file

@ -25,14 +25,15 @@ from litellm.integrations.otel.mappers import resolve_mappers
from litellm.integrations.otel.model.metadata import (
LLMCallEvent,
RequestIdentity,
guardrail_entries_from_request_data,
model_from_request_data,
)
from litellm.integrations.otel.model.payloads import (
GuardrailSpanData,
LLMCallSpanData,
MCPToolCallSpanData,
ServiceSpanData,
SpanError,
is_mcp_tool_call,
)
from litellm.integrations.otel.plumbing.providers import (
build_tracer_provider,
@ -43,7 +44,10 @@ from litellm.integrations.otel.model.spans import SpanRole, span_role_for_servic
from litellm.integrations.otel.model.utils import to_ns
if TYPE_CHECKING:
from litellm.types.utils import StandardLoggingGuardrailInformation
from litellm.types.utils import (
StandardLoggingGuardrailInformation,
StandardLoggingPayload,
)
LITELLM_TRACER_NAME = "litellm"
@ -200,11 +204,53 @@ class OpenTelemetryV2(CustomLogger):
self._open_llm_calls.popitem(last=False)
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
if self._emit_mcp_tool_call(kwargs, start_time, end_time):
return
self._close_llm_call(kwargs, start_time, end_time)
async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time):
if self._emit_mcp_tool_call(kwargs, start_time, end_time):
return
self._close_llm_call(kwargs, start_time, end_time)
def _emit_mcp_tool_call(
self,
kwargs: Mapping[str, Any],
start_time: datetime | float | None,
end_time: datetime | float | None,
) -> bool:
"""Emit an MCP tool-call span when the closed request was a tool call.
MCP tool calls reach the success/failure callbacks like any other request
(with ``call_type`` ``call_mcp_tool``), but they are not LLM calls and have
no ``pre_call`` carrier — so they get their own CLIENT span here, parented
to the request's server span. Returns whether it handled the event, so the
caller skips the LLM-call path. The whole span is emitted at once (there is
no boundary to open it at), deduped on the call id by the emitter.
"""
raw_payload = kwargs.get("standard_logging_object")
if not raw_payload or not is_mcp_tool_call(
cast(Mapping[str, object], raw_payload)
):
return False
payload = cast("StandardLoggingPayload", raw_payload)
data = MCPToolCallSpanData.from_standard_logging_payload(
payload, capture_content=self.config.capture_span_content
)
# A stray LLM carrier from a ``pre_call`` that mis-fired for this id would
# otherwise linger until evicted; drop it so it's neither leaked nor closed
# as a phantom LLM span.
if data.identity.call_id:
self._open_llm_calls.pop(data.identity.call_id, None)
self._emitter.emit(
SpanRole.MCP_TOOL_CALL,
data,
parent_context=resolve_request_span_context(),
start_time_ns=to_ns(start_time),
end_time_ns=to_ns(end_time),
)
return True
def _close_llm_call(
self,
kwargs: Mapping[str, Any],
@ -421,46 +467,27 @@ class OpenTelemetryV2(CustomLogger):
)
return data
async def async_post_call_success_hook(
self,
data: Mapping[str, Any],
user_api_key_dict: Any,
response: Any,
) -> Any:
self._emit_guardrail_spans(data)
return response
async def async_post_call_failure_hook(
self,
request_data: Mapping[str, Any],
original_exception: BaseException | None,
user_api_key_dict: Any,
traceback_str: str | None = None,
) -> None:
self._emit_guardrail_spans(request_data)
def _emit_guardrail_spans(self, request_data: Mapping[str, Any]) -> None:
def emit_guardrail_span(self, entry: "StandardLoggingGuardrailInformation") -> None:
# Emitted by the guardrail-recording code the moment a guardrail finishes,
# not from a post-call hook — that hook does not fire on every path (a
# pass-through request that passes its guardrails never reaches it), which
# left passing guardrails without a span.
#
# A guardrail is a sibling of the LLM call under the request's root span,
# so parent it to the explicit anchor — not the active span, which on the
# failure path can be the live ``auth`` phase span (post-call failure hooks
# run from inside it on an auth rejection). Emit with the guardrail's actual
# execution window so a pre_call guardrail is placed before the LLM call
# rather than at post-call emission time.
guardrails = guardrail_entries_from_request_data(request_data)
if not guardrails:
return
parent_ctx = resolve_request_span_context()
for entry in guardrails:
data = GuardrailSpanData.from_logging_entry(
cast("StandardLoggingGuardrailInformation", entry)
)
self._emitter.emit(
SpanRole.GUARDRAIL,
data,
parent_context=parent_ctx,
start_time_ns=to_ns(data.start_time),
end_time_ns=to_ns(data.end_time),
)
# so parent it to the explicit anchor — never the active span, which during
# a pre_call guardrail can be the live ``auth`` phase span. Emit with the
# guardrail's actual execution window so a pre_call guardrail is placed
# before the LLM call rather than at emission time. One entry in, one span
# out — the module-level entry point routes each entry to this single
# registered logger so a guardrail is never emitted more than once.
data = GuardrailSpanData.from_logging_entry(entry)
self._emitter.emit(
SpanRole.GUARDRAIL,
data,
parent_context=resolve_request_span_context(),
start_time_ns=to_ns(data.start_time),
end_time_ns=to_ns(data.end_time),
)
def create_litellm_proxy_request_started_span(
self, start_time: datetime, headers: Mapping[str, str] | None
@ -481,6 +508,26 @@ def _registered_v2_logger() -> "OpenTelemetryV2 | None":
return logger if isinstance(logger, OpenTelemetryV2) else None
def emit_guardrail_span(entry: "StandardLoggingGuardrailInformation") -> None:
"""Emit a guardrail span on the registered v2 OTel logger.
Called by the guardrail-recording code the moment a guardrail finishes, so a
span is produced regardless of whether a post-call hook later runs (it does
not on the pass-through allow path). Routes through the single canonical
logger — the same one every other v2 entry point uses — so a guardrail
recorded once yields exactly one span; fanning out across every reachable
``OpenTelemetryV2`` instance double-emits the same entry. Best-effort: span
emission must never break guardrail evaluation.
"""
logger = _registered_v2_logger()
if logger is None:
return
try:
logger.emit_guardrail_span(entry)
except Exception:
pass
def seed_request_identity(user_api_key_dict: Any, model: Any = None) -> None:
logger = _registered_v2_logger()
if logger is not None:

View file

@ -7,6 +7,7 @@ from typing_extensions import Protocol, runtime_checkable
from litellm.integrations.otel.model.payloads import (
GuardrailSpanData,
LLMCallSpanData,
MCPToolCallSpanData,
ServiceSpanData,
)
@ -21,7 +22,7 @@ AttributeMap = dict[str, AttrValue]
# The closed set of span-data types the engine routes through the mapper chain.
# Server spans (PROXY_REQUEST + management routes) belong to the mounted FastAPI
# instrumentor, not the mapper chain.
SpanData = LLMCallSpanData | GuardrailSpanData | ServiceSpanData
SpanData = LLMCallSpanData | MCPToolCallSpanData | GuardrailSpanData | ServiceSpanData
@runtime_checkable

View file

@ -14,10 +14,18 @@ from litellm.integrations.otel.mappers.utils import collect, drop_none
from litellm.integrations.otel.model.payloads import (
GuardrailSpanData,
LLMCallSpanData,
MCPToolCallSpanData,
ServiceSpanData,
ToolDefinition,
)
from litellm.integrations.otel.model.semconv import DB, Error, GenAI, LiteLLM, Server
from litellm.integrations.otel.model.semconv import (
DB,
MCP,
Error,
GenAI,
LiteLLM,
Server,
)
from litellm.integrations.otel.model.spans import db_system
@ -64,6 +72,18 @@ class GenAIMapper:
"parameters": lambda t: t.parameters_json or None,
}
_MCP_ATTRS: dict[str, Callable[[MCPToolCallSpanData], AttrValue | None]] = {
GenAI.OPERATION_NAME: lambda d: d.operation.value,
MCP.METHOD_NAME: lambda d: d.method,
MCP.SESSION_ID: lambda d: d.session_id,
GenAI.TOOL_NAME: lambda d: d.tool_name or None,
GenAI.TOOL_CALL_ARGUMENTS: lambda d: d.arguments_json,
GenAI.TOOL_CALL_RESULT: lambda d: d.result_json,
LiteLLM.MCP_SERVER_NAME: lambda d: d.server_name,
LiteLLM.CALL_ID: lambda d: d.identity.call_id or None,
f"{LiteLLM.COST_PREFIX}total": lambda d: d.response_cost,
}
_GUARDRAIL_ATTRS: dict[str, Callable[[GuardrailSpanData], AttrValue | None]] = {
LiteLLM.GUARDRAIL_NAME: lambda d: d.guardrail_name,
LiteLLM.GUARDRAIL_MODE: lambda d: d.mode,
@ -92,6 +112,8 @@ class GenAIMapper:
match data:
case LLMCallSpanData():
return self._llm_call(data)
case MCPToolCallSpanData():
return collect(self._MCP_ATTRS, data)
case GuardrailSpanData():
return self._guardrail(data)
case ServiceSpanData():

View file

@ -255,26 +255,6 @@ def model_from_request_data(data: object) -> str | None:
return None
def guardrail_entries_from_request_data(
request_data: Mapping[str, Any],
) -> list[dict]:
"""The guardrail-information dicts buried in ``metadata`` of a post-call dict.
``standard_logging_guardrail_information`` is stored as either a single dict
or a list of them; normalize to a list of dicts (dropping non-dict noise) so
the caller just iterates. Empty list when none are present.
"""
metadata = request_data.get("metadata")
if not isinstance(metadata, Mapping):
return []
info = metadata.get("standard_logging_guardrail_information")
if isinstance(info, Mapping):
return [cast(dict, info)]
if isinstance(info, list):
return [entry for entry in info if isinstance(entry, dict)]
return []
def resolve_provider_model(payload: "StandardLoggingPayload") -> str | None:
"""The model litellm dispatched to the provider, from the payload.

View file

@ -14,6 +14,7 @@ from litellm.integrations.otel.model.metadata import (
)
from litellm.integrations.otel.model.semconv import (
GenAIOperation,
MCPMethod,
resolve_operation,
resolve_provider,
)
@ -35,11 +36,13 @@ __all__ = [
"LLMCallSpanData",
"LLMRequestParams",
"LLMUsage",
"MCPToolCallSpanData",
"ProxyRequestSpanData",
"ServerInfo",
"ServiceSpanData",
"SpanError",
"ToolDefinition",
"is_mcp_tool_call",
]
if TYPE_CHECKING:
@ -309,6 +312,77 @@ class LLMCallSpanData:
)
# --- the MCP tool-call model ------------------------------------------------- #
@dataclass(frozen=True)
class MCPToolCallSpanData:
"""One MCP ``tools/call`` execution, parsed from a closed request's payload.
The proxy is an MCP *client* to the upstream server it forwards the call to,
so this is a CLIENT span. ``arguments_json``/``result_json`` are the tool's
input/output — sensitive content, so they're only retained when content
capture is enabled, mirroring ``LLMCallSpanData``'s message bodies.
"""
operation: GenAIOperation
method: str
tool_name: str
server_name: str | None
session_id: str | None
arguments_json: str | None
result_json: str | None
error: SpanError | None
response_cost: float | None
identity: RequestIdentity
@classmethod
def from_standard_logging_payload(
cls, payload: "StandardLoggingPayload", capture_content: bool = False
) -> "MCPToolCallSpanData":
meta = _mcp_tool_call_metadata(cast(Mapping[str, object], payload))
return cls(
operation=resolve_operation(as_str(payload.get("call_type"))),
method=MCPMethod.TOOLS_CALL.value,
tool_name=as_str(meta.get("name")) or "",
server_name=as_str(meta.get("mcp_server_name")),
session_id=as_str(meta.get("mcp_session_id")),
arguments_json=(
_json_or_none(meta.get("arguments"))
if capture_content and meta.get("arguments") is not None
else None
),
result_json=(
_json_or_none(meta.get("result"))
if capture_content and meta.get("result") is not None
else None
),
error=_parse_error(payload),
response_cost=as_float(payload.get("response_cost")),
identity=RequestContext.from_standard_logging_payload(payload).identity,
)
def _mcp_tool_call_metadata(payload: Mapping[str, object]) -> Mapping[str, object]:
"""The MCP gateway's tool-call metadata, which lives under
``StandardLoggingPayload.metadata`` (a ``StandardLoggingMetadata`` key), not
at the payload's top level."""
metadata = payload.get("metadata")
if not isinstance(metadata, Mapping):
return {}
meta = metadata.get("mcp_tool_call_metadata")
return meta if isinstance(meta, Mapping) else {}
def is_mcp_tool_call(payload: Mapping[str, object]) -> bool:
"""Whether a closed request's payload is an MCP tool call rather than an LLM
call — true when the MCP gateway stamped its tool-call metadata, or the call
type says so on a path that hasn't populated the metadata yet."""
return bool(_mcp_tool_call_metadata(payload)) or (
payload.get("call_type") == "call_mcp_tool"
)
# --- service event_metadata sanitization ------------------------------------ #
# Substrings (case-insensitive) of keys that must never reach a span: secrets,

View file

@ -16,7 +16,7 @@ class GenAIOperation(str, Enum):
GENERATE_CONTENT = "generate_content"
CREATE_AGENT = "create_agent" # reserved for future agent spans
INVOKE_AGENT = "invoke_agent" # reserved for future agent spans
EXECUTE_TOOL = "execute_tool" # reserved for future tool spans
EXECUTE_TOOL = "execute_tool" # MCP tool-call spans
class GenAIProvider(str, Enum):
@ -38,6 +38,16 @@ class GenAIProvider(str, Enum):
IBM_WATSONX_AI = "ibm.watsonx.ai"
class MCPMethod(str, Enum):
"""Well-known values for ``mcp.method.name`` that litellm's MCP gateway
serves. The value is the JSON-RPC method exactly as it travels on the wire."""
TOOLS_CALL = "tools/call"
TOOLS_LIST = "tools/list"
PROMPTS_GET = "prompts/get"
PROMPTS_LIST = "prompts/list"
class GenAI:
"""Canonical OTel GenAI span-attribute keys."""
@ -68,11 +78,68 @@ class GenAI:
SYSTEM_INSTRUCTIONS: Final = "gen_ai.system_instructions"
OUTPUT_TYPE: Final = "gen_ai.output.type"
CONVERSATION_ID: Final = "gen_ai.conversation.id"
# agent / tool (reserved)
# agent (reserved)
AGENT_ID: Final = "gen_ai.agent.id"
AGENT_NAME: Final = "gen_ai.agent.name"
# tool / tool-call (stamped on MCP tool-call spans). Arguments and result are
# the tool's input/output payloads — sensitive, so they're opt-in and gated by
# the same content-capture mode as prompt/response content.
TOOL_NAME: Final = "gen_ai.tool.name"
TOOL_CALL_ID: Final = "gen_ai.tool.call.id"
TOOL_CALL_ARGUMENTS: Final = "gen_ai.tool.call.arguments"
TOOL_CALL_RESULT: Final = "gen_ai.tool.call.result"
# prompt (MCP ``prompts/get`` etc.)
PROMPT_NAME: Final = "gen_ai.prompt.name"
class MCP:
"""OTel GenAI MCP (Model Context Protocol) span-attribute keys.
``METHOD_NAME`` is the only key litellm populates from a closed request today;
the rest are part of the convention's vocabulary and are stamped when the
corresponding signal (session, protocol version, resource) is available.
"""
METHOD_NAME: Final = "mcp.method.name"
SESSION_ID: Final = "mcp.session.id"
PROTOCOL_VERSION: Final = "mcp.protocol.version"
RESOURCE_URI: Final = "mcp.resource.uri"
class JsonRpc:
"""JSON-RPC keys carried on MCP spans. The error/status code lives in the
``rpc.*`` namespace per semconv, not ``jsonrpc.*``."""
REQUEST_ID: Final = "jsonrpc.request.id"
PROTOCOL_VERSION: Final = "jsonrpc.protocol.version"
RESPONSE_STATUS_CODE: Final = "rpc.response.status_code"
class NetworkTransport(str, Enum):
"""Well-known values for ``network.transport``."""
TCP = "tcp"
UDP = "udp"
QUIC = "quic"
UNIX = "unix"
PIPE = "pipe"
class Network:
"""OTel network keys, recommended on MCP spans to describe the transport
carrying the JSON-RPC messages (stdio pipe, HTTP, websocket, …)."""
PROTOCOL_NAME: Final = "network.protocol.name"
PROTOCOL_VERSION: Final = "network.protocol.version"
TRANSPORT: Final = "network.transport"
class Client:
"""Peer (client) network keys, stamped on MCP *server* spans the same way
``server.*`` is stamped on client spans."""
ADDRESS: Final = "client.address"
PORT: Final = "client.port"
class Error:
@ -137,6 +204,10 @@ class LiteLLM:
SERVICE_NAME: Final = "litellm.service.name"
SERVICE_CALL_TYPE: Final = "litellm.service.call_type"
PREPROCESSING_MS: Final = "litellm.preprocessing.duration_ms"
# The logical name of the MCP server a tool call was routed to. There is no
# semconv key for an MCP server's *name* (the convention uses ``server.address``
# for its network location), so it lives under the vendor namespace.
MCP_SERVER_NAME: Final = "litellm.mcp.server.name"
class Metric:
@ -179,6 +250,7 @@ _OPERATION_BY_CALL_TYPE: dict[str, GenAIOperation] = {
"aembedding": GenAIOperation.EMBEDDINGS,
"responses": GenAIOperation.CHAT,
"aresponses": GenAIOperation.CHAT,
"call_mcp_tool": GenAIOperation.EXECUTE_TOOL,
}

View file

@ -46,6 +46,7 @@ if TYPE_CHECKING:
from litellm.integrations.otel.model.payloads import (
GuardrailSpanData,
LLMCallSpanData,
MCPToolCallSpanData,
ProxyRequestSpanData,
ServiceSpanData,
)
@ -54,6 +55,7 @@ if TYPE_CHECKING:
class SpanRole(str, Enum):
PROXY_REQUEST = "proxy_request"
LLM_CALL = "llm_call"
MCP_TOOL_CALL = "mcp_tool_call"
GUARDRAIL = "guardrail"
DB_CALL = "db_call"
SERVICE = "service"
@ -81,6 +83,11 @@ SPAN_REGISTRY: dict[SpanRole, SpanSpec] = {
SpanRole.LLM_CALL: SpanSpec(
SpanRole.LLM_CALL, LiteLLMSpanKind.CLIENT, parent=SpanRole.PROXY_REQUEST
),
# The proxy is an MCP client to the upstream server it dispatches the tool
# call to, so this is a CLIENT span, sibling of the LLM call under the request.
SpanRole.MCP_TOOL_CALL: SpanSpec(
SpanRole.MCP_TOOL_CALL, LiteLLMSpanKind.CLIENT, parent=SpanRole.PROXY_REQUEST
),
SpanRole.GUARDRAIL: SpanSpec(
SpanRole.GUARDRAIL, LiteLLMSpanKind.INTERNAL, parent=SpanRole.PROXY_REQUEST
),
@ -165,6 +172,11 @@ def llm_call_span_name(data: "LLMCallSpanData") -> str:
return f"{data.operation.value} {model}".strip()
def mcp_tool_call_span_name(data: "MCPToolCallSpanData") -> str:
"""``"{mcp.method.name} {tool}"`` e.g. ``"tools/call get-weather"`` (MCP semconv)."""
return f"{data.method} {data.tool_name}".strip()
def proxy_request_span_name(data: "ProxyRequestSpanData") -> str:
"""``"{method} {route}"`` (HTTP semconv)."""
return f"{data.http_method} {data.route}".strip()

View file

@ -1,3 +1,4 @@
import re
from typing import Optional, Tuple
from urllib.parse import urlparse
@ -71,6 +72,25 @@ def _is_azure_claude_model(model: str) -> bool:
return False
_CLAUDE_PATTERN = re.compile(r"^claude-[a-z]+-\d+-\d+(?:-\d{8})?$", re.IGNORECASE)
def _matches_claude_model_pattern(model: str) -> bool:
"""
Check if a model string matches the Claude model naming pattern.
Matches patterns like:
- claude-opus-4-7
- claude-sonnet-4-6
- claude-haiku-4-5
- claude-opus-5-1-20270101 (with optional date suffix)
This allows future Claude models to be routed to the Anthropic provider
without requiring updates to model_prices_and_context_window.json.
"""
return _CLAUDE_PATTERN.match(model) is not None
def handle_cohere_chat_model_custom_llm_provider(
model: str, custom_llm_provider: Optional[str] = None
) -> Tuple[str, Optional[str]]:
@ -353,6 +373,9 @@ def get_llm_provider( # noqa: PLR0915
elif endpoint == "https://api.lambda.ai/v1":
custom_llm_provider = "lambda_ai"
dynamic_api_key = get_secret_str("LAMBDA_API_KEY")
elif endpoint == "https://api.inceptionlabs.ai/v1":
custom_llm_provider = "inception"
dynamic_api_key = get_secret_str("INCEPTION_API_KEY")
elif endpoint == "https://api.hyperbolic.xyz/v1":
custom_llm_provider = "hyperbolic"
dynamic_api_key = get_secret_str("HYPERBOLIC_API_KEY")
@ -398,6 +421,9 @@ def get_llm_provider( # noqa: PLR0915
custom_llm_provider = "anthropic_text"
else:
custom_llm_provider = "anthropic"
## anthropic - pattern-based matching for future Claude models
elif _matches_claude_model_pattern(model):
custom_llm_provider = "anthropic"
## cohere
elif model in litellm.cohere_models or model in litellm.cohere_embedding_models:
custom_llm_provider = "cohere"
@ -931,6 +957,13 @@ def _get_openai_compatible_provider_info( # noqa: PLR0915
) = litellm.LambdaAIChatConfig()._get_openai_compatible_provider_info(
api_base, api_key
)
elif custom_llm_provider == "inception":
(
api_base,
dynamic_api_key,
) = litellm.InceptionChatConfig()._get_openai_compatible_provider_info(
api_base, api_key
)
elif custom_llm_provider == "hyperbolic":
(
api_base,

View file

@ -5300,8 +5300,12 @@ class StandardLoggingPayloadSetup:
tb_lines[:MAXIMUM_TRACEBACK_LINES_TO_LOG]
) # Limit to first 100 lines
# Get additional error details
error_message = str(original_exception)
explicit_message = getattr(original_exception, "message", None)
error_message = (
explicit_message
if isinstance(explicit_message, str) and explicit_message
else str(original_exception)
)
return StandardLoggingPayloadErrorInformation(
error_code=error_status,

View file

@ -34,6 +34,14 @@ _IMAGE_RESPONSE_CALL_TYPES = frozenset(
_VALID_DATA_RESIDENCIES = frozenset(r.value for r in DataResidency)
def _get_token_detail_value(details: object, key: str) -> Optional[int]:
if isinstance(details, dict):
value = details.get(key)
else:
value = getattr(details, key, None)
return value if isinstance(value, int) else None
def _is_above_128k(tokens: float) -> bool:
if tokens > 128000:
return True
@ -870,17 +878,47 @@ def calculate_image_response_cost_from_usage(
cached_tokens=0,
)
output_tokens_details = getattr(usage, "completion_tokens_details", None)
if output_tokens_details is None:
output_tokens_details = getattr(usage, "output_tokens_details", None)
if output_tokens_details is None:
completion_tokens_details = CompletionTokensDetailsWrapper(
text_tokens=0,
image_tokens=completion_tokens,
reasoning_tokens=0,
audio_tokens=0,
)
else:
text_tokens = _get_token_detail_value(output_tokens_details, "text_tokens") or 0
image_tokens = (
_get_token_detail_value(output_tokens_details, "image_tokens") or 0
)
audio_tokens = (
_get_token_detail_value(output_tokens_details, "audio_tokens") or 0
)
reasoning_tokens = (
_get_token_detail_value(output_tokens_details, "reasoning_tokens") or 0
)
known_output_tokens = (
text_tokens + image_tokens + audio_tokens + reasoning_tokens
)
if completion_tokens > known_output_tokens:
text_tokens += completion_tokens - known_output_tokens
completion_tokens_details = CompletionTokensDetailsWrapper(
text_tokens=text_tokens,
image_tokens=image_tokens,
reasoning_tokens=reasoning_tokens,
audio_tokens=audio_tokens,
)
normalized_usage = Usage(
prompt_tokens=prompt_tokens,
completion_tokens=completion_tokens,
total_tokens=total_tokens,
prompt_tokens_details=prompt_tokens_details,
completion_tokens_details=CompletionTokensDetailsWrapper(
text_tokens=0,
image_tokens=completion_tokens,
reasoning_tokens=0,
audio_tokens=0,
),
completion_tokens_details=completion_tokens_details,
)
prompt_cost, completion_cost = generic_cost_per_token(

View file

@ -3,6 +3,7 @@
import asyncio
import contextvars
import logging
from typing import Coroutine, Optional
import atexit
from typing_extensions import TypedDict
@ -494,31 +495,43 @@ class LoggingWorker:
processed = 0
start_time = loop.time()
while not self._queue.empty() and processed < MAX_ITERATIONS_TO_CLEAR_QUEUE:
if loop.time() - start_time >= MAX_TIME_TO_CLEAR_QUEUE:
self._safe_log(
"warning",
f"[LoggingWorker] atexit: Reached time limit ({MAX_TIME_TO_CLEAR_QUEUE}s), stopping flush",
)
break
# logging.raiseExceptions is a process-wide global; scope the
# suppression to just the drain loop, where shutdown callbacks may
# log to already-closed handler streams, so other threads keep their
# logging error reporting for as little of the window as possible.
previous_raise_exceptions = logging.raiseExceptions
logging.raiseExceptions = False
try:
while (
not self._queue.empty()
and processed < MAX_ITERATIONS_TO_CLEAR_QUEUE
):
if loop.time() - start_time >= MAX_TIME_TO_CLEAR_QUEUE:
self._safe_log(
"warning",
f"[LoggingWorker] atexit: Reached time limit ({MAX_TIME_TO_CLEAR_QUEUE}s), stopping flush",
)
break
try:
task = self._queue.get_nowait()
except asyncio.QueueEmpty:
break
try:
task = self._queue.get_nowait()
except asyncio.QueueEmpty:
break
# Run the coroutine synchronously in new loop
# Note: We run the coroutine directly, not via create_task,
# since we're in a new event loop context
try:
loop.run_until_complete(task["coroutine"])
processed += 1
except Exception:
# Silent failure to not break user's program
pass
finally:
# Clear reference to prevent memory leaks
task = None
# Run the coroutine synchronously in new loop
# Note: We run the coroutine directly, not via create_task,
# since we're in a new event loop context
try:
loop.run_until_complete(task["coroutine"])
processed += 1
except Exception:
# Silent failure to not break user's program
pass
finally:
# Clear reference to prevent memory leaks
task = None
finally:
logging.raiseExceptions = previous_raise_exceptions
self._safe_log(
"info",

View file

@ -1670,15 +1670,15 @@ def convert_to_gemini_tool_call_result( # noqa: PLR0915
if gemini_call_id:
_function_response["id"] = gemini_call_id
# Create part with function_response, and optionally inline_data for images (Computer Use)
_part: VertexPartType = {"function_response": _function_response}
# For Computer Use, if we have images/files, we need separate parts:
# - One part with function_response
# - One part per inline_data item
# Gemini's PartType is a oneof, so we can't have both in the same part
# For multimodal function responses, Gemini expects media parts nested
# inside functionResponse.parts instead of sibling content parts.
if inline_data_list:
return [_part] + [{"inline_data": d} for d in inline_data_list]
_function_response["parts"] = [
{"inline_data": inline_data} for inline_data in inline_data_list
]
return [_part]
return _part

View file

@ -6,10 +6,16 @@ from pydantic import BaseModel
from litellm.constants import DEFAULT_MAX_RECURSE_DEPTH
def strip_null_bytes(value: str) -> str:
"""Strip NUL bytes, which PostgreSQL text/jsonb columns reject (error 22P05)."""
return value.replace("\x00", "")
def safe_dumps(data: Any, max_depth: int = DEFAULT_MAX_RECURSE_DEPTH) -> str:
"""
Recursively serialize data while detecting circular references.
If a circular reference is detected then a marker string is returned.
NUL bytes are stripped from strings to prevent PostgreSQL 22P05 errors.
"""
def _serialize(obj: Any, seen: set, depth: int) -> Any:
@ -17,7 +23,9 @@ def safe_dumps(data: Any, max_depth: int = DEFAULT_MAX_RECURSE_DEPTH) -> str:
if depth > max_depth:
return "MaxDepthExceeded"
# Base-case: if it is a primitive, simply return it.
if isinstance(obj, (str, int, float, bool, type(None))):
if isinstance(obj, str):
return strip_null_bytes(obj)
if isinstance(obj, (int, float, bool, type(None))):
return obj
# Check for circular reference.
if id(obj) in seen:
@ -28,7 +36,7 @@ def safe_dumps(data: Any, max_depth: int = DEFAULT_MAX_RECURSE_DEPTH) -> str:
result = {}
for k, v in obj.items():
if isinstance(k, (str)):
result[k] = _serialize(v, seen, depth + 1)
result[strip_null_bytes(k)] = _serialize(v, seen, depth + 1)
seen.remove(id(obj))
return result
elif isinstance(obj, list):
@ -51,7 +59,7 @@ def safe_dumps(data: Any, max_depth: int = DEFAULT_MAX_RECURSE_DEPTH) -> str:
else:
# Fall back to string conversion for non-serializable objects.
try:
return str(obj)
return strip_null_bytes(str(obj))
except Exception:
return "Unserializable Object"

View file

@ -256,7 +256,10 @@ class LiteLLMAiohttpTransport(AiohttpTransport):
from yarl import URL as YarlURL
try:
data = request.content
# Coerce an empty body to None so aiohttp does not attach a
# `Content-Type: application/octet-stream` header for bodyless
# requests (e.g. DELETE /responses/{id}), which upstream APIs reject.
data = request.content or None
except httpx.RequestNotRead:
data = request.stream # type: ignore
request.headers.pop("transfer-encoding", None) # handled by aiohttp

View file

@ -1,6 +1,8 @@
import base64
import datetime
from typing import Any, Dict, List, Optional, Union
import json
import math
from typing import Any, Dict, List, Optional, Sequence, Union
import httpx
@ -12,6 +14,245 @@ from litellm.secret_managers.main import get_secret_str
from litellm.types.llms.openai import AllMessageValues
from litellm.types.utils import TokenCountResponse
GEMINI_IMAGE_ASPECT_RATIOS: Dict[str, float] = {
"1:1": 1 / 1,
"1:4": 1 / 4,
"1:8": 1 / 8,
"2:3": 2 / 3,
"3:2": 3 / 2,
"3:4": 3 / 4,
"4:1": 4 / 1,
"4:3": 4 / 3,
"4:5": 4 / 5,
"5:4": 5 / 4,
"8:1": 8 / 1,
"9:16": 9 / 16,
"16:9": 16 / 9,
"21:9": 21 / 9,
}
# Supported aspect ratio dimensions from Google Gemini image generation docs:
# https://ai.google.dev/gemini-api/docs/image-generation#aspect_ratios_and_image_size
GEMINI_IMAGE_SIZE_TO_ASPECT_RATIO: Dict[tuple[int, int], str] = {
(512, 512): "1:1",
(1024, 1024): "1:1",
(2048, 2048): "1:1",
(4096, 4096): "1:1",
(256, 1024): "1:4",
(512, 2048): "1:4",
(1024, 4096): "1:4",
(2048, 8192): "1:4",
(192, 1536): "1:8",
(384, 3072): "1:8",
(768, 6144): "1:8",
(1536, 12288): "1:8",
(424, 632): "2:3",
(848, 1264): "2:3",
(1696, 2528): "2:3",
(3392, 5056): "2:3",
(632, 424): "3:2",
(1264, 848): "3:2",
(2528, 1696): "3:2",
(5056, 3392): "3:2",
(448, 600): "3:4",
(896, 1200): "3:4",
(1792, 2400): "3:4",
(3584, 4800): "3:4",
(1024, 256): "4:1",
(2048, 512): "4:1",
(4096, 1024): "4:1",
(8192, 2048): "4:1",
(600, 448): "4:3",
(1200, 896): "4:3",
(2400, 1792): "4:3",
(4800, 3584): "4:3",
(464, 576): "4:5",
(928, 1152): "4:5",
(1856, 2304): "4:5",
(3712, 4608): "4:5",
(576, 464): "5:4",
(1152, 928): "5:4",
(2304, 1856): "5:4",
(4608, 3712): "5:4",
(1536, 192): "8:1",
(3072, 384): "8:1",
(6144, 768): "8:1",
(12288, 1536): "8:1",
(384, 688): "9:16",
(768, 1376): "9:16",
(1536, 2752): "9:16",
(3072, 5504): "9:16",
(688, 384): "16:9",
(1376, 768): "16:9",
(2752, 1536): "16:9",
(5504, 3072): "16:9",
(792, 336): "21:9",
(1584, 672): "21:9",
(3168, 1344): "21:9",
(6336, 2688): "21:9",
(1280, 896): "4:3",
(896, 1280): "3:4",
}
def map_openai_size_to_gemini_image_config(
size: str, model: str
) -> Optional[Dict[str, str]]:
dimensions = _parse_openai_image_size(size)
if dimensions is None:
return None
width, height = dimensions
image_config = {
"aspectRatio": _map_dimensions_to_gemini_aspect_ratio(width, height)
}
image_size = _map_dimensions_to_gemini_image_size(width, height)
if is_gemini_image_model(model):
if supports_gemini_image_size(model):
image_config["imageSize"] = image_size
else:
image_config["imageSize"] = image_size
return image_config
def supports_gemini_image_size(model: str) -> bool:
try:
model_info = litellm.get_model_info(model=model)
value = model_info.get("supports_image_size")
if value is not None:
return bool(value)
except Exception:
pass
return "2.5-flash" not in model
def is_gemini_image_model(model: str) -> bool:
base_model = model.split("/", 1)[-1]
return "gemini" in base_model
def map_openai_image_params_to_gemini(
params: Dict[str, Any],
model: str,
supported_params: Sequence[str],
optional_params: Optional[Dict[str, Any]] = None,
parse_image_config_string: bool = False,
) -> Dict[str, Any]:
optional_params = optional_params or {}
filtered_params = {
key: value for key, value in params.items() if key in supported_params
}
mapped_params: Dict[str, Any] = {}
if "n" in filtered_params and "n" not in optional_params:
mapped_params["sampleCount"] = filtered_params["n"]
if "size" in filtered_params and "size" not in optional_params:
image_config = map_openai_size_to_gemini_image_config(
filtered_params["size"],
model,
)
if image_config is not None:
if is_gemini_image_model(model):
mapped_params["imageConfig"] = image_config
else:
mapped_params["aspectRatio"] = image_config["aspectRatio"]
if "imageSize" in image_config:
mapped_params["imageSize"] = image_config["imageSize"]
image_config_param = filtered_params.get("imageConfig")
if isinstance(image_config_param, str) and parse_image_config_string:
try:
image_config_param = json.loads(image_config_param)
except json.JSONDecodeError as exc:
raise litellm.UnsupportedParamsError(
model=model,
message="`imageConfig` must be valid JSON when provided as a string.",
) from exc
if isinstance(image_config_param, dict):
mapped_params["imageConfig"] = image_config_param
for key, value in filtered_params.items():
if key not in ("n", "size", "imageConfig") and key not in optional_params:
mapped_params[key] = value
return mapped_params
def get_gemini_image_generation_config(
model: str,
optional_params: Dict[str, Any],
) -> Dict[str, Any]:
generation_config: Dict[str, Any] = {"response_modalities": ["IMAGE", "TEXT"]}
image_config: Dict[str, Any] = {}
if isinstance(optional_params.get("imageConfig"), dict):
image_config.update(optional_params["imageConfig"])
if not supports_gemini_image_size(model):
image_config.pop("imageSize", None)
if image_config:
generation_config["imageConfig"] = image_config
candidate_count = next(
(
optional_params[key]
for key in ("candidateCount", "candidate_count", "sampleCount", "n")
if optional_params.get(key) is not None
),
None,
)
if candidate_count is not None:
generation_config["candidateCount"] = candidate_count
return generation_config
def _parse_openai_image_size(size: str) -> Optional[tuple[int, int]]:
if size == "auto":
return None
width_str, separator, height_str = size.lower().partition("x")
if not separator:
return None
try:
width = int(width_str)
height = int(height_str)
except ValueError:
return None
if width <= 0 or height <= 0:
return None
return width, height
def _map_dimensions_to_gemini_aspect_ratio(width: int, height: int) -> str:
if (width, height) in GEMINI_IMAGE_SIZE_TO_ASPECT_RATIO:
return GEMINI_IMAGE_SIZE_TO_ASPECT_RATIO[(width, height)]
requested_ratio = width / height
return min(
GEMINI_IMAGE_ASPECT_RATIOS,
key=lambda aspect_ratio: abs(
math.log(GEMINI_IMAGE_ASPECT_RATIOS[aspect_ratio] / requested_ratio)
),
)
def _map_dimensions_to_gemini_image_size(width: int, height: int) -> str:
effective_square_side = math.sqrt(width * height)
if effective_square_side < 768:
return "512"
if effective_square_side < 1536:
return "1K"
if effective_square_side < 3072:
return "2K"
return "4K"
class GeminiError(BaseLLMException):
pass

View file

@ -4,8 +4,9 @@ Gemini Image Edit Cost Calculator
from typing import Any
import litellm
from litellm.types.utils import ImageResponse
from litellm.llms.gemini.image_generation.cost_calculator import (
cost_calculator as image_generation_cost_calculator,
)
def cost_calculator(
@ -15,20 +16,10 @@ def cost_calculator(
"""
Gemini image edit cost calculator.
Mirrors image generation pricing: charge per returned image based on
model metadata (`output_cost_per_image`).
Gemini image edits and generations share image response billing behavior:
use provider token usage when present, otherwise fall back to per-image pricing.
"""
model_info = litellm.get_model_info(
return image_generation_cost_calculator(
model=model,
custom_llm_provider="gemini",
image_response=image_response,
)
output_cost_per_image: float = model_info.get("output_cost_per_image") or 0.0
if not isinstance(image_response, ImageResponse):
raise ValueError(
f"image_response must be of type ImageResponse got type={type(image_response)}"
)
num_images = len(image_response.data or [])
return output_cost_per_image * num_images

View file

@ -7,10 +7,22 @@ from httpx._types import RequestFiles
from litellm.images.utils import ImageEditRequestUtils
from litellm.llms.base_llm.image_edit.transformation import BaseImageEditConfig
from litellm.llms.gemini.common_utils import (
get_gemini_image_generation_config,
map_openai_image_params_to_gemini,
)
from litellm.llms.gemini.image_usage_transformation import (
transform_gemini_image_usage,
)
from litellm.secret_managers.main import get_secret_str
from litellm.types.images.main import ImageEditOptionalRequestParams
from litellm.types.router import GenericLiteLLMParams
from litellm.types.utils import FileTypes, ImageObject, ImageResponse, OpenAIImage
from litellm.types.utils import (
FileTypes,
ImageObject,
ImageResponse,
OpenAIImage,
)
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
@ -22,7 +34,7 @@ else:
class GeminiImageEditConfig(BaseImageEditConfig):
DEFAULT_BASE_URL: str = "https://generativelanguage.googleapis.com/v1beta"
SUPPORTED_PARAMS: List[str] = ["size"]
SUPPORTED_PARAMS: List[str] = ["n", "size", "imageConfig"]
def get_supported_openai_params(self, model: str) -> List[str]:
return list(self.SUPPORTED_PARAMS)
@ -33,21 +45,12 @@ class GeminiImageEditConfig(BaseImageEditConfig):
model: str,
drop_params: bool,
) -> Dict[str, Any]:
supported_params = self.get_supported_openai_params(model)
filtered_params = {
key: value
for key, value in image_edit_optional_params.items()
if key in supported_params
}
mapped_params: Dict[str, Any] = {}
if "size" in filtered_params:
mapped_params["aspectRatio"] = self._map_size_to_aspect_ratio(
filtered_params["size"] # type: ignore[arg-type]
)
return mapped_params
return map_openai_image_params_to_gemini(
params=image_edit_optional_params, # type: ignore[arg-type]
model=model,
supported_params=self.get_supported_openai_params(model),
parse_image_config_string=True,
)
def validate_environment(
self,
@ -107,18 +110,10 @@ class GeminiImageEditConfig(BaseImageEditConfig):
request_body: Dict[str, Any] = {"contents": contents}
generation_config: Dict[str, Any] = {}
if "aspectRatio" in image_edit_optional_request_params:
# Move aspectRatio into imageConfig inside generationConfig
if "imageConfig" not in generation_config:
generation_config["imageConfig"] = {}
generation_config["imageConfig"]["aspectRatio"] = (
image_edit_optional_request_params["aspectRatio"]
)
if generation_config:
request_body["generationConfig"] = generation_config
request_body["generationConfig"] = get_gemini_image_generation_config(
model=model,
optional_params=image_edit_optional_request_params,
)
empty_files = cast(RequestFiles, [])
return request_body, empty_files
@ -156,18 +151,12 @@ class GeminiImageEditConfig(BaseImageEditConfig):
)
model_response.data = cast(List[OpenAIImage], data_list)
if "usageMetadata" in response_json:
model_response.usage = transform_gemini_image_usage(
response_json["usageMetadata"]
)
return model_response
def _map_size_to_aspect_ratio(self, size: str) -> str:
aspect_ratio_map = {
"1024x1024": "1:1",
"1792x1024": "16:9",
"1024x1792": "9:16",
"1280x896": "4:3",
"896x1280": "3:4",
}
return aspect_ratio_map.get(size, "1:1")
def _prepare_inline_image_parts(
self, image: Union[FileTypes, List[FileTypes]]
) -> List[Dict[str, Any]]:

View file

@ -5,18 +5,21 @@ import httpx
from litellm.llms.base_llm.image_generation.transformation import (
BaseImageGenerationConfig,
)
from litellm.llms.gemini.common_utils import (
get_gemini_image_generation_config,
is_gemini_image_model,
map_openai_image_params_to_gemini,
)
from litellm.llms.gemini.image_usage_transformation import (
transform_gemini_image_usage,
)
from litellm.secret_managers.main import get_secret_str
from litellm.types.llms.gemini import GeminiImageGenerationRequest
from litellm.types.llms.openai import (
AllMessageValues,
OpenAIImageGenerationOptionalParams,
)
from litellm.types.utils import (
ImageObject,
ImageResponse,
ImageUsage,
ImageUsageInputTokensDetails,
)
from litellm.types.utils import ImageObject, ImageResponse
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
@ -36,7 +39,10 @@ class GoogleImageGenConfig(BaseImageGenerationConfig):
Google AI Imagen API supported parameters
https://ai.google.dev/gemini-api/docs/imagen
"""
return ["n", "size"]
supported_params = ["n", "size"]
if is_gemini_image_model(model):
supported_params.append("imageConfig")
return supported_params # type: ignore[return-value]
def map_openai_params(
self,
@ -45,64 +51,11 @@ class GoogleImageGenConfig(BaseImageGenerationConfig):
model: str,
drop_params: bool,
) -> dict:
supported_params = self.get_supported_openai_params(model)
mapped_params = {}
for k, v in non_default_params.items():
if k not in optional_params.keys():
if k in supported_params:
# Map OpenAI parameters to Google format
if k == "n":
mapped_params["sampleCount"] = v
elif k == "size":
# Map OpenAI size format to Google aspectRatio
mapped_params["aspectRatio"] = self._map_size_to_aspect_ratio(v)
else:
mapped_params[k] = v
return mapped_params
def _map_size_to_aspect_ratio(self, size: str) -> str:
"""
https://ai.google.dev/gemini-api/docs/image-generation
"""
aspect_ratio_map = {
"1024x1024": "1:1",
"1792x1024": "16:9",
"1024x1792": "9:16",
"1280x896": "4:3",
"896x1280": "3:4",
}
return aspect_ratio_map.get(size, "1:1")
def _transform_image_usage(self, usage_metadata: dict) -> ImageUsage:
"""
Transform Gemini usageMetadata to ImageUsage format
"""
input_tokens_details = ImageUsageInputTokensDetails(
image_tokens=0,
text_tokens=0,
)
# Extract detailed token counts from promptTokensDetails
tokens_details = usage_metadata.get("promptTokensDetails", [])
for details in tokens_details:
if isinstance(details, dict):
modality = str(details.get("modality", "")).upper()
raw_token_count = details.get(
"tokenCount", details.get("token_count", 0)
)
token_count = raw_token_count if isinstance(raw_token_count, int) else 0
if modality == "TEXT":
input_tokens_details.text_tokens += token_count
elif modality == "IMAGE":
input_tokens_details.image_tokens += token_count
return ImageUsage(
input_tokens=usage_metadata.get("promptTokenCount", 0),
input_tokens_details=input_tokens_details,
output_tokens=usage_metadata.get("candidatesTokenCount", 0),
total_tokens=usage_metadata.get("totalTokenCount", 0),
return map_openai_image_params_to_gemini(
params=non_default_params,
model=model,
supported_params=self.get_supported_openai_params(model),
optional_params=optional_params,
)
def get_complete_url(
@ -127,7 +80,7 @@ class GoogleImageGenConfig(BaseImageGenerationConfig):
complete_url = complete_url.rstrip("/")
# Gemini Flash Image Preview models use generateContent endpoint
if "gemini" in model:
if is_gemini_image_model(model):
complete_url = f"{complete_url}/models/{model}:generateContent"
else:
# All other Imagen models use predict endpoint
@ -179,10 +132,13 @@ class GoogleImageGenConfig(BaseImageGenerationConfig):
}
"""
# For Gemini Flash Image Preview models, use standard Gemini format
if "gemini" in model:
if is_gemini_image_model(model):
request_body: dict = {
"contents": [{"parts": [{"text": prompt}]}],
"generationConfig": {"response_modalities": ["IMAGE", "TEXT"]},
"generationConfig": get_gemini_image_generation_config(
model=model,
optional_params=optional_params,
),
}
return request_body
else:
@ -200,6 +156,9 @@ class GoogleImageGenConfig(BaseImageGenerationConfig):
)
return request_body_obj.model_dump(exclude_none=True)
def _transform_image_usage(self, usage_metadata: dict):
return transform_gemini_image_usage(usage_metadata)
def transform_image_generation_response(
self,
model: str,
@ -229,7 +188,7 @@ class GoogleImageGenConfig(BaseImageGenerationConfig):
model_response.data = []
# Handle different response formats based on model
if "gemini" in model:
if is_gemini_image_model(model):
# Gemini Flash Image Preview models return in candidates format
candidates = response_data.get("candidates", [])
for candidate in candidates:
@ -255,7 +214,7 @@ class GoogleImageGenConfig(BaseImageGenerationConfig):
# Extract usage metadata for Gemini models
if "usageMetadata" in response_data:
model_response.usage = self._transform_image_usage(
model_response.usage = transform_gemini_image_usage(
response_data["usageMetadata"]
)
else:

View file

@ -0,0 +1,73 @@
from typing import Any
from litellm.types.utils import ImageUsage, ImageUsageInputTokensDetails
def _get_token_count(details: dict) -> int:
raw_token_count = details.get("tokenCount", details.get("token_count", 0))
return raw_token_count if isinstance(raw_token_count, int) else 0
def _get_modality_token_details(usage_metadata: dict, *details_keys: str) -> list:
for details_key in details_keys:
details = usage_metadata.get(details_key)
if isinstance(details, list):
return details
return []
def _sum_modality_token_details(
usage_metadata: dict, *details_keys: str
) -> ImageUsageInputTokensDetails:
tokens_details = ImageUsageInputTokensDetails(
image_tokens=0,
text_tokens=0,
)
for details in _get_modality_token_details(usage_metadata, *details_keys):
if isinstance(details, dict):
modality = str(details.get("modality", "")).upper()
token_count = _get_token_count(details)
if modality == "TEXT":
tokens_details.text_tokens += token_count
elif modality == "IMAGE":
tokens_details.image_tokens += token_count
return tokens_details
def transform_gemini_image_usage(usage_metadata: dict) -> ImageUsage:
"""
Transform Gemini usageMetadata to ImageUsage format.
"""
input_tokens_details = _sum_modality_token_details(
usage_metadata, "promptTokensDetails", "prompt_tokens_details"
)
output_tokens = usage_metadata.get("candidatesTokenCount", 0)
output_tokens_details = _sum_modality_token_details(
usage_metadata, "candidatesTokensDetails", "candidates_tokens_details"
)
if not _get_modality_token_details(
usage_metadata, "candidatesTokensDetails", "candidates_tokens_details"
):
output_tokens_details.image_tokens = output_tokens
else:
known_output_tokens = (
output_tokens_details.text_tokens + output_tokens_details.image_tokens
)
if output_tokens > known_output_tokens:
output_tokens_details.text_tokens += output_tokens - known_output_tokens
usage_payload: dict[str, Any] = {
"input_tokens": usage_metadata.get("promptTokenCount", 0),
"input_tokens_details": input_tokens_details,
"output_tokens": output_tokens,
"total_tokens": usage_metadata.get("totalTokenCount", 0),
"prompt_tokens": usage_metadata.get("promptTokenCount", 0),
"prompt_tokens_details": input_tokens_details.model_dump(),
"completion_tokens": output_tokens,
"completion_tokens_details": output_tokens_details.model_dump(),
"output_tokens_details": output_tokens_details.model_dump(),
}
return ImageUsage(**usage_payload)

View file

View file

View file

@ -0,0 +1,54 @@
"""
Translate from OpenAI's `/v1/chat/completions` to Inception's `/v1/chat/completions`
Inception Labs (https://www.inceptionlabs.ai) serves the Mercury family of
diffusion LLMs through an OpenAI-compatible API, so we only need to point the
OpenAI-like handler at the Inception API base and pick up the Inception API key.
"""
from typing import List, Optional, Tuple
import litellm
from litellm.secret_managers.main import get_secret_str
from ...openai_like.chat.transformation import OpenAILikeChatConfig
class InceptionChatConfig(OpenAILikeChatConfig):
"""
Inception is OpenAI-compatible with standard endpoints
"""
@property
def custom_llm_provider(self) -> Optional[str]:
return "inception"
def get_supported_openai_params(self, model: str) -> List:
return [
"max_tokens",
"max_completion_tokens",
"temperature",
"stop",
"tools",
"tool_choice",
"stream",
"stream_options",
"response_format",
"reasoning_effort",
"reasoning_summary",
"reasoning_summary_wait",
"diffusing",
"realtime",
]
def _get_openai_compatible_provider_info(
self, api_base: Optional[str], api_key: Optional[str]
) -> Tuple[Optional[str], Optional[str]]:
passed_api_base = api_base
api_base = api_base or get_secret_str("INCEPTION_API_BASE") or "https://api.inceptionlabs.ai/v1" # type: ignore
dynamic_api_key = api_key
if passed_api_base is None or api_key:
dynamic_api_key = (
api_key or litellm.inception_key or get_secret_str("INCEPTION_API_KEY")
)
return api_base, dynamic_api_key

View file

@ -0,0 +1,43 @@
"""
Inception fill-in-the-middle (FIM) completions.
Inception's FIM endpoint is OpenAI text-completion compatible: it takes a
`prompt` (prefix) plus an optional `suffix` and returns standard
`choices[].text`. It is served at `/v1/fim/completions` rather than
`/v1/completions`, so routing points the OpenAI client at the `/v1/fim` base
(see the `text-completion-inception` branch in `main.py`).
"""
from typing import List
from litellm.llms.openai.completion.transformation import OpenAITextCompletionConfig
class InceptionTextCompletionConfig(OpenAITextCompletionConfig):
def get_supported_openai_params(self, model: str) -> List:
return [
"suffix",
"max_tokens",
"max_completion_tokens",
"top_p",
"frequency_penalty",
"presence_penalty",
"stop",
"stream",
"stream_options",
]
def map_openai_params(
self,
non_default_params: dict,
optional_params: dict,
model: str,
drop_params: bool,
) -> dict:
supported_params = self.get_supported_openai_params(model)
for param, value in non_default_params.items():
if param == "max_completion_tokens":
optional_params["max_tokens"] = value
elif param in supported_params:
optional_params[param] = value
return optional_params

View file

@ -0,0 +1 @@
"""LangFlow LLM provider for LiteLLM."""

View file

@ -0,0 +1,37 @@
import hashlib
from typing import Any, Dict, Optional
def get_session_id_from_a2a_params(params: Dict[str, Any]) -> Optional[str]:
message = params.get("message", {})
if isinstance(message, dict):
return message.get("contextId")
return getattr(message, "contextId", None)
def scope_session_to_principal(session_id: str, principal: Optional[str]) -> str:
"""
Bind a client-supplied A2A contextId to the authenticated principal.
Without this, two distinct keys authorized for the same LangFlow agent could
set the same contextId and read/append to each other's LangFlow memory. The
principal is hashed (it is already a hashed token) so the raw value is never
sent to the LangFlow backend, while the original contextId is kept as a
suffix for operator-side correlation.
"""
if not principal:
return session_id
principal_prefix = hashlib.sha256(principal.encode("utf-8")).hexdigest()[:16]
return f"{principal_prefix}-{session_id}"
def merge_a2a_session_into_litellm_params(
litellm_params: Dict[str, Any],
params: Dict[str, Any],
principal: Optional[str] = None,
) -> Dict[str, Any]:
merged = dict(litellm_params)
session_id = get_session_id_from_a2a_params(params)
if session_id and "session_id" not in merged:
merged["session_id"] = scope_session_to_principal(session_id, principal)
return merged

View file

@ -0,0 +1 @@
"""LangFlow chat transformation."""

View file

@ -0,0 +1,327 @@
"""LangFlow run API: POST {api_base}/api/v1/run/{flow_id}"""
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union
from urllib.parse import quote
import httpx
from litellm._logging import verbose_logger
from litellm.litellm_core_utils.prompt_templates.common_utils import (
convert_content_list_to_str,
)
from litellm.llms.base_llm.chat.transformation import BaseConfig, BaseLLMException
from litellm.types.llms.openai import AllMessageValues
from litellm.types.utils import Choices, Message, ModelResponse, Usage
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
from litellm.utils import CustomStreamWrapper
LiteLLMLoggingObj = _LiteLLMLoggingObj
else:
LiteLLMLoggingObj = Any
HTTPHandler = Any
AsyncHTTPHandler = Any
CustomStreamWrapper = Any
class LangFlowError(BaseLLMException):
"""Exception class for LangFlow API errors."""
pass
class LangFlowConfig(BaseConfig):
"""
Configuration for the LangFlow API.
LangFlow is a visual, low-code platform for building AI agents and pipelines.
Each flow has a unique flow_id and is invoked via a simple HTTP endpoint.
"""
def __init__(self, **kwargs):
super().__init__(**kwargs)
def _get_openai_compatible_provider_info(
self,
api_base: Optional[str],
api_key: Optional[str],
) -> Tuple[Optional[str], Optional[str]]:
from litellm.secret_managers.main import get_secret_str
api_base = (
api_base or get_secret_str("LANGFLOW_API_BASE") or "http://localhost:7860"
)
api_key = api_key or get_secret_str("LANGFLOW_API_KEY")
return api_base, api_key
def get_supported_openai_params(self, model: str) -> List[str]:
return ["stream"]
def map_openai_params(
self,
non_default_params: dict,
optional_params: dict,
model: str,
drop_params: bool,
) -> dict:
return optional_params
def _get_flow_id(self, model: str, optional_params: dict) -> str:
"""
Extract flow_id from the authorized model name only.
Model format: "langflow/{flow_id}". Request kwargs must not override
flow_id (would allow calling another flow with the same API key).
"""
if optional_params.get("flow_id") is not None:
raise LangFlowError(
status_code=400,
message=(
"flow_id cannot be set via request parameters; "
"use model langflow/{flow_id}"
),
)
flow_id = (model.split("/", 1)[1] if "/" in model else model).strip()
if not flow_id:
raise LangFlowError(
status_code=400,
message="flow_id is required; use model langflow/{flow_id}",
)
return flow_id
def get_complete_url(
self,
api_base: Optional[str],
api_key: Optional[str],
model: str,
optional_params: dict,
litellm_params: dict,
stream: Optional[bool] = None,
) -> str:
if api_base is None:
raise ValueError(
"api_base is required for LangFlow. Set it via LANGFLOW_API_BASE env var or api_base parameter."
)
api_base = api_base.rstrip("/")
flow_id = quote(self._get_flow_id(model, optional_params), safe="")
return f"{api_base}/api/v1/run/{flow_id}"
def _get_last_user_message(self, messages: List[AllMessageValues]) -> str:
"""Extract the text of the last user message to use as input_value."""
for msg in reversed(messages):
if msg.get("role") == "user":
content = msg.get("content", "")
if isinstance(content, list):
content = convert_content_list_to_str(msg)
if not isinstance(content, str):
content = str(content)
return content
# Fallback: use last message regardless of role
if messages:
content = messages[-1].get("content", "")
if isinstance(content, list):
content = convert_content_list_to_str(messages[-1])
if not isinstance(content, str):
content = str(content)
return content
return ""
def _reject_caller_tweaks(self, params: dict) -> None:
if params.get("tweaks") is not None:
raise LangFlowError(
status_code=400,
message=(
"tweaks cannot be set via request parameters; they would "
"override the operator-configured LangFlow flow components"
),
)
def transform_request(
self,
model: str,
messages: List[AllMessageValues],
optional_params: dict,
litellm_params: dict,
headers: dict,
) -> dict:
"""
Transform the request to LangFlow format.
LangFlow request format:
{
"input_value": "<last user message>",
"input_type": "chat",
"output_type": "chat",
"session_id": "<session_id>"
}
"""
self._reject_caller_tweaks(optional_params)
input_value = self._get_last_user_message(messages)
payload: Dict[str, Any] = {
"input_value": input_value,
"input_type": optional_params.get("input_type", "chat"),
"output_type": optional_params.get("output_type", "chat"),
}
session_id = optional_params.get("session_id")
if session_id:
payload["session_id"] = session_id
verbose_logger.debug(f"LangFlow request payload: {payload}")
return payload
def _extract_content_from_response(self, response_json: dict) -> Optional[str]:
"""
Extract the assistant text from a LangFlow run response.
Expected structure:
{"outputs": [{"outputs": [{"results": {"message": {"text": "..."}}}]}]}
Returns None when no message text is present so the caller can surface an
explicit error instead of forwarding a raw JSON blob as the answer.
"""
outputs = response_json.get("outputs", [])
if not (isinstance(outputs, list) and outputs):
return None
first_output = outputs[0]
if not isinstance(first_output, dict):
return None
inner_outputs = first_output.get("outputs", [])
if not (isinstance(inner_outputs, list) and inner_outputs):
return None
first_inner = inner_outputs[0]
if not isinstance(first_inner, dict):
return None
results = first_inner.get("results", {})
if isinstance(results, dict):
message = results.get("message", {})
if isinstance(message, dict) and message.get("text"):
return message["text"]
outputs_dict = first_inner.get("outputs", {})
if isinstance(outputs_dict, dict):
for val in outputs_dict.values():
if isinstance(val, dict):
msg = val.get("message", {})
if isinstance(msg, dict) and msg.get("text"):
return msg["text"]
return None
def transform_response(
self,
model: str,
raw_response: httpx.Response,
model_response: ModelResponse,
logging_obj: LiteLLMLoggingObj,
request_data: dict,
messages: List[AllMessageValues],
optional_params: dict,
litellm_params: dict,
encoding: Any,
api_key: Optional[str] = None,
json_mode: Optional[bool] = None,
) -> ModelResponse:
try:
response_json = raw_response.json()
except Exception as e:
raise LangFlowError(
message=f"LangFlow returned a non-JSON response: {e}",
status_code=raw_response.status_code,
)
verbose_logger.debug(f"LangFlow response: {response_json}")
content = self._extract_content_from_response(response_json)
if content is None:
raise LangFlowError(
message=(
"Could not extract a message from the LangFlow response; "
"ensure the flow ends in a Chat Output component"
),
status_code=500,
)
message = Message(content=content, role="assistant")
choice = Choices(finish_reason="stop", index=0, message=message)
model_response.choices = [choice]
model_response.model = model
try:
from litellm.utils import token_counter
prompt_tokens = token_counter(model=model, messages=messages)
completion_tokens = token_counter(
model=model, text=content, count_response_tokens=True
)
usage = Usage(
prompt_tokens=prompt_tokens,
completion_tokens=completion_tokens,
total_tokens=prompt_tokens + completion_tokens,
)
setattr(model_response, "usage", usage)
except Exception as e:
verbose_logger.warning(f"Failed to calculate token usage: {e}")
return model_response
def sign_request(
self,
headers: dict,
optional_params: dict,
request_data: dict,
api_base: str,
api_key: Optional[str] = None,
model: Optional[str] = None,
stream: Optional[bool] = None,
fake_stream: Optional[bool] = None,
) -> Tuple[dict, Optional[bytes]]:
self._reject_caller_tweaks(request_data)
return headers, None
def validate_environment(
self,
headers: dict,
model: str,
messages: List[AllMessageValues],
optional_params: dict,
litellm_params: dict,
api_key: Optional[str] = None,
api_base: Optional[str] = None,
) -> dict:
headers["Content-Type"] = "application/json"
if api_key:
headers["x-api-key"] = api_key
return headers
def get_error_class(
self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers]
) -> BaseLLMException:
return LangFlowError(status_code=status_code, message=error_message)
@property
def supports_stream_param_in_request_body(self) -> bool:
return False
def should_fake_stream(
self,
model: Optional[str],
stream: Optional[bool],
custom_llm_provider: Optional[str] = None,
) -> bool:
return stream is True

View file

@ -299,11 +299,8 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig):
or litellm.openai_key
or get_secret_str("OPENAI_API_KEY")
)
headers.update(
{
"Authorization": f"Bearer {api_key}",
}
)
headers.setdefault("Content-Type", "application/json")
headers["Authorization"] = f"Bearer {api_key}"
return headers
def get_complete_url(

View file

@ -996,7 +996,19 @@ def _gemini_convert_messages_with_history( # noqa: PLR0915
excluded_keys=["thoughtSignature"],
):
assistant_content.append(gemini_tool_call_part)
last_message_with_tool_calls = assistant_msg
# Only record this as the active tool-call message when it actually
# carries tool calls. The `if` guard above is also entered for a
# text-only assistant message (`assistant_msg.get("tool_calls", [])
# is not None` is True for an empty list), so without this check a
# later assistant message with no tool calls would clobber the
# reference. The following tool result would then be matched against
# an assistant message that has no tool_calls, raising "Missing
# corresponding tool call for tool response message".
if (
assistant_msg.get("tool_calls")
or assistant_msg.get("function_call") is not None
):
last_message_with_tool_calls = assistant_msg
## HANDLE SERVER-SIDE TOOL INVOCATIONS (context circulation)
_psf = assistant_msg.get("provider_specific_fields")

View file

@ -3,6 +3,7 @@ from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union
import httpx
from litellm import get_model_info
from litellm.exceptions import BadRequestError
from litellm.litellm_core_utils.url_utils import encode_url_path_segment
from litellm.llms.base_llm.vector_store.transformation import BaseVectorStoreConfig
from litellm.llms.vertex_ai.vertex_llm_base import VertexBase
@ -16,6 +17,8 @@ from litellm.types.vector_stores import (
VectorStoreSearchOptionalRequestParams,
VectorStoreSearchResponse,
VectorStoreSearchResult,
VertexSearchDataStoreExtraBody,
VertexSearchEngineExtraBody,
)
if TYPE_CHECKING:
@ -26,6 +29,31 @@ else:
LiteLLMLoggingObj = Any
# Fields that select which data store / serving config to search. These are
# always determined by the request URL path (vector_store_id / vertex_engine_id),
# so allowing them per request could silently redirect the search to a different
# target. Rejected in both data-store and engine/app modes.
VERTEX_SEARCH_TARGET_SELECTING_FIELDS = frozenset(
{
"branch",
"servingConfig",
"entity",
}
)
# Allowlists of native Discovery Engine SearchRequest fields callers may forward
# via extra_body, derived from the TypedDicts so the type is the source of truth.
# Engine/app mode is a superset (adds dataStoreSpecs, numResultsPerDataStore),
# since an app fans out across multiple member data stores.
VERTEX_SEARCH_DATASTORE_EXTRA_BODY_FIELDS = frozenset(
VertexSearchDataStoreExtraBody.__annotations__
)
VERTEX_SEARCH_ENGINE_EXTRA_BODY_FIELDS = frozenset(
VertexSearchEngineExtraBody.__annotations__
)
class VertexSearchAPIVectorStoreConfig(BaseVectorStoreConfig, VertexBase):
"""
Configuration for Vertex AI Search API Vector Store
@ -36,6 +64,66 @@ class VertexSearchAPIVectorStoreConfig(BaseVectorStoreConfig, VertexBase):
def __init__(self):
super().__init__()
@staticmethod
def get_supported_extra_body_fields(is_engine: bool = False) -> frozenset:
"""
Native SearchRequest fields callers may forward via ``extra_body``.
The set depends on which serving config the request targets:
- engine/app mode (``is_engine=True``): includes multi-store fields such
as ``dataStoreSpecs`` and ``numResultsPerDataStore``.
- data-store mode: the engine-only fields are excluded.
"""
if is_engine:
return VERTEX_SEARCH_ENGINE_EXTRA_BODY_FIELDS
return VERTEX_SEARCH_DATASTORE_EXTRA_BODY_FIELDS
@classmethod
def _filter_extra_body(
cls, extra_body: Dict[str, Any], is_engine: bool = False
) -> Dict[str, Any]:
"""
Validate ``extra_body`` against the supported-field allowlist for the
active serving config (engine/app vs data store).
Raises ``BadRequestError`` (HTTP 400) if the caller includes a
target-selecting field (e.g. ``servingConfig``) or any field not
supported for the active mode, so the request fails loudly instead of
silently searching the wrong target. Engine-only fields
(``dataStoreSpecs``, ``numResultsPerDataStore``) are rejected in
data-store mode where they are meaningless.
"""
supported = cls.get_supported_extra_body_fields(is_engine=is_engine)
filtered = {
key: value for key, value in extra_body.items() if value is not None
}
target_selecting = set(filtered) & VERTEX_SEARCH_TARGET_SELECTING_FIELDS
if target_selecting:
raise BadRequestError(
message=(
"Vertex AI Search extra_body may not set target-selecting fields "
f"{sorted(target_selecting)}: the data store is scoped by "
"vector_store_id / vertex_engine_id and cannot be overridden per request."
),
model="vertex_ai/search_api",
llm_provider="vertex_ai",
)
unsupported = set(filtered) - supported
if unsupported:
mode = "engine/app" if is_engine else "data store"
raise BadRequestError(
message=(
f"Unsupported Vertex AI Search extra_body fields {sorted(unsupported)} "
f"for {mode} mode. Supported fields: {sorted(supported)}."
),
model="vertex_ai/search_api",
llm_provider="vertex_ai",
)
return filtered
def get_auth_credentials(
self, litellm_params: dict
) -> BaseVectorStoreAuthCredentials:
@ -133,23 +221,41 @@ class VertexSearchAPIVectorStoreConfig(BaseVectorStoreConfig, VertexBase):
extra_body: Optional[Dict[str, Any]] = None,
) -> Tuple[str, Dict[str, Any]]:
"""
Transform search request for Vertex AI RAG API
Transform a search request for the Vertex AI Search (Discovery Engine) API.
Per-request params pass through to the engine: max_num_results maps to
pageSize, and extra_body fields on the supported allowlist
(`get_supported_extra_body_fields`) are merged in with precedence, so
callers can send native Discovery Engine tuning fields such as filter,
boostSpec, or contentSearchSpec.
The allowlist depends on the serving config: engine/app mode (when
`vertex_engine_id` is set) additionally accepts multi-store fields like
`dataStoreSpecs` and `numResultsPerDataStore`, while data-store mode
rejects them. Target-selecting fields (e.g. servingConfig, branch) are
rejected in both modes: the target is scoped by the URL path
(vector_store_id / vertex_engine_id) and must not be overridable per
request.
"""
# Convert query to string if it's a list
if isinstance(query, list):
query = " ".join(query)
# Vertex AI RAG API endpoint for retrieving contexts
url = f"{api_base}:search"
# Construct full rag corpus path
# Build the request body for Vertex AI Search API
request_body = {"query": query, "pageSize": 10}
is_engine = bool(litellm_params.get("vertex_engine_id"))
#########################################################
# Update logging object with details of the request
#########################################################
litellm_logging_obj.model_call_details["query"] = query
request_body: Dict[str, Any] = {"query": query, "pageSize": 10}
max_num_results = vector_store_search_optional_params.get("max_num_results")
if max_num_results is not None:
request_body["pageSize"] = max_num_results
if isinstance(extra_body, dict):
request_body.update(
self._filter_extra_body(extra_body, is_engine=is_engine)
)
litellm_logging_obj.model_call_details["query"] = request_body.get(
"query", query
)
return url, request_body

View file

@ -114,33 +114,18 @@ class VertexAIModelGardenModels(VertexBase):
openai_like_chat_completions = OpenAILikeChatHandler()
## CONSTRUCT API BASE
# Skip _check_custom_proxy: its ":verb" URL construction corrupts a
# user-supplied api_base (e.g. Vertex MG dedicated endpoint), and
# OpenAILikeChatHandler already appends "/chat/completions".
stream: bool = optional_params.get("stream", False) or False
optional_params["stream"] = stream
default_api_base = create_vertex_url(
vertex_location=vertex_location or "us-central1",
vertex_project=vertex_project or project_id,
stream=stream,
model=model,
)
if len(default_api_base.split(":")) > 1:
endpoint = default_api_base.split(":")[-1]
else:
endpoint = ""
_, api_base = self._check_custom_proxy(
api_base=api_base,
custom_llm_provider="vertex_ai",
gemini_api_key=None,
endpoint=endpoint,
stream=stream,
auth_header=None,
url=default_api_base,
model=model,
vertex_project=vertex_project or project_id,
vertex_location=vertex_location or "us-central1",
vertex_api_version="v1beta1",
)
if api_base is None:
api_base = create_vertex_url(
vertex_location=vertex_location or "us-central1",
vertex_project=vertex_project or project_id,
stream=stream,
model=model,
)
# Publisher/catalog models: model id must be sent in the JSON body (OpenAPI route).
# Single-segment endpoint ids: model is encoded in the URL path; body model stays empty.
if not _vertex_model_garden_model_id_in_json_body(model):

View file

@ -641,6 +641,7 @@ async def acompletion( # noqa: PLR0915
if (
custom_llm_provider == "text-completion-openai"
or custom_llm_provider == "text-completion-codestral"
or custom_llm_provider == "text-completion-inception"
) and isinstance(response, TextCompletionResponse):
response = litellm.OpenAITextCompletionConfig().convert_to_chat_model_response_object(
response_object=response,
@ -3803,6 +3804,67 @@ def completion( # type: ignore # noqa: PLR0915
):
return _model_response
response = _model_response
elif custom_llm_provider == "text-completion-inception":
passed_api_base = (
api_base
or optional_params.pop("api_base", None)
or optional_params.pop("base_url", None)
)
api_base = (
passed_api_base
or get_secret_str("INCEPTION_API_BASE")
or "https://api.inceptionlabs.ai/v1"
)
# FIM is served at `/v1/fim/completions`; the OpenAI client appends
# `/completions`, so point it at the `/v1/fim` base.
api_base = api_base.rstrip("/")
if not api_base.endswith("/fim"):
api_base += "/fim"
# Don't forward the server-managed Inception key to a caller-supplied
# api_base; only resolve it for the default/server base, or when the
# caller passes their own key.
if passed_api_base is None or api_key:
api_key = (
api_key
or litellm.inception_key
or get_secret_str("INCEPTION_API_KEY")
)
_response = openai_text_completions.completion(
model=model,
messages=messages,
model_response=model_response,
print_verbose=print_verbose,
api_key=api_key, # type: ignore[arg-type]
custom_llm_provider="text-completion-inception",
api_base=api_base,
acompletion=acompletion,
client=client,
logging_obj=logging,
optional_params=optional_params,
litellm_params=litellm_params,
logger_fn=logger_fn,
timeout=timeout, # type: ignore
)
if (
optional_params.get("stream", False) is False
and acompletion is False
and text_completion is False
):
_response = litellm.OpenAITextCompletionConfig().convert_to_chat_model_response_object(
response_object=_response, model_response_object=model_response
)
if optional_params.get("stream", False) or acompletion is True:
logging.post_call(
input=messages,
api_key=api_key,
original_response=_response,
additional_args={"headers": headers},
)
response = _response
elif custom_llm_provider in ("sagemaker_chat", "sagemaker_nova"):
# boto3 reads keys from .env
# sagemaker_chat: HF Messages API endpoints
@ -4503,6 +4565,39 @@ def completion( # type: ignore # noqa: PLR0915
client=client,
)
elif custom_llm_provider == "langflow":
# LangFlow - Visual AI Agent Platform
from litellm.llms.langflow.chat.transformation import LangFlowConfig
(
api_base,
api_key,
) = LangFlowConfig()._get_openai_compatible_provider_info(
api_base=api_base or litellm.api_base,
api_key=api_key or litellm.api_key,
)
headers = headers or litellm.headers
response = base_llm_http_handler.completion(
model=model,
stream=stream,
messages=messages,
acompletion=acompletion,
api_base=api_base,
model_response=model_response,
optional_params=optional_params,
litellm_params=litellm_params,
shared_session=shared_session,
custom_llm_provider=custom_llm_provider,
timeout=timeout,
headers=headers,
encoding=_get_encoding(),
api_key=api_key,
logging_obj=logging,
client=client,
)
else:
raise LiteLLMUnknownProvider(
model=model, custom_llm_provider=custom_llm_provider

File diff suppressed because it is too large Load diff

View file

@ -7,6 +7,7 @@ from starlette.requests import Request
from starlette.types import Scope
from litellm._logging import verbose_logger
from litellm.constants import DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL
from litellm.proxy._types import (
LiteLLM_TeamTable,
ProxyException,
@ -14,6 +15,88 @@ from litellm.proxy._types import (
UserAPIKeyAuth,
)
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.auth.ip_address_utils import IPAddressUtils
def _parse_mcp_server_names_from_path(
path: str, mcp_servers_header: Optional[List[str]] = None
) -> Optional[List[str]]:
"""Resolve the single MCP server name a cold-start passthrough bypass may
target. Delegates parsing to
:meth:`MCPRequestHandler._extract_target_server_names_from_path` so the
names used here always match the names downstream routing uses; returns
``None`` whenever the bypass must not activate (aggregate ``/mcp``,
multi-server CSV paths, or any other unrecognized path).
Also fails closed when the ``x-mcp-servers`` header introduces any server
not present in the path-derived target set. Downstream routing for
``/mcp/...`` paths overrides the header with path-derived names, but a
header/path mismatch here is a sign of a confused or hostile caller —
refuse the cold-start bypass rather than admit anonymously based on the
path while the header advertises a stricter, non-passthrough target."""
servers = MCPRequestHandler._extract_target_server_names_from_path(path)
if len(servers) != 1:
verbose_logger.debug(
"MCP cold-start: path %r resolved to %r; passthrough 401 bypass "
"requires exactly one target and will not activate",
path,
servers,
)
return None
if mcp_servers_header is not None and (set(mcp_servers_header) - set(servers)):
verbose_logger.debug(
"MCP cold-start: x-mcp-servers header %r introduces target(s) not "
"in path-derived set %r; passthrough 401 bypass will not activate",
mcp_servers_header,
servers,
)
return None
return servers
def _is_mcp_passthrough_cold_start(
mcp_servers: Optional[List[str]], client_ip: Optional[str]
) -> bool:
"""True only when EVERY targeted server is a pass-through server with no
auth headers — the cold-start OAuth discovery case per RFC 9728 / MCP
Authorization spec. Lets the route handler's 401 emitter produce the
spec-compliant WWW-Authenticate challenge instead of surfacing a generic
admission error.
Uses "all" semantics (mirrors :meth:`MCPRequestHandler._target_servers_use_oauth2`):
one non-passthrough target in a co-targeted set must not flip the bypass
open for the others. Fails closed when any target cannot be resolved."""
if not mcp_servers:
return False
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
global_mcp_server_manager,
)
for name in mcp_servers:
server = global_mcp_server_manager.get_mcp_server_by_name(
name, client_ip=client_ip
)
if server is None or not getattr(server, "is_oauth_passthrough", False):
return False
return True
def _is_litellm_auth_admission_error(exc: Exception) -> bool:
if isinstance(exc, HTTPException):
return exc.status_code == 401
if isinstance(exc, ProxyException):
try:
return int(exc.code) == 401
except (TypeError, ValueError):
return False
return False
def _has_client_supplied_mcp_auth(
mcp_auth_header: Optional[str],
mcp_server_auth_headers: Optional[Dict[str, Dict[str, str]]],
) -> bool:
return bool(mcp_auth_header) or bool(mcp_server_auth_headers)
class MCPRequestHandler:
@ -37,7 +120,7 @@ class MCPRequestHandler:
LITELLM_MCP_ACCESS_GROUPS_HEADER_NAME = SpecialHeaders.mcp_access_groups.value
@staticmethod
async def process_mcp_request(
async def process_mcp_request( # noqa: PLR0915
scope: Scope,
) -> Tuple[
UserAPIKeyAuth,
@ -130,7 +213,9 @@ class MCPRequestHandler:
elif (
not litellm_api_key
and MCPRequestHandler._target_servers_delegate_auth_to_upstream( # noqa: E501
path=request_route, mcp_servers=mcp_servers
path=request_route,
mcp_servers=mcp_servers,
client_ip=IPAddressUtils.get_mcp_client_ip(request),
)
):
# Operator opted this oauth2 server into upstream-delegated auth
@ -172,25 +257,87 @@ class MCPRequestHandler:
# than coercing (``int("None")`` would raise ValueError and
# rewrite the auth error as a 500).
status = e.status_code if isinstance(e, HTTPException) else e.code
if status in (
401,
403,
"401",
"403",
) and MCPRequestHandler._target_servers_use_oauth2(
path=request_route, mcp_servers=mcp_servers
is_auth_error = status in (401, 403, "401", "403")
is_unauthenticated = status in (401, "401")
client_ip = IPAddressUtils.get_mcp_client_ip(request)
if is_auth_error and MCPRequestHandler._target_servers_use_oauth2(
path=request_route,
mcp_servers=mcp_servers,
client_ip=client_ip,
):
verbose_logger.debug(
"MCP OAuth2: target server is OAuth2-mode, treating "
"Authorization as upstream OAuth2 token passthrough"
)
validated_user_api_key_auth = UserAPIKeyAuth()
elif is_unauthenticated:
# Pass-through cold-start return: per RFC 9728 / MCP
# Authorization spec the client completes upstream OAuth
# discovery and returns with ``Authorization: Bearer
# <upstream-token>``. For ``auth_type=none`` passthrough
# servers that bearer is not a LiteLLM key (auth above
# failed) but is meant to be forwarded upstream
# unchanged. Fall back to anonymous admission so the
# caller is not rejected for following the discovery
# flow without also setting ``x-litellm-api-key``.
# Only trigger on 401 (token unrecognized); a 403 means
# the key WAS recognized but is forbidden (e.g. over
# budget / rate limited) and must propagate so those
# controls are not bypassed via anonymous admission.
mcp_servers_from_path = _parse_mcp_server_names_from_path(
request_route, mcp_servers
)
if (
mcp_servers_from_path is not None
and not _has_client_supplied_mcp_auth(
mcp_auth_header,
mcp_server_auth_headers,
)
and _is_mcp_passthrough_cold_start(
mcp_servers_from_path, client_ip=client_ip
)
):
verbose_logger.debug(
"MCP pass-through return: target server is "
"passthrough, treating Authorization as "
"upstream OAuth token for delegated auth"
)
validated_user_api_key_auth = UserAPIKeyAuth()
else:
raise
else:
raise
else:
validated_user_api_key_auth = await user_api_key_auth(
api_key=litellm_api_key, request=request
)
try:
validated_user_api_key_auth = await user_api_key_auth(
api_key=litellm_api_key, request=request
)
except (HTTPException, ProxyException) as exc:
# Cold-start MCP OAuth discovery: RFC 9728 / MCP Authorization spec
# require unauthenticated requests to protected resources to receive
# 401 + WWW-Authenticate. Defer to _raise_preemptive_401_for_unauthenticated_servers
# for pass-through servers instead of surfacing a generic admission error.
mcp_servers_from_path = _parse_mcp_server_names_from_path(
request_route, mcp_servers
)
client_ip = IPAddressUtils.get_mcp_client_ip(request)
if (
mcp_servers_from_path is not None
and not _has_client_supplied_mcp_auth(
mcp_auth_header,
mcp_server_auth_headers,
)
and _is_litellm_auth_admission_error(exc)
and _is_mcp_passthrough_cold_start(
mcp_servers_from_path, client_ip=client_ip
)
):
verbose_logger.debug(
"MCP pass-through cold start: deferring admission to route 401 emitter"
)
validated_user_api_key_auth = UserAPIKeyAuth()
else:
raise
return (
validated_user_api_key_auth,
@ -262,7 +409,9 @@ class MCPRequestHandler:
return [servers_and_path]
@staticmethod
def _target_servers_use_oauth2(path: str, mcp_servers: Optional[List[str]]) -> bool:
def _target_servers_use_oauth2(
path: str, mcp_servers: Optional[List[str]], client_ip: Optional[str]
) -> bool:
"""
True only when EVERY MCP server the request targets is configured for
``auth_type == oauth2``. If any target is non-OAuth2 — or if the target
@ -291,14 +440,16 @@ class MCPRequestHandler:
return False
for name in target_names:
server = global_mcp_server_manager.get_mcp_server_by_name(name)
server = global_mcp_server_manager.get_mcp_server_by_name(
name, client_ip=client_ip
)
if server is None or server.auth_type != MCPAuth.oauth2:
return False
return True
@staticmethod
def _target_servers_delegate_auth_to_upstream(
path: str, mcp_servers: Optional[List[str]]
path: str, mcp_servers: Optional[List[str]], client_ip: Optional[str]
) -> bool:
"""
True only when EVERY MCP server the request targets is configured for
@ -328,7 +479,9 @@ class MCPRequestHandler:
return False
for name in target_names:
server = global_mcp_server_manager.get_mcp_server_by_name(name)
server = global_mcp_server_manager.get_mcp_server_by_name(
name, client_ip=client_ip
)
if server is None or server.auth_type != MCPAuth.oauth2:
return False
# `is True` is intentional: opt-in must be an explicit boolean
@ -1090,22 +1243,21 @@ class MCPRequestHandler:
)
return []
# Sentinel stored in cache when an org has no object_permission, so we
# don't re-query the DB on every MCP request for that org.
_ORG_NO_PERMISSION_SENTINEL = "__org_no_mcp_permission__"
@staticmethod
async def _get_org_object_permission(
user_api_key_auth: Optional[UserAPIKeyAuth] = None,
):
"""
Get org object_permission, using user_api_key_cache to avoid DB hits on every request.
Caches both positive results and the absence of an object_permission so that orgs
with no MCP permissions configured (the common default) do not trigger a DB query
on every request.
Get org object_permission via the established ``get_org_object`` /
``get_object_permission`` helpers so MCP requests share the same
``user_api_key_cache`` entries as the rest of the proxy.
"""
from litellm.proxy.proxy_server import prisma_client, user_api_key_cache
from litellm.proxy.auth.auth_checks import get_object_permission, get_org_object
from litellm.proxy.proxy_server import (
prisma_client,
proxy_logging_obj,
user_api_key_cache,
)
if not user_api_key_auth or not user_api_key_auth.org_id:
return None
@ -1114,45 +1266,25 @@ class MCPRequestHandler:
verbose_logger.debug("prisma_client is None")
return None
org_id = user_api_key_auth.org_id
cache_key = f"org_object_permission:{org_id}"
from litellm.proxy._types import LiteLLM_ObjectPermissionTable
try:
cached = await user_api_key_cache.async_get_cache(key=cache_key)
if cached is not None:
# Sentinel means the DB confirmed no object_permission for this org
if cached == MCPRequestHandler._ORG_NO_PERMISSION_SENTINEL:
return None
# Redis deserialises to a plain dict; reconstruct the Pydantic model
# so callers can access .mcp_servers / .mcp_tool_permissions as attrs.
if isinstance(cached, dict):
return LiteLLM_ObjectPermissionTable(**cached)
return cached
org_row = await prisma_client.db.litellm_organizationtable.find_unique(
where={"organization_id": org_id},
include={"object_permission": True},
org_obj = await get_org_object(
org_id=user_api_key_auth.org_id,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
parent_otel_span=user_api_key_auth.parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
)
if org_row is None or org_row.object_permission is None:
# Cache the negative result so subsequent calls skip the DB
await user_api_key_cache.async_set_cache(
key=cache_key,
value=MCPRequestHandler._ORG_NO_PERMISSION_SENTINEL,
)
if org_obj is None or not org_obj.object_permission_id:
return None
# Convert raw Prisma model → Pydantic before caching. Caching the
# Pydantic .dict() ensures the value survives a Redis JSON round-trip
# as a plain dict that we can reconstruct above (same pattern used by
# get_end_user_object / get_team_object in auth_checks.py).
obj_perm = LiteLLM_ObjectPermissionTable(**org_row.object_permission.dict())
await user_api_key_cache.async_set_cache(
key=cache_key, value=obj_perm.dict()
return await get_object_permission(
object_permission_id=org_obj.object_permission_id,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
parent_otel_span=user_api_key_auth.parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
)
return obj_perm
except Exception as e:
verbose_logger.warning(f"Failed to get org object permission: {str(e)}")
return None
@ -1273,16 +1405,26 @@ class MCPRequestHandler:
)
return []
# Sentinel stored in cache when an agent has no object_permission, so we
# don't re-query the DB on every MCP request for that agent.
_AGENT_NO_PERMISSION_SENTINEL = "__agent_no_mcp_permission__"
@staticmethod
async def _get_agent_object_permission(
user_api_key_auth: Optional[UserAPIKeyAuth] = None,
):
"""
Fetch the agent's object_permission from the DB (single query).
Returns the object_permission object or None.
Get agent object_permission via the established ``get_object_permission``
helper. Caches the ``agent_id -> object_permission_id`` mapping so we
avoid re-reading the agent row on every request, and reuses the shared
``object_permission_id`` cache populated by the org / team / key paths.
"""
from litellm.proxy.proxy_server import prisma_client
from litellm.proxy.auth.auth_checks import get_object_permission
from litellm.proxy.proxy_server import (
prisma_client,
proxy_logging_obj,
user_api_key_cache,
)
if not user_api_key_auth or not user_api_key_auth.agent_id:
return None
@ -1291,15 +1433,42 @@ class MCPRequestHandler:
verbose_logger.debug("prisma_client is None")
return None
agent_id = user_api_key_auth.agent_id
cache_key = f"agent_object_permission_id:{agent_id}"
try:
agent_row = await prisma_client.db.litellm_agentstable.find_unique(
where={"agent_id": user_api_key_auth.agent_id},
include={"object_permission": True},
object_permission_id: Optional[str] = (
await user_api_key_cache.async_get_cache(key=cache_key)
)
if agent_row is None or agent_row.object_permission is None:
if object_permission_id == MCPRequestHandler._AGENT_NO_PERMISSION_SENTINEL:
return None
return agent_row.object_permission
if object_permission_id is None:
agent_row = await prisma_client.db.litellm_agentstable.find_unique(
where={"agent_id": agent_id},
)
object_permission_id = (
getattr(agent_row, "object_permission_id", None)
if agent_row is not None
else None
)
await user_api_key_cache.async_set_cache(
key=cache_key,
value=object_permission_id
or MCPRequestHandler._AGENT_NO_PERMISSION_SENTINEL,
ttl=DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL,
)
if not object_permission_id:
return None
return await get_object_permission(
object_permission_id=object_permission_id,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
parent_otel_span=user_api_key_auth.parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
)
except Exception as e:
verbose_logger.warning(f"Failed to get agent object permission: {str(e)}")
return None

View file

@ -1,8 +1,11 @@
import asyncio
import html as _html
import json
from typing import Any, Dict, Optional
import time
from typing import Any, Dict, Optional, Tuple
from urllib.parse import parse_qsl, urlencode, urlparse, urlunparse
import httpx
from fastapi import APIRouter, Form, HTTPException, Request
from fastapi.responses import HTMLResponse, JSONResponse, RedirectResponse
@ -26,11 +29,54 @@ from litellm.proxy.utils import get_server_root_path
from litellm.types.mcp import MCPAuth
from litellm.types.mcp_server.mcp_server_manager import MCPServer
# TTL cache for upstream OAuth metadata fetched from pass-through MCP servers.
# Keeps us from hammering the upstream IdP on each discovery request.
# Keyed by (server_id, resource_url) → (expires_at_epoch, payload).
# A payload of ``None`` is a negative-result entry that prevents repeated
# upstream fetches when the IdP consistently has no metadata to serve.
_OAUTH_METADATA_CACHE: Dict[Tuple[str, str], Tuple[float, Optional[dict]]] = {}
_OAUTH_METADATA_CACHE_TTL_SECONDS = 300
_OAUTH_METADATA_NEGATIVE_CACHE_TTL_SECONDS = 60
_OAUTH_METADATA_CACHE_MAX_SIZE = 128
# Per-(server_id, resource_url) async locks so concurrent discovery requests
# coalesce onto a single upstream fetch instead of issuing N parallel calls.
_OAUTH_METADATA_FETCH_LOCKS: Dict[Tuple[str, str], asyncio.Lock] = {}
router = APIRouter(
tags=["mcp"],
)
def _prune_oauth_metadata_cache(now: Optional[float] = None) -> None:
now = now if now is not None else time.time()
expired_cache_keys = [
cache_key
for cache_key, (expires_at, _payload) in _OAUTH_METADATA_CACHE.items()
if expires_at <= now
]
for cache_key in expired_cache_keys:
_OAUTH_METADATA_CACHE.pop(cache_key, None)
if len(_OAUTH_METADATA_CACHE) > _OAUTH_METADATA_CACHE_MAX_SIZE:
overflow = len(_OAUTH_METADATA_CACHE) - _OAUTH_METADATA_CACHE_MAX_SIZE
cache_keys_by_expiry = sorted(
_OAUTH_METADATA_CACHE,
key=lambda cache_key: _OAUTH_METADATA_CACHE[cache_key][0],
)
for cache_key in cache_keys_by_expiry[:overflow]:
_OAUTH_METADATA_CACHE.pop(cache_key, None)
# Drop locks whose cache entry has been evicted and that aren't currently
# held; held locks stay so in-flight callers continue to coalesce.
for cache_key in list(_OAUTH_METADATA_FETCH_LOCKS):
if cache_key in _OAUTH_METADATA_CACHE:
continue
lock = _OAUTH_METADATA_FETCH_LOCKS.get(cache_key)
if lock is None or lock.locked():
continue
_OAUTH_METADATA_FETCH_LOCKS.pop(cache_key, None)
def encode_state_with_base_url(
base_url: str,
original_state: str,
@ -125,6 +171,17 @@ def _resolve_oauth2_server_for_root_endpoints(
return None
def _normalize_for_token_comparison(value: Any) -> str:
"""Stringify ``value`` for token-rule comparison.
Booleans are lower-cased so Python's ``True`` / ``False`` line up with
JSON-style ``"true"`` / ``"false"`` rules from admin config.
"""
if isinstance(value, bool):
return "true" if value else "false"
return str(value)
def _validate_token_response(
token_response: Dict[str, Any],
validation_rules: Dict[str, Any],
@ -136,7 +193,9 @@ def _validate_token_response(
``token_response["team"]["enterprise_id"]``). Top-level keys are tried first,
then dot-split traversal. All comparisons are string-coerced so that numeric
values in the response (e.g. ``"org_id": 12345``) match string rules
(``"org_id": "12345"``).
(``"org_id": "12345"``). Booleans are normalised to JSON-style ``"true"`` /
``"false"`` so admin rules written as ``{"verified": "true"}`` match upstream
responses of ``{"verified": true}``.
"""
for key, expected in validation_rules.items():
actual: Any = token_response.get(key)
@ -163,7 +222,9 @@ def _validate_token_response(
),
},
)
if str(actual) != str(expected):
if _normalize_for_token_comparison(actual) != _normalize_for_token_comparison(
expected
):
raise HTTPException(
status_code=403,
detail={
@ -400,6 +461,11 @@ async def exchange_token_with_server(
headers={"Accept": "application/json"},
data=token_data,
)
if response is None:
raise HTTPException(
status_code=502,
detail="MCP upstream token endpoint returned no response",
)
response.raise_for_status()
token_response = response.json()
@ -505,6 +571,11 @@ async def register_client_with_server(
headers=headers,
json=register_data,
)
if response is None:
raise HTTPException(
status_code=502,
detail="MCP upstream registration endpoint returned no response",
)
response.raise_for_status()
token_response = response.json()
@ -766,7 +837,119 @@ async def callback(
"""
def _build_oauth_protected_resource_response(
async def fetch_upstream_oauth_protected_resource(
mcp_server: MCPServer,
) -> Optional[dict]:
"""Fetch the upstream MCP server's ``.well-known/oauth-protected-resource``
metadata for a pass-through server.
Tries host-only first, then falls back to the RFC 9728 §3.1 path-suffix
form (e.g. ``https://host/.well-known/oauth-protected-resource/mcp``) to
cover upstreams that scope metadata per resource path.
Responses are cached in-process for ~5 minutes keyed on
``(server_id, resource_url)`` so we do not hammer the IdP.
Returns the parsed JSON dict on success, or ``None`` if neither form
responds with a 2xx JSON payload. Raises on network/connection errors so
the caller can emit HTTP 502 rather than fabricate a gateway response.
"""
if not mcp_server.url:
return None
upstream = urlparse(mcp_server.url)
if not upstream.scheme or not upstream.netloc:
return None
cache_key = (mcp_server.server_id, mcp_server.url)
now = time.time()
_prune_oauth_metadata_cache(now)
cached = _OAUTH_METADATA_CACHE.get(cache_key)
if cached is not None and cached[0] > now:
return cached[1]
lock = _OAUTH_METADATA_FETCH_LOCKS.setdefault(cache_key, asyncio.Lock())
async with lock:
now = time.time()
cached = _OAUTH_METADATA_CACHE.get(cache_key)
if cached is not None and cached[0] > now:
return cached[1]
host_base = f"{upstream.scheme}://{upstream.netloc}"
candidates = [f"{host_base}/.well-known/oauth-protected-resource"]
# RFC 9728 §3.1 path fallback
if upstream.path and upstream.path not in ("", "/"):
candidates.append(
f"{host_base}/.well-known/oauth-protected-resource"
f"{upstream.path.rstrip('/')}"
)
async_client = get_async_httpx_client(
llm_provider=httpxSpecialProvider.Oauth2Check
)
network_errors: list[Exception] = []
for candidate in candidates:
try:
response = await async_client.get(
candidate,
headers={"Accept": "application/json"},
)
except Exception as exc:
if is_network_error(exc):
network_errors.append(exc)
else:
verbose_logger.warning(
"MCP OAuth metadata fetch for %s raised non-transport "
"%s: %s — treating as no metadata for this candidate",
candidate,
type(exc).__name__,
exc,
)
continue
if response.status_code == 200:
try:
payload = response.json()
except Exception as exc:
verbose_logger.warning(
"MCP OAuth metadata at %s returned 200 but JSON "
"decode failed (%s: %s) — treating as no metadata",
candidate,
type(exc).__name__,
exc,
)
continue
if isinstance(payload, dict):
now = time.time()
_OAUTH_METADATA_CACHE[cache_key] = (
now + _OAUTH_METADATA_CACHE_TTL_SECONDS,
payload,
)
_prune_oauth_metadata_cache(now)
return payload
if len(network_errors) == len(candidates):
raise network_errors[-1]
# Negative-result caching: when no candidate yielded a usable payload,
# remember that for a shorter TTL so we don't re-fetch on every
# subsequent discovery request (and so the per-key lock can be pruned).
now = time.time()
_OAUTH_METADATA_CACHE[cache_key] = (
now + _OAUTH_METADATA_NEGATIVE_CACHE_TTL_SECONDS,
None,
)
_prune_oauth_metadata_cache(now)
return None
def is_network_error(exc: Exception) -> bool:
"""True for transport-layer failures (connection refused, DNS, TLS, timeout)
as opposed to HTTP protocol errors (4xx/5xx with a valid response)."""
return isinstance(exc, httpx.TransportError)
async def _build_oauth_protected_resource_response(
request: Request,
mcp_server_name: Optional[str],
use_standard_pattern: bool,
@ -774,6 +957,12 @@ def _build_oauth_protected_resource_response(
"""
Build OAuth protected resource response with the appropriate URL pattern.
For pass-through MCP servers (``MCPServer.is_oauth_passthrough``), the
gateway proxies the upstream's own ``oauth-protected-resource`` metadata
so that standards-compliant MCP clients discover the **upstream** IdP
instead of the gateway. The ``resource`` field is rewritten to the
gateway's own URL so clients present the bearer token back to the gateway.
Args:
request: FastAPI Request object
mcp_server_name: Name of the MCP server
@ -813,6 +1002,46 @@ def _build_oauth_protected_resource_response(
else:
resource_url = f"{request_base_url}/mcp"
# Pass-through branch: proxy the upstream's own metadata so discovery
# directs the client at the real IdP (Okta, Keycloak, …) instead of us.
if mcp_server is not None and mcp_server.is_oauth_passthrough:
try:
upstream_metadata = await fetch_upstream_oauth_protected_resource(
mcp_server
)
except Exception as exc:
verbose_logger.warning(
"Failed to fetch upstream oauth-protected-resource metadata "
f"for pass-through MCP server {mcp_server.name!r}: {exc}"
)
raise HTTPException(
status_code=502,
detail=(
"Failed to fetch upstream oauth-protected-resource "
f"metadata for MCP server {mcp_server.name!r}"
),
)
if upstream_metadata is not None:
response = {**upstream_metadata, "resource": resource_url}
return response
# Upstream responded but with non-200 or non-dict payload. For
# pass-through servers the gateway is NOT the authorization server,
# so we must not fall through to the default gateway metadata —
# that would point clients at the wrong IdP.
verbose_logger.warning(
"Upstream oauth-protected-resource metadata unavailable for "
f"pass-through MCP server {mcp_server.name!r}"
)
raise HTTPException(
status_code=502,
detail=(
"Upstream oauth-protected-resource metadata unavailable "
f"for MCP server {mcp_server.name!r}"
),
)
return {
"authorization_servers": [
(
@ -843,7 +1072,7 @@ async def oauth_protected_resource_mcp_standard(request: Request, mcp_server_nam
This endpoint is compliant with MCP specification and works with standard
MCP clients like mcp-inspector and VSCode Copilot.
"""
return _build_oauth_protected_resource_response(
return await _build_oauth_protected_resource_response(
request=request,
mcp_server_name=mcp_server_name,
use_standard_pattern=True,
@ -868,36 +1097,22 @@ async def oauth_protected_resource_mcp(
This endpoint is kept for backward compatibility. New integrations should
use the standard MCP pattern (/mcp/{server_name}) instead.
"""
return _build_oauth_protected_resource_response(
return await _build_oauth_protected_resource_response(
request=request,
mcp_server_name=mcp_server_name,
use_standard_pattern=False,
)
"""
https://datatracker.ietf.org/doc/html/rfc8414#section-3.1
RFC 8414: Path-aware OAuth discovery
If the issuer identifier value contains a path component, any
terminating "/" MUST be removed before inserting "/.well-known/" and
the well-known URI suffix between the host component and the path(include root path)
component.
"""
def _build_oauth_authorization_server_response(
request: Request,
mcp_server_name: Optional[str],
) -> dict:
"""
Build OAuth authorization server metadata response.
"""Build OAuth authorization server metadata response (gateway-as-AS shape).
Args:
request: FastAPI Request object
mcp_server_name: Name of the MCP server
Returns:
OAuth authorization server metadata dict
Synchronous because the body only does dict construction and synchronous
registry lookups; unlike :func:`_build_oauth_protected_resource_response`
it does not need to await any upstream IO.
"""
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
global_mcp_server_manager,

View file

@ -0,0 +1,80 @@
"""Exceptions raised by the LiteLLM MCP proxy."""
from typing import Optional
from fastapi import HTTPException
class MCPUpstreamAuthError(Exception):
"""Raised when an upstream MCP server returns an authentication failure
(typically HTTP 401) and the gateway should surface it transparently to
the client instead of swallowing it.
Only relevant for pass-through MCP servers (see
``MCPServer.is_oauth_passthrough``). The gateway converts this exception
into an HTTP 401 response on single-server routes, preserving any
``WWW-Authenticate`` challenge emitted by the upstream so standards-
compliant MCP clients can trigger the upstream OAuth flow.
"""
def __init__(
self,
status_code: int,
www_authenticate: Optional[str],
server_name: str,
) -> None:
self.status_code = status_code
self.www_authenticate = www_authenticate
self.server_name = server_name
super().__init__(f"Upstream MCP server {server_name!r} returned {status_code}")
def to_http_exception(
self,
base_url: Optional[str] = None,
request_path: Optional[str] = None,
) -> HTTPException:
"""Convert this upstream-auth error into an ``HTTPException`` that
preserves the upstream status code and any ``WWW-Authenticate``
challenge, so standards-compliant MCP clients can trigger the
upstream OAuth flow.
When the upstream 401 omits ``WWW-Authenticate`` (non-compliant per
RFC 7235 §3.1) we fabricate a ``Bearer resource_metadata=`` challenge
that points at the gateway's well-known endpoint for this server, so
MCP clients can still initiate RFC 9728 discovery against the upstream
IdP via the gateway's proxied metadata. Callers must pass ``base_url``
(the gateway origin, no trailing slash) so the fabricated URI is
absolute as RFC 9728 §3.2 requires; if ``base_url`` is missing we
skip fabrication entirely rather than emit a relative URI that strict
clients reject in the Bearer challenge.
When ``request_path`` is supplied and matches the legacy
``/{server_name}/mcp`` MCP transport route, the fabricated URI uses
the matching legacy well-known form
``/.well-known/oauth-protected-resource/{server_name}/mcp``. Otherwise
we default to the standard form
``/.well-known/oauth-protected-resource/mcp/{server_name}``. This
keeps the ``resource_metadata`` URI aligned with the resource pattern
the client originally targeted, matching the path-aware behaviour of
``_get_passthrough_resource_metadata_url`` in ``server.py``.
"""
challenge: Optional[str] = self.www_authenticate
if challenge is None and self.status_code == 401 and base_url:
prefix = base_url.rstrip("/")
if request_path and request_path.startswith(f"/{self.server_name}/mcp"):
resource_metadata_url = (
f"{prefix}/.well-known/oauth-protected-resource/"
f"{self.server_name}/mcp"
)
else:
resource_metadata_url = (
f"{prefix}/.well-known/oauth-protected-resource/"
f"mcp/{self.server_name}"
)
challenge = f'Bearer resource_metadata="{resource_metadata_url}"'
detail = "Forbidden" if self.status_code == 403 else "Unauthorized"
return HTTPException(
status_code=self.status_code,
detail=detail,
headers={"www-authenticate": challenge} if challenge else None,
)

View file

@ -48,6 +48,7 @@ from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import (
MCPRequestHandler,
)
from litellm.proxy._experimental.mcp_server.exceptions import MCPUpstreamAuthError
from litellm.proxy._experimental.mcp_server.oauth2_token_cache import resolve_mcp_auth
from litellm.proxy._experimental.mcp_server.utils import (
MCP_TOOL_PREFIX_SEPARATOR,
@ -118,6 +119,103 @@ _AZURE_ENTRA_HOSTS = {
}
def _should_strip_caller_authorization(
mcp_server: MCPServer,
raw_headers: Optional[Dict[str, str]],
user_api_key_auth: Optional[UserAPIKeyAuth],
) -> bool:
"""Decide whether the caller's ``Authorization`` header must NOT be
forwarded upstream when populating ``extra_headers`` for an MCP server.
Centralized so ``_call_regular_mcp_tool`` (this module) and
``_prepare_mcp_server_headers`` (``server.py``) cannot drift apart on
this security-sensitive decision.
Strip rules:
- **M2M (client_credentials) servers**: never forward the caller's
``Authorization`` — the proxy fetches its own upstream token.
- **OAuth pass-through servers**: strip when the ``Authorization``
header is actually the LiteLLM API key — either because admission
validated it (``user_api_key_auth.api_key`` is set) and the caller
did NOT also supply ``x-litellm-api-key`` to disambiguate, or
because the legacy ``user_api_key_auth is None`` call sites did
not supply an explicit admission header. In the anonymous /
pass-through cold-start case (RFC 9728) the bearer in
``Authorization`` is the upstream OAuth token and must be
forwarded, so we keep it.
"""
if mcp_server.has_client_credentials:
return True
if not mcp_server.is_oauth_passthrough:
return False
normalized_raw_headers = {
str(k).lower(): v for k, v in (raw_headers or {}).items() if isinstance(k, str)
}
has_explicit_litellm_admission_header = (
normalized_raw_headers.get("x-litellm-api-key") is not None
)
admission_consumed_authorization_as_litellm_key = (
user_api_key_auth is not None
and bool(getattr(user_api_key_auth, "api_key", None))
and not has_explicit_litellm_admission_header
)
return admission_consumed_authorization_as_litellm_key or (
user_api_key_auth is None and not has_explicit_litellm_admission_header
)
def _extract_upstream_auth_failure(
exc: BaseException,
) -> Optional[Tuple[int, Optional[str]]]:
"""Walk the exception tree looking for an HTTP 401/403 response from the
upstream MCP server.
The MCP SDK wraps transport errors in anyio ``ExceptionGroup`` objects and
may chain through ``__cause__`` / ``__context__``. We inspect all of those
layers for an ``httpx.Response``-bearing exception (typically
``httpx.HTTPStatusError``) and extract the status code and any upstream
``WWW-Authenticate`` header.
Returns ``(status_code, www_authenticate)`` on match, else ``None``.
"""
seen: Set[int] = set()
stack: List[BaseException] = [exc]
while stack:
current = stack.pop()
if id(current) in seen:
continue
seen.add(id(current))
response = getattr(current, "response", None)
if response is not None:
status_code = getattr(response, "status_code", None)
if isinstance(status_code, int) and status_code in (401, 403):
www_authenticate: Optional[str] = None
headers = getattr(response, "headers", None)
if headers is not None:
try:
www_authenticate = headers.get("www-authenticate")
except Exception:
www_authenticate = None
return status_code, www_authenticate
# anyio / PEP 654 ExceptionGroup
sub_exceptions = getattr(current, "exceptions", None)
if sub_exceptions:
stack.extend(sub_exceptions)
if current.__cause__ is not None:
stack.append(current.__cause__)
if (
current.__context__ is not None
and current.__context__ is not current.__cause__
):
stack.append(current.__context__)
return None
def _warn_on_server_name_fields(
*,
server_id: str,
@ -483,6 +581,7 @@ class MCPServerManager:
delegate_auth_to_upstream=bool(
server_config.get("delegate_auth_to_upstream", False)
),
oauth_passthrough=bool(server_config.get("oauth_passthrough", False)),
# AWS SigV4 fields
aws_access_key_id=server_config.get("aws_access_key_id", None),
aws_secret_access_key=server_config.get("aws_secret_access_key", None),
@ -881,6 +980,7 @@ class MCPServerManager:
delegate_auth_to_upstream=bool(
getattr(mcp_server, "delegate_auth_to_upstream", False)
),
oauth_passthrough=bool(getattr(mcp_server, "oauth_passthrough", False)),
created_at=getattr(mcp_server, "created_at", None),
updated_at=getattr(mcp_server, "updated_at", None),
tool_name_to_display_name=_deserialize_json_dict(
@ -1599,7 +1699,9 @@ class MCPServerManager:
]
return tools
else:
tools = await self._fetch_tools_with_timeout(client, server.name)
tools = await self._fetch_tools_with_timeout(
client, server.name, server=server
)
self._remember_upstream_initialize_instructions(server, client)
prefixed_or_original_tools = self._create_prefixed_tools(
@ -1608,6 +1710,11 @@ class MCPServerManager:
return prefixed_or_original_tools
except MCPUpstreamAuthError:
# Pass-through 401 must surface to single-server routes so the
# client triggers the upstream OAuth flow. The multi-server
# aggregator catches this explicitly to keep absorbing.
raise
except Exception as e:
verbose_logger.warning(
f"Failed to get tools from server {server.name}: {str(e)}"
@ -2209,7 +2316,10 @@ class MCPServerManager:
return None
async def _fetch_tools_with_timeout(
self, client: MCPClient, server_name: str
self,
client: MCPClient,
server_name: str,
server: Optional[MCPServer] = None,
) -> List[MCPTool]:
"""
Fetch tools from MCP client with timeout and error handling.
@ -2217,16 +2327,28 @@ class MCPServerManager:
Uses anyio.fail_after() instead of asyncio.wait_for() to avoid conflicts
with the MCP SDK's anyio TaskGroup. See GitHub issue #20715 for details.
For pass-through MCP servers (``MCPServer.is_oauth_passthrough``) an
upstream HTTP 401 is converted into :class:`MCPUpstreamAuthError`
instead of being swallowed to an empty tool list. That lets the
single-server HTTP routes surface a proper 401 + ``WWW-Authenticate``
challenge so standards-compliant MCP clients trigger the upstream
OAuth flow. Non-pass-through servers keep today's swallow-and-log
behaviour so the multi-server ``/mcp`` aggregator doesn't get
tainted by a single bad server.
Args:
client: MCP client instance
server_name: Name of the server for logging
server: Optional MCPServer; when pass-through, auth errors are
re-raised as :class:`MCPUpstreamAuthError`.
Returns:
List of tools from the server
"""
is_passthrough = bool(server is not None and server.is_oauth_passthrough)
try:
with anyio.fail_after(MCP_TOOL_LISTING_TIMEOUT):
tools = await client.list_tools()
tools = await client.list_tools(raise_on_error=is_passthrough)
verbose_logger.debug(f"Tools from {server_name}: {tools}")
return tools
except TimeoutError:
@ -2243,6 +2365,19 @@ class MCPServerManager:
)
return []
except Exception as e:
if is_passthrough:
auth_info = _extract_upstream_auth_failure(e)
if auth_info is not None:
status_code, www_authenticate = auth_info
verbose_logger.info(
f"Upstream auth failure from pass-through MCP server "
f"{server_name}: HTTP {status_code}"
)
raise MCPUpstreamAuthError(
status_code=status_code,
www_authenticate=www_authenticate,
server_name=server_name,
) from e
verbose_logger.warning(f"Error listing tools from {server_name}: {str(e)}")
return []
@ -2650,6 +2785,9 @@ class MCPServerManager:
"name": name,
"arguments": arguments,
"server_name": server_name,
"mcp_rate_limit_server_name": server.alias
or server.server_name
or server.name,
"user_api_key_auth": user_api_key_auth,
"user_api_key_user_id": (
getattr(user_api_key_auth, "user_id", None)
@ -2768,6 +2906,7 @@ class MCPServerManager:
proxy_logging_obj: Optional[ProxyLogging],
host_progress_callback: Optional[Callable] = None,
hook_extra_headers: Optional[Dict[str, str]] = None,
user_api_key_auth: Optional[UserAPIKeyAuth] = None,
) -> CallToolResult:
"""
Call a regular MCP tool using the MCP client.
@ -2833,13 +2972,16 @@ class MCPServerManager:
normalized_raw_headers = {
str(k).lower(): v for k, v in raw_headers.items() if isinstance(k, str)
}
strip_caller_authorization = _should_strip_caller_authorization(
mcp_server=mcp_server,
raw_headers=raw_headers,
user_api_key_auth=user_api_key_auth,
)
for header in mcp_server.extra_headers:
if not isinstance(header, str):
continue
if (
mcp_server.has_client_credentials
and header.lower() == "authorization"
):
if header.lower() == "authorization" and strip_caller_authorization:
continue
header_value = normalized_raw_headers.get(header.lower())
if header_value is None:
@ -3131,6 +3273,7 @@ class MCPServerManager:
proxy_logging_obj=proxy_logging_obj,
host_progress_callback=host_progress_callback,
hook_extra_headers=hook_result.get("extra_headers"),
user_api_key_auth=user_api_key_auth,
)
return await self._gather_openapi_tool_tasks(tasks, proxy_logging_obj)
@ -3162,7 +3305,23 @@ class MCPServerManager:
if server.needs_user_oauth_token:
# Skip OAuth2 servers that rely on user-provided tokens
continue
tools = await self._get_tools_from_server(server)
try:
tools = await self._get_tools_from_server(server)
except MCPUpstreamAuthError as e:
# Pass-through servers expect a user-supplied bearer token;
# at startup we have none, so an upstream 401 is normal.
# Swallow it so we keep mapping the remaining servers.
verbose_logger.debug(
f"Skipping tool name mapping for server {server.name} "
f"due to upstream auth error: {str(e)}"
)
continue
except Exception as e:
verbose_logger.warning(
f"Failed to get tools from server {server.name} during "
f"tool name mapping initialization: {str(e)}"
)
continue
for tool in tools:
# The tool.name here is already prefixed from _get_tools_from_server
# Extract original name for mapping
@ -3754,6 +3913,7 @@ class MCPServerManager:
allow_all_keys=server.allow_all_keys,
available_on_public_internet=server.available_on_public_internet,
delegate_auth_to_upstream=server.delegate_auth_to_upstream,
oauth_passthrough=getattr(server, "oauth_passthrough", False),
is_byok=server.is_byok,
byok_description=server.byok_description,
byok_api_key_help_url=server.byok_api_key_help_url,

View file

@ -16,6 +16,7 @@ from typing import (
from fastapi import APIRouter, Depends, HTTPException, Query, Request, status
from litellm._logging import verbose_logger
from litellm.proxy._experimental.mcp_server.exceptions import MCPUpstreamAuthError
from litellm.proxy._experimental.mcp_server.ui_session_utils import (
build_effective_auth_contexts,
)
@ -46,6 +47,9 @@ if MCP_AVAILABLE:
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
global_mcp_server_manager,
)
from litellm.proxy._experimental.mcp_server.oauth_utils import (
get_request_base_url,
)
from litellm.proxy._experimental.mcp_server.server import (
ListMCPToolsRestAPIResponseObject,
MCPServer,
@ -423,101 +427,6 @@ if MCP_AVAILABLE:
allowed_mcp_servers.append(server)
return allowed_mcp_servers
async def _list_tools_for_single_server(
server_id: str,
allowed_server_ids: List[str],
rest_client_ip: Optional[str],
mcp_server_auth_headers: dict,
mcp_auth_header: Optional[str],
raw_headers_from_request: dict,
user_api_key_dict: "UserAPIKeyAuth",
) -> dict:
"""
Resolve and fetch tools for a single specified MCP server.
Returns the full REST response dict (tools / error / message).
Raises HTTPException on access / IP-filter errors.
"""
# Resolve a server name to its UUID if needed
_name_resolved = None
if server_id not in allowed_server_ids:
_name_resolved = global_mcp_server_manager.get_mcp_server_by_name(server_id)
if _name_resolved is not None and _name_resolved.server_id in set(
allowed_server_ids
):
server_id = _name_resolved.server_id
if server_id not in allowed_server_ids:
_server = (
global_mcp_server_manager.get_mcp_server_by_id(server_id)
or _name_resolved
)
if (
_server is not None
and rest_client_ip is not None
and not global_mcp_server_manager._is_server_accessible_from_ip(
_server, rest_client_ip
)
):
raise HTTPException(
status_code=403,
detail={
"error": "ip_filtering",
"message": (
f"MCP server '{server_id}' is not accessible from your IP address "
f"({rest_client_ip}). This server is restricted to internal "
"networks only. To make it externally accessible, set "
"'available_on_public_internet: true' in the server configuration."
),
},
)
raise HTTPException(
status_code=403,
detail={
"error": "access_denied",
"message": f"The key is not allowed to access server {server_id}",
},
)
server = global_mcp_server_manager.get_mcp_server_by_id(server_id)
if server is None:
return {
"tools": [],
"error": "server_not_found",
"message": f"Server with id {server_id} not found",
}
server_auth_header = _get_server_auth_header(
server, mcp_server_auth_headers, mcp_auth_header
)
user_oauth_extra_headers = await _get_user_oauth_extra_headers(
server, user_api_key_dict
)
try:
tools = await _get_tools_for_single_server(
server,
server_auth_header,
raw_headers_from_request,
user_api_key_dict,
extra_headers=user_oauth_extra_headers,
)
except Exception as e:
verbose_logger.exception(f"Error getting tools from {server.name}: {e}")
return {
"tools": [],
"error": "server_error",
"message": f"Failed to get tools from server {server.name}: {str(e)}",
}
return {
"tools": tools,
"error": None,
"message": "Successfully retrieved tools",
}
########################################################
async def _list_tools_for_single_server(
server_id: str,
allowed_server_ids: List[str],
@ -591,6 +500,11 @@ if MCP_AVAILABLE:
user_api_key_dict,
extra_headers=user_oauth_extra_headers,
)
except MCPUpstreamAuthError:
# Surface the upstream 401/403 to the caller so it can emit the
# matching status code and WWW-Authenticate challenge; that is what
# lets standards-compliant MCP clients run the upstream OAuth flow.
raise
except Exception as e:
verbose_logger.exception(f"Error getting tools from {server.name}: {e}")
return {
@ -757,6 +671,24 @@ if MCP_AVAILABLE:
),
}
except MCPUpstreamAuthError as e:
# Surface upstream pass-through 401/403 challenges to the client so
# standards-compliant MCP clients can run the upstream OAuth flow.
raise e.to_http_exception(
base_url=get_request_base_url(request),
request_path=request.scope.get("_original_path") or request.url.path,
)
except HTTPException as http_exc:
# Internal access/IP 403s keep the legacy error-dict response shape
# so the existing contract stays intact.
verbose_logger.exception(
"HTTPException in list_tool_rest_api: %s", str(http_exc)
)
return {
"tools": [],
"error": "unexpected_error",
"message": (f"An unexpected error occurred: {http_exc.detail}"),
}
except Exception as e:
verbose_logger.exception(
"Unexpected error in list_tool_rest_api: %s", str(e)

View file

@ -20,6 +20,7 @@ from typing import (
Dict,
List,
Optional,
Set,
Tuple,
Union,
cast,
@ -38,6 +39,7 @@ from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging
from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import (
MCPRequestHandler,
)
from litellm.proxy._experimental.mcp_server.exceptions import MCPUpstreamAuthError
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
get_request_base_url,
)
@ -145,6 +147,19 @@ _SESSION_MANAGERS_INITIALIZED = False
_INITIALIZATION_LOCK = asyncio.Lock()
def _mcp_session_id_from_headers(
raw_headers: Optional[Dict[str, str]],
) -> Optional[str]:
"""The ``mcp-session-id`` of a stateful MCP session, read case-insensitively
from the request headers. ``None`` for stateless calls (no such header)."""
if not raw_headers:
return None
for key, value in raw_headers.items():
if isinstance(key, str) and key.lower() == "mcp-session-id":
return value or None
return None
if MCP_AVAILABLE:
from mcp.server import Server
from mcp.server.lowlevel.server import NotificationOptions
@ -174,6 +189,7 @@ if MCP_AVAILABLE:
)
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
MCPServerManager,
_should_strip_caller_authorization,
global_mcp_server_manager,
)
from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import (
@ -1086,6 +1102,42 @@ if MCP_AVAILABLE:
return allowed_mcp_servers
def _client_has_passthrough_authorization(
server: MCPServer,
oauth2_headers: Optional[Dict[str, str]],
mcp_server_auth_headers: Optional[Dict[str, Dict[str, str]]],
) -> bool:
"""True if the incoming request already carries an ``Authorization``
header the gateway will forward to this pass-through server.
The client may supply the bearer as either the top-level
``Authorization`` header (surfaced via ``oauth2_headers``) or a
per-server ``x-mcp-auth-<alias>`` style header (surfaced via
``mcp_server_auth_headers``). Either form skips the pre-emptive 401.
"""
if oauth2_headers:
for k in oauth2_headers.keys():
if k.lower() == "authorization":
return True
if mcp_server_auth_headers:
for key in (server.alias, server.server_name, server.name):
if not key:
continue
server_headers = None
for k, v in mcp_server_auth_headers.items():
if k.lower() == key.lower():
server_headers = v
break
if server_headers is None:
continue
if isinstance(server_headers, str) and server_headers.strip():
return True
if isinstance(server_headers, dict):
for hk in server_headers.keys():
if hk.lower() == "authorization":
return True
return False
async def _get_user_oauth_extra_headers_from_db(
server: MCPServer,
user_api_key_auth: Optional[UserAPIKeyAuth],
@ -1266,6 +1318,7 @@ if MCP_AVAILABLE:
mcp_auth_header: Optional[str],
oauth2_headers: Optional[Dict[str, str]],
raw_headers: Optional[Dict[str, str]],
user_api_key_auth: Optional[UserAPIKeyAuth] = None,
) -> Tuple[Optional[Union[Dict[str, str], str]], Optional[Dict[str, str]]]:
"""Build auth and extra headers for a server."""
server_auth_header: Optional[Union[Dict[str, str], str]] = None
@ -1298,10 +1351,20 @@ if MCP_AVAILABLE:
str(k).lower(): v for k, v in raw_headers.items() if isinstance(k, str)
}
# Centralized strip decision shared with
# ``MCPServerManager._call_regular_mcp_tool`` so the two
# code paths cannot drift on this security-sensitive choice.
# See ``_should_strip_caller_authorization`` for the rules.
strip_caller_authorization = _should_strip_caller_authorization(
mcp_server=server,
raw_headers=raw_headers,
user_api_key_auth=user_api_key_auth,
)
for header in server.extra_headers:
if not isinstance(header, str):
continue
if server.has_client_credentials and header.lower() == "authorization":
if header.lower() == "authorization" and strip_caller_authorization:
continue
header_value = normalized_raw_headers.get(header.lower())
if header_value is None:
@ -1510,6 +1573,7 @@ if MCP_AVAILABLE:
mcp_auth_header=mcp_auth_header,
oauth2_headers=oauth2_headers,
raw_headers=raw_headers,
user_api_key_auth=user_api_key_auth,
)
# Prefer server-stored per-user OAuth when configured, so a stale
@ -1561,6 +1625,13 @@ if MCP_AVAILABLE:
f"Successfully fetched {len(tools)} tools from server {server.name}, {len(filtered_tools)} after filtering"
)
return filtered_tools
except MCPUpstreamAuthError:
# Surface upstream 401/403 to the outer handler so the
# client receives a proper WWW-Authenticate challenge
# instead of a silently empty tool list. Without this
# re-raise the broad ``except Exception`` below would
# swallow the auth error.
raise
except Exception as e:
verbose_logger.exception(
f"Error getting tools from server {server.name}: {str(e)}"
@ -1684,6 +1755,7 @@ if MCP_AVAILABLE:
mcp_auth_header=mcp_auth_header,
oauth2_headers=oauth2_headers,
raw_headers=raw_headers,
user_api_key_auth=user_api_key_auth,
)
try:
@ -1741,6 +1813,7 @@ if MCP_AVAILABLE:
mcp_auth_header=mcp_auth_header,
oauth2_headers=oauth2_headers,
raw_headers=raw_headers,
user_api_key_auth=user_api_key_auth,
)
try:
@ -1796,6 +1869,7 @@ if MCP_AVAILABLE:
mcp_auth_header=mcp_auth_header,
oauth2_headers=oauth2_headers,
raw_headers=raw_headers,
user_api_key_auth=user_api_key_auth,
)
try:
@ -2324,6 +2398,7 @@ if MCP_AVAILABLE:
name=original_tool_name, # Use original name for logging
arguments=arguments,
server_name=server_name,
session_id=_mcp_session_id_from_headers(raw_headers),
)
)
litellm_logging_obj: Optional[LiteLLMLoggingObj] = kwargs.get(
@ -2631,6 +2706,7 @@ if MCP_AVAILABLE:
mcp_auth_header=mcp_auth_header,
oauth2_headers=oauth2_headers,
raw_headers=raw_headers,
user_api_key_auth=user_api_key_auth,
)
return await global_mcp_server_manager.get_prompt_from_server(
@ -2681,6 +2757,7 @@ if MCP_AVAILABLE:
mcp_auth_header=mcp_auth_header,
oauth2_headers=oauth2_headers,
raw_headers=raw_headers,
user_api_key_auth=user_api_key_auth,
)
return await global_mcp_server_manager.read_resource_from_server(
@ -2695,8 +2772,10 @@ if MCP_AVAILABLE:
name: str,
arguments: Dict[str, Any],
server_name: Optional[str],
session_id: Optional[str] = None,
) -> StandardLoggingMCPToolCall:
mcp_server = global_mcp_server_manager._get_mcp_server_from_tool_name(name)
namespaced_tool_name = f"{server_name}/{name}" if server_name else name
if mcp_server:
mcp_info = mcp_server.mcp_info or {}
return StandardLoggingMCPToolCall(
@ -2704,13 +2783,15 @@ if MCP_AVAILABLE:
arguments=arguments,
mcp_server_name=mcp_info.get("server_name"),
mcp_server_logo_url=mcp_info.get("logo_url"),
namespaced_tool_name=f"{server_name}/{name}" if server_name else name,
namespaced_tool_name=namespaced_tool_name,
mcp_session_id=session_id,
)
else:
return StandardLoggingMCPToolCall(
name=name,
arguments=arguments,
namespaced_tool_name=f"{server_name}/{name}" if server_name else name,
namespaced_tool_name=namespaced_tool_name,
mcp_session_id=session_id,
)
async def _handle_managed_mcp_tool(
@ -3136,6 +3217,117 @@ if MCP_AVAILABLE:
)
return user_api_key_auth.model_copy(update={"object_permission": updated_op})
def _get_passthrough_resource_metadata_url(scope: Scope, server_name: str) -> str:
request = StarletteRequest(scope)
base_url = get_request_base_url(request)
_path = scope.get("_original_path") or scope.get("path", "") or ""
if _path.startswith(f"/{server_name}/mcp"):
return f"{base_url}/.well-known/oauth-protected-resource/{server_name}/mcp"
return f"{base_url}/.well-known/oauth-protected-resource/mcp/{server_name}"
def _get_passthrough_www_authenticate(
scope: Scope,
server_name: str,
invalid_token: bool = False,
) -> str:
resource_metadata_url = _get_passthrough_resource_metadata_url(
scope=scope,
server_name=server_name,
)
params = []
if invalid_token:
params.append('error="invalid_token"')
params.append(f'resource_metadata="{resource_metadata_url}"')
return "Bearer " + ", ".join(params)
async def _raise_preemptive_401_for_unauthenticated_servers(
scope: Scope,
mcp_servers: Optional[List[str]],
oauth2_headers: Optional[Dict[str, str]],
mcp_server_auth_headers: Optional[Dict[str, Dict[str, str]]],
user_api_key_auth: Optional[UserAPIKeyAuth],
client_ip: Optional[str],
allowed_server_ids: Optional[Set[str]] = None,
) -> None:
"""Fail fast with HTTP 401 for MCP servers that need user auth but
didn't receive it on this request. Covers both gateway-managed OAuth2
(points clients at the gateway AS metadata) and pass-through OAuth
(points clients at the upstream resource-metadata via our well-known).
``allowed_server_ids`` may be passed by callers that have already
narrowed the authorized server set (e.g. toolset scoping); servers
not in that set are skipped so a client targeting a toolset that
excludes a passthrough server is not pushed into an OAuth flow for
a server it will be 403'd on immediately after authentication.
"""
for server_name in mcp_servers or []:
server = global_mcp_server_manager.get_mcp_server_by_name(
server_name, client_ip=client_ip
)
if (
server is not None
and allowed_server_ids is not None
and server.server_id not in allowed_server_ids
):
# Caller's narrowed scope excludes this server — skip the
# preemptive challenge and let downstream authorization
# return 403.
continue
if server and server.auth_type == MCPAuth.oauth2 and not oauth2_headers:
# For per-user OAuth servers, only skip the pre-emptive 401 when
# a stored token actually exists for this user+server pair.
# If no stored token exists, fail fast with 401 so clients can
# kick off PKCE/interactive OAuth flow immediately.
if server.needs_user_oauth_token:
stored_oauth_headers = await _get_user_oauth_extra_headers_from_db(
server=server,
user_api_key_auth=user_api_key_auth,
)
if stored_oauth_headers:
continue
request = StarletteRequest(scope)
base_url = get_request_base_url(request)
_path = scope.get("_original_path") or scope.get("path", "") or ""
# Pick the well-known AS-metadata form that matches the inbound route
# so strict RFC 9728 §3.2 clients can resolve it correctly.
if _path.startswith(f"/mcp/{server_name}"):
_as_url = f"{base_url}/.well-known/oauth-authorization-server/mcp/{server_name}"
else:
_as_url = f"{base_url}/.well-known/oauth-authorization-server/{server_name}"
authorization_uri = f'Bearer authorization_uri="{_as_url}"'
raise HTTPException(
status_code=401,
detail="Unauthorized",
headers={"www-authenticate": authorization_uri},
)
# Pass-through OAuth: when the admin has opted a server into
# forwarding the client's bearer token (is_oauth_passthrough) and
# the client hasn't supplied one, fail fast with 401 and point
# them at the gateway's oauth-protected-resource well-known URL.
# That endpoint proxies the upstream's metadata so the client
# kicks off OAuth against the real upstream IdP, not the gateway.
if (
server
and server.is_oauth_passthrough
and not _client_has_passthrough_authorization(
server, oauth2_headers, mcp_server_auth_headers
)
):
www_authenticate = _get_passthrough_www_authenticate(
scope=scope,
server_name=server_name,
)
raise HTTPException(
status_code=401,
detail="Unauthorized",
headers={"www-authenticate": www_authenticate},
)
def _get_forwarded_auth_from_scope(scope: Scope) -> Optional[str]:
"""Return the upstream-bound ``Authorization`` header value, or None.
@ -3248,12 +3440,15 @@ if MCP_AVAILABLE:
passthrough_servers = [
srv
for srv in allowed_servers
if srv.extra_headers
and any(h.lower() == "authorization" for h in srv.extra_headers)
# Exclude M2M servers: _prepare_mcp_server_headers skips caller
# Authorization when has_client_credentials is set, so probing
# those with the caller's token would send the wrong credential.
and not srv.has_client_credentials
# Restrict to genuine OAuth pass-through servers (auth_type none +
# Authorization in extra_headers). Gateway-managed OAuth2 servers
# must not receive the ``resource_metadata=`` challenge emitted
# below — they require ``authorization_uri=`` pointing at the
# gateway AS metadata. ``is_oauth_passthrough`` already requires
# ``auth_type in (None, MCPAuth.none)``, which is mutually
# exclusive with ``has_client_credentials`` (oauth2 + M2M flow),
# so M2M servers are implicitly excluded here.
if srv.is_oauth_passthrough
]
if not passthrough_servers:
return
@ -3264,19 +3459,20 @@ if MCP_AVAILABLE:
for srv in passthrough_servers
]
)
request = StarletteRequest(scope)
base_url = get_request_base_url(request)
for srv, (probe_status, _) in zip(passthrough_servers, probe_results):
if probe_status == 401:
# Token is missing or expired — direct the client to re-authorize.
authorization_uri = (
f"Bearer authorization_uri="
f"{base_url}/.well-known/oauth-authorization-server/{srv.name}"
# Token is missing or expired: keep pass-through clients on the
# protected-resource discovery flow so they re-authorize against
# the upstream IdP metadata proxied by LiteLLM.
www_authenticate = _get_passthrough_www_authenticate(
scope=scope,
server_name=srv.name,
invalid_token=True,
)
raise HTTPException(
status_code=401,
detail="Unauthorized",
headers={"WWW-Authenticate": authorization_uri},
headers={"www-authenticate": www_authenticate},
)
if probe_status == 403:
# Token is valid but the caller lacks permission — do not hint
@ -3311,39 +3507,6 @@ if MCP_AVAILABLE:
verbose_logger.debug(
f"MCP server auth headers: {list(mcp_server_auth_headers.keys()) if mcp_server_auth_headers else None}"
)
# https://datatracker.ietf.org/doc/html/rfc9728#name-www-authenticate-response
for server_name in mcp_servers or []:
server = global_mcp_server_manager.get_mcp_server_by_name(
server_name, client_ip=_client_ip
)
if server and server.auth_type == MCPAuth.oauth2 and not oauth2_headers:
# For per-user OAuth servers, only skip the pre-emptive 401 when
# a stored token actually exists for this user+server pair.
# If no stored token exists, fail fast with 401 so clients can
# kick off PKCE/interactive OAuth flow immediately.
if server.needs_user_oauth_token:
stored_oauth_headers = (
await _get_user_oauth_extra_headers_from_db(
server=server,
user_api_key_auth=user_api_key_auth,
)
)
if stored_oauth_headers:
continue
request = StarletteRequest(scope)
base_url = get_request_base_url(request)
authorization_uri = (
f"Bearer authorization_uri="
f"{base_url}/.well-known/oauth-authorization-server/{server_name}"
)
raise HTTPException(
status_code=401,
detail="Unauthorized",
headers={"www-authenticate": authorization_uri},
)
# Strip any client-supplied x-mcp-toolset-id to prevent forgery.
scope["headers"] = [
@ -3355,10 +3518,28 @@ if MCP_AVAILABLE:
# Apply toolset scope if set server-side via ContextVar (set by
# /toolset/{name}/mcp and /{name}/mcp route handlers in proxy_server.py).
active_toolset_id = _mcp_active_toolset_id.get()
toolset_allowed_server_ids: Optional[Set[str]] = None
if active_toolset_id and user_api_key_auth is not None:
user_api_key_auth = await _apply_toolset_scope(
user_api_key_auth, active_toolset_id
)
op = user_api_key_auth.object_permission
toolset_allowed_server_ids = set(op.mcp_servers or []) if op else set()
# https://datatracker.ietf.org/doc/html/rfc9728#name-www-authenticate-response
# Must run after toolset scoping so the challenge set is derived
# from the fully-authorized server set: a passthrough server that
# the active toolset excludes should not trigger an OAuth flow
# for a server the caller will be 403'd on after authentication.
await _raise_preemptive_401_for_unauthenticated_servers(
scope=scope,
mcp_servers=mcp_servers,
oauth2_headers=oauth2_headers,
mcp_server_auth_headers=mcp_server_auth_headers,
user_api_key_auth=user_api_key_auth,
client_ip=_client_ip,
allowed_server_ids=toolset_allowed_server_ids,
)
# Pre-flight auth check for pass-through servers. Must run after
# toolset scoping so the probe list is derived from the fully-authorized
@ -3589,6 +3770,13 @@ if MCP_AVAILABLE:
not in _stateful_session_auth_contexts
):
_stateful_session_locks.pop(active_request_session_id, None)
except MCPUpstreamAuthError as e:
# Pass-through server returned 401 — surface it to the client so
# standards-compliant MCP clients trigger the upstream OAuth flow.
raise e.to_http_exception(
base_url=get_request_base_url(StarletteRequest(scope)),
request_path=scope.get("_original_path") or scope.get("path"),
)
except HTTPException:
# Re-raise HTTP exceptions to preserve status codes and details
raise
@ -3632,6 +3820,50 @@ if MCP_AVAILABLE:
verbose_logger.debug(
f"MCP server auth headers: {list(mcp_server_auth_headers.keys()) if mcp_server_auth_headers else None}"
)
# Strip any client-supplied x-mcp-toolset-id to prevent forgery.
scope["headers"] = [
(k, v)
for k, v in scope.get("headers", [])
if k.lower() != b"x-mcp-toolset-id"
]
# Apply toolset scope if set server-side via ContextVar so the
# downstream probe list matches the fully-authorized server set
# (mirrors the streamable HTTP handler).
active_toolset_id = _mcp_active_toolset_id.get()
toolset_allowed_server_ids: Optional[Set[str]] = None
if active_toolset_id and user_api_key_auth is not None:
user_api_key_auth = await _apply_toolset_scope(
user_api_key_auth, active_toolset_id
)
op = user_api_key_auth.object_permission
toolset_allowed_server_ids = set(op.mcp_servers or []) if op else set()
# https://datatracker.ietf.org/doc/html/rfc9728#name-www-authenticate-response
# Must run after toolset scoping so the challenge set is derived
# from the fully-authorized server set: a passthrough server that
# the active toolset excludes should not trigger an OAuth flow
# for a server the caller will be 403'd on after authentication.
await _raise_preemptive_401_for_unauthenticated_servers(
scope=scope,
mcp_servers=mcp_servers,
oauth2_headers=oauth2_headers,
mcp_server_auth_headers=mcp_server_auth_headers,
user_api_key_auth=user_api_key_auth,
client_ip=_sse_client_ip,
allowed_server_ids=toolset_allowed_server_ids,
)
# Pre-flight auth check for pass-through servers: surface upstream
# 401/403 as a proper challenge before the SSE session commits 200
# headers, so clients can refresh their OAuth token instead of
# being stuck with a silently empty tool list. Must run after
# toolset scoping so the probe list is derived from the fully-
# authorized server set, not the raw user-supplied names.
await _check_passthrough_upstream_auth(
scope, user_api_key_auth, mcp_servers, _sse_client_ip
)
set_auth_context(
user_api_key_auth=user_api_key_auth,
mcp_auth_header=mcp_auth_header,
@ -3652,9 +3884,20 @@ if MCP_AVAILABLE:
_sse_client_ip,
):
await sse_session_manager.handle_request(scope, receive, send)
except MCPUpstreamAuthError as e:
# Pass-through server returned 401 — surface it to the client so
# standards-compliant MCP clients trigger the upstream OAuth flow.
raise e.to_http_exception(
base_url=get_request_base_url(StarletteRequest(scope)),
request_path=scope.get("_original_path") or scope.get("path"),
)
except HTTPException:
# Re-raise HTTP exceptions to preserve status codes and details
# (e.g. 401 + WWW-Authenticate challenges from OAuth pass-through).
raise
except Exception as e:
verbose_logger.exception(f"Error handling MCP request: {e}")
# Instead of re-raising, try to send a graceful error response
# Try to send a graceful error response for non-HTTP exceptions
try:
# Send a proper HTTP error response instead of letting the exception bubble up
from starlette.responses import JSONResponse

View file

@ -701,6 +701,10 @@ class LiteLLMRoutes(enum.Enum):
"/v2/guardrails/list",
"/project/list",
"/project/info",
# Read-only search tool routes power the Search Tools UI page.
# Create/update/delete and test_connection stay admin-only.
"/search_tools/list",
"/search_tools/ui/available_providers",
]
+ spend_tracking_routes
+ key_management_routes
@ -1050,6 +1054,7 @@ class GenerateRequestBase(LiteLLMPydanticObjectBase):
model_config = ConfigDict(protected_namespaces=())
model_rpm_limit: Optional[dict] = None
model_tpm_limit: Optional[dict] = None
mcp_rpm_limit: Optional[Dict[str, int]] = None
guardrails: Optional[List[str]] = None
policies: Optional[List[str]] = None
prompts: Optional[List[str]] = None
@ -1290,6 +1295,7 @@ class NewMCPServerRequest(LiteLLMPydanticObjectBase):
allow_all_keys: bool = False
available_on_public_internet: bool = True
delegate_auth_to_upstream: bool = False
oauth_passthrough: bool = False
is_byok: bool = False
byok_description: List[str] = Field(default_factory=list)
byok_api_key_help_url: Optional[str] = None
@ -1373,6 +1379,7 @@ class UpdateMCPServerRequest(LiteLLMPydanticObjectBase):
allow_all_keys: bool = False
available_on_public_internet: bool = True
delegate_auth_to_upstream: bool = False
oauth_passthrough: bool = False
is_byok: bool = False
byok_description: List[str] = Field(default_factory=list)
byok_api_key_help_url: Optional[str] = None
@ -1445,6 +1452,7 @@ class LiteLLM_MCPServerTable(LiteLLMPydanticObjectBase):
allow_all_keys: bool = False
available_on_public_internet: bool = True
delegate_auth_to_upstream: bool = False
oauth_passthrough: bool = False
is_byok: bool = False
byok_description: List[str] = Field(default_factory=list)
byok_api_key_help_url: Optional[str] = None
@ -1851,6 +1859,7 @@ class NewTeamRequest(TeamBase):
] = None # raise an error if 'guaranteed_throughput' is set and we're overallocating tpm
model_tpm_limit: Optional[Dict[str, int]] = None
mcp_rpm_limit: Optional[Dict[str, int]] = None
team_member_budget: Optional[float] = (
None # allow user to set a budget for all team members
)
@ -1920,6 +1929,7 @@ class UpdateTeamRequest(LiteLLMPydanticObjectBase):
prompts: Optional[List[str]] = None
model_rpm_limit: Optional[Dict[str, int]] = None
model_tpm_limit: Optional[Dict[str, int]] = None
mcp_rpm_limit: Optional[Dict[str, int]] = None
allowed_vector_store_indexes: Optional[List[AllowedVectorStoreIndexItem]] = None
enforced_batch_output_expires_after: Optional[dict] = None
enforced_file_expires_after: Optional[dict] = None
@ -2516,7 +2526,7 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase):
)
mcp_trusted_proxy_ranges: Optional[List[str]] = Field(
None,
description="CIDR ranges of trusted reverse proxies. When set, X-Forwarded-For headers are only trusted from these IPs.",
description="CIDR ranges of trusted reverse proxies. When set, X-Forwarded-For and X-Forwarded-* origin headers are only trusted from these IPs.",
)
trusted_proxy_ranges: Optional[List[str]] = Field(
None,
@ -3667,6 +3677,7 @@ class ProxyException(Exception):
provider_specific_fields: Optional[dict] = None,
):
self.message = str(message)
super().__init__(self.message)
self.type = type
self.param = param
self.openai_code = openai_code or code
@ -4285,6 +4296,7 @@ class PassThroughEndpointLoggingTypedDict(TypedDict):
LiteLLM_ManagementEndpoint_MetadataFields = [
"model_rpm_limit",
"model_tpm_limit",
"mcp_rpm_limit",
"rpm_limit_type",
"tpm_limit_type",
"enforced_params",
@ -4436,6 +4448,72 @@ class JWTRoutingOverride(BaseModel):
}
class JWTIssuerConfig(BaseModel):
"""
Issuer-bound JWT validation configuration.
When a token's unverified `iss` claim matches an entry in
``LiteLLM_JWTAuth.issuers``, LiteLLM validates it only against that
issuer's JWKS and audience. Tokens whose `iss` does not match any
configured issuer fall back to the global JWT_AUDIENCE/JWT_ISSUER
validation path; `issuers` is additive routing, not an allow-list.
"""
issuer: str = Field(description="Exact expected JWT issuer (`iss`) value.")
jwks_url: Optional[str] = Field(
default=None,
description="Issuer JWKS URL. If omitted, LiteLLM uses the issuer's OIDC discovery document.",
)
audience: Optional[Union[str, List[str]]] = Field(
default=None,
description="Expected token audience for this issuer.",
)
disable_audience_validation: bool = Field(
default=False,
description="Explicitly disable audience validation for this issuer. Use only when the issuer cannot provide an audience suitable for LiteLLM.",
)
user_id_jwt_field: Optional[str] = Field(
default=None,
description="Issuer-specific claim path to normalize into LiteLLM's user id.",
)
user_email_jwt_field: Optional[str] = Field(
default=None,
description="Issuer-specific claim path to normalize into LiteLLM's user email.",
)
team_id_jwt_field: Optional[str] = Field(
default=None,
description="Issuer-specific claim path to normalize into LiteLLM's team id.",
)
team_ids_jwt_field: Optional[str] = Field(
default=None,
description="Issuer-specific claim path to normalize into LiteLLM's team ids.",
)
org_id_jwt_field: Optional[str] = Field(
default=None,
description="Issuer-specific claim path to normalize into LiteLLM's organization id.",
)
end_user_id_jwt_field: Optional[str] = Field(
default=None,
description="Issuer-specific claim path to normalize into LiteLLM's end-user id.",
)
model_config = {
"extra": "forbid",
}
@model_validator(mode="after")
def validate_audience_configured(self) -> "JWTIssuerConfig":
if self.audience is None and not self.disable_audience_validation:
raise ValueError(
f"JWT issuer {self.issuer} must configure audience or set disable_audience_validation=True"
)
if self.audience is not None and self.disable_audience_validation:
raise ValueError(
f"JWT issuer {self.issuer} cannot set audience and disable_audience_validation=True together"
)
return self
class LiteLLM_JWTAuth(LiteLLMPydanticObjectBase):
"""
A class to define the roles and permissions for a LiteLLM Proxy w/ JWT Auth.
@ -4540,6 +4618,10 @@ class LiteLLM_JWTAuth(LiteLLMPydanticObjectBase):
default=None,
description="Optional claim-based routing overrides for JWT-shaped tokens. Matching rules route requests to oauth2 before default JWT flow.",
)
issuers: Optional[List[JWTIssuerConfig]] = Field(
default=None,
description="Optional issuer-bound JWT validation rules. When a token's `iss` matches a configured issuer, validation uses that issuer's JWKS, audience, and claim mappings. Tokens with an unlisted `iss` fall back to the global JWT_AUDIENCE/JWT_ISSUER validation path — this is additive routing, not an allow-list.",
)
#########################################################
def __init__(self, **kwargs: Any) -> None:

View file

@ -6,12 +6,14 @@ The A2A SDK can point to LiteLLM's URL and invoke agents registered with LiteLLM
"""
import json
from typing import Any, Dict, List, Optional
from typing import Any, AsyncGenerator, Dict, List, Optional
from urllib.parse import urlparse
from fastapi import APIRouter, Depends, HTTPException, Request, Response
from fastapi.responses import JSONResponse, StreamingResponse
from litellm._logging import verbose_proxy_logger
from litellm.litellm_core_utils.url_utils import SSRFError, validate_url
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.agent_endpoints.utils import merge_agent_headers
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
@ -19,9 +21,66 @@ from litellm.types.utils import all_litellm_params
router = APIRouter()
_PASCAL_TO_WIRE: Dict[str, str] = {
"GetTask": "tasks/get",
"ListTasks": "tasks/list",
"CancelTask": "tasks/cancel",
"SubscribeToTask": "tasks/resubscribe",
"CreateTaskPushNotificationConfig": "tasks/pushNotificationConfig/set",
"GetTaskPushNotificationConfig": "tasks/pushNotificationConfig/get",
"ListTaskPushNotificationConfigs": "tasks/pushNotificationConfig/list",
"DeleteTaskPushNotificationConfig": "tasks/pushNotificationConfig/delete",
"GetExtendedAgentCard": "agent/getAuthenticatedExtendedCard",
}
def _validate_push_notification_url(url: str) -> None:
parsed = urlparse(url)
if parsed.scheme != "https":
raise HTTPException(
status_code=400,
detail="Push notification URL must use HTTPS",
)
try:
validate_url(url)
except (SSRFError, ValueError) as e:
raise HTTPException(status_code=400, detail=str(e)) from e
def _caller_identity_headers(user_api_key_dict: UserAPIKeyAuth) -> Dict[str, str]:
headers: Dict[str, str] = {}
if user_api_key_dict.user_id:
headers["X-LiteLLM-User-Id"] = user_api_key_dict.user_id
if user_api_key_dict.team_id:
headers["X-LiteLLM-Team-Id"] = user_api_key_dict.team_id
return headers
def _forwarding_headers(
user_api_key_dict: UserAPIKeyAuth,
request_data: dict,
agent_extra_headers: Optional[Dict[str, str]],
) -> Optional[Dict[str, str]]:
sanitized = (
{
k: v
for k, v in agent_extra_headers.items()
if not k.lower().startswith("x-litellm-")
}
if agent_extra_headers
else None
)
merged = merge_agent_headers(dynamic_headers=sanitized, static_headers=None) or {}
identity = _caller_identity_headers(user_api_key_dict)
trace_id = request_data.get("litellm_trace_id")
if trace_id:
identity["X-LiteLLM-Trace-Id"] = str(trace_id)
merged.update(identity)
return merged or None
def _jsonrpc_error(
request_id: Optional[str],
request_id: Optional[Any],
code: int,
message: str,
status_code: int = 400,
@ -67,9 +126,158 @@ def _enforce_inbound_trace_id(agent: Any, request: Request) -> None:
)
async def _forward_jsonrpc(
agent_url: str,
body: dict,
extra_headers: Optional[Dict[str, str]] = None,
) -> dict:
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
from litellm.types.llms.custom_http import httpxSpecialProvider
headers = {"Content-Type": "application/json", **(extra_headers or {})}
handler = get_async_httpx_client(
llm_provider=httpxSpecialProvider.A2A,
params={"timeout": 60.0},
)
resp = await handler.post(agent_url, json=body, headers=headers)
try:
result = resp.json()
except Exception:
resp.raise_for_status()
raise
if not resp.is_success and "error" not in result:
resp.raise_for_status()
return result
async def _a2a_sse_event_source(
agent_url: str,
body: dict,
request_id: Optional[Any] = None,
extra_headers: Optional[Dict[str, str]] = None,
) -> AsyncGenerator[dict, None]:
"""Stream an upstream A2A SSE response as parsed JSON-RPC event dicts.
Upstream HTTP/JSON-RPC errors are surfaced as a single JSON-RPC error event
so the caller can relay them instead of breaking the stream.
"""
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
from litellm.types.agents import _normalize_a2a_jsonrpc_response
from litellm.types.llms.custom_http import httpxSpecialProvider
headers = {
"Content-Type": "application/json",
"Accept": "text/event-stream",
**(extra_headers or {}),
}
handler = get_async_httpx_client(
llm_provider=httpxSpecialProvider.A2A,
params={"timeout": None},
)
async_client = handler.client
req = async_client.build_request("POST", agent_url, json=body, headers=headers)
resp = await async_client.send(req, stream=True)
try:
if not resp.is_success:
error_body = await resp.aread()
error_event: Optional[dict] = None
try:
parsed = json.loads(error_body)
if isinstance(parsed, dict) and "error" in parsed:
error_event = _normalize_a2a_jsonrpc_response(
parsed, request_id=request_id
)
except Exception:
error_event = None
yield error_event or {
"jsonrpc": "2.0",
"id": request_id,
"error": {"code": -32603, "message": resp.reason_phrase},
}
return
async for line in resp.aiter_lines():
stripped = line.strip()
if not stripped.startswith("data:"):
continue
payload = stripped[len("data:") :].strip()
if not payload:
continue
try:
yield json.loads(payload)
except Exception:
continue
finally:
await resp.aclose()
async def _forward_jsonrpc_sse(
agent_url: str,
body: dict,
request_id: Optional[Any] = None,
extra_headers: Optional[Dict[str, str]] = None,
proxy_logging_obj: Optional[Any] = None,
user_api_key_dict: Optional[Any] = None,
request_data: Optional[dict] = None,
) -> StreamingResponse:
event_source = _a2a_sse_event_source(
agent_url, body, request_id=request_id, extra_headers=extra_headers
)
def _serialize_chunk(chunk: Any) -> str:
return f"data: {json.dumps(chunk)}\n\n"
def _serialize_error(proxy_exc: Any) -> str:
return (
"data: "
+ json.dumps(
{
"jsonrpc": "2.0",
"id": request_id,
"error": {
"code": -32603,
"message": getattr(proxy_exc, "message", str(proxy_exc)),
},
}
)
+ "\n\n"
)
if (
proxy_logging_obj is not None
and user_api_key_dict is not None
and request_data is not None
):
# Route streamed events through the shared streaming generator so the
# post-call streaming hook (and therefore agent guardrails) inspects
# tasks/resubscribe output the same way message/stream does.
from litellm.proxy.common_request_processing import (
ProxyBaseLLMRequestProcessing,
)
generator: AsyncGenerator[str, None] = (
ProxyBaseLLMRequestProcessing.async_streaming_data_generator(
response=event_source,
user_api_key_dict=user_api_key_dict,
request_data=request_data,
proxy_logging_obj=proxy_logging_obj,
serialize_chunk=_serialize_chunk,
serialize_error=_serialize_error,
)
)
else:
async def _passthrough() -> AsyncGenerator[str, None]:
async for chunk in event_source:
yield _serialize_chunk(chunk)
generator = _passthrough()
return StreamingResponse(generator, media_type="text/event-stream")
async def _handle_stream_message(
api_base: Optional[str],
request_id: str,
request_id: Any,
params: dict,
litellm_params: Optional[dict] = None,
agent_id: Optional[str] = None,
@ -310,8 +518,6 @@ async def invoke_agent_a2a( # noqa: PLR0915
- message/send: Send a message and get a response
- message/stream: Send a message and stream the response
"""
from litellm.a2a_protocol import asend_message
from litellm.a2a_protocol.main import A2A_SDK_AVAILABLE
from litellm.proxy.agent_endpoints.auth.agent_permission_handler import (
AgentRequestHandler,
)
@ -322,9 +528,11 @@ async def invoke_agent_a2a( # noqa: PLR0915
version,
)
body = {}
body: Dict[str, Any] = {}
request_data: Dict[str, Any] = body
try:
body = await request.json()
request_data = body
verbose_proxy_logger.debug(f"A2A request for agent '{agent_id}': {body}")
@ -334,11 +542,14 @@ async def invoke_agent_a2a( # noqa: PLR0915
body.get("id"), -32600, "Invalid Request: jsonrpc must be '2.0'"
)
request_id = body.get("id")
method = body.get("method")
request_id: Optional[Any] = body.get("id")
method: Optional[str] = body.get("method")
params = body.get("params", {})
if params:
if method:
method = _PASCAL_TO_WIRE.get(method, method)
if isinstance(params, dict):
# extract any litellm params from the params - eg. 'guardrails'
# ``metadata`` is intentionally excluded: it's a first-class A2A
# ``MessageSendParams`` field that the completion bridge forwards
@ -347,20 +558,12 @@ async def invoke_agent_a2a( # noqa: PLR0915
# silently drop the caller's A2A request-level metadata.
params_to_remove = []
for key, value in params.items():
if key in all_litellm_params and key != "metadata":
if key in all_litellm_params and key not in {"id", "metadata"}:
params_to_remove.append(key)
body[key] = value
for key in params_to_remove:
params.pop(key)
if not A2A_SDK_AVAILABLE:
return _jsonrpc_error(
request_id,
-32603,
"Server error: 'a2a' package not installed. Please install 'a2a-sdk'.",
500,
)
# Find the agent
agent = _get_agent(agent_id)
if agent is None:
@ -389,6 +592,19 @@ async def invoke_agent_a2a( # noqa: PLR0915
litellm_params = agent.litellm_params or {}
custom_llm_provider = litellm_params.get("custom_llm_provider")
# Hand the authenticated key hash to the completion bridge so provider
# configs can scope provider-side session state per key (e.g. LangFlow
# session memory) instead of trusting the client-supplied A2A contextId.
if custom_llm_provider and user_api_key_dict.api_key:
from litellm.a2a_protocol.litellm_completion_bridge.handler import (
A2A_USER_API_KEY_HASH_PARAM,
)
litellm_params = {
**litellm_params,
A2A_USER_API_KEY_HASH_PARAM: user_api_key_dict.api_key,
}
# URL is required unless using completion bridge with a provider that derives endpoint from model
# (e.g., bedrock/agentcore derives endpoint from ARN in model string)
if not agent_url and not custom_llm_provider:
@ -428,6 +644,7 @@ async def invoke_agent_a2a( # noqa: PLR0915
route_type="asend_message",
version=version,
)
request_data = data
# Build merged headers for the backend agent
static_headers: Dict[str, str] = dict(agent.static_headers or {})
@ -440,9 +657,10 @@ async def invoke_agent_a2a( # noqa: PLR0915
# 1. Admin-configured extra_headers: forward named headers from client request
if agent.extra_headers:
for header_name in agent.extra_headers:
val = normalized.get(header_name.lower())
header_name_str = str(header_name)
val = normalized.get(header_name_str.lower())
if val is not None:
dynamic_headers[header_name] = val
dynamic_headers[header_name_str] = val
# 2. Convention-based forwarding: x-a2a-{agent_id_or_name}-{header_name}
# Matches both agent_id (UUID) and agent_name (alias), case-insensitive.
@ -476,10 +694,20 @@ async def invoke_agent_a2a( # noqa: PLR0915
# Route through SDK functions
if method == "message/send":
from litellm.a2a_protocol import asend_message
from litellm.a2a_protocol.main import A2A_SDK_AVAILABLE
if not A2A_SDK_AVAILABLE:
return _jsonrpc_error(
request_id,
-32603,
"Server error: 'a2a' package not installed. Please install 'a2a-sdk'.",
500,
)
from a2a.types import MessageSendParams, SendMessageRequest
a2a_request = SendMessageRequest(
id=request_id,
id=request_id if request_id is not None else "",
params=MessageSendParams(**params),
)
# Defer spend-log until after post_call_success_hook so guardrail
@ -519,7 +747,7 @@ async def invoke_agent_a2a( # noqa: PLR0915
elif method == "message/stream":
return await _handle_stream_message(
api_base=agent_url,
request_id=request_id,
request_id=request_id if request_id is not None else "",
params=params,
litellm_params=litellm_params,
agent_id=agent.agent_id,
@ -530,6 +758,106 @@ async def invoke_agent_a2a( # noqa: PLR0915
request_data=data,
proxy_logging_obj=proxy_logging_obj,
)
elif method in {
"tasks/get",
"tasks/list",
"tasks/cancel",
"tasks/pushNotificationConfig/set",
"tasks/pushNotificationConfig/get",
"tasks/pushNotificationConfig/list",
"tasks/pushNotificationConfig/delete",
"agent/getAuthenticatedExtendedCard",
}:
if not agent_url:
return _jsonrpc_error(
request_id, -32000, f"Agent '{agent_id}' has no URL configured", 500
)
if method == "tasks/pushNotificationConfig/set":
if not isinstance(params, dict):
raise HTTPException(
status_code=400,
detail="params must be an object",
)
push_config = params.get("pushNotificationConfig", {})
if "pushNotificationConfig" in params and not isinstance(
push_config, dict
):
raise HTTPException(
status_code=400,
detail="pushNotificationConfig must be an object",
)
for callback_url in (params.get("url"), push_config.get("url")):
if not callback_url:
continue
if not isinstance(callback_url, str):
raise HTTPException(
status_code=400,
detail="Push notification URL must be a string",
)
_validate_push_notification_url(callback_url)
forward_body = {
"jsonrpc": "2.0",
"id": request_id,
"method": method,
"params": params,
}
caller_headers = _forwarding_headers(
user_api_key_dict=user_api_key_dict,
request_data=data,
agent_extra_headers=agent_extra_headers,
)
result = await _forward_jsonrpc(
agent_url, forward_body, extra_headers=caller_headers
)
if method == "agent/getAuthenticatedExtendedCard":
if isinstance(result.get("result"), dict) and "url" in result["result"]:
result["result"][
"url"
] = f"{str(request.base_url).rstrip('/')}/a2a/{agent_id}"
from litellm.types.agents import LiteLLMSendMessageResponse
response = LiteLLMSendMessageResponse.from_dict(
result, request_id=request_id
)
response = await proxy_logging_obj.post_call_success_hook(
user_api_key_dict=user_api_key_dict,
data=data,
response=response,
)
return JSONResponse(
content=(
response.model_dump(mode="json", exclude_none=True)
if hasattr(response, "model_dump")
else response
)
)
elif method == "tasks/resubscribe":
if not agent_url:
return _jsonrpc_error(
request_id, -32000, f"Agent '{agent_id}' has no URL configured", 500
)
forward_body = {
"jsonrpc": "2.0",
"id": request_id,
"method": method,
"params": params,
}
sse_caller_headers = _forwarding_headers(
user_api_key_dict=user_api_key_dict,
request_data=data,
agent_extra_headers=agent_extra_headers,
)
return await _forward_jsonrpc_sse(
agent_url,
forward_body,
request_id=request_id,
extra_headers=sse_caller_headers,
proxy_logging_obj=proxy_logging_obj,
user_api_key_dict=user_api_key_dict,
request_data=data,
)
else:
return _jsonrpc_error(request_id, -32601, f"Method '{method}' not found")
@ -537,4 +865,12 @@ async def invoke_agent_a2a( # noqa: PLR0915
raise
except Exception as e:
verbose_proxy_logger.exception(f"Error invoking agent: {e}")
try:
await proxy_logging_obj.post_call_failure_hook(
user_api_key_dict=user_api_key_dict,
original_exception=e,
request_data=request_data,
)
except Exception:
pass
return _jsonrpc_error(body.get("id"), -32603, f"Internal error: {str(e)}", 500)

View file

@ -319,36 +319,34 @@ async def create_agent(
Example Request:
```bash
curl -X POST "http://localhost:4000/agents" \\
curl -X POST "http://localhost:4000/v1/agents" \\
-H "Authorization: Bearer <your_api_key>" \\
-H "Content-Type: application/json" \\
-d '{
"agent": {
"agent_name": "my-custom-agent",
"agent_card_params": {
"protocolVersion": "1.0",
"name": "Hello World Agent",
"description": "Just a hello world agent",
"url": "http://localhost:9999/",
"version": "1.0.0",
"defaultInputModes": ["text"],
"defaultOutputModes": ["text"],
"capabilities": {
"streaming": true
},
"skills": [
{
"id": "hello_world",
"name": "Returns hello world",
"description": "just returns hello world",
"tags": ["hello world"],
"examples": ["hi", "hello world"]
}
]
"agent_name": "my-custom-agent",
"agent_card_params": {
"protocolVersion": "1.0",
"name": "Hello World Agent",
"description": "Just a hello world agent",
"url": "http://localhost:9999/",
"version": "1.0.0",
"defaultInputModes": ["text"],
"defaultOutputModes": ["text"],
"capabilities": {
"streaming": true
},
"litellm_params": {
"make_public": true
}
"skills": [
{
"id": "hello_world",
"name": "Returns hello world",
"description": "just returns hello world",
"tags": ["hello world"],
"examples": ["hi", "hello world"]
}
]
},
"litellm_params": {
"make_public": true
}
}'
```
@ -441,7 +439,7 @@ async def get_agent_by_id(
Example Request:
```bash
curl -X GET "http://localhost:4000/agents/123e4567-e89b-12d3-a456-426614174000" \\
curl -X GET "http://localhost:4000/v1/agents/123e4567-e89b-12d3-a456-426614174000" \\
-H "Authorization: Bearer <your_api_key>"
```
"""
@ -535,28 +533,26 @@ async def update_agent(
Example Request:
```bash
curl -X PUT "http://localhost:4000/agents/123e4567-e89b-12d3-a456-426614174000" \\
curl -X PUT "http://localhost:4000/v1/agents/123e4567-e89b-12d3-a456-426614174000" \\
-H "Authorization: Bearer <your_api_key>" \\
-H "Content-Type: application/json" \\
-d '{
"agent": {
"agent_name": "updated-agent",
"agent_card_params": {
"protocolVersion": "1.0",
"name": "Updated Agent",
"description": "Updated description",
"url": "http://localhost:9999/",
"version": "1.1.0",
"defaultInputModes": ["text"],
"defaultOutputModes": ["text"],
"capabilities": {
"streaming": true
},
"skills": []
"agent_name": "updated-agent",
"agent_card_params": {
"protocolVersion": "1.0",
"name": "Updated Agent",
"description": "Updated description",
"url": "http://localhost:9999/",
"version": "1.1.0",
"defaultInputModes": ["text"],
"defaultOutputModes": ["text"],
"capabilities": {
"streaming": true
},
"litellm_params": {
"make_public": false
}
"skills": []
},
"litellm_params": {
"make_public": false
}
}'
```
@ -645,28 +641,26 @@ async def patch_agent(
Example Request:
```bash
curl -X PUT "http://localhost:4000/agents/123e4567-e89b-12d3-a456-426614174000" \\
curl -X PATCH "http://localhost:4000/v1/agents/123e4567-e89b-12d3-a456-426614174000" \\
-H "Authorization: Bearer <your_api_key>" \\
-H "Content-Type: application/json" \\
-d '{
"agent": {
"agent_name": "updated-agent",
"agent_card_params": {
"protocolVersion": "1.0",
"name": "Updated Agent",
"description": "Updated description",
"url": "http://localhost:9999/",
"version": "1.1.0",
"defaultInputModes": ["text"],
"defaultOutputModes": ["text"],
"capabilities": {
"streaming": true
},
"skills": []
"agent_name": "updated-agent",
"agent_card_params": {
"protocolVersion": "1.0",
"name": "Updated Agent",
"description": "Updated description",
"url": "http://localhost:9999/",
"version": "1.1.0",
"defaultInputModes": ["text"],
"defaultOutputModes": ["text"],
"capabilities": {
"streaming": true
},
"litellm_params": {
"make_public": false
}
"skills": []
},
"litellm_params": {
"make_public": false
}
}'
```
@ -753,7 +747,7 @@ async def delete_agent(
Example Request:
```bash
curl -X DELETE "http://localhost:4000/agents/123e4567-e89b-12d3-a456-426614174000" \\
curl -X DELETE "http://localhost:4000/v1/agents/123e4567-e89b-12d3-a456-426614174000" \\
-H "Authorization: Bearer <your_api_key>"
```

View file

@ -546,6 +546,7 @@ async def common_checks( # noqa: PLR0915
route=route,
request_headers=_safe_get_request_headers(request=request),
request_query_params=_safe_get_request_query_params(request=request),
llm_router=llm_router,
)
if route in MODEL_DISCOVERY_ROUTES:
@ -1126,7 +1127,7 @@ async def get_end_user_object(
end_user_id: Optional[str],
prisma_client: Optional[PrismaClient],
user_api_key_cache: UserApiKeyCache,
route: str,
route: Optional[str] = "",
parent_otel_span: Optional[Span] = None,
proxy_logging_obj: Optional[ProxyLogging] = None,
) -> Optional[LiteLLM_EndUserTable]:
@ -1170,9 +1171,6 @@ async def get_end_user_object(
parent_otel_span=parent_otel_span,
)
# Check budget limits
await _check_end_user_budget(end_user_obj=return_obj, route=route)
return return_obj
# Fetch from database
@ -1203,14 +1201,9 @@ async def get_end_user_object(
model_type=LiteLLM_EndUserTable,
)
# Check budget limits
await _check_end_user_budget(end_user_obj=_response, route=route)
return _response
except Exception as e:
if isinstance(e, litellm.BudgetExceededError):
raise e
except Exception:
return None
@ -1307,8 +1300,6 @@ async def _end_user_id_exists_in_db(
)
if end_user_obj is not None:
return True
except litellm.BudgetExceededError:
raise
except Exception as e:
verbose_proxy_logger.debug(
f"end_user validation: get_end_user_object lookup failed: {e}"
@ -3428,6 +3419,29 @@ async def can_team_call_search_tool(
)
async def can_user_view_search_tool(
search_tool_name: str,
valid_token: UserAPIKeyAuth,
team_object: Optional[LiteLLM_TeamTable],
) -> bool:
"""
Boolean variant of the key + team authorization enforced on /search, used to
scope /search_tools/list so a non-admin caller only sees tools it may invoke.
"""
try:
await can_key_call_search_tool(
search_tool_name=search_tool_name,
valid_token=valid_token,
)
await can_team_call_search_tool(
search_tool_name=search_tool_name,
team_object=team_object,
)
except ProxyException:
return False
return True
async def is_valid_fallback_model(
model: str,
llm_router: Optional[Router],

View file

@ -522,14 +522,20 @@ def get_request_route(request: Request) -> str:
if not isinstance(scope, dict):
return str(request.url.path)
raw_path: str = str(scope.get("path", request.url.path))
root_path: str = str(scope.get("app_root_path", scope.get("root_path", "")))
root_path: str = str(
scope.get("app_root_path", scope.get("root_path", ""))
).rstrip("/")
if not isinstance(raw_path, str):
return str(request.url.path)
# Only strip root_path when it is a meaningful prefix (not bare "/").
# Stripping bare "/" would remove the leading slash from every path
# e.g. "/team/new" → "team/new", breaking route matching.
if root_path and root_path != "/" and raw_path.startswith(root_path):
return raw_path[len(root_path) :]
# Strip root_path only when it matches whole path segments — guarding
# against sibling paths like "/apifoo" being truncated under
# root_path="/api". Trailing slashes on root_path are stripped above,
# so bare "/" or "/prefix/" still leave the leading "/" intact.
if root_path and (
raw_path == root_path or raw_path.startswith(root_path + "/")
):
stripped = raw_path[len(root_path) :]
return stripped or "/"
return raw_path
except Exception as e:
verbose_proxy_logger.debug(
@ -934,6 +940,40 @@ def get_team_model_tpm_limit(
return None
def get_key_mcp_rpm_limit(
user_api_key_dict: UserAPIKeyAuth,
) -> Optional[Dict[str, int]]:
"""
Get the per-MCP-server rpm limit for a given api key.
Priority order (returns first found):
1. Key metadata (mcp_rpm_limit)
2. Team metadata (mcp_rpm_limit)
The returned dict is keyed by MCP server name (alias if set, else the
configured server name).
"""
if user_api_key_dict.metadata:
result = user_api_key_dict.metadata.get("mcp_rpm_limit")
if result is not None:
return result
if user_api_key_dict.team_metadata:
team_limit = user_api_key_dict.team_metadata.get("mcp_rpm_limit")
if team_limit is not None:
return team_limit
return None
def get_team_mcp_rpm_limit(
user_api_key_dict: UserAPIKeyAuth,
) -> Optional[Dict[str, int]]:
if user_api_key_dict.team_metadata:
return user_api_key_dict.team_metadata.get("mcp_rpm_limit")
return None
def get_project_model_rpm_limit(
user_api_key_dict: UserAPIKeyAuth,
) -> Optional[Dict[str, int]]:
@ -1244,7 +1284,9 @@ def _route_uses_model_routing_sources(route: str) -> bool:
def _extract_models_from_managed_resource_id(
resource_id: Any, resource_id_field: Optional[str] = None
resource_id: Any,
resource_id_field: Optional[str] = None,
llm_router: Optional[Router] = None,
) -> List[str]:
if not isinstance(resource_id, str) or not resource_id:
return []
@ -1301,16 +1343,18 @@ def _extract_models_from_managed_resource_id(
)
if resource_id_field == "video_id":
model_id = decode_video_id_with_provider(resource_id).get("model_id")
_append_model_candidates(
candidates=candidates,
value=decode_video_id_with_provider(resource_id).get("model_id"),
value=_resolve_model_id_with_router(model_id, llm_router),
)
else:
model_id = decode_character_id_with_provider(resource_id).get(
"model_id"
)
_append_model_candidates(
candidates=candidates,
value=decode_character_id_with_provider(resource_id).get(
"model_id"
),
value=_resolve_model_id_with_router(model_id, llm_router),
)
except Exception as e:
verbose_proxy_logger.debug(
@ -1320,11 +1364,26 @@ def _extract_models_from_managed_resource_id(
return _dedupe_model_candidates(candidates)
def _resolve_model_id_with_router(
model_id: Optional[str], llm_router: Optional[Router]
) -> Optional[str]:
if model_id is None or llm_router is None:
return model_id
try:
return llm_router.resolve_model_name_from_model_id(model_id) or model_id
except Exception as e:
verbose_proxy_logger.debug(
"Unable to resolve model_id from managed resource ID: %s", str(e)
)
return model_id
def _extract_model_candidates_from_request(
request_data: dict,
route: str,
request_headers: Optional[Mapping[str, Any]] = None,
request_query_params: Optional[Mapping[str, Any]] = None,
llm_router: Optional[Router] = None,
) -> List[str]:
candidates: List[str] = []
uses_model_routing_sources = _route_uses_model_routing_sources(route=route)
@ -1374,7 +1433,9 @@ def _extract_model_candidates_from_request(
_append_model_candidates(
candidates,
_extract_models_from_managed_resource_id(
request_data.get(field), resource_id_field=field
request_data.get(field),
resource_id_field=field,
llm_router=llm_router,
),
)
@ -1396,12 +1457,14 @@ def get_model_from_request(
route: str,
request_headers: Optional[Mapping[str, Any]] = None,
request_query_params: Optional[Mapping[str, Any]] = None,
llm_router: Optional[Router] = None,
) -> Optional[Union[str, List[str]]]:
candidates = _extract_model_candidates_from_request(
request_data=request_data,
route=route,
request_headers=request_headers,
request_query_params=request_query_params,
llm_router=llm_router,
)
model = _format_model_candidates(candidates)

View file

@ -12,12 +12,12 @@ import fnmatch
import hashlib
import os
import re
from typing import Any, List, Literal, Optional, Set, Tuple, cast
from typing import Any, List, Literal, Optional, Set, Tuple, Union, cast
from cryptography import x509
from cryptography.hazmat.backends import default_backend
from cryptography.hazmat.primitives import serialization
from fastapi import HTTPException
from fastapi import HTTPException, status
import jwt
from jwt.api_jwk import PyJWK
@ -29,6 +29,7 @@ from litellm.proxy._types import (
RBAC_ROLES,
JWKKeyValue,
JWTAuthBuilderResult,
JWTIssuerConfig,
JWTKeyItem,
LiteLLM_EndUserTable,
LiteLLM_JWTAuth,
@ -66,6 +67,10 @@ from .auth_checks import (
)
class NoMatchingJWTPublicKeyError(Exception):
"""Raised when a JWKS endpoint returns no key matching the requested ``kid``."""
class JWTHandler:
"""
- treat the sub id passed in as the user id
@ -91,6 +96,22 @@ class JWTHandler:
"ES512",
"EdDSA",
]
LITELLM_JWT_ISSUER_CLAIM = "_litellm_jwt_issuer"
LITELLM_USER_ID_CLAIM = "_litellm_user_id"
LITELLM_USER_EMAIL_CLAIM = "_litellm_user_email"
LITELLM_TEAM_ID_CLAIM = "_litellm_team_id"
LITELLM_TEAM_IDS_CLAIM = "_litellm_team_ids"
LITELLM_ORG_ID_CLAIM = "_litellm_org_id"
LITELLM_END_USER_ID_CLAIM = "_litellm_end_user_id"
LITELLM_INTERNAL_CLAIMS = (
LITELLM_JWT_ISSUER_CLAIM,
LITELLM_USER_ID_CLAIM,
LITELLM_USER_EMAIL_CLAIM,
LITELLM_TEAM_ID_CLAIM,
LITELLM_TEAM_IDS_CLAIM,
LITELLM_ORG_ID_CLAIM,
LITELLM_END_USER_ID_CLAIM,
)
def __init__(
self,
@ -213,7 +234,33 @@ class JWTHandler:
return True
return False
def _is_trusted_issuer_normalized_token(self, token: dict) -> bool:
issuer = token.get(self.LITELLM_JWT_ISSUER_CLAIM)
if not isinstance(issuer, str) or not issuer:
return False
litellm_jwtauth = getattr(self, "litellm_jwtauth", None)
issuer_configs = getattr(litellm_jwtauth, "issuers", None) or []
return any(issuer_config.issuer == issuer for issuer_config in issuer_configs)
def _has_trusted_issuer_normalized_claim(self, token: dict, claim: str) -> bool:
return self._is_trusted_issuer_normalized_token(token=token) and claim in token
def get_team_ids_from_jwt(self, token: dict) -> List[str]:
if self._has_trusted_issuer_normalized_claim(
token=token, claim=self.LITELLM_TEAM_IDS_CLAIM
):
issuer_team_ids = token.get(self.LITELLM_TEAM_IDS_CLAIM)
if isinstance(issuer_team_ids, list):
return issuer_team_ids
if isinstance(issuer_team_ids, str):
return [issuer_team_ids]
# Issuer-scoped claim exists but has an unexpected type
# (e.g. int/dict from an unusual upstream mapping). Don't silently
# fall through to the global ``team_ids_jwt_field`` path — that
# would read a semantically unrelated claim on the same token.
return []
if self.litellm_jwtauth.team_ids_jwt_field is not None:
team_ids: Optional[List[str]] = get_nested_value(
data=token,
@ -242,12 +289,18 @@ class JWTHandler:
default-team behavior should still go through ``get_team_id``.
"""
team_ids: List[str] = list(self.get_team_ids_from_jwt(token))
if self.litellm_jwtauth.team_id_jwt_field is not None:
singular: Any = None
if self._has_trusted_issuer_normalized_claim(
token=token, claim=self.LITELLM_TEAM_ID_CLAIM
):
singular = token.get(self.LITELLM_TEAM_ID_CLAIM)
elif self.litellm_jwtauth.team_id_jwt_field is not None:
singular = get_nested_value(
data=token,
key_path=self.litellm_jwtauth.team_id_jwt_field,
default=None,
)
if singular is not None:
if isinstance(singular, list):
for item in singular:
if item is None:
@ -262,6 +315,11 @@ class JWTHandler:
def get_end_user_id(
self, token: dict, default_value: Optional[str]
) -> Optional[str]:
if self._has_trusted_issuer_normalized_claim(
token=token, claim=self.LITELLM_END_USER_ID_CLAIM
):
return token.get(self.LITELLM_END_USER_ID_CLAIM)
try:
if self.litellm_jwtauth.end_user_id_jwt_field is not None:
user_id = get_nested_value(
@ -303,6 +361,14 @@ class JWTHandler:
return False
def get_team_id(self, token: dict, default_value: Optional[str]) -> Optional[str]:
if self._has_trusted_issuer_normalized_claim(
token=token, claim=self.LITELLM_TEAM_ID_CLAIM
):
team_id = token.get(self.LITELLM_TEAM_ID_CLAIM)
if isinstance(team_id, list):
return team_id[0] if team_id else default_value
return team_id
try:
if self.litellm_jwtauth.team_id_jwt_field is not None:
# Use a sentinel value to detect if the path actually exists
@ -376,6 +442,11 @@ class JWTHandler:
return self.litellm_jwtauth.user_id_upsert
def get_user_id(self, token: dict, default_value: Optional[str]) -> Optional[str]:
if self._has_trusted_issuer_normalized_claim(
token=token, claim=self.LITELLM_USER_ID_CLAIM
):
return token.get(self.LITELLM_USER_ID_CLAIM)
try:
if self.litellm_jwtauth.user_id_jwt_field is not None:
user_id = get_nested_value(
@ -467,6 +538,11 @@ class JWTHandler:
def get_user_email(
self, token: dict, default_value: Optional[str]
) -> Optional[str]:
if self._has_trusted_issuer_normalized_claim(
token=token, claim=self.LITELLM_USER_EMAIL_CLAIM
):
return token.get(self.LITELLM_USER_EMAIL_CLAIM)
try:
if self.litellm_jwtauth.user_email_jwt_field is not None:
user_email = get_nested_value(
@ -495,6 +571,11 @@ class JWTHandler:
return object_id
def get_org_id(self, token: dict, default_value: Optional[str]) -> Optional[str]:
if self._has_trusted_issuer_normalized_claim(
token=token, claim=self.LITELLM_ORG_ID_CLAIM
):
return token.get(self.LITELLM_ORG_ID_CLAIM)
try:
if self.litellm_jwtauth.org_id_jwt_field is not None:
org_id = get_nested_value(
@ -590,55 +671,77 @@ class JWTHandler:
await self.user_api_key_cache.async_set_cache(
key=cache_key,
value=jwks_uri,
ttl=self.litellm_jwtauth.public_key_ttl,
ttl=self._get_public_key_cache_ttl(),
)
return jwks_uri
def _get_public_key_cache_ttl(self) -> float:
litellm_jwtauth = getattr(self, "litellm_jwtauth", None)
if litellm_jwtauth is None:
return 600
return litellm_jwtauth.public_key_ttl
async def _get_public_key_from_jwks_url(
self, jwks_url: str, kid: Optional[str]
) -> dict:
resolved_jwks_url = await self._resolve_jwks_url(jwks_url)
cache_key = f"litellm_jwt_auth_keys_{resolved_jwks_url}"
cached_keys = await self.user_api_key_cache.async_get_cache(cache_key)
if cached_keys is None:
response = await self.http_handler.get(resolved_jwks_url)
try:
response_json = response.json()
except Exception as e:
verbose_proxy_logger.error(
f"Error parsing response: {e}. Original Response: {response.text}"
)
raise Exception(
f"Error parsing response: {e}. Check server logs for original response."
)
if "keys" in response_json:
keys: JWKKeyValue = response_json["keys"]
else:
keys = response_json
await self.user_api_key_cache.async_set_cache(
key=cache_key,
value=keys,
ttl=self._get_public_key_cache_ttl(),
)
else:
keys = cached_keys
public_key = self.parse_keys(keys=keys, kid=kid)
if public_key is not None:
return cast(dict, public_key)
raise NoMatchingJWTPublicKeyError(
f"No matching public key found. keys={resolved_jwks_url}, kid={kid}"
)
async def get_public_key(self, kid: Optional[str]) -> dict:
keys_url = os.getenv("JWT_PUBLIC_KEY_URL")
if keys_url is None:
raise Exception("Missing JWT Public Key URL from environment.")
keys_url_list = [url.strip() for url in keys_url.split(",")]
keys_url_list = [url.strip() for url in keys_url.split(",") if url.strip()]
for key_url in keys_url_list:
key_url = await self._resolve_jwks_url(key_url)
cache_key = f"litellm_jwt_auth_keys_{key_url}"
cached_keys = await self.user_api_key_cache.async_get_cache(cache_key)
if cached_keys is None:
response = await self.http_handler.get(key_url)
try:
response_json = response.json()
except Exception as e:
verbose_proxy_logger.error(
f"Error parsing response: {e}. Original Response: {response.text}"
)
raise Exception(
f"Error parsing response: {e}. Check server logs for original response."
)
if "keys" in response_json:
keys: JWKKeyValue = response.json()["keys"]
else:
keys = response_json
await self.user_api_key_cache.async_set_cache(
key=cache_key,
value=keys,
ttl=self.litellm_jwtauth.public_key_ttl, # cache for 10 mins
try:
return await self._get_public_key_from_jwks_url(
jwks_url=key_url, kid=kid
)
except NoMatchingJWTPublicKeyError as e:
verbose_proxy_logger.debug(
"JWT Auth: No matching public key found at %s: %s", key_url, e
)
else:
keys = cached_keys
public_key = self.parse_keys(keys=keys, kid=kid)
if public_key is not None:
return cast(dict, public_key)
raise Exception(
raise NoMatchingJWTPublicKeyError(
f"No matching public key found. keys={keys_url_list}, kid={kid}"
)
@ -753,6 +856,11 @@ class JWTHandler:
minted by other applications that share the same IdP signing keys.
When both are unset PyJWT only checks the signature and expiry, which
is preserved for backward compatibility but logged once as a warning.
The warning fires even in mixed deployments that also configure
``LiteLLM_JWTAuth.issuers``: tokens whose ``iss`` does not match any
configured issuer fall through to this global path, and if env-var
scoping is absent that fallback is itself unscoped.
"""
audience = os.getenv("JWT_AUDIENCE")
issuer = os.getenv("JWT_ISSUER")
@ -782,77 +890,230 @@ class JWTHandler:
"options": options or None,
}
async def auth_jwt(self, token: str) -> dict:
decode_kwargs = self._build_decode_kwargs()
def _get_configured_issuer(self, token: str) -> Optional[JWTIssuerConfig]:
litellm_jwtauth = getattr(self, "litellm_jwtauth", None)
if litellm_jwtauth is None:
return None
issuer_configs = litellm_jwtauth.issuers
if not issuer_configs:
return None
claims = self.get_unverified_claims(token=token)
if claims is None:
return None
issuer = claims.get("iss")
if not isinstance(issuer, str) or not issuer:
return None
for issuer_config in issuer_configs:
if issuer_config.issuer == issuer:
return issuer_config
return None
def _get_jwks_url_for_issuer(self, issuer_config: JWTIssuerConfig) -> str:
if issuer_config.jwks_url:
return issuer_config.jwks_url
# _resolve_jwks_url fetches this OIDC discovery document and follows
# its jwks_uri, matching JWTIssuerConfig.jwks_url's documented fallback.
return f"{issuer_config.issuer.rstrip('/')}/.well-known/openid-configuration"
def _get_claim_value_for_issuer_mapping(self, token: dict, claim_field: str) -> Any:
"""Resolve a mapped claim from ``token``.
Returns ``None`` when the field is absent or empty so that mapped claims
behave like the global ``litellm_jwtauth`` path — present claims override
the normalised value, missing ones simply leave it ``None``.
"""
sentinel = object()
claim_value = get_nested_value(
data=token,
key_path=claim_field,
default=sentinel,
)
if claim_value is sentinel or claim_value is None or claim_value == "":
return None
return claim_value
def _apply_issuer_claim_mappings(
self, token: dict, issuer_config: JWTIssuerConfig
) -> dict:
normalized: dict = {
k: v for k, v in token.items() if k not in self.LITELLM_INTERNAL_CLAIMS
}
normalized[self.LITELLM_JWT_ISSUER_CLAIM] = issuer_config.issuer
claim_mappings = [
(issuer_config.user_id_jwt_field, self.LITELLM_USER_ID_CLAIM),
(issuer_config.user_email_jwt_field, self.LITELLM_USER_EMAIL_CLAIM),
(issuer_config.team_id_jwt_field, self.LITELLM_TEAM_ID_CLAIM),
(issuer_config.team_ids_jwt_field, self.LITELLM_TEAM_IDS_CLAIM),
(issuer_config.org_id_jwt_field, self.LITELLM_ORG_ID_CLAIM),
(issuer_config.end_user_id_jwt_field, self.LITELLM_END_USER_ID_CLAIM),
]
for source_claim, normalized_claim in claim_mappings:
if source_claim is None:
continue
claim_value = self._get_claim_value_for_issuer_mapping(
token=token,
claim_field=source_claim,
)
if claim_value is not None:
normalized[normalized_claim] = claim_value
return normalized
def _get_jwk_from_public_key(self, public_key: dict) -> dict:
jwk = {}
for key in ["kty", "kid", "n", "e", "x", "y", "crv"]:
if key in public_key:
jwk[key] = public_key[key]
return jwk
def _get_decode_options(
self,
audience: Optional[Union[str, List[str]]],
issuer: Optional[str] = None,
disable_audience_validation: bool = False,
) -> Optional[dict]:
# Disabling audience verification must be an explicit choice — never
# an implicit consequence of ``audience`` being None. Otherwise a
# caller that accidentally constructs a config with ``audience=None``
# (bypassing the model validator) would silently lose audience
# validation. Require callers to opt in via
# ``disable_audience_validation=True``.
if audience is None and not disable_audience_validation:
raise ValueError(
"audience must be provided unless disable_audience_validation=True"
)
options: dict = {}
if audience is None:
options["verify_aud"] = False
if issuer is None:
options["verify_iss"] = False
return options or None
def _decode_jwt_with_public_key(
self,
token: str,
public_key: Union[dict, str],
audience: Optional[Union[str, List[str]]],
issuer: Optional[str] = None,
options: Optional[dict] = None,
disable_audience_validation: bool = False,
) -> dict:
decode_options = (
options
if options is not None
else self._get_decode_options(
audience=audience,
issuer=issuer,
disable_audience_validation=disable_audience_validation,
)
)
if isinstance(public_key, dict):
public_key_obj = PyJWK.from_dict(
self._get_jwk_from_public_key(public_key=public_key)
).key
return jwt.decode(
token,
public_key_obj, # type: ignore
algorithms=self.SUPPORTED_JWT_ALGORITHMS,
options=decode_options, # type: ignore[arg-type]
audience=audience,
issuer=issuer,
leeway=self.leeway,
)
cert = x509.load_pem_x509_certificate(public_key.encode(), default_backend())
key = cert.public_key().public_bytes(
serialization.Encoding.PEM,
serialization.PublicFormat.SubjectPublicKeyInfo,
)
return jwt.decode(
token,
key,
algorithms=self.SUPPORTED_JWT_ALGORITHMS,
audience=audience,
issuer=issuer,
options=decode_options, # type: ignore[arg-type]
leeway=self.leeway,
)
async def _auth_jwt_with_issuer(
self, token: str, issuer_config: JWTIssuerConfig, kid: Optional[str]
) -> dict:
public_key = await self._get_public_key_from_jwks_url(
jwks_url=self._get_jwks_url_for_issuer(issuer_config=issuer_config),
kid=kid,
)
try:
payload = self._decode_jwt_with_public_key(
token=token,
public_key=public_key,
audience=issuer_config.audience,
issuer=issuer_config.issuer,
disable_audience_validation=issuer_config.disable_audience_validation,
)
except jwt.ExpiredSignatureError:
raise ProxyException(
message="Token Expired",
type=ProxyErrorTypes.expired_key,
param=None,
code=status.HTTP_401_UNAUTHORIZED,
)
except Exception as e:
raise Exception(f"Validation fails: {str(e)}")
return self._apply_issuer_claim_mappings(
token=payload,
issuer_config=issuer_config,
)
async def auth_jwt(self, token: str) -> dict:
header = jwt.get_unverified_header(token)
verbose_proxy_logger.debug("header: %s", header)
kid = header.get("kid", None)
issuer_config = self._get_configured_issuer(token=token)
if issuer_config is not None:
return await self._auth_jwt_with_issuer(
token=token,
issuer_config=issuer_config,
kid=kid,
)
decode_kwargs = self._build_decode_kwargs()
public_key = await self.get_public_key(kid=kid)
if public_key is not None and isinstance(public_key, dict):
jwk = {}
if "kty" in public_key:
jwk["kty"] = public_key["kty"]
if "kid" in public_key:
jwk["kid"] = public_key["kid"]
if "n" in public_key:
jwk["n"] = public_key["n"]
if "e" in public_key:
jwk["e"] = public_key["e"]
if "x" in public_key:
jwk["x"] = public_key["x"]
if "y" in public_key:
jwk["y"] = public_key["y"]
if "crv" in public_key:
jwk["crv"] = public_key["crv"]
# parse RSA/EC/OKP keys
public_key_obj = PyJWK.from_dict(jwk).key
if public_key is not None:
try:
# decode the token using the public key
payload = jwt.decode(
token,
public_key_obj, # type: ignore
algorithms=self.SUPPORTED_JWT_ALGORITHMS,
leeway=self.leeway, # allow testing of expired tokens
**decode_kwargs,
payload = self._decode_jwt_with_public_key(
token=token,
public_key=public_key,
audience=decode_kwargs["audience"],
issuer=decode_kwargs["issuer"],
options=decode_kwargs["options"],
)
return payload
return {
k: v
for k, v in payload.items()
if k not in self.LITELLM_INTERNAL_CLAIMS
}
except jwt.ExpiredSignatureError:
# the token is expired, do something to refresh it
raise Exception("Token Expired")
except Exception as e:
raise Exception(f"Validation fails: {str(e)}")
elif public_key is not None and isinstance(public_key, str):
try:
cert = x509.load_pem_x509_certificate(
public_key.encode(), default_backend()
raise ProxyException(
message="Token Expired",
type=ProxyErrorTypes.expired_key,
param=None,
code=status.HTTP_401_UNAUTHORIZED,
)
# Extract public key
key = cert.public_key().public_bytes(
serialization.Encoding.PEM,
serialization.PublicFormat.SubjectPublicKeyInfo,
)
# decode the token using the public key
payload = jwt.decode(
token,
key,
algorithms=self.SUPPORTED_JWT_ALGORITHMS,
**decode_kwargs,
)
return payload
except jwt.ExpiredSignatureError:
# the token is expired, do something to refresh it
raise Exception("Token Expired")
except Exception as e:
raise Exception(f"Validation fails: {str(e)}")

View file

@ -153,8 +153,9 @@ class IPAddressUtils:
verbose_proxy_logger.warning(
"use_x_forwarded_for is enabled but mcp_trusted_proxy_ranges "
"is not configured. X-Forwarded-* headers will NOT be "
"trusted, so MCP OAuth discovery URLs will use the proxy's "
"literal base URL. Set mcp_trusted_proxy_ranges in "
"trusted, so MCP OAuth discovery URLs and access-control "
"client IPs will use the proxy's literal request values. "
"Set mcp_trusted_proxy_ranges in "
"general_settings to your reverse-proxy CIDR(s) to allow "
"X-Forwarded-* through."
)
@ -199,17 +200,19 @@ class IPAddressUtils:
# If XFF is enabled, validate the request comes from a trusted proxy
if use_xff and "x-forwarded-for" in request.headers:
trusted_ranges = general_settings.get("mcp_trusted_proxy_ranges")
if trusted_ranges:
# Validate direct connection is from trusted proxy
if not IPAddressUtils.is_request_from_trusted_proxy(
request, general_settings=general_settings
):
direct_ip = request.client.host if request.client else None
trusted_networks = IPAddressUtils.parse_trusted_proxy_networks(
trusted_ranges
)
if not IPAddressUtils.is_trusted_proxy(direct_ip, trusted_networks):
# Untrusted source trying to set XFF - ignore XFF, use direct IP
if general_settings.get("mcp_trusted_proxy_ranges"):
# Direct connection isn't in any configured trusted CIDR.
verbose_proxy_logger.warning(
"XFF header from untrusted IP %s, ignoring", direct_ip
)
return direct_ip
# XFF enabled but no trusted proxy ranges configured: the direct
# peer is typically the reverse proxy's own (private) IP, so
# returning it would mis-classify external callers as internal.
# Fail closed for access control.
return ""
return _get_request_ip_address(request, use_x_forwarded_for=use_xff)

View file

@ -30,6 +30,7 @@ from litellm.proxy._types import *
from litellm.proxy.auth.auth_checks import (
ExperimentalUIJWTToken,
_cache_key_object,
_check_end_user_budget,
_delete_cache_key_object,
_get_user_role,
_is_model_cost_zero,
@ -146,12 +147,14 @@ def _get_model_from_request_context(
request_data: dict,
route: str,
request: Optional[Request],
llm_router: Optional[Any] = None,
) -> Optional[Union[str, List[str]]]:
return get_model_from_request(
request_data=request_data,
route=route,
request_headers=_safe_get_request_headers(request=request),
request_query_params=_safe_get_request_query_params(request=request),
llm_router=llm_router,
)
@ -1034,6 +1037,7 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
request_data=request_data,
route=route,
request=request,
llm_router=llm_router,
)
skip_budget_checks = False
if model is not None and llm_router is not None:
@ -1451,6 +1455,7 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
request_data=request_data,
route=route,
request=request,
llm_router=llm_router,
)
skip_budget_checks = False
if model is not None and llm_router is not None:
@ -1579,6 +1584,7 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
request_data=request_data,
route=route,
request=request,
llm_router=llm_router,
)
current_models = _get_model_names_for_budget_checks(
model=current_model
@ -1757,8 +1763,7 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
async def _safe_fetch(label: str, awaitable):
"""Run an awaitable and return its result. Re-raises authentication /
authorization failures (HTTPException, ProxyException,
BudgetExceededError — which ``get_end_user_object`` raises for
end-user budget violations) so they propagate to the caller.
BudgetExceededError) so they propagate to the caller.
Other exceptions (e.g. transient DB errors fetching context) are
swallowed with a debug log and ``None`` is returned so
``common_checks`` can still run against whatever limits are recorded
@ -2159,6 +2164,7 @@ def _should_skip_budget_checks(
request_data=request_data,
route=route,
request=request,
llm_router=llm_router,
)
if model is not None and llm_router is not None:
return _is_model_cost_zero(model=model, llm_router=llm_router)
@ -2475,6 +2481,7 @@ async def _enforce_key_and_fallback_model_access(
request_data=request_data,
route=route,
request=request,
llm_router=llm_router,
)
if model is not None:
@ -2577,6 +2584,14 @@ async def _run_post_custom_auth_checks(
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
)
# common_checks() enforces the end-user budget, but the centralized
# gate skips it for custom-auth deployments unless
# custom_auth_run_common_checks is set. Enforce it here on that path
# so an over-budget end user can't keep making requests.
if end_user_object is not None and not general_settings.get(
"custom_auth_run_common_checks", False
):
await _check_end_user_budget(end_user_obj=end_user_object, route=route)
# 2. Check token expiry
if valid_token.expires is not None:
@ -2616,6 +2631,7 @@ async def _run_post_custom_auth_checks(
request_data=request_data,
route=route,
request=request,
llm_router=llm_router,
)
current_models = _get_model_names_for_budget_checks(model=current_model)

View file

@ -2,6 +2,7 @@ import asyncio
import copy
import logging
import os
import secrets
import time
import traceback
from datetime import datetime, timedelta
@ -39,6 +40,7 @@ from litellm.proxy.health_check import (
from litellm.proxy.middleware.in_flight_requests_middleware import (
get_in_flight_requests,
)
from litellm.proxy.shutdown.graceful_shutdown_manager import GracefulShutdownManager
#### Health ENDPOINTS ####
@ -1551,6 +1553,50 @@ def _allow_public_health_readiness_details() -> bool:
return general_settings.get("allow_public_health_readiness_details") is True
def _drain_endpoint_enabled() -> bool:
from litellm.proxy.proxy_server import general_settings
return general_settings.get("enable_drain_endpoint") is True
def _drain_endpoint_token() -> Optional[str]:
"""
Shared secret required on the X-Drain-Token header to call /health/drain.
Falls back to the ``DRAIN_ENDPOINT_TOKEN`` env var when unset in
general_settings so the kubelet preStop hook can supply it via
``valueFrom.secretKeyRef`` without a config reload.
"""
from litellm.proxy.proxy_server import general_settings
token = general_settings.get("drain_endpoint_token")
if isinstance(token, str) and token:
return token
env_token = os.getenv("DRAIN_ENDPOINT_TOKEN")
if env_token:
return env_token
return None
def _authorize_drain_request(request: Request) -> None:
"""
Reject /health/drain calls that don't carry the configured X-Drain-Token.
When no token is configured the endpoint is treated as already opted-in
(the ``enable_drain_endpoint`` flag is the only gate). Comparison uses
``secrets.compare_digest`` to avoid timing leaks.
"""
expected = _drain_endpoint_token()
if expected is None:
return
supplied = request.headers.get("x-drain-token") or ""
if not secrets.compare_digest(supplied, expected):
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
detail="Invalid or missing X-Drain-Token",
)
async def _resolve_public_readiness_db(response: Response) -> str:
"""
Return the db status string for the public probe and flip the response to
@ -1580,6 +1626,10 @@ async def health_readiness(response: Response):
credential. Admins can opt into the legacy detailed payload with
general_settings.allow_public_health_readiness_details.
"""
if GracefulShutdownManager.is_shutting_down():
response.status_code = status.HTTP_503_SERVICE_UNAVAILABLE
return {"status": "shutting_down"}
if _allow_public_health_readiness_details():
return await _get_health_readiness_details(response=response)
@ -1616,6 +1666,54 @@ async def health_backlog():
return {"in_flight_requests": get_in_flight_requests()}
@router.get(
"/health/drain",
tags=["health"],
)
async def health_drain(request: Request):
"""
Graceful-drain probe for Kubernetes ``preStop`` hooks.
Disabled by default and returns 404 unless ``general_settings`` sets
``enable_drain_endpoint: true``. Calling it flips a process-wide
shutting-down flag, so a successful call permanently takes the worker out
of rotation until the pod restarts.
Because the kubelet calls preStop hooks without proxy credentials, the
endpoint does not require ``user_api_key_auth``. To prevent any
pod-reachable caller from triggering shutdown, set
``general_settings.drain_endpoint_token`` (or the ``DRAIN_ENDPOINT_TOKEN``
env var) and supply the same value on the ``X-Drain-Token`` header from
the preStop hook. Calls without the header (or with a wrong value) get a
401 and have no side effect.
When enabled, it marks the worker as shutting down (so /health/readiness
and /health/liveliness immediately start returning 503, removing the pod
from service) and blocks until the in-flight request counter drains to
zero or ``GRACEFUL_SHUTDOWN_TIMEOUT`` elapses. Unlike a fixed ``sleep``,
this returns as soon as real in-flight work is done.
Wire it up as:
```yaml
lifecycle:
preStop:
httpGet:
path: /health/drain
port: 4000
httpHeaders:
- name: X-Drain-Token
value: <same value as drain_endpoint_token>
```
"""
if not _drain_endpoint_enabled():
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Not Found")
_authorize_drain_request(request)
GracefulShutdownManager.start_shutdown()
drained = await GracefulShutdownManager.wait_for_drain(exclude_self=True)
return {"status": "drained", "drained_requests": drained}
@router.get(
"/health/liveliness", # Historical LiteLLM name; doesn't match k8s terminology but kept for backwards compatibility
tags=["health"],
@ -1624,10 +1722,16 @@ async def health_backlog():
"/health/liveness", # Kubernetes has "liveness" probes (https://kubernetes.io/docs/tasks/configure-pod-container/configure-liveness-readiness-startup-probes/#define-a-liveness-command)
tags=["health"],
)
async def health_liveliness():
async def health_liveliness(response: Response):
"""
Unprotected endpoint for checking if worker is alive
Unprotected endpoint for checking if worker is alive.
Returns 503 once graceful shutdown has begun so Kubernetes stops counting
the draining pod as live and terminates it on schedule.
"""
if GracefulShutdownManager.is_shutting_down():
response.status_code = status.HTTP_503_SERVICE_UNAVAILABLE
return {"status": "shutting_down"}
return "I'm alive!"

View file

@ -36,7 +36,7 @@ from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.auth.auth_utils import get_model_rate_limit_from_metadata
from litellm.types.caching import RedisPipelineIncrementOperation
from litellm.types.llms.openai import BaseLiteLLMOpenAIResponseObject
from litellm.types.utils import ModelResponse, Usage
from litellm.types.utils import CallTypes, ModelResponse, Usage
if TYPE_CHECKING:
from opentelemetry.trace import Span as _Span
@ -1375,6 +1375,79 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
)
)
def _add_mcp_per_key_rate_limit_descriptor(
self,
user_api_key_dict: UserAPIKeyAuth,
mcp_server_name: Optional[str],
descriptors: List[RateLimitDescriptor],
) -> None:
"""
Add a per-MCP-server rpm descriptor for the API key, if a limit is
configured for the server being called.
MCP tool calls have no token usage, so only requests_per_unit is set;
tokens_per_unit stays None so the TPM reservation path is never engaged.
"""
from litellm.proxy.auth.auth_utils import get_key_mcp_rpm_limit
if not mcp_server_name or not user_api_key_dict.api_key:
return
mcp_rpm_limit = get_key_mcp_rpm_limit(user_api_key_dict)
if not mcp_rpm_limit:
return
server_rpm_limit = mcp_rpm_limit.get(mcp_server_name)
if server_rpm_limit is None:
return
descriptors.append(
RateLimitDescriptor(
key="mcp_per_key",
value=f"{user_api_key_dict.api_key}:{mcp_server_name}",
rate_limit={
"requests_per_unit": server_rpm_limit,
"tokens_per_unit": None,
"window_size": self.window_size,
},
)
)
def _add_mcp_per_team_rate_limit_descriptor(
self,
user_api_key_dict: UserAPIKeyAuth,
mcp_server_name: Optional[str],
descriptors: List[RateLimitDescriptor],
) -> None:
"""
Add a per-MCP-server rpm descriptor for the team, if a limit is
configured for the server being called.
"""
from litellm.proxy.auth.auth_utils import get_team_mcp_rpm_limit
if not mcp_server_name or not user_api_key_dict.team_id:
return
mcp_rpm_limit = get_team_mcp_rpm_limit(user_api_key_dict)
if not mcp_rpm_limit:
return
server_rpm_limit = mcp_rpm_limit.get(mcp_server_name)
if server_rpm_limit is None:
return
descriptors.append(
RateLimitDescriptor(
key="mcp_per_team",
value=f"{user_api_key_dict.team_id}:{mcp_server_name}",
rate_limit={
"requests_per_unit": server_rpm_limit,
"tokens_per_unit": None,
"window_size": self.window_size,
},
)
)
def _should_enforce_rate_limit(
self,
limit_type: Optional[str],
@ -1533,6 +1606,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
rpm_limit_type: Optional[str],
tpm_limit_type: Optional[str],
model_has_failures: bool,
call_type: Optional[str] = None,
) -> List[RateLimitDescriptor]:
"""
Create all rate limit descriptors for the request.
@ -1653,6 +1727,21 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
descriptors=descriptors,
)
# REST MCP calls pass the raw body through this hook before server
# resolution; only the later synthetic hook payload may carry this key.
if call_type == CallTypes.call_mcp_tool.value and "server_id" not in data:
mcp_server_name = data.get("mcp_server_name", None)
self._add_mcp_per_key_rate_limit_descriptor(
user_api_key_dict=user_api_key_dict,
mcp_server_name=mcp_server_name,
descriptors=descriptors,
)
self._add_mcp_per_team_rate_limit_descriptor(
user_api_key_dict=user_api_key_dict,
mcp_server_name=mcp_server_name,
descriptors=descriptors,
)
if (
get_team_model_rpm_limit(user_api_key_dict) is not None
or get_team_model_tpm_limit(user_api_key_dict) is not None
@ -1983,6 +2072,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
rpm_limit_type=rpm_limit_type,
tpm_limit_type=tpm_limit_type,
model_has_failures=model_has_failures,
call_type=call_type,
)
# Add team model rate limits from team_metadata

View file

@ -386,6 +386,7 @@ async def new_user(
- soft_budget: Optional[float] - Get alerts when user crosses given budget, doesn't block requests.
- model_max_budget: Optional[dict] - Model-specific max budget for user. [Docs](https://docs.litellm.ai/docs/proxy/users#add-model-specific-budgets-to-keys)
- model_rpm_limit: Optional[float] - Model-specific rpm limit for user. [Docs](https://docs.litellm.ai/docs/proxy/users#add-model-specific-limits-to-keys)
- mcp_rpm_limit: Optional[dict] - Per-MCP-server rpm limit, keyed by MCP server name {"github": 100, "slack": 200}. Enforced for keys and teams only; values set on a user are stored but not enforced per user.
- model_tpm_limit: Optional[float] - Model-specific tpm limit for user. [Docs](https://docs.litellm.ai/docs/proxy/users#add-model-specific-limits-to-keys)
- spend: Optional[float] - Amount spent by user. Default is 0. Will be updated by proxy whenever user is used. You can set duration as seconds ("30s"), minutes ("30m"), hours ("30h"), days ("30d"), months ("1mo").
- agent_id: Optional[str] - The agent id associated with the user.
@ -1427,6 +1428,7 @@ async def user_update(
- soft_budget: Optional[float] - Get alerts when user crosses given budget, doesn't block requests.
- model_max_budget: Optional[dict] - Model-specific max budget for user. [Docs](https://docs.litellm.ai/docs/proxy/users#add-model-specific-budgets-to-keys)
- model_rpm_limit: Optional[float] - Model-specific rpm limit for user. [Docs](https://docs.litellm.ai/docs/proxy/users#add-model-specific-limits-to-keys)
- mcp_rpm_limit: Optional[dict] - Per-MCP-server rpm limit, keyed by MCP server name {"github": 100, "slack": 200}. Enforced for keys and teams only; values set on a user are stored but not enforced per user.
- model_tpm_limit: Optional[float] - Model-specific tpm limit for user. [Docs](https://docs.litellm.ai/docs/proxy/users#add-model-specific-limits-to-keys)
- spend: Optional[float] - Amount spent by user. Default is 0. Will be updated by proxy whenever user is used. You can set duration as seconds ("30s"), minutes ("30m"), hours ("30h"), days ("30d"), months ("1mo").
- agent_id: Optional[str] - The agent id associated with the user.

View file

@ -894,7 +894,12 @@ async def _common_key_generation_helper( # noqa: PLR0915
user_api_key_dict.user_role is not None
and user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN.value
)
if not _is_proxy_admin:
_org_inherited_from_team = (
team_table is not None
and team_table.organization_id is not None
and data.organization_id == team_table.organization_id
)
if not _is_proxy_admin and not _org_inherited_from_team:
await _validate_caller_can_assign_key_org(
user_api_key_dict=user_api_key_dict,
organization_id=data.organization_id,
@ -1388,6 +1393,7 @@ async def generate_key_fn(
- model_max_budget: Optional[Dict[str, BudgetConfig]] - Model-specific budgets {"gpt-4": {"budget_limit": 0.0005, "time_period": "30d"}}}. IF null or {} then no model specific budget.
- model_rpm_limit: Optional[dict] - key-specific model rpm limit. Example - {"text-davinci-002": 1000, "gpt-3.5-turbo": 1000}. IF null or {} then no model specific rpm limit.
- model_tpm_limit: Optional[dict] - key-specific model tpm limit. Example - {"text-davinci-002": 1000, "gpt-3.5-turbo": 1000}. IF null or {} then no model specific tpm limit.
- mcp_rpm_limit: Optional[dict] - key-specific per-MCP-server rpm limit, keyed by MCP server name (alias if set, else the configured name). Example - {"github": 100, "slack": 200}. IF null or {} then no MCP-specific rpm limit.
- tpm_limit_type: Optional[str] - Type of tpm limit. Options: "best_effort_throughput" (no error if we're overallocating tpm), "guaranteed_throughput" (raise an error if we're overallocating tpm), "dynamic" (dynamically exceed limit when no 429 errors). Defaults to "best_effort_throughput".
- rpm_limit_type: Optional[str] - Type of rpm limit. Options: "best_effort_throughput" (no error if we're overallocating rpm), "guaranteed_throughput" (raise an error if we're overallocating rpm), "dynamic" (dynamically exceed limit when no 429 errors). Defaults to "best_effort_throughput".
- allowed_cache_controls: Optional[list] - List of allowed cache control values. Example - ["no-cache", "no-store"]. See all values - https://docs.litellm.ai/docs/proxy/caching#turn-on--off-caching-per-request
@ -1606,6 +1612,7 @@ async def generate_service_account_key_fn(
- model_max_budget: Optional[Dict[str, BudgetConfig]] - Model-specific budgets {"gpt-4": {"budget_limit": 0.0005, "time_period": "30d"}}}. IF null or {} then no model specific budget.
- model_rpm_limit: Optional[dict] - key-specific model rpm limit. Example - {"text-davinci-002": 1000, "gpt-3.5-turbo": 1000}. IF null or {} then no model specific rpm limit.
- model_tpm_limit: Optional[dict] - key-specific model tpm limit. Example - {"text-davinci-002": 1000, "gpt-3.5-turbo": 1000}. IF null or {} then no model specific tpm limit.
- mcp_rpm_limit: Optional[dict] - key-specific per-MCP-server rpm limit, keyed by MCP server name (alias if set, else the configured name). Example - {"github": 100, "slack": 200}. IF null or {} then no MCP-specific rpm limit.
- tpm_limit_type: Optional[str] - TPM rate limit type - "best_effort_throughput", "guaranteed_throughput", or "dynamic"
- rpm_limit_type: Optional[str] - RPM rate limit type - "best_effort_throughput", "guaranteed_throughput", or "dynamic"
- allowed_cache_controls: Optional[list] - List of allowed cache control values. Example - ["no-cache", "no-store"]. See all values - https://docs.litellm.ai/docs/proxy/caching#turn-on--off-caching-per-request
@ -2422,6 +2429,7 @@ async def update_key_fn( # noqa: PLR0915
- tpm_limit: Optional[int] - Tokens per minute limit
- rpm_limit: Optional[int] - Requests per minute limit
- model_rpm_limit: Optional[dict] - Model-specific RPM limits {"gpt-4": 100, "claude-v1": 200}
- mcp_rpm_limit: Optional[dict] - Per-MCP-server RPM limits, keyed by MCP server name {"github": 100, "slack": 200}
- model_tpm_limit: Optional[dict] - Model-specific TPM limits {"gpt-4": 100000, "claude-v1": 200000}
- tpm_limit_type: Optional[str] - TPM rate limit type - "best_effort_throughput", "guaranteed_throughput", or "dynamic"
- rpm_limit_type: Optional[str] - RPM rate limit type - "best_effort_throughput", "guaranteed_throughput", or "dynamic"
@ -3401,6 +3409,7 @@ async def generate_key_helper_fn( # noqa: PLR0915
model_max_budget: Optional[dict] = {},
model_rpm_limit: Optional[dict] = None,
model_tpm_limit: Optional[dict] = None,
mcp_rpm_limit: Optional[dict] = None,
guardrails: Optional[list] = None,
policies: Optional[list] = None,
prompts: Optional[list] = None,
@ -3479,6 +3488,9 @@ async def generate_key_helper_fn( # noqa: PLR0915
if model_tpm_limit is not None:
metadata = metadata or {}
metadata["model_tpm_limit"] = model_tpm_limit
if mcp_rpm_limit is not None:
metadata = metadata or {}
metadata["mcp_rpm_limit"] = mcp_rpm_limit
if guardrails is not None:
metadata = metadata or {}
metadata["guardrails"] = guardrails

View file

@ -1542,7 +1542,7 @@ if MCP_AVAILABLE:
master_key,
algorithms=["HS256"],
# UI session cookies may omit exp; don't require it.
options={"verify_exp": False},
options={"verify_exp": False, "verify_aud": False},
)
if decoded.get("login_method") in ("sso", "username_password"):
cookie_key = decoded.get("key", "")

View file

@ -863,8 +863,9 @@ async def new_team( # noqa: PLR0915
- members_with_roles: List[{"role": "admin" or "user", "user_id": "<user-id>"}] - A list of users and their roles in the team. Get user_id when making a new user via `/user/new`.
- team_member_permissions: Optional[List[str]] - A list of routes that non-admin team members can access. example: ["/key/generate", "/key/update", "/key/delete"]
- metadata: Optional[dict] - Metadata for team, store information for team. Example metadata = {"extra_info": "some info"}
- model_rpm_limit: Optional[Dict[str, int]] - The RPM (Requests Per Minute) limit for this team - applied across all keys for this team.
- model_rpm_limit: Optional[Dict[str, int]] - The RPM (Requests Per Minute) limit for this team - applied across all keys for this team.
- model_tpm_limit: Optional[Dict[str, int]] - The TPM (Tokens Per Minute) limit for this team - applied across all keys for this team.
- mcp_rpm_limit: Optional[Dict[str, int]] - Per-MCP-server RPM limit for this team, keyed by MCP server name (alias if set, else the configured name). Example: {"github": 100, "slack": 200}. Applied across all keys for this team.
- tpm_limit: Optional[int] - The TPM (Tokens Per Minute) limit for this team - all keys with this team_id will have at max this TPM limit
- rpm_limit: Optional[int] - The RPM (Requests Per Minute) limit for this team - all keys associated with this team_id will have at max this RPM limit
- rpm_limit_type: Optional[Literal["guaranteed_throughput", "best_effort_throughput"]] - The type of RPM limit enforcement. Use "guaranteed_throughput" to raise an error if overallocating RPM, or "best_effort_throughput" for best effort enforcement.

View file

@ -667,6 +667,34 @@ class HttpPassThroughEndpointHelpers(BasePassthroughUtils):
return stream
def _carry_guardrail_logging_info(
request_data: dict, guardrail_data: Optional[dict]
) -> None:
"""Copy guardrail logging entries from ``guardrail_data`` onto ``request_data``.
Post-call guardrails run against a throwaway ``hook_data`` dict (its
``metadata`` is what ``_init_kwargs_for_pass_through_endpoint`` already
stripped off ``_parsed_body``), so a block records the
``standard_logging_guardrail_information`` there and not on the dict the
failure handler forwards to ``post_call_failure_hook``. Without this the
otel guardrail span is emitted on allow but missing on block. Carry the
entries over so the failure path matches the unified path.
"""
if guardrail_data is None:
return
source_metadata = guardrail_data.get("metadata")
if not isinstance(source_metadata, dict):
return
entries = source_metadata.get("standard_logging_guardrail_information")
if not entries:
return
metadata = request_data.get("metadata")
if not isinstance(metadata, dict):
metadata = request_data["metadata"] = {}
metadata.setdefault("standard_logging_guardrail_information", list(entries))
async def pass_through_request( # noqa: PLR0915
request: Request,
target: str,
@ -718,6 +746,9 @@ async def pass_through_request( # noqa: PLR0915
# kwargs for pass through endpoint, contains metadata, litellm_params, call_type, litellm_call_id, passthrough_logging_payload
kwargs: Optional[dict] = None
logging_obj: Optional[Logging] = None
# the dict post-call guardrails wrote their logging info into; the failure
# handler reuses it so a guardrail block still surfaces its span/logs
post_call_guardrail_data: Optional[dict] = None
#########################################################
try:
@ -1160,6 +1191,7 @@ async def pass_through_request( # noqa: PLR0915
**existing_metadata,
"guardrails": guardrails_to_run,
}
post_call_guardrail_data = hook_data
response_body = await proxy_logging_obj.post_call_success_hook(
data=hook_data,
user_api_key_dict=user_api_key_dict,
@ -1343,6 +1375,8 @@ async def pass_through_request( # noqa: PLR0915
if "custom_llm_provider" not in request_payload and custom_llm_provider:
request_payload["custom_llm_provider"] = custom_llm_provider
_carry_guardrail_logging_info(request_payload, post_call_guardrail_data)
await proxy_logging_obj.post_call_failure_hook(
user_api_key_dict=user_api_key_dict,
original_exception=e,

View file

@ -417,6 +417,7 @@ from litellm.proxy.middleware.prometheus_auth_middleware import PrometheusAuthMi
from litellm.proxy.middleware.request_size_limit_middleware import (
RequestSizeLimitMiddleware,
)
from litellm.proxy.shutdown.graceful_shutdown_manager import GracefulShutdownManager
from litellm.proxy.ocr_endpoints.endpoints import router as ocr_router
from litellm.proxy.openai_files_endpoints.files_endpoints import (
router as openai_files_router,
@ -973,6 +974,11 @@ async def proxy_startup_event(app: FastAPI): # noqa: PLR0915
# End of startup event
yield
# Shutdown event - drain in-flight requests before tearing down dependencies
# so SIGTERM (rolling update, scale-down, liveness kill) doesn't drop them.
GracefulShutdownManager.start_shutdown()
await GracefulShutdownManager.wait_for_drain()
# Shutdown event - close shared aiohttp session
if shared_aiohttp_session is not None:
try:
@ -1265,7 +1271,7 @@ async def openai_exception_handler(request: Request, exc: ProxyException):
headers = exc.headers
error_dict = exc.to_dict()
status_code = int(exc.code) if exc.code else status.HTTP_500_INTERNAL_SERVER_ERROR
_close_dangling_otel_server_span(request, status_code)
_close_dangling_otel_server_span(request, status_code, exc=exc)
return JSONResponse(
status_code=status_code,
content={"error": error_dict},
@ -1273,7 +1279,9 @@ async def openai_exception_handler(request: Request, exc: ProxyException):
)
def _close_dangling_otel_server_span(request: Request, status_code: int) -> None:
def _close_dangling_otel_server_span(
request: Request, status_code: int, exc: Optional[Exception] = None
) -> None:
parent_otel_span = getattr(request.state, "parent_otel_span", None)
if parent_otel_span is None:
return
@ -1296,6 +1304,10 @@ def _close_dangling_otel_server_span(request: Request, status_code: int) -> None
open_telemetry_logger.set_response_status_code_attribute(
parent_otel_span, status_code
)
if status_code >= 400:
open_telemetry_logger.record_error_attributes_on_span(
parent_otel_span, exc, status_code
)
parent_otel_span.set_status(
Status(StatusCode.ERROR if status_code >= 400 else StatusCode.OK)
)
@ -1312,7 +1324,7 @@ def _close_dangling_otel_server_span(request: Request, status_code: int) -> None
async def otel_request_validation_exception_handler(
request: Request, exc: RequestValidationError
):
_close_dangling_otel_server_span(request, 422)
_close_dangling_otel_server_span(request, 422, exc=exc)
return JSONResponse(
status_code=422,
content={"detail": jsonable_encoder(exc.errors())},
@ -1326,7 +1338,7 @@ async def otel_unhandled_exception_handler(request: Request, exc: Exception):
verbose_proxy_logger.exception(
"Unhandled exception in request: %s", type(exc).__name__
)
_close_dangling_otel_server_span(request, 500)
_close_dangling_otel_server_span(request, 500, exc=exc)
return JSONResponse(
status_code=500,
content={
@ -15845,6 +15857,8 @@ async def _mcp_forward_as_path(path_segment: str, request: Request):
)
scope = dict(request.scope)
# Preserve the public request path for OAuth challenge URL selection.
scope["_original_path"] = scope.get("path", "")
scope["path"] = f"/mcp/{path_segment}"
return await _stream_mcp_asgi_response(
handle_streamable_http_mcp, scope, request.receive
@ -15992,6 +16006,7 @@ async def dynamic_mcp_route(mcp_server_name: str, request: Request):
)
if toolset is not None:
scope = dict(request.scope)
scope["_original_path"] = scope.get("path", "")
scope["path"] = "/mcp"
token = _mcp_active_toolset_id.set(toolset.toolset_id)
try:

View file

@ -7,6 +7,48 @@
"credential_fields": [],
"litellm_params_template": {}
},
{
"agent_type": "langflow",
"agent_type_display_name": "LangFlow",
"description": "Connect to LangFlow AI agents via the LangFlow Platform API",
"logo_url": "/ui/assets/logos/langflow.svg",
"model_template": "langflow/{flow_id}",
"credential_fields": [
{
"key": "flow_id",
"label": "Flow ID",
"placeholder": "your-flow-id",
"tooltip": "The Flow ID from your LangFlow deployment (found in the flow URL or settings)",
"required": true,
"field_type": "text",
"default_value": null,
"include_in_litellm_params": false
},
{
"key": "api_base",
"label": "LangFlow API Base",
"placeholder": "http://localhost:7860",
"tooltip": "The base URL for your LangFlow server (e.g., http://localhost:7860 or your deployed LangFlow URL)",
"required": true,
"field_type": "text",
"default_value": "http://localhost:7860",
"include_in_litellm_params": true
},
{
"key": "api_key",
"label": "LangFlow API Key",
"placeholder": null,
"tooltip": "API key for authenticating with your LangFlow server (x-api-key header)",
"required": false,
"field_type": "password",
"default_value": null,
"include_in_litellm_params": true
}
],
"litellm_params_template": {
"custom_llm_provider": "langflow"
}
},
{
"agent_type": "langgraph",
"agent_type_display_name": "LangGraph",

View file

@ -325,6 +325,7 @@ model LiteLLM_MCPServerTable {
allow_all_keys Boolean @default(false)
available_on_public_internet Boolean @default(true)
delegate_auth_to_upstream Boolean @default(false)
oauth_passthrough Boolean @default(false)
is_byok Boolean @default(false)
byok_description String[] @default([])
byok_api_key_help_url String?

View file

@ -3,12 +3,17 @@ CRUD ENDPOINTS FOR SEARCH TOOLS
"""
from datetime import datetime
from typing import Any, Dict, List, Union
from typing import Any, Dict, List, Optional, Union
from fastapi import APIRouter, Depends, HTTPException
from pydantic import BaseModel
from litellm._logging import verbose_proxy_logger
from litellm.proxy._types import (
LiteLLM_TeamTable,
LitellmUserRoles,
UserAPIKeyAuth,
)
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.search_endpoints.search_tool_registry import SearchToolRegistry
from litellm.types.search import (
@ -41,13 +46,61 @@ def _convert_datetime_to_str(value: Union[datetime, str, None]) -> Union[str, No
return value
async def _filter_visible_search_tools(
search_tools: List[SearchToolInfoResponse],
user_api_key_dict: UserAPIKeyAuth,
) -> List[SearchToolInfoResponse]:
"""
Drop search tools the caller is not authorized to invoke, applying the same
key/team object_permission allowlists enforced on /search. Admins see all tools.
"""
if user_api_key_dict.user_role in (
LitellmUserRoles.PROXY_ADMIN,
LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY,
):
return search_tools
from litellm.proxy.auth.auth_checks import (
can_user_view_search_tool,
get_team_object,
)
from litellm.proxy.proxy_server import (
prisma_client,
proxy_logging_obj,
user_api_key_cache,
)
team_object: Optional[LiteLLM_TeamTable] = None
if user_api_key_dict.team_id:
team_object = await get_team_object(
team_id=user_api_key_dict.team_id,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
parent_otel_span=user_api_key_dict.parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
)
visible: List[SearchToolInfoResponse] = []
for tool in search_tools:
tool_name = tool.get("search_tool_name")
if tool_name and await can_user_view_search_tool(
search_tool_name=tool_name,
valid_token=user_api_key_dict,
team_object=team_object,
):
visible.append(tool)
return visible
@router.get(
"/search_tools/list",
tags=["Search Tools"],
dependencies=[Depends(user_api_key_auth)],
response_model=ListSearchToolsResponse,
)
async def list_search_tools():
async def list_search_tools(
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
):
"""
List all search tools that are available in the database and config file.
@ -114,22 +167,25 @@ async def list_search_tools():
f"Could not get config-defined search tools: {e}"
)
for search_tool in config_search_tools:
tool_name = search_tool.get("search_tool_name")
for config_search_tool in config_search_tools:
tool_name = config_search_tool.get("search_tool_name")
if tool_name:
litellm_params_dict = dict(search_tool.get("litellm_params", {}))
litellm_params_dict = dict(config_search_tool.get("litellm_params", {}))
masked_litellm_params_dict = _get_masked_values(
litellm_params_dict,
unmasked_length=4,
number_of_asterisks=4,
)
config_tool_info = config_search_tool.get("search_tool_info")
search_tool_configs.append(
SearchToolInfoResponse(
search_tool_id=None,
search_tool_name=tool_name,
litellm_params=masked_litellm_params_dict,
search_tool_info=search_tool.get("search_tool_info"),
search_tool_info=(
dict(config_tool_info) if config_tool_info else None
),
created_at=None,
updated_at=None,
is_from_config=True,
@ -142,8 +198,8 @@ async def list_search_tools():
if tool.get("search_tool_name") not in db_tool_names
]
for search_tool in search_tools_from_db:
litellm_params_dict = dict(search_tool.get("litellm_params", {}))
for db_search_tool in search_tools_from_db:
litellm_params_dict = dict(db_search_tool.get("litellm_params", {}))
masked_litellm_params_dict = _get_masked_values(
litellm_params_dict,
unmasked_length=4,
@ -152,17 +208,25 @@ async def list_search_tools():
search_tool_configs.append(
SearchToolInfoResponse(
search_tool_id=search_tool.get("search_tool_id"),
search_tool_name=search_tool.get("search_tool_name", ""),
search_tool_id=db_search_tool.get("search_tool_id"),
search_tool_name=db_search_tool.get("search_tool_name", ""),
litellm_params=masked_litellm_params_dict,
search_tool_info=search_tool.get("search_tool_info"),
created_at=_convert_datetime_to_str(search_tool.get("created_at")),
updated_at=_convert_datetime_to_str(search_tool.get("updated_at")),
search_tool_info=db_search_tool.get("search_tool_info"),
created_at=_convert_datetime_to_str(
db_search_tool.get("created_at")
),
updated_at=_convert_datetime_to_str(
db_search_tool.get("updated_at")
),
is_from_config=False,
)
)
return ListSearchToolsResponse(search_tools=search_tool_configs)
visible_search_tools = await _filter_visible_search_tools(
search_tool_configs, user_api_key_dict
)
return ListSearchToolsResponse(search_tools=visible_search_tools)
except Exception as e:
verbose_proxy_logger.exception(f"Error getting search tools: {e}")
raise HTTPException(status_code=500, detail=str(e))

View file

View file

@ -0,0 +1,174 @@
"""
Application-level graceful shutdown coordination for the LiteLLM proxy.
Kubernetes terminates a pod by sending ``SIGTERM`` and, after
``terminationGracePeriodSeconds``, ``SIGKILL``. By default LiteLLM delegates
the signal to uvicorn and tears down immediately, dropping any in-flight
requests (streaming, batch inference, long-lived calls).
A fixed ``preStop`` sleep can not solve this: it has to be sized for the
*worst-case* request, so it either wastes time on every routine shutdown or is
too short for a long-running request. This manager instead drains based on the
*actual* in-flight request counter (already tracked by
``InFlightRequestsMiddleware``), so a pod terminates as soon as its real
in-flight work is done — and never waits longer than ``GRACEFUL_SHUTDOWN_TIMEOUT``.
The state is process-scoped (class-level), matching the per-uvicorn-worker
granularity of ``InFlightRequestsMiddleware``.
"""
import asyncio
import os
import time
from typing import Callable, Optional
from litellm._logging import verbose_proxy_logger
from litellm.proxy.middleware.in_flight_requests_middleware import (
get_in_flight_requests,
)
# Keep below terminationGracePeriodSeconds so the process exits before SIGKILL.
DEFAULT_GRACEFUL_SHUTDOWN_TIMEOUT = 30.0
_DRAIN_POLL_INTERVAL = 0.1
_DRAIN_LOG_INTERVAL = 5.0
class GracefulShutdownManager:
"""
Process-scoped singleton that tracks whether the worker is draining and
blocks until in-flight requests reach zero (or a timeout elapses).
"""
_is_shutting_down: bool = False
_shutdown_started_at: Optional[float] = None
_drain_performed: bool = False
@classmethod
def is_shutting_down(cls) -> bool:
"""Whether this worker has begun graceful shutdown."""
return cls._is_shutting_down
@classmethod
def get_timeout(cls) -> float:
"""
Read GRACEFUL_SHUTDOWN_TIMEOUT (seconds) from the environment on each
call so deployments can tune it without code changes. Falls back to the
default on an unset or malformed value.
"""
raw = os.getenv("GRACEFUL_SHUTDOWN_TIMEOUT")
if raw is None:
return DEFAULT_GRACEFUL_SHUTDOWN_TIMEOUT
try:
return float(raw)
except (TypeError, ValueError):
verbose_proxy_logger.warning(
"GRACEFUL_SHUTDOWN_TIMEOUT=%r is not a number; using default %ss",
raw,
DEFAULT_GRACEFUL_SHUTDOWN_TIMEOUT,
)
return DEFAULT_GRACEFUL_SHUTDOWN_TIMEOUT
@classmethod
def start_shutdown(cls) -> None:
"""
Mark the worker as draining. Idempotent — repeated calls (e.g. SIGTERM
followed by a preStop hit on /health/drain) do not reset the clock.
"""
if cls._is_shutting_down:
return
cls._is_shutting_down = True
cls._shutdown_started_at = time.monotonic()
verbose_proxy_logger.info(
"graceful_shutdown_started in_flight_requests=%s",
get_in_flight_requests(),
)
@classmethod
async def wait_for_drain(
cls,
timeout: Optional[float] = None,
exclude_self: bool = False,
count_fn: Optional[Callable[[], int]] = None,
poll_interval: float = _DRAIN_POLL_INTERVAL,
log_interval: float = _DRAIN_LOG_INTERVAL,
) -> int:
"""
Poll the in-flight request counter until it reaches the drain target or
``timeout`` seconds elapse.
Args:
timeout: Max seconds to wait. Defaults to ``get_timeout()``.
exclude_self: When the caller is itself an in-flight HTTP request
(the /health/drain endpoint), set this so the caller's own
request is not counted as outstanding work.
count_fn: Source of the current in-flight count. Defaults to the
live ``InFlightRequestsMiddleware`` counter; injectable for tests.
poll_interval: Seconds between counter polls.
log_interval: Minimum seconds between ``drain_waiting`` log lines.
Returns:
Number of requests that drained while waiting (>= 0).
"""
# A preStop /health/drain hook and the lifespan SIGTERM handler both
# drain; once one has run, the other must not wait again, otherwise the
# effective window is 2x the timeout and terminationGracePeriodSeconds
# has to be doubled to avoid a mid-drain SIGKILL.
if cls._drain_performed:
return 0
cls._drain_performed = True
if timeout is None:
timeout = cls.get_timeout()
if count_fn is None:
count_fn = get_in_flight_requests
# The /health/drain HTTP request flows through InFlightRequestsMiddleware
# and so counts itself; treat <=1 as "drained" in that case.
target = 1 if exclude_self else 0
start = time.monotonic()
initial = count_fn()
last_log = start
if timeout <= 0:
return max(0, initial - target)
while True:
current = count_fn()
if current <= target:
drained = max(0, initial - current)
verbose_proxy_logger.info(
"graceful_shutdown_complete drained_requests=%s elapsed_s=%.2f",
drained,
time.monotonic() - start,
)
return drained
elapsed = time.monotonic() - start
if elapsed >= timeout:
verbose_proxy_logger.warning(
"graceful_shutdown_timeout in_flight_requests=%s elapsed_s=%.2f "
"timeout_s=%s — proceeding with teardown",
current,
elapsed,
timeout,
)
return max(0, initial - current)
now = time.monotonic()
if now - last_log >= log_interval:
verbose_proxy_logger.info(
"drain_waiting in_flight_requests=%s elapsed_s=%.2f",
current,
elapsed,
)
last_log = now
await asyncio.sleep(poll_interval)
@classmethod
def reset(cls) -> None:
"""Reset state. Intended for use in tests."""
cls._is_shutting_down = False
cls._shutdown_started_at = None
cls._drain_performed = False

View file

@ -72,7 +72,7 @@ async def reserve_budget_for_request(
return None
if route in {"/models", "/v1/models", "/utils/token_counter"}:
return None
if get_model_from_request(request_body, route) is None:
if get_model_from_request(request_body, route, llm_router=llm_router) is None:
return None
counters = await _get_budget_counters(
@ -797,7 +797,7 @@ def estimate_request_max_cost(
route: str,
llm_router: Optional[Router],
) -> Optional[float]:
model = get_model_from_request(request_body, route)
model = get_model_from_request(request_body, route, llm_router=llm_router)
if model is None:
return None

View file

@ -24,7 +24,7 @@ from litellm.litellm_core_utils.core_helpers import (
get_litellm_metadata_from_kwargs,
reconstruct_model_name,
)
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps, strip_null_bytes
from litellm.proxy._types import SpendLogsMetadata, SpendLogsPayload
from litellm.proxy.spend_tracking.spend_log_error_logger import spend_log_error
from litellm.proxy.utils import PrismaClient, hash_token
@ -304,7 +304,7 @@ def get_logging_payload( # noqa: PLR0915
# BUG FIX: Don't overwrite api_key when standard_logging_payload is None
# The api_key was already extracted from metadata (line 243) and hashed (lines 256-259)
request_tags = (
json.dumps(metadata.get("tags", []))
safe_dumps(metadata.get("tags", []))
if isinstance(metadata.get("tags", []), list)
else "[]"
)
@ -312,7 +312,7 @@ def get_logging_payload( # noqa: PLR0915
standard_logging_payload is not None
and standard_logging_payload.get("request_tags") is not None
): # use 'tags' from standard logging payload instead
request_tags = json.dumps(standard_logging_payload["request_tags"])
request_tags = safe_dumps(standard_logging_payload["request_tags"])
_model_id = metadata.get("model_info", {}).get("id", "")
_model_group = metadata.get("model_group", "")
@ -606,7 +606,7 @@ def _get_messages_for_spend_logs_payload(
messages = standard_logging_payload.get("messages")
if messages is not None:
try:
return json.dumps(messages, default=str)
return safe_dumps(messages)
except Exception:
return "{}"
return "{}"
@ -976,7 +976,7 @@ def _get_proxy_server_request_for_spend_logs_payload(
perform_redaction(model_call_details=_request_body, result=None)
_request_body = _sanitize_request_body_for_spend_logs_payload(_request_body)
_request_body_json_str = json.dumps(_request_body, default=str)
_request_body_json_str = safe_dumps(_request_body)
if LITELLM_TRUNCATED_PAYLOAD_FIELD in _request_body_json_str:
verbose_proxy_logger.info(
"Spend Log: request body was truncated before storing in DB. %s",
@ -1059,7 +1059,7 @@ def _get_response_for_spend_logs_payload(
if sanitized_response is None:
return "{}"
if isinstance(sanitized_response, str):
result_str = sanitized_response
result_str = strip_null_bytes(sanitized_response)
else:
result_str = safe_dumps(sanitized_response)
if LITELLM_TRUNCATED_PAYLOAD_FIELD in result_str:

View file

@ -643,6 +643,7 @@ class ProxyLogging:
"user_api_key_request_route": kwargs.get("user_api_key_request_route"),
"mcp_tool_name": request_obj.tool_name, # Keep original for reference
"mcp_arguments": request_obj.arguments, # Keep original for reference
"mcp_server_name": kwargs.get("mcp_rate_limit_server_name"),
# Raw Bearer token from the original HTTP request — allows guardrails
# (e.g. MCPJWTSigner) to independently verify the caller's identity
# before re-signing an outbound token (FR-5 verify+re-sign).

View file

@ -9100,7 +9100,10 @@ class Router:
except Exception:
pass
# Three mutually exclusive scenarios for the model's metadata:
if custom_model_info is not None and litellm_model_name_model_info is not None:
# (1) It has both custom model_info set and exists in the built-in map
# merge with custom overriding built-in
model_info = cast(
ModelInfo,
_update_dictionary(
@ -9109,7 +9112,12 @@ class Router:
),
)
elif litellm_model_name_model_info is not None:
# (2) Built-in only — no custom pricing to merge
model_info = litellm_model_name_model_info
elif custom_model_info is not None:
# (3) Custom only — model not in built-in cost map yet
# custom_model_info already includes base_model defaults at this point, if applicable
model_info = cast(ModelInfo, custom_model_info)
return model_info

View file

@ -298,6 +298,23 @@ class MakeAgentsPublicRequest(BaseModel):
agent_ids: List[str]
def _normalize_a2a_jsonrpc_response(
response_dict: Dict[str, Any],
request_id: Optional[Any] = None,
) -> Dict[str, Any]:
"""
Ensure JSON-RPC responses include ``id`` when the caller supplied one.
The a2a SDK may omit ``id`` on error payloads even when the upstream agent
returned it. Backfill from the outbound request id so LiteLLM can surface the
agent error instead of failing Pydantic validation.
"""
normalized = dict(response_dict)
if normalized.get("id") is None and request_id is not None:
normalized["id"] = str(request_id)
return normalized
class LiteLLMSendMessageResponse(LiteLLMPydanticObjectBase):
"""
LiteLLM wrapper for A2A SendMessageResponse.
@ -322,31 +339,42 @@ class LiteLLMSendMessageResponse(LiteLLMPydanticObjectBase):
@classmethod
def from_a2a_response(
cls, response: "SendMessageResponse"
cls,
response: "SendMessageResponse",
request_id: Optional[Any] = None,
) -> "LiteLLMSendMessageResponse":
"""
Create a LiteLLMSendMessageResponse from an a2a SDK SendMessageResponse.
Args:
response: The a2a SDK SendMessageResponse
request_id: JSON-RPC request id to backfill when the SDK omits it on errors
Returns:
LiteLLMSendMessageResponse with _hidden_params support
"""
# Convert the a2a response to a dict
response_dict = response.model_dump(mode="json", exclude_none=True)
response_dict = _normalize_a2a_jsonrpc_response(
response_dict, request_id=request_id
)
return cls(**response_dict)
@classmethod
def from_dict(cls, response_dict: Dict[str, Any]) -> "LiteLLMSendMessageResponse":
def from_dict(
cls,
response_dict: Dict[str, Any],
request_id: Optional[Any] = None,
) -> "LiteLLMSendMessageResponse":
"""
Create a LiteLLMSendMessageResponse from a dict.
Args:
response_dict: Dict with A2A response structure
request_id: JSON-RPC request id to backfill when missing on error payloads
Returns:
LiteLLMSendMessageResponse with _hidden_params support
"""
return cls(**response_dict)
return cls(
**_normalize_a2a_jsonrpc_response(response_dict, request_id=request_id)
)

View file

@ -20,6 +20,7 @@ class ImageEditOptionalRequestParams(TypedDict, total=False):
response_format: Optional[Literal["url", "b64_json"]]
size: Optional[str]
user: Optional[str]
imageConfig: Optional[Dict[str, Any]]
class ImageEditRequestParams(ImageEditOptionalRequestParams, total=False):

View file

@ -203,6 +203,8 @@ class Status1(Enum):
completed = "completed"
failed = "failed"
cancelled = "cancelled"
incomplete = "incomplete"
budget_exceeded = "budget_exceeded"
class InteractionStatusUpdate(BaseModel):
@ -386,13 +388,13 @@ class ResponseModality(Enum):
class Status3(Enum):
UNSPECIFIED = "UNSPECIFIED"
IN_PROGRESS = "IN_PROGRESS"
REQUIRES_ACTION = "REQUIRES_ACTION"
COMPLETED = "COMPLETED"
FAILED = "FAILED"
CANCELLED = "CANCELLED"
INCOMPLETE = "INCOMPLETE"
IN_PROGRESS = "in_progress"
REQUIRES_ACTION = "requires_action"
COMPLETED = "completed"
FAILED = "failed"
CANCELLED = "cancelled"
INCOMPLETE = "incomplete"
BUDGET_EXCEEDED = "budget_exceeded"
class ModelOption(RootModel[str]):

View file

@ -1,7 +1,7 @@
from enum import Enum
from typing import Any, Dict, Iterable, List, Literal, Optional, Union
from typing import Any, Dict, List, Literal, Optional
from typing_extensions import Required, TypedDict
from typing_extensions import TypedDict
from .vertex_ai import (
GenerationConfig,
@ -171,6 +171,9 @@ class GeminiImageGenerationParameters(BaseModel):
aspectRatio: Optional[str] = None
"""Aspect ratio for generated images (e.g., '1:1', '16:9', '9:16', '4:3', '3:4')"""
imageSize: Optional[str] = None
"""Image size for generated images (e.g., '1K', '2K')"""
personGeneration: Optional[str] = None
"""Controls person generation in images"""

View file

@ -1084,6 +1084,7 @@ OpenAIImageGenerationOptionalParams = Literal[
"image_url",
"image_prompt_strength",
"aspect_ratio",
"imageConfig",
]
OpenAIImageEditOptionalParams = Literal[

View file

@ -20,6 +20,7 @@ class FunctionResponse(TypedDict, total=False):
id: str
name: Required[str]
response: Optional[dict]
parts: List["FunctionResponsePartType"]
class FunctionCall(TypedDict, total=False):
@ -40,6 +41,11 @@ class BlobType(TypedDict, total=False):
data: Required[str]
class FunctionResponsePartType(TypedDict, total=False):
inline_data: BlobType
file_data: FileDataType
class PartType(TypedDict, total=False):
text: str
inline_data: BlobType

View file

@ -68,12 +68,29 @@ class MCPServer(BaseModel):
access_groups: Optional[List[str]] = None
allow_all_keys: bool = False
available_on_public_internet: bool = True
# When True AND auth_type == oauth2, MCP requests targeting this server
# Explicit opt-in to upstream-delegated authentication for ``oauth2``
# servers. When ``auth_type == oauth2`` and this is ``True``, MCP requests
# bypass LiteLLM API-key/SSO auth (and the pre-emptive 401) so the client
# completes PKCE directly with the upstream MCP server. Honored only for
# auth_type=oauth2; ignored for any other auth_type. See
# MCPRequestHandler._target_servers_delegate_auth_to_upstream.
# completes PKCE directly with the upstream MCP server. See
# ``MCPRequestHandler._target_servers_delegate_auth_to_upstream``.
#
# Honored only for ``auth_type == oauth2``; ignored for any other
# ``auth_type``. OAuth pass-through for non-oauth2 servers
# (``auth_type in (None, MCPAuth.none)``) is a separate, explicit opt-in —
# see ``oauth_passthrough`` / ``is_oauth_passthrough``.
delegate_auth_to_upstream: bool = False
# Explicit opt-in to OAuth pass-through for non-oauth2 servers. When this
# is ``True`` AND ``auth_type in (None, MCPAuth.none)`` AND ``extra_headers``
# contains ``Authorization``, the gateway proxies upstream
# ``/.well-known/oauth-protected-resource`` metadata, emits spec-compliant
# 401 challenges when no bearer is supplied, and propagates upstream
# 401/403 responses instead of swallowing them. See ``is_oauth_passthrough``.
#
# Intentionally distinct from ``delegate_auth_to_upstream`` (oauth2-only):
# reusing that flag would silently change behavior for servers that forward
# ``Authorization`` for non-OAuth reasons (e.g. static bearer tokens). Must
# be set explicitly to avoid regressing servers that did not opt in.
oauth_passthrough: bool = False
is_byok: bool = False
byok_description: List[str] = []
byok_api_key_help_url: Optional[str] = None
@ -139,6 +156,42 @@ class MCPServer(BaseModel):
return False
@property
def is_oauth_passthrough(self) -> bool:
"""True iff the gateway should transparently forward upstream OAuth
(discovery + 401s) rather than participating as an authorization
server itself.
A server is pass-through for OAuth purposes when ALL three conditions
hold:
1. ``auth_type`` is ``None`` or ``MCPAuth.none`` (the gateway does
not manage OAuth for this server).
2. ``extra_headers`` includes ``Authorization`` — the admin has
opted this server into forwarding the client's bearer token
straight to the upstream MCP server.
3. ``oauth_passthrough`` is ``True`` — the admin has
explicitly opted into upstream-delegated OAuth semantics for
this server. This is the explicit detection flag: without it,
a server that merely forwards ``Authorization`` (e.g. for
static bearer tokens or custom auth schemes) keeps the
pre-PR behavior and is not treated as OAuth pass-through.
This is deliberately a separate flag from
``delegate_auth_to_upstream`` (which is oauth2-only) so enabling
pass-through here never changes behavior for oauth2 servers.
This is intentionally narrower than ``requires_per_user_auth``,
which also covers PATs (``x-api-key``, ``api-key``, ``apikey``).
Those are static credentials, not OAuth bearer tokens, so they
must not trigger upstream OAuth discovery or 401 propagation.
"""
if self.auth_type not in (None, MCPAuth.none):
return False
if not self.extra_headers:
return False
if self.oauth_passthrough is not True:
return False
return any(h.lower() == "authorization" for h in self.extra_headers)
@property
def has_token_exchange_config(self) -> bool:
"""True if this server is configured for OAuth2 token exchange (OBO / RFC 8693)."""

View file

@ -148,6 +148,7 @@ class ProviderSpecificModelInfo(TypedDict, total=False):
supports_xhigh_reasoning_effort: Optional[bool]
supports_max_reasoning_effort: Optional[bool]
supports_output_config: Optional[bool]
supports_image_size: Optional[bool]
bedrock_output_config_effort_ceiling: Optional[
Literal["low", "medium", "high", "max", "xhigh"]
]
@ -2580,6 +2581,12 @@ class StandardLoggingMCPToolCall(TypedDict, total=False):
Cost per query for the MCP server tool call
"""
mcp_session_id: Optional[str]
"""
The MCP `mcp-session-id` of the stateful session this tool call ran in, when
the client is driving a stateful session. Absent for stateless calls.
"""
class StandardLoggingVectorStoreRequest(TypedDict, total=False):
"""
@ -3294,6 +3301,8 @@ class LlmProviders(str, Enum):
V0 = "v0"
MORPH = "morph"
LAMBDA_AI = "lambda_ai"
INCEPTION = "inception"
TEXT_COMPLETION_INCEPTION = "text-completion-inception"
DEEPSEEK = "deepseek"
SAMBANOVA = "sambanova"
MARITALK = "maritalk"
@ -3358,6 +3367,7 @@ class LlmProviders(str, Enum):
AMAZON_NOVA = "amazon_nova"
A2A_AGENT = "a2a_agent"
LANGGRAPH = "langgraph"
LANGFLOW = "langflow"
MINIMAX = "minimax"
SYNTHETIC = "synthetic"
APERTIS = "apertis"

View file

@ -112,6 +112,66 @@ class VectorStoreSearchRequest(VectorStoreSearchOptionalRequestParams, total=Fal
query: Union[str, List[str]]
class VertexSearchDataStoreExtraBody(TypedDict, total=False):
"""
Native Discovery Engine ``SearchRequest`` fields callers may forward via
``extra_body`` when searching a Vertex AI Search **data store** serving
config (``.../dataStores/{id}/servingConfigs/default_config``).
The data store is scoped by the request URL path, so target-selecting
fields (``servingConfig``, ``branch``, ``entity``) are intentionally
omitted and rejected by the transformation layer. Engine/app-only fields
such as ``dataStoreSpecs`` and ``numResultsPerDataStore`` live on
``VertexSearchEngineExtraBody`` instead.
"""
query: str
pageSize: int
pageToken: str
offset: int
oneBoxPageSize: int
pageCategories: List[str]
imageQuery: Dict[str, Any]
filter: str
canonicalFilter: str
orderBy: str
userInfo: Dict[str, Any]
languageCode: str
facetSpecs: List[Dict[str, Any]]
boostSpec: Dict[str, Any]
params: Dict[str, Any]
queryExpansionSpec: Dict[str, Any]
spellCorrectionSpec: Dict[str, Any]
userPseudoId: str
contentSearchSpec: Dict[str, Any]
rankingExpression: str
rankingExpressionBackend: str
safeSearch: bool
userLabels: Dict[str, str]
naturalLanguageQueryUnderstandingSpec: Dict[str, Any]
searchAsYouTypeSpec: Dict[str, Any]
displaySpec: Dict[str, Any]
crowdingSpecs: List[Dict[str, Any]]
relevanceThreshold: str
relevanceScoreSpec: Dict[str, Any]
customRankingParams: Dict[str, Any]
class VertexSearchEngineExtraBody(VertexSearchDataStoreExtraBody, total=False):
"""
Native Discovery Engine ``SearchRequest`` fields callers may forward via
``extra_body`` when searching a Vertex AI Search **engine/app** serving
config (``.../engines/{id}/servingConfigs/default_serving_config``).
Inherits every data-store field and adds fields that only make sense when
an app fans out across multiple member data stores, e.g. ``dataStoreSpecs``
(per-store scoping/filtering) and ``numResultsPerDataStore``.
"""
dataStoreSpecs: List[Dict[str, Any]]
numResultsPerDataStore: int
# Vector Store Creation Types
class VectorStoreExpirationPolicy(TypedDict, total=False):
"""The expiration policy for a vector store"""

View file

@ -3147,6 +3147,7 @@ def get_optional_params_image_gen(
size: Optional[str] = None,
style: Optional[str] = None,
user: Optional[str] = None,
imageConfig: Optional[dict] = None,
custom_llm_provider: Optional[str] = None,
additional_drop_params: Optional[list] = None,
provider_config: Optional[BaseImageGenerationConfig] = None,
@ -3183,6 +3184,7 @@ def get_optional_params_image_gen(
"size": None,
"style": None,
"user": None,
"imageConfig": None,
}
non_default_params = _get_non_default_params(
@ -4547,6 +4549,18 @@ def get_optional_params( # noqa: PLR0915
),
)
elif custom_llm_provider == "text-completion-inception":
optional_params = litellm.InceptionTextCompletionConfig().map_openai_params(
non_default_params=non_default_params,
optional_params=optional_params,
model=model,
drop_params=(
drop_params
if drop_params is not None and isinstance(drop_params, bool)
else False
),
)
elif custom_llm_provider == "databricks":
optional_params = litellm.DatabricksConfig().map_openai_params(
non_default_params=non_default_params,
@ -6083,6 +6097,7 @@ def _get_model_info_helper( # noqa: PLR0915
"provider_specific_entry", None
),
uses_embed_content=_model_info.get("uses_embed_content", None),
supports_image_size=_model_info.get("supports_image_size", None),
)
except Exception as e:
verbose_logger.debug(f"Error getting model info: {e}")
@ -6637,6 +6652,14 @@ def validate_environment( # noqa: PLR0915
keys_in_environment = True
else:
missing_keys.append("CODESTRAL_API_KEY")
elif (
custom_llm_provider == "inception"
or custom_llm_provider == "text-completion-inception"
):
if "INCEPTION_API_KEY" in os.environ:
keys_in_environment = True
else:
missing_keys.append("INCEPTION_API_KEY")
elif custom_llm_provider == "deepseek":
if "DEEPSEEK_API_KEY" in os.environ:
keys_in_environment = True
@ -8291,6 +8314,7 @@ class ProviderConfigManager:
LlmProviders.XAI: (lambda: litellm.XAIChatConfig(), False),
LlmProviders.ZAI: (lambda: litellm.ZAIChatConfig(), False),
LlmProviders.LAMBDA_AI: (lambda: litellm.LambdaAIChatConfig(), False),
LlmProviders.INCEPTION: (lambda: litellm.InceptionChatConfig(), False),
LlmProviders.LLAMA: (lambda: litellm.LlamaAPIConfig(), False),
LlmProviders.TEXT_COMPLETION_OPENAI: (
lambda: litellm.OpenAITextCompletionConfig(),
@ -8356,6 +8380,10 @@ class ProviderConfigManager:
lambda: litellm.CodestralTextCompletionConfig(),
False,
),
LlmProviders.TEXT_COMPLETION_INCEPTION: (
lambda: litellm.InceptionTextCompletionConfig(),
False,
),
LlmProviders.SAMBANOVA: (lambda: litellm.SambanovaConfig(), False),
LlmProviders.MARITALK: (lambda: litellm.MaritalkConfig(), False),
LlmProviders.VLLM: (lambda: litellm.VLLMConfig(), False),
@ -8394,6 +8422,10 @@ class ProviderConfigManager:
lambda: ProviderConfigManager._get_langgraph_config(),
False,
),
LlmProviders.LANGFLOW: (
lambda: ProviderConfigManager._get_langflow_config(),
False,
),
}
@staticmethod
@ -8465,6 +8497,13 @@ class ProviderConfigManager:
return LangGraphConfig()
@staticmethod
def _get_langflow_config() -> BaseConfig:
"""Get LangFlow config."""
from litellm.llms.langflow.chat.transformation import LangFlowConfig
return LangFlowConfig()
@staticmethod
def get_provider_chat_config( # noqa: PLR0915
model: str,
@ -8917,6 +8956,8 @@ class ProviderConfigManager:
return litellm.FireworksAITextCompletionConfig()
elif LlmProviders.TOGETHER_AI == provider:
return litellm.TogetherAITextCompletionConfig()
elif LlmProviders.TEXT_COMPLETION_INCEPTION == provider:
return litellm.InceptionTextCompletionConfig()
return litellm.OpenAITextCompletionConfig()
@staticmethod

View file

@ -1075,6 +1075,7 @@
},
"eu.anthropic.claude-opus-4-6-v1": {
"cache_creation_input_token_cost": 6.875e-06,
"cache_creation_input_token_cost_above_1hr": 1.1e-05,
"cache_read_input_token_cost": 5.5e-07,
"input_cost_per_token": 5.5e-06,
"litellm_provider": "bedrock_converse",
@ -1104,6 +1105,7 @@
},
"au.anthropic.claude-opus-4-6-v1": {
"cache_creation_input_token_cost": 6.875e-06,
"cache_creation_input_token_cost_above_1hr": 1.1e-05,
"cache_read_input_token_cost": 5.5e-07,
"input_cost_per_token": 5.5e-06,
"litellm_provider": "bedrock_converse",
@ -1241,6 +1243,7 @@
},
"eu.anthropic.claude-opus-4-7": {
"cache_creation_input_token_cost": 6.875e-06,
"cache_creation_input_token_cost_above_1hr": 1.1e-05,
"cache_read_input_token_cost": 5.5e-07,
"input_cost_per_token": 5.5e-06,
"litellm_provider": "bedrock_converse",
@ -1271,6 +1274,7 @@
},
"au.anthropic.claude-opus-4-7": {
"cache_creation_input_token_cost": 6.875e-06,
"cache_creation_input_token_cost_above_1hr": 1.1e-05,
"cache_read_input_token_cost": 5.5e-07,
"input_cost_per_token": 5.5e-06,
"litellm_provider": "bedrock_converse",
@ -1543,6 +1547,7 @@
},
"eu.anthropic.claude-sonnet-4-6": {
"cache_creation_input_token_cost": 4.125e-06,
"cache_creation_input_token_cost_above_1hr": 6.6e-06,
"cache_read_input_token_cost": 3.3e-07,
"input_cost_per_token": 3.3e-06,
"litellm_provider": "bedrock_converse",
@ -1571,6 +1576,7 @@
},
"au.anthropic.claude-sonnet-4-6": {
"cache_creation_input_token_cost": 4.125e-06,
"cache_creation_input_token_cost_above_1hr": 6.6e-06,
"cache_read_input_token_cost": 3.3e-07,
"input_cost_per_token": 3.3e-06,
"litellm_provider": "bedrock_converse",
@ -1599,6 +1605,7 @@
},
"jp.anthropic.claude-sonnet-4-6": {
"cache_creation_input_token_cost": 4.125e-06,
"cache_creation_input_token_cost_above_1hr": 6.6e-06,
"cache_read_input_token_cost": 3.3e-07,
"input_cost_per_token": 3.3e-06,
"litellm_provider": "bedrock_converse",
@ -1995,11 +2002,13 @@
},
"au.anthropic.claude-sonnet-4-5-20250929-v1:0": {
"cache_creation_input_token_cost": 4.125e-06,
"cache_creation_input_token_cost_above_1hr": 6.6e-06,
"cache_read_input_token_cost": 3.3e-07,
"input_cost_per_token": 3.3e-06,
"input_cost_per_token_above_200k_tokens": 6.6e-06,
"output_cost_per_token_above_200k_tokens": 2.475e-05,
"cache_creation_input_token_cost_above_200k_tokens": 8.25e-06,
"cache_creation_input_token_cost_above_1hr_above_200k_tokens": 1.32e-05,
"cache_read_input_token_cost_above_200k_tokens": 6.6e-07,
"litellm_provider": "bedrock_converse",
"max_input_tokens": 200000,
@ -7494,6 +7503,27 @@
"supports_video_input": true,
"supports_vision": true
},
"azure_ai/kimi-k2.6": {
"input_cost_per_token": 9.5e-07,
"litellm_provider": "azure_ai",
"max_input_tokens": 262144,
"max_output_tokens": 262144,
"max_tokens": 262144,
"mode": "chat",
"output_cost_per_token": 4e-06,
"source": "https://techcommunity.microsoft.com/blog/azure-ai-foundry-blog/introducing-kimi-k2-6-in-microsoft-foundry/4513125",
"supported_modalities": [
"text",
"image"
],
"supported_output_modalities": [
"text"
],
"supports_function_calling": true,
"supports_reasoning": true,
"supports_tool_choice": true,
"supports_vision": true
},
"azure_ai/ministral-3b": {
"input_cost_per_token": 4e-08,
"litellm_provider": "azure_ai",
@ -12682,7 +12712,8 @@
"litellm_provider": "deepinfra",
"mode": "chat",
"supports_tool_choice": true,
"supports_function_calling": true
"supports_function_calling": true,
"supports_image_size": false
},
"deepinfra/google/gemini-2.5-pro": {
"max_tokens": 1000000,
@ -13583,6 +13614,7 @@
},
"eu.anthropic.claude-haiku-4-5-20251001-v1:0": {
"cache_creation_input_token_cost": 1.375e-06,
"cache_creation_input_token_cost_above_1hr": 2.2e-06,
"cache_read_input_token_cost": 1.1e-07,
"input_cost_per_token": 1.1e-06,
"deprecation_date": "2026-10-15",
@ -13787,11 +13819,13 @@
},
"eu.anthropic.claude-sonnet-4-5-20250929-v1:0": {
"cache_creation_input_token_cost": 4.125e-06,
"cache_creation_input_token_cost_above_1hr": 6.6e-06,
"cache_read_input_token_cost": 3.3e-07,
"input_cost_per_token": 3.3e-06,
"input_cost_per_token_above_200k_tokens": 6.6e-06,
"output_cost_per_token_above_200k_tokens": 2.475e-05,
"cache_creation_input_token_cost_above_200k_tokens": 8.25e-06,
"cache_creation_input_token_cost_above_1hr_above_200k_tokens": 1.32e-05,
"cache_read_input_token_cost_above_200k_tokens": 6.6e-07,
"litellm_provider": "bedrock_converse",
"max_input_tokens": 200000,
@ -15006,7 +15040,8 @@
"search_context_size_medium": 0.035,
"search_context_size_high": 0.035
},
"supports_service_tier": true
"supports_service_tier": true,
"supports_image_size": false
},
"gemini-2.5-flash-image": {
"cache_read_input_token_cost": 3e-08,
@ -15056,7 +15091,8 @@
"supports_vision": true,
"supports_web_search": false,
"tpm": 8000000,
"supports_service_tier": true
"supports_service_tier": true,
"supports_image_size": false
},
"gemini-3-pro-image-preview": {
"input_cost_per_image": 0.0011,
@ -15345,7 +15381,8 @@
"search_context_size_medium": 0.035,
"search_context_size_high": 0.035
},
"supports_service_tier": true
"supports_service_tier": true,
"supports_image_size": false
},
"gemini-2.5-flash-lite-preview-09-2025": {
"cache_read_input_token_cost": 1e-08,
@ -15395,7 +15432,8 @@
"search_context_size_low": 0.035,
"search_context_size_medium": 0.035,
"search_context_size_high": 0.035
}
},
"supports_image_size": false
},
"gemini-2.5-flash-preview-09-2025": {
"cache_read_input_token_cost": 7.5e-08,
@ -15445,7 +15483,8 @@
"search_context_size_low": 0.035,
"search_context_size_medium": 0.035,
"search_context_size_high": 0.035
}
},
"supports_image_size": false
},
"gemini-live-2.5-flash-preview-native-audio-09-2025": {
"cache_read_input_token_cost": 7.5e-08,
@ -15596,7 +15635,8 @@
"search_context_size_low": 0.035,
"search_context_size_medium": 0.035,
"search_context_size_high": 0.035
}
},
"supports_image_size": false
},
"gemini-2.5-pro": {
"cache_read_input_token_cost": 1.25e-07,
@ -16606,7 +16646,8 @@
"search_context_size_medium": 0.035,
"search_context_size_high": 0.035
},
"supports_service_tier": true
"supports_service_tier": true,
"supports_image_size": false
},
"gemini/gemini-2.5-flash-image": {
"cache_read_input_token_cost": 3e-08,
@ -16662,7 +16703,8 @@
"search_context_size_medium": 0.035,
"search_context_size_high": 0.035
},
"supports_service_tier": true
"supports_service_tier": true,
"supports_image_size": false
},
"gemini/gemini-3-pro-image-preview": {
"input_cost_per_image": 0.0011,
@ -16841,7 +16883,8 @@
"search_context_size_medium": 0.035,
"search_context_size_high": 0.035
},
"supports_service_tier": true
"supports_service_tier": true,
"supports_image_size": false
},
"gemini/gemini-2.5-flash-lite-preview-09-2025": {
"cache_read_input_token_cost": 1e-08,
@ -16893,7 +16936,8 @@
"search_context_size_low": 0.035,
"search_context_size_medium": 0.035,
"search_context_size_high": 0.035
}
},
"supports_image_size": false
},
"gemini/gemini-2.5-flash-preview-09-2025": {
"cache_read_input_token_cost": 7.5e-08,
@ -16945,7 +16989,8 @@
"search_context_size_low": 0.035,
"search_context_size_medium": 0.035,
"search_context_size_high": 0.035
}
},
"supports_image_size": false
},
"gemini/gemini-flash-latest": {
"cache_read_input_token_cost": 7.5e-08,
@ -17102,7 +17147,8 @@
"search_context_size_low": 0.035,
"search_context_size_medium": 0.035,
"search_context_size_high": 0.035
}
},
"supports_image_size": false
},
"gemini/gemini-2.5-flash-preview-tts": {
"input_cost_per_token": 3e-07,
@ -23086,11 +23132,13 @@
},
"jp.anthropic.claude-sonnet-4-5-20250929-v1:0": {
"cache_creation_input_token_cost": 4.125e-06,
"cache_creation_input_token_cost_above_1hr": 6.6e-06,
"cache_read_input_token_cost": 3.3e-07,
"input_cost_per_token": 3.3e-06,
"input_cost_per_token_above_200k_tokens": 6.6e-06,
"output_cost_per_token_above_200k_tokens": 2.475e-05,
"cache_creation_input_token_cost_above_200k_tokens": 8.25e-06,
"cache_creation_input_token_cost_above_1hr_above_200k_tokens": 1.32e-05,
"cache_read_input_token_cost_above_200k_tokens": 6.6e-07,
"litellm_provider": "bedrock_converse",
"max_input_tokens": 200000,
@ -23116,6 +23164,7 @@
},
"jp.anthropic.claude-haiku-4-5-20251001-v1:0": {
"cache_creation_input_token_cost": 1.375e-06,
"cache_creation_input_token_cost_above_1hr": 2.2e-06,
"cache_read_input_token_cost": 1.1e-07,
"input_cost_per_token": 1.1e-06,
"litellm_provider": "bedrock_converse",
@ -23228,6 +23277,31 @@
"supports_system_messages": true,
"supports_tool_choice": true
},
"inception/mercury-2": {
"cache_read_input_token_cost": 2.5e-08,
"input_cost_per_token": 2.5e-07,
"litellm_provider": "inception",
"max_input_tokens": 128000,
"max_output_tokens": 50000,
"max_tokens": 50000,
"mode": "chat",
"output_cost_per_token": 7.5e-07,
"supports_function_calling": true,
"supports_prompt_caching": true,
"supports_response_schema": true,
"supports_system_messages": true,
"supports_tool_choice": true
},
"text-completion-inception/mercury-edit-2": {
"cache_read_input_token_cost": 2.5e-08,
"input_cost_per_token": 2.5e-07,
"litellm_provider": "text-completion-inception",
"max_input_tokens": 32000,
"max_output_tokens": 8192,
"max_tokens": 8192,
"mode": "completion",
"output_cost_per_token": 7.5e-07
},
"lambda_ai/deepseek-llama3.3-70b": {
"input_cost_per_token": 2e-07,
"litellm_provider": "lambda_ai",
@ -26449,7 +26523,8 @@
"supports_function_calling": true,
"supports_response_schema": true,
"supports_vision": true,
"supports_native_streaming": true
"supports_native_streaming": true,
"supports_image_size": false
},
"oci/google.gemini-2.5-pro": {
"input_cost_per_token": 1.25e-06,
@ -26477,7 +26552,8 @@
"supports_function_calling": true,
"supports_response_schema": true,
"supports_vision": true,
"supports_native_streaming": true
"supports_native_streaming": true,
"supports_image_size": false
},
"oci/cohere.command-a-vision": {
"input_cost_per_token": 1.56e-06,
@ -27513,7 +27589,8 @@
"supports_response_schema": true,
"supports_system_messages": true,
"supports_tool_choice": true,
"supports_vision": true
"supports_vision": true,
"supports_image_size": false
},
"openrouter/google/gemini-2.5-pro": {
"input_cost_per_audio_token": 7e-07,
@ -29383,7 +29460,8 @@
"mode": "responses",
"supports_web_search": true,
"supports_reasoning": false,
"supports_function_calling": true
"supports_function_calling": true,
"supports_image_size": false
},
"perplexity/xai/grok-4-1-fast-non-reasoning": {
"litellm_provider": "perplexity",
@ -29965,7 +30043,8 @@
"supports_vision": true,
"supports_system_messages": true,
"supports_tool_choice": true,
"supports_response_schema": true
"supports_response_schema": true,
"supports_image_size": false
},
"replicate/openai/gpt-oss-120b": {
"input_cost_per_token": 1.8e-07,
@ -31784,6 +31863,7 @@
},
"au.anthropic.claude-haiku-4-5-20251001-v1:0": {
"cache_creation_input_token_cost": 1.375e-06,
"cache_creation_input_token_cost_above_1hr": 2.2e-06,
"cache_read_input_token_cost": 1.1e-07,
"input_cost_per_token": 1.1e-06,
"litellm_provider": "bedrock_converse",
@ -32663,7 +32743,8 @@
"supports_vision": true,
"supports_function_calling": true,
"supports_tool_choice": true,
"supports_response_schema": true
"supports_response_schema": true,
"supports_image_size": false
},
"vercel_ai_gateway/google/gemini-2.5-pro": {
"input_cost_per_token": 2.5e-06,
@ -33430,6 +33511,7 @@
},
"vertex_ai/claude-haiku-4-5": {
"cache_creation_input_token_cost": 1.25e-06,
"cache_creation_input_token_cost_above_1hr": 2e-06,
"cache_read_input_token_cost": 1e-07,
"input_cost_per_token": 1e-06,
"litellm_provider": "vertex_ai-anthropic_models",
@ -33451,6 +33533,7 @@
},
"vertex_ai/claude-haiku-4-5@20251001": {
"cache_creation_input_token_cost": 1.25e-06,
"cache_creation_input_token_cost_above_1hr": 2e-06,
"cache_read_input_token_cost": 1e-07,
"input_cost_per_token": 1e-06,
"litellm_provider": "vertex_ai-anthropic_models",
@ -33501,6 +33584,7 @@
},
"vertex_ai/claude-3-7-sonnet@20250219": {
"cache_creation_input_token_cost": 3.75e-06,
"cache_creation_input_token_cost_above_1hr": 6e-06,
"cache_read_input_token_cost": 3e-07,
"deprecation_date": "2026-05-11",
"input_cost_per_token": 3e-06,
@ -33600,6 +33684,7 @@
},
"vertex_ai/claude-opus-4": {
"cache_creation_input_token_cost": 1.875e-05,
"cache_creation_input_token_cost_above_1hr": 3e-05,
"cache_read_input_token_cost": 1.5e-06,
"input_cost_per_token": 1.5e-05,
"litellm_provider": "vertex_ai-anthropic_models",
@ -33625,6 +33710,7 @@
},
"vertex_ai/claude-opus-4-1": {
"cache_creation_input_token_cost": 1.875e-05,
"cache_creation_input_token_cost_above_1hr": 3e-05,
"cache_read_input_token_cost": 1.5e-06,
"input_cost_per_token": 1.5e-05,
"input_cost_per_token_batches": 7.5e-06,
@ -33642,6 +33728,7 @@
},
"vertex_ai/claude-opus-4-1@20250805": {
"cache_creation_input_token_cost": 1.875e-05,
"cache_creation_input_token_cost_above_1hr": 3e-05,
"cache_read_input_token_cost": 1.5e-06,
"input_cost_per_token": 1.5e-05,
"input_cost_per_token_batches": 7.5e-06,
@ -33659,6 +33746,7 @@
},
"vertex_ai/claude-opus-4-5": {
"cache_creation_input_token_cost": 6.25e-06,
"cache_creation_input_token_cost_above_1hr": 1e-05,
"cache_read_input_token_cost": 5e-07,
"input_cost_per_token": 5e-06,
"litellm_provider": "vertex_ai-anthropic_models",
@ -33685,6 +33773,7 @@
},
"vertex_ai/claude-opus-4-5@20251101": {
"cache_creation_input_token_cost": 6.25e-06,
"cache_creation_input_token_cost_above_1hr": 1e-05,
"cache_read_input_token_cost": 5e-07,
"input_cost_per_token": 5e-06,
"litellm_provider": "vertex_ai-anthropic_models",
@ -33712,6 +33801,7 @@
},
"vertex_ai/claude-opus-4-6": {
"cache_creation_input_token_cost": 6.25e-06,
"cache_creation_input_token_cost_above_1hr": 1e-05,
"cache_read_input_token_cost": 5e-07,
"input_cost_per_token": 5e-06,
"litellm_provider": "vertex_ai-anthropic_models",
@ -33739,6 +33829,7 @@
},
"vertex_ai/claude-opus-4-6@default": {
"cache_creation_input_token_cost": 6.25e-06,
"cache_creation_input_token_cost_above_1hr": 1e-05,
"cache_read_input_token_cost": 5e-07,
"input_cost_per_token": 5e-06,
"litellm_provider": "vertex_ai-anthropic_models",
@ -33766,6 +33857,7 @@
},
"vertex_ai/claude-opus-4-7": {
"cache_creation_input_token_cost": 6.25e-06,
"cache_creation_input_token_cost_above_1hr": 1e-05,
"cache_read_input_token_cost": 5e-07,
"input_cost_per_token": 5e-06,
"litellm_provider": "vertex_ai-anthropic_models",
@ -33793,6 +33885,7 @@
},
"vertex_ai/claude-opus-4-7@default": {
"cache_creation_input_token_cost": 6.25e-06,
"cache_creation_input_token_cost_above_1hr": 1e-05,
"cache_read_input_token_cost": 5e-07,
"input_cost_per_token": 5e-06,
"litellm_provider": "vertex_ai-anthropic_models",
@ -33876,6 +33969,7 @@
},
"vertex_ai/claude-sonnet-4-5": {
"cache_creation_input_token_cost": 3.75e-06,
"cache_creation_input_token_cost_above_1hr": 6e-06,
"cache_read_input_token_cost": 3e-07,
"input_cost_per_token": 3e-06,
"input_cost_per_token_above_200k_tokens": 6e-06,
@ -33902,6 +33996,7 @@
},
"vertex_ai/claude-sonnet-4-6": {
"cache_creation_input_token_cost": 3.75e-06,
"cache_creation_input_token_cost_above_1hr": 6e-06,
"cache_read_input_token_cost": 3e-07,
"input_cost_per_token": 3e-06,
"litellm_provider": "vertex_ai-anthropic_models",
@ -33929,6 +34024,7 @@
},
"vertex_ai/claude-sonnet-4-5@20250929": {
"cache_creation_input_token_cost": 3.75e-06,
"cache_creation_input_token_cost_above_1hr": 6e-06,
"cache_read_input_token_cost": 3e-07,
"input_cost_per_token": 3e-06,
"input_cost_per_token_above_200k_tokens": 6e-06,
@ -33956,6 +34052,7 @@
},
"vertex_ai/claude-opus-4@20250514": {
"cache_creation_input_token_cost": 1.875e-05,
"cache_creation_input_token_cost_above_1hr": 3e-05,
"cache_read_input_token_cost": 1.5e-06,
"input_cost_per_token": 1.5e-05,
"litellm_provider": "vertex_ai-anthropic_models",
@ -33981,6 +34078,7 @@
},
"vertex_ai/claude-sonnet-4": {
"cache_creation_input_token_cost": 3.75e-06,
"cache_creation_input_token_cost_above_1hr": 6e-06,
"cache_read_input_token_cost": 3e-07,
"input_cost_per_token": 3e-06,
"input_cost_per_token_above_200k_tokens": 6e-06,
@ -34010,6 +34108,7 @@
},
"vertex_ai/claude-sonnet-4@20250514": {
"cache_creation_input_token_cost": 3.75e-06,
"cache_creation_input_token_cost_above_1hr": 6e-06,
"cache_read_input_token_cost": 3e-07,
"input_cost_per_token": 3e-06,
"input_cost_per_token_above_200k_tokens": 6e-06,
@ -34217,7 +34316,8 @@
"supports_url_context": true,
"supports_vision": true,
"supports_web_search": false,
"tpm": 8000000
"tpm": 8000000,
"supports_image_size": false
},
"vertex_ai/gemini-3-pro-image-preview": {
"input_cost_per_image": 0.0011,
@ -41035,6 +41135,7 @@
},
"vertex_ai/claude-sonnet-4-6@default": {
"cache_creation_input_token_cost": 3.75e-06,
"cache_creation_input_token_cost_above_1hr": 6e-06,
"cache_read_input_token_cost": 3e-07,
"input_cost_per_token": 3e-06,
"litellm_provider": "vertex_ai-anthropic_models",

View file

@ -1273,6 +1273,24 @@
"interactions": true
}
},
"inception": {
"display_name": "Inception (`inception`)",
"url": "https://docs.litellm.ai/docs/providers/inception",
"endpoints": {
"chat_completions": true,
"messages": true,
"responses": true,
"embeddings": false,
"image_generations": false,
"audio_transcriptions": false,
"audio_speech": false,
"moderations": false,
"batches": false,
"rerank": false,
"a2a": true,
"interactions": true
}
},
"infinity": {
"display_name": "Infinity (`infinity`)",
"url": "https://docs.litellm.ai/docs/providers/infinity",
@ -2430,6 +2448,24 @@
"interactions": true
}
},
"langflow": {
"display_name": "LangFlow (`langflow`)",
"url": "https://docs.litellm.ai/docs/providers/langflow",
"endpoints": {
"chat_completions": true,
"messages": false,
"responses": false,
"embeddings": false,
"image_generations": false,
"audio_transcriptions": false,
"audio_speech": false,
"moderations": false,
"batches": false,
"rerank": false,
"a2a": true,
"interactions": false
}
},
"vertex_ai/agent_engine": {
"display_name": "Vertex AI Agent Engine (`vertex_ai/agent_engine`)",
"url": "https://docs.litellm.ai/docs/providers/vertex_ai_agent_engine",

View file

@ -325,6 +325,7 @@ model LiteLLM_MCPServerTable {
allow_all_keys Boolean @default(false)
available_on_public_internet Boolean @default(true)
delegate_auth_to_upstream Boolean @default(false)
oauth_passthrough Boolean @default(false)
is_byok Boolean @default(false)
byok_description String[] @default([])
byok_api_key_help_url String?

View file

@ -53,6 +53,7 @@ CASSETTE_CACHE_HIGH_WATER_FRACTION = 0.85
SAFE_BODY_MATCHER_NAME = "safe_body"
KEY_FINGERPRINT_MATCHER_NAME = "key_fingerprint"
TOLERANT_QUERY_MATCHER_NAME = "tolerant_query"
TOLERANT_PATH_MATCHER_NAME = "tolerant_path"
KEY_FINGERPRINT_HEADER = "x-litellm-key-fp"
VCR_DIAG_DIR_ENV = "LITELLM_VCR_DIAG_DIR"
@ -411,6 +412,7 @@ def _canonical_body(request) -> tuple[bytes, str]:
_VCR_UUID_RE = re.compile(
rb"[0-9a-fA-F]{8}-[0-9a-fA-F]{4}-[0-9a-fA-F]{4}-[0-9a-fA-F]{4}-[0-9a-fA-F]{12}"
)
_VCR_LITELLM_BATCH_JOB_RE = re.compile(rb"litellm-batch-[0-9a-fA-F]{8}")
# ISO-8601 timestamps, e.g. ``2026-05-25T03:40:37.262045Z`` /
# ``2026-05-25T03:40:37+00:00``.
_VCR_ISO_TS_RE = re.compile(
@ -436,6 +438,7 @@ def _normalize_volatile_tokens(body: bytes) -> bytes:
if not body:
return body
body = _VCR_UUID_RE.sub(b"<vcr-uuid>", body)
body = _VCR_LITELLM_BATCH_JOB_RE.sub(b"litellm-batch-<vcr-id>", body)
body = _VCR_ISO_TS_RE.sub(b"<vcr-iso-ts>", body)
body = _VCR_UNIX_MS_RE.sub(b"<vcr-unix-ms>", body)
body = _VCR_UNIX_FLOAT_RE.sub(b"<vcr-unix-float>", body)
@ -1059,6 +1062,54 @@ def _tolerant_query_matcher(r1, r2) -> None:
_vcr_matchers.query(r1, r2)
_BEDROCK_MANAGED_S3_PATH_RE = re.compile(
r"(?P<prefix>(?:^|/)(?:litellm-bedrock-files/[^/?#]+-|litellm-bedrock-files-[^/?#]+-))"
r"[0-9a-fA-F]{8}-[0-9a-fA-F]{4}-[0-9a-fA-F]{4}-[0-9a-fA-F]{4}-[0-9a-fA-F]{12}"
r"(?P<suffix>\.jsonl)"
)
def _request_path_for_matcher(request) -> str:
path = getattr(request, "path", None)
if path is not None:
return str(path)
uri = getattr(request, "uri", None) or getattr(request, "url", "") or ""
uri = str(uri)
if not uri:
return ""
if "//" in uri:
rest = uri.split("//", 1)[1]
path_part = "/" + rest.split("/", 1)[1] if "/" in rest else "/"
else:
path_part = uri
return path_part.split("?", 1)[0]
def _normalize_volatile_path(path: str) -> str:
return _BEDROCK_MANAGED_S3_PATH_RE.sub(
lambda match: f"{match.group('prefix')}<vcr-uuid>{match.group('suffix')}",
path,
)
def _tolerant_path_matcher(r1, r2) -> None:
"""vcrpy's ``path`` matcher, plus LiteLLM-managed Bedrock S3 upload UUIDs.
Bedrock batch file uploads use object keys like
``litellm-bedrock-files-{model}-{uuid}.jsonl`` (and older cassettes may
contain ``litellm-bedrock-files/{model}-{uuid}.jsonl``). The UUID is
generated client-side before the S3 PUT, so strict path matching makes
every replay miss even when the JSONL body and all provider semantics are
identical.
"""
path1 = _normalize_volatile_path(_request_path_for_matcher(r1))
path2 = _normalize_volatile_path(_request_path_for_matcher(r2))
if path1 == path2:
return
_vcr_matchers.path(r1, r2)
def vcr_config_dict() -> dict:
return {
"decode_compressed_response": True,
@ -1069,7 +1120,7 @@ def vcr_config_dict() -> dict:
"scheme",
"host",
"port",
"path",
TOLERANT_PATH_MATCHER_NAME,
TOLERANT_QUERY_MATCHER_NAME,
KEY_FINGERPRINT_MATCHER_NAME,
SAFE_BODY_MATCHER_NAME,
@ -1136,6 +1187,7 @@ def register_persister_if_enabled(vcr) -> None:
vcr.register_matcher(SAFE_BODY_MATCHER_NAME, _safe_body_matcher)
vcr.register_matcher(KEY_FINGERPRINT_MATCHER_NAME, _key_fingerprint_matcher)
vcr.register_matcher(TOLERANT_QUERY_MATCHER_NAME, _tolerant_query_matcher)
vcr.register_matcher(TOLERANT_PATH_MATCHER_NAME, _tolerant_path_matcher)
patch_vcrpy_aiohttp_record_path()
patch_vcrpy_cassette_load_guard()
global _atexit_banner_registered
@ -1647,13 +1699,18 @@ def _is_live_call_host(host: str) -> bool:
return False
if any(host.endswith(suffix) for suffix in _LIVE_CALL_HOST_SUFFIXES):
return True
# AWS Bedrock endpoints are ``bedrock-runtime[-fips].{region}.amazonaws.com``
# (region between ``bedrock-runtime`` and ``amazonaws.com``), so plain
# suffix matching can't catch them.
if host.endswith(".amazonaws.com") and host.split(".", 1)[0].startswith(
"bedrock-runtime"
):
return True
if host.endswith(".amazonaws.com"):
first_label = host.split(".", 1)[0]
# AWS Bedrock control/runtime endpoints are
# ``bedrock[-runtime][-fips].{region}.amazonaws.com`` (region between
# the service label and ``amazonaws.com``), so plain suffix matching
# can't catch them.
if first_label.startswith("bedrock"):
return True
# Bedrock batch file upload/download uses real S3. Treat those as part
# of the paid provider path so unmarked batch tests surface as leaks.
if first_label in {"s3", "s3-fips"} or ".s3." in host or ".s3-" in host:
return True
return False

Some files were not shown because too many files have changed in this diff Show more