mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
Merge branch 'BerriAI:litellm_internal_staging' into fix/realtime-usage-detail-keys
This commit is contained in:
commit
8dd404244d
247 changed files with 25784 additions and 1795 deletions
|
|
@ -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
|
||||
|
|
|
|||
2
Makefile
2
Makefile
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -0,0 +1,2 @@
|
|||
-- AlterTable
|
||||
ALTER TABLE "LiteLLM_MCPServerTable" ADD COLUMN IF NOT EXISTS "oauth_passthrough" BOOLEAN NOT NULL DEFAULT false;
|
||||
|
|
@ -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?
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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", {})
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
0
litellm/a2a_protocol/providers/langflow/__init__.py
Normal file
0
litellm/a2a_protocol/providers/langflow/__init__.py
Normal file
62
litellm/a2a_protocol/providers/langflow/config.py
Normal file
62
litellm/a2a_protocol/providers/langflow/config.py
Normal 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
|
||||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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 []
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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)}"
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
]
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
}
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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]]:
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
73
litellm/llms/gemini/image_usage_transformation.py
Normal file
73
litellm/llms/gemini/image_usage_transformation.py
Normal 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)
|
||||
0
litellm/llms/inception/__init__.py
Normal file
0
litellm/llms/inception/__init__.py
Normal file
0
litellm/llms/inception/chat/__init__.py
Normal file
0
litellm/llms/inception/chat/__init__.py
Normal file
54
litellm/llms/inception/chat/transformation.py
Normal file
54
litellm/llms/inception/chat/transformation.py
Normal 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
|
||||
0
litellm/llms/inception/completion/__init__.py
Normal file
0
litellm/llms/inception/completion/__init__.py
Normal file
43
litellm/llms/inception/completion/transformation.py
Normal file
43
litellm/llms/inception/completion/transformation.py
Normal 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
|
||||
1
litellm/llms/langflow/__init__.py
Normal file
1
litellm/llms/langflow/__init__.py
Normal file
|
|
@ -0,0 +1 @@
|
|||
"""LangFlow LLM provider for LiteLLM."""
|
||||
37
litellm/llms/langflow/a2a.py
Normal file
37
litellm/llms/langflow/a2a.py
Normal 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
|
||||
1
litellm/llms/langflow/chat/__init__.py
Normal file
1
litellm/llms/langflow/chat/__init__.py
Normal file
|
|
@ -0,0 +1 @@
|
|||
"""LangFlow chat transformation."""
|
||||
327
litellm/llms/langflow/chat/transformation.py
Normal file
327
litellm/llms/langflow/chat/transformation.py
Normal 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
|
||||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
80
litellm/proxy/_experimental/mcp_server/exceptions.py
Normal file
80
litellm/proxy/_experimental/mcp_server/exceptions.py
Normal 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,
|
||||
)
|
||||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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>"
|
||||
```
|
||||
|
||||
|
|
|
|||
|
|
@ -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],
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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)}")
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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!"
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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", "")
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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?
|
||||
|
|
|
|||
|
|
@ -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))
|
||||
|
|
|
|||
0
litellm/proxy/shutdown/__init__.py
Normal file
0
litellm/proxy/shutdown/__init__.py
Normal file
174
litellm/proxy/shutdown/graceful_shutdown_manager.py
Normal file
174
litellm/proxy/shutdown/graceful_shutdown_manager.py
Normal 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
|
||||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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).
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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]):
|
||||
|
|
|
|||
|
|
@ -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"""
|
||||
|
||||
|
|
|
|||
|
|
@ -1084,6 +1084,7 @@ OpenAIImageGenerationOptionalParams = Literal[
|
|||
"image_url",
|
||||
"image_prompt_strength",
|
||||
"aspect_ratio",
|
||||
"imageConfig",
|
||||
]
|
||||
|
||||
OpenAIImageEditOptionalParams = Literal[
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)."""
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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"""
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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?
|
||||
|
|
|
|||
|
|
@ -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
Loading…
Add table
Reference in a new issue