From b98a6562544f6dafb9734ecebf1f186cd5371d44 Mon Sep 17 00:00:00 2001 From: Yassin Kortam Date: Tue, 2 Jun 2026 11:45:36 -0700 Subject: [PATCH 01/27] Add MCP semantic conventions to otelv2 (#29468) * Add MCP semantic conventions to otelv2 Emit OpenTelemetry GenAI MCP tool-call spans from the v2 logger. A closed call_mcp_tool request now produces a CLIENT span named "tools/call {tool}" carrying mcp.method.name, gen_ai.operation.name=execute_tool, gen_ai.tool.name, the upstream server name, and (opt-in, content-gated) tool arguments/result. Adds the MCP and JSON-RPC attribute vocabulary to the semconv module, an MCPToolCallSpanData payload built from StandardLoggingMCPToolCall, an MCP_TOOL_CALL span role, and mapper support. * Complete the MCP span-attribute vocabulary in otelv2 semconv Add the remaining OTel GenAI MCP semconv attribute keys: gen_ai.prompt.name, the network.* transport keys with their well-known NetworkTransport values, and the client.* peer keys for MCP server spans. A test pins the full vocabulary so a dropped or renamed key fails loudly. * Populate mcp.session.id on MCP tool-call spans Capture the mcp-session-id header (case-insensitively) at the tool-call entry point and thread it through StandardLoggingMCPToolCall into the span, so spans for stateful MCP sessions carry mcp.session.id. Stateless calls have no such header and the attribute is simply absent. * Test that stateless MCP calls omit mcp.session.id --------- Co-authored-by: Claude --- litellm/integrations/otel/__init__.py | 18 ++- litellm/integrations/otel/emitter.py | 25 +++- litellm/integrations/otel/logger.py | 49 +++++- litellm/integrations/otel/mappers/base.py | 3 +- litellm/integrations/otel/mappers/genai.py | 24 ++- litellm/integrations/otel/model/payloads.py | 63 ++++++++ litellm/integrations/otel/model/semconv.py | 76 +++++++++- litellm/integrations/otel/model/spans.py | 12 ++ .../proxy/_experimental/mcp_server/server.py | 22 ++- litellm/types/utils.py | 6 + .../integrations/otel/test_otel_v2_logger.py | 122 +++++++++++++++ .../otel/test_otel_v2_sources_of_truth.py | 140 +++++++++++++++++- .../mcp_server/test_mcp_session_logging.py | 19 +++ 13 files changed, 563 insertions(+), 16 deletions(-) create mode 100644 tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_session_logging.py diff --git a/litellm/integrations/otel/__init__.py b/litellm/integrations/otel/__init__.py index 42a84a85fbd..da3ce4af3e7 100644 --- a/litellm/integrations/otel/__init__.py +++ b/litellm/integrations/otel/__init__.py @@ -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", ] diff --git a/litellm/integrations/otel/emitter.py b/litellm/integrations/otel/emitter.py index cae6514efdf..7fb7be7ab84 100644 --- a/litellm/integrations/otel/emitter.py +++ b/litellm/integrations/otel/emitter.py @@ -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): diff --git a/litellm/integrations/otel/logger.py b/litellm/integrations/otel/logger.py index d7058b34d50..de5e7a7b8ab 100644 --- a/litellm/integrations/otel/logger.py +++ b/litellm/integrations/otel/logger.py @@ -31,8 +31,10 @@ from litellm.integrations.otel.model.metadata import ( 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 +45,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 +205,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], diff --git a/litellm/integrations/otel/mappers/base.py b/litellm/integrations/otel/mappers/base.py index e8fb5af9797..dfdaf77a83e 100644 --- a/litellm/integrations/otel/mappers/base.py +++ b/litellm/integrations/otel/mappers/base.py @@ -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 diff --git a/litellm/integrations/otel/mappers/genai.py b/litellm/integrations/otel/mappers/genai.py index 57fa51ea1fb..6c61feced4d 100644 --- a/litellm/integrations/otel/mappers/genai.py +++ b/litellm/integrations/otel/mappers/genai.py @@ -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(): diff --git a/litellm/integrations/otel/model/payloads.py b/litellm/integrations/otel/model/payloads.py index 65b50d0fc12..08fce09868b 100644 --- a/litellm/integrations/otel/model/payloads.py +++ b/litellm/integrations/otel/model/payloads.py @@ -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,66 @@ 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 = cast(Mapping[str, object], payload.get("mcp_tool_call_metadata") or {}) + 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 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(payload.get("mcp_tool_call_metadata")) or ( + payload.get("call_type") == "call_mcp_tool" + ) + + # --- service event_metadata sanitization ------------------------------------ # # Substrings (case-insensitive) of keys that must never reach a span: secrets, diff --git a/litellm/integrations/otel/model/semconv.py b/litellm/integrations/otel/model/semconv.py index 1c6c30eda0d..7df07f30a01 100644 --- a/litellm/integrations/otel/model/semconv.py +++ b/litellm/integrations/otel/model/semconv.py @@ -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, } diff --git a/litellm/integrations/otel/model/spans.py b/litellm/integrations/otel/model/spans.py index e4876f4ee58..1adc1d68dde 100644 --- a/litellm/integrations/otel/model/spans.py +++ b/litellm/integrations/otel/model/spans.py @@ -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() diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index c17ec13d3ef..dd36e32291d 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -145,6 +145,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 @@ -2324,6 +2337,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( @@ -2695,8 +2709,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 +2720,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( diff --git a/litellm/types/utils.py b/litellm/types/utils.py index d208eb29ac1..d3c2c8c18fe 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -2580,6 +2580,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): """ diff --git a/tests/test_litellm/integrations/otel/test_otel_v2_logger.py b/tests/test_litellm/integrations/otel/test_otel_v2_logger.py index 7fb0e10a247..971bb9340b2 100644 --- a/tests/test_litellm/integrations/otel/test_otel_v2_logger.py +++ b/tests/test_litellm/integrations/otel/test_otel_v2_logger.py @@ -238,6 +238,128 @@ def test_idempotent_on_repeat_callback(): assert len(exporter.get_finished_spans()) == 1 +# --------------------------------------------------------------------------- # +# MCP tool-call spans +# --------------------------------------------------------------------------- # + + +def _mcp_payload(**overrides): + payload = { + "call_type": "call_mcp_tool", + "status": "success", + "litellm_call_id": "mcp_1", + "response_cost": 0.01, + "metadata": {"user_api_key_team_id": "t1"}, + "hidden_params": {}, + "mcp_tool_call_metadata": { + "name": "get_weather", + "arguments": {"city": "Paris"}, + "result": {"temp_c": 21}, + "mcp_server_name": "weather-mcp", + "mcp_session_id": "sess-abc123", + }, + } + payload.update(overrides) + return payload + + +def _logger_capturing(): + from litellm.integrations.otel.model.config import CaptureMessageContent + + cfg = OpenTelemetryV2Config( + exporter="in_memory", + legacy_compat=False, + capture_message_content=CaptureMessageContent.SPAN_ONLY, + ) + exporter = InMemorySpanExporter() + tracer_provider = providers.build_tracer_provider(cfg, exporter=exporter) + return OpenTelemetryV2(config=cfg, tracer_provider=tracer_provider), exporter + + +def test_mcp_tool_call_emits_client_span(): + """A closed MCP tool call becomes a CLIENT span named ``tools/call {tool}``, + carrying the MCP semconv method/operation and the vendor server name.""" + logger, exporter = _logger() + kwargs = {"standard_logging_object": _mcp_payload()} + asyncio.run(logger.async_log_success_event(kwargs, None, None, None)) + (span,) = exporter.get_finished_spans() + assert span.name == "tools/call get_weather" + assert span.kind is SpanKind.CLIENT + assert span.attributes["mcp.method.name"] == "tools/call" + assert span.attributes["mcp.session.id"] == "sess-abc123" + assert span.attributes[GenAI.OPERATION_NAME] == "execute_tool" + assert span.attributes["gen_ai.tool.name"] == "get_weather" + assert span.attributes[LiteLLM.MCP_SERVER_NAME] == "weather-mcp" + assert span.attributes[LiteLLM.CALL_ID] == "mcp_1" + assert span.status.status_code is StatusCode.UNSET + # Tool I/O is content: withheld while capture is off (the default). + assert "gen_ai.tool.call.arguments" not in span.attributes + assert "gen_ai.tool.call.result" not in span.attributes + + +def test_mcp_tool_call_stateless_omits_session_id(): + """A stateless MCP call carries no ``mcp-session-id``, so the span must omit + ``mcp.session.id`` rather than stamping an empty or ``None`` value.""" + logger, exporter = _logger() + payload = _mcp_payload() + del payload["mcp_tool_call_metadata"]["mcp_session_id"] + asyncio.run( + logger.async_log_success_event( + {"standard_logging_object": payload}, None, None, None + ) + ) + (span,) = exporter.get_finished_spans() + assert "mcp.session.id" not in span.attributes + assert span.attributes["mcp.method.name"] == "tools/call" + + +def test_mcp_tool_call_is_not_logged_as_llm_call(): + """The MCP branch must short-circuit the LLM-call path: even if ``pre_call`` + opened a stray carrier for this id, the result is one MCP span, never an LLM + ``chat`` span.""" + logger, exporter = _logger() + kwargs = {"standard_logging_object": _mcp_payload()} + logger.log_pre_api_call(model="MCP: get_weather", messages=[], kwargs=kwargs) + asyncio.run(logger.async_log_success_event(kwargs, None, None, None)) + (span,) = exporter.get_finished_spans() + assert span.attributes["mcp.method.name"] == "tools/call" + assert "gen_ai.request.model" not in span.attributes + + +def test_mcp_tool_call_captures_io_when_enabled(): + logger, exporter = _logger_capturing() + kwargs = {"standard_logging_object": _mcp_payload()} + asyncio.run(logger.async_log_success_event(kwargs, None, None, None)) + (span,) = exporter.get_finished_spans() + assert '"Paris"' in span.attributes["gen_ai.tool.call.arguments"] + assert "21" in span.attributes["gen_ai.tool.call.result"] + + +def test_mcp_tool_call_failure_marks_error(): + logger, exporter = _logger() + payload = _mcp_payload( + status="failure", + error_information={"error_class": "MCPError", "error_message": "upstream 500"}, + ) + asyncio.run( + logger.async_log_failure_event( + {"standard_logging_object": payload}, None, None, None + ) + ) + (span,) = exporter.get_finished_spans() + assert span.name == "tools/call get_weather" + assert span.status.status_code is StatusCode.ERROR + assert span.attributes["error.type"] == "MCPError" + + +def test_mcp_tool_call_deduped_on_repeat(): + logger, exporter = _logger() + kwargs = {"standard_logging_object": _mcp_payload()} + asyncio.run(logger.async_log_success_event(kwargs, None, None, None)) + asyncio.run(logger.async_log_success_event(kwargs, None, None, None)) + assert len(exporter.get_finished_spans()) == 1 + + def test_pre_call_idempotent_keeps_first_span(): """A retried call may re-enter ``pre_call`` with the same call id; the first span (with the true start time) is kept, not replaced.""" diff --git a/tests/test_litellm/integrations/otel/test_otel_v2_sources_of_truth.py b/tests/test_litellm/integrations/otel/test_otel_v2_sources_of_truth.py index c4e80145c70..ecd39c1c936 100644 --- a/tests/test_litellm/integrations/otel/test_otel_v2_sources_of_truth.py +++ b/tests/test_litellm/integrations/otel/test_otel_v2_sources_of_truth.py @@ -1,8 +1,6 @@ """Tests for the OTel v2 sources of truth: span registry, semconv keys, config, and the typed StandardLoggingPayload adapter. These need no OTel SDK.""" -import pytest - from litellm.integrations.otel import ( BAGGAGE_PROMOTED_KEYS, DB, @@ -92,11 +90,14 @@ def test_registry_hierarchy_shape(): # guardrail runs before the LLM call exists, so it's a sibling of it. assert set(child_roles(SpanRole.PROXY_REQUEST)) == { SpanRole.LLM_CALL, + SpanRole.MCP_TOOL_CALL, SpanRole.GUARDRAIL, SpanRole.DB_CALL, SpanRole.SERVICE, } assert SPAN_REGISTRY[SpanRole.LLM_CALL].kind is LiteLLMSpanKind.CLIENT + # The proxy is an MCP client to the upstream tool server: CLIENT span. + assert SPAN_REGISTRY[SpanRole.MCP_TOOL_CALL].kind is LiteLLMSpanKind.CLIENT assert SPAN_REGISTRY[SpanRole.PROXY_REQUEST].kind is LiteLLMSpanKind.SERVER assert SPAN_REGISTRY[SpanRole.GUARDRAIL].parent is SpanRole.PROXY_REQUEST # An outbound datastore call is a CLIENT span; an internal service is INTERNAL. @@ -121,14 +122,52 @@ def _all_constants(cls): def test_attribute_keys_are_unique_across_namespaces(): + from litellm.integrations.otel import MCP, Client, JsonRpc, Network + # prefixes are allowed to be substrings; exact keys must not collide. exact = set() - for cls in (GenAI, Error, Server, HTTP, DB): + for cls in (GenAI, Error, Server, HTTP, DB, MCP, JsonRpc, Network, Client): for key in _all_constants(cls): assert key not in exact, f"duplicate attribute key {key}" exact.add(key) +def test_mcp_attribute_vocabulary_is_complete(): + """Every span-attribute key the OTel GenAI MCP semconv defines has a constant. + + Pins the vocabulary so a dropped or renamed key fails here rather than + silently emitting a non-conformant attribute name. + """ + from litellm.integrations.otel import MCP, Client, JsonRpc, Network + + defined = set() + for cls in (GenAI, Error, Server, MCP, JsonRpc, Network, Client): + defined |= _all_constants(cls) + required = { + "mcp.method.name", + "mcp.session.id", + "mcp.protocol.version", + "mcp.resource.uri", + "jsonrpc.request.id", + "jsonrpc.protocol.version", + "rpc.response.status_code", + "gen_ai.operation.name", + "gen_ai.tool.name", + "gen_ai.tool.call.arguments", + "gen_ai.tool.call.result", + "gen_ai.prompt.name", + "error.type", + "server.address", + "server.port", + "client.address", + "client.port", + "network.protocol.name", + "network.protocol.version", + "network.transport", + } + assert required <= defined, f"missing MCP semconv keys: {required - defined}" + + def test_provider_resolution(): assert resolve_provider("openai") == "openai" assert resolve_provider("bedrock") == "aws.bedrock" @@ -143,6 +182,101 @@ def test_operation_resolution(): assert resolve_operation("aembedding") is GenAIOperation.EMBEDDINGS assert resolve_operation("atext_completion") is GenAIOperation.TEXT_COMPLETION assert resolve_operation(None) is GenAIOperation.CHAT + # An MCP tool call is an ``execute_tool`` operation, not a chat completion. + assert resolve_operation("call_mcp_tool") is GenAIOperation.EXECUTE_TOOL + + +# --- MCP tool-call (source of truth #1/#2/#3) ------------------------------- # + + +def _mcp_payload(capture=False, **overrides): + payload = { + "call_type": "call_mcp_tool", + "status": "success", + "litellm_call_id": "mcp_call_1", + "response_cost": 0.01, + "metadata": {"user_api_key_team_id": "t1"}, + "hidden_params": {}, + "mcp_tool_call_metadata": { + "name": "get_weather", + "arguments": {"city": "Paris"}, + "result": {"temp_c": 21}, + "mcp_server_name": "weather-mcp", + "mcp_session_id": "sess-abc123", + }, + } + payload.update(overrides) + return payload + + +def test_mcp_method_values_match_wire_format(): + from litellm.integrations.otel import MCP, MCPMethod + + assert MCPMethod.TOOLS_CALL.value == "tools/call" + assert MCPMethod.TOOLS_LIST.value == "tools/list" + assert MCP.METHOD_NAME == "mcp.method.name" + + +def test_mcp_tool_call_adapter_extracts_fields(): + from litellm.integrations.otel import MCPToolCallSpanData + + data = MCPToolCallSpanData.from_standard_logging_payload(_mcp_payload()) + assert data.operation is GenAIOperation.EXECUTE_TOOL + assert data.method == "tools/call" + assert data.tool_name == "get_weather" + assert data.server_name == "weather-mcp" + assert data.session_id == "sess-abc123" + assert data.response_cost == 0.01 + assert data.identity.call_id == "mcp_call_1" + assert data.identity.team_id == "t1" + assert data.error is None + + +def test_mcp_tool_call_content_gated_off_by_default(): + # Arguments and result are sensitive tool I/O: withheld unless content capture + # is explicitly enabled, exactly like prompt/response bodies. + from litellm.integrations.otel import MCPToolCallSpanData + + off = MCPToolCallSpanData.from_standard_logging_payload(_mcp_payload()) + assert off.arguments_json is None and off.result_json is None + + on = MCPToolCallSpanData.from_standard_logging_payload( + _mcp_payload(), capture_content=True + ) + assert on.arguments_json is not None and '"Paris"' in on.arguments_json + assert on.result_json is not None and "21" in on.result_json + + +def test_mcp_tool_call_failure_path(): + from litellm.integrations.otel import MCPToolCallSpanData + + data = MCPToolCallSpanData.from_standard_logging_payload( + _mcp_payload( + status="failure", + error_information={"error_class": "MCPError", "error_message": "boom"}, + ) + ) + assert data.error is not None + assert data.error.error_type == "MCPError" + assert data.error.message == "boom" + + +def test_is_mcp_tool_call_detection(): + from litellm.integrations.otel import is_mcp_tool_call + + assert is_mcp_tool_call(_mcp_payload()) is True + # call_type alone is enough even before the gateway stamps its metadata. + assert is_mcp_tool_call({"call_type": "call_mcp_tool"}) is True + assert is_mcp_tool_call({"call_type": "acompletion"}) is False + assert is_mcp_tool_call({}) is False + + +def test_mcp_tool_call_span_name(): + from litellm.integrations.otel import MCPToolCallSpanData + from litellm.integrations.otel.model.spans import mcp_tool_call_span_name + + data = MCPToolCallSpanData.from_standard_logging_payload(_mcp_payload()) + assert mcp_tool_call_span_name(data) == "tools/call get_weather" # --- typed adapter (source of truth #3) ------------------------------------- # diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_session_logging.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_session_logging.py new file mode 100644 index 00000000000..790937cc1de --- /dev/null +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_session_logging.py @@ -0,0 +1,19 @@ +"""The MCP ``mcp-session-id`` is captured for tool-call logging so the otel span +can carry ``mcp.session.id``. Guards the header read against casing and absence.""" + +from litellm.proxy._experimental.mcp_server.server import _mcp_session_id_from_headers + + +def test_reads_session_id_case_insensitively(): + # Clients send varied casing (``Mcp-Session-Id``, ``mcp-session-id``); all resolve. + assert _mcp_session_id_from_headers({"mcp-session-id": "s1"}) == "s1" + assert _mcp_session_id_from_headers({"Mcp-Session-Id": "s2"}) == "s2" + assert _mcp_session_id_from_headers({"MCP-SESSION-ID": "s3"}) == "s3" + + +def test_stateless_call_has_no_session_id(): + # No header (stateless request) and an empty value both yield None, not "". + assert _mcp_session_id_from_headers({"authorization": "Bearer x"}) is None + assert _mcp_session_id_from_headers({"mcp-session-id": ""}) is None + assert _mcp_session_id_from_headers(None) is None + assert _mcp_session_id_from_headers({}) is None From ce7b1fd29dba06b3784152fb562975fa84478164 Mon Sep 17 00:00:00 2001 From: Yassin Kortam Date: Tue, 2 Jun 2026 11:46:25 -0700 Subject: [PATCH 02/27] fix(passthrough): emit otel guardrail span when a guardrail blocks (#29470) * fix(passthrough): emit otel guardrail span when a guardrail blocks The otel_v2 logger emits guardrail spans from its post-call hooks by reading standard_logging_guardrail_information off the top-level metadata of the dict handed to those hooks. On passthrough, post-call guardrails run against a throwaway hook_data dict (metadata was already stripped off _parsed_body by _init_kwargs_for_pass_through_endpoint), so a deny that raises a non ModifyResponseException records its logging info on hook_data and then the generic failure handler forwards _parsed_body, which no longer carries it. The span was therefore present on allow but missing on block; the unified path keeps metadata on the same dict it passes to the failure hook, so its span always shows. Carry the guardrail logging entries recorded on hook_data over to the request_data forwarded to post_call_failure_hook so the failure path matches the unified path. Resolves LIT-3510 * test(passthrough): cover guardrail-logging carry helper; simplify helper Address review feedback on the guardrail-block span fix. Simplify _carry_guardrail_logging_info: the realistic failure path always builds fresh metadata on request_data, so the merge-into-existing-list branch was dead code. Use setdefault with a shallow-copied list so the carried entries never share the source hook_data list reference. Drop the module-level sys.modules proxy_server mock from the otel span test; pass_through_endpoints imports proxy_server lazily, so it is unnecessary and avoided the test-isolation risk of registering a mock under that key. Add pure unit tests for _carry_guardrail_logging_info (no otel dependency) that pin its contract: carries entries, copies the list, populates existing metadata without clobbering prior guardrail entries, and no-ops when there is nothing to carry. * test(passthrough): cover deny-path guardrail logging forwarding without otel The otel span regression test skips in coverage jobs that lack the optional opentelemetry package, leaving the failure-handler wiring (capturing hook_data and carrying its guardrail logging info) uncovered. Add an otel-independent regression that drives the real pass_through_request through a post-call deny and asserts post_call_failure_hook receives request_data carrying the standard_logging_guardrail_information. Fails on the pre-fix code. --------- Co-authored-by: Claude --- .../pass_through_endpoints.py | 34 ++++ .../test_carry_guardrail_logging_info.py | 68 +++++++ ...t_passthrough_guardrail_block_otel_span.py | 189 ++++++++++++++++++ .../test_passthrough_post_call_guardrails.py | 48 +++++ 4 files changed, 339 insertions(+) create mode 100644 tests/test_litellm/proxy/pass_through_endpoints/test_carry_guardrail_logging_info.py create mode 100644 tests/test_litellm/proxy/pass_through_endpoints/test_passthrough_guardrail_block_otel_span.py diff --git a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py index 9d68132b37d..985785ad77e 100644 --- a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py @@ -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, diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_carry_guardrail_logging_info.py b/tests/test_litellm/proxy/pass_through_endpoints/test_carry_guardrail_logging_info.py new file mode 100644 index 00000000000..3071812e117 --- /dev/null +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_carry_guardrail_logging_info.py @@ -0,0 +1,68 @@ +"""Unit tests for ``_carry_guardrail_logging_info``. + +This is the helper that lets a passthrough guardrail block still surface its otel +span: it copies ``standard_logging_guardrail_information`` from the post-call +guardrail's (otherwise discarded) ``hook_data`` onto the dict the failure handler +forwards to ``post_call_failure_hook``. No otel dependency here, so these run +everywhere and pin the helper's contract directly. +""" + +from litellm.proxy.pass_through_endpoints.pass_through_endpoints import ( + _carry_guardrail_logging_info, +) + +_ENTRY = {"guardrail_name": "block-demo", "guardrail_status": "guardrail_intervened"} + + +def _source(entries): + return {"metadata": {"standard_logging_guardrail_information": entries}} + + +def test_carries_entries_onto_request_without_metadata(): + request_data: dict = {} + _carry_guardrail_logging_info(request_data, _source([_ENTRY])) + assert request_data["metadata"]["standard_logging_guardrail_information"] == [ + _ENTRY + ] + + +def test_carried_list_is_copied_not_shared(): + source = _source([_ENTRY]) + request_data: dict = {} + _carry_guardrail_logging_info(request_data, source) + carried = request_data["metadata"]["standard_logging_guardrail_information"] + assert carried is not source["metadata"]["standard_logging_guardrail_information"] + carried.append({"guardrail_name": "other"}) + assert source["metadata"]["standard_logging_guardrail_information"] == [_ENTRY] + + +def test_existing_metadata_without_guardrail_key_is_populated(): + request_data: dict = {"metadata": {"user_api_key": "sk-x"}} + _carry_guardrail_logging_info(request_data, _source([_ENTRY])) + assert request_data["metadata"]["user_api_key"] == "sk-x" + assert request_data["metadata"]["standard_logging_guardrail_information"] == [ + _ENTRY + ] + + +def test_existing_guardrail_entries_are_not_clobbered(): + existing = [{"guardrail_name": "already-logged"}] + request_data = {"metadata": {"standard_logging_guardrail_information": existing}} + _carry_guardrail_logging_info(request_data, _source([_ENTRY])) + assert ( + request_data["metadata"]["standard_logging_guardrail_information"] is existing + ) + + +def test_noop_when_guardrail_data_is_none(): + request_data: dict = {} + _carry_guardrail_logging_info(request_data, None) + assert request_data == {} + + +def test_noop_when_no_guardrail_entries(): + request_data: dict = {} + _carry_guardrail_logging_info(request_data, {"metadata": {}}) + _carry_guardrail_logging_info(request_data, _source([])) + _carry_guardrail_logging_info(request_data, {}) + assert request_data == {} diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_passthrough_guardrail_block_otel_span.py b/tests/test_litellm/proxy/pass_through_endpoints/test_passthrough_guardrail_block_otel_span.py new file mode 100644 index 00000000000..9009da88afa --- /dev/null +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_passthrough_guardrail_block_otel_span.py @@ -0,0 +1,189 @@ +"""Regression: a guardrail block on a passthrough endpoint must still emit the +otel guardrail span. + +Before the fix the post-call guardrail recorded its +``standard_logging_guardrail_information`` onto a throwaway ``hook_data`` dict, +which the failure handler discarded. So ``pass_through_request`` forwarded a +``request_data`` without it to ``post_call_failure_hook`` and the otel guardrail +span (emitted from that hook) was present on allow but missing on block. The +unified path always has it. These tests drive the real ``pass_through_request`` +with a real ``ProxyLogging`` + a real otel V2 logger and assert the span is +emitted on both allow and block. +""" + +import json +from contextlib import ExitStack +from unittest.mock import AsyncMock, MagicMock, patch + +import httpx +import pytest +from fastapi import HTTPException + +pytest.importorskip("opentelemetry") + +from opentelemetry.sdk.trace.export.in_memory_span_exporter import ( # noqa: E402 + InMemorySpanExporter, +) + +import litellm # noqa: E402 +from litellm.caching.dual_cache import DualCache # noqa: E402 +from litellm.integrations.custom_guardrail import ( # noqa: E402 + CustomGuardrail, + log_guardrail_information, +) +from litellm.integrations.otel.logger import OpenTelemetryV2 # noqa: E402 +from litellm.integrations.otel.model.config import OpenTelemetryV2Config # noqa: E402 +from litellm.integrations.otel.plumbing import providers # noqa: E402 +from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache # noqa: E402 +from litellm.proxy.utils import ProxyLogging # noqa: E402 +from litellm.types.guardrails import GuardrailEventHooks # noqa: E402 + +_PT_MOD = "litellm.proxy.pass_through_endpoints.pass_through_endpoints" +_COLLECT = ( + "litellm.proxy.pass_through_endpoints.passthrough_guardrails." + "PassthroughGuardrailHandler.collect_guardrails" +) +_GUARDRAIL_SPAN = "execute_guardrail block-demo" +_TRIGGER = "BLOCKME" + +# pass_through_endpoints imports proxy_server lazily (inside the request +# function), so importing this at module scope does not require the real +# proxy_server and does not mutate sys.modules. +from litellm.proxy.pass_through_endpoints.pass_through_endpoints import ( # noqa: E402 + pass_through_request, +) + + +class _BlockOnTextGuardrail(CustomGuardrail): + """Denies (HTTP 400) when the response carries the trigger word; records its + standard guardrail logging info on both allow and block via the decorator.""" + + @log_guardrail_information + async def async_post_call_success_hook(self, data, user_api_key_dict, response): + if _TRIGGER in json.dumps(response): + raise HTTPException( + status_code=400, detail={"error": "blocked by block-demo guardrail"} + ) + return response + + +def _user_api_key_dict(): + d = MagicMock() + d.api_key = "sk-test" + d.user_id = "user-1" + d.team_id = "team-1" + d.org_id = None + d.metadata = {} + d.team_metadata = {} + d.parent_otel_span = None + d.request_route = "/mock/echo" + return d + + +def _mock_request(): + r = MagicMock() + r.method = "POST" + r.query_params = {} + r.url = "http://testserver/mock/echo" + headers = MagicMock() + headers.copy.return_value = {} + r.headers = headers + return r + + +def _httpx_response(text: str) -> httpx.Response: + body = {"candidates": [{"content": {"role": "model", "parts": [{"text": text}]}}]} + return httpx.Response( + status_code=200, + headers={"content-type": "application/json"}, + content=json.dumps(body).encode("utf-8"), + request=httpx.Request("POST", "https://upstream.example/echo"), + ) + + +def _otel_logger_with_exporter(): + cfg = OpenTelemetryV2Config(exporter="in_memory") + exporter = InMemorySpanExporter() + tracer_provider = providers.build_tracer_provider(cfg, exporter=exporter) + return OpenTelemetryV2(config=cfg, tracer_provider=tracer_provider), exporter + + +def _guardrail_span_names(exporter): + return [ + s.name + for s in exporter.get_finished_spans() + if s.name.startswith("execute_guardrail") + ] + + +async def _drive(response_text: str): + """Run the real pass_through_request with the block-demo guardrail + otel V2 + logger registered, returning (status_code, guardrail_span_names).""" + otel, exporter = _otel_logger_with_exporter() + guardrail = _BlockOnTextGuardrail( + guardrail_name="block-demo", event_hook=[GuardrailEventHooks.post_call] + ) + proxy_logging = ProxyLogging(user_api_key_cache=UserApiKeyCache(DualCache())) + + saved_callbacks = list(litellm.callbacks) + litellm.callbacks = [guardrail, otel] + + mock_async_client_obj = MagicMock() + mock_async_client_obj.client = AsyncMock() + mock_pt_logging = MagicMock() + mock_pt_logging.pass_through_async_success_handler = AsyncMock() + + patches = [ + patch( + f"{_PT_MOD}.HttpPassThroughEndpointHelpers.non_streaming_http_request_handler", + new_callable=AsyncMock, + return_value=_httpx_response(response_text), + ), + patch(f"{_PT_MOD}._is_streaming_response", return_value=False), + patch("litellm.proxy.proxy_server.proxy_logging_obj", proxy_logging), + patch("litellm.proxy.proxy_server.llm_router", None), + patch(f"{_PT_MOD}.pass_through_endpoint_logging", mock_pt_logging), + patch(f"{_PT_MOD}.get_async_httpx_client", return_value=mock_async_client_obj), + patch(f"{_PT_MOD}._read_request_body", new_callable=AsyncMock, return_value={}), + patch(f"{_PT_MOD}._safe_get_request_headers", return_value={}), + patch(_COLLECT, return_value=["block-demo"]), + ] + try: + with ExitStack() as stack: + for p in patches: + stack.enter_context(p) + try: + result = await pass_through_request( + request=_mock_request(), + target="https://upstream.example/echo", + custom_headers={"Content-Type": "application/json"}, + user_api_key_dict=_user_api_key_dict(), + stream=False, + ) + # A deny (HTTP 4xx) re-raises as ProxyException; an allow returns + # the upstream Response. + status_code = result.status_code + except Exception as e: + status_code = getattr(e, "code", None) or getattr( + e, "status_code", None + ) + return int(status_code), _guardrail_span_names(exporter) + finally: + litellm.callbacks = saved_callbacks + + +@pytest.mark.asyncio +async def test_guardrail_block_emits_otel_guardrail_span(): + status_code, span_names = await _drive(f"{_TRIGGER} please") + assert status_code == 400 + assert span_names == [_GUARDRAIL_SPAN], ( + "guardrail span must be emitted when a passthrough guardrail blocks, " + f"got spans: {span_names}" + ) + + +@pytest.mark.asyncio +async def test_guardrail_allow_emits_otel_guardrail_span(): + status_code, span_names = await _drive("hello world") + assert status_code == 200 + assert span_names == [_GUARDRAIL_SPAN] diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_passthrough_post_call_guardrails.py b/tests/test_litellm/proxy/pass_through_endpoints/test_passthrough_post_call_guardrails.py index f061434a971..eafe71e1063 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/test_passthrough_post_call_guardrails.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_passthrough_post_call_guardrails.py @@ -12,6 +12,7 @@ from unittest.mock import AsyncMock, MagicMock, patch import httpx import pytest +from fastapi import HTTPException from litellm.integrations.custom_guardrail import ( CustomGuardrail, @@ -216,6 +217,53 @@ class TestPassthroughPostCallGuardrails: assert body["error"]["guardrail_name"] == "rubrik" assert body["error"]["model"] == "gemini-2.0-flash" + @patch(_COLLECT, return_value=["rubrik"]) + async def test_deny_forwards_guardrail_logging_info_to_failure_hook( + self, + mock_collect, + ): + """A post-call guardrail deny (non-ModifyResponseException) records its + standard_logging_guardrail_information on the hook_data dict; the failure + handler must forward that info to post_call_failure_hook so downstream + loggers (e.g. the otel guardrail span) still see it. Regression for the + block path dropping it.""" + mock_response = _make_httpx_response(_GEMINI_RESPONSE) + + def _block(*, data, user_api_key_dict, response): + metadata = data.setdefault("metadata", {}) + metadata.setdefault("standard_logging_guardrail_information", []).append( + {"guardrail_name": "rubrik", "guardrail_status": "guardrail_intervened"} + ) + raise HTTPException(status_code=400, detail={"error": "blocked"}) + + captured = {} + + async def _capture_failure(**kwargs): + captured.update(kwargs) + + mock_proxy_logging = MagicMock() + mock_proxy_logging.pre_call_hook = AsyncMock(return_value={}) + mock_proxy_logging.post_call_success_hook = AsyncMock(side_effect=_block) + mock_proxy_logging.post_call_failure_hook = AsyncMock( + side_effect=_capture_failure + ) + + with _common_patches(mock_proxy_logging, mock_response): + with pytest.raises(Exception): + await pass_through_request( + request=_make_mock_request(), + target="https://example.com/v1/generateContent", + custom_headers={"Content-Type": "application/json"}, + user_api_key_dict=_make_user_api_key_dict(), + stream=False, + ) + + mock_proxy_logging.post_call_failure_hook.assert_awaited_once() + entries = captured["request_data"]["metadata"][ + "standard_logging_guardrail_information" + ] + assert any(e.get("guardrail_name") == "rubrik" for e in entries) + @pytest.mark.asyncio class TestUnifiedGuardrailCallTypeResolution: From efaafbbd025c55657f0783012c827bf7693f00d5 Mon Sep 17 00:00:00 2001 From: milan-berri Date: Tue, 2 Jun 2026 22:07:11 +0300 Subject: [PATCH 03/27] fix(proxy): strip NUL bytes from spend log payloads to prevent PostgreSQL 22P05 (#29515) A raw NUL byte (\x00) in request/response content is serialized by json.dumps into the \u0000 JSON escape. When update_spend_logs writes this to the LiteLLM_SpendLogs jsonb columns, Postgres rejects the whole batch with error 22P05 ("unsupported Unicode escape sequence ... cannot be converted to text"), crashing the periodic update_spend job and dropping the spend-log batch. Centralize stripping in safe_dumps (covers metadata/response paths and any future caller) and route the messages, proxy_server_request, request_tags, and response (string branch) payloads through it instead of json.dumps. Dict keys are stripped too. Adds regression tests for safe_dumps and the spend-log message, response, and request_tags payload builders. Co-authored-by: Cursor --- litellm/litellm_core_utils/safe_json_dumps.py | 14 +++- .../spend_tracking/spend_tracking_utils.py | 12 ++-- .../test_safe_json_dumps.py | 36 ++++++++++- .../test_spend_tracking_utils.py | 64 +++++++++++++++++++ 4 files changed, 116 insertions(+), 10 deletions(-) diff --git a/litellm/litellm_core_utils/safe_json_dumps.py b/litellm/litellm_core_utils/safe_json_dumps.py index 051aa2f27a5..154306d01b8 100644 --- a/litellm/litellm_core_utils/safe_json_dumps.py +++ b/litellm/litellm_core_utils/safe_json_dumps.py @@ -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" diff --git a/litellm/proxy/spend_tracking/spend_tracking_utils.py b/litellm/proxy/spend_tracking/spend_tracking_utils.py index e2881faca0d..d215294fd04 100644 --- a/litellm/proxy/spend_tracking/spend_tracking_utils.py +++ b/litellm/proxy/spend_tracking/spend_tracking_utils.py @@ -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: diff --git a/tests/test_litellm/litellm_core_utils/test_safe_json_dumps.py b/tests/test_litellm/litellm_core_utils/test_safe_json_dumps.py index c71a229cca5..74574370e46 100644 --- a/tests/test_litellm/litellm_core_utils/test_safe_json_dumps.py +++ b/tests/test_litellm/litellm_core_utils/test_safe_json_dumps.py @@ -8,7 +8,7 @@ sys.path.insert( 0, os.path.abspath("../../..") ) # Adds the parent directory to the system path -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 def test_primitive_types(): @@ -140,6 +140,40 @@ def test_non_standard_dict_keys_complex(): raise e +def test_strip_null_bytes_helper(): + assert strip_null_bytes("hello\x00world") == "helloworld" + assert strip_null_bytes("\x00\x00abc\x00") == "abc" + assert strip_null_bytes("no null here") == "no null here" + + +def test_null_byte_stripped_from_string(): + out = safe_dumps("hello\x00world") + assert "\\u0000" not in out + assert json.loads(out) == "helloworld" + + +def test_null_byte_stripped_in_nested_structure(): + data = { + "messages": [{"role": "user", "content": "bad\x00content"}], + "nested": {"k\x00ey": "v\x00alue"}, + } + out = safe_dumps(data) + assert "\\u0000" not in out + result = json.loads(out) + assert result["messages"][0]["content"] == "badcontent" + assert result["nested"] == {"key": "value"} + + +def test_null_byte_stripped_in_fallback_str(): + class WithNullStr: + def __str__(self): + return "obj\x00repr" + + out = safe_dumps({"obj": WithNullStr()}) + assert "\\u0000" not in out + assert json.loads(out)["obj"] == "objrepr" + + def test_pydantic_base_model(): from pydantic import BaseModel diff --git a/tests/test_litellm/proxy/spend_tracking/test_spend_tracking_utils.py b/tests/test_litellm/proxy/spend_tracking/test_spend_tracking_utils.py index 5ca058fc8d9..0c7511589de 100644 --- a/tests/test_litellm/proxy/spend_tracking/test_spend_tracking_utils.py +++ b/tests/test_litellm/proxy/spend_tracking/test_spend_tracking_utils.py @@ -300,6 +300,25 @@ def test_get_messages_for_spend_logs_realtime_returns_messages(mock_should_store assert parsed[1]["content"] == "What is the weather today?" +@patch( + "litellm.proxy.spend_tracking.spend_tracking_utils._should_store_prompts_and_responses_in_spend_logs" +) +def test_get_messages_for_spend_logs_strips_null_bytes(mock_should_store): + """Regression for PostgreSQL 22P05: NUL bytes must be stripped from messages.""" + mock_should_store.return_value = True + payload = cast( + StandardLoggingPayload, + { + "call_type": "_arealtime", + "messages": [{"role": "user", "content": "hello\x00world"}], + }, + ) + result = _get_messages_for_spend_logs_payload(payload) + assert "\\u0000" not in result + parsed = json.loads(result) + assert parsed[0]["content"] == "helloworld" + + @patch( "litellm.proxy.spend_tracking.spend_tracking_utils._should_store_prompts_and_responses_in_spend_logs" ) @@ -370,6 +389,21 @@ def test_get_response_for_spend_logs_payload_truncates_large_base64(mock_should_ assert parsed["data"][0]["other_field"] == "value" +@patch( + "litellm.proxy.spend_tracking.spend_tracking_utils._should_store_prompts_and_responses_in_spend_logs" +) +def test_get_response_for_spend_logs_payload_strips_null_bytes(mock_should_store): + """Regression for PostgreSQL 22P05: NUL bytes must be stripped from response.""" + mock_should_store.return_value = True + payload = cast( + StandardLoggingPayload, + {"response": {"content": "answer\x00here"}}, + ) + response_json = _get_response_for_spend_logs_payload(payload) + assert "\\u0000" not in response_json + assert json.loads(response_json)["content"] == "answerhere" + + @patch( "litellm.proxy.spend_tracking.spend_tracking_utils._should_store_prompts_and_responses_in_spend_logs" ) @@ -936,6 +970,36 @@ def test_get_logging_payload_includes_overhead_in_spend_logs_metadata(): ), f"Expected overhead '{test_overhead_ms}', got '{metadata.get('litellm_overhead_time_ms')}'" +@patch("litellm.proxy.proxy_server.master_key", None) +@patch("litellm.proxy.proxy_server.general_settings", {}) +def test_get_logging_payload_strips_null_bytes_from_request_tags(): + """Regression for PostgreSQL 22P05: NUL bytes must be stripped from request_tags.""" + kwargs = { + "model": "gpt-3.5-turbo", + "litellm_params": { + "metadata": { + "user_api_key": "sk-test-key", + "tags": ["clean-tag", "bad\x00tag"], + } + }, + } + + start_time = datetime.datetime.now(timezone.utc) + end_time = datetime.datetime.now(timezone.utc) + + payload = get_logging_payload( + kwargs=kwargs, + response_obj={}, + start_time=start_time, + end_time=end_time, + ) + + request_tags = payload.get("request_tags") + assert request_tags is not None + assert "\\u0000" not in request_tags + assert json.loads(request_tags) == ["clean-tag", "badtag"] + + @patch("litellm.proxy.proxy_server.master_key", None) @patch("litellm.proxy.proxy_server.general_settings", {}) def test_get_logging_payload_handles_missing_overhead_gracefully(): From 6d6eda8101478c59bd92183aaf1db174cc336769 Mon Sep 17 00:00:00 2001 From: Mateo Wang <277851410+mateo-berri@users.noreply.github.com> Date: Tue, 2 Jun 2026 12:22:04 -0700 Subject: [PATCH 04/27] [internal copy of #28008] Support MCP OAuth passthrough and issuer-scoped JWT auth (#28356) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit * fix(proxy): point /metrics 401 at the opt-out flag Operators upgrading past 35bbca60b0 (which made /metrics auth default-on) see "Malformed API Key passed in. Ensure Key has 'Bearer ' prefix." with no hint that litellm_settings.require_auth_for_metrics_endpoint: false restores the previous unauthenticated behavior. Append that discovery hint to the existing 401 body so a Prometheus scraper that breaks after upgrade has a clear migration path. No behavior change. * fix(proxy): bound budget reservation per request instead of pinning to remaining headroom reserve_budget_for_request fell back to reserving the entire remaining team/key/user headroom whenever a request omitted max_tokens, which pinned the spend counter at max_budget for the duration of the in-flight request and false-positive-blocked every concurrent or back-to-back request until the success callback reconciled. Surfaced as an integration-test team being budget-blocked at its $2000 cap while DB spend was $0.144. Switch the missing-max_tokens path to a fixed default of 16384 output tokens (mirrors parallel_request_limiter_v3's DEFAULT_MAX_TOKENS_ESTIMATE precedent), and clamp explicit max_tokens at the model's max_output_tokens for reservation accounting only. The outbound request body is unchanged, so providers see whatever the caller actually sent; only the local integer used to compute reservation cost is bounded. This also prevents a hostile max_tokens=999999999 from inflating one request's reservation up to the entire team headroom. For Opus 4.7 (output $25/M, max_output 128K) on a $2000 budget the worst-case per-request reservation drops from "everything left" to $3.20, raising admittable concurrency from 1 to ~625. * fix(proxy): reserve per-image cost for image-generation requests Image-generation routes (dall-e-3, flux, etc.) have no per-token output cost so they fell through to the no-reservation read-time-only path. Concurrent image requests against a depleted budget could all pass common_checks (counter exactly at max_budget passes the strict-`>` gate) and reach the provider before reconciliation caught up. Add per-image reservation in _estimate_request_max_cost_for_model: when the model has a per-image cost field, reserve `n × cost_per_image` upfront. The atomic counter increment serializes concurrent admissions, so the second request sees the post-first-reservation counter and raises BudgetExceededError instead of silently leaking through. Both `output_cost_per_image` and `input_cost_per_image` are honored — naming is inconsistent across providers (OpenAI dall-e-3 uses input_cost_per_image, aiml/dall-e-3 uses output_cost_per_image for the same per-generated-image price). Per-pixel pricing (DALL-E 2 size variants) and TTS/STT routes still fall through to read-time enforcement; those are follow-ups. * fix(proxy): gate image-gen reservation strictly on model mode The previous detection treated any model with input_cost_per_image or output_cost_per_image as image generation. Several chat and embedding models carry those fields to price multimodal vision input, not generated images: - gemini-3.1-pro-preview (mode=chat) has output_cost_per_image=0.00012 alongside input/output token pricing. - azure/gpt-realtime-* (mode=chat) has input_cost_per_image=5e-6. - amazon.titan-embed-image-v1 (mode=embedding) has input_cost_per_image=6e-5. For these models the image-gen branch fired first and reserved a fraction of a cent per request, short-circuiting the token-priced path entirely. Long Gemini chats reserved 1 × $0.00012 instead of the true token cost. Gate strictly on mode in {"image_generation", "image_edit"}. All 197 real image_generation entries and all 31 image_edit entries (Flux Kontext, Stability inpaint/outpaint, etc.) carry the right mode, so the field-presence fallback was unnecessary. Adds regression tests for the chat-model-with-image-cost-field case and for image_edit reservation. * build(packaging): relax core runtime pins to ranges Backport of #27241 onto litellm_1.84.0rc2. The 12 entries in `[project.dependencies]` were exact `==` pins, a side effect of the Poetry -> uv migration. This forces every downstream package that lists litellm as a dependency to downgrade common runtime libraries (openai, pydantic, aiohttp, click, jsonschema, ...) to the exact versions we ship. Switch to lower-bounded ranges with upper bounds where the upstream package is pre-1.0 or has a known breaking-major-version policy. Reproducibility for our Docker proxy and CI continues to come from `uv.lock`, which is regenerated here as a metadata-only diff. Conflict resolution vs upstream merge: - The upstream merge commit also surfaced unrelated context entries (nvidia-riva-client, soundfile/stt-nvidia-riva extra) that exist in staging but not in rc2. Those are not part of #27241's intent and were dropped from the resolution; the rc2 uv.lock keeps its existing entry set, only the 12 specifier strings changed. - `uv lock --check` passes (392 packages resolved, no drift). * build(packaging): raise jinja2 floor to 3.1.6 Our `uv.lock` already resolves jinja2 to 3.1.6, so Docker / CI installs get that version. The `pyproject.toml` floor was lagging at 3.1.0, which means downstream consumers using `--resolution=lowest-direct` or older constraint files can land on 3.1.0-3.1.5 instead of the version we actually test against. Aligns the declared floor with the resolved version so external installers see the same baseline our test matrix exercises. `uv lock` diff is metadata-only (no resolved-version drift). * fix(mcp): forward extra_headers for OpenAPI MCP tools OpenAPI-generated tools only applied static closure headers and BYOK Authorization via ContextVar. Copy MCPServer.extra_headers from the incoming MCP request into _request_extra_headers (set in server.py before local tool dispatch), merge in openapi_to_mcp_generator via a small helper. OAuth2 M2M: do not forward caller Authorization from raw_headers (same rule as _prepare_mcp_server_headers for managed MCP). Adds TestRequestExtraHeaders and clarifies mcp_server_manager registration comment. Fixes #26794 Co-authored-by: Cursor * refactor(mcp): access has_client_credentials on MCPServer directly Greptile: getattr default was redundant; property exists on MCPServer and mcp_server is non-None inside the extra_headers forwarding block. Co-authored-by: Cursor * fix(mcp): static headers win over forwarded headers in OpenAPI MCP Match the existing MCP invariant in merge_mcp_headers and the managed MCP path: operator-configured static headers always override caller-forwarded headers on name conflict, with case-insensitive comparison so different casing cannot bypass the precedence. _request_auth_header (BYOK) still overrides Authorization last. Addresses Veria review on PR #27383. Co-authored-by: Mateo Wang * fix(proxy): always merge caller-supplied tags into request metadata Caller-supplied tags (`x-litellm-tags` header, body `tags`, `metadata.tags`) were silently dropped unless the key/team had `metadata.allow_client_tags: true` set. Restore the documented behavior: tags from the request always flow into `metadata.tags` and union with any admin-configured static tags from key/team/project metadata. Removes the `allow_client_tags` opt-in flag from the pre-call pipeline. The flag was only ever read here; it has no schema or endpoint footprint, so leftover values in existing key metadata are inert. Test cleanup mirrors the simplification: drop the three tests that verified the strip-when-not-opted-in path, drop the `allow_client_tags` fixture lines from the merge/union tests. * docs(proxy): refresh stale comments referencing removed tag strip The tag-strip block was removed in the parent commit but two surrounding comments still referenced "tags without opt-in" and "runs AFTER the strip". Update them to describe the remaining user_api_key_* and _pipeline_managed_guardrails strip that the snapshot/merge ordering actually protects against. * chore: reject bare str at file-input sinks to prevent local-file read (#27762) Cherry-pick of #27762 onto litellm_1.84.0rc2. * chore: reject bare str at file-input sinks to prevent local-file read (#27667) * fix: use os.PathLike in ocr sink and check truthy reasoningSummary for bridge - ocr/main.py: widen Path check to os.PathLike for consistency with other sinks - main.py: bridge condition checks truthiness of reasoning_summary, not just None * fix: remove unused pathlib.Path import in ocr/main.py Co-authored-by: yuneng-jiang Co-authored-by: ryan-crabbe-berri Co-authored-by: stuxf <70670632+stuxf@users.noreply.github.com> * Strip SERVER_ROOT_PATH before lazy-feature prefix match LazyFeatureMiddleware compared the raw scope path against registered prefixes (e.g. /policies), so requests under a server root path like /api/v1/policies/... never matched, the feature never loaded, and the endpoint returned 404. Strip the configured root path before matching, normalizing trailing slashes and enforcing a component boundary so /api does not falsely match /apiv2. * Cache normalized SERVER_ROOT_PATH at middleware init SERVER_ROOT_PATH is a process-startup env var. Read it once in __init__ instead of calling get_server_root_path() + rstrip on every request that arrives before all lazy features have loaded. * chore(proxy): backport /key/regenerate ownership-rebind + premium-gate guards (#27793) Backport of #27793 onto litellm_1.84.0rc2. A non-admin caller could rebind their own key's user_id via /key/regenerate. _execute_virtual_key_regeneration had org/team guards but no user_id guard, and prepare_key_update_data did not strip the field — it survived model_dump(exclude_unset=True) into the Prisma update. On the next request, _return_user_api_key_auth_obj resolved the rebound user_id against litellm_usertable and returned PROXY_ADMIN whenever the target row's user_role was admin. /key/update had the equivalent guard inline at _validate_update_key_data; extract it to a shared helper _validate_caller_can_change_key_ownership and call from both /key/update and _execute_virtual_key_regeneration. Also tighten the premium gate that allowed the master-key rotation branch to skip the enterprise check. The previous predicate was a field-presence test, not an identity check. Verify the caller actually holds the master key via _is_master_key before allowing the non-premium path. Block explicit-null user_id and empty-string user_id as removal attempts; both 403-reject for non-admin callers. * fix(proxy): expose db status on public /health/readiness Backport of #27866 onto litellm_1.84.0rc2. External readiness probes consumed the legacy detailed payload's `db` field to drive alerting and pod-rotation decisions. Stripping the body to {"status": "healthy"} broke those probes silently — the HTTP code still flipped to 503, but probes checking body.db == "connected" treated the response as healthy. Add `db` back to the unauthenticated payload. The rest of the diagnostic fields (litellm_version, callbacks, cache, log_level) stay behind /health/readiness/details so the recon-leak gate from #26912 holds. Values match the legacy contract: "connected", "disconnected", "Not connected". The 503-on-DB-disconnect behavior from LIT-2607 is preserved. * fix(ui): fetch version + debug flag from /health/readiness/details The proxy moved `litellm_version`, `is_detailed_debug`, and other diagnostic fields off the public `/health/readiness` payload behind an auth-gated `/health/readiness/details` endpoint. The navbar version tag and the detailed-debug-mode banner stopped working because they were still reading those fields from the unauthed response, which no longer contains them. Replace `useHealthReadiness` with a `useHealthReadinessDetails` hook that takes an `accessToken` argument and sends a Bearer header to the auth-gated endpoint. The hook stays disabled while `accessToken` is falsy, so the navbar can keep rendering on the public model hub (where the token is null) without triggering an auth redirect or a 401-loop. * fix(ui): disable retries on readiness/details + cover token forwarding Two small follow-ups on the readiness/details migration: - Set `retry: false` on the query. The payload feeds a passive navbar tag and a debug banner; a 401 from an expired token shouldn't fan out into three retries against the proxy. - Add navbar specs that assert the `accessToken` prop is forwarded into the hook (matches the DebugWarningBanner spec). Without this, the navbar could silently regress to passing `undefined` and the existing tests wouldn't catch it. * chore: update Next.js build artifacts (2026-05-14 03:52 UTC, node v20.20.2) * Merge pull request #27898 from stuxf/chore/banned-params-extra-body-cover chore(proxy): cover extra_body + azure_ad_token in banned-params check (cherry picked from commit a6a9d8edf024a7d808ba18df4aace4815e5f5925) * Merge pull request #27801 from stuxf/chore/get-instance-fn-runtime-s3-gate chore(proxy): refuse remote-URL instance-fn loads outside config-file path (cherry picked from commit e3e5209f51a605d49f4c1ef9b010ed5fdd1812c6) * fix: block client-side pricing injection via request body Authenticated clients could supply CustomPricingLiteLLMParams fields (input_cost_per_token, output_cost_per_token, etc.) in the request body. These were forwarded to register_model() in main.py, permanently mutating the shared global litellm.model_cost dict for all users on the instance. Adds all CustomPricingLiteLLMParams fields to _BANNED_REQUEST_BODY_PARAMS so is_request_body_safe() rejects them before they reach completion(). New pricing fields added to CustomPricingLiteLLMParams are auto-covered. Admin opt-in via allow_client_side_credentials or configurable_clientside_auth_params still works as before. Co-Authored-By: Claude Sonnet 4.6 (1M context) * fix: block SSRF fields in RAG ingest vector_store config aws_sts_endpoint, aws_web_identity_token, and aws_bedrock_runtime_endpoint in ingest_options.vector_store were passed directly to the Bedrock ingestion class, which reads them into boto3 STS client construction. Any authenticated caller could redirect AssumeRole calls to an attacker-controlled server, leaking the proxy's instance profile credentials. Calls is_request_body_safe() on ingest_options["vector_store"] before forwarding to litellm.aingest(). Same banned-params list and admin opt-in escape hatch (allow_client_side_credentials) as the /chat/completions path. ValueError from the safety check is caught and re-raised as HTTP 400. Co-Authored-By: Claude Sonnet 4.6 (1M context) * fix: harden /key/update authorization checks (#27878) * fix: patch Host-header auth bypass in get_request_route Starlette reconstructs request.url from the Host header. A malformed Host like `localhost/?x=1` causes Starlette to build the full URL as `http://localhost/?x=1/health`, which url-parses to path="/". Since "/" is in LiteLLMRoutes.public_routes, all protected routes became reachable without authentication. Fix: read scope["path"] (set by uvicorn from the HTTP request line, not derivable from headers) instead of request.url.path. Sub-path deployments are handled via scope["app_root_path"] / scope["root_path"], mirroring Starlette's own base_url construction logic. Affected variants confirmed fixed: Host: localhost/?x=1 Host: localhost:4000/?x=1 Host: localhost/#test Host: localhost:4000/#test Co-Authored-By: Claude Sonnet 4.6 (1M context) * style: reduce comments in route fix Co-Authored-By: Claude Sonnet 4.6 (1M context) * fix: block credential fields in RAG ingest vector_store options Credential fields (vertex_credentials, aws_access_key_id, api_key, etc.) in ingest_options.vector_store are now rejected at the API boundary with a 400 error. Credentials must be configured server-side. Previously any authenticated user could supply a vertex_credentials dict with type=external_account pointing credential_source.file at an arbitrary path (e.g. /proc/1/environ) and token_url at an attacker-controlled server. google-auth's identity_pool.Credentials refresh() would read the file and POST its contents to the attacker. Co-Authored-By: Claude Sonnet 4.6 (1M context) * fix: block /key/update self-escalation by assigned users Non-admin users who were assigned a key (created_by != caller) could update any non-budget field — models, rpm_limit, guardrails, etc. — without admin authorization, allowing privilege self-escalation. Gate: only the key creator (created_by == caller) may edit their own key without admin check; budget changes always require admin regardless of creator status. All other callers must pass _check_key_admin_access. Co-Authored-By: Claude Sonnet 4.6 (1M context) * fix: block user-controlled api_base in RAG ingest vector_store options A user-supplied api_base in ingest_options.vector_store caused the server to forward its configured provider credentials (Gemini, OpenAI) to an attacker-controlled endpoint via SSRF. Add api_base to the blocked credential params set alongside api_key and the existing credential fields. Co-Authored-By: Claude Sonnet 4.6 (1M context) * fix: restrict /utils/transform_request to PROXY_ADMIN and apply body safety check Any authenticated internal_user could POST arbitrary provider config (aws_sts_endpoint, api_base, etc.) to /utils/transform_request and have the server forward its credentials to an attacker-controlled endpoint. - Gate the endpoint on PROXY_ADMIN role (403 for all other roles) - Call is_request_body_safe() to reject banned params even for admins - Convert ValueError from safety check to HTTP 400 Co-Authored-By: Claude Sonnet 4.6 (1M context) * fix: apply banned-param check to /utils/transform_request Without is_request_body_safe(), any authenticated user could pass aws_sts_endpoint, api_base, or aws_web_identity_token to /utils/transform_request and have the server forward its configured provider credentials to an attacker-controlled endpoint during SDK credential resolution. Applies the same banned-param blocklist already used by LLM endpoints. Endpoint remains accessible to all authenticated users. Co-Authored-By: Claude Sonnet 4.6 (1M context) * fix: block SSRF via api_base in /prompts/test dotprompt YAML frontmatter Any frontmatter key not in ["model","input","output"] flowed into optional_params and was merged into the LLM call data dict, bypassing is_request_body_safe. An attacker with any bearer key could set api_base in YAML to redirect the outbound LLM request — including the provider API key — to an attacker-controlled host. Fix: call is_request_body_safe on the constructed data dict after optional_params are merged, before invoking ProxyBaseLLMRequestProcessing. ValueError from the banned-param check is surfaced as HTTP 400. Co-Authored-By: Claude Sonnet 4.6 (1M context) * Update litellm/proxy/rag_endpoints/endpoints.py Co-authored-by: veria-ai[bot] <224490171+veria-ai[bot]@users.noreply.github.com> * fix: coerce nested config strings before banned-param check _NESTED_CONFIG_KEYS descent used isinstance(nested, dict) which silently skipped litellm_embedding_config when delivered as a JSON string via multipart/form-data. Banned params (api_base, aws_sts_endpoint, etc.) nested inside the stringified value were invisible to is_request_body_safe. _NESTED_METADATA_KEYS already used _coerce_metadata_to_dict which parses JSON strings before checking. Apply the same coercion to _NESTED_CONFIG_KEYS. Co-Authored-By: Claude Sonnet 4.6 (1M context) * fix: replace substring match with prefix match in is_llm_api_route mapped_pass_through_routes used `_llm_passthrough_route in route` (substring) so any admin-only path whose URL contained a provider name (openai, anthropic, azure, bedrock, etc.) was misclassified as an LLM API route and bypassed the admin gate in non_proxy_admin_allowed_routes_check. Confirmed live: non-admin key could GET /credentials/by_name/openai (read masked provider API key) and DELETE /credentials/openai (delete credential). Fix: use exact match or startswith(prefix + "/") — the same pattern used everywhere else in RouteChecks — so only routes that actually start with a passthrough prefix are allowed through. Co-Authored-By: Claude Sonnet 4.6 (1M context) * fix: stabilize PR #27878 test failures - key_management_endpoints: extend can_skip_admin_check to team keys so team members with /key/update permission can update non-budget fields. can_team_member_execute_key_management_endpoint already validates team membership + permission and raises if unauthorized; reaching the admin check on a team key means the caller was authorized. - test: set created_by on mock key in test_update_key_non_budget_fields_allowed_for_internal_user so caller_is_creator resolves correctly (MagicMock default ≠ user_id). - auth_utils.get_request_route: guard against non-dict request.scope (e.g. MagicMock in unit tests) to prevent a MagicMock leaking into UserAPIKeyAuth.request_route and failing Pydantic validation. - ci: assign test_multipart_bypass_repro.py to the proxy-runtime shard in test-unit-proxy-db.yml to satisfy the shard-coverage check. Co-Authored-By: Claude Sonnet 4.6 (1M context) * fix(lint): add explicit str() cast in get_request_route for MyPy scope.get() returns Any|None which MyPy cannot coerce to str implicitly. Wrap both scope.get() calls in str() to satisfy the type checker. Co-Authored-By: Claude Sonnet 4.6 (1M context) * fix: guard bare-/ root_path strip + make total_spend migration idempotent auth_utils.get_request_route: when Starlette sets scope["app_root_path"] to "/" (e.g. behind some middleware), the old stripping logic would remove the leading slash from every path ("/team/new" → "team/new"), breaking route matching and causing auth to misclassify protected routes. Skip stripping when root_path is bare "/". migration: add IF NOT EXISTS to total_spend ALTER TABLE so the migration is safe to replay when a prior partial run already created the column. Without this guard, prisma migrate deploy fails on CI DBs that were partially migrated, causing all subsequent DB operations (including /team/new) to 500. Co-Authored-By: Claude Sonnet 4.6 (1M context) * fix: require creator still owns key for personal-key bypass in /key/update caller_is_creator now requires both created_by == caller AND user_id == caller. Previously checking only created_by let a demoted admin who originally created a key for another user continue editing non-budget fields on it after reassignment, bypassing _check_key_admin_access. Adds regression test: creator whose key was reassigned is blocked (403). Co-Authored-By: Claude Sonnet 4.6 (1M context) * fix: extract auth checks to fix PLR0915 + broaden max_budget assertion internal_user_endpoints._update_single_user_helper exceeded 50 statements (PLR0915). Extract authorization checks into _check_user_update_authz helper to bring statement count under the limit. test_validate_max_budget: assert "negative" (substring of both the local "cannot be negative" and the CI "non-negative finite number" messages) so the test is stable regardless of which exact wording the function uses. Co-Authored-By: Claude Sonnet 4.6 (1M context) --------- Co-authored-by: Claude Sonnet 4.6 (1M context) Co-authored-by: veria-ai[bot] <224490171+veria-ai[bot]@users.noreply.github.com> * bump: version 0.4.71 → 0.4.72 * uv lock * feat(mcp): support OAuth passthrough discovery * fix(mcp): support OAuth browser auth * fix(mcp): refine upstream OAuth metadata fallback * feat(proxy): support issuer-scoped JWT auth * fix(mcp): validate oauth callback redirect sink * feat(proxy): support issuer-scoped JWT auth * test(mcp): align trusted proxy fixtures * style(mcp): satisfy black formatting * chore(ui): bump next to 16.2.6 * fix(mcp): address oauth passthrough review findings * test(mcp): split oauth passthrough regressions * fix(interactions): align openapi response fields * security: prevent forwarding litellm api keys to upstream mcp servers - Strip Authorization header from extra_headers for pass-through servers - Pass-through servers (auth_type=None with extra_headers: [Authorization]) must not receive the user's LiteLLM API key - Only OAuth2 M2M and pass-through servers skip Authorization header - Other headers (x-request-id, x-trace-id) are still forwarded normally - Fixes credential leakage / authentication bypass in MCP pass-through mode * fix(interactions): remove steps field not in google openapi spec The steps field was added but is not present in the current Google Interactions OpenAPI specification. Revert to using only the fields that are actually defined in the spec. * fix(mcp): forward Authorization in pass-through when x-litellm-api-key is admission Commit 3753970cc9 widened the Authorization strip to cover all is_oauth_passthrough servers — protecting against the LiteLLM admission key leaking upstream when the caller used Authorization for admission, but also silently stripping legitimate upstream OAuth bearers when the caller used x-litellm-api-key for admission. That broke transparent OAuth pass-through (EAI-506 V5/V6): standards- compliant MCP clients (OpenCode, Claude Code, mcp-inspector) complete PKCE against the upstream IdP and send the resulting token as plain Authorization: Bearer per the MCP spec — with the wider strip in place, that token never reaches the upstream and tools/list returns empty. Narrow the strip: skip Authorization for pass-through servers only when the caller did NOT supply x-litellm-api-key. When x-litellm-api-key is present, admission is unambiguous and Authorization is free to carry the upstream OAuth bearer. The original security guarantee is preserved — a client that sends only Authorization (no x-litellm-api-key) still has it stripped, so the LiteLLM key cannot leak upstream via that path. Tests: - new: forwards Authorization when x-litellm-api-key is present - new: still strips Authorization when only Authorization is present - existing pass-through + M2M tests unchanged Co-Authored-By: Claude Sonnet 4.6 * fix(interactions): align status enum with openapi spec * fix(mcp,jwt): address greptile review concerns - Cache _get_agent_object_permission via user_api_key_cache (sentinel for no-permission rows) so MCP requests from agent keys don't hit the DB on every tool-list / tool-call. - Re-raise HTTPException in handle_sse_mcp so 401 + WWW-Authenticate challenges (and other HTTP errors) propagate to SSE clients instead of being swallowed as 500. - Normalise booleans in _validate_token_response so admin rules written as JSON-style "true" / "false" match upstream responses that return Python True / False. - Treat configured JWT issuer claim mappings as advisory: when a mapped field is absent or empty, leave the normalised claim unset instead of raising, matching the global litellm_jwtauth path. Co-authored-by: Claude * test: replace dall-e-3 with gpt-image-1 in health check and router tests (#27813) OpenAI returns 'The model dall-e-3 does not exist' for the test account, breaking test_openai_img_gen_health_check and test_image_generation. Switch to gpt-image-1, matching the existing TestOpenAIGPTImage1 pattern. (cherry picked from commit aee58db88057b274eab70388dce72eac31ea014f) * fix(tests): drop dall-e-only test classes; route live image tests via gpt-image-1 Second wave of failures from the 2026-05-12 DALL-E shutdown: - tests/image_gen_tests/test_image_edits.py::TestOpenAIImageEditDallE2 and tests/image_gen_tests/test_image_generation.py::TestOpenAIDalle3 are explicitly named for the deprecated models and can't pass; remove. gpt-image-1 coverage already exists in sibling classes. - tests/local_testing/test_router.py image gen tests use dall-e-3 only as a routing example; swap to gpt-image-1. - tests/local_testing/test_custom_callback_input.py image_generation success/failure paths swapped to gpt-image-1. (cherry picked from commit 945b10ded467e53fc3c9b8df0329dbc55591a56e) * test(fireworks): replace deprecated llama-v3p3-70b-instruct model Fireworks removed llama-v3p3-70b-instruct from serverless, so every live test using it now fails with NotFoundError ("Model not found, inaccessible, and/or not deployed"). Swap the 6 references (3 files) to the currently-served accounts/fireworks/models/deepseek-v3p1 — the canonical model in Fireworks' current docs examples and present in LiteLLM's cost map. test_get_model_params_fireworks_ai is a pure pricing-heuristic test (no network) asserting the >16b branch, so it uses llama-v3p1-70b- instruct instead to keep the "fireworks-ai-above-16b" assertion and branch coverage intact. (cherry picked from commit 39a1d438f23f88d1c88f3e74930ab221b3e450de) * test(fireworks): mock remaining live smoke tests test_completion_fireworks_ai and test_completion_cost_fireworks_ai made real Fireworks calls and broke whenever Fireworks rotated its serverless catalog (no externally-verifiable model list exists). They also asserted nothing — just printed. Mock the HTTP post and assert real behavior instead: the request is built with the right model/messages and the OpenAI-compatible response parses back; the cost path yields a non-zero cost against the local cost map. No network, no model dependency, stronger than the old smoke checks. (cherry picked from commit b5db7ed37da21818c4defe030e3762447fe62e15) * fix(tests): replace shut-down gpt-4o-audio-preview with gpt-audio-1.5 (#28281) * fix(tests): replace shut-down gpt-4o-audio-preview with gpt-audio-1.5 OpenAI shut down gpt-4o-audio-preview on 2026-05-07, so the live audio calls in test_stream_chunk_builder_openai_audio_output_usage and test_standard_logging_payload_audio now hard-fail with a model-not-found error on every PR. The error was not "openai-internal", so the except block swallowed it and execution fell through to an unbound completion/response (UnboundLocalError). Switch both tests to gpt-audio-1.5, OpenAI's recommended successor (GA, not deprecated, already present in the litellm cost map so the response_cost assertion still resolves). Also broaden the except to skip with the real error in the reason instead of crashing, so a transient upstream blip can't reintroduce the UnboundLocalError. * fix(tests): narrow audio-test skip to model-not-found, re-raise the rest Address review feedback: an unconditional skip on any exception would silently mask a litellm-internal regression in the audio path (broken param transformation, serialization, bad header) instead of failing CI. Skip only on the upstream-unavailable class (model_not_found / "does not exist" / openai-internal) and re-raise everything else, so genuine regressions still fail loudly. The UnboundLocalError is still fixed because the handler either skips or raises - it never falls through. * fix(tests): add budget_exceeded to expected Interaction status enum Staging added budget_exceeded to the Interaction OpenAPI status enum; the staging merge into this branch picked up the spec change but not the matching test update, so test_status_enum_values failed in CI. Align the test's expected list (exact-match by design) with the live spec. * fix(tests): mock HTTP fetch in test_img_url_token_counter The test parameterized a live third-party image URL (blog.purpureus.net) which now 404s, causing get_image_dimensions to fall through to its base64 decode path and crash with 'not enough values to unpack' on every PR run. Mock safe_get with a tiny 1x1 PNG so the URL branch is still exercised without any network dependency. * fix(tests): swap gpt-4o-audio-preview to gpt-audio-1.5 in test_gpt4o_audio OpenAI shut down gpt-4o-audio-preview on 2026-05-07, so both live tests in test_gpt4o_audio.py (test_audio_output_from_model and test_audio_input_to_model) hard-fail model_not_found on every PR. Swap the hardcoded model to OpenAI's successor gpt-audio-1.5 (same chat-completions audio surface; already in the litellm cost map). Mirror the narrowed-skip pattern from the prior audio fixes: skip on model_not_found / does-not-exist / openai-internal, re-raise everything else so genuine litellm regressions still fail CI loudly. (cherry picked from commit 92de7423efca5756a2cb1bcf3228812628f91960) * fix(tests): migrate realtime + rerank tests off shut-down upstream models (#28191) * fix(tests): use gpt-realtime in realtime guardrails test OpenAI shut down gpt-4o-realtime-preview-2024-12-17 on 2026-05-07, so the live OpenAI realtime guardrails integration test now fails with model_not_found (session.created never arrives, _wait_for_event times out). Point OPENAI_REALTIME_URL at the current GA model, gpt-realtime. Scope limited to this test: the pricing-catalog JSON keeps the retired entries intentionally (historical cost calc + separate Azure timeline), and the Azure realtime cost-calc test is unaffected. * fix(tests): mock nvidia_nim rerank instead of hitting EOL'd endpoint NVIDIA reached end-of-life for the hosted nvidia/llama-3.2-nv-rerankqa-1b-v2 rerank API on 2026-05-18 with no published replacement, so the live BaseLLMRerankTest.test_basic_rerank for nvidia_nim now returns HTTP 410 ("Gone"). NVIDIA's hosted catalog rotates on a schedule, so swapping in another live model would only defer the failure. Override test_basic_rerank in TestNvidiaNim to mock the sync/async HTTP transport (same pattern as test_nvidia_nim_rerank_ranking_endpoint in this file) and inject a fake NVIDIA_NIM_API_KEY via monkeypatch. The request/response transformation and cost calculation stay covered offline. Scope limited to nvidia_nim; other BaseLLMRerankTest providers untouched. * fix(tests): migrate remaining realtime tests off shut-down gpt-4o-realtime-preview OpenAI's 2026-05-07 shutdown removed the entire gpt-4o-realtime-preview family, including the undated 'gpt-4o-realtime-preview' alias (not just the dated snapshot fixed earlier). Three live tests still connected with the dead alias and failed with messages_received=1 (an error event instead of session.created): - test_openai_realtime_simple.py: get_model() -> gpt-realtime (drives TestOpenAIRealtime.test_realtime_connection / test_realtime_with_query_params) - test_openai_realtime.py: test_openai_realtime_direct_call_no_intent and test_openai_realtime_direct_call_with_intent -> openai/gpt-realtime (the with_intent test shares the same dead alias even though it was not in the failing set this run) Mocked unit tests (test_realtime_query_params_construction, test_realtime_query_params_use_normalized_model_name) are left as-is: they never hit the network and assert string plumbing only. Also fixes test_text_message_blocked_by_guardrail_no_ai_response, which now connects (the earlier URL swap worked) but tripped a model-wording-brittle assertion. The guardrail flow asks the model to voice the block message verbatim; gpt-4o-realtime-preview complied (output contained 'blocked'), gpt-realtime refuses verbatim-repeat instructions ('I'm sorry, but I can't repeat that message.'). Since the original user message is blocked before it reaches OpenAI, the refusal is still a safe outcome. Assertion #3 now accepts both voicing and refusal, and adds a hard check that the blocked phrase never leaks into AI output. (cherry picked from commit ce87c411bfb33a8b37acaa630a39e4e4c8685add) * fix(model_prices): register mistral/ministral-8b-2512 Mistral's API now returns model='ministral-8b-2512' when 'mistral-tiny' is requested, so test_completion_mistral_api fails with 'This model isn't mapped yet'. Adding the entry so completion_cost can resolve the cost for that response. Author: Claude * fix(mcp,auth): address greptile review concerns - handle_sse_mcp now calls _raise_preemptive_401_for_unauthenticated_servers so SSE clients to pass-through OAuth MCP servers receive the RFC 9728 401 + WWW-Authenticate challenge that the streamable-HTTP path already emits. - get_request_route strips a trailing slash from root_path before length-based prefix removal so non-canonical ASGI root_path values like "/litellm/" don't strip the leading slash from the returned route. - _mcp_oauth_user_api_key_auth's cookie JWT decode now passes options={"verify_aud": False} so a future revision of the UI session JWT containing an aud claim cannot silently downgrade the request to unauthenticated. Co-authored-by: Claude * fix(tests): backfill local model_cost into remote-fetched map litellm.model_cost is loaded at import time from LITELLM_MODEL_COST_MAP_URL (pinned to main), so pricing entries that exist only in this branch (e.g. mistral/ministral-8b-2512, freshly added because Mistral's API now returns this id from mistral-tiny) are absent at test time and completion_cost lookups raise 'This model isn't mapped yet'. Backfill the in-tree backup into litellm.model_cost in the local_testing conftest so cassette-driven cost calculations resolve against the entries that ship with the branch under test. Fixes local_testing_part1 failures on test_completion_mistral_api and test_completion_mistral_api_modified_input. * fix(mcp,jwt): address greptile concurrency and code-quality concerns - _apply_issuer_claim_mappings now builds a new dict and reads from the original token, rather than mutating its input. The change is behaviour-preserving (caller passes a fresh jwt.decode result), but avoids the surprise-mutation pattern flagged by greptile. - is_network_error uses isinstance(exc, httpx.TransportError) instead of matching type(exc).__name__ against a hand-maintained string set, so ReadError / WriteError / ProxyError / etc. are also treated as transport-level failures and surfaced as HTTP 502. - fetch_upstream_oauth_protected_resource now coalesces concurrent discovery requests per (server_id, resource_url) through an asyncio.Lock so concurrent .well-known calls share a single upstream fetch + cache write. - Drop the redundant 'if trusted_ranges:' branch in get_mcp_client_ip; it is always true on the path that reaches it (the prior 'if not trusted_ranges:' early-returns). Co-authored-by: Claude * fix(jwt,mcp): fall back to global JWKS on unknown issuer; prune fetch locks - handle_jwt._get_configured_issuer now returns None for tokens whose 'iss' is not in the configured issuers list, letting auth_jwt fall through to the legacy JWT_PUBLIC_KEY_URL path instead of hard-raising. This keeps existing tokens from non-configured IdPs working when an operator adds the new 'issuers' list to a live deployment. - discoverable_endpoints._prune_oauth_metadata_cache now also prunes entries in _OAUTH_METADATA_FETCH_LOCKS whose cache entry has been evicted and whose lock isn't currently held, bounding the locks dict to match the cache it guards. Co-authored-by: Claude * fix(mcp,auth): restore client_ip in oauth2 target check, drop from delegate check The merge of staging into the PR branch (d42a66adb6) misplaced the client_ip=client_ip kwarg: it landed inside _target_servers_delegate_auth_to_upstream (which never accepted client_ip and isn't called with it), while the sibling _target_servers_use_oauth2 has client_ip in its signature but stopped passing it through to get_mcp_server_by_name. That left ruff flagging F821 on the undefined name and lint failing. Move client_ip back into _target_servers_use_oauth2's lookup (matching the call site that already forwards IPAddressUtils.get_mcp_client_ip) and drop it from _target_servers_delegate_auth_to_upstream so its body matches its signature again. * fix(mcp): respect client ip for delegated auth * fix(auth): address remaining greptile style findings - get_request_route: require root_path to match whole path segments before stripping, so '/apifoo' isn't truncated to 'foo' when root_path='/api'. - get_mcp_client_ip: collapse the two trusted-proxy validation branches into a single is_request_from_trusted_proxy call so the return value drives control flow instead of being discarded for the side-effect warning. Co-authored-by: Claude * fix(jwt): strip internal _litellm_* claims in global JWKS auth path Prevents identity spoofing where a token signed by the global JWKS could inject _litellm_jwt_issuer and other _litellm_* claims that downstream getters trust. The issuer-scoped path already strips these via _apply_issuer_claim_mappings; mirror that behavior for the global fallback path. Co-authored-by: Yassin Kortam * fix(mcp): surface MCPUpstreamAuthError as 401 in SSE/HTTP transport handlers Both handle_sse_mcp and handle_streamable_http_mcp only caught HTTPException to preserve 401 + WWW-Authenticate challenges, but MCPUpstreamAuthError (raised when a pass-through server's upstream rejects a bearer token mid-session) inherits from Exception. It was falling through to the generic handler and surfacing as an opaque 500. Mirror the REST endpoint behavior: translate MCPUpstreamAuthError into an HTTPException(status_code=e.status_code) with the upstream www-authenticate header so standards-compliant MCP clients trigger the upstream OAuth flow. Co-authored-by: Yassin Kortam * fix(mcp): add upstream auth pre-flight in SSE handler Mirror handle_streamable_http_mcp by calling _check_passthrough_upstream_auth after the cold-start 401 emitter so expired/invalid upstream tokens surface a proper 401 + WWW-Authenticate challenge before the SSE session commits 200 headers, instead of letting list_tools silently return [] when the upstream rejects the token. Co-authored-by: Claude * fix(mcp): tighten cold-start bypass against CSV paths + dedupe upstream auth probe - Return None from _parse_mcp_server_names_from_path for CSV multi-server paths (/mcp/a,b). The regex previously truncated at the first comma and silently passed a single server name to the cold-start gate. - Switch _is_mcp_passthrough_cold_start to all-targets semantics, matching _target_servers_use_oauth2: one non-passthrough target in a co-targeted set must not flip the anonymous-admission bypass open for the others. - Drop the redundant HTTPStatusError block in _extract_upstream_auth_failure - any HTTPStatusError carries a .response, so the preceding generic block already handles 401/403 detection. Co-authored-by: Yassin Kortam * fix(mcp,tests): sync stubs and cold-start assertions with delegate-check The merge of base-branch _target_servers_delegate_auth_to_upstream into process_mcp_request inserts an additional get_mcp_server_by_name(name) lookup ahead of the cold-start path, which breaks two test patterns: 1. lookup_by_name(name) side-effect stubs in TestMCPDelegateAuthToUpstream are called positionally by the delegate check, then again by the cold-start path with client_ip=... — raising TypeError: unexpected keyword argument 'client_ip'. Accept **_kwargs to match the real signature. 2. TestMCPPassthroughColdStartAdmission assertions count the lookup exactly once with client_ip=..., but the delegate check now adds a positional-only call ahead of it. Switch assert_called_once_with to assert_any_call for the cold-start invocation, and assert client_ip was *not* passed for the aggregate /mcp test where cold-start must not fire. Both updates align with CLAUDE.md guidance to keep monkeypatch stubs in sync with the real signature when an optional parameter is added. Co-authored-by: Claude * fix(mcp): correct passthrough probe 401 + slashed-name cold start parser - _check_passthrough_upstream_auth now emits 'Bearer resource_metadata="..."' pointing at the gateway's oauth-protected-resource well-known URL, mirroring the pre-emptive 401 path. Pass-through servers don't use the gateway as an authorization server, so the previous 'authorization_uri=' challenge sent clients to the wrong metadata endpoint. - _parse_mcp_server_names_from_path now accepts server names that contain a single slash (e.g. custom_solutions/user_123), mirroring MCPRequestHandler._extract_target_server_names_from_path. Without this, the cold-start bypass missed slashed-name servers and the generic admission error propagated instead of the spec-compliant 401 challenge. - _is_mcp_passthrough_cold_start drops the unused scope parameter from its signature. Co-authored-by: Yassin Kortam * style(mcp): format discoverable endpoints * refactor(mcp): dedupe MCPUpstreamAuthError->HTTPException + thread client_ip into delegate-auth gate Co-authored-by: Yassin Kortam * fix(mcp): handle passthrough OAuth metadata and startup auth errors - discoverable_endpoints: For pass-through MCP servers, when upstream oauth-protected-resource returns a non-200/non-dict response, raise HTTP 502 instead of falling through to default gateway metadata. Falling through would direct MCP clients at the gateway, which is not the authorization server for pass-through configs. - mcp_server_manager: Wrap _get_tools_from_server in startup tool name mapping with try/except. Since _get_tools_from_server now re-raises MCPUpstreamAuthError, an upstream 401 from a pass-through server at startup (when no user token is present) would otherwise abort the loop and leave subsequent servers unmapped. Co-authored-by: Yassin Kortam * fix(mcp): restrict passthrough probe challenge to OAuth passthrough servers The probe filter previously matched any server with Authorization in extra_headers, including gateway-managed OAuth2 servers. Those would then receive the resource_metadata= WWW-Authenticate challenge meant for pass-through servers, instead of the authorization_uri= challenge pointing at the gateway AS metadata. Use srv.is_oauth_passthrough so only genuine pass-through servers get the resource-metadata challenge. Co-authored-by: Yassin Kortam * test(proxy): cover issuer-scoped JWT auth * fix(mcp): use resource metadata for passthrough reauth * fix(mcp,tests): assert cold-start helper directly for aggregate /mcp Threading client_ip into _target_servers_delegate_auth_to_upstream made get_mcp_server_by_name(name, client_ip=...) also fire from the delegate-auth check, so the call_args_list assertion on client_ip-in-kwargs no longer uniquely signals a cold-start lookup. Patch _is_mcp_passthrough_cold_start and assert it is not invoked, which is the actual contract the test is pinning. * fix(mcp,jwt): drop unneeded async helper + suppress misleading unscoped JWT warning - _build_oauth_authorization_server_response: revert to sync (no awaits in body). The function only does dict construction and synchronous registry lookups; async added coroutine creation overhead per discovery call without need. - _build_decode_kwargs: accept has_issuer_config so the global path's 'JWT auth is unscoped' warning is suppressed when LiteLLM_JWTAuth.issuers provides per-issuer scoping. Previously the warning fired spuriously for admins who intentionally use only the new issuers config. * fix(jwt,mcp): clarify issuers fallthrough + add TTL on mcp permission cache - LiteLLM_JWTAuth.issuers docs now state explicitly that unlisted issuers fall back to the global JWT_AUDIENCE/JWT_ISSUER path; the field is additive routing, not an allow-list. Matches actual control flow in handle_jwt.auth_jwt and the regression tests asserting backwards compatibility with the global JWKS path. - MCPRequestHandler._get_{org,agent}_object_permission now pass ttl=DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL on async_set_cache, mirroring the auth_checks.py pattern so the cache TTL is explicit on both DualCache layers. * fix(tests): align merged JWT and MCP cold-start assertions Update the tests carried over from PR #28008 to match the assertions on the staging branch: - tests/test_litellm/proxy/auth/test_handle_jwt.py: unknown issuers now fall back to the legacy JWT_PUBLIC_KEY_URL path (per litellm_feat/v1.84.0-mcp-gateway-jwt-auth's '\''fall back to global JWKS on unknown issuer'\''), and mapped issuer claims that are absent no longer fail closed — they simply leave the normalised LiteLLM internal claim absent. - tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py: the aggregate '\''/mcp'\'' route still triggers the delegate-auth-to-upstream lookup once for the header-supplied server name; cold-start admission must NOT fire on top of that. Tighten the assertion to assert_called_once_with so a future regression that re-enters cold-start is caught. Co-authored-by: Mateo Wang * fix(jwt): guard litellm_jwtauth access in auth_jwt global path JWTHandler() can be constructed without update_environment() being called (tests do this directly), in which case self.litellm_jwtauth does not exist. Accessing it raises AttributeError before getattr can fall back. Use the same safe pattern other call sites use. * Gate MCP OAuth pass-through on delegate_auth_to_upstream flag Sameer's review on #28356/#28008 flagged that the new pass-through behaviors (preemptive 401 challenges, /.well-known/oauth-protected- resource proxying, upstream 401/403 propagation as MCPUpstreamAuthError, and Authorization-stripping when no x-litellm-api-key is supplied) were implicitly enabled for every server with auth_type=none plus Authorization in extra_headers. Existing users doing static bearer pass-through for non-OAuth reasons would have silently regressed. Make the detection rule explicit: extend the existing delegate_auth_to_upstream flag (previously oauth2-only) to also gate is_oauth_passthrough. Now requires flag + auth_type=None + Authorization in extra_headers, per Sameer's suggested detection rule. The UI toggle now appears for both modes (oauth2 PKCE passthrough and auth_type=none OAuth pass-through) with mode-appropriate copy. Update test fixtures to set the flag where the test intent is to exercise OAuth pass-through behavior, and add negative tests covering the new default-false case. * fix(mcp): route org object_permission lookup through shared auth helpers Replace the bespoke litellm_organizationtable.find_unique + dedicated cache key in _get_org_object_permission with get_org_object + get_object_permission so MCP requests share the same user_api_key_cache entries as the rest of the proxy and no longer fragment org-row caching. * fix(mcp): wrap get_object_permission call in shared try/except Ensure exceptions from get_object_permission in _get_org_object_permission are caught and return None, preserving the original fail-safe semantics. Co-authored-by: Yassin Kortam * fix(jwt): validate issuer audience at config load + dedicated key-miss exception - Move JWTIssuerConfig audience-required guard into a Pydantic model_validator so misconfiguration fails at startup instead of on the first request. - Replace the string-match `No matching public key found` filter in get_public_key's multi-URL fallback with a dedicated NoMatchingJWTPublicKeyError; only that specific exception triggers continuation, every other error still surfaces. * fix(mcp): admit and forward Authorization for passthrough OAuth return For pass-through MCP servers (auth_type=none with delegate_auth_to_upstream) the RFC 9728 cold-start flow sends the client back with only "Authorization: Bearer " after upstream OAuth discovery. Previously this path 1) was rejected in process_mcp_request because the oauth2_headers fallback only covered auth_type=oauth2 targets, and 2) had the Authorization header stripped by _prepare_mcp_server_headers when no x-litellm-api-key was present, treating the upstream token as a potential LiteLLM key leak. - Extend the elif oauth2_headers fallback to also admit anonymously when every target is a pass-through server. - Pass user_api_key_auth into _prepare_mcp_server_headers so it can forward Authorization for pass-through servers when admission did not consume the bearer as a LiteLLM key (api_key is unset). Co-authored-by: Yassin Kortam * fix(mcp): consistent www-authenticate casing + SSE toolset scoping - Normalize the WWW-Authenticate header key emitted by _check_passthrough_upstream_auth to lowercase to match the other 401 emitters in the OAuth pass-through flow. - Mirror the streamable HTTP handler's toolset scoping in handle_sse_mcp: strip client-supplied x-mcp-toolset-id and apply _apply_toolset_scope before _check_passthrough_upstream_auth so the upstream probe list is derived from the fully-authorized server set. - Tighten _has_client_supplied_mcp_auth signature so mcp_server_auth_headers is Optional, matching its caller in process_mcp_request. Co-authored-by: Yassin Kortam * security(mcp): strip Authorization in call_tool when LiteLLM admission used legacy header Mirror the OAuth pass-through admission check from _prepare_mcp_server_headers (list-tools path) in _call_regular_mcp_tool (tool-call path): when the server is OAuth pass-through and the caller did not supply x-litellm-api-key, Authorization on the inbound request may itself be the LiteLLM API key — so strip it before forwarding instead of leaking the gateway credential upstream. When x-litellm-api-key is present, admission is unambiguous and Authorization continues to carry the upstream OAuth bearer (transparent pass-through). * refactor(mcp): centralize caller Authorization strip decision Extracted the security-sensitive logic that decides whether the caller's Authorization header is forwarded to (or stripped from) an outgoing MCP request into a single helper, _should_strip_caller_authorization, in mcp_server_manager.py. Previously the same condition was duplicated across _call_regular_mcp_tool (mcp_server_manager.py) and _prepare_mcp_server_headers (server.py). Keeping two copies of this check risked future divergence and credential-leak / broken-passthrough bugs. Both call sites now share the helper, preserving exact behavior. Co-authored-by: Yassin Kortam * log MCP OAuth discovery diagnostics for unmatched paths and non-transport upstream errors * fix(jwt): include issuer-normalized team id in get_all_jwt_team_ids The aggregator for team IDs only consulted the issuer-normalized claim for the plural (team_ids) path and fell back to the global config for the singular path. When an operator configures team_id_jwt_field only at the issuer level, get_team_id correctly returned the mapped value but get_all_jwt_team_ids silently dropped it, causing membership reconciliation to disagree with request routing. Co-authored-by: Yassin Kortam * fix(mcp/jwt): dedupe cold-start path parser; reject conflicting audience flags - _parse_mcp_server_names_from_path now delegates to MCPRequestHandler._extract_target_server_names_from_path so the names used by the cold-start passthrough bypass cannot drift from the names used by downstream routing. - JWTIssuerConfig now rejects the combination of audience and disable_audience_validation=True at validation time instead of silently ignoring the flag. * fix(mcp): restrict passthrough cold-start bypass to 401 only The new elif passthrough cold-start branch reused is_auth_error which matches both 401 and 403. A 403 from user_api_key_auth indicates the LiteLLM key WAS recognized but is forbidden (e.g. over budget / rate limited); falling through to anonymous UserAPIKeyAuth() in that case bypasses spend and rate-limit controls on passthrough servers. Only trigger the cold-start anonymous admission on 401, which is the signal that the bearer is an upstream OAuth token rather than a recognized LiteLLM key. Co-authored-by: Yassin Kortam * fix(jwt/mcp): warn on unscoped JWT fallback; route agent permission lookup through shared helper - _build_decode_kwargs no longer suppresses the unscoped-fallback warning when LiteLLM_JWTAuth.issuers is set: tokens whose iss does not match any configured issuer still fall through to the global path, and that fallback is itself unscoped when JWT_AUDIENCE/JWT_ISSUER are absent. - _get_agent_object_permission now caches the agent_id -> object_permission_id mapping and delegates the permission lookup to the shared get_object_permission helper, so the agent path reuses the same cache entries as the org / team / key paths. * fix(mcp): fabricate resource_metadata challenge when upstream 401 omits WWW-Authenticate When an upstream pass-through MCP server returns 401 without a WWW-Authenticate header (non-compliant per RFC 7235 §3.1), to_http_exception() now produces a synthetic Bearer challenge pointing at the gateway's standard-pattern oauth-protected-resource well-known endpoint for that server. This keeps MCP clients on the RFC 9728 discovery flow instead of receiving a bare 401 with no recovery hint. * fix(jwt): make _get_decode_options explicitly control verify_iss Previously, _get_decode_options only set verify_aud based on whether audience was provided. The issuer JWT path relied on always passing issuer=issuer_config.issuer to trigger PyJWT's default verify_iss=True, making the helper's behavior implicitly dependent on caller behavior. Now _get_decode_options accepts issuer as well, mirroring the verify_aud handling and matching the dimensions handled by _build_decode_kwargs. Co-authored-by: Yassin Kortam * fix(mcp): emit absolute resource_metadata URI in fabricated 401 challenge Per RFC 9728 §3.2 the resource_metadata Bearer challenge must be an absolute URI; strict MCP clients reject relative URIs and fail to initiate discovery. MCPUpstreamAuthError.to_http_exception now accepts the gateway base URL and prepends it when the upstream omitted WWW-Authenticate, and all four call sites (streamable HTTP, SSE, and the two REST tool-list paths) supply it. * fix(mcp): correct 403 detail text and remove dead _list_tools_for_single_server duplicate - MCPUpstreamAuthError.to_http_exception() now returns detail='Forbidden' for 403 upstream responses (and 'Unauthorized' for 401), matching the _check_passthrough_upstream_auth pre-flight probe. - Remove the shadowed first definition of _list_tools_for_single_server in rest_endpoints.py; the second definition was the live one and the dead copy was a maintenance trap. Co-authored-by: Yassin Kortam * fix: address potential bugs in auth_utils, mcp discoverable endpoints, and mcp auth - auth_utils.get_request_route: return '/' instead of empty string when raw_path exactly equals root_path so downstream route allowlist checks still see a leading slash - discoverable_endpoints.fetch_upstream_oauth_protected_resource: also cache negative results (no upstream metadata) for a shorter TTL so we don't re-fetch on every discovery request and so the per-key fetch lock can be pruned - user_api_key_auth_mcp: guard the oauth2_headers 401 cold-start passthrough bypass with _has_client_supplied_mcp_auth, matching the parallel bypass in the no-Authorization branch so MCP-auth-bearing requests don't silently downgrade to anonymous admission Co-authored-by: Yassin Kortam * test(vertex): tolerate transient InternalServerError in google maps tool test test_gemini_google_maps_tool_simple makes live calls to Vertex AI's Google Maps grounding backend, which intermittently returns 500 INTERNAL ("Please retry") — a transient upstream failure, not a LiteLLM bug. The test already passes on RateLimitError; treat InternalServerError the same way so transient Vertex-side failures don't fail CI. * refactor(mcp): drop redundant has_client_credentials filter on passthrough probe is_oauth_passthrough already requires auth_type in (None, MCPAuth.none), which is mutually exclusive with has_client_credentials (auth_type == MCPAuth.oauth2), so the extra guard was always True and only added confusion about whether a server could be both passthrough and M2M. Co-authored-by: Yassin Kortam * fix: restore unreachable InternalServerError skip handler in vertex test Co-authored-by: Yassin Kortam * feat(mcp): add dedicated oauth_passthrough flag for non-oauth2 pass-through Previously is_oauth_passthrough reused delegate_auth_to_upstream — a flag scoped to oauth2 servers (PKCE bypass) — to gate OAuth pass-through for auth_type=none servers. Overloading it risked regressing existing deployments that set delegate_auth_to_upstream, since the same flag would silently start driving pass-through (discovery proxying, 401 challenges, upstream 401/403 propagation) on non-oauth2 servers. Introduce a separate oauth_passthrough opt-in so the two behaviors never imply each other: - MCPServer.is_oauth_passthrough now requires oauth_passthrough (not delegate_auth_to_upstream). - Persist oauth_passthrough on LiteLLM_MCPServerTable (new column + migration) and wire it through config/DB load and API responses. - UI splits the single toggle into two: "Delegate auth to upstream (PKCE passthrough)" for oauth2 and "OAuth pass-through" for auth_type=none servers forwarding Authorization. Adds backend tests (property, round-trip, and a regression guard that delegate_auth_to_upstream alone never enables pass-through) and UI tests for the toggle split. * fix(mcp): reconcile cold-start bypass with x-mcp-servers header and skip non-absolute WWW-Authenticate fabrication - _parse_mcp_server_names_from_path now fails closed when the x-mcp-servers header introduces any target not present in the path-derived target set, closing a header/path mismatch where the cold-start passthrough bypass could otherwise admit anonymously while the header advertises a non-passthrough server. - MCPUpstreamAuthError.to_http_exception no longer emits a relative resource_metadata URI when base_url is missing; per RFC 9728 3.2 the URI must be absolute, so we skip fabrication entirely rather than send a challenge strict MCP clients will reject. Co-authored-by: Yassin Kortam * fix(mcp): fabricate path-aware resource_metadata URI for upstream 401 When MCPUpstreamAuthError.to_http_exception fabricates a `WWW-Authenticate: Bearer resource_metadata=...` challenge (because the upstream 401 omitted one), the URL now matches the inbound MCP transport pattern the client originally used: - /mcp/{server_name} -> /.well-known/oauth-protected-resource/mcp/{server_name} - /{server_name}/mcp -> /.well-known/oauth-protected-resource/{server_name}/mcp This mirrors the path-aware behaviour of _get_passthrough_resource_metadata_url in server.py so strict RFC 9728 \xA73.2 clients on legacy routes get a resource_metadata URI aligned with the resource pattern they originally targeted. Co-authored-by: Yassin Kortam * fix(jwt+mcp): tighten issuer-scoped claim type handling, RFC-quote authorization_uri, surface MCP upstream auth errors, defense-in-depth on decode options - handle_jwt: when an issuer-scoped _litellm_team_ids claim exists but has an unexpected type, return [] instead of falling through to the global team_ids_jwt_field path (different claim semantically). - handle_jwt: _get_decode_options/_decode_jwt_with_public_key now take an explicit disable_audience_validation flag; passing audience=None without it raises, so audience checks can't silently disappear if the model validator is ever bypassed. _auth_jwt_with_issuer forwards the flag from JWTIssuerConfig. - mcp_server: quote the authorization_uri WWW-Authenticate parameter value (RFC 6750 / 9728 auth-param must be quoted-string), matching the pass-through path. - mcp_server: in _fetch_and_filter_server_tools, re-raise MCPUpstreamAuthError so the outer streamable-HTTP handler can surface a proper 401 + WWW-Authenticate challenge instead of returning an empty tool list. Co-authored-by: Yassin Kortam * chore(docker): align Dockerfile.non_root/Dockerfile.database to current wolfi-base SHA The older sha256:3258be... pin has been intermittently returning 500/not-found from cgr.dev, breaking the test-server-root-path GitHub Action and the build_docker_database_image CircleCI job. Move both Dockerfiles onto the same sha256:31da65... digest already in use by Dockerfile, gateway/Dockerfile, backend/Dockerfile, and migrations/Dockerfile so the base image is consistent across the repo. * ci(docker): bump wolfi-base pin to current working digest The previously aligned sha256:31da6565f35a... and the older sha256:3258be... both return HTTP 500 from cgr.dev's manifest endpoint, breaking the build_docker_database_image CircleCI job and test-server-root-path GitHub Action. The current 'latest' tag resolves to sha256:5743937d521c... which serves manifests normally, so move docker/Dockerfile.database and docker/Dockerfile.non_root onto that digest. * ci(docker): retry apk add in Dockerfile.database for apk.cgr.dev flakes Mirror the retry-loop pattern from #28888 (which fixed backend/Dockerfile, gateway/Dockerfile, and migrations/Dockerfile) into docker/Dockerfile.database. The build_docker_database_image CI job has been intermittently failing with "remote server returned error (try 'apk update')" when apk.cgr.dev flakes mid-fetch; bumping the wolfi-base SHA doesn't address the mirror, only a retry does. Same explicit-failure form as #28888: exit non-zero on the 3rd miss instead of silently succeeding because `sleep 5` was the last command in the `&& break || sleep 5` chain. * fix(mcp): scope preemptive 401 to toolset-narrowed server set Move _raise_preemptive_401_for_unauthenticated_servers after toolset scoping in both the StreamableHTTP and SSE handlers, and add an optional allowed_server_ids parameter so passthrough/oauth2 servers that the active toolset excludes no longer trigger a spurious 401 challenge. Without this, a client targeting a toolset whose scope excludes a passthrough server could be pushed into an OAuth flow for a server it would be 403'd on immediately after authentication. Co-authored-by: Yassin Kortam * revert(docker): drop unrelated Wolfi bump and apk retry loop from MCP/JWT PR These Docker changes are out of scope for the MCP OAuth passthrough + JWT auth work and duplicate the build-reliability fix already merged to litellm_internal_staging in #28888, which adds the same apk retry loop on the componentized backend/gateway/migrations Dockerfiles and also fixes the underlying nodeenv/libatomic root cause. Restoring docker/Dockerfile.database and docker/Dockerfile.non_root to the base so this PR is purely the MCP/JWT change. * fix(mcp): surface upstream 403 challenges from REST tools/list The single-server pass-through path converted an upstream MCPUpstreamAuthError into an HTTPException, but list_tool_rest_api only re-raised 401s; an upstream 403 (valid token, insufficient scope) collapsed into a 200 response with error=unexpected_error, so clients never saw the status or WWW-Authenticate challenge needed to refresh scopes. Let MCPUpstreamAuthError propagate and convert it once in list_tool_rest_api so both 401 and 403 reach the client, while internal access/IP 403s keep the legacy error-dict shape. * fix(mcp): fail closed for IP access control when XFF trusted ranges unset When use_x_forwarded_for is enabled but mcp_trusted_proxy_ranges is not configured, get_mcp_client_ip previously fell back to the direct peer IP. Behind an internal reverse proxy that peer is the proxy's private address, so every external caller was classified as internal and could reach MCP servers with available_on_public_internet=false. Return an empty string in that case so is_internal_ip treats the caller as external. --------- Co-authored-by: Yuneng Jiang Co-authored-by: Milan Co-authored-by: Cursor Co-authored-by: Mateo Wang Co-authored-by: Krrish Dholakia Co-authored-by: ryan-crabbe-berri Co-authored-by: stuxf <70670632+stuxf@users.noreply.github.com> Co-authored-by: Claude Sonnet 4.6 (1M context) Co-authored-by: veria-ai[bot] <224490171+veria-ai[bot]@users.noreply.github.com> Co-authored-by: gym-cmd <186399764+gym-cmd@users.noreply.github.com> Co-authored-by: Artem Dudarev Co-authored-by: Claude Co-authored-by: Yassin Kortam --- .../migration.sql | 2 + .../litellm_proxy_extras/schema.prisma | 1 + litellm/experimental_mcp_client/client.py | 14 +- .../mcp_server/auth/user_api_key_auth_mcp.py | 305 +++++++-- .../mcp_server/discoverable_endpoints.py | 263 +++++++- .../_experimental/mcp_server/exceptions.py | 80 +++ .../mcp_server/mcp_server_manager.py | 173 ++++- .../mcp_server/rest_endpoints.py | 122 +--- .../proxy/_experimental/mcp_server/server.py | 321 +++++++-- litellm/proxy/_types.py | 75 ++- litellm/proxy/auth/auth_utils.py | 18 +- litellm/proxy/auth/handle_jwt.py | 440 ++++++++++--- litellm/proxy/auth/ip_address_utils.py | 23 +- .../mcp_management_endpoints.py | 2 +- litellm/proxy/proxy_server.py | 3 + litellm/proxy/schema.prisma | 1 + litellm/types/interactions/generated.py | 16 +- .../types/mcp_server/mcp_server_manager.py | 61 +- schema.prisma | 1 + .../mcp_server/test_discoverable_endpoints.py | 178 ++++- tests/local_testing/conftest.py | 14 +- tests/mcp_tests/test_mcp_server.py | 63 ++ tests/mcp_tests/test_per_user_oauth_cache.py | 25 + tests/proxy_unit_tests/test_jwt.py | 528 ++++++++++++++- .../auth/test_user_api_key_auth_mcp.py | 502 ++++++++++++++- .../mcp_server/test_discoverable_endpoints.py | 4 +- .../mcp_server/test_mcp_oauth_passthrough.py | 474 ++++++++++++++ .../test_mcp_oauth_passthrough_cold_start.py | 156 +++++ .../test_mcp_oauth_passthrough_tools.py | 197 ++++++ .../mcp_server/test_mcp_server.py | 120 +++- .../mcp_server/test_mcp_server_manager.py | 161 +++++ .../mcp_server/test_rest_endpoints.py | 72 +++ .../proxy/auth/test_handle_jwt.py | 608 ++++++++++++++++++ .../proxy/auth/test_mcp_ip_filtering.py | 74 ++- .../test_mcp_management_endpoints.py | 31 + .../MCPPermissionManagement.test.tsx | 62 ++ .../mcp_tools/MCPPermissionManagement.tsx | 55 +- .../mcp_tools/create_mcp_server.tsx | 2 + .../mcp_tools/mcp_server_edit.test.tsx | 35 + .../components/mcp_tools/mcp_server_edit.tsx | 35 +- .../components/mcp_tools/mcp_server_view.tsx | 21 + .../src/components/mcp_tools/types.tsx | 1 + 42 files changed, 4925 insertions(+), 414 deletions(-) create mode 100644 litellm-proxy-extras/litellm_proxy_extras/migrations/20260526120000_add_oauth_passthrough_to_mcp_servers/migration.sql create mode 100644 litellm/proxy/_experimental/mcp_server/exceptions.py create mode 100644 tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_oauth_passthrough.py create mode 100644 tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_oauth_passthrough_cold_start.py create mode 100644 tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_oauth_passthrough_tools.py diff --git a/litellm-proxy-extras/litellm_proxy_extras/migrations/20260526120000_add_oauth_passthrough_to_mcp_servers/migration.sql b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260526120000_add_oauth_passthrough_to_mcp_servers/migration.sql new file mode 100644 index 00000000000..3c387891a5e --- /dev/null +++ b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260526120000_add_oauth_passthrough_to_mcp_servers/migration.sql @@ -0,0 +1,2 @@ +-- AlterTable +ALTER TABLE "LiteLLM_MCPServerTable" ADD COLUMN IF NOT EXISTS "oauth_passthrough" BOOLEAN NOT NULL DEFAULT false; diff --git a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma index 78143fe0411..c4754ef6117 100644 --- a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma +++ b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma @@ -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? diff --git a/litellm/experimental_mcp_client/client.py b/litellm/experimental_mcp_client/client.py index 0dc56b6a3bc..7559fe142c4 100644 --- a/litellm/experimental_mcp_client/client.py +++ b/litellm/experimental_mcp_client/client.py @@ -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 [] diff --git a/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py b/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py index 2aacab80f57..863e6acd41e 100644 --- a/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py +++ b/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py @@ -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 + # ``. 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 diff --git a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py index 8324ba641a4..ed374635fea 100644 --- a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py @@ -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, diff --git a/litellm/proxy/_experimental/mcp_server/exceptions.py b/litellm/proxy/_experimental/mcp_server/exceptions.py new file mode 100644 index 00000000000..fd8fc3d5e58 --- /dev/null +++ b/litellm/proxy/_experimental/mcp_server/exceptions.py @@ -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, + ) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index b4678a50b2c..129dfc102f9 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -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 [] @@ -2768,6 +2903,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 +2969,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 +3270,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 +3302,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 +3910,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, diff --git a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py index 693ca5a7642..e20c9f3a082 100644 --- a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py @@ -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) diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index dd36e32291d..6e33a105ec8 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -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, ) @@ -187,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 ( @@ -1099,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-`` 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], @@ -1279,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 @@ -1311,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: @@ -1523,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 @@ -1574,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)}" @@ -1697,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: @@ -1754,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: @@ -1809,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: @@ -2645,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( @@ -2695,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( @@ -3154,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. @@ -3266,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 @@ -3282,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 @@ -3329,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"] = [ @@ -3373,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 @@ -3607,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 @@ -3650,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, @@ -3670,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 diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 98a17e4be95..9f89cae1a41 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -1290,6 +1290,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 +1374,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 +1447,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 @@ -2516,7 +2519,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, @@ -4436,6 +4439,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 +4609,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: diff --git a/litellm/proxy/auth/auth_utils.py b/litellm/proxy/auth/auth_utils.py index 86265270357..4e5169d8d84 100644 --- a/litellm/proxy/auth/auth_utils.py +++ b/litellm/proxy/auth/auth_utils.py @@ -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( diff --git a/litellm/proxy/auth/handle_jwt.py b/litellm/proxy/auth/handle_jwt.py index 9838b4ba49b..654e7e0ff28 100644 --- a/litellm/proxy/auth/handle_jwt.py +++ b/litellm/proxy/auth/handle_jwt.py @@ -12,7 +12,7 @@ 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 @@ -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,73 +890,217 @@ 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 Exception("Token Expired") + 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 - - 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() - ) - - # 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 + 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 diff --git a/litellm/proxy/auth/ip_address_utils.py b/litellm/proxy/auth/ip_address_utils.py index 39d3282942f..be0d83dfcdc 100644 --- a/litellm/proxy/auth/ip_address_utils.py +++ b/litellm/proxy/auth/ip_address_utils.py @@ -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) diff --git a/litellm/proxy/management_endpoints/mcp_management_endpoints.py b/litellm/proxy/management_endpoints/mcp_management_endpoints.py index b35e2b6e3fd..4d67df16b0b 100644 --- a/litellm/proxy/management_endpoints/mcp_management_endpoints.py +++ b/litellm/proxy/management_endpoints/mcp_management_endpoints.py @@ -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", "") diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index e0f139dee57..61df62ba2be 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -15845,6 +15845,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 +15994,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: diff --git a/litellm/proxy/schema.prisma b/litellm/proxy/schema.prisma index 78143fe0411..c4754ef6117 100644 --- a/litellm/proxy/schema.prisma +++ b/litellm/proxy/schema.prisma @@ -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? diff --git a/litellm/types/interactions/generated.py b/litellm/types/interactions/generated.py index d546e897891..b38cd8f58b9 100644 --- a/litellm/types/interactions/generated.py +++ b/litellm/types/interactions/generated.py @@ -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]): diff --git a/litellm/types/mcp_server/mcp_server_manager.py b/litellm/types/mcp_server/mcp_server_manager.py index 13e325838dc..6aa62c35106 100644 --- a/litellm/types/mcp_server/mcp_server_manager.py +++ b/litellm/types/mcp_server/mcp_server_manager.py @@ -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).""" diff --git a/schema.prisma b/schema.prisma index 78143fe0411..c4754ef6117 100644 --- a/schema.prisma +++ b/schema.prisma @@ -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? diff --git a/tests/litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py b/tests/litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py index 8785e450a4b..2a8768df722 100644 --- a/tests/litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py +++ b/tests/litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py @@ -1,8 +1,32 @@ """Tests for MCP OAuth discoverable endpoints""" import pytest +from fastapi import HTTPException from unittest.mock import AsyncMock, MagicMock, patch +TRUSTED_PROXY_IP = "10.0.0.5" +TRUSTED_PROXY_RANGES = ["10.0.0.0/8"] + + +def set_request_from_trusted_proxy(mock_request): + mock_request.client = MagicMock() + mock_request.client.host = TRUSTED_PROXY_IP + + +@pytest.fixture +def trusted_proxy_origin_headers(): + with ( + patch( + "litellm.proxy._experimental.mcp_server.discoverable_endpoints.IPAddressUtils.is_request_from_trusted_proxy", + return_value=True, + ), + patch( + "litellm.proxy._experimental.mcp_server.oauth_utils.IPAddressUtils.is_request_from_trusted_proxy", + return_value=True, + ), + ): + yield + @pytest.mark.asyncio async def test_authorize_endpoint_includes_response_type(): @@ -56,7 +80,7 @@ async def test_authorize_endpoint_includes_response_type(): request=mock_request, client_id="test_client_id", mcp_server_name="test_oauth", - redirect_uri="https://client.example.com/callback", + redirect_uri="http://127.0.0.1:60108/callback", state="test_state", ) @@ -154,7 +178,6 @@ async def test_token_endpoint_forwards_code_verifier(): from litellm.types.mcp_server.mcp_server_manager import MCPServer from litellm.proxy._types import MCPTransport from fastapi import Request - import httpx except ImportError: pytest.skip("MCP discoverable endpoints not available") @@ -244,10 +267,15 @@ async def test_register_client_without_mcp_server_name_returns_dummy(): from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( register_client, ) + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + global_mcp_server_manager, + ) from fastapi import Request except ImportError: pytest.skip("MCP discoverable endpoints not available") + global_mcp_server_manager.registry.clear() + mock_request = MagicMock(spec=Request) mock_request.base_url = "https://proxy.litellm.example/" mock_request.headers = {} @@ -410,7 +438,9 @@ async def test_register_client_remote_registration_success(): @pytest.mark.asyncio -async def test_authorize_endpoint_respects_x_forwarded_proto(): +async def test_authorize_endpoint_respects_x_forwarded_proto( + trusted_proxy_origin_headers, +): """Test that authorize endpoint uses X-Forwarded-Proto header to construct correct redirect_uri""" try: from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( @@ -449,6 +479,7 @@ async def test_authorize_endpoint_respects_x_forwarded_proto(): mock_request = MagicMock(spec=Request) mock_request.base_url = "http://litellm.example.com/" # HTTP mock_request.headers = {"X-Forwarded-Proto": "https"} # Behind HTTPS proxy + set_request_from_trusted_proxy(mock_request) # Mock the encryption functions with patch( @@ -461,7 +492,7 @@ async def test_authorize_endpoint_respects_x_forwarded_proto(): request=mock_request, client_id="test_client_id", mcp_server_name="test_oauth", - redirect_uri="https://client.example.com/callback", + redirect_uri="http://127.0.0.1:60108/callback", state="test_state", ) @@ -476,7 +507,9 @@ async def test_authorize_endpoint_respects_x_forwarded_proto(): @pytest.mark.asyncio -async def test_token_endpoint_respects_x_forwarded_proto(): +async def test_token_endpoint_respects_x_forwarded_proto( + trusted_proxy_origin_headers, +): """Test that token endpoint uses X-Forwarded-Proto header for redirect_uri""" try: from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( @@ -515,6 +548,7 @@ async def test_token_endpoint_respects_x_forwarded_proto(): mock_request = MagicMock(spec=Request) mock_request.base_url = "http://litellm-proxy.example.com/" # HTTP mock_request.headers = {"X-Forwarded-Proto": "https"} # Behind HTTPS proxy + set_request_from_trusted_proxy(mock_request) # Mock httpx client response mock_response = MagicMock() @@ -535,7 +569,7 @@ async def test_token_endpoint_respects_x_forwarded_proto(): mock_get_client.return_value = mock_async_client # Call token endpoint - response = await token_endpoint( + await token_endpoint( request=mock_request, grant_type="authorization_code", code="test_code", @@ -666,7 +700,9 @@ async def test_oauth_protected_resource_legacy_pattern(): @pytest.mark.asyncio -async def test_oauth_protected_resource_respects_x_forwarded_proto(): +async def test_oauth_protected_resource_respects_x_forwarded_proto( + trusted_proxy_origin_headers, +): """Test that oauth_protected_resource_mcp uses X-Forwarded-Proto for URLs""" try: from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( @@ -704,6 +740,7 @@ async def test_oauth_protected_resource_respects_x_forwarded_proto(): mock_request = MagicMock(spec=Request) mock_request.base_url = "http://litellm.example.com/" # HTTP mock_request.headers = {"X-Forwarded-Proto": "https"} # Behind HTTPS proxy + set_request_from_trusted_proxy(mock_request) # Call the endpoint response = await oauth_protected_resource_mcp( @@ -719,7 +756,9 @@ async def test_oauth_protected_resource_respects_x_forwarded_proto(): @pytest.mark.asyncio -async def test_oauth_authorization_server_respects_x_forwarded_proto(): +async def test_oauth_authorization_server_respects_x_forwarded_proto( + trusted_proxy_origin_headers, +): """Test that oauth_authorization_server_mcp uses X-Forwarded-Proto for URLs""" try: from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( @@ -757,6 +796,7 @@ async def test_oauth_authorization_server_respects_x_forwarded_proto(): mock_request = MagicMock(spec=Request) mock_request.base_url = "http://litellm.example.com/" # HTTP mock_request.headers = {"X-Forwarded-Proto": "https"} # Behind HTTPS proxy + set_request_from_trusted_proxy(mock_request) # Call the endpoint response = await oauth_authorization_server_mcp( @@ -773,20 +813,28 @@ async def test_oauth_authorization_server_respects_x_forwarded_proto(): @pytest.mark.asyncio -async def test_register_client_respects_x_forwarded_proto(): +async def test_register_client_respects_x_forwarded_proto( + trusted_proxy_origin_headers, +): """Test that register_client uses X-Forwarded-Proto for redirect_uris""" try: from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( register_client, ) + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + global_mcp_server_manager, + ) from fastapi import Request except ImportError: pytest.skip("MCP discoverable endpoints not available") + global_mcp_server_manager.registry.clear() + # Mock request with http base_url but X-Forwarded-Proto: https mock_request = MagicMock(spec=Request) mock_request.base_url = "http://proxy.litellm.example/" # HTTP mock_request.headers = {"X-Forwarded-Proto": "https"} # Behind HTTPS proxy + set_request_from_trusted_proxy(mock_request) with patch( "litellm.proxy._experimental.mcp_server.discoverable_endpoints._read_request_body", @@ -803,7 +851,9 @@ async def test_register_client_respects_x_forwarded_proto(): @pytest.mark.asyncio -async def test_authorize_endpoint_respects_x_forwarded_host(): +async def test_authorize_endpoint_respects_x_forwarded_host( + trusted_proxy_origin_headers, +): """Test that authorize endpoint uses X-Forwarded-Host and X-Forwarded-Proto to construct correct redirect_uri""" try: from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( @@ -847,6 +897,7 @@ async def test_authorize_endpoint_respects_x_forwarded_host(): "X-Forwarded-Proto": "https", "X-Forwarded-Host": "proxy.example.com", } + set_request_from_trusted_proxy(mock_request) # Mock the encryption functions with patch( @@ -859,7 +910,7 @@ async def test_authorize_endpoint_respects_x_forwarded_host(): request=mock_request, client_id="test_client_id", mcp_server_name="test_oauth", - redirect_uri="https://client.example.com/callback", + redirect_uri="http://127.0.0.1:60108/callback", state="test_state", ) @@ -875,7 +926,9 @@ async def test_authorize_endpoint_respects_x_forwarded_host(): @pytest.mark.asyncio -async def test_token_endpoint_respects_x_forwarded_host(): +async def test_token_endpoint_respects_x_forwarded_host( + trusted_proxy_origin_headers, +): """Test that token endpoint uses X-Forwarded-Host and X-Forwarded-Proto for redirect_uri""" try: from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( @@ -917,6 +970,7 @@ async def test_token_endpoint_respects_x_forwarded_host(): "X-Forwarded-Proto": "https", "X-Forwarded-Host": "proxy.example.com", } + set_request_from_trusted_proxy(mock_request) # Mock httpx client response mock_response = MagicMock() @@ -937,7 +991,7 @@ async def test_token_endpoint_respects_x_forwarded_host(): mock_get_client.return_value = mock_async_client # Call token endpoint - response = await token_endpoint( + await token_endpoint( request=mock_request, grant_type="authorization_code", code="test_code", @@ -1075,7 +1129,12 @@ async def test_token_endpoint_respects_x_forwarded_host(): ], ) def test_get_request_base_url_comprehensive( - base_url, x_forwarded_proto, x_forwarded_host, x_forwarded_port, expected_url + base_url, + x_forwarded_proto, + x_forwarded_host, + x_forwarded_port, + expected_url, + trusted_proxy_origin_headers, ): """Comprehensive test for get_request_base_url with various header combinations""" try: @@ -1089,6 +1148,7 @@ def test_get_request_base_url_comprehensive( # Create mock request mock_request = MagicMock(spec=Request) mock_request.base_url = base_url + set_request_from_trusted_proxy(mock_request) # Build headers dict headers = {} @@ -1116,3 +1176,93 @@ def test_get_request_base_url_comprehensive( f"X-Forwarded-Host={x_forwarded_host}, " f"X-Forwarded-Port={x_forwarded_port}" ) + + +def test_get_request_base_url_ignores_forwarded_headers_from_untrusted_client(): + try: + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + get_request_base_url, + ) + from fastapi import Request + except ImportError: + pytest.skip("MCP discoverable endpoints not available") + + mock_request = MagicMock(spec=Request) + mock_request.base_url = "https://gateway.example.com/mcp" + mock_request.headers = { + "X-Forwarded-Proto": "https", + "X-Forwarded-Host": "attacker.example.com", + "X-Forwarded-Port": "443", + } + mock_request.client = MagicMock() + mock_request.client.host = "203.0.113.10" + + with patch( + "litellm.proxy.proxy_server.general_settings", + { + "use_x_forwarded_for": True, + "mcp_trusted_proxy_ranges": TRUSTED_PROXY_RANGES, + }, + create=True, + ): + assert get_request_base_url(mock_request) == "https://gateway.example.com/mcp" + + +def test_validate_trusted_redirect_uri_rejects_spoofed_forwarded_host(): + try: + from litellm.proxy._experimental.mcp_server.oauth_utils import ( + validate_trusted_redirect_uri, + ) + from fastapi import Request + except ImportError: + pytest.skip("MCP OAuth utilities not available") + + mock_request = MagicMock(spec=Request) + mock_request.base_url = "https://gateway.example.com/" + mock_request.headers = { + "X-Forwarded-Proto": "https", + "X-Forwarded-Host": "attacker.example.com", + } + mock_request.client = MagicMock() + mock_request.client.host = "203.0.113.10" + + with ( + patch( + "litellm.proxy.proxy_server.general_settings", + { + "use_x_forwarded_for": True, + "mcp_trusted_proxy_ranges": TRUSTED_PROXY_RANGES, + }, + create=True, + ), + pytest.raises(HTTPException), + ): + validate_trusted_redirect_uri( + mock_request, + "https://attacker.example.com/callback", + ) + + +def test_validate_trusted_redirect_uri_allows_forwarded_origin_from_trusted_proxy( + trusted_proxy_origin_headers, +): + try: + from litellm.proxy._experimental.mcp_server.oauth_utils import ( + validate_trusted_redirect_uri, + ) + from fastapi import Request + except ImportError: + pytest.skip("MCP OAuth utilities not available") + + mock_request = MagicMock(spec=Request) + mock_request.base_url = "http://localhost:4000/" + mock_request.headers = { + "X-Forwarded-Proto": "https", + "X-Forwarded-Host": "proxy.example.com", + } + set_request_from_trusted_proxy(mock_request) + + validate_trusted_redirect_uri( + mock_request, + "https://proxy.example.com/callback", + ) diff --git a/tests/local_testing/conftest.py b/tests/local_testing/conftest.py index abb871789c3..0831313c136 100644 --- a/tests/local_testing/conftest.py +++ b/tests/local_testing/conftest.py @@ -22,13 +22,13 @@ sys.path.insert( ) # Adds the parent directory to the system path import litellm -# ``litellm.model_cost`` is loaded at import time from the URL pinned to -# ``main`` (``LITELLM_MODEL_COST_MAP_URL``). The in-tree backup ships with -# this branch and can include pricing entries that main has not yet picked -# up (e.g. an upstream provider rotates a model id and the test cassette -# records the new name). Backfill any entries that are missing from the -# remote-fetched map so cost-calculator lookups in tests succeed against -# the cassette state the branch is being tested with. +# ``litellm.model_cost`` is loaded at import time from the URL pinned to ``main`` +# (``LITELLM_MODEL_COST_MAP_URL``). The in-tree backup ships with this branch +# and can include pricing entries that ``main`` has not yet picked up (e.g. +# Mistral now returns ``ministral-8b-2512`` from ``mistral-tiny`` and the entry +# was added on this branch). Backfill any entries that are missing from the +# remote-fetched map so cost-calculator lookups in tests succeed against the +# cassette state the branch is being tested with. from litellm.litellm_core_utils.get_model_cost_map import GetModelCostMap for _k, _v in GetModelCostMap.load_local_model_cost_map().items(): diff --git a/tests/mcp_tests/test_mcp_server.py b/tests/mcp_tests/test_mcp_server.py index e65b45fb38b..eea2f2721ab 100644 --- a/tests/mcp_tests/test_mcp_server.py +++ b/tests/mcp_tests/test_mcp_server.py @@ -514,6 +514,69 @@ async def test_sse_mcp_handler_mock(): ) +@pytest.mark.asyncio +async def test_sse_mcp_handler_propagates_passthrough_401(): + """SSE handler must raise 401 + WWW-Authenticate when the upstream + pass-through probe rejects the client's bearer token, instead of letting + the SSE session start and silently return empty tool lists.""" + from fastapi import HTTPException + + from litellm.proxy._types import UserAPIKeyAuth + + mock_scope = { + "type": "http", + "method": "GET", + "path": "/mcp/sse", + "headers": [(b"accept", b"text/event-stream")], + "query_string": b"", + "server": ("localhost", 8000), + "scheme": "http", + } + mock_receive = AsyncMock() + mock_send = AsyncMock() + + mock_auth_result = (UserAPIKeyAuth(), None, None, {}, {}, []) + + challenge = HTTPException( + status_code=401, + detail="Unauthorized", + headers={"WWW-Authenticate": "Bearer authorization_uri=https://example/"}, + ) + + with ( + patch( + "litellm.proxy._experimental.mcp_server.server._SESSION_MANAGERS_INITIALIZED", + True, + ), + patch( + "litellm.proxy._experimental.mcp_server.server.sse_session_manager", + AsyncMock(), + ), + patch( + "litellm.proxy._experimental.mcp_server.server.extract_mcp_auth_context", + new=AsyncMock(return_value=mock_auth_result), + ), + patch( + "litellm.proxy._experimental.mcp_server.server.set_auth_context", + ), + patch( + "litellm.proxy._experimental.mcp_server.server._raise_preemptive_401_for_unauthenticated_servers", + new=AsyncMock(), + ), + patch( + "litellm.proxy._experimental.mcp_server.server._check_passthrough_upstream_auth", + new=AsyncMock(side_effect=challenge), + ), + ): + from litellm.proxy._experimental.mcp_server.server import handle_sse_mcp + + with pytest.raises(HTTPException) as excinfo: + await handle_sse_mcp(mock_scope, mock_receive, mock_send) + + assert excinfo.value.status_code == 401 + assert excinfo.value.headers and "WWW-Authenticate" in excinfo.value.headers + + def test_generate_stable_server_id(): """ Test the _generate_stable_server_id method to ensure hash stability across releases. diff --git a/tests/mcp_tests/test_per_user_oauth_cache.py b/tests/mcp_tests/test_per_user_oauth_cache.py index 43e514b32ae..141b906fce9 100644 --- a/tests/mcp_tests/test_per_user_oauth_cache.py +++ b/tests/mcp_tests/test_per_user_oauth_cache.py @@ -183,6 +183,31 @@ class TestValidateTokenResponse: server_id="atlassian", ) + def test_boolean_value_matches_lowercase_string_rule(self): + """Boolean ``True`` in token response must match the JSON-style rule ``"true"``. + + Admin config is typically written as ``{"verified": "true"}`` (lower-case + from JSON / YAML), but the OAuth response returns ``{"verified": true}`` + (Python ``True``). The normaliser must align them. + """ + _validate_token_response = _import_validate() + token_response = {"access_token": "tok", "verified": True} + # Should not raise + _validate_token_response( + token_response=token_response, + validation_rules={"verified": "true"}, + server_id="test", + ) + + def test_boolean_false_matches_lowercase_string_rule(self): + _validate_token_response = _import_validate() + token_response = {"access_token": "tok", "is_admin": False} + _validate_token_response( + token_response=token_response, + validation_rules={"is_admin": "false"}, + server_id="test", + ) + # ── _compute_per_user_token_ttl ────────────────────────────────────────────── diff --git a/tests/proxy_unit_tests/test_jwt.py b/tests/proxy_unit_tests/test_jwt.py index 9a8d6d37020..92209e11315 100644 --- a/tests/proxy_unit_tests/test_jwt.py +++ b/tests/proxy_unit_tests/test_jwt.py @@ -2,6 +2,8 @@ # Unit tests for JWT-Auth import asyncio +import base64 +import logging import os import random import sys @@ -21,6 +23,9 @@ from datetime import datetime, timedelta from unittest.mock import AsyncMock, MagicMock, patch import pytest +import jwt +from cryptography.hazmat.primitives import serialization +from cryptography.hazmat.primitives.asymmetric import rsa from fastapi import Request, HTTPException from fastapi.routing import APIRoute from fastapi.responses import Response @@ -35,7 +40,7 @@ from litellm.proxy._types import ( from litellm.proxy.auth.handle_jwt import JWTHandler, JWTAuthManager from litellm.proxy.management_endpoints.team_endpoints import new_team from litellm.proxy.proxy_server import chat_completion -from typing import Literal +from typing import Literal, Optional public_key = { "kty": "RSA", @@ -1584,3 +1589,524 @@ async def test_auth_jwt_mismatched_key_fails(monkeypatch): with pytest.raises(Exception) as exc: await h.auth_jwt(token) assert "Validation fails" in str(exc.value) + + +def _base64url_encode_bytes(value: bytes) -> str: + return base64.urlsafe_b64encode(value).rstrip(b"=").decode() + + +def _base64url_encode_int(value: int) -> str: + value_bytes = value.to_bytes((value.bit_length() + 7) // 8, "big") + return _base64url_encode_bytes(value=value_bytes) + + +def _get_rsa_key_and_jwk(kid: str): + private_key = rsa.generate_private_key(public_exponent=65537, key_size=2048) + public_numbers = private_key.public_key().public_numbers() + jwk = { + "kty": "RSA", + "n": _base64url_encode_int(value=public_numbers.n), + "e": _base64url_encode_int(value=public_numbers.e), + "kid": kid, + "alg": "RS256", + "use": "sig", + } + return private_key, jwk + + +def _encode_rsa_jwt( + private_key, + issuer: str, + audience: str, + kid: str, + extra_claims: Optional[dict] = None, +) -> str: + private_key_pem = private_key.private_bytes( + encoding=serialization.Encoding.PEM, + format=serialization.PrivateFormat.PKCS8, + encryption_algorithm=serialization.NoEncryption(), + ) + current_time = int(time.time()) + claims = { + "sub": "test-subject", + "iss": issuer, + "aud": audience, + "iat": current_time, + "exp": current_time + 300, + } + if extra_claims: + claims.update(extra_claims) + + return jwt.encode( + claims, + private_key_pem, + algorithm="RS256", + headers={"kid": kid}, + ) + + +def _get_jwt_handler_with_issuer_keys(issuers: list, keys_by_url: dict) -> JWTHandler: + cache = DualCache() + for jwks_url, keys in keys_by_url.items(): + cache.set_cache( + key=f"litellm_jwt_auth_keys_{jwks_url}", + value=keys, + ) + + jwt_handler = JWTHandler() + jwt_handler.update_environment( + prisma_client=None, + user_api_key_cache=cache, + litellm_jwtauth=LiteLLM_JWTAuth(issuers=issuers), + ) + return jwt_handler + + +@pytest.mark.asyncio +async def test_multi_issuer_jwt_validates_selected_issuer_and_maps_claims( + monkeypatch, +): + monkeypatch.delenv("JWT_AUDIENCE", raising=False) + monkeypatch.delenv("JWT_PUBLIC_KEY_URL", raising=False) + + issuer_one = "https://issuer-one.example.com" + issuer_two = "https://issuer-two.example.com" + issuer_one_jwks_url = f"{issuer_one}/keys" + issuer_two_jwks_url = f"{issuer_two}/keys" + shared_kid = "shared-kid" + + _, issuer_one_jwk = _get_rsa_key_and_jwk(kid=shared_kid) + issuer_two_private_key, issuer_two_jwk = _get_rsa_key_and_jwk(kid=shared_kid) + + jwt_handler = _get_jwt_handler_with_issuer_keys( + issuers=[ + { + "issuer": issuer_one, + "jwks_url": issuer_one_jwks_url, + "audience": "audience-one", + "user_id_jwt_field": "email", + "user_email_jwt_field": "email", + }, + { + "issuer": issuer_two, + "jwks_url": issuer_two_jwks_url, + "audience": "audience-two", + "user_id_jwt_field": "repository_owner", + "team_id_jwt_field": "repository", + }, + ], + keys_by_url={ + issuer_one_jwks_url: [issuer_one_jwk], + issuer_two_jwks_url: [issuer_two_jwk], + }, + ) + + token = _encode_rsa_jwt( + private_key=issuer_two_private_key, + issuer=issuer_two, + audience="audience-two", + kid=shared_kid, + extra_claims={ + "repository_owner": "example-org", + "repository": "example-org/litellm-fork", + }, + ) + + claims = await jwt_handler.auth_jwt(token=token) + + assert claims[JWTHandler.LITELLM_JWT_ISSUER_CLAIM] == issuer_two + assert jwt_handler.get_user_id(token=claims, default_value=None) == ("example-org") + assert jwt_handler.get_team_id(token=claims, default_value=None) == ( + "example-org/litellm-fork" + ) + + +@pytest.mark.asyncio +async def test_multi_issuer_jwt_maps_kubernetes_namespace_claim(monkeypatch): + monkeypatch.delenv("JWT_AUDIENCE", raising=False) + monkeypatch.delenv("JWT_PUBLIC_KEY_URL", raising=False) + + issuer = "https://oidc.eks.eu-west-1.amazonaws.com/id/test-cluster" + jwks_url = f"{issuer}/keys" + private_key, jwk = _get_rsa_key_and_jwk(kid="k8s-key") + jwt_handler = _get_jwt_handler_with_issuer_keys( + issuers=[ + { + "issuer": issuer, + "jwks_url": jwks_url, + "audience": None, + "disable_audience_validation": True, + "user_id_jwt_field": "kubernetes\\.io.namespace", + } + ], + keys_by_url={jwks_url: [jwk]}, + ) + token = _encode_rsa_jwt( + private_key=private_key, + issuer=issuer, + audience="kubernetes.default.svc", + kid="k8s-key", + extra_claims={"kubernetes.io": {"namespace": "example-namespace"}}, + ) + + claims = await jwt_handler.auth_jwt(token=token) + + assert ( + jwt_handler.get_user_id(token=claims, default_value=None) == "example-namespace" + ) + + +@pytest.mark.asyncio +async def test_multi_issuer_jwt_falls_back_to_global_jwks_for_unknown_issuer( + monkeypatch, +): + """Unknown ``iss`` claims fall through to the global ``JWT_PUBLIC_KEY_URL`` + path so adding the new ``issuers`` config to a live deployment doesn't + break tokens minted by issuers that still rely on the legacy global JWKS. + """ + monkeypatch.delenv("JWT_AUDIENCE", raising=False) + monkeypatch.delenv("JWT_PUBLIC_KEY_URL", raising=False) + + configured_issuer = "https://issuer.example.com" + unknown_issuer = "https://unknown-issuer.example.com" + global_jwks_url = "https://global.example.com/keys" + monkeypatch.setenv("JWT_PUBLIC_KEY_URL", global_jwks_url) + + configured_private_key, configured_jwk = _get_rsa_key_and_jwk(kid="configured-key") + unknown_private_key, unknown_jwk = _get_rsa_key_and_jwk(kid="global-key") + jwt_handler = _get_jwt_handler_with_issuer_keys( + issuers=[ + { + "issuer": configured_issuer, + "jwks_url": f"{configured_issuer}/keys", + "audience": "expected-audience", + } + ], + keys_by_url={ + f"{configured_issuer}/keys": [configured_jwk], + global_jwks_url: [unknown_jwk], + }, + ) + token = _encode_rsa_jwt( + private_key=unknown_private_key, + issuer=unknown_issuer, + audience="expected-audience", + kid="global-key", + ) + + claims = await jwt_handler.auth_jwt(token=token) + + assert claims["iss"] == unknown_issuer + + +@pytest.mark.asyncio +async def test_multi_issuer_jwt_unknown_issuer_without_global_jwks_rejected( + monkeypatch, +): + """When there is no ``JWT_PUBLIC_KEY_URL`` to fall back to, an unknown + ``iss`` claim still fails — the fallback path raises ``Missing JWT + Public Key URL`` rather than the legacy ``Unsupported JWT issuer``. + """ + monkeypatch.delenv("JWT_AUDIENCE", raising=False) + monkeypatch.delenv("JWT_PUBLIC_KEY_URL", raising=False) + + configured_issuer = "https://issuer.example.com" + private_key, jwk = _get_rsa_key_and_jwk(kid="issuer-key") + jwt_handler = _get_jwt_handler_with_issuer_keys( + issuers=[ + { + "issuer": configured_issuer, + "jwks_url": f"{configured_issuer}/keys", + "audience": "expected-audience", + } + ], + keys_by_url={f"{configured_issuer}/keys": [jwk]}, + ) + token = _encode_rsa_jwt( + private_key=private_key, + issuer="https://unknown-issuer.example.com", + audience="expected-audience", + kid="issuer-key", + ) + + with pytest.raises(Exception) as exc: + await jwt_handler.auth_jwt(token=token) + + assert "Missing JWT Public Key URL" in str(exc.value) + + +@pytest.mark.asyncio +async def test_multi_issuer_jwt_rejects_wrong_audience(monkeypatch): + monkeypatch.delenv("JWT_AUDIENCE", raising=False) + monkeypatch.delenv("JWT_PUBLIC_KEY_URL", raising=False) + + issuer = "https://issuer.example.com" + jwks_url = f"{issuer}/keys" + private_key, jwk = _get_rsa_key_and_jwk(kid="issuer-key") + jwt_handler = _get_jwt_handler_with_issuer_keys( + issuers=[ + { + "issuer": issuer, + "jwks_url": jwks_url, + "audience": "expected-audience", + } + ], + keys_by_url={jwks_url: [jwk]}, + ) + token = _encode_rsa_jwt( + private_key=private_key, + issuer=issuer, + audience="wrong-audience", + kid="issuer-key", + ) + + with pytest.raises(Exception) as exc: + await jwt_handler.auth_jwt(token=token) + + assert "Validation fails" in str(exc.value) + + +@pytest.mark.asyncio +async def test_multi_issuer_jwt_same_kid_does_not_cross_issuer_keys(monkeypatch): + monkeypatch.delenv("JWT_AUDIENCE", raising=False) + monkeypatch.delenv("JWT_PUBLIC_KEY_URL", raising=False) + + issuer_one = "https://issuer-one.example.com" + issuer_two = "https://issuer-two.example.com" + issuer_one_jwks_url = f"{issuer_one}/keys" + issuer_two_jwks_url = f"{issuer_two}/keys" + shared_kid = "shared-kid" + issuer_one_private_key, issuer_one_jwk = _get_rsa_key_and_jwk(kid=shared_kid) + _, issuer_two_jwk = _get_rsa_key_and_jwk(kid=shared_kid) + jwt_handler = _get_jwt_handler_with_issuer_keys( + issuers=[ + { + "issuer": issuer_one, + "jwks_url": issuer_one_jwks_url, + "audience": "audience-one", + }, + { + "issuer": issuer_two, + "jwks_url": issuer_two_jwks_url, + "audience": "audience-two", + }, + ], + keys_by_url={ + issuer_one_jwks_url: [issuer_one_jwk], + issuer_two_jwks_url: [issuer_two_jwk], + }, + ) + token = _encode_rsa_jwt( + private_key=issuer_one_private_key, + issuer=issuer_two, + audience="audience-two", + kid=shared_kid, + ) + + with pytest.raises(Exception) as exc: + await jwt_handler.auth_jwt(token=token) + + assert "Validation fails" in str(exc.value) + + +@pytest.mark.asyncio +async def test_multi_issuer_jwt_missing_mapped_claim_is_optional(monkeypatch): + """Configured issuer claim mappings are advisory, not mandatory. + + When the token simply omits a mapped field (e.g. a service-to-service token + with no ``email`` claim), JWT auth still succeeds and the normalized claim + is just absent — matching the global ``litellm_jwtauth`` behaviour. + """ + monkeypatch.delenv("JWT_AUDIENCE", raising=False) + monkeypatch.delenv("JWT_PUBLIC_KEY_URL", raising=False) + + issuer = "https://issuer.example.com" + jwks_url = f"{issuer}/keys" + private_key, jwk = _get_rsa_key_and_jwk(kid="issuer-key") + jwt_handler = _get_jwt_handler_with_issuer_keys( + issuers=[ + { + "issuer": issuer, + "jwks_url": jwks_url, + "audience": "expected-audience", + "user_id_jwt_field": "email", + } + ], + keys_by_url={jwks_url: [jwk]}, + ) + token = _encode_rsa_jwt( + private_key=private_key, + issuer=issuer, + audience="expected-audience", + kid="issuer-key", + ) + + claims = await jwt_handler.auth_jwt(token=token) + + assert claims[JWTHandler.LITELLM_JWT_ISSUER_CLAIM] == issuer + assert JWTHandler.LITELLM_USER_ID_CLAIM not in claims + + +def test_multi_issuer_jwt_requires_audience_unless_explicitly_disabled( + monkeypatch, +): + monkeypatch.delenv("JWT_AUDIENCE", raising=False) + monkeypatch.delenv("JWT_PUBLIC_KEY_URL", raising=False) + + issuer = "https://issuer.example.com" + jwks_url = f"{issuer}/keys" + + with pytest.raises(Exception) as exc: + LiteLLM_JWTAuth( + issuers=[ + { + "issuer": issuer, + "jwks_url": jwks_url, + } + ] + ) + + assert "must configure audience" in str(exc.value) + + +@pytest.mark.asyncio +async def test_global_jwt_ignores_user_supplied_internal_claims(monkeypatch): + monkeypatch.delenv("JWT_AUDIENCE", raising=False) + monkeypatch.delenv("JWT_ISSUER", raising=False) + + jwks_url = "https://global-issuer.example.com/keys" + monkeypatch.setenv("JWT_PUBLIC_KEY_URL", jwks_url) + + private_key, jwk = _get_rsa_key_and_jwk(kid="global-key") + cache = DualCache() + cache.set_cache(key=f"litellm_jwt_auth_keys_{jwks_url}", value=[jwk]) + + jwt_handler = JWTHandler() + jwt_handler.update_environment( + prisma_client=None, + user_api_key_cache=cache, + litellm_jwtauth=LiteLLM_JWTAuth( + user_id_jwt_field="email", + user_email_jwt_field="email", + team_id_jwt_field="team.id", + team_ids_jwt_field="teams", + org_id_jwt_field="org.id", + end_user_id_jwt_field="end_user.id", + ), + ) + token = _encode_rsa_jwt( + private_key=private_key, + issuer="https://global-issuer.example.com", + audience="some-other-client", + kid="global-key", + extra_claims={ + "email": "real-user@example.com", + "team": {"id": "real-team"}, + "teams": ["real-team", "secondary-team"], + "org": {"id": "real-org"}, + "end_user": {"id": "real-end-user"}, + JWTHandler.LITELLM_JWT_ISSUER_CLAIM: "https://issuer.example.com", + JWTHandler.LITELLM_USER_ID_CLAIM: "victim-user", + JWTHandler.LITELLM_USER_EMAIL_CLAIM: "victim@example.com", + JWTHandler.LITELLM_TEAM_ID_CLAIM: "victim-team", + JWTHandler.LITELLM_TEAM_IDS_CLAIM: ["victim-team"], + JWTHandler.LITELLM_ORG_ID_CLAIM: "victim-org", + JWTHandler.LITELLM_END_USER_ID_CLAIM: "victim-end-user", + }, + ) + + claims = await jwt_handler.auth_jwt(token=token) + + assert jwt_handler.get_user_id(token=claims, default_value=None) == ( + "real-user@example.com" + ) + assert jwt_handler.get_user_email(token=claims, default_value=None) == ( + "real-user@example.com" + ) + assert jwt_handler.get_team_id(token=claims, default_value=None) == "real-team" + assert jwt_handler.get_team_ids_from_jwt(token=claims) == [ + "real-team", + "secondary-team", + ] + assert jwt_handler.get_org_id(token=claims, default_value=None) == "real-org" + assert jwt_handler.get_end_user_id(token=claims, default_value=None) == ( + "real-end-user" + ) + + +@pytest.mark.asyncio +async def test_multi_issuer_jwt_strips_unmapped_internal_claims(monkeypatch): + monkeypatch.delenv("JWT_AUDIENCE", raising=False) + monkeypatch.delenv("JWT_PUBLIC_KEY_URL", raising=False) + + issuer = "https://issuer.example.com" + jwks_url = f"{issuer}/keys" + private_key, jwk = _get_rsa_key_and_jwk(kid="issuer-key") + jwt_handler = _get_jwt_handler_with_issuer_keys( + issuers=[ + { + "issuer": issuer, + "jwks_url": jwks_url, + "audience": "expected-audience", + "user_email_jwt_field": "email", + } + ], + keys_by_url={jwks_url: [jwk]}, + ) + token = _encode_rsa_jwt( + private_key=private_key, + issuer=issuer, + audience="expected-audience", + kid="issuer-key", + extra_claims={ + "email": "real-user@example.com", + JWTHandler.LITELLM_USER_ID_CLAIM: "victim-user", + JWTHandler.LITELLM_TEAM_ID_CLAIM: "victim-team", + }, + ) + + claims = await jwt_handler.auth_jwt(token=token) + + assert JWTHandler.LITELLM_USER_ID_CLAIM not in claims + assert JWTHandler.LITELLM_TEAM_ID_CLAIM not in claims + assert jwt_handler.get_user_id(token=claims, default_value=None) is None + assert jwt_handler.get_team_id(token=claims, default_value=None) is None + assert jwt_handler.get_user_email(token=claims, default_value=None) == ( + "real-user@example.com" + ) + + +@pytest.mark.asyncio +async def test_multi_issuer_jwt_does_not_emit_unscoped_global_warning( + monkeypatch, caplog +): + monkeypatch.delenv("JWT_AUDIENCE", raising=False) + monkeypatch.delenv("JWT_ISSUER", raising=False) + monkeypatch.delenv("JWT_PUBLIC_KEY_URL", raising=False) + JWTHandler._unscoped_jwt_warning_emitted = False + + issuer = "https://issuer.example.com" + jwks_url = f"{issuer}/keys" + private_key, jwk = _get_rsa_key_and_jwk(kid="issuer-key") + jwt_handler = _get_jwt_handler_with_issuer_keys( + issuers=[ + { + "issuer": issuer, + "jwks_url": jwks_url, + "audience": "expected-audience", + } + ], + keys_by_url={jwks_url: [jwk]}, + ) + token = _encode_rsa_jwt( + private_key=private_key, + issuer=issuer, + audience="expected-audience", + kid="issuer-key", + ) + + with caplog.at_level(logging.WARNING): + await jwt_handler.auth_jwt(token=token) + + assert "Tokens minted by any application" not in caplog.text + assert JWTHandler._unscoped_jwt_warning_emitted is False diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py b/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py index 95c826daa8e..7753378ab4f 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py @@ -1,12 +1,9 @@ import json import os import sys -from unittest import mock -from unittest.mock import AsyncMock, MagicMock, patch +from unittest.mock import AsyncMock, MagicMock, call as mock_call, patch -import orjson import pytest -from fastapi import FastAPI, Request from fastapi.testclient import TestClient sys.path.insert( @@ -19,7 +16,6 @@ from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import ( MCPRequestHandler, ) from litellm.proxy._types import SpecialHeaders, UserAPIKeyAuth -from litellm.proxy.auth.user_api_key_auth import user_api_key_auth @pytest.mark.asyncio @@ -453,7 +449,7 @@ class TestMCPRequestHandler: with patch( "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth", side_effect=mock_user_api_key_auth, - ) as mock_auth: + ): # Call the method ( auth_result, @@ -998,6 +994,284 @@ class TestMCPPublicRouteGuard: assert isinstance(auth_result, UserAPIKeyAuth) +@pytest.mark.asyncio +class TestMCPPassthroughColdStartAdmission: + @staticmethod + def _make_passthrough_server(): + server = MagicMock() + server.is_oauth_passthrough = True + return server + + async def test_cold_start_ignores_header_without_path_target(self): + from fastapi import HTTPException + + scope = { + "type": "http", + "method": "POST", + "path": "/mcp", + "headers": [(b"x-mcp-servers", b"passthrough_server")], + } + + async def mock_user_api_key_auth_fails(api_key, request): + raise HTTPException(status_code=401, detail="Invalid API key") + + with ( + patch( + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth", + side_effect=mock_user_api_key_auth_fails, + ), + patch( + "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager" + ) as mock_mgr, + patch( + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp._is_mcp_passthrough_cold_start" + ) as mock_cold_start, + ): + mock_mgr.get_mcp_server_by_name.return_value = ( + TestMCPPassthroughColdStartAdmission._make_passthrough_server() + ) + with pytest.raises(HTTPException) as exc_info: + await MCPRequestHandler.process_mcp_request(scope) + + assert exc_info.value.status_code == 401 + # Cold-start admission must not fire for the aggregate ``/mcp`` + # route — only path-targeted routes are eligible for OAuth + # discovery admission. + mock_cold_start.assert_not_called() + + async def test_cold_start_rejects_server_specific_authorization_header(self): + from fastapi import HTTPException + + scope = { + "type": "http", + "method": "POST", + "path": "/mcp/passthrough_server", + "headers": [ + ( + b"x-mcp-passthrough_server-authorization", + b"Bearer upstream-token", + ) + ], + } + + async def mock_user_api_key_auth_fails(api_key, request): + raise HTTPException(status_code=401, detail="Invalid API key") + + with ( + patch( + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth", + side_effect=mock_user_api_key_auth_fails, + ), + patch( + "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager" + ) as mock_mgr, + ): + mock_mgr.get_mcp_server_by_name.return_value = ( + TestMCPPassthroughColdStartAdmission._make_passthrough_server() + ) + with pytest.raises(HTTPException) as exc_info: + await MCPRequestHandler.process_mcp_request(scope) + + assert exc_info.value.status_code == 401 + + async def test_cold_start_rejects_legacy_mcp_auth_header(self): + from fastapi import HTTPException + + scope = { + "type": "http", + "method": "POST", + "path": "/mcp/passthrough_server", + "headers": [(b"x-mcp-auth", b"Bearer upstream-token")], + } + + async def mock_user_api_key_auth_fails(api_key, request): + raise HTTPException(status_code=401, detail="Invalid API key") + + with ( + patch( + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth", + side_effect=mock_user_api_key_auth_fails, + ), + patch( + "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager" + ) as mock_mgr, + ): + mock_mgr.get_mcp_server_by_name.return_value = ( + TestMCPPassthroughColdStartAdmission._make_passthrough_server() + ) + with pytest.raises(HTTPException) as exc_info: + await MCPRequestHandler.process_mcp_request(scope) + + assert exc_info.value.status_code == 401 + + async def test_cold_start_fails_closed_when_client_ip_hides_server(self): + from fastapi import HTTPException + + scope = { + "type": "http", + "method": "POST", + "path": "/mcp/passthrough_server", + "headers": [], + } + + async def mock_user_api_key_auth_fails(api_key, request): + raise HTTPException(status_code=401, detail="Invalid API key") + + with ( + patch( + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth", + side_effect=mock_user_api_key_auth_fails, + ), + patch( + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.IPAddressUtils.get_mcp_client_ip", + return_value="203.0.113.10", + ), + patch( + "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager" + ) as mock_mgr, + ): + mock_mgr.get_mcp_server_by_name.return_value = None + with pytest.raises(HTTPException) as exc_info: + await MCPRequestHandler.process_mcp_request(scope) + + assert exc_info.value.status_code == 401 + mock_mgr.get_mcp_server_by_name.assert_any_call( + "passthrough_server", client_ip="203.0.113.10" + ) + + async def test_cold_start_propagates_non_401_http_error(self): + from fastapi import HTTPException + + scope = { + "type": "http", + "method": "POST", + "path": "/mcp/passthrough_server", + "headers": [], + } + + async def mock_user_api_key_auth_forbidden(api_key, request): + raise HTTPException(status_code=403, detail="Forbidden") + + with ( + patch( + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth", + side_effect=mock_user_api_key_auth_forbidden, + ), + patch( + "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager" + ) as mock_mgr, + ): + mock_mgr.get_mcp_server_by_name.return_value = ( + TestMCPPassthroughColdStartAdmission._make_passthrough_server() + ) + with pytest.raises(HTTPException) as exc_info: + await MCPRequestHandler.process_mcp_request(scope) + + assert exc_info.value.status_code == 403 + + async def test_cold_start_propagates_non_auth_proxy_exception(self): + from litellm.proxy._types import ProxyException + + scope = { + "type": "http", + "method": "POST", + "path": "/mcp/passthrough_server", + "headers": [], + } + + async def mock_user_api_key_auth_server_error(api_key, request): + raise ProxyException( + message="Internal error", + type="server_error", + param=None, + code=500, + ) + + with ( + patch( + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth", + side_effect=mock_user_api_key_auth_server_error, + ), + patch( + "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager" + ) as mock_mgr, + ): + mock_mgr.get_mcp_server_by_name.return_value = ( + TestMCPPassthroughColdStartAdmission._make_passthrough_server() + ) + with pytest.raises(ProxyException): + await MCPRequestHandler.process_mcp_request(scope) + + async def test_cold_start_allows_401_for_path_passthrough_target(self): + from fastapi import HTTPException + + scope = { + "type": "http", + "method": "POST", + "path": "/mcp/passthrough_server", + "headers": [], + } + + async def mock_user_api_key_auth_fails(api_key, request): + raise HTTPException(status_code=401, detail="Invalid API key") + + with ( + patch( + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth", + side_effect=mock_user_api_key_auth_fails, + ), + patch( + "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager" + ) as mock_mgr, + ): + mock_mgr.get_mcp_server_by_name.return_value = ( + TestMCPPassthroughColdStartAdmission._make_passthrough_server() + ) + auth_result, *_rest = await MCPRequestHandler.process_mcp_request(scope) + + assert isinstance(auth_result, UserAPIKeyAuth) + mock_mgr.get_mcp_server_by_name.assert_any_call( + "passthrough_server", client_ip="" + ) + + async def test_cold_start_allows_proxy_exception_401_for_path_target(self): + from litellm.proxy._types import ProxyException + + scope = { + "type": "http", + "method": "POST", + "path": "/mcp/passthrough_server", + "headers": [], + } + + async def mock_user_api_key_auth_fails(api_key, request): + raise ProxyException( + message="Authentication Error", + type="auth_error", + param="api_key", + code=401, + ) + + with ( + patch( + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth", + side_effect=mock_user_api_key_auth_fails, + ), + patch( + "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager" + ) as mock_mgr, + ): + mock_mgr.get_mcp_server_by_name.return_value = ( + TestMCPPassthroughColdStartAdmission._make_passthrough_server() + ) + auth_result, *_rest = await MCPRequestHandler.process_mcp_request(scope) + + assert isinstance(auth_result, UserAPIKeyAuth) + mock_mgr.get_mcp_server_by_name.assert_any_call( + "passthrough_server", client_ip="" + ) + + @pytest.mark.asyncio class TestMCPOAuth2FallbackTargetGating: """ @@ -1009,9 +1283,14 @@ class TestMCPOAuth2FallbackTargetGating: """ @staticmethod - def _make_server(auth_type): + def _make_server(auth_type, is_oauth_passthrough=False): server = MagicMock() server.auth_type = auth_type + # MagicMock would otherwise auto-create truthy stand-ins for any + # attribute access (including ``is_oauth_passthrough``), which + # would silently flip the passthrough fallback gate on. Pin the + # boolean explicitly so non-passthrough fixtures stay non-passthrough. + server.is_oauth_passthrough = is_oauth_passthrough return server async def test_fallback_blocked_when_target_is_not_oauth2(self): @@ -1113,6 +1392,88 @@ class TestMCPOAuth2FallbackTargetGating: auth_result, *_rest = await MCPRequestHandler.process_mcp_request(scope) assert isinstance(auth_result, UserAPIKeyAuth) + async def test_fallback_allowed_when_target_is_passthrough(self): + """ + Cold-start return per RFC 9728 / MCP Authorization spec: client + discovered the upstream IdP via the gateway's protected-resource + metadata, completed OAuth, and is returning with + ``Authorization: Bearer ``. The bearer is not a + LiteLLM key but the target is a pass-through server, so admission + falls back to anonymous and forwards the bearer upstream. + """ + from fastapi import HTTPException + + from litellm.types.mcp import MCPAuth + + scope = { + "type": "http", + "method": "POST", + "path": "/mcp/passthrough_server", + "headers": [(b"authorization", b"Bearer upstream-token-xyz")], + } + + async def mock_user_api_key_auth_fails(api_key, request): + raise HTTPException(status_code=401, detail="Invalid API key") + + with ( + patch( + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth", + side_effect=mock_user_api_key_auth_fails, + ), + patch( + "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager" + ) as mock_mgr, + ): + mock_mgr.get_mcp_server_by_name.return_value = ( + TestMCPOAuth2FallbackTargetGating._make_server( + auth_type=MCPAuth.none, + is_oauth_passthrough=True, + ) + ) + auth_result, *_rest = await MCPRequestHandler.process_mcp_request(scope) + assert isinstance(auth_result, UserAPIKeyAuth) + assert auth_result.api_key is None + + async def test_fallback_blocked_when_client_ip_hides_oauth2_target(self): + from fastapi import HTTPException + + scope = { + "type": "http", + "method": "POST", + "path": "/mcp/hidden_oauth2_server", + "headers": [(b"authorization", b"Bearer upstream-token")], + } + + async def mock_user_api_key_auth_fails(api_key, request): + raise HTTPException(status_code=401, detail="Invalid API key") + + with ( + patch( + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.user_api_key_auth", + side_effect=mock_user_api_key_auth_fails, + ), + patch( + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp.IPAddressUtils.get_mcp_client_ip", + return_value="203.0.113.10", + ), + patch( + "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager" + ) as mock_mgr, + ): + mock_mgr.get_mcp_server_by_name.return_value = None + with pytest.raises(HTTPException) as exc_info: + await MCPRequestHandler.process_mcp_request(scope) + + assert exc_info.value.status_code == 401 + # Lookup may run twice — once for the oauth2-target fallback gate + # and once for the passthrough-target fallback gate. Both must + # resolve to ``None`` (hidden by client IP) so neither bypass + # opens. Use ``assert_any_call`` to assert the IP-scoped lookup + # happened without locking the count. + mock_mgr.get_mcp_server_by_name.assert_any_call( + "hidden_oauth2_server", client_ip="203.0.113.10" + ) + async def test_fallback_blocked_when_any_target_in_header_is_not_oauth2(self): """ x-mcp-servers can list multiple targets. If ANY of them is non-OAuth2, @@ -1241,6 +1602,39 @@ class TestMCPDelegateAuthToUpstream: is False ) + def test_build_mcp_server_table_preserves_oauth_passthrough(self): + """Registry → API list rows must expose oauth_passthrough for the UI. + + ``oauth_passthrough`` is the dedicated non-oauth2 pass-through opt-in, + distinct from ``delegate_auth_to_upstream`` (oauth2-only). Both must + round-trip independently so neither flag silently implies the other. + """ + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + MCPServerManager, + ) + from litellm.types.mcp import MCPAuth + from litellm.types.mcp_server.mcp_server_manager import MCPServer + + manager = MCPServerManager() + passthrough = MCPServer( + server_id="passthrough-1", + name="passthrough", + transport="http", + auth_type=MCPAuth.none, + extra_headers=["Authorization"], + oauth_passthrough=True, + available_on_public_internet=True, + ) + row = manager._build_mcp_server_table(passthrough) + assert row.oauth_passthrough is True + # The oauth2-only flag must remain independent and default off. + assert row.delegate_auth_to_upstream is False + + not_passthrough = passthrough.model_copy(update={"oauth_passthrough": False}) + assert ( + manager._build_mcp_server_table(not_passthrough).oauth_passthrough is False + ) + async def test_delegate_skips_litellm_auth_with_no_authorization(self): """ oauth2 + delegate_auth_to_upstream=True, no Authorization header at @@ -1806,9 +2200,12 @@ class TestMCPDelegateAuthToUpstream: delegate_auth_to_upstream=True, ) - def lookup_by_name(name): + def lookup_by_name(name, **_kwargs): # Only the *exact* delegated name resolves. Anything else (e.g. # ``delegated_server/extra``) returns None so the bypass fails. + # ``**_kwargs`` accepts the ``client_ip`` kwarg the cold-start + # admission path now forwards (real signature: + # ``get_mcp_server_by_name(name, client_ip=None)``). if name == "delegated_server": return delegate_server return None @@ -1869,7 +2266,10 @@ class TestMCPDelegateAuthToUpstream: auth_type=MCPAuth.api_key, ) - def lookup_by_name(name): + def lookup_by_name(name, **_kwargs): + # ``**_kwargs`` accepts the ``client_ip`` kwarg the cold-start + # admission path now forwards (real signature: + # ``get_mcp_server_by_name(name, client_ip=None)``). return { "delegated_server": delegate_server, "non_delegate_server": non_delegate, @@ -2342,7 +2742,6 @@ class TestMCPAccessGroupsE2E: mock_auth.assert_called_once() -@pytest.mark.asyncio def test_mcp_path_based_server_segregation(monkeypatch): # Import the MCP server FastAPI app and context getter from litellm.proxy._experimental.mcp_server.server import app, get_auth_context @@ -2956,6 +3355,89 @@ class TestAgentMCPPermissions: ) assert sorted(result) == ["tool_a", "tool_b"] + async def test_get_agent_object_permission_uses_shared_helper(self): + """``_get_agent_object_permission`` must resolve the agent's + ``object_permission_id`` and then defer to the shared + ``get_object_permission`` helper so cache entries are shared with the + org / team / key paths.""" + from litellm.caching.dual_cache import DualCache + + cache = DualCache() + agent_row = MagicMock() + agent_row.object_permission_id = "perm-xyz" + prisma_client = MagicMock() + prisma_client.db.litellm_agentstable.find_unique = AsyncMock( + return_value=agent_row + ) + user_api_key_auth = UserAPIKeyAuth( + api_key="test-key", + user_id="test-user", + agent_id="agent-shared", + ) + expected_perm = MagicMock() + + with ( + patch("litellm.proxy.proxy_server.prisma_client", prisma_client), + patch("litellm.proxy.proxy_server.user_api_key_cache", cache), + patch("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock()), + patch( + "litellm.proxy.auth.auth_checks.get_object_permission", + new_callable=AsyncMock, + return_value=expected_perm, + ) as mock_get_perm, + ): + result = await MCPRequestHandler._get_agent_object_permission( + user_api_key_auth + ) + assert result is expected_perm + mock_get_perm.assert_awaited_once() + assert mock_get_perm.await_args.kwargs["object_permission_id"] == "perm-xyz" + + # Second call: the agent_id -> object_permission_id mapping is + # cached, so the agent row is not re-fetched. + prisma_client.db.litellm_agentstable.find_unique.reset_mock() + await MCPRequestHandler._get_agent_object_permission(user_api_key_auth) + prisma_client.db.litellm_agentstable.find_unique.assert_not_called() + + async def test_get_agent_object_permission_caches_missing_permission(self): + """When the agent has no ``object_permission_id`` the sentinel must be + cached so subsequent requests do not hit the DB again.""" + from litellm.caching.dual_cache import DualCache + + cache = DualCache() + agent_row = MagicMock() + agent_row.object_permission_id = None + prisma_client = MagicMock() + prisma_client.db.litellm_agentstable.find_unique = AsyncMock( + return_value=agent_row + ) + user_api_key_auth = UserAPIKeyAuth( + api_key="test-key", + user_id="test-user", + agent_id="agent-no-perm", + ) + + with ( + patch("litellm.proxy.proxy_server.prisma_client", prisma_client), + patch("litellm.proxy.proxy_server.user_api_key_cache", cache), + patch("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock()), + patch( + "litellm.proxy.auth.auth_checks.get_object_permission", + new_callable=AsyncMock, + ) as mock_get_perm, + ): + assert ( + await MCPRequestHandler._get_agent_object_permission(user_api_key_auth) + is None + ) + assert ( + await MCPRequestHandler._get_agent_object_permission(user_api_key_auth) + is None + ) + + mock_get_perm.assert_not_awaited() + prisma_client.db.litellm_agentstable.find_unique.assert_awaited_once() + @pytest.mark.asyncio async def test_tool_permission_servers_included_in_allowed_servers(): diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py index c8789e0b0a6..da66d60aed8 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py @@ -1515,7 +1515,7 @@ async def test_oauth_protected_resource_returns_empty_scopes_when_none(): mock_request.headers = {} try: - response = _build_oauth_protected_resource_response( + response = await _build_oauth_protected_resource_response( request=mock_request, mcp_server_name="atlassian_mcp", use_standard_pattern=False, @@ -2005,7 +2005,7 @@ async def test_discovery_root_does_not_expose_private_server_for_external_client request=mock_request, mcp_server_name=None, ) - resource_response = _build_oauth_protected_resource_response( + resource_response = await _build_oauth_protected_resource_response( request=mock_request, mcp_server_name=None, use_standard_pattern=False, diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_oauth_passthrough.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_oauth_passthrough.py new file mode 100644 index 00000000000..ad78609ee18 --- /dev/null +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_oauth_passthrough.py @@ -0,0 +1,474 @@ +"""Unit tests for MCP OAuth passthrough metadata behavior. + +Covers: +- `MCPServer.is_oauth_passthrough` property semantics. +- `/.well-known/oauth-protected-resource/...` pass-through branch (proxies + upstream metadata, normalizes the `resource` field, caches, and surfaces + network errors as HTTP 502). +""" + +import asyncio +import sys +import time +from unittest.mock import AsyncMock, MagicMock, patch + +import httpx +import pytest +from fastapi import HTTPException, Request + +sys.path.insert(0, "../../../../../") + + +from litellm.proxy._experimental.mcp_server import discoverable_endpoints +from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + _OAUTH_METADATA_CACHE, + _OAUTH_METADATA_FETCH_LOCKS, + _build_oauth_protected_resource_response, +) +from litellm.proxy._types import MCPTransport +from litellm.types.mcp import MCPAuth +from litellm.types.mcp_server.mcp_server_manager import MCPServer + + +@pytest.fixture(autouse=True) +def _mock_mcp_client_ip(): + """Bypass IP-based access control in tests.""" + with patch( + "litellm.proxy._experimental.mcp_server.discoverable_endpoints" + ".IPAddressUtils.get_mcp_client_ip", + return_value=None, + ): + yield + + +@pytest.fixture(autouse=True) +def _clear_metadata_cache(): + """Prevent cross-test cache bleed for the oauth-protected-resource TTL cache.""" + _OAUTH_METADATA_CACHE.clear() + _OAUTH_METADATA_FETCH_LOCKS.clear() + yield + _OAUTH_METADATA_CACHE.clear() + _OAUTH_METADATA_FETCH_LOCKS.clear() + + +def _make_request(base_url: str = "https://gateway.example.com/") -> Request: + request = MagicMock(spec=Request) + request.base_url = base_url + request.headers = {} + return request + + +# -------------------------------------------------------------------------- +# is_oauth_passthrough property +# -------------------------------------------------------------------------- + + +def test_is_oauth_passthrough_true_when_none_auth_and_authorization_header(): + server = MCPServer( + server_id="s1", + name="s1", + transport=MCPTransport.http, + auth_type=MCPAuth.none, + extra_headers=["Authorization"], + oauth_passthrough=True, + ) + assert server.is_oauth_passthrough is True + + +def test_is_oauth_passthrough_true_when_auth_type_none_and_mixed_case_header(): + server = MCPServer( + server_id="s1", + name="s1", + transport=MCPTransport.http, + auth_type=None, + extra_headers=["authorization", "x-request-id"], + oauth_passthrough=True, + ) + assert server.is_oauth_passthrough is True + + +def test_is_oauth_passthrough_false_for_oauth2_server(): + server = MCPServer( + server_id="s1", + name="s1", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + extra_headers=["Authorization"], + oauth_passthrough=True, + ) + assert server.is_oauth_passthrough is False + + +def test_is_oauth_passthrough_false_without_authorization_header(): + server = MCPServer( + server_id="s1", + name="s1", + transport=MCPTransport.http, + auth_type=MCPAuth.none, + extra_headers=["x-api-key"], + oauth_passthrough=True, + ) + assert server.is_oauth_passthrough is False + + +def test_is_oauth_passthrough_false_without_extra_headers(): + server = MCPServer( + server_id="s1", + name="s1", + transport=MCPTransport.http, + auth_type=MCPAuth.none, + oauth_passthrough=True, + ) + assert server.is_oauth_passthrough is False + + +def test_is_oauth_passthrough_false_without_oauth_passthrough_flag(): + """The detection flag must be set explicitly. Without it, the legacy + behavior is preserved for servers that forward Authorization for + non-OAuth reasons (static bearer tokens, custom auth schemes).""" + server = MCPServer( + server_id="s1", + name="s1", + transport=MCPTransport.http, + auth_type=MCPAuth.none, + extra_headers=["Authorization"], + # oauth_passthrough defaults to False + ) + assert server.is_oauth_passthrough is False + + +def test_is_oauth_passthrough_false_when_oauth_passthrough_explicitly_false(): + server = MCPServer( + server_id="s1", + name="s1", + transport=MCPTransport.http, + auth_type=MCPAuth.none, + extra_headers=["Authorization"], + oauth_passthrough=False, + ) + assert server.is_oauth_passthrough is False + + +def test_is_oauth_passthrough_false_when_only_delegate_auth_to_upstream_set(): + """Regression guard: ``delegate_auth_to_upstream`` is the oauth2-only + PKCE-bypass flag and must NOT, on its own, turn a non-oauth2 server into + an OAuth pass-through server. Pass-through requires the dedicated + ``oauth_passthrough`` opt-in. This protects existing deployments that set + ``delegate_auth_to_upstream`` from silently gaining pass-through behavior. + """ + server = MCPServer( + server_id="s1", + name="s1", + transport=MCPTransport.http, + auth_type=MCPAuth.none, + extra_headers=["Authorization"], + delegate_auth_to_upstream=True, + # oauth_passthrough intentionally left at its default (False) + ) + assert server.is_oauth_passthrough is False + + +# -------------------------------------------------------------------------- +# _build_oauth_protected_resource_response: pass-through branch +# -------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_oauth_protected_resource_passthrough_proxies_upstream_metadata(): + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + global_mcp_server_manager, + ) + + global_mcp_server_manager.registry.clear() + passthrough_server = MCPServer( + server_id="passthrough-1", + name="sample_docs", + server_name="sample_docs", + alias="sample_docs", + url="https://upstream.example.com/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.none, + extra_headers=["Authorization"], + oauth_passthrough=True, + ) + global_mcp_server_manager.registry[passthrough_server.server_id] = ( + passthrough_server + ) + + upstream_payload = { + "resource": "https://upstream.example.com/mcp", + "authorization_servers": ["https://okta.example.com/oauth2/default"], + "scopes_supported": ["openid", "profile"], + "bearer_methods_supported": ["header"], + } + mock_response = MagicMock() + mock_response.status_code = 200 + mock_response.json.return_value = upstream_payload + mock_client = MagicMock() + mock_client.get = AsyncMock(return_value=mock_response) + + with patch.object( + discoverable_endpoints, "get_async_httpx_client", return_value=mock_client + ): + result = await _build_oauth_protected_resource_response( + request=_make_request(), + mcp_server_name="sample_docs", + use_standard_pattern=True, + ) + + assert result["authorization_servers"] == [ + "https://okta.example.com/oauth2/default" + ] + # resource is normalized to the gateway URL so bearers are sent back to us + assert result["resource"].endswith("/mcp/sample_docs") + assert result["scopes_supported"] == ["openid", "profile"] + + +@pytest.mark.asyncio +async def test_oauth_protected_resource_passthrough_cache_hit(): + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + global_mcp_server_manager, + ) + + global_mcp_server_manager.registry.clear() + passthrough_server = MCPServer( + server_id="passthrough-2", + name="sample_docs", + server_name="sample_docs", + alias="sample_docs", + url="https://upstream.example.com/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.none, + extra_headers=["Authorization"], + oauth_passthrough=True, + ) + global_mcp_server_manager.registry[passthrough_server.server_id] = ( + passthrough_server + ) + + mock_response = MagicMock() + mock_response.status_code = 200 + mock_response.json.return_value = { + "authorization_servers": ["https://okta.example.com"], + } + mock_client = MagicMock() + mock_client.get = AsyncMock(return_value=mock_response) + + with patch.object( + discoverable_endpoints, "get_async_httpx_client", return_value=mock_client + ): + await _build_oauth_protected_resource_response( + request=_make_request(), + mcp_server_name="sample_docs", + use_standard_pattern=True, + ) + await _build_oauth_protected_resource_response( + request=_make_request(), + mcp_server_name="sample_docs", + use_standard_pattern=True, + ) + + assert mock_client.get.await_count == 1 + + +def test_oauth_metadata_cache_prunes_to_max_size(): + now = 1_000_000.0 + max_size = discoverable_endpoints._OAUTH_METADATA_CACHE_MAX_SIZE + + for index in range(max_size + 10): + _OAUTH_METADATA_CACHE[(f"server-{index}", f"https://upstream/{index}")] = ( + now + index + 1, + {"index": index}, + ) + + discoverable_endpoints._prune_oauth_metadata_cache(now) + + assert len(_OAUTH_METADATA_CACHE) == max_size + assert ("server-0", "https://upstream/0") not in _OAUTH_METADATA_CACHE + assert ( + f"server-{max_size + 9}", + f"https://upstream/{max_size + 9}", + ) in _OAUTH_METADATA_CACHE + + +def test_oauth_metadata_fetch_locks_pruned_alongside_cache(): + now = 1_000_000.0 + cached_key = ("server-active", "https://upstream/active") + expired_key = ("server-expired", "https://upstream/expired") + orphan_key = ("server-orphan", "https://upstream/orphan") + + _OAUTH_METADATA_CACHE[cached_key] = (now + 100, {"index": 0}) + _OAUTH_METADATA_CACHE[expired_key] = (now - 1, {"index": 1}) + + _OAUTH_METADATA_FETCH_LOCKS[cached_key] = asyncio.Lock() + _OAUTH_METADATA_FETCH_LOCKS[expired_key] = asyncio.Lock() + _OAUTH_METADATA_FETCH_LOCKS[orphan_key] = asyncio.Lock() + + discoverable_endpoints._prune_oauth_metadata_cache(now) + + assert cached_key in _OAUTH_METADATA_FETCH_LOCKS + assert expired_key not in _OAUTH_METADATA_FETCH_LOCKS + assert orphan_key not in _OAUTH_METADATA_FETCH_LOCKS + + +@pytest.mark.asyncio +async def test_oauth_metadata_fetch_locks_held_lock_preserved_during_prune(): + held_key = ("server-busy", "https://upstream/busy") + held_lock = asyncio.Lock() + _OAUTH_METADATA_FETCH_LOCKS[held_key] = held_lock + + async with held_lock: + discoverable_endpoints._prune_oauth_metadata_cache(time.time()) + assert held_key in _OAUTH_METADATA_FETCH_LOCKS + + +@pytest.mark.asyncio +async def test_oauth_metadata_cache_expired_entry_is_refetched(): + passthrough_server = MCPServer( + server_id="expired-cache-server", + name="sample_docs", + server_name="sample_docs", + alias="sample_docs", + url="https://upstream.example.com/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.none, + extra_headers=["Authorization"], + oauth_passthrough=True, + ) + _OAUTH_METADATA_CACHE[(passthrough_server.server_id, passthrough_server.url)] = ( + 0, + {"authorization_servers": ["https://stale.example.com"]}, + ) + + mock_response = MagicMock() + mock_response.status_code = 200 + mock_response.json.return_value = { + "authorization_servers": ["https://fresh.example.com"], + } + mock_client = MagicMock() + mock_client.get = AsyncMock(return_value=mock_response) + + with patch.object( + discoverable_endpoints, "get_async_httpx_client", return_value=mock_client + ): + result = await discoverable_endpoints.fetch_upstream_oauth_protected_resource( + passthrough_server + ) + + assert result == {"authorization_servers": ["https://fresh.example.com"]} + assert mock_client.get.await_count == 1 + + +@pytest.mark.asyncio +async def test_oauth_protected_resource_passthrough_network_error_returns_502(): + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + global_mcp_server_manager, + ) + + global_mcp_server_manager.registry.clear() + passthrough_server = MCPServer( + server_id="passthrough-3", + name="sample_docs", + server_name="sample_docs", + alias="sample_docs", + url="https://upstream.example.com/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.none, + extra_headers=["Authorization"], + oauth_passthrough=True, + ) + global_mcp_server_manager.registry[passthrough_server.server_id] = ( + passthrough_server + ) + + mock_client = MagicMock() + mock_client.get = AsyncMock(side_effect=httpx.ConnectError("boom")) + + with patch.object( + discoverable_endpoints, "get_async_httpx_client", return_value=mock_client + ): + with pytest.raises(HTTPException) as exc_info: + await _build_oauth_protected_resource_response( + request=_make_request(), + mcp_server_name="sample_docs", + use_standard_pattern=True, + ) + + assert exc_info.value.status_code == 502 + + +@pytest.mark.asyncio +async def test_fetch_upstream_metadata_returns_none_when_not_all_candidates_network_fail(): + passthrough_server = MCPServer( + server_id="passthrough-partial-network", + name="sample_docs", + server_name="sample_docs", + alias="sample_docs", + url="https://upstream.example.com/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.none, + extra_headers=["Authorization"], + oauth_passthrough=True, + ) + + not_found_response = MagicMock() + not_found_response.status_code = 404 + mock_client = MagicMock() + mock_client.get = AsyncMock( + side_effect=[not_found_response, httpx.ConnectError("path fallback failed")] + ) + + with patch.object( + discoverable_endpoints, "get_async_httpx_client", return_value=mock_client + ): + result = await discoverable_endpoints.fetch_upstream_oauth_protected_resource( + passthrough_server + ) + + assert result is None + assert mock_client.get.await_count == 2 + + +@pytest.mark.asyncio +async def test_oauth_protected_resource_gateway_managed_unchanged(): + """Regression guard: OAuth2 servers still advertise the gateway as AS.""" + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + global_mcp_server_manager, + ) + + global_mcp_server_manager.registry.clear() + oauth2_server = MCPServer( + server_id="oauth2-1", + name="keycloak_whoami", + server_name="keycloak_whoami", + alias="keycloak_whoami", + url="https://upstream.example.com/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + client_id="cid", + client_secret="cs", + authorization_url="https://keycloak/auth", + token_url="https://keycloak/token", + scopes=["read"], + ) + global_mcp_server_manager.registry[oauth2_server.server_id] = oauth2_server + + # If the code mistakenly fetched upstream metadata for a gateway-managed + # server, this spy would catch it. + mock_client = MagicMock() + mock_client.get = AsyncMock() + + with patch.object( + discoverable_endpoints, "get_async_httpx_client", return_value=mock_client + ): + result = await _build_oauth_protected_resource_response( + request=_make_request(), + mcp_server_name="keycloak_whoami", + use_standard_pattern=True, + ) + + mock_client.get.assert_not_awaited() + assert result["authorization_servers"] == [ + "https://gateway.example.com/keycloak_whoami" + ] + assert result["scopes_supported"] == ["read"] diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_oauth_passthrough_cold_start.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_oauth_passthrough_cold_start.py new file mode 100644 index 00000000000..3e934577a66 --- /dev/null +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_oauth_passthrough_cold_start.py @@ -0,0 +1,156 @@ +"""Unit tests for MCP OAuth passthrough cold-start route behavior.""" + +import sys + +import pytest + +sys.path.insert(0, "../../../../../") + +from litellm.proxy._types import MCPTransport +from litellm.types.mcp import MCPAuth +from litellm.types.mcp_server.mcp_server_manager import MCPServer + + +def _make_scope(path: str, headers: list = None) -> dict: + """Build a minimal ASGI HTTP scope for testing.""" + raw_headers = [(key.encode(), value.encode()) for key, value in (headers or [])] + return { + "type": "http", + "method": "POST", + "path": path, + "headers": raw_headers, + "query_string": b"", + "server": ("localhost", 4000), + "scheme": "http", + } + + +@pytest.mark.parametrize( + "route,expected_metadata_path", + [ + ( + "/mcp/sample_docs", + "/.well-known/oauth-protected-resource/mcp/sample_docs", + ), + ( + "/sample_docs/mcp", + "/.well-known/oauth-protected-resource/sample_docs/mcp", + ), + ], +) +def test_passthrough_cold_start_emits_401_with_matching_resource_metadata( + route, expected_metadata_path +): + """No auth headers on a passthrough server route emits matching metadata.""" + from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import ( + _is_mcp_passthrough_cold_start, + _parse_mcp_server_names_from_path, + ) + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + global_mcp_server_manager, + ) + + global_mcp_server_manager.registry.clear() + passthrough_server = MCPServer( + server_id="pt-cold-start", + name="sample_docs", + server_name="sample_docs", + alias="sample_docs", + url="https://upstream.example.com/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.none, + extra_headers=["Authorization"], + oauth_passthrough=True, + ) + global_mcp_server_manager.registry[passthrough_server.server_id] = ( + passthrough_server + ) + + if route.startswith("/mcp/"): + scope = _make_scope(route) + else: + scope = _make_scope("/mcp/sample_docs") + scope["_original_path"] = route + + servers = _parse_mcp_server_names_from_path(scope.get("path", "")) + assert _is_mcp_passthrough_cold_start(servers, client_ip=None) is True + + server_name = "sample_docs" + base_url = "http://localhost:4000" + path = scope.get("_original_path") or scope.get("path", "") or "" + if path.startswith(f"/{server_name}/mcp"): + resource_metadata_url = ( + f"{base_url}/.well-known/oauth-protected-resource/{server_name}/mcp" + ) + else: + resource_metadata_url = ( + f"{base_url}/.well-known/oauth-protected-resource/mcp/{server_name}" + ) + + assert resource_metadata_url == f"{base_url}{expected_metadata_path}", ( + f"resource_metadata_url {resource_metadata_url!r} does not match " + f"expected {base_url + expected_metadata_path!r}" + ) + + +def test_is_mcp_passthrough_cold_start_false_for_oauth2_server(): + """Gateway-managed OAuth2 servers must not trigger the cold-start bypass.""" + from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import ( + _is_mcp_passthrough_cold_start, + ) + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + global_mcp_server_manager, + ) + + global_mcp_server_manager.registry.clear() + oauth2_server = MCPServer( + server_id="oauth2-cold", + name="keycloak_whoami", + server_name="keycloak_whoami", + alias="keycloak_whoami", + url="https://upstream.example.com/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + client_id="cid", + client_secret="cs", + authorization_url="https://keycloak/auth", + token_url="https://keycloak/token", + ) + global_mcp_server_manager.registry[oauth2_server.server_id] = oauth2_server + + result = _is_mcp_passthrough_cold_start(["keycloak_whoami"], client_ip=None) + assert result is False + + +def test_is_mcp_passthrough_cold_start_false_for_empty_servers(): + """Aggregate /mcp route (no server list) must not trigger bypass.""" + from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import ( + _is_mcp_passthrough_cold_start, + ) + + assert _is_mcp_passthrough_cold_start(None, client_ip=None) is False + assert _is_mcp_passthrough_cold_start([], client_ip=None) is False + + +@pytest.mark.parametrize( + "path,expected", + [ + ("/mcp/sample_docs", ["sample_docs"]), + # Server names may contain at most one slash (mirrors + # ``_extract_target_server_names_from_path``), so when more than two + # segments follow ``/mcp/`` the first two are treated as the name. + ("/mcp/sample_docs/tools/list", ["sample_docs/tools"]), + ("/mcp/custom_solutions/user_123", ["custom_solutions/user_123"]), + ("/sample_docs/mcp", ["sample_docs"]), + ("/sample_docs/mcp/tools/list", ["sample_docs"]), + ("/mcp", None), + ("/mcp/", None), + ("/other/path", None), + ], +) +def test_parse_mcp_server_names_from_path(path, expected): + from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import ( + _parse_mcp_server_names_from_path, + ) + + assert _parse_mcp_server_names_from_path(path) == expected diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_oauth_passthrough_tools.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_oauth_passthrough_tools.py new file mode 100644 index 00000000000..d900f690c57 --- /dev/null +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_oauth_passthrough_tools.py @@ -0,0 +1,197 @@ +"""Unit tests for MCP OAuth passthrough tool-fetch behavior.""" + +import sys +from unittest.mock import AsyncMock, MagicMock + +import httpx +import pytest + +sys.path.insert(0, "../../../../../") + +from litellm.proxy._experimental.mcp_server.exceptions import MCPUpstreamAuthError +from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + MCPServerManager, + _extract_upstream_auth_failure, +) +from litellm.proxy._types import MCPTransport +from litellm.types.mcp import MCPAuth +from litellm.types.mcp_server.mcp_server_manager import MCPServer + + +def test_extract_upstream_auth_failure_finds_401_in_http_status_error(): + response = httpx.Response( + status_code=401, + headers={"www-authenticate": 'Bearer resource_metadata="https://x"'}, + request=httpx.Request("GET", "https://upstream/mcp"), + ) + exc = httpx.HTTPStatusError("401", request=response.request, response=response) + + result = _extract_upstream_auth_failure(exc) + assert result == (401, 'Bearer resource_metadata="https://x"') + + +def test_extract_upstream_auth_failure_walks_exception_group(): + response = httpx.Response( + status_code=401, + headers={"www-authenticate": "Bearer"}, + request=httpx.Request("GET", "https://upstream/mcp"), + ) + inner = httpx.HTTPStatusError("401", request=response.request, response=response) + + try: + raise ExceptionGroup("wrapped", [inner]) # noqa: F821 (PEP 654, py3.11+) + except Exception as group: + result = _extract_upstream_auth_failure(group) + + assert result == (401, "Bearer") + + +def test_extract_upstream_auth_failure_returns_none_for_non_auth(): + assert _extract_upstream_auth_failure(RuntimeError("boom")) is None + + +@pytest.mark.asyncio +async def test_fetch_tools_from_passthrough_raises_on_upstream_401(): + manager = MCPServerManager() + passthrough_server = MCPServer( + server_id="p1", + name="sample_docs", + url="https://upstream/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.none, + extra_headers=["Authorization"], + oauth_passthrough=True, + ) + + response = httpx.Response( + status_code=401, + headers={"www-authenticate": 'Bearer resource_metadata="https://upstream"'}, + request=httpx.Request("GET", "https://upstream/mcp"), + ) + upstream_error = httpx.HTTPStatusError( + "401", request=response.request, response=response + ) + + mock_client = MagicMock() + mock_client.list_tools = AsyncMock(side_effect=upstream_error) + + with pytest.raises(MCPUpstreamAuthError) as exc_info: + await manager._fetch_tools_with_timeout( + mock_client, passthrough_server.name, server=passthrough_server + ) + + assert exc_info.value.status_code == 401 + assert exc_info.value.www_authenticate == ( + 'Bearer resource_metadata="https://upstream"' + ) + assert exc_info.value.server_name == "sample_docs" + mock_client.list_tools.assert_awaited_with(raise_on_error=True) + + +@pytest.mark.asyncio +async def test_fetch_tools_from_passthrough_returns_tools_on_success(): + manager = MCPServerManager() + passthrough_server = MCPServer( + server_id="p1", + name="sample_docs", + url="https://upstream/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.none, + extra_headers=["Authorization"], + oauth_passthrough=True, + ) + + tool = MagicMock() + tool.name = "list_documents" + mock_client = MagicMock() + mock_client.list_tools = AsyncMock(return_value=[tool]) + + tools = await manager._fetch_tools_with_timeout( + mock_client, passthrough_server.name, server=passthrough_server + ) + assert tools == [tool] + + +def test_to_http_exception_preserves_upstream_www_authenticate(): + err = MCPUpstreamAuthError( + status_code=401, + www_authenticate='Bearer resource_metadata="https://upstream/.well-known/oauth-protected-resource"', + server_name="sample_docs", + ) + + http_exc = err.to_http_exception() + assert http_exc.status_code == 401 + assert http_exc.headers == { + "www-authenticate": 'Bearer resource_metadata="https://upstream/.well-known/oauth-protected-resource"' + } + + +def test_to_http_exception_skips_fabrication_when_base_url_missing(): + """Without ``base_url`` we cannot build an RFC 9728 §3.2-compliant absolute + URI, so we omit the fabricated ``WWW-Authenticate`` challenge entirely + instead of emitting a relative URI strict clients reject.""" + err = MCPUpstreamAuthError( + status_code=401, + www_authenticate=None, + server_name="sample_docs", + ) + + http_exc = err.to_http_exception() + assert http_exc.status_code == 401 + assert http_exc.headers is None + + +def test_to_http_exception_fabricates_absolute_resource_metadata_with_base_url(): + err = MCPUpstreamAuthError( + status_code=401, + www_authenticate=None, + server_name="sample_docs", + ) + + http_exc = err.to_http_exception(base_url="https://gateway.example.com/") + assert http_exc.status_code == 401 + assert http_exc.headers == { + "www-authenticate": 'Bearer resource_metadata="https://gateway.example.com/.well-known/oauth-protected-resource/mcp/sample_docs"' + } + + +def test_to_http_exception_skips_challenge_for_non_401_status(): + err = MCPUpstreamAuthError( + status_code=403, + www_authenticate=None, + server_name="sample_docs", + ) + + http_exc = err.to_http_exception() + assert http_exc.status_code == 403 + assert http_exc.headers is None + + +@pytest.mark.asyncio +async def test_fetch_tools_from_gateway_managed_swallows_errors(): + """Regression guard: non-pass-through servers keep returning [] on errors.""" + manager = MCPServerManager() + oauth2_server = MCPServer( + server_id="o1", + name="keycloak_whoami", + url="https://upstream/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + ) + + response = httpx.Response( + status_code=401, + headers={}, + request=httpx.Request("GET", "https://upstream/mcp"), + ) + upstream_error = httpx.HTTPStatusError( + "401", request=response.request, response=response + ) + mock_client = MagicMock() + mock_client.list_tools = AsyncMock(side_effect=upstream_error) + + tools = await manager._fetch_tools_with_timeout( + mock_client, oauth2_server.name, server=oauth2_server + ) + assert tools == [] + mock_client.list_tools.assert_awaited_with(raise_on_error=False) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py index bb0cc860375..6b6c7bc37d5 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py @@ -131,13 +131,129 @@ def test_prepare_mcp_server_headers_case_insensitive_extra_headers(): mcp_server_auth_headers=None, mcp_auth_header=None, oauth2_headers=None, - raw_headers={"authorization": "Bearer token"}, + raw_headers={ + "x-litellm-api-key": "Bearer sk-litellm-key", + "authorization": "Bearer token", + }, ) assert server_auth_header is None assert extra_headers == {"Authorization": "Bearer token"} +def test_prepare_mcp_server_headers_passthrough_strips_authorization_without_admission_header(): + try: + from litellm.proxy._experimental.mcp_server.server import ( + _prepare_mcp_server_headers, + ) + except ImportError: + pytest.skip("MCP server not available") + + server = MCPServer( + server_id="server-passthrough-no-admission", + name="server", + transport=MCPTransport.http, + auth_type=MCPAuth.none, + extra_headers=["Authorization", "x-request-id"], + oauth_passthrough=True, + ) + + server_auth_header, extra_headers = _prepare_mcp_server_headers( + server=server, + mcp_server_auth_headers=None, + mcp_auth_header=None, + oauth2_headers=None, + raw_headers={ + "authorization": "Bearer sk-litellm-key", + "x-request-id": "req-789", + }, + ) + + assert server_auth_header is None + assert extra_headers == {"x-request-id": "req-789"} + + +def test_prepare_mcp_server_headers_passthrough_forwards_authorization_for_anonymous_admission(): + """Cold-start return per RFC 9728: client admits anonymously through + the pass-through fallback in :meth:`MCPRequestHandler.process_mcp_request` + (``user_api_key_auth.api_key is None``) and the ``Authorization`` bearer + is the upstream OAuth token — it must be forwarded, not stripped.""" + try: + from litellm.proxy._experimental.mcp_server.server import ( + _prepare_mcp_server_headers, + ) + except ImportError: + pytest.skip("MCP server not available") + + from litellm.proxy._types import UserAPIKeyAuth + + server = MCPServer( + server_id="server-passthrough-anon-admission", + name="server", + transport=MCPTransport.http, + auth_type=MCPAuth.none, + extra_headers=["Authorization", "x-request-id"], + oauth_passthrough=True, + ) + + server_auth_header, extra_headers = _prepare_mcp_server_headers( + server=server, + mcp_server_auth_headers=None, + mcp_auth_header=None, + oauth2_headers=None, + raw_headers={ + "authorization": "Bearer upstream-oauth-token", + "x-request-id": "req-790", + }, + user_api_key_auth=UserAPIKeyAuth(), + ) + + assert server_auth_header is None + assert extra_headers == { + "Authorization": "Bearer upstream-oauth-token", + "x-request-id": "req-790", + } + + +def test_prepare_mcp_server_headers_passthrough_strips_authorization_for_authenticated_admission(): + """When admission validated ``Authorization`` as a LiteLLM key + (``user_api_key_auth.api_key`` is set, no explicit ``x-litellm-api-key``), + the bearer must still be stripped to avoid leaking the gateway key + upstream.""" + try: + from litellm.proxy._experimental.mcp_server.server import ( + _prepare_mcp_server_headers, + ) + except ImportError: + pytest.skip("MCP server not available") + + from litellm.proxy._types import UserAPIKeyAuth + + server = MCPServer( + server_id="server-passthrough-authenticated", + name="server", + transport=MCPTransport.http, + auth_type=MCPAuth.none, + extra_headers=["Authorization", "x-request-id"], + oauth_passthrough=True, + ) + + server_auth_header, extra_headers = _prepare_mcp_server_headers( + server=server, + mcp_server_auth_headers=None, + mcp_auth_header=None, + oauth2_headers=None, + raw_headers={ + "authorization": "Bearer sk-litellm-key", + "x-request-id": "req-791", + }, + user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-key"), + ) + + assert server_auth_header is None + assert extra_headers == {"x-request-id": "req-791"} + + def test_prepare_mcp_server_headers_oauth2_m2m_omits_litellm_caller_authorization(): """M2M OAuth must not put caller Bearer (LiteLLM API key) into extra_headers (#23652).""" try: @@ -514,6 +630,7 @@ async def test_mcp_get_prompt_success(): mcp_auth_header=None, oauth2_headers=None, raw_headers=None, + user_api_key_auth=user_api_key_auth, ) mock_manager.get_prompt_from_server.assert_awaited_once_with( server=server, @@ -575,6 +692,7 @@ async def test_mcp_read_resource_success(): mcp_auth_header=None, oauth2_headers=None, raw_headers=None, + user_api_key_auth=user_api_key_auth, ) mock_manager.read_resource_from_server.assert_awaited_once_with( server=server, diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py index 2db9845c765..ec690aef629 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -480,6 +480,167 @@ class TestMCPServerManager: assert captured_extra_headers == {"Authorization": "Bearer token"} assert isinstance(result, CallToolResult) + @pytest.mark.asyncio + async def test_call_regular_mcp_tool_passthrough_strips_authorization_when_admission_consumed_litellm_key( + self, + ): + """OAuth pass-through must not forward the caller's Authorization to upstream + when LiteLLM admission consumed the bearer as its API key — otherwise the + LiteLLM key the caller used for admission would leak upstream.""" + from litellm.proxy._types import UserAPIKeyAuth + + manager = MCPServerManager() + server = MCPServer( + server_id="server-passthrough-call", + name="passthrough-server", + url="https://example.com", + transport=MCPTransport.http, + auth_type=MCPAuth.none, + extra_headers=["Authorization", "x-request-id"], + oauth_passthrough=True, + ) + + mock_client = AsyncMock() + mock_client.call_tool = AsyncMock( + return_value=CallToolResult(content=[], isError=False) + ) + captured_extra_headers = None + + async def capture_create_mcp_client( + server, mcp_auth_header, extra_headers, stdio_env, subject_token=None + ): # pragma: no cover - helper + nonlocal captured_extra_headers + captured_extra_headers = extra_headers + return mock_client + + manager._create_mcp_client = AsyncMock(side_effect=capture_create_mcp_client) + + await manager._call_regular_mcp_tool( + mcp_server=server, + original_tool_name="tool", + arguments={}, + tasks=[], + mcp_auth_header=None, + mcp_server_auth_headers=None, + oauth2_headers=None, + raw_headers={ + "authorization": "Bearer sk-litellm-key", + "x-request-id": "req-123", + }, + proxy_logging_obj=None, + user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-key"), + ) + + assert captured_extra_headers == {"x-request-id": "req-123"} + + @pytest.mark.asyncio + async def test_call_regular_mcp_tool_passthrough_forwards_authorization_with_admission_header( + self, + ): + """OAuth pass-through forwards Authorization upstream when x-litellm-api-key + provides admission — in that case Authorization carries the upstream OAuth + bearer, not the LiteLLM key.""" + from litellm.proxy._types import UserAPIKeyAuth + + manager = MCPServerManager() + server = MCPServer( + server_id="server-passthrough-call-admission", + name="passthrough-server", + url="https://example.com", + transport=MCPTransport.http, + auth_type=MCPAuth.none, + extra_headers=["Authorization"], + oauth_passthrough=True, + ) + + mock_client = AsyncMock() + mock_client.call_tool = AsyncMock( + return_value=CallToolResult(content=[], isError=False) + ) + captured_extra_headers = None + + async def capture_create_mcp_client( + server, mcp_auth_header, extra_headers, stdio_env, subject_token=None + ): # pragma: no cover - helper + nonlocal captured_extra_headers + captured_extra_headers = extra_headers + return mock_client + + manager._create_mcp_client = AsyncMock(side_effect=capture_create_mcp_client) + + await manager._call_regular_mcp_tool( + mcp_server=server, + original_tool_name="tool", + arguments={}, + tasks=[], + mcp_auth_header=None, + mcp_server_auth_headers=None, + oauth2_headers=None, + raw_headers={ + "x-litellm-api-key": "Bearer sk-litellm-key", + "authorization": "Bearer upstream-oauth-bearer", + }, + proxy_logging_obj=None, + user_api_key_auth=UserAPIKeyAuth(api_key="sk-litellm-key"), + ) + + assert captured_extra_headers == { + "Authorization": "Bearer upstream-oauth-bearer" + } + + @pytest.mark.asyncio + async def test_call_regular_mcp_tool_passthrough_forwards_authorization_for_anonymous_admission( + self, + ): + """OAuth pass-through cold-start return (RFC 9728): the caller's only + credential is the upstream bearer in Authorization, and LiteLLM admission + is anonymous (no api_key on user_api_key_auth). Authorization must be + forwarded so the delegated flow can complete.""" + from litellm.proxy._types import UserAPIKeyAuth + + manager = MCPServerManager() + server = MCPServer( + server_id="server-passthrough-call-anon", + name="passthrough-server", + url="https://example.com", + transport=MCPTransport.http, + auth_type=MCPAuth.none, + extra_headers=["Authorization"], + oauth_passthrough=True, + ) + + mock_client = AsyncMock() + mock_client.call_tool = AsyncMock( + return_value=CallToolResult(content=[], isError=False) + ) + captured_extra_headers = None + + async def capture_create_mcp_client( + server, mcp_auth_header, extra_headers, stdio_env, subject_token=None + ): # pragma: no cover - helper + nonlocal captured_extra_headers + captured_extra_headers = extra_headers + return mock_client + + manager._create_mcp_client = AsyncMock(side_effect=capture_create_mcp_client) + + await manager._call_regular_mcp_tool( + mcp_server=server, + original_tool_name="tool", + arguments={}, + tasks=[], + mcp_auth_header=None, + mcp_server_auth_headers=None, + oauth2_headers=None, + raw_headers={"authorization": "Bearer upstream-oauth-bearer"}, + proxy_logging_obj=None, + user_api_key_auth=UserAPIKeyAuth(api_key=None), + ) + + assert captured_extra_headers == { + "Authorization": "Bearer upstream-oauth-bearer" + } + @pytest.mark.asyncio async def test_get_prompts_from_server_success(self): """Ensure prompts are fetched and prefixed when requested.""" diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py index 593facd9279..6433e0f6360 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py @@ -544,6 +544,78 @@ class TestListToolsRestAPI: assert result["error"] is None assert result["message"] == "Successfully retrieved tools" + @pytest.mark.parametrize("upstream_status", [401, 403]) + async def test_upstream_auth_failure_surfaces_status_and_challenge( + self, monkeypatch, upstream_status + ): + """A single-server pass-through request whose upstream rejects the token + must surface the upstream status (401 or 403) plus its WWW-Authenticate + challenge, not collapse into a 200 ``unexpected_error`` body.""" + from litellm.proxy._experimental.mcp_server.exceptions import ( + MCPUpstreamAuthError, + ) + + class StubServer: + alias = "server-1" + server_name = "server-1" + name = "passthrough" + allowed_tools = None + mcp_info = {"server_name": "passthrough"} + available_on_public_internet = True + + stub_server = StubServer() + + async def fake_contexts(user_api_key_auth): + return [user_api_key_auth] + + async def fake_get_allowed_mcp_servers(*args, **kwargs): + return ["server-1"] + + challenge = 'Bearer resource_metadata="https://upstream/.well-known"' + + async def fake_get_tools(*args, **kwargs): + raise MCPUpstreamAuthError( + status_code=upstream_status, + www_authenticate=challenge, + server_name="passthrough", + ) + + monkeypatch.setattr( + rest_endpoints, + "build_effective_auth_contexts", + fake_contexts, + raising=False, + ) + monkeypatch.setattr( + rest_endpoints.global_mcp_server_manager, + "get_allowed_mcp_servers", + fake_get_allowed_mcp_servers, + raising=False, + ) + monkeypatch.setattr( + rest_endpoints.global_mcp_server_manager, + "get_mcp_server_by_id", + lambda server_id: stub_server if server_id == "server-1" else None, + raising=False, + ) + monkeypatch.setattr( + rest_endpoints, + "_get_tools_for_single_server", + fake_get_tools, + raising=False, + ) + + request = _build_request(path="/mcp-rest/tools/list", method="GET") + with pytest.raises(HTTPException) as exc_info: + await rest_endpoints.list_tool_rest_api( + request, + server_id="server-1", + user_api_key_dict=UserAPIKeyAuth(), + ) + + assert exc_info.value.status_code == upstream_status + assert exc_info.value.headers == {"www-authenticate": challenge} + async def test_name_resolution_finds_server_by_uuid(self, monkeypatch): """When server_id is a name string, it should be resolved to its UUID and used for the tools lookup when the UUID is in allowed_server_ids.""" diff --git a/tests/test_litellm/proxy/auth/test_handle_jwt.py b/tests/test_litellm/proxy/auth/test_handle_jwt.py index 14119f7ad4e..cca6754956f 100644 --- a/tests/test_litellm/proxy/auth/test_handle_jwt.py +++ b/tests/test_litellm/proxy/auth/test_handle_jwt.py @@ -3179,3 +3179,611 @@ def test_build_decode_kwargs_no_warning_when_scoped( if "neither JWT_AUDIENCE nor JWT_ISSUER" in r.getMessage() ] assert matching == [] + + +def _base64url_encode_int(value: int) -> str: + import base64 + + value_bytes = value.to_bytes((value.bit_length() + 7) // 8, "big") + return base64.urlsafe_b64encode(value_bytes).decode("utf-8").rstrip("=") + + +def _get_rsa_key_and_jwk(kid: str): + from cryptography.hazmat.primitives.asymmetric import rsa + + private_key = rsa.generate_private_key(public_exponent=65537, key_size=2048) + public_numbers = private_key.public_key().public_numbers() + jwk = { + "kty": "RSA", + "n": _base64url_encode_int(value=public_numbers.n), + "e": _base64url_encode_int(value=public_numbers.e), + "kid": kid, + "alg": "RS256", + "use": "sig", + } + return private_key, jwk + + +def _encode_rsa_jwt( + private_key, + issuer: str, + audience: str, + kid: str, + extra_claims: Optional[dict] = None, +) -> str: + import time + + import jwt + from cryptography.hazmat.primitives import serialization + + private_key_pem = private_key.private_bytes( + encoding=serialization.Encoding.PEM, + format=serialization.PrivateFormat.PKCS8, + encryption_algorithm=serialization.NoEncryption(), + ) + current_time = int(time.time()) + claims = { + "sub": "test-subject", + "iss": issuer, + "aud": audience, + "iat": current_time, + "exp": current_time + 300, + } + if extra_claims: + claims.update(extra_claims) + + return jwt.encode( + claims, + private_key_pem, + algorithm="RS256", + headers={"kid": kid}, + ) + + +def _get_jwt_handler_with_issuer_keys(issuers: list, keys_by_url: dict) -> JWTHandler: + from litellm.caching.dual_cache import DualCache + + cache = DualCache() + for jwks_url, keys in keys_by_url.items(): + cache.set_cache( + key=f"litellm_jwt_auth_keys_{jwks_url}", + value=keys, + ) + + jwt_handler = JWTHandler() + jwt_handler.update_environment( + prisma_client=None, + user_api_key_cache=cache, + litellm_jwtauth=LiteLLM_JWTAuth(issuers=issuers), + ) + return jwt_handler + + +@pytest.mark.asyncio +async def test_get_public_key_fetches_and_caches_jwks_response(): + from unittest.mock import AsyncMock, MagicMock + + from litellm.caching.dual_cache import DualCache + + jwt_handler = JWTHandler() + cache = DualCache() + jwt_handler.update_environment( + prisma_client=None, + user_api_key_cache=cache, + litellm_jwtauth=LiteLLM_JWTAuth(public_key_ttl=123), + ) + expected_key_id = "cached-key" + _, jwk = _get_rsa_key_and_jwk(kid=expected_key_id) + mock_response = MagicMock() + mock_response.json.return_value = {"keys": [jwk]} + jwt_handler.http_handler.get = AsyncMock(return_value=mock_response) + + public_key = await jwt_handler._get_public_key_from_jwks_url( + jwks_url="https://issuer.example.com/keys", + kid=expected_key_id, + ) + + assert public_key == jwk + cached_keys = await cache.async_get_cache( + key="litellm_jwt_auth_keys_https://issuer.example.com/keys" + ) + assert cached_keys == [jwk] + + +@pytest.mark.asyncio +async def test_get_public_key_tries_next_jwks_url_when_kid_missing(monkeypatch): + from litellm.caching.dual_cache import DualCache + + first_jwks_url = "https://first.example.com/keys" + second_jwks_url = "https://second.example.com/keys" + monkeypatch.setenv("JWT_PUBLIC_KEY_URL", f"{first_jwks_url}, {second_jwks_url},,") + _, first_jwk = _get_rsa_key_and_jwk(kid="first-key") + _, second_jwk = _get_rsa_key_and_jwk(kid="second-key") + cache = DualCache() + cache.set_cache(key=f"litellm_jwt_auth_keys_{first_jwks_url}", value=[first_jwk]) + cache.set_cache(key=f"litellm_jwt_auth_keys_{second_jwks_url}", value=[second_jwk]) + jwt_handler = JWTHandler() + jwt_handler.update_environment( + prisma_client=None, + user_api_key_cache=cache, + litellm_jwtauth=LiteLLM_JWTAuth(), + ) + + public_key = await jwt_handler.get_public_key(kid="second-key") + + assert public_key == second_jwk + + +def test_get_jwks_url_for_issuer_falls_back_to_discovery_document(): + jwt_handler = JWTHandler() + issuer_config = LiteLLM_JWTAuth( + issuers=[ + { + "issuer": "https://issuer.example.com/tenant/", + "disable_audience_validation": True, + } + ] + ).issuers[0] + + jwks_url = jwt_handler._get_jwks_url_for_issuer(issuer_config=issuer_config) + + assert ( + jwks_url == "https://issuer.example.com/tenant/.well-known/openid-configuration" + ) + + +@pytest.mark.asyncio +async def test_multi_issuer_jwt_validates_selected_issuer_and_maps_claims( + monkeypatch, +): + monkeypatch.delenv("JWT_AUDIENCE", raising=False) + monkeypatch.delenv("JWT_PUBLIC_KEY_URL", raising=False) + + issuer_one = "https://issuer-one.example.com" + issuer_two = "https://issuer-two.example.com" + issuer_one_jwks_url = f"{issuer_one}/keys" + issuer_two_jwks_url = f"{issuer_two}/keys" + shared_kid = "shared-kid" + + _, issuer_one_jwk = _get_rsa_key_and_jwk(kid=shared_kid) + issuer_two_private_key, issuer_two_jwk = _get_rsa_key_and_jwk(kid=shared_kid) + + jwt_handler = _get_jwt_handler_with_issuer_keys( + issuers=[ + { + "issuer": issuer_one, + "jwks_url": issuer_one_jwks_url, + "audience": "audience-one", + "user_id_jwt_field": "email", + "user_email_jwt_field": "email", + }, + { + "issuer": issuer_two, + "jwks_url": issuer_two_jwks_url, + "audience": "audience-two", + "user_id_jwt_field": "repository_owner", + "team_id_jwt_field": "repository", + }, + ], + keys_by_url={ + issuer_one_jwks_url: [issuer_one_jwk], + issuer_two_jwks_url: [issuer_two_jwk], + }, + ) + + token = _encode_rsa_jwt( + private_key=issuer_two_private_key, + issuer=issuer_two, + audience="audience-two", + kid=shared_kid, + extra_claims={ + "repository_owner": "example-org", + "repository": "example-org/litellm-fork", + }, + ) + + claims = await jwt_handler.auth_jwt(token=token) + + assert claims[JWTHandler.LITELLM_JWT_ISSUER_CLAIM] == issuer_two + assert jwt_handler.get_user_id(token=claims, default_value=None) == "example-org" + assert jwt_handler.get_team_id(token=claims, default_value=None) == ( + "example-org/litellm-fork" + ) + + +@pytest.mark.asyncio +async def test_multi_issuer_jwt_maps_kubernetes_namespace_claim(monkeypatch): + monkeypatch.delenv("JWT_AUDIENCE", raising=False) + monkeypatch.delenv("JWT_PUBLIC_KEY_URL", raising=False) + + issuer = "https://oidc.eks.eu-west-1.amazonaws.com/id/test-cluster" + jwks_url = f"{issuer}/keys" + private_key, jwk = _get_rsa_key_and_jwk(kid="k8s-key") + jwt_handler = _get_jwt_handler_with_issuer_keys( + issuers=[ + { + "issuer": issuer, + "jwks_url": jwks_url, + "audience": None, + "disable_audience_validation": True, + "user_id_jwt_field": "kubernetes\\.io.namespace", + } + ], + keys_by_url={jwks_url: [jwk]}, + ) + token = _encode_rsa_jwt( + private_key=private_key, + issuer=issuer, + audience="kubernetes.default.svc", + kid="k8s-key", + extra_claims={"kubernetes.io": {"namespace": "example-namespace"}}, + ) + + claims = await jwt_handler.auth_jwt(token=token) + + assert ( + jwt_handler.get_user_id(token=claims, default_value=None) == "example-namespace" + ) + + +@pytest.mark.asyncio +async def test_multi_issuer_jwt_unknown_issuer_falls_back_to_global_jwks(monkeypatch): + """Tokens whose ``iss`` is not in the configured issuers list fall through + to the legacy ``JWT_PUBLIC_KEY_URL`` path so operators can add the new + ``issuers`` list to a live deployment without breaking existing tokens + minted by non-configured IdPs. With no global JWKS configured, the legacy + path surfaces a ``Missing JWT Public Key URL from environment.`` error. + """ + monkeypatch.delenv("JWT_AUDIENCE", raising=False) + monkeypatch.delenv("JWT_PUBLIC_KEY_URL", raising=False) + + configured_issuer = "https://issuer.example.com" + private_key, jwk = _get_rsa_key_and_jwk(kid="issuer-key") + jwt_handler = _get_jwt_handler_with_issuer_keys( + issuers=[ + { + "issuer": configured_issuer, + "jwks_url": f"{configured_issuer}/keys", + "audience": "expected-audience", + } + ], + keys_by_url={f"{configured_issuer}/keys": [jwk]}, + ) + token = _encode_rsa_jwt( + private_key=private_key, + issuer="https://unknown-issuer.example.com", + audience="expected-audience", + kid="issuer-key", + ) + + with pytest.raises(Exception) as exc: + await jwt_handler.auth_jwt(token=token) + + assert "Missing JWT Public Key URL from environment." in str(exc.value) + + +@pytest.mark.asyncio +async def test_multi_issuer_jwt_rejects_wrong_audience(monkeypatch): + monkeypatch.delenv("JWT_AUDIENCE", raising=False) + monkeypatch.delenv("JWT_PUBLIC_KEY_URL", raising=False) + + issuer = "https://issuer.example.com" + jwks_url = f"{issuer}/keys" + private_key, jwk = _get_rsa_key_and_jwk(kid="issuer-key") + jwt_handler = _get_jwt_handler_with_issuer_keys( + issuers=[ + { + "issuer": issuer, + "jwks_url": jwks_url, + "audience": "expected-audience", + } + ], + keys_by_url={jwks_url: [jwk]}, + ) + token = _encode_rsa_jwt( + private_key=private_key, + issuer=issuer, + audience="wrong-audience", + kid="issuer-key", + ) + + with pytest.raises(Exception) as exc: + await jwt_handler.auth_jwt(token=token) + + assert "Validation fails" in str(exc.value) + + +@pytest.mark.asyncio +async def test_multi_issuer_jwt_same_kid_does_not_cross_issuer_keys(monkeypatch): + monkeypatch.delenv("JWT_AUDIENCE", raising=False) + monkeypatch.delenv("JWT_PUBLIC_KEY_URL", raising=False) + + issuer_one = "https://issuer-one.example.com" + issuer_two = "https://issuer-two.example.com" + issuer_one_jwks_url = f"{issuer_one}/keys" + issuer_two_jwks_url = f"{issuer_two}/keys" + shared_kid = "shared-kid" + issuer_one_private_key, issuer_one_jwk = _get_rsa_key_and_jwk(kid=shared_kid) + _, issuer_two_jwk = _get_rsa_key_and_jwk(kid=shared_kid) + jwt_handler = _get_jwt_handler_with_issuer_keys( + issuers=[ + { + "issuer": issuer_one, + "jwks_url": issuer_one_jwks_url, + "audience": "audience-one", + }, + { + "issuer": issuer_two, + "jwks_url": issuer_two_jwks_url, + "audience": "audience-two", + }, + ], + keys_by_url={ + issuer_one_jwks_url: [issuer_one_jwk], + issuer_two_jwks_url: [issuer_two_jwk], + }, + ) + token = _encode_rsa_jwt( + private_key=issuer_one_private_key, + issuer=issuer_two, + audience="audience-two", + kid=shared_kid, + ) + + with pytest.raises(Exception) as exc: + await jwt_handler.auth_jwt(token=token) + + assert "Validation fails" in str(exc.value) + + +@pytest.mark.asyncio +async def test_multi_issuer_jwt_missing_mapped_claim_leaves_user_id_unset( + monkeypatch, +): + """Mapped issuer claims behave like the global ``litellm_jwtauth`` path — + present claims override the normalised value, missing ones simply leave + the corresponding LiteLLM-internal claim absent (rather than failing the + JWT outright). This keeps multi-issuer auth tolerant of tokens that omit + optional fields like email or org id. + """ + monkeypatch.delenv("JWT_AUDIENCE", raising=False) + monkeypatch.delenv("JWT_PUBLIC_KEY_URL", raising=False) + + issuer = "https://issuer.example.com" + jwks_url = f"{issuer}/keys" + private_key, jwk = _get_rsa_key_and_jwk(kid="issuer-key") + jwt_handler = _get_jwt_handler_with_issuer_keys( + issuers=[ + { + "issuer": issuer, + "jwks_url": jwks_url, + "audience": "expected-audience", + "user_id_jwt_field": "email", + } + ], + keys_by_url={jwks_url: [jwk]}, + ) + token = _encode_rsa_jwt( + private_key=private_key, + issuer=issuer, + audience="expected-audience", + kid="issuer-key", + ) + + claims = await jwt_handler.auth_jwt(token=token) + + assert claims[jwt_handler.LITELLM_JWT_ISSUER_CLAIM] == issuer + assert jwt_handler.LITELLM_USER_ID_CLAIM not in claims + + +def test_multi_issuer_jwt_requires_audience_unless_explicitly_disabled( + monkeypatch, +): + monkeypatch.delenv("JWT_AUDIENCE", raising=False) + monkeypatch.delenv("JWT_PUBLIC_KEY_URL", raising=False) + + issuer = "https://issuer.example.com" + jwks_url = f"{issuer}/keys" + + with pytest.raises(Exception) as exc: + LiteLLM_JWTAuth( + issuers=[ + { + "issuer": issuer, + "jwks_url": jwks_url, + } + ] + ) + + assert "must configure audience" in str(exc.value) + + +def test_multi_issuer_jwt_rejects_audience_with_disable_audience_validation(): + issuer = "https://issuer.example.com" + jwks_url = f"{issuer}/keys" + + with pytest.raises(Exception) as exc: + LiteLLM_JWTAuth( + issuers=[ + { + "issuer": issuer, + "jwks_url": jwks_url, + "audience": "some-audience", + "disable_audience_validation": True, + } + ] + ) + + assert "cannot set audience and disable_audience_validation=True together" in str( + exc.value + ) + + +@pytest.mark.asyncio +async def test_global_jwt_ignores_user_supplied_internal_claims(monkeypatch): + from litellm.caching.dual_cache import DualCache + + monkeypatch.delenv("JWT_AUDIENCE", raising=False) + monkeypatch.delenv("JWT_ISSUER", raising=False) + + jwks_url = "https://global-issuer.example.com/keys" + monkeypatch.setenv("JWT_PUBLIC_KEY_URL", jwks_url) + + private_key, jwk = _get_rsa_key_and_jwk(kid="global-key") + cache = DualCache() + cache.set_cache(key=f"litellm_jwt_auth_keys_{jwks_url}", value=[jwk]) + + jwt_handler = JWTHandler() + jwt_handler.update_environment( + prisma_client=None, + user_api_key_cache=cache, + litellm_jwtauth=LiteLLM_JWTAuth( + user_id_jwt_field="email", + user_email_jwt_field="email", + team_id_jwt_field="team.id", + team_ids_jwt_field="teams", + org_id_jwt_field="org.id", + end_user_id_jwt_field="end_user.id", + ), + ) + token = _encode_rsa_jwt( + private_key=private_key, + issuer="https://global-issuer.example.com", + audience="some-other-client", + kid="global-key", + extra_claims={ + "email": "real-user@example.com", + "team": {"id": "real-team"}, + "teams": ["real-team", "secondary-team"], + "org": {"id": "real-org"}, + "end_user": {"id": "real-end-user"}, + JWTHandler.LITELLM_JWT_ISSUER_CLAIM: "https://issuer.example.com", + JWTHandler.LITELLM_USER_ID_CLAIM: "victim-user", + JWTHandler.LITELLM_USER_EMAIL_CLAIM: "victim@example.com", + JWTHandler.LITELLM_TEAM_ID_CLAIM: "victim-team", + JWTHandler.LITELLM_TEAM_IDS_CLAIM: ["victim-team"], + JWTHandler.LITELLM_ORG_ID_CLAIM: "victim-org", + JWTHandler.LITELLM_END_USER_ID_CLAIM: "victim-end-user", + }, + ) + + claims = await jwt_handler.auth_jwt(token=token) + + assert jwt_handler.get_user_id(token=claims, default_value=None) == ( + "real-user@example.com" + ) + assert jwt_handler.get_user_email(token=claims, default_value=None) == ( + "real-user@example.com" + ) + assert jwt_handler.get_team_id(token=claims, default_value=None) == "real-team" + assert jwt_handler.get_team_ids_from_jwt(token=claims) == [ + "real-team", + "secondary-team", + ] + assert jwt_handler.get_org_id(token=claims, default_value=None) == "real-org" + assert jwt_handler.get_end_user_id(token=claims, default_value=None) == ( + "real-end-user" + ) + + +@pytest.mark.asyncio +async def test_multi_issuer_jwt_strips_unmapped_internal_claims(monkeypatch): + monkeypatch.delenv("JWT_AUDIENCE", raising=False) + monkeypatch.delenv("JWT_PUBLIC_KEY_URL", raising=False) + + issuer = "https://issuer.example.com" + jwks_url = f"{issuer}/keys" + private_key, jwk = _get_rsa_key_and_jwk(kid="issuer-key") + jwt_handler = _get_jwt_handler_with_issuer_keys( + issuers=[ + { + "issuer": issuer, + "jwks_url": jwks_url, + "audience": "expected-audience", + "user_email_jwt_field": "email", + } + ], + keys_by_url={jwks_url: [jwk]}, + ) + token = _encode_rsa_jwt( + private_key=private_key, + issuer=issuer, + audience="expected-audience", + kid="issuer-key", + extra_claims={ + "email": "real-user@example.com", + JWTHandler.LITELLM_USER_ID_CLAIM: "victim-user", + JWTHandler.LITELLM_TEAM_ID_CLAIM: "victim-team", + }, + ) + + claims = await jwt_handler.auth_jwt(token=token) + + assert JWTHandler.LITELLM_USER_ID_CLAIM not in claims + assert JWTHandler.LITELLM_TEAM_ID_CLAIM not in claims + assert jwt_handler.get_user_id(token=claims, default_value=None) is None + assert jwt_handler.get_team_id(token=claims, default_value=None) is None + assert jwt_handler.get_user_email(token=claims, default_value=None) == ( + "real-user@example.com" + ) + + +@pytest.mark.asyncio +async def test_multi_issuer_jwt_does_not_emit_unscoped_global_warning( + monkeypatch, caplog +): + import logging + + monkeypatch.delenv("JWT_AUDIENCE", raising=False) + monkeypatch.delenv("JWT_ISSUER", raising=False) + monkeypatch.delenv("JWT_PUBLIC_KEY_URL", raising=False) + JWTHandler._unscoped_jwt_warning_emitted = False + + issuer = "https://issuer.example.com" + jwks_url = f"{issuer}/keys" + private_key, jwk = _get_rsa_key_and_jwk(kid="issuer-key") + jwt_handler = _get_jwt_handler_with_issuer_keys( + issuers=[ + { + "issuer": issuer, + "jwks_url": jwks_url, + "audience": "expected-audience", + } + ], + keys_by_url={jwks_url: [jwk]}, + ) + token = _encode_rsa_jwt( + private_key=private_key, + issuer=issuer, + audience="expected-audience", + kid="issuer-key", + ) + + with caplog.at_level(logging.WARNING): + await jwt_handler.auth_jwt(token=token) + + assert "Tokens minted by any application" not in caplog.text + assert JWTHandler._unscoped_jwt_warning_emitted is False + + +def test_build_decode_kwargs_warns_for_unscoped_global_fallback_in_mixed_deployment( + monkeypatch, _reset_unscoped_warning_flag, caplog +): + """The unscoped-fallback warning must fire even when per-issuer configs + are set. In mixed deployments, tokens whose ``iss`` does not match any + configured issuer fall through to the global path; if env-var scoping is + absent that fallback IS unscoped, and the operator needs to be told.""" + import logging + + monkeypatch.delenv("JWT_AUDIENCE", raising=False) + monkeypatch.delenv("JWT_ISSUER", raising=False) + caplog.set_level(logging.WARNING) + + JWTHandler._build_decode_kwargs() + + matching = [ + r + for r in caplog.records + if "neither JWT_AUDIENCE nor JWT_ISSUER" in r.getMessage() + ] + assert len(matching) == 1 diff --git a/tests/test_litellm/proxy/auth/test_mcp_ip_filtering.py b/tests/test_litellm/proxy/auth/test_mcp_ip_filtering.py index 3b13ef3641f..9444e4ebd2d 100644 --- a/tests/test_litellm/proxy/auth/test_mcp_ip_filtering.py +++ b/tests/test_litellm/proxy/auth/test_mcp_ip_filtering.py @@ -5,8 +5,9 @@ Tests that internal callers see all MCP servers while external callers only see servers with available_on_public_internet=True. """ -import ipaddress -from unittest.mock import patch +from unittest.mock import MagicMock, patch + +from fastapi import Request from litellm.proxy.auth.ip_address_utils import IPAddressUtils from litellm.types.mcp_server.mcp_server_manager import MCPServer @@ -58,6 +59,75 @@ class TestIsInternalIp: assert IPAddressUtils.is_internal_ip("not-an-ip") is False +class TestMCPClientIPExtraction: + def test_fails_closed_when_xff_enabled_without_trusted_proxy_ranges(self): + request = MagicMock(spec=Request) + request.client = MagicMock() + request.client.host = "203.0.113.5" + request.headers = {"x-forwarded-for": "10.0.0.1"} + + result = IPAddressUtils.get_mcp_client_ip( + request, + general_settings={"use_x_forwarded_for": True}, + ) + + # XFF is untrusted (no mcp_trusted_proxy_ranges) so it must be ignored, + # and we must not trust the direct peer either: fail closed so the caller + # is classified as external and is_internal_ip("") is False. + assert result == "" + assert IPAddressUtils.is_internal_ip(result) is False + + def test_private_proxy_peer_does_not_grant_internal_access(self): + # Regression: behind an internal reverse proxy with use_x_forwarded_for + # enabled but mcp_trusted_proxy_ranges unset, the direct peer is the + # proxy's private IP. Returning it would mis-classify an external caller + # as internal and expose available_on_public_internet=false servers. + request = MagicMock(spec=Request) + request.client = MagicMock() + request.client.host = "10.0.0.7" + request.headers = {"x-forwarded-for": "8.8.8.8"} + + result = IPAddressUtils.get_mcp_client_ip( + request, + general_settings={"use_x_forwarded_for": True}, + ) + + assert result == "" + assert IPAddressUtils.is_internal_ip(result) is False + + def test_honours_xff_from_trusted_proxy(self): + request = MagicMock(spec=Request) + request.client = MagicMock() + request.client.host = "10.0.0.5" + request.headers = {"x-forwarded-for": "192.168.1.10"} + + result = IPAddressUtils.get_mcp_client_ip( + request, + general_settings={ + "use_x_forwarded_for": True, + "mcp_trusted_proxy_ranges": ["10.0.0.0/8"], + }, + ) + + assert result == "192.168.1.10" + + def test_ignores_xff_from_untrusted_direct_caller(self): + request = MagicMock(spec=Request) + request.client = MagicMock() + request.client.host = "203.0.113.5" + request.headers = {"x-forwarded-for": "10.0.0.1"} + + result = IPAddressUtils.get_mcp_client_ip( + request, + general_settings={ + "use_x_forwarded_for": True, + "mcp_trusted_proxy_ranges": ["10.0.0.0/8"], + }, + ) + + assert result == "203.0.113.5" + + class TestMCPServerIPFiltering: """Tests that external callers only see public MCP servers.""" diff --git a/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py index 5d66c184495..a5c8320a5cc 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py @@ -1747,6 +1747,37 @@ class TestTemporaryMCPSessionEndpoints: assert isinstance(result, UserAPIKeyAuth) auth_builder_mock.assert_not_called() + def test_mcp_oauth_authorize_token_routes_use_browser_auth_dependency(self): + from fastapi.routing import APIRoute + + from litellm.proxy.management_endpoints.mcp_management_endpoints import ( + _mcp_oauth_user_api_key_auth, + router, + ) + + oauth_routes = { + route.path: route + for route in router.routes + if isinstance(route, APIRoute) + and route.path + in { + "/v1/mcp/server/oauth/{server_id}/authorize", + "/v1/mcp/server/oauth/{server_id}/token", + } + } + + assert set(oauth_routes) == { + "/v1/mcp/server/oauth/{server_id}/authorize", + "/v1/mcp/server/oauth/{server_id}/token", + } + for route in oauth_routes.values(): + dependency_names = { + dependant.name + for dependant in route.dependant.dependencies + if dependant.call is _mcp_oauth_user_api_key_auth + } + assert dependency_names == {None, "user_api_key_dict"} + @pytest.mark.asyncio async def test_mcp_authorize_proxies_to_discoverable_endpoint(self): from litellm.proxy.management_endpoints.mcp_management_endpoints import ( diff --git a/ui/litellm-dashboard/src/components/mcp_tools/MCPPermissionManagement.test.tsx b/ui/litellm-dashboard/src/components/mcp_tools/MCPPermissionManagement.test.tsx index 393c9e4a619..2cbfce320af 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/MCPPermissionManagement.test.tsx +++ b/ui/litellm-dashboard/src/components/mcp_tools/MCPPermissionManagement.test.tsx @@ -51,6 +51,68 @@ const renderWithForm = (props = {}) => { expect(toggle).toHaveAttribute("aria-checked", "false"); }); + const renderWithInitialValues = ( + initialValues: Record, + props = {}, + ) => { + const Wrapper: React.FC = ({ children }) => { + const [form] = Form.useForm(); + return ( +
+ {/* In the real app auth_type is registered by the parent form; the + component only watches it. Register a hidden field here so + Form.useWatch("auth_type") resolves the initial value. */} + + {children} +
+ ); + }; + return render( + + + , + ); + }; + + it("shows only the oauth2 PKCE-delegation toggle for oauth2 servers", async () => { + renderWithInitialValues({ allow_all_keys: false, auth_type: "oauth2" }); + await expandPanel(); + expect( + screen.getByText("Delegate auth to upstream (PKCE passthrough)"), + ).toBeInTheDocument(); + // The non-oauth2 pass-through toggle must NOT appear for oauth2 servers. + expect(screen.queryByText("OAuth pass-through")).not.toBeInTheDocument(); + }); + + it("shows only the OAuth pass-through toggle for none-auth servers forwarding Authorization", async () => { + renderWithInitialValues({ + allow_all_keys: false, + auth_type: "none", + extra_headers: ["Authorization"], + }); + await expandPanel(); + expect(screen.getByText("OAuth pass-through")).toBeInTheDocument(); + // The oauth2-only PKCE delegation toggle must NOT appear here. + expect( + screen.queryByText("Delegate auth to upstream (PKCE passthrough)"), + ).not.toBeInTheDocument(); + }); + + it("hides both upstream-auth toggles for none-auth servers without an Authorization header", async () => { + renderWithInitialValues({ + allow_all_keys: false, + auth_type: "none", + extra_headers: ["x-api-key"], + }); + await expandPanel(); + expect(screen.queryByText("OAuth pass-through")).not.toBeInTheDocument(); + expect( + screen.queryByText("Delegate auth to upstream (PKCE passthrough)"), + ).not.toBeInTheDocument(); + }); + it("should reflect allow_all_keys when editing an existing server", async () => { renderWithForm({ mcpServer: { diff --git a/ui/litellm-dashboard/src/components/mcp_tools/MCPPermissionManagement.tsx b/ui/litellm-dashboard/src/components/mcp_tools/MCPPermissionManagement.tsx index 58848df39a0..b5f0fa2e7eb 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/MCPPermissionManagement.tsx +++ b/ui/litellm-dashboard/src/components/mcp_tools/MCPPermissionManagement.tsx @@ -25,6 +25,21 @@ const MCPPermissionManagement: React.FC = ({ const form = Form.useFormInstance(); const watchedAuthType = Form.useWatch("auth_type", form); const isOAuth2 = watchedAuthType === AUTH_TYPE.OAUTH2; + const isNoneAuth = watchedAuthType === AUTH_TYPE.NONE || watchedAuthType == null; + const watchedExtraHeaders = Form.useWatch("extra_headers", form); + const hasAuthorizationHeader = Array.isArray(watchedExtraHeaders) + && watchedExtraHeaders.some( + (h) => typeof h === "string" && h.toLowerCase() === "authorization", + ); + // Two distinct, independent opt-ins: + // - delegate_auth_to_upstream: oauth2 servers only (PKCE passthrough — + // bypass LiteLLM admission). + // - oauth_passthrough: auth_type=none + Authorization in extra_headers + // (OAuth pass-through: proxy upstream oauth-protected-resource, emit 401 + // challenges, propagate upstream 401/403). + // Kept as separate flags so neither silently implies the other and existing + // oauth2 servers can't regress into pass-through behavior. + const canEnableOAuthPassthrough = isNoneAuth && hasAuthorizationHeader; const watchedDelegateAuth = Form.useWatch("delegate_auth_to_upstream", form); const watchedPublicInternet = Form.useWatch("available_on_public_internet", form); const showInternalDelegatePkceWarning = @@ -51,22 +66,34 @@ const MCPPermissionManagement: React.FC = ({ if (typeof mcpServer.delegate_auth_to_upstream === "boolean") { form.setFieldValue("delegate_auth_to_upstream", mcpServer.delegate_auth_to_upstream); } + if (typeof mcpServer.oauth_passthrough === "boolean") { + form.setFieldValue("oauth_passthrough", mcpServer.oauth_passthrough); + } } else { form.setFieldValue("allow_all_keys", false); form.setFieldValue("available_on_public_internet", true); form.setFieldValue("delegate_auth_to_upstream", false); + form.setFieldValue("oauth_passthrough", false); } }, [mcpServer, form]); - // delegate_auth_to_upstream is only honored server-side when auth_type=oauth2. + // delegate_auth_to_upstream is only honored server-side for oauth2 servers. // Force it back to false whenever the user switches away from oauth2 so a - // stale toggle value doesn't get persisted with another auth type. + // stale toggle value doesn't get persisted unexpectedly. useEffect(() => { if (!isOAuth2) { form.setFieldValue("delegate_auth_to_upstream", false); } }, [isOAuth2, form]); + // oauth_passthrough is only honored for auth_type=none servers that forward + // Authorization upstream. Force it back to false otherwise. + useEffect(() => { + if (!canEnableOAuthPassthrough) { + form.setFieldValue("oauth_passthrough", false); + } + }, [canEnableOAuthPassthrough, form]); + return ( = ({ )} + {canEnableOAuthPassthrough && ( +
+
+ + OAuth pass-through + + + + +

+ Forward upstream OAuth discovery and 401 challenges so clients negotiate OAuth directly with the upstream MCP server. +

+
+ + + +
+ )} + {showInternalDelegatePkceWarning && ( = ({ allow_all_keys: allowAllKeysRaw, available_on_public_internet: availableOnPublicInternetRaw, delegate_auth_to_upstream: delegateAuthToUpstreamRaw, + oauth_passthrough: oauthPassthroughRaw, token_validation_json: rawTokenValidationJson, ...restValues } = values; @@ -399,6 +400,7 @@ const CreateMCPServer: React.FC = ({ allow_all_keys: Boolean(allowAllKeysRaw), available_on_public_internet: Boolean(availableOnPublicInternetRaw), delegate_auth_to_upstream: Boolean(delegateAuthToUpstreamRaw), + oauth_passthrough: Boolean(oauthPassthroughRaw), static_headers: staticHeaders, ...(tokenValidation !== null && { token_validation: tokenValidation }), }; diff --git a/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_edit.test.tsx b/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_edit.test.tsx index ed3b22a569e..4070a5dd1af 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_edit.test.tsx +++ b/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_edit.test.tsx @@ -243,6 +243,41 @@ describe("MCPServerEdit (delegate auth)", () => { expect(payload.auth_type).toBe("none"); expect(payload.delegate_auth_to_upstream).toBe(false); }); + + it("does not enable oauth_passthrough for an oauth2 server", async () => { + vi.mocked(networking.updateMCPServer).mockResolvedValue({ + ...interactiveOAuthServer, + oauth_passthrough: false, + }); + + render( + , + ); + + const saveButtons = screen.getAllByRole("button", { name: "Save Changes" }); + await act(async () => { + fireEvent.click(saveButtons[0]); + }); + + await waitFor(() => { + expect(networking.updateMCPServer).toHaveBeenCalledTimes(1); + }); + + const [, payload] = vi.mocked(networking.updateMCPServer).mock.calls[0]; + expect(payload.auth_type).toBe("oauth2"); + // oauth_passthrough is non-oauth2 only — must be forced false here. + expect(payload.oauth_passthrough).toBe(false); + }); }); describe("MCPServerEdit (tool allowlist)", () => { diff --git a/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_edit.tsx b/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_edit.tsx index ab9c9ed6689..222de54f3cc 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_edit.tsx +++ b/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_edit.tsx @@ -396,6 +396,7 @@ const MCPServerEdit: React.FC = ({ allow_all_keys: allowAllKeysRaw, available_on_public_internet: availableOnPublicInternetRaw, delegate_auth_to_upstream: delegateAuthToUpstreamRaw, + oauth_passthrough: oauthPassthroughRaw, token_validation_json: rawTokenValidationJson, ...restValues } = values; @@ -573,14 +574,34 @@ const MCPServerEdit: React.FC = ({ allow_all_keys: Boolean(allowAllKeysRaw ?? mcpServer.allow_all_keys), available_on_public_internet: Boolean(availableOnPublicInternetRaw ?? mcpServer.available_on_public_internet), // ``delegate_auth_to_upstream`` is only honored server-side for - // ``auth_type=oauth2``. The Form.Item is conditionally rendered so the - // value drops out of the form on auth_type change; force false for any - // non-oauth2 server to avoid persisting a stale ``true`` that would - // silently re-activate if auth_type is later switched back to oauth2. - delegate_auth_to_upstream: - restValues.auth_type === AUTH_TYPE.OAUTH2 + // ``auth_type=oauth2`` (PKCE passthrough). The Form.Item is + // conditionally rendered so the value drops out of the form on + // auth_type change; force false for any other configuration to avoid + // persisting a stale ``true`` that would silently re-activate if the + // configuration is later switched back. + delegate_auth_to_upstream: (() => { + const isOauth2 = restValues.auth_type === AUTH_TYPE.OAUTH2; + return isOauth2 ? Boolean(delegateAuthToUpstreamRaw ?? mcpServer.delegate_auth_to_upstream) - : false, + : false; + })(), + // ``oauth_passthrough`` is the dedicated, non-oauth2 opt-in. It is only + // honored for ``auth_type=none`` servers that forward ``Authorization`` + // upstream. Kept separate from ``delegate_auth_to_upstream`` so enabling + // pass-through never regresses oauth2 servers. Force false otherwise. + oauth_passthrough: (() => { + const isNoneAuth = + restValues.auth_type === AUTH_TYPE.NONE || restValues.auth_type == null; + const extraHeaders = Array.isArray(restValues.extra_headers) + ? restValues.extra_headers + : []; + const hasAuthorizationHeader = extraHeaders.some( + (h: unknown) => typeof h === "string" && h.toLowerCase() === "authorization", + ); + return isNoneAuth && hasAuthorizationHeader + ? Boolean(oauthPassthroughRaw ?? mcpServer.oauth_passthrough) + : false; + })(), // Include token_validation when it is set (non-null) or when clearing an existing value ...(tokenValidation !== null || mcpServer.token_validation ? { token_validation: tokenValidation } : {}), }; diff --git a/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_view.tsx b/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_view.tsx index 5a8035d4e0b..87e8b77837c 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_view.tsx +++ b/ui/litellm-dashboard/src/components/mcp_tools/mcp_server_view.tsx @@ -290,6 +290,27 @@ export const MCPServerView: React.FC = ({ )} + {handleAuth(mcpServer.auth_type) !== "oauth2" && + Array.isArray(mcpServer.extra_headers) && + mcpServer.extra_headers.some( + (h) => typeof h === "string" && h.toLowerCase() === "authorization", + ) && ( +
+ OAuth Pass-through +
+ {mcpServer.oauth_passthrough ? ( + + + Enabled + + ) : ( + + Disabled + + )} +
+
+ )}
Access Groups
diff --git a/ui/litellm-dashboard/src/components/mcp_tools/types.tsx b/ui/litellm-dashboard/src/components/mcp_tools/types.tsx index 511f20baef5..3fa27afe1b2 100644 --- a/ui/litellm-dashboard/src/components/mcp_tools/types.tsx +++ b/ui/litellm-dashboard/src/components/mcp_tools/types.tsx @@ -211,6 +211,7 @@ export interface MCPServer { allow_all_keys?: boolean; available_on_public_internet?: boolean; delegate_auth_to_upstream?: boolean; + oauth_passthrough?: boolean; /** Stdio-only fields (present when transport === 'stdio') */ command?: string | null; From ebbc5cc787e64141d609fd13d474f0abc916de35 Mon Sep 17 00:00:00 2001 From: ryan-crabbe-berri Date: Tue, 2 Jun 2026 12:51:20 -0700 Subject: [PATCH 05/27] feat(vector-stores): forward per-request params to Vertex AI Search (#29459) * feat(vector-stores): forward per-request params to Vertex AI Search The vertex_ai/search_api search transform hardcoded the request body to query plus pageSize 10, dropping max_num_results and extra_body. Map max_num_results to pageSize and merge extra_body through with precedence, so callers can send native Discovery Engine fields such as dataStoreSpecs. Resolves LIT-3506 * fix(vector-stores): log effective query when extra_body overrides it When a caller passes a query inside extra_body, the outbound Vertex Search request used that value but model_call_details recorded the original, so the echoed search_query was stale. Log the effective query from the request body. * fix(vector-stores): allowlist Vertex AI Search extra_body fields Raw-merging extra_body let callers set dataStoreSpecs/branch to search a different Discovery Engine data store with the proxy's Vertex credentials, bypassing the vector_store_id path authorization. Reject target-selecting fields and forward only allowlisted per-request tuning fields. Resolves LIT-3506 * refactor(vector-stores): split Vertex AI Search extra_body allowlists by mode Data-store and engine/app serving configs accept different SearchRequest fields, so derive two TypedDicts (VertexSearchDataStoreExtraBody and VertexSearchEngineExtraBody) in types/vector_stores.py and make _filter_extra_body mode-aware via vertex_engine_id. dataStoreSpecs and numResultsPerDataStore now pass through in engine/app mode (where an app fans out across stores) and are rejected in data-store mode. branch/servingConfig/entity remain rejected in both modes. * fix(vector-stores): raise BadRequestError (400) for invalid Vertex Search extra_body Rejecting unsupported or target-selecting extra_body fields previously raised a bare ValueError, which the vector store error path mapped to a generic APIConnectionError (HTTP 500). Raise litellm.BadRequestError so invalid per-request input surfaces as HTTP 400 with a clear message. --- .../search_api/transformation.py | 126 ++++++++++++- litellm/types/vector_stores.py | 60 ++++++ ...x_ai_search_vector_store_transformation.py | 171 ++++++++++++++++++ 3 files changed, 347 insertions(+), 10 deletions(-) diff --git a/litellm/llms/vertex_ai/vector_stores/search_api/transformation.py b/litellm/llms/vertex_ai/vector_stores/search_api/transformation.py index 14a0a406dff..46dedb3d0a4 100644 --- a/litellm/llms/vertex_ai/vector_stores/search_api/transformation.py +++ b/litellm/llms/vertex_ai/vector_stores/search_api/transformation.py @@ -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 diff --git a/litellm/types/vector_stores.py b/litellm/types/vector_stores.py index ce247fc900f..6adfbf4fd35 100644 --- a/litellm/types/vector_stores.py +++ b/litellm/types/vector_stores.py @@ -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""" diff --git a/tests/test_litellm/llms/vertex_ai/test_vertex_ai_search_vector_store_transformation.py b/tests/test_litellm/llms/vertex_ai/test_vertex_ai_search_vector_store_transformation.py index 5ca71dc08c3..034f85f5a0b 100644 --- a/tests/test_litellm/llms/vertex_ai/test_vertex_ai_search_vector_store_transformation.py +++ b/tests/test_litellm/llms/vertex_ai/test_vertex_ai_search_vector_store_transformation.py @@ -1,5 +1,8 @@ +from types import SimpleNamespace + import pytest +from litellm.exceptions import BadRequestError from litellm.llms.vertex_ai.vector_stores.search_api.transformation import ( VertexSearchAPIVectorStoreConfig, ) @@ -126,3 +129,171 @@ def test_should_raise_when_neither_engine_id_nor_vector_store_id_provided(): "vertex_location": "global", }, ) + + +_ENGINE_BASE = ( + "https://discoveryengine.googleapis.com/v1/projects/p/locations/global/" + "collections/default_collection/engines/app-2/servingConfigs/default_serving_config" +) + +_DATASTORE_BASE = ( + "https://discoveryengine.googleapis.com/v1/projects/p/locations/global/" + "collections/default_collection/dataStores/ds-1/servingConfigs/default_config" +) + + +def _search_request(**overrides): + """Engine/app-mode search request (vertex_engine_id set).""" + kwargs = dict( + vector_store_id="vs", + query="hello", + vector_store_search_optional_params={}, + api_base=_ENGINE_BASE, + litellm_logging_obj=SimpleNamespace(model_call_details={}), + litellm_params={"vertex_engine_id": "app-2"}, + ) + kwargs.update(overrides) + return VertexSearchAPIVectorStoreConfig().transform_search_vector_store_request( + **kwargs + ) + + +def _datastore_search_request(**overrides): + """Data-store-mode search request (no vertex_engine_id).""" + kwargs = dict( + vector_store_id="ds-1", + query="hello", + vector_store_search_optional_params={}, + api_base=_DATASTORE_BASE, + litellm_logging_obj=SimpleNamespace(model_call_details={}), + litellm_params={}, + ) + kwargs.update(overrides) + return VertexSearchAPIVectorStoreConfig().transform_search_vector_store_request( + **kwargs + ) + + +def test_search_request_defaults_to_query_and_pagesize_10(): + url, body = _search_request() + + assert url == _ENGINE_BASE + ":search" + assert body == {"query": "hello", "pageSize": 10} + + +def test_search_request_maps_max_num_results_to_pagesize(): + _, body = _search_request( + vector_store_search_optional_params={"max_num_results": 25} + ) + + assert body["pageSize"] == 25 + + +def test_engine_search_request_forwards_datastorespecs(): + specs = [ + { + "dataStore": "projects/p/locations/global/collections/default_collection/dataStores/ds-beta" + } + ] + + _, body = _search_request(extra_body={"dataStoreSpecs": specs}) + + assert body["dataStoreSpecs"] == specs + + +def test_engine_search_request_forwards_num_results_per_data_store(): + _, body = _search_request(extra_body={"numResultsPerDataStore": 3}) + + assert body["numResultsPerDataStore"] == 3 + + +def test_datastore_search_request_rejects_datastorespecs(): + specs = [{"dataStore": "projects/p/.../dataStores/ds-beta"}] + + with pytest.raises(BadRequestError, match="data store mode"): + _datastore_search_request(extra_body={"dataStoreSpecs": specs}) + + +def test_datastore_search_request_rejects_num_results_per_data_store(): + with pytest.raises(BadRequestError, match="data store mode"): + _datastore_search_request(extra_body={"numResultsPerDataStore": 3}) + + +@pytest.mark.parametrize("field", ["branch", "servingConfig", "entity"]) +def test_search_request_rejects_target_selecting_fields(field): + with pytest.raises(BadRequestError, match="target-selecting"): + _search_request(extra_body={field: "x"}) + + +@pytest.mark.parametrize("field", ["branch", "servingConfig", "entity"]) +def test_datastore_search_request_rejects_target_selecting_fields(field): + with pytest.raises(BadRequestError, match="target-selecting"): + _datastore_search_request(extra_body={field: "x"}) + + +def test_search_request_rejects_unsupported_extra_body_field(): + with pytest.raises(BadRequestError, match="Unsupported Vertex AI Search extra_body"): + _search_request(extra_body={"notARealField": True}) + + +def test_rejected_extra_body_raises_http_400(): + with pytest.raises(BadRequestError) as exc_info: + _search_request(extra_body={"notARealField": True}) + + assert exc_info.value.status_code == 400 + + +def test_search_request_forwards_supported_extra_body_fields(): + _, body = _search_request( + extra_body={ + "filter": 'category: ANY("docs")', + "boostSpec": {"conditionBoostSpecs": []}, + } + ) + + assert body["filter"] == 'category: ANY("docs")' + assert body["boostSpec"] == {"conditionBoostSpecs": []} + assert body["query"] == "hello" + + +def test_datastore_search_request_forwards_supported_extra_body_fields(): + _, body = _datastore_search_request( + extra_body={"filter": 'category: ANY("docs")'} + ) + + assert body["filter"] == 'category: ANY("docs")' + + +def test_search_request_ignores_none_valued_extra_body_fields(): + _, body = _search_request(extra_body={"filter": None}) + + assert "filter" not in body + + +def test_search_request_extra_body_takes_precedence_over_defaults(): + _, body = _search_request( + vector_store_search_optional_params={"max_num_results": 5}, + extra_body={"pageSize": 50, "filter": 'category: ANY("docs")'}, + ) + + assert body["pageSize"] == 50 + assert body["filter"] == 'category: ANY("docs")' + + +def test_search_request_joins_list_query(): + _, body = _search_request(query=["foo", "bar"]) + + assert body["query"] == "foo bar" + + +def test_search_request_logs_effective_query_when_extra_body_overrides_query(): + log = SimpleNamespace(model_call_details={}) + + _, body = _search_request( + query="original", + extra_body={"query": "from-extra-body"}, + litellm_logging_obj=log, + ) + + assert body["query"] == "from-extra-body" + assert log.model_call_details["query"] == "from-extra-body" From 4a81ec49824c8584b6110e2deb0cc5e8af70f714 Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Wed, 3 Jun 2026 01:22:10 +0530 Subject: [PATCH 06/27] feat(proxy): add per-MCP-server RPM rate limiting for keys and teams (#29482) * feat(proxy): add per-MCP-server RPM rate limiting for keys and teams Adds mcp_rpm_limit, a dict keyed by MCP server name (alias if set, else the configured name) that caps requests per minute per server for a key or team. The v3 rate limiter builds a per-server descriptor only when a limit is configured for the server being called, so other servers stay uncapped and no TPM reservation is engaged. Server identity is surfaced into the request data via mcp_rate_limit_server_name so the limiter can resolve it. * fix(proxy): gate MCP rpm descriptors on call_mcp_tool; document mcp_rpm_limit param Only honor mcp_server_name when the call is an actual MCP tool call. Without this, a normal LLM request could inject mcp_server_name in its body to consume a target server's MCP quota and 429 legitimate tool calls. Also adds the mcp_rpm_limit parameter docstring to update_key, new_user, and user_update so the API docs validator passes. * Fix MCP rate limit quota handling * Delete scripts/test_mcp_rpm_limit.sh * docs(proxy): clarify mcp_rpm_limit is enforced for keys and teams, not per user * fix(proxy): accept mcp_rpm_limit in generate_key_helper_fn NewUserRequest and GenerateKeyRequest inherit mcp_rpm_limit from GenerateRequestBase, so /user/new and /key/generate forwarded the field to generate_key_helper_fn, which did not accept it and returned a 500 ("unexpected keyword argument 'mcp_rpm_limit'"). Accept the param and store it in metadata, matching model_rpm_limit/model_tpm_limit, so the limit is persisted where get_key_mcp_rpm_limit reads it. --------- Co-authored-by: Cursor Agent Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com> --- .../mcp_server/mcp_server_manager.py | 3 + litellm/proxy/_types.py | 4 + litellm/proxy/auth/auth_utils.py | 34 +++ .../hooks/parallel_request_limiter_v3.py | 92 ++++++- .../internal_user_endpoints.py | 2 + .../key_management_endpoints.py | 7 + .../management_endpoints/team_endpoints.py | 3 +- litellm/proxy/utils.py | 1 + .../mcp_server/test_mcp_hook_extra_headers.py | 84 +++++++ .../proxy/auth/test_auth_utils.py | 17 ++ .../hooks/test_parallel_request_limiter_v3.py | 227 ++++++++++++++++++ .../management_endpoints/test_common_utils.py | 27 +++ 12 files changed, 499 insertions(+), 2 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 129dfc102f9..739dc4a2f88 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -2785,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) diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 9f89cae1a41..09ee88239f3 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -1050,6 +1050,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 @@ -1854,6 +1855,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 ) @@ -1923,6 +1925,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 @@ -4288,6 +4291,7 @@ class PassThroughEndpointLoggingTypedDict(TypedDict): LiteLLM_ManagementEndpoint_MetadataFields = [ "model_rpm_limit", "model_tpm_limit", + "mcp_rpm_limit", "rpm_limit_type", "tpm_limit_type", "enforced_params", diff --git a/litellm/proxy/auth/auth_utils.py b/litellm/proxy/auth/auth_utils.py index 4e5169d8d84..80840c27425 100644 --- a/litellm/proxy/auth/auth_utils.py +++ b/litellm/proxy/auth/auth_utils.py @@ -940,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]]: diff --git a/litellm/proxy/hooks/parallel_request_limiter_v3.py b/litellm/proxy/hooks/parallel_request_limiter_v3.py index d03ad70562a..4343747d104 100644 --- a/litellm/proxy/hooks/parallel_request_limiter_v3.py +++ b/litellm/proxy/hooks/parallel_request_limiter_v3.py @@ -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 diff --git a/litellm/proxy/management_endpoints/internal_user_endpoints.py b/litellm/proxy/management_endpoints/internal_user_endpoints.py index 75eb5cd55ef..7b8f0f72e13 100644 --- a/litellm/proxy/management_endpoints/internal_user_endpoints.py +++ b/litellm/proxy/management_endpoints/internal_user_endpoints.py @@ -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. diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index 0e645013b92..80ded0bdd16 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -1388,6 +1388,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 +1607,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 +2424,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 +3404,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 +3483,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 diff --git a/litellm/proxy/management_endpoints/team_endpoints.py b/litellm/proxy/management_endpoints/team_endpoints.py index 8a8e703831b..ae7da0d29f2 100644 --- a/litellm/proxy/management_endpoints/team_endpoints.py +++ b/litellm/proxy/management_endpoints/team_endpoints.py @@ -863,8 +863,9 @@ async def new_team( # noqa: PLR0915 - members_with_roles: List[{"role": "admin" or "user", "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. diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index 0e72f47e224..8bd50a50a38 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -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). diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_hook_extra_headers.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_hook_extra_headers.py index cbea386a69c..04ff1e4be20 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_hook_extra_headers.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_hook_extra_headers.py @@ -826,3 +826,87 @@ class TestUserAPIKeyAuthJwtClaims: auth.jwt_claims = claims assert auth.jwt_claims == claims assert auth.jwt_claims["groups"] == ["admin"] + + +class TestMcpRateLimitServerNameSurfacing: + """ + The per-MCP-server rate limiter only sees the request `data` dict, so the + server identity must be surfaced into it. These tests pin the contract + between pre_call_tool_check, _convert_mcp_to_llm_format, and the limiter. + """ + + def setup_method(self): + self.proxy_logging = ProxyLogging(user_api_key_cache=MagicMock()) + + def test_convert_mcp_to_llm_format_surfaces_rate_limit_server_name(self): + request_obj = MagicMock() + request_obj.tool_name = "list_repos" + request_obj.arguments = {"org": "acme"} + + result = self.proxy_logging._convert_mcp_to_llm_format( + request_obj, {"mcp_rate_limit_server_name": "github"} + ) + + assert result["mcp_server_name"] == "github" + + def test_convert_mcp_to_llm_format_server_name_none_when_absent(self): + request_obj = MagicMock() + request_obj.tool_name = "list_repos" + request_obj.arguments = {} + + result = self.proxy_logging._convert_mcp_to_llm_format(request_obj, {}) + + assert result["mcp_server_name"] is None + + @pytest.mark.asyncio + async def test_pre_call_tool_check_resolves_alias_for_rate_limit(self): + """ + The rate-limit server key must be the alias when set (falling back to + server_name), matching how an admin keys mcp_rpm_limit in config. + """ + manager = MCPServerManager() + server = MCPServer( + server_id="test-id", + name="gh", + alias="gh", + server_name="github_full_name", + url="https://example.com", + transport=MCPTransport.http, + auth_type=MCPAuth.none, + ) + + captured = {} + + def capture_convert(request_obj, kwargs): + captured["kwargs"] = kwargs + return {"model": "fake"} + + proxy_logging = MagicMock(spec=ProxyLogging) + proxy_logging._create_mcp_request_object_from_kwargs = MagicMock( + return_value=MagicMock() + ) + proxy_logging._convert_mcp_to_llm_format = MagicMock( + side_effect=capture_convert + ) + proxy_logging.pre_call_hook = AsyncMock(return_value=None) + proxy_logging._convert_mcp_hook_response_to_kwargs = MagicMock( + return_value={"arguments": {}} + ) + + with patch.object(manager, "check_allowed_or_banned_tools", return_value=True): + with patch.object( + manager, + "check_tool_permission_for_key_team", + new_callable=AsyncMock, + ): + with patch.object(manager, "validate_allowed_params"): + await manager.pre_call_tool_check( + name="list_repos", + arguments={}, + server_name="github_full_name", + user_api_key_auth=None, + proxy_logging_obj=proxy_logging, + server=server, + ) + + assert captured["kwargs"]["mcp_rate_limit_server_name"] == "gh" diff --git a/tests/test_litellm/proxy/auth/test_auth_utils.py b/tests/test_litellm/proxy/auth/test_auth_utils.py index 2d40db9017e..60cf50efc75 100644 --- a/tests/test_litellm/proxy/auth/test_auth_utils.py +++ b/tests/test_litellm/proxy/auth/test_auth_utils.py @@ -14,6 +14,7 @@ from litellm.proxy.auth.auth_utils import ( abbreviate_api_key, check_complete_credentials, get_end_user_id_from_request_body, + get_key_mcp_rpm_limit, get_key_model_rpm_limit, get_key_model_tpm_limit, get_model_from_request, @@ -92,6 +93,22 @@ class TestGetKeyModelRpmLimit: assert result == {} +class TestGetKeyMcpRpmLimit: + def test_empty_dict_limits_are_returned(self): + key_override = UserAPIKeyAuth( + api_key="sk-123", + metadata={"mcp_rpm_limit": {}}, + team_metadata={"mcp_rpm_limit": {"github": 50}}, + ) + assert get_key_mcp_rpm_limit(key_override) == {} + + team_empty = UserAPIKeyAuth( + api_key="sk-123", + team_metadata={"mcp_rpm_limit": {}}, + ) + assert get_key_mcp_rpm_limit(team_empty) == {} + + class TestGetKeyModelTpmLimit: """Tests for get_key_model_tpm_limit function.""" diff --git a/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py b/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py index 3e2eb4b02c2..676f623a5dd 100644 --- a/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py +++ b/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py @@ -2893,3 +2893,230 @@ async def test_pre_call_hook_rejects_caller_supplied_stash_values(): ): leaked = [k for k in _LITELLM_STASH_KEYS if k in channel] assert not leaked, f"caller-supplied stash survived in {channel!r}: {leaked}" + + +# ----------------------- Per-MCP-server rate limiting (v3) ----------------------- + + +def _make_mcp_handler(): + local_cache = DualCache() + handler = _PROXY_MaxParallelRequestsHandler( + internal_usage_cache=InternalUsageCache(local_cache) + ) + return handler, local_cache + + +def _find_descriptor(descriptors, key): + return next((d for d in descriptors if d["key"] == key), None) + + +def _build_mcp_descriptors(handler, user_api_key_dict, data, call_type="call_mcp_tool"): + return handler._create_rate_limit_descriptors( + user_api_key_dict=user_api_key_dict, + data=data, + rpm_limit_type=None, + tpm_limit_type=None, + model_has_failures=False, + call_type=call_type, + ) + + +def test_mcp_per_key_descriptor_created_for_matching_server_v3(): + handler, _ = _make_mcp_handler() + api_key = hash_token("sk-mcp-key") + user_api_key_dict = UserAPIKeyAuth( + api_key=api_key, + metadata={"mcp_rpm_limit": {"github": 5}}, + ) + + descriptors = _build_mcp_descriptors( + handler, user_api_key_dict, {"mcp_server_name": "github"} + ) + + descriptor = _find_descriptor(descriptors, "mcp_per_key") + assert descriptor is not None + assert descriptor["value"] == f"{api_key}:github" + assert descriptor["rate_limit"]["requests_per_unit"] == 5 + # MCP tool calls have no token usage; tokens_per_unit must stay None so the + # TPM reservation path is never engaged (otherwise budget would leak). + assert descriptor["rate_limit"]["tokens_per_unit"] is None + + +def test_mcp_per_key_descriptor_skipped_for_non_matching_server_v3(): + handler, _ = _make_mcp_handler() + user_api_key_dict = UserAPIKeyAuth( + api_key=hash_token("sk-mcp-key"), + metadata={"mcp_rpm_limit": {"github": 5}}, + ) + + descriptors = _build_mcp_descriptors( + handler, user_api_key_dict, {"mcp_server_name": "slack"} + ) + + assert _find_descriptor(descriptors, "mcp_per_key") is None + + +def test_mcp_descriptor_skipped_for_non_mcp_request_v3(): + """A non-MCP request must not create an MCP descriptor even if the caller + injects mcp_server_name in the body; otherwise an LLM call could consume a + target server's MCP quota and 429 legitimate tool calls.""" + handler, _ = _make_mcp_handler() + user_api_key_dict = UserAPIKeyAuth( + api_key=hash_token("sk-mcp-key"), + metadata={"mcp_rpm_limit": {"github": 5}}, + ) + + descriptors = _build_mcp_descriptors( + handler, + user_api_key_dict, + {"model": "gpt-4", "mcp_server_name": "github"}, + call_type="completion", + ) + + assert _find_descriptor(descriptors, "mcp_per_key") is None + + +def test_mcp_descriptor_skipped_for_raw_rest_body_v3(): + handler, _ = _make_mcp_handler() + user_api_key_dict = UserAPIKeyAuth( + api_key=hash_token("sk-mcp-key"), + team_id="team-1", + metadata={"mcp_rpm_limit": {"github": 5}}, + team_metadata={"mcp_rpm_limit": {"github": 3}}, + ) + + descriptors = _build_mcp_descriptors( + handler, + user_api_key_dict, + { + "server_id": "slack", + "name": "demo-tool", + "arguments": {}, + "mcp_server_name": "github", + }, + ) + + assert _find_descriptor(descriptors, "mcp_per_key") is None + assert _find_descriptor(descriptors, "mcp_per_team") is None + + +def test_mcp_per_team_descriptor_created_from_team_metadata_v3(): + handler, _ = _make_mcp_handler() + user_api_key_dict = UserAPIKeyAuth( + api_key=hash_token("sk-mcp-key"), + team_id="team-1", + team_metadata={"mcp_rpm_limit": {"github": 3}}, + ) + + descriptors = _build_mcp_descriptors( + handler, user_api_key_dict, {"mcp_server_name": "github"} + ) + + descriptor = _find_descriptor(descriptors, "mcp_per_team") + assert descriptor is not None + assert descriptor["value"] == "team-1:github" + assert descriptor["rate_limit"]["requests_per_unit"] == 3 + assert descriptor["rate_limit"]["tokens_per_unit"] is None + + +@pytest.mark.asyncio +async def test_mcp_per_key_rpm_enforced_v3(monkeypatch): + """ + A key configured with mcp_rpm_limit={"github": 2} must allow 2 calls to the + github MCP server within the window and reject the 3rd with a 429, while + calls to a different MCP server are unaffected. + """ + monkeypatch.setenv("LITELLM_RATE_LIMIT_WINDOW_SIZE", "60") + api_key = hash_token("sk-mcp-enforce") + local_cache = DualCache() + handler = _PROXY_MaxParallelRequestsHandler( + internal_usage_cache=InternalUsageCache(local_cache) + ) + + window_starts: Dict[str, int] = {} + request_counts: Dict[str, int] = {} + + async def mock_batch_rate_limiter(*args, **kwargs): + keys = kwargs.get("keys") if kwargs else args[0] + args_list = kwargs.get("args") if kwargs else args[1] + now = args_list[0] + window_size = args_list[1] + results = [] + for i in range(0, len(keys), 2): + window_key = keys[i] + counter_key = keys[i + 1] + prev_window = window_starts.get(window_key) + prev_counter = request_counts.get(counter_key, 0) + if prev_window is None or (now - prev_window) >= window_size: + window_starts[window_key] = now + new_counter = 1 + else: + new_counter = prev_counter + 1 + request_counts[counter_key] = new_counter + results.append(now) + results.append(new_counter) + return results + + handler.batch_rate_limiter_script = mock_batch_rate_limiter + + user_api_key_dict = UserAPIKeyAuth( + api_key=api_key, + metadata={"mcp_rpm_limit": {"github": 2}}, + ) + + for _ in range(2): + await handler.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=local_cache, + data={"mcp_server_name": "github"}, + call_type="call_mcp_tool", + ) + + with pytest.raises(HTTPException) as exc_info: + await handler.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=local_cache, + data={"mcp_server_name": "github"}, + call_type="call_mcp_tool", + ) + assert exc_info.value.status_code == 429 + + # A different server has no configured limit -> not rate limited. + for _ in range(5): + await handler.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=local_cache, + data={"mcp_server_name": "slack"}, + call_type="call_mcp_tool", + ) + + # The TPM counter must never be created for an MCP descriptor. + assert not any(":tokens" in key and "github" in key for key in request_counts) + + +def test_get_key_mcp_rpm_limit_precedence(): + from litellm.proxy.auth.auth_utils import ( + get_key_mcp_rpm_limit, + get_team_mcp_rpm_limit, + ) + + # Key metadata takes precedence over team metadata. + key_first = UserAPIKeyAuth( + api_key=hash_token("sk-mcp-key"), + metadata={"mcp_rpm_limit": {"github": 10}}, + team_metadata={"mcp_rpm_limit": {"github": 99}}, + ) + assert get_key_mcp_rpm_limit(key_first) == {"github": 10} + + # Falls back to team metadata when key has none. + team_only = UserAPIKeyAuth( + api_key=hash_token("sk-mcp-key"), + team_metadata={"mcp_rpm_limit": {"github": 7}}, + ) + assert get_key_mcp_rpm_limit(team_only) == {"github": 7} + assert get_team_mcp_rpm_limit(team_only) == {"github": 7} + + # No configuration anywhere. + none_set = UserAPIKeyAuth(api_key=hash_token("sk-mcp-key")) + assert get_key_mcp_rpm_limit(none_set) is None + assert get_team_mcp_rpm_limit(none_set) is None diff --git a/tests/test_litellm/proxy/management_endpoints/test_common_utils.py b/tests/test_litellm/proxy/management_endpoints/test_common_utils.py index f898763d2cb..d53ea6fa34d 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_common_utils.py +++ b/tests/test_litellm/proxy/management_endpoints/test_common_utils.py @@ -482,6 +482,33 @@ class TestSetObjectMetadataField: _set_object_metadata_field(team, "model_rpm_limit", {"x": 1}) assert team.metadata == {"model_rpm_limit": {"x": 1}} + def test_mcp_rpm_limit_is_hoisted_into_metadata(self): + """ + Per-MCP-server rpm limits are stored in the metadata JSON column, not a + dedicated DB column. The key/team management endpoints rely on + LiteLLM_ManagementEndpoint_MetadataFields to move the request field into + metadata; this regression guards that mcp_rpm_limit is in that list and + round-trips through the same loop the endpoints use. + """ + from litellm.proxy._types import LiteLLM_ManagementEndpoint_MetadataFields + + assert "mcp_rpm_limit" in LiteLLM_ManagementEndpoint_MetadataFields + + from types import SimpleNamespace + + team = LiteLLM_TeamTable(team_id="t1", metadata={}) + mcp_rpm_limit = {"github": 100} + data = SimpleNamespace(mcp_rpm_limit=mcp_rpm_limit) + + with patch( + "litellm.proxy.management_endpoints.common_utils._premium_user_check" + ): + for field in LiteLLM_ManagementEndpoint_MetadataFields: + if getattr(data, field, None) is not None: + _set_object_metadata_field(team, field, getattr(data, field)) + + assert team.metadata["mcp_rpm_limit"] == mcp_rpm_limit + class TestRequireCallerUserIdForNonAdmin: """ From c1602587c1da679ee47bb5f47cac7b492a32b7f4 Mon Sep 17 00:00:00 2001 From: Mateo Wang <277851410+mateo-berri@users.noreply.github.com> Date: Tue, 2 Jun 2026 13:07:05 -0700 Subject: [PATCH 07/27] fix(tests): drop module-level test calls that break local_testing collection (#29520) * fix(tests): drop module-level test calls that break local_testing collection Several files in tests/local_testing invoked their test functions at module scope (e.g. test_register_model.py ran test_update_model_cost_via_completion() at the bottom of the file). Those calls execute during pytest collection, so they fire real network requests at import time. test_register_model.py's call hit an OpenAI 429 and raised, turning into a collection error. A collection error aborts the whole session for every job that globs tests/local_testing/**/test_*.py, which is why unrelated jobs like langfuse_logging_unit_tests (-k langfuse) and litellm_assistants_api_testing (-k assistants) both failed even though neither touches register_model; the -k filter only applies after collection. pytest discovers and runs these test_* functions on its own, so the top-level calls were dead and harmful. Removes them from test_register_model.py, test_wandb.py, test_lunary.py, and test_multiple_deployments.py, and adds a regression test that scans the directory for module-level test invocations. * test(local_testing): skip unparseable files in module-scope invocation guardrail A syntax error in any tests/local_testing file would make ast.parse raise an unhandled SyntaxError, so the guardrail itself would crash with a confusing traceback instead of its assertion message. Such a file already fails pytest collection on its own, which is the clearer signal, so the guardrail now skips files it cannot parse and stays focused on detecting module-scope test calls. Reads files as utf-8 for deterministic behavior across platforms. --- tests/local_testing/test_lunary.py | 3 -- .../test_multiple_deployments.py | 3 -- .../test_no_top_level_test_invocations.py | 36 +++++++++++++++++++ tests/local_testing/test_register_model.py | 3 -- tests/local_testing/test_wandb.py | 3 -- 5 files changed, 36 insertions(+), 12 deletions(-) create mode 100644 tests/local_testing/test_no_top_level_test_invocations.py diff --git a/tests/local_testing/test_lunary.py b/tests/local_testing/test_lunary.py index d181d24c782..0dbae1b817f 100644 --- a/tests/local_testing/test_lunary.py +++ b/tests/local_testing/test_lunary.py @@ -26,9 +26,6 @@ def test_lunary_logging(): print(e) -test_lunary_logging() - - def test_lunary_template(): import lunary diff --git a/tests/local_testing/test_multiple_deployments.py b/tests/local_testing/test_multiple_deployments.py index f7276d4f14e..72bfd5012c1 100644 --- a/tests/local_testing/test_multiple_deployments.py +++ b/tests/local_testing/test_multiple_deployments.py @@ -49,6 +49,3 @@ def test_multiple_deployments(): except Exception as e: traceback.print_exc() pytest.fail(f"An exception occurred: {e}") - - -test_multiple_deployments() diff --git a/tests/local_testing/test_no_top_level_test_invocations.py b/tests/local_testing/test_no_top_level_test_invocations.py new file mode 100644 index 00000000000..eb1d836a18d --- /dev/null +++ b/tests/local_testing/test_no_top_level_test_invocations.py @@ -0,0 +1,36 @@ +import ast +from pathlib import Path + +LOCAL_TESTING_DIR = Path(__file__).parent + + +def _top_level_test_invocations(tree): + invocations = [] + for node in tree.body: + if not isinstance(node, ast.Expr) or not isinstance(node.value, ast.Call): + continue + func = node.value.func + name = getattr(func, "id", None) or getattr(func, "attr", None) + if name and name.startswith("test_"): + invocations.append((name, node.lineno)) + return invocations + + +def test_no_module_level_test_invocations(): + offenders = [] + for path in sorted(LOCAL_TESTING_DIR.rglob("*.py")): + try: + tree = ast.parse(path.read_text(encoding="utf-8"), filename=str(path)) + except SyntaxError: + continue + for name, lineno in _top_level_test_invocations(tree): + offenders.append( + f"{path.relative_to(LOCAL_TESTING_DIR)}:{lineno} calls {name}()" + ) + + assert not offenders, ( + "Test functions are invoked at module scope, so they run during pytest " + "collection (making network calls and erroring collection for every job " + "that globs this directory). Remove these calls; pytest collects test " + "functions automatically:\n" + "\n".join(offenders) + ) diff --git a/tests/local_testing/test_register_model.py b/tests/local_testing/test_register_model.py index 6b170798874..635fd79abff 100644 --- a/tests/local_testing/test_register_model.py +++ b/tests/local_testing/test_register_model.py @@ -60,6 +60,3 @@ def test_update_model_cost_via_completion(): assert litellm.model_cost["gpt-3.5-turbo"]["output_cost_per_token"] == 0.4 except Exception as e: pytest.fail(f"An error occurred: {e}") - - -test_update_model_cost_via_completion() diff --git a/tests/local_testing/test_wandb.py b/tests/local_testing/test_wandb.py index 6cdca40492f..58a9c9f5ddf 100644 --- a/tests/local_testing/test_wandb.py +++ b/tests/local_testing/test_wandb.py @@ -51,9 +51,6 @@ def test_wandb_logging_async(): pass -test_wandb_logging_async() - - def test_wandb_logging(): try: response = completion( From ae7ac72331ef21d990fece0cb8cd66ab1d594af2 Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Wed, 3 Jun 2026 03:15:56 +0530 Subject: [PATCH 08/27] feat(agents): add LangFlow agent provider with A2A session bridging (#28963) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit * feat(agents): add LangFlow agent provider with A2A session bridging Register LangFlow as a completion provider and agent type (UI + /api/v1/run), and map A2A contextId to LangFlow session_id for multi-turn conversations. Co-authored-by: Cursor * docs(providers): document langflow in provider_endpoints_support.json Co-authored-by: Cursor * fix(agents): address Greptile review for LangFlow integration Move A2A contextId→session_id mapping into LangFlow A2A provider config, add langflow.svg logo, remove live integration test, use model for token count. Co-authored-by: Cursor * fix(langflow): prevent flow_id override via request optional_params Derive flow_id only from the authorized model name and reject flow_id kwargs so callers cannot invoke a different LangFlow run endpoint. Co-authored-by: Cursor * refactor(langflow): remove redundant flow_id branch in _get_flow_id * fix(langflow): surface an error when the run response has no extractable message Previously the response parser returned the raw JSON blob as the assistant message when it could not find message text, silently presenting an unparseable payload as a valid answer. It now returns None and the caller raises a LangFlowError so the failure is visible to the client. * fix(langflow): URL-encode flow_id path segment to prevent path injection flow_id is taken from the model suffix and interpolated into /api/v1/run/{flow_id}. Without path-segment encoding a model such as langflow/../../x (or one containing ?) could move the request off the run endpoint to another path on the configured LangFlow server using the operator x-api-key. Encode the segment with quote(safe="") so it always stays a single path segment. * fix(langflow): reject empty flow_id from model name * fix(langflow): return stripped flow_id so validation matches URL path * fix(langflow): reject caller-supplied tweaks to prevent flow component override * fix(langflow): reject caller-supplied tweaks injected via extra_body The transform_request guard only inspected optional_params, but extra_body is popped before transform_request runs and merged into the request body afterward, letting a caller reintroduce tweaks and override the operator-configured LangFlow flow components. Validate the final request body in sign_request so tweaks cannot reach LangFlow through extra_body. * test(langflow): move provider tests into mirrored coverage path The langflow tests lived under tests/llm_translation/, whose CircleCI job runs without --cov and uploads nothing to Codecov, so none of the new langflow code counted toward patch coverage (codecov/patch reported 9.78% of the diff hit against a 70.83% target). Relocate them to tests/test_litellm/llms/langflow/, which the GitHub Actions provider job runs with --cov=./litellm and uploads, and add regression tests for the previously untested happy paths (transform_response building the ModelResponse with usage, non-JSON body handling, last-user message extraction, outputs-dict response shape, sign_request pass-through, error class and stream flags). Patch coverage on the diff is now ~88%. * fix(langflow): require litellm_params in A2A config instead of silent empty fallback * fix(langflow): scope A2A session_id to the authenticated key The LangFlow A2A bridge used the LangFlow session_id verbatim from the client-controlled A2A contextId, so two distinct virtual keys authorized for the same agent could read or append to each other's LangFlow conversation memory by reusing a contextId. Hand the authenticated key hash to the completion bridge through litellm_params and namespace the forwarded session_id with it. The same key keeps a stable session across turns, while different keys can no longer collide on a shared contextId. The principal is hashed before it is embedded in the session_id, so the stored token is never sent to the LangFlow backend; the original contextId is preserved as a suffix for operator-side correlation. * fix(langflow): wire authenticated key hash through A2A bridge and tests Define A2A_USER_API_KEY_HASH_PARAM in the completion bridge handler, strip it before litellm.acompletion, inject the authenticated key hash at the proxy A2A endpoint, and add regression tests for per-key LangFlow session scoping. --------- Co-authored-by: Cursor Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com> --- .../litellm_completion_bridge/handler.py | 83 ++-- .../a2a_protocol/providers/config_manager.py | 5 + .../providers/langflow/__init__.py | 0 .../a2a_protocol/providers/langflow/config.py | 62 +++ litellm/llms/langflow/__init__.py | 1 + litellm/llms/langflow/a2a.py | 37 ++ litellm/llms/langflow/chat/__init__.py | 1 + litellm/llms/langflow/chat/transformation.py | 327 ++++++++++++++ litellm/main.py | 33 ++ .../proxy/agent_endpoints/a2a_endpoints.py | 13 + .../public_endpoints/agent_create_fields.json | 42 ++ litellm/types/utils.py | 1 + litellm/utils.py | 11 + provider_endpoints_support.json | 18 + .../chat/test_langflow_chat_transformation.py | 398 ++++++++++++++++++ .../llms/langflow/test_langflow_a2a.py | 159 +++++++ .../agent_endpoints/test_a2a_endpoints.py | 102 +++++ .../public/assets/logos/langflow.svg | 5 + .../src/components/agents/agent_type_utils.ts | 2 + 19 files changed, 1265 insertions(+), 35 deletions(-) create mode 100644 litellm/a2a_protocol/providers/langflow/__init__.py create mode 100644 litellm/a2a_protocol/providers/langflow/config.py create mode 100644 litellm/llms/langflow/__init__.py create mode 100644 litellm/llms/langflow/a2a.py create mode 100644 litellm/llms/langflow/chat/__init__.py create mode 100644 litellm/llms/langflow/chat/transformation.py create mode 100644 tests/test_litellm/llms/langflow/chat/test_langflow_chat_transformation.py create mode 100644 tests/test_litellm/llms/langflow/test_langflow_a2a.py create mode 100644 ui/litellm-dashboard/public/assets/logos/langflow.svg diff --git a/litellm/a2a_protocol/litellm_completion_bridge/handler.py b/litellm/a2a_protocol/litellm_completion_bridge/handler.py index 67ffcf4f8f7..52e471ff702 100644 --- a/litellm/a2a_protocol/litellm_completion_bridge/handler.py +++ b/litellm/a2a_protocol/litellm_completion_bridge/handler.py @@ -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", {}) diff --git a/litellm/a2a_protocol/providers/config_manager.py b/litellm/a2a_protocol/providers/config_manager.py index ecb8f66bdeb..a421afec184 100644 --- a/litellm/a2a_protocol/providers/config_manager.py +++ b/litellm/a2a_protocol/providers/config_manager.py @@ -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, diff --git a/litellm/a2a_protocol/providers/langflow/__init__.py b/litellm/a2a_protocol/providers/langflow/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/litellm/a2a_protocol/providers/langflow/config.py b/litellm/a2a_protocol/providers/langflow/config.py new file mode 100644 index 00000000000..9302c38126b --- /dev/null +++ b/litellm/a2a_protocol/providers/langflow/config.py @@ -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 diff --git a/litellm/llms/langflow/__init__.py b/litellm/llms/langflow/__init__.py new file mode 100644 index 00000000000..d1270fc91f5 --- /dev/null +++ b/litellm/llms/langflow/__init__.py @@ -0,0 +1 @@ +"""LangFlow LLM provider for LiteLLM.""" diff --git a/litellm/llms/langflow/a2a.py b/litellm/llms/langflow/a2a.py new file mode 100644 index 00000000000..dbe3e02401d --- /dev/null +++ b/litellm/llms/langflow/a2a.py @@ -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 diff --git a/litellm/llms/langflow/chat/__init__.py b/litellm/llms/langflow/chat/__init__.py new file mode 100644 index 00000000000..286b12e31f1 --- /dev/null +++ b/litellm/llms/langflow/chat/__init__.py @@ -0,0 +1 @@ +"""LangFlow chat transformation.""" diff --git a/litellm/llms/langflow/chat/transformation.py b/litellm/llms/langflow/chat/transformation.py new file mode 100644 index 00000000000..f898163ad02 --- /dev/null +++ b/litellm/llms/langflow/chat/transformation.py @@ -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": "", + "input_type": "chat", + "output_type": "chat", + "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 diff --git a/litellm/main.py b/litellm/main.py index 09c70998cf7..3ef094042e8 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -4503,6 +4503,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 diff --git a/litellm/proxy/agent_endpoints/a2a_endpoints.py b/litellm/proxy/agent_endpoints/a2a_endpoints.py index 993d30e3811..7b56155982c 100644 --- a/litellm/proxy/agent_endpoints/a2a_endpoints.py +++ b/litellm/proxy/agent_endpoints/a2a_endpoints.py @@ -389,6 +389,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: diff --git a/litellm/proxy/public_endpoints/agent_create_fields.json b/litellm/proxy/public_endpoints/agent_create_fields.json index e58bd97cce7..36484cc1065 100644 --- a/litellm/proxy/public_endpoints/agent_create_fields.json +++ b/litellm/proxy/public_endpoints/agent_create_fields.json @@ -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", diff --git a/litellm/types/utils.py b/litellm/types/utils.py index d3c2c8c18fe..63c2513aed2 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -3364,6 +3364,7 @@ class LlmProviders(str, Enum): AMAZON_NOVA = "amazon_nova" A2A_AGENT = "a2a_agent" LANGGRAPH = "langgraph" + LANGFLOW = "langflow" MINIMAX = "minimax" SYNTHETIC = "synthetic" APERTIS = "apertis" diff --git a/litellm/utils.py b/litellm/utils.py index 68982ea2b35..6188206148f 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -8394,6 +8394,10 @@ class ProviderConfigManager: lambda: ProviderConfigManager._get_langgraph_config(), False, ), + LlmProviders.LANGFLOW: ( + lambda: ProviderConfigManager._get_langflow_config(), + False, + ), } @staticmethod @@ -8465,6 +8469,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, diff --git a/provider_endpoints_support.json b/provider_endpoints_support.json index 1e8357a8137..3a01541060f 100644 --- a/provider_endpoints_support.json +++ b/provider_endpoints_support.json @@ -2430,6 +2430,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", diff --git a/tests/test_litellm/llms/langflow/chat/test_langflow_chat_transformation.py b/tests/test_litellm/llms/langflow/chat/test_langflow_chat_transformation.py new file mode 100644 index 00000000000..c03919a0659 --- /dev/null +++ b/tests/test_litellm/llms/langflow/chat/test_langflow_chat_transformation.py @@ -0,0 +1,398 @@ +from unittest.mock import MagicMock, patch + +import httpx +import pytest + +from litellm.llms.langflow.chat.transformation import LangFlowConfig, LangFlowError +from litellm.types.utils import LlmProviders, ModelResponse +from litellm.utils import ProviderConfigManager + + +def test_flow_id_cannot_be_overridden_via_optional_params(): + config = LangFlowConfig() + url = config.get_complete_url( + api_base="http://localhost:7860", + api_key=None, + model="langflow/authorized-flow", + optional_params={}, + litellm_params={}, + stream=False, + ) + assert url.endswith("/api/v1/run/authorized-flow") + + with pytest.raises(LangFlowError): + config.get_complete_url( + api_base="http://localhost:7860", + api_key=None, + model="langflow/authorized-flow", + optional_params={"flow_id": "malicious-flow"}, + litellm_params={}, + stream=False, + ) + + +def test_langflow_config_get_complete_url(): + config = LangFlowConfig() + url = config.get_complete_url( + api_base="http://localhost:7860", + api_key=None, + model="langflow/my-flow-id", + optional_params={}, + litellm_params={}, + stream=False, + ) + assert url == "http://localhost:7860/api/v1/run/my-flow-id" + + +def test_langflow_config_get_complete_url_requires_api_base(): + config = LangFlowConfig() + with pytest.raises(ValueError): + config.get_complete_url( + api_base=None, + api_key=None, + model="langflow/my-flow-id", + optional_params={}, + litellm_params={}, + stream=False, + ) + + +def test_langflow_config_flow_id_is_path_segment_encoded(): + config = LangFlowConfig() + url = config.get_complete_url( + api_base="http://localhost:7860", + api_key=None, + model="langflow/../../secret?x=1", + optional_params={}, + litellm_params={}, + stream=False, + ) + assert url == "http://localhost:7860/api/v1/run/..%2F..%2Fsecret%3Fx%3D1" + assert "/api/v1/run/" in url + assert url.rsplit("/api/v1/run/", 1)[1] not in ("..", "../..") + + +@pytest.mark.parametrize("model", ["langflow/", "langflow/ "]) +def test_langflow_config_rejects_empty_flow_id(model): + config = LangFlowConfig() + with pytest.raises(LangFlowError): + config.get_complete_url( + api_base="http://localhost:7860", + api_key=None, + model=model, + optional_params={}, + litellm_params={}, + stream=False, + ) + + +def test_langflow_config_strips_flow_id_whitespace(): + config = LangFlowConfig() + url = config.get_complete_url( + api_base="http://localhost:7860", + api_key=None, + model="langflow/ my-flow-id ", + optional_params={}, + litellm_params={}, + stream=False, + ) + assert url == "http://localhost:7860/api/v1/run/my-flow-id" + + +def test_langflow_config_transform_request_includes_session_id(): + config = LangFlowConfig() + request = config.transform_request( + model="langflow/my-flow-id", + messages=[{"role": "user", "content": "hello"}], + optional_params={"session_id": "sess-abc"}, + litellm_params={}, + headers={}, + ) + + assert request["input_value"] == "hello" + assert request["input_type"] == "chat" + assert request["output_type"] == "chat" + assert request["session_id"] == "sess-abc" + + +def test_langflow_config_transform_request_uses_last_user_message(): + config = LangFlowConfig() + request = config.transform_request( + model="langflow/my-flow-id", + messages=[ + {"role": "system", "content": "be helpful"}, + {"role": "user", "content": "first"}, + {"role": "assistant", "content": "ok"}, + {"role": "user", "content": [{"type": "text", "text": "second"}]}, + ], + optional_params={}, + litellm_params={}, + headers={}, + ) + + assert request["input_value"] == "second" + assert "session_id" not in request + + +def test_langflow_config_transform_request_falls_back_to_last_message(): + config = LangFlowConfig() + request = config.transform_request( + model="langflow/my-flow-id", + messages=[{"role": "assistant", "content": "only assistant"}], + optional_params={}, + litellm_params={}, + headers={}, + ) + + assert request["input_value"] == "only assistant" + + +def test_langflow_config_transform_request_empty_messages(): + config = LangFlowConfig() + request = config.transform_request( + model="langflow/my-flow-id", + messages=[], + optional_params={}, + litellm_params={}, + headers={}, + ) + + assert request["input_value"] == "" + + +def test_langflow_config_rejects_tweaks_from_request_params(): + config = LangFlowConfig() + with pytest.raises(LangFlowError): + config.transform_request( + model="langflow/my-flow-id", + messages=[{"role": "user", "content": "hi"}], + optional_params={"tweaks": {"HttpComponent": {"url": "http://attacker"}}}, + litellm_params={}, + headers={}, + ) + + +def test_langflow_config_rejects_tweaks_from_request_body(): + config = LangFlowConfig() + with pytest.raises(LangFlowError): + config.sign_request( + headers={}, + optional_params={}, + request_data={ + "input_value": "hi", + "tweaks": {"HttpComponent": {"url": "http://attacker"}}, + }, + api_base="http://localhost:7860", + ) + + +def test_langflow_config_sign_request_passes_through_without_tweaks(): + config = LangFlowConfig() + headers, body = config.sign_request( + headers={"x-api-key": "secret"}, + optional_params={}, + request_data={"input_value": "hi"}, + api_base="http://localhost:7860", + ) + assert headers == {"x-api-key": "secret"} + assert body is None + + +def test_langflow_config_validate_environment_sets_api_key_header(): + config = LangFlowConfig() + headers = config.validate_environment( + headers={}, + model="langflow/my-flow-id", + messages=[{"role": "user", "content": "hi"}], + optional_params={}, + litellm_params={}, + api_key="secret", + ) + assert headers["Content-Type"] == "application/json" + assert headers["x-api-key"] == "secret" + + +def test_langflow_extra_body_cannot_inject_tweaks_into_run_payload(): + import json + + import litellm + from litellm.llms.custom_httpx.http_handler import HTTPHandler + + posted_bodies = [] + + def fake_post(*args, **kwargs): + body = kwargs.get("data") + posted_bodies.append(json.loads(body) if isinstance(body, str) else body) + resp = MagicMock(spec=httpx.Response) + resp.status_code = 200 + resp.json.return_value = { + "outputs": [{"outputs": [{"results": {"message": {"text": "hi"}}}]}] + } + resp.headers = {} + resp.text = "{}" + return resp + + with patch.object(HTTPHandler, "post", side_effect=fake_post): + with pytest.raises(Exception): + litellm.completion( + model="langflow/my-flow", + messages=[{"role": "user", "content": "hello"}], + api_base="http://example.com", + api_key="sk-test", + extra_body={"tweaks": {"HttpComponent": {"url": "http://attacker"}}}, + ) + + assert all("tweaks" not in (body or {}) for body in posted_bodies) + + +def test_langflow_config_extract_response(): + config = LangFlowConfig() + content = config._extract_content_from_response( + { + "session_id": "sess-abc", + "outputs": [ + { + "outputs": [ + { + "results": { + "message": {"text": "Hello from LangFlow"}, + } + } + ] + } + ], + } + ) + assert content == "Hello from LangFlow" + + +def test_langflow_config_extract_response_from_outputs_dict(): + config = LangFlowConfig() + content = config._extract_content_from_response( + { + "outputs": [ + { + "outputs": [ + { + "results": {}, + "outputs": { + "message": {"message": {"text": "via outputs dict"}} + }, + } + ] + } + ], + } + ) + assert content == "via outputs dict" + + +def test_langflow_extract_response_returns_none_when_no_message(): + config = LangFlowConfig() + assert config._extract_content_from_response({"outputs": []}) is None + assert config._extract_content_from_response({"detail": "flow failed"}) is None + assert config._extract_content_from_response({"outputs": ["not-a-dict"]}) is None + assert ( + config._extract_content_from_response({"outputs": [{"outputs": ["bad"]}]}) + is None + ) + assert ( + config._extract_content_from_response( + {"outputs": [{"outputs": [{"results": {"message": {"text": ""}}}]}]} + ) + is None + ) + + +def test_langflow_transform_response_builds_model_response_with_usage(): + config = LangFlowConfig() + raw_response = httpx.Response( + status_code=200, + json={ + "session_id": "sess-abc", + "outputs": [ + {"outputs": [{"results": {"message": {"text": "Hello from LangFlow"}}}]} + ], + }, + ) + + result = config.transform_response( + model="langflow/my-flow-id", + raw_response=raw_response, + model_response=ModelResponse(), + logging_obj=None, + request_data={}, + messages=[{"role": "user", "content": "hi"}], + optional_params={}, + litellm_params={}, + encoding=None, + ) + + assert result.choices[0].message.content == "Hello from LangFlow" + assert result.choices[0].finish_reason == "stop" + assert result.model == "langflow/my-flow-id" + assert result.usage.completion_tokens > 0 + assert result.usage.total_tokens == ( + result.usage.prompt_tokens + result.usage.completion_tokens + ) + + +def test_langflow_transform_response_raises_on_unparseable_body(): + config = LangFlowConfig() + raw_response = httpx.Response(status_code=200, json={"detail": "flow failed"}) + + with pytest.raises(LangFlowError): + config.transform_response( + model="langflow/my-flow-id", + raw_response=raw_response, + model_response=ModelResponse(), + logging_obj=None, + request_data={}, + messages=[{"role": "user", "content": "hi"}], + optional_params={}, + litellm_params={}, + encoding=None, + ) + + +def test_langflow_transform_response_raises_on_non_json_body(): + config = LangFlowConfig() + raw_response = httpx.Response( + status_code=200, content=b"not json", headers={"content-type": "text/plain"} + ) + + with pytest.raises(LangFlowError): + config.transform_response( + model="langflow/my-flow-id", + raw_response=raw_response, + model_response=ModelResponse(), + logging_obj=None, + request_data={}, + messages=[{"role": "user", "content": "hi"}], + optional_params={}, + litellm_params={}, + encoding=None, + ) + + +def test_langflow_config_get_error_class(): + config = LangFlowConfig() + err = config.get_error_class(error_message="boom", status_code=503, headers={}) + assert isinstance(err, LangFlowError) + assert err.status_code == 503 + + +def test_langflow_config_stream_behavior_flags(): + config = LangFlowConfig() + assert config.supports_stream_param_in_request_body is False + assert config.should_fake_stream(model="langflow/x", stream=True) is True + assert config.should_fake_stream(model="langflow/x", stream=False) is False + + +def test_langflow_provider_config_registered(): + cfg = ProviderConfigManager.get_provider_chat_config( + model="langflow/flow-1", + provider=LlmProviders.LANGFLOW, + ) + assert cfg is not None + assert cfg.__class__.__name__ == "LangFlowConfig" diff --git a/tests/test_litellm/llms/langflow/test_langflow_a2a.py b/tests/test_litellm/llms/langflow/test_langflow_a2a.py new file mode 100644 index 00000000000..c49ec8d87c2 --- /dev/null +++ b/tests/test_litellm/llms/langflow/test_langflow_a2a.py @@ -0,0 +1,159 @@ +from unittest.mock import AsyncMock, patch + +import pytest + +from litellm.a2a_protocol.litellm_completion_bridge.handler import ( + A2A_USER_API_KEY_HASH_PARAM, +) +from litellm.a2a_protocol.providers.config_manager import A2AProviderConfigManager +from litellm.llms.langflow.a2a import merge_a2a_session_into_litellm_params + + +def test_merge_a2a_session_into_litellm_params(): + merged = merge_a2a_session_into_litellm_params( + {"custom_llm_provider": "langflow", "model": "langflow/flow-1"}, + {"message": {"contextId": "shared-session-99"}}, + ) + assert merged["session_id"] == "shared-session-99" + + +def test_merge_a2a_session_is_scoped_per_principal(): + """The LangFlow session must be bound to the authenticated key so two + distinct keys cannot share memory by reusing the same A2A contextId, while + the same key keeps a stable session across turns.""" + base = {"custom_llm_provider": "langflow", "model": "langflow/flow-1"} + params = {"message": {"contextId": "ctx-1"}} + + key_a = merge_a2a_session_into_litellm_params(base, params, "hash-a")["session_id"] + key_a_again = merge_a2a_session_into_litellm_params(base, params, "hash-a")[ + "session_id" + ] + key_b = merge_a2a_session_into_litellm_params(base, params, "hash-b")["session_id"] + + assert key_a == key_a_again, "same key + contextId must stay on one session" + assert key_a != key_b, "different keys must not collide on the same contextId" + assert key_a != "ctx-1", "raw client contextId must not be used verbatim" + assert key_a.endswith("-ctx-1"), "original contextId kept for correlation" + assert "hash-a" not in key_a, "raw principal must not be sent to LangFlow" + + +def test_merge_a2a_session_without_context_id_is_noop(): + merged = merge_a2a_session_into_litellm_params( + {"custom_llm_provider": "langflow", "model": "langflow/flow-1"}, + {"message": {"role": "user"}}, + ) + assert "session_id" not in merged + + +def test_langflow_a2a_provider_config_registered(): + cfg = A2AProviderConfigManager.get_provider_config( + custom_llm_provider="langflow", + model="langflow/flow-1", + ) + assert cfg is not None + assert cfg.__class__.__name__ == "LangFlowA2AConfig" + + +@pytest.mark.asyncio +async def test_langflow_a2a_config_passes_session_id_to_completion(): + from litellm.a2a_protocol.providers.langflow.config import LangFlowA2AConfig + + mock_response = type( + "R", + (), + { + "choices": [ + type( + "C", + (), + {"message": type("M", (), {"content": "ok"})()}, + )() + ] + }, + )() + + with patch("litellm.acompletion", new_callable=AsyncMock) as mock_acompletion: + mock_acompletion.return_value = mock_response + + await LangFlowA2AConfig().handle_non_streaming( + request_id="req-1", + params={ + "message": { + "role": "user", + "parts": [{"kind": "text", "text": "hi"}], + "contextId": "shared-session-99", + } + }, + litellm_params={ + "custom_llm_provider": "langflow", + "model": "langflow/flow-1", + "api_base": "http://localhost:7860", + }, + api_base="http://localhost:7860", + ) + + assert ( + mock_acompletion.call_args.kwargs.get("session_id") == "shared-session-99" + ) + + +@pytest.mark.asyncio +async def test_langflow_a2a_config_scopes_session_by_authenticated_key(): + from litellm.a2a_protocol.providers.langflow.config import LangFlowA2AConfig + + mock_response = type( + "R", + (), + {"choices": [type("C", (), {"message": type("M", (), {"content": "ok"})()})()]}, + )() + + with patch("litellm.acompletion", new_callable=AsyncMock) as mock_acompletion: + mock_acompletion.return_value = mock_response + + await LangFlowA2AConfig().handle_non_streaming( + request_id="req-1", + params={ + "message": { + "role": "user", + "parts": [{"kind": "text", "text": "hi"}], + "contextId": "ctx-1", + } + }, + litellm_params={ + "custom_llm_provider": "langflow", + "model": "langflow/flow-1", + "api_base": "http://localhost:7860", + A2A_USER_API_KEY_HASH_PARAM: "hashed-key-1", + }, + api_base="http://localhost:7860", + ) + + forwarded = mock_acompletion.call_args.kwargs + assert forwarded.get("session_id") != "ctx-1" + assert forwarded.get("session_id").endswith("-ctx-1") + assert ( + A2A_USER_API_KEY_HASH_PARAM not in forwarded + ), "internal principal param must not leak to the LLM call" + + +@pytest.mark.asyncio +async def test_langflow_a2a_config_requires_litellm_params_non_streaming(): + from litellm.a2a_protocol.providers.langflow.config import LangFlowA2AConfig + + with pytest.raises(ValueError, match="litellm_params is required"): + await LangFlowA2AConfig().handle_non_streaming( + request_id="req-1", + params={"message": {"contextId": "shared-session-99"}}, + ) + + +@pytest.mark.asyncio +async def test_langflow_a2a_config_requires_litellm_params_streaming(): + from litellm.a2a_protocol.providers.langflow.config import LangFlowA2AConfig + + with pytest.raises(ValueError, match="litellm_params is required"): + async for _ in LangFlowA2AConfig().handle_streaming( + request_id="req-1", + params={"message": {"contextId": "shared-session-99"}}, + ): + pass diff --git a/tests/test_litellm/proxy/agent_endpoints/test_a2a_endpoints.py b/tests/test_litellm/proxy/agent_endpoints/test_a2a_endpoints.py index 268e6d2dc13..a32f2eadb99 100644 --- a/tests/test_litellm/proxy/agent_endpoints/test_a2a_endpoints.py +++ b/tests/test_litellm/proxy/agent_endpoints/test_a2a_endpoints.py @@ -246,3 +246,105 @@ async def test_invoke_agent_a2a_handles_none_agent_card_params(): assert body["jsonrpc"] == "2.0" assert body["error"]["code"] == -32000 assert "no URL configured" in body["error"]["message"] + + +@pytest.mark.asyncio +async def test_invoke_agent_a2a_injects_authenticated_key_hash_for_bridge(): + """Completion-bridge agents must receive the authenticated key hash in + litellm_params so provider configs (e.g. LangFlow) can scope provider-side + session memory per key. Regression for cross-key A2A session bleed.""" + from litellm.a2a_protocol.litellm_completion_bridge.handler import ( + A2A_USER_API_KEY_HASH_PARAM, + ) + from litellm.proxy._types import UserAPIKeyAuth + + captured = {} + + async def mock_add_litellm_data(data, **kwargs): + data["proxy_server_request"] = { + "url": "http://localhost:4000/a2a/lf-agent", + "method": "POST", + "headers": {}, + "body": {}, + } + data.setdefault("metadata", {}) + return data + + async def capture_asend_message(**kwargs): + captured.update(kwargs) + resp = MagicMock() + resp.model_dump.return_value = {"jsonrpc": "2.0", "id": "test-id", "result": {}} + return resp + + mock_agent = MagicMock() + mock_agent.agent_id = "lf-agent" + mock_agent.agent_name = "lf-agent" + # No URL: the bridge derives the endpoint from the LangFlow agent config. + mock_agent.agent_card_params = {"name": "LF Agent"} + mock_agent.litellm_params = { + "custom_llm_provider": "langflow", + "model": "langflow/flow-1", + } + mock_agent.static_headers = None + mock_agent.extra_headers = None + + mock_request = MagicMock() + mock_request.headers = {} + mock_request.json = AsyncMock( + return_value={ + "jsonrpc": "2.0", + "id": "test-id", + "method": "message/send", + "params": { + "message": { + "role": "user", + "parts": [{"kind": "text", "text": "hi"}], + "messageId": "msg-1", + "contextId": "ctx-1", + } + }, + } + ) + + mock_user_api_key_dict = UserAPIKeyAuth( + api_key="sk-hashed-123", + user_id="test-user", + team_id="test-team", + ) + + with ( + patch( + "litellm.proxy.agent_endpoints.a2a_endpoints._get_agent", + return_value=mock_agent, + ), + patch( + "litellm.proxy.common_request_processing.add_litellm_data_to_request", + side_effect=mock_add_litellm_data, + ), + patch( + "litellm.proxy.agent_endpoints.auth.agent_permission_handler.AgentRequestHandler.is_agent_allowed", + new=AsyncMock(return_value=True), + ), + patch( + "litellm.a2a_protocol.asend_message", + new=AsyncMock(side_effect=capture_asend_message), + ), + patch("litellm.proxy.proxy_server.general_settings", {}), + patch("litellm.proxy.proxy_server.proxy_config", MagicMock()), + patch("litellm.proxy.proxy_server.version", "1.0.0"), + patch("litellm.a2a_protocol.main.A2A_SDK_AVAILABLE", True), + patch.dict(sys.modules, {"a2a": MagicMock(), "a2a.types": MagicMock()}), + ): + from litellm.proxy.agent_endpoints.a2a_endpoints import invoke_agent_a2a + + await invoke_agent_a2a( + agent_id="lf-agent", + request=mock_request, + fastapi_response=MagicMock(), + user_api_key_dict=mock_user_api_key_dict, + ) + + assert ( + captured.get("litellm_params", {}).get(A2A_USER_API_KEY_HASH_PARAM) + == mock_user_api_key_dict.api_key + ), "authenticated key hash was not forwarded to the completion bridge" diff --git a/ui/litellm-dashboard/public/assets/logos/langflow.svg b/ui/litellm-dashboard/public/assets/logos/langflow.svg new file mode 100644 index 00000000000..1c7b36c4dd6 --- /dev/null +++ b/ui/litellm-dashboard/public/assets/logos/langflow.svg @@ -0,0 +1,5 @@ + + + + + diff --git a/ui/litellm-dashboard/src/components/agents/agent_type_utils.ts b/ui/litellm-dashboard/src/components/agents/agent_type_utils.ts index fd04aa4c26c..8a78a7fabb7 100644 --- a/ui/litellm-dashboard/src/components/agents/agent_type_utils.ts +++ b/ui/litellm-dashboard/src/components/agents/agent_type_utils.ts @@ -10,11 +10,13 @@ export const detectAgentType = (agent: Agent): string => { const customProvider = agent.litellm_params?.custom_llm_provider; // Check by custom_llm_provider first + if (customProvider === "langflow") return "langflow"; if (customProvider === "langgraph") return "langgraph"; if (customProvider === "azure_ai") return "azure_ai_foundry"; if (customProvider === "bedrock") return "bedrock_agentcore"; // Check by model prefix + if (model.startsWith("langflow/")) return "langflow"; if (model.startsWith("langgraph/")) return "langgraph"; if (model.startsWith("azure_ai/agents/")) return "azure_ai_foundry"; if (model.startsWith("bedrock/agentcore/")) return "bedrock_agentcore"; From d991c47018aa2ef7317d4ce654c585a1f00e83c2 Mon Sep 17 00:00:00 2001 From: ryan-crabbe-berri Date: Tue, 2 Jun 2026 14:57:30 -0700 Subject: [PATCH 09/27] fix(ui/agents): make A2A skill tags enterable and validated (#29512) * fix(ui/agents): make A2A skill tags enterable and validated Skill tags were marked required but rendered as a comma-split text input that couldn't surface validation and let empty values save. Switch tags and examples to Select tag inputs, drop the misleading "Required" skills label (the API allows zero skills), and validate the full configure step so an added skill must be complete before advancing. Resolves LIT-3153 * fix(ui/agents): allow Enter to create skill tags/examples Drop open={false} from the tags and examples Select inputs. With the dropdown forced closed, AntD suppresses the "create from input" option, so pressing Enter (as the placeholder instructs) did nothing. Matches the existing extra_headers Select. --- .../src/components/agents/add_agent_form.tsx | 2 +- .../src/components/agents/agent_config.ts | 8 ++++---- .../components/agents/agent_form_fields.tsx | 20 ++++++++++++------- 3 files changed, 18 insertions(+), 12 deletions(-) diff --git a/ui/litellm-dashboard/src/components/agents/add_agent_form.tsx b/ui/litellm-dashboard/src/components/agents/add_agent_form.tsx index 3929a9e1832..eee19171fb6 100644 --- a/ui/litellm-dashboard/src/components/agents/add_agent_form.tsx +++ b/ui/litellm-dashboard/src/components/agents/add_agent_form.tsx @@ -197,7 +197,7 @@ const AddAgentForm: React.FC = ({ const handleNext = async () => { try { if (currentStep === 0) { - await form.validateFields(["agent_name"]); + await form.validateFields(); const agentName = form.getFieldValue("agent_name"); if (agentName && !newKeyName) { setNewKeyName(`${agentName}-key`); diff --git a/ui/litellm-dashboard/src/components/agents/agent_config.ts b/ui/litellm-dashboard/src/components/agents/agent_config.ts index e87b191b19c..14b6729bf93 100644 --- a/ui/litellm-dashboard/src/components/agents/agent_config.ts +++ b/ui/litellm-dashboard/src/components/agents/agent_config.ts @@ -212,14 +212,14 @@ export const SKILL_FIELD_CONFIG = { }, tags: { name: "tags", - label: "Tags (comma-separated)", + label: "Tags", required: true, - placeholder: "e.g., hello world, greeting", + placeholder: "Type a tag and press Enter", }, examples: { name: "examples", - label: "Examples (comma-separated)", - placeholder: "e.g., hi, hello world", + label: "Examples", + placeholder: "Type an example and press Enter", }, }; diff --git a/ui/litellm-dashboard/src/components/agents/agent_form_fields.tsx b/ui/litellm-dashboard/src/components/agents/agent_form_fields.tsx index 42e55b8c56f..27b93838ce6 100644 --- a/ui/litellm-dashboard/src/components/agents/agent_form_fields.tsx +++ b/ui/litellm-dashboard/src/components/agents/agent_form_fields.tsx @@ -56,7 +56,7 @@ const AgentFormFields: React.FC = ({ showAgentName = true, {/* Skills */} {shouldShow(AGENT_FORM_CONFIG.skills.key) && ( - + {(fields, { add, remove }) => ( <> @@ -94,20 +94,26 @@ const AgentFormFields: React.FC = ({ showAgentName = true, label={SKILL_FIELD_CONFIG.tags.label} name={[field.name, 'tags']} rules={[{ required: SKILL_FIELD_CONFIG.tags.required, message: 'Required' }]} - getValueFromEvent={(e) => e.target.value.split(',').map((s: string) => s.trim())} - getValueProps={(value) => ({ value: Array.isArray(value) ? value.join(', ') : value })} > - + +