From 391b8a265b7c5285b2640361268ec67b72c143f3 Mon Sep 17 00:00:00 2001 From: Atharva-Kanherkar <142440039+Atharva-Kanherkar@users.noreply.github.com> Date: Mon, 7 Sep 2026 18:40:43 +0530 Subject: [PATCH 01/44] fix(responses): stream one lifecycle across MCP auto-execute rounds Each auto-executed MCP round is a distinct upstream response, but the client reads one stream. The follow-up round's response.created, response.in_progress, and the interim response.completed were forwarded as-is, so one SSE body carried two lifecycles and output_index restarted at zero, which aborts accumulating clients such as the OpenAI SDK's responses.stream() before the final answer arrives. Fold the rounds into one public lifecycle: drop the openers of follow-up rounds, hold back the completed event of a round whose tool calls the gateway executes, shift later output indexes past the items already emitted, and list every round's items on the single final response.completed. Each mcp_call item is announced with output_item.added, keeps one item id across its events, and owns its own output_index. Sequence numbers stay strictly increasing when a round or a gateway event restarts numbering. The final round's response id is kept on response.completed so previous_response_id continuation still works. --- .../responses/mcp/mcp_streaming_iterator.py | 212 ++++++++++++++---- 1 file changed, 166 insertions(+), 46 deletions(-) diff --git a/litellm/responses/mcp/mcp_streaming_iterator.py b/litellm/responses/mcp/mcp_streaming_iterator.py index ca12b3e7cc3..ac158a32371 100644 --- a/litellm/responses/mcp/mcp_streaming_iterator.py +++ b/litellm/responses/mcp/mcp_streaming_iterator.py @@ -34,6 +34,16 @@ else: MAX_MCP_TOOL_CALL_ROUNDS: Final = 5 +def _output_items(response: ResponsesAPIResponse) -> Sequence[object]: + """Read a response's output items as plain objects; the field is a wide union of item models.""" + return tuple(cast("Sequence[object]", response.output)) # cast-ok: items are only carried, never inspected + + +def _set_event_field(event: ResponsesAPIStreamingResponse, name: str, value: object) -> None: + """Events are pydantic models with extra fields allowed, so any event type can carry the field.""" + setattr(event, name, value) + + async def create_mcp_list_tools_events( mcp_tools_with_litellm_proxy: Sequence[Mapping[str, object]], user_api_key_auth: "UserAPIKeyAuth | None", @@ -170,6 +180,7 @@ def create_mcp_call_events( result: str | None = None, base_item_id: str | None = None, sequence_start: int = 1, + output_index: int = 0, ) -> list[ResponsesAPIStreamingResponse]: """Create MCP call events following OpenAI's specification""" events: Final[list[ResponsesAPIStreamingResponse]] = [] @@ -179,7 +190,7 @@ def create_mcp_call_events( in_progress_event: Final = MCPCallInProgressEvent( type=ResponsesAPIStreamEvents.MCP_CALL_IN_PROGRESS, sequence_number=sequence_start, - output_index=0, + output_index=output_index, item_id=item_id, ) events.append(in_progress_event) @@ -187,7 +198,7 @@ def create_mcp_call_events( # MCP call arguments delta event (streaming the arguments) arguments_delta_event: Final = MCPCallArgumentsDeltaEvent( type=ResponsesAPIStreamEvents.MCP_CALL_ARGUMENTS_DELTA, - output_index=0, + output_index=output_index, item_id=item_id, delta=arguments, # JSON string with arguments sequence_number=sequence_start + 1, @@ -197,7 +208,7 @@ def create_mcp_call_events( # MCP call arguments done event arguments_done_event: Final = MCPCallArgumentsDoneEvent( type=ResponsesAPIStreamEvents.MCP_CALL_ARGUMENTS_DONE, - output_index=0, + output_index=output_index, item_id=item_id, arguments=arguments, # Complete JSON string with finalized arguments sequence_number=sequence_start + 2, @@ -210,7 +221,7 @@ def create_mcp_call_events( type=ResponsesAPIStreamEvents.MCP_CALL_COMPLETED, sequence_number=sequence_start + 3, item_id=item_id, - output_index=0, + output_index=output_index, ) events.append(completed_event) @@ -219,7 +230,7 @@ def create_mcp_call_events( output_item_done_event: Final = OutputItemDoneEvent( type=ResponsesAPIStreamEvents.OUTPUT_ITEM_DONE, - output_index=0, + output_index=output_index, item=BaseLiteLLMOpenAIResponseObject( **{ "id": item_id, @@ -239,7 +250,7 @@ def create_mcp_call_events( type=ResponsesAPIStreamEvents.MCP_CALL_FAILED, sequence_number=sequence_start + 3, item_id=item_id, - output_index=0, + output_index=output_index, ) events.append(failed_event) @@ -330,6 +341,17 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator): self._error_event_emitted = False self._last_sequence_number = 0 + # Every auto-execute round is a distinct upstream response, but the + # client is reading one stream. Fold the rounds into one public + # lifecycle: one response.created, one response.completed whose + # output holds every round's items, and output indexes that are + # never reused for a different item. + self._round_index = 0 + self._output_index_offset = 0 + self._round_max_output_index = -1 + self._composed_output: list[object] = [] # mutable-ok: grows as each round finishes + self._pending_mcp_call_items: list[dict[str, object]] = [] # mutable-ok: grows per executed tool + def _extract_mcp_headers_from_params(self) -> None: """Extract MCP headers from original request params to pass to tool calls""" @@ -415,8 +437,14 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator): async def __anext__(self) -> ResponsesAPIStreamingResponse: chunk: Final = await self._anext_impl() sequence_number: Final = getattr(chunk, "sequence_number", None) - if isinstance(sequence_number, int) and sequence_number > self._last_sequence_number: - self._last_sequence_number = sequence_number + if isinstance(sequence_number, int): + # Follow-up rounds and gateway events restart their numbering. + # Keep the public stream strictly increasing. + if sequence_number <= self._last_sequence_number and self._last_sequence_number > 0: + self._last_sequence_number += 1 + _set_event_field(chunk, "sequence_number", self._last_sequence_number) + else: + self._last_sequence_number = max(self._last_sequence_number, sequence_number) return chunk async def _anext_impl(self) -> ResponsesAPIStreamingResponse: @@ -472,7 +500,7 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator): await self._create_follow_up_iterator() if self.base_iterator is not None: self.phase = "continue_initial_response" - return await self.__anext__() + return await self._anext_impl() self.phase = "finished" if self._stream_error is not None and not self._error_event_emitted: self._error_event_emitted = True @@ -530,17 +558,11 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator): if chunk_type == ResponsesAPIStreamEvents.OUTPUT_ITEM_ADDED: self.initial_events_emitted = True self.phase = "mcp_discovery" - return chunk + return await self._compose_round_chunk(chunk) - # If auto-execution is enabled, check for completed responses - if self.should_auto_execute and self._is_response_completed(chunk): - response_obj = getattr(chunk, "response", None) - if isinstance(response_obj, ResponsesAPIResponse): - self.collected_response = response_obj - self.phase = "tool_execution" - await self._generate_tool_execution_events() - - return chunk + # None means the chunk was folded into the single public + # lifecycle; fall through so phase 4 runs the follow-up. + return await self._compose_round_chunk(chunk) except StopAsyncIteration: if self.should_auto_execute and self.collected_response: self.phase = "tool_execution" @@ -566,6 +588,69 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator): chunk_type: Final[object] = getattr(chunk, "type", None) return chunk_type == ResponsesAPIStreamEvents.RESPONSE_COMPLETED + def _follow_up_pending(self) -> bool: + """True when the current round's tool calls were executed and a follow-up round will run.""" + return self.collected_response is not None and self.collected_response is self._tool_results_for_response + + def _round_output_width(self, response: ResponsesAPIResponse) -> int: + """How many output indexes this round used, counting items it streamed but never listed.""" + return max(len(_output_items(response)), self._round_max_output_index + 1) + + def _absorb_round(self, response: ResponsesAPIResponse) -> None: + """Bank a finished round's items so the final response.completed can list them.""" + width: Final = self._round_output_width(response) + self._composed_output.extend(_output_items(response)) + self._composed_output.extend(self._pending_mcp_call_items) + self._output_index_offset += width + len(self._pending_mcp_call_items) + self._pending_mcp_call_items = [] + self._round_max_output_index = -1 + + async def _compose_round_chunk(self, chunk: ResponsesAPIStreamingResponse) -> ResponsesAPIStreamingResponse | None: + """ + Fold one round's event into the single public lifecycle. + + Returns None when the event must not reach the client: the lifecycle + openers of a follow-up round, and the response.completed of a round + whose tool calls the gateway executes itself. Shifts output_index on + follow-up rounds past the items already emitted, and lists every + round's items on the final response.completed. + """ + chunk_type: Final[object] = getattr(chunk, "type", None) + if self._round_index > 0 and chunk_type in ( + ResponsesAPIStreamEvents.RESPONSE_CREATED, + ResponsesAPIStreamEvents.RESPONSE_IN_PROGRESS, + ): + return None + + output_index: Final[object] = getattr(chunk, "output_index", None) + if isinstance(output_index, int): + self._round_max_output_index = max(self._round_max_output_index, output_index) + if self._output_index_offset: + _set_event_field(chunk, "output_index", output_index + self._output_index_offset) + + if not (self.should_auto_execute and self._is_response_completed(chunk)): + return chunk + + response_obj: Final[object] = getattr(chunk, "response", None) + if isinstance(response_obj, ResponsesAPIResponse): + self.collected_response = response_obj + # Move to tool execution phase after this chunk + self.phase = "tool_execution" + await self._generate_tool_execution_events() + + if not isinstance(response_obj, ResponsesAPIResponse): + return chunk + if self._follow_up_pending(): + self._absorb_round(response_obj) + return None + if self._composed_output: + merged_output: Final[list[object]] = [ # mutable-ok: the response model declares output as a list + *self._composed_output, + *_output_items(response_obj), + ] + _set_event_field(chunk, "response", response_obj.model_copy(update={"output": merged_output})) + return chunk + async def _process_base_iterator_chunk(self) -> ResponsesAPIStreamingResponse: """ Process a chunk from the base iterator with response ID consistency enforcement. @@ -593,17 +678,11 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator): ) response_obj.id = self._cached_response_id - # If auto-execution is enabled, check for completed responses - if self.should_auto_execute and self._is_response_completed(chunk): - # Collect the response for tool execution - response_obj = getattr(chunk, "response", None) - if isinstance(response_obj, ResponsesAPIResponse): - self.collected_response = response_obj - # Move to tool execution phase after emitting this chunk - self.phase = "tool_execution" - await self._generate_tool_execution_events() - - return chunk + composed: Final = await self._compose_round_chunk(chunk) + if composed is None: + # The chunk stays internal; hand the next public event back instead. + return await self._anext_impl() + return composed async def _create_initial_response_iterator(self) -> None: """Create the initial response iterator by making the first LLM call""" @@ -667,6 +746,16 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator): return self.tool_call_round += 1 + # Each executed tool is one mcp_call output item of the single + # public response. Announce it at an output_index past the items + # this round already streamed, and keep that item id for the + # completion events below. + from litellm.types.llms.openai import OutputItemAddedEvent + + next_output_index = self._output_index_offset + self._round_output_width( # rebind-ok: advances per item + self.collected_response + ) + call_items: Final[dict[str, tuple[str, int]]] = {} # mutable-ok: filled per tool call as events queue for tool_call in tool_calls: ( tool_name, @@ -674,14 +763,36 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator): tool_call_id, ) = LiteLLM_Proxy_MCP_Handler._extract_tool_call_details(tool_call) if tool_name and tool_call_id: + item_id = f"mcp_{uuid.uuid4().hex[:8]}" + output_index = next_output_index + next_output_index += 1 + call_items[tool_call_id] = (item_id, output_index) + self.tool_execution_events.append( + OutputItemAddedEvent.model_validate( + { + "type": ResponsesAPIStreamEvents.OUTPUT_ITEM_ADDED, + "sequence_number": len(self.tool_execution_events) + 1, + "output_index": output_index, + "item": { + "id": item_id, + "type": "mcp_call", + "status": "in_progress", + "arguments": tool_arguments or "{}", + "name": tool_name, + "server_label": "litellm", + }, + } + ) + ) # Create MCP call events for this tool execution call_events = create_mcp_call_events( tool_name=tool_name, tool_call_id=tool_call_id, arguments=tool_arguments or "{}", # JSON string with arguments result=None, # Will be set after execution - base_item_id=f"mcp_{uuid.uuid4().hex[:8]}", + base_item_id=item_id, sequence_start=len(self.tool_execution_events) + 1, + output_index=output_index, ) # Add the in_progress and arguments events (not the completed event yet) self.tool_execution_events.extend(call_events[:-1]) @@ -719,37 +830,45 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator): tool_arguments = args or "{}" break - item_id = f"mcp_{uuid.uuid4().hex[:8]}" + if tool_call_id in call_items: + item_id, output_index = call_items[tool_call_id] + else: + item_id = f"mcp_{uuid.uuid4().hex[:8]}" + output_index = next_output_index + next_output_index += 1 # Create the completion event completed_event = MCPCallCompletedEvent( type=ResponsesAPIStreamEvents.MCP_CALL_COMPLETED, sequence_number=len(self.tool_execution_events) + 1, item_id=item_id, - output_index=0, + output_index=output_index, ) self.tool_execution_events.append(completed_event) # Create output_item.done event with the tool call result from litellm.types.llms.openai import OutputItemDoneEvent + mcp_call_item = BaseLiteLLMOpenAIResponseObject( + **{ + "id": item_id, + "type": "mcp_call", + "approval_request_id": f"mcpr_{uuid.uuid4().hex[:8]}", + "arguments": tool_arguments, + "error": None, + "name": tool_name, + "output": result_text, + "server_label": "litellm", # or extract from tool config + } + ) output_item_done_event = OutputItemDoneEvent( type=ResponsesAPIStreamEvents.OUTPUT_ITEM_DONE, - output_index=0, - item=BaseLiteLLMOpenAIResponseObject( - **{ - "id": item_id, - "type": "mcp_call", - "approval_request_id": f"mcpr_{uuid.uuid4().hex[:8]}", - "arguments": tool_arguments, - "error": None, - "name": tool_name, - "output": result_text, - "server_label": "litellm", # or extract from tool config - } - ), + output_index=output_index, + item=mcp_call_item, ) self.tool_execution_events.append(output_item_done_event) + # The response model accepts output items as dicts, not as the generic event object. + self._pending_mcp_call_items.append(mcp_call_item.model_dump()) # Store tool results for follow-up call self.tool_results = tool_results @@ -824,6 +943,7 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator): self.base_iterator = follow_up_response self.collected_response = None self._cached_response_id = None + self._round_index += 1 except Exception as e: verbose_logger.error("Error creating follow-up iterator: %s", e) From 13ce30488ff6eed68c213a80a41a38979a279e59 Mon Sep 17 00:00:00 2001 From: Atharva-Kanherkar <142440039+Atharva-Kanherkar@users.noreply.github.com> Date: Mon, 7 Sep 2026 18:40:43 +0530 Subject: [PATCH 02/44] test(responses): cover the single public lifecycle for MCP rounds Assert one response.created and one response.completed across rounds, dense output indexes for the function call, the mcp_call item, and the final message, stable mcp_call item ids, strictly increasing sequence numbers, a serializable merged output, and that a stream without auto-execution is forwarded unchanged. Update the two existing tests that counted one completed event per internal round. --- .../mcp/test_mcp_streaming_iterator.py | 128 +++++++++++++++++- 1 file changed, 122 insertions(+), 6 deletions(-) diff --git a/tests/test_litellm/responses/mcp/test_mcp_streaming_iterator.py b/tests/test_litellm/responses/mcp/test_mcp_streaming_iterator.py index 5001589ce54..57ccbbfa493 100644 --- a/tests/test_litellm/responses/mcp/test_mcp_streaming_iterator.py +++ b/tests/test_litellm/responses/mcp/test_mcp_streaming_iterator.py @@ -57,6 +57,10 @@ def _text_message(text: str): return {"type": "message", "role": "assistant", "content": [{"type": "output_text", "text": text}]} +def _item_type(item) -> str: + return item["type"] if isinstance(item, dict) else item.type + + def _tool_call_stream(call_id: str, tool_name: str, response_id: str = "resp-1") -> _FakeAsyncStream: return _FakeAsyncStream([_completed_chunk([_function_call(call_id, tool_name)], response_id=response_id)]) @@ -134,11 +138,19 @@ async def test_second_round_tool_call_is_executed_and_reaches_final_text(monkeyp assert iterator.tool_call_round == 2 # The stream reached round 3 and produced the final text response instead - # of stopping after round 1 or round 2. + # of stopping after round 1 or round 2. The client sees one lifecycle whose + # final output lists every round's items in order. completed_chunks = [c for c in chunks if getattr(c, "type", None) == ResponsesAPIStreamEvents.RESPONSE_COMPLETED] - assert len(completed_chunks) == 3 + assert len(completed_chunks) == 1 final_output = completed_chunks[-1].response.output - assert final_output[0]["content"][0]["text"] == "Here's what I found after retrying." + assert [_item_type(item) for item in final_output] == [ + "function_call", + "mcp_call", + "function_call", + "mcp_call", + "message", + ] + assert final_output[-1]["content"][0]["text"] == "Here's what I found after retrying." @pytest.mark.asyncio @@ -207,7 +219,7 @@ async def test_continuation_id_is_final_round_not_interim_tool_call(monkeypatch) chunks = [chunk async for chunk in iterator] completed = [c for c in chunks if getattr(c, "type", None) == ResponsesAPIStreamEvents.RESPONSE_COMPLETED] - assert completed[-1].response.output[0]["content"][0]["text"] == "The first item is Alpha." + assert completed[-1].response.output[-1]["content"][0]["text"] == "The first item is Alpha." assert completed[-1].response.id == "resp-final" assert completed[-1].response.id != "resp-interim" @@ -280,7 +292,9 @@ async def test_streaming_follow_up_replays_reasoning_when_store_is_false(monkeyp base_iterator=_FakeAsyncStream( [ _output_item_added_chunk(), - _completed_chunk([_reasoning_item("gAAAAA-opaque-blob"), _function_call("call_1", "read_wiki_contents")]), + _completed_chunk( + [_reasoning_item("gAAAAA-opaque-blob"), _function_call("call_1", "read_wiki_contents")] + ), ] ), mcp_events=[], @@ -316,7 +330,9 @@ async def test_streaming_follow_up_keeps_previous_response_id_when_stored(monkey base_iterator=_FakeAsyncStream( [ _output_item_added_chunk(), - _completed_chunk([_reasoning_item("gAAAAA-opaque-blob"), _function_call("call_1", "read_wiki_contents")]), + _completed_chunk( + [_reasoning_item("gAAAAA-opaque-blob"), _function_call("call_1", "read_wiki_contents")] + ), ] ), mcp_events=[], @@ -336,3 +352,103 @@ async def test_streaming_follow_up_keeps_previous_response_id_when_stored(monkey follow_up_kwargs = aresponses_mock.call_args_list[0].kwargs assert follow_up_kwargs["previous_response_id"] == "resp_prev" assert not [item for item in follow_up_kwargs["input"] if item.get("type") == "reasoning"] + + +def _event(event_type, **fields): + return SimpleNamespace(type=event_type, **fields) + + +def _lifecycle_round(response_id: str, item: dict, sequence_start: int = 0): + """One upstream Responses round as a provider streams it: its own id, indexes from 0, numbering from 0.""" + return [ + _event( + ResponsesAPIStreamEvents.RESPONSE_CREATED, + response=ResponsesAPIResponse(id=response_id, created_at=0, output=[]), + sequence_number=sequence_start, + ), + _event( + ResponsesAPIStreamEvents.OUTPUT_ITEM_ADDED, output_index=0, item=item, sequence_number=sequence_start + 1 + ), + _event( + ResponsesAPIStreamEvents.OUTPUT_ITEM_DONE, output_index=0, item=item, sequence_number=sequence_start + 2 + ), + _completed_chunk([item], response_id=response_id), + ] + + +@pytest.mark.asyncio +async def test_auto_execute_rounds_share_one_public_lifecycle(monkeypatch): + """ + Every auto-execute round is a distinct upstream response, but the client + reads one stream. It must see one response.created, one response.completed, + and no output_index reused for a different item, otherwise accumulating + clients such as the OpenAI SDK's responses.stream() abort mid-stream. + """ + _mock_mcp_environment(monkeypatch) + + follow_up = _FakeAsyncStream(_lifecycle_round("resp-final", _text_message("Alpha."))) + monkeypatch.setattr(responses_main_module, "aresponses", AsyncMock(side_effect=[follow_up])) + + iterator = _make_iterator(_lifecycle_round("resp-interim", _function_call("call_1", "read_wiki_contents"))) + chunks = [chunk async for chunk in iterator] + types = [chunk.type for chunk in chunks] + + assert types.count(ResponsesAPIStreamEvents.RESPONSE_CREATED) == 1 + assert types.count(ResponsesAPIStreamEvents.RESPONSE_COMPLETED) == 1 + assert types[-1] == ResponsesAPIStreamEvents.RESPONSE_COMPLETED + + # The function call, the gateway's mcp_call, and the final message each own an index. + added = [c for c in chunks if c.type == ResponsesAPIStreamEvents.OUTPUT_ITEM_ADDED] + assert [(c.output_index, _item_type(c.item)) for c in added] == [ + (0, "function_call"), + (1, "mcp_call"), + (2, "message"), + ] + mcp_item_ids = {c.item_id for c in chunks if c.type == ResponsesAPIStreamEvents.MCP_CALL_IN_PROGRESS} + mcp_done = [ + c for c in chunks if c.type == ResponsesAPIStreamEvents.OUTPUT_ITEM_DONE and _item_type(c.item) == "mcp_call" + ] + assert [c.output_index for c in mcp_done] == [1] + assert {c.item.id for c in mcp_done} == mcp_item_ids + round_two = [ + c for c in chunks if getattr(c, "item_id", None) is None and getattr(c, "output_index", None) is not None + ] + assert max(c.output_index for c in round_two) == 2 + + # The single completed event lists every round's items and keeps the final round's id for continuation. + completed = chunks[-1] + assert completed.response.id == "resp-final" + assert [_item_type(item) for item in completed.response.output] == ["function_call", "mcp_call", "message"] + assert completed.response.output[-1]["content"][0]["text"] == "Alpha." + # The proxy serializes every chunk; the merged output must still be a valid response. + assert '"type":"mcp_call"' in completed.response.model_dump_json(exclude_none=True, exclude_unset=True) + + # Numbering stays strictly increasing across rounds and gateway events. + sequence_numbers = [c.sequence_number for c in chunks if getattr(c, "sequence_number", None) is not None] + assert sequence_numbers == sorted(sequence_numbers) + assert len(set(sequence_numbers)) == len(sequence_numbers) + + +@pytest.mark.asyncio +async def test_stream_without_auto_execute_is_forwarded_unchanged(monkeypatch): + """With approval required there is one round, and it passes through untouched.""" + _mock_mcp_environment(monkeypatch) + aresponses_mock = AsyncMock() + monkeypatch.setattr(responses_main_module, "aresponses", aresponses_mock) + + upstream = _lifecycle_round("resp-1", _function_call("call_1", "read_wiki_contents")) + iterator = MCPEnhancedStreamingIterator( + base_iterator=_FakeAsyncStream(list(upstream)), + mcp_events=[], + tool_server_map={"read_wiki_contents": "deepwiki"}, + mcp_tools_with_litellm_proxy=[{"require_approval": "always"}], + user_api_key_auth=None, + original_request_params={"model": "gpt-4", "input": "hi", "tools": [{"type": "mcp"}]}, + ) + + chunks = [chunk async for chunk in iterator] + + assert chunks == upstream + assert [c.output_index for c in chunks if hasattr(c, "output_index")] == [0, 0] + assert [c.sequence_number for c in chunks if hasattr(c, "sequence_number")] == [0, 1, 2] + aresponses_mock.assert_not_called() From 146e5c492b1bee31a57c77a8fafff74122a13df9 Mon Sep 17 00:00:00 2001 From: Atharva-Kanherkar <142440039+Atharva-Kanherkar@users.noreply.github.com> Date: Mon, 7 Sep 2026 18:49:04 +0530 Subject: [PATCH 03/44] fix(responses): keep the MCP lifecycle change within the type-discipline budget Mark the dict literals handed straight to model constructors, clear the pending mcp_call list in place, and bind the merged response before setting it on the event, so LIT002 stays at or below its base count. --- litellm/responses/mcp/mcp_streaming_iterator.py | 13 ++++++++----- 1 file changed, 8 insertions(+), 5 deletions(-) diff --git a/litellm/responses/mcp/mcp_streaming_iterator.py b/litellm/responses/mcp/mcp_streaming_iterator.py index ac158a32371..1c892eec803 100644 --- a/litellm/responses/mcp/mcp_streaming_iterator.py +++ b/litellm/responses/mcp/mcp_streaming_iterator.py @@ -602,7 +602,7 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator): self._composed_output.extend(_output_items(response)) self._composed_output.extend(self._pending_mcp_call_items) self._output_index_offset += width + len(self._pending_mcp_call_items) - self._pending_mcp_call_items = [] + self._pending_mcp_call_items.clear() self._round_max_output_index = -1 async def _compose_round_chunk(self, chunk: ResponsesAPIStreamingResponse) -> ResponsesAPIStreamingResponse | None: @@ -648,7 +648,10 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator): *self._composed_output, *_output_items(response_obj), ] - _set_event_field(chunk, "response", response_obj.model_copy(update={"output": merged_output})) + merged_response: Final = response_obj.model_copy( + update={"output": merged_output} # mutable-ok: pydantic's update argument must be a dict + ) + _set_event_field(chunk, "response", merged_response) return chunk async def _process_base_iterator_chunk(self) -> ResponsesAPIStreamingResponse: @@ -769,11 +772,11 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator): call_items[tool_call_id] = (item_id, output_index) self.tool_execution_events.append( OutputItemAddedEvent.model_validate( - { + { # mutable-ok: consumed once by model_validate "type": ResponsesAPIStreamEvents.OUTPUT_ITEM_ADDED, "sequence_number": len(self.tool_execution_events) + 1, "output_index": output_index, - "item": { + "item": { # mutable-ok: consumed once by model_validate "id": item_id, "type": "mcp_call", "status": "in_progress", @@ -850,7 +853,7 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator): from litellm.types.llms.openai import OutputItemDoneEvent mcp_call_item = BaseLiteLLMOpenAIResponseObject( - **{ + **{ # mutable-ok: consumed once by the model constructor "id": item_id, "type": "mcp_call", "approval_request_id": f"mcpr_{uuid.uuid4().hex[:8]}", From fa966ca2d0bc9bca963112c8befbda5997f306ef Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Tue, 8 Sep 2026 08:10:30 +0000 Subject: [PATCH 04/44] refactor(types): replace Any with real types across 20 backend files Fifth round of basedpyright Any reduction. Every change is typing-only and leaves runtime behavior identical. Prisma table access now goes through the PrismaTableRepository and TableActions protocols the repo already has, instead of reading untyped attributes off prisma_client.db. Payload and parameter annotations move from dict[str, Any] to dict[str, object] or Mapping[str, object]. Guardrail constructors that took **kwargs: Any now take Unpack of a PEP 728 TypedDict, the same _CustomGuardrailOptions shape three other guardrails already use. Calls into the OpenAI and Azure assistants SDKs pass explicit keywords rather than splatting a dict, so the arguments are checked against the real SDK signatures. --- litellm/llms/azure/assistants.py | 9 +-- litellm/llms/ollama/chat/transformation.py | 2 +- .../llms/ollama/completion/transformation.py | 2 +- litellm/llms/openai/openai.py | 61 ++++++++++++++----- litellm/llms/vertex_ai/fine_tuning/handler.py | 5 +- .../vertex_imagen_transformation.py | 25 ++++---- .../vertex_gemini_transformation.py | 6 +- litellm/proxy/_experimental/mcp_server/db.py | 8 ++- .../proxy/common_utils/reset_budget_job.py | 15 +++-- .../cisco_ai_defense/cisco_ai_defense.py | 7 ++- .../guardrail_hooks/deepkeep/deepkeep.py | 14 +++-- .../guardrail_hooks/headroom/headroom.py | 13 ++-- .../model_armor/model_armor.py | 6 +- .../auto_router_endpoints.py | 10 +-- .../spend_tracking/ptu_flat_cost_rollup.py | 30 +++++++-- litellm/proxy/utils.py | 11 ++-- litellm/realtime_api/main.py | 2 +- litellm/repositories/model_repository.py | 20 ++---- litellm/repositories/team_repository.py | 42 ++++++++++++- .../router_utils/fallback_event_handlers.py | 11 ++-- 20 files changed, 199 insertions(+), 100 deletions(-) diff --git a/litellm/llms/azure/assistants.py b/litellm/llms/azure/assistants.py index f7b419405ac..a4742a25a87 100644 --- a/litellm/llms/azure/assistants.py +++ b/litellm/llms/azure/assistants.py @@ -1,5 +1,5 @@ from collections.abc import Coroutine, Iterable -from typing import Any, Final, Literal, TypedDict +from typing import Final, Literal, TypedDict import httpx from openai import AsyncAzureOpenAI, AzureOpenAI @@ -715,7 +715,8 @@ class AzureAssistantsAPI(BaseAzureLLM): event_handler: AssistantEventHandler | None, litellm_params: dict | None = None, ) -> AsyncAssistantStreamManager[AsyncAssistantEventHandler]: - data: Final[dict[str, Any]] = { + stream_fn: Final = client.beta.threads.runs.stream + base_data: Final[_RunThreadStreamData] = { "thread_id": thread_id, "assistant_id": assistant_id, "additional_instructions": additional_instructions, @@ -725,8 +726,8 @@ class AzureAssistantsAPI(BaseAzureLLM): "tools": tools, } if event_handler is not None: - data["event_handler"] = event_handler - return client.beta.threads.runs.stream(**data) + return stream_fn(**base_data, event_handler=event_handler) + return stream_fn(**base_data) def run_thread_stream( self, diff --git a/litellm/llms/ollama/chat/transformation.py b/litellm/llms/ollama/chat/transformation.py index 181894646e3..257ebf921d9 100644 --- a/litellm/llms/ollama/chat/transformation.py +++ b/litellm/llms/ollama/chat/transformation.py @@ -124,7 +124,7 @@ class OllamaChatConfig(BaseConfig): setattr(self.__class__, key, value) @classmethod - def get_config(cls): + def get_config(cls) -> dict[str, object]: return super().get_config() def get_supported_openai_params(self, model: str): diff --git a/litellm/llms/ollama/completion/transformation.py b/litellm/llms/ollama/completion/transformation.py index dccc83efed4..a9bbfedafd6 100644 --- a/litellm/llms/ollama/completion/transformation.py +++ b/litellm/llms/ollama/completion/transformation.py @@ -227,7 +227,7 @@ class OllamaConfig(BaseConfig): model: str, api_base: str | None = None, api_key: str | None = None, - ) -> Any: + ) -> dict[str, object] | None: """ curl http://localhost:11434/api/show -d '{ "name": "mistral" diff --git a/litellm/llms/openai/openai.py b/litellm/llms/openai/openai.py index edc8d64d9c2..15c543b94f0 100644 --- a/litellm/llms/openai/openai.py +++ b/litellm/llms/openai/openai.py @@ -1,7 +1,7 @@ import time import types from collections.abc import AsyncIterator, Callable, Coroutine, Iterable, Iterator, Mapping -from typing import TYPE_CHECKING, Any, Final, Literal, Optional, cast +from typing import TYPE_CHECKING, Final, Literal, Optional, cast import httpx @@ -2754,7 +2754,12 @@ class OpenAIAssistantsAPI(BaseLLM): message_thread: Final = await openai_client.beta.threads.create(**data) - return Thread(**message_thread.dict()) + return Thread( + id=message_thread.id, + created_at=message_thread.created_at, + metadata=message_thread.metadata, + object=message_thread.object, + ) # fmt: off @@ -2840,7 +2845,12 @@ class OpenAIAssistantsAPI(BaseLLM): message_thread: Final = openai_client.beta.threads.create(**data) - return Thread(**message_thread.dict()) + return Thread( + id=message_thread.id, + created_at=message_thread.created_at, + metadata=message_thread.metadata, + object=message_thread.object, + ) async def async_get_thread( self, @@ -2863,7 +2873,12 @@ class OpenAIAssistantsAPI(BaseLLM): response: Final = await openai_client.beta.threads.retrieve(thread_id=thread_id) - return Thread(**response.dict()) + return Thread( + id=response.id, + created_at=response.created_at, + metadata=response.metadata, + object=response.object, + ) # fmt: off @@ -2929,7 +2944,12 @@ class OpenAIAssistantsAPI(BaseLLM): response: Final = openai_client.beta.threads.retrieve(thread_id=thread_id) - return Thread(**response.dict()) + return Thread( + id=response.id, + created_at=response.created_at, + metadata=response.metadata, + object=response.object, + ) def delete_thread(self): pass @@ -2986,18 +3006,27 @@ class OpenAIAssistantsAPI(BaseLLM): tools: Iterable[AssistantToolParam] | None, event_handler: AssistantEventHandler | None, ) -> AsyncAssistantStreamManager[AsyncAssistantEventHandler]: - data: Final[dict[str, Any]] = { - "thread_id": thread_id, - "assistant_id": assistant_id, - "additional_instructions": additional_instructions, - "instructions": instructions, - "metadata": metadata, - "model": model, - "tools": tools, - } + runs_stream: Final = client.beta.threads.runs.stream if event_handler is not None: - data["event_handler"] = event_handler - return client.beta.threads.runs.stream(**data) + return runs_stream( + thread_id=thread_id, + assistant_id=assistant_id, + additional_instructions=additional_instructions, + instructions=instructions, + metadata=metadata, + model=model, + tools=tools, + event_handler=event_handler, + ) + return runs_stream( + thread_id=thread_id, + assistant_id=assistant_id, + additional_instructions=additional_instructions, + instructions=instructions, + metadata=metadata, + model=model, + tools=tools, + ) def run_thread_stream( self, diff --git a/litellm/llms/vertex_ai/fine_tuning/handler.py b/litellm/llms/vertex_ai/fine_tuning/handler.py index 7ecc5e8ff3d..c79b6ffce43 100644 --- a/litellm/llms/vertex_ai/fine_tuning/handler.py +++ b/litellm/llms/vertex_ai/fine_tuning/handler.py @@ -280,7 +280,7 @@ class VertexFineTuningAPI(VertexLLM): vertex_location: str, vertex_credentials: str, request_route: str, - ): + ) -> object: _auth_header, vertex_project = await self._ensure_access_token_async( credentials=vertex_credentials, project_id=vertex_project, @@ -341,5 +341,4 @@ class VertexFineTuningAPI(VertexLLM): f"Error creating fine tuning job. Status code: {response.status_code}. Response: {response.text}" ) - response_json: Final = response.json() - return response_json + return response.json() diff --git a/litellm/llms/vertex_ai/image_edit/vertex_imagen_transformation.py b/litellm/llms/vertex_ai/image_edit/vertex_imagen_transformation.py index c6ad5928b74..fddc075bfc6 100644 --- a/litellm/llms/vertex_ai/image_edit/vertex_imagen_transformation.py +++ b/litellm/llms/vertex_ai/image_edit/vertex_imagen_transformation.py @@ -1,6 +1,7 @@ import base64 import json import os +from collections.abc import Mapping from io import BufferedRandom, BufferedReader, BytesIO from pathlib import Path from typing import TYPE_CHECKING, Any, Final, cast @@ -47,11 +48,11 @@ class VertexAIImagenImageEditConfig(BaseImageEditConfig, VertexLLM): image_edit_optional_params: ImageEditOptionalRequestParams, model: str, drop_params: bool, - ) -> dict[str, Any]: + ) -> dict[str, object]: supported_params: Final = self.get_supported_openai_params(model) filtered_params = {key: value for key, value in image_edit_optional_params.items() if key in supported_params} - mapped_params: Final[dict[str, Any]] = {} + mapped_params: Final[dict[str, object]] = {} # Map OpenAI parameters to Imagen format if "n" in filtered_params: @@ -148,10 +149,10 @@ class VertexAIImagenImageEditConfig(BaseImageEditConfig, VertexLLM): model: str, prompt: str | None, image: FileTypes | None, - image_edit_optional_request_params: dict[str, Any], + image_edit_optional_request_params: Mapping[str, object], litellm_params: GenericLiteLLMParams, headers: dict, - ) -> tuple[dict[str, Any], RequestFiles | None]: + ) -> tuple[dict[str, object], RequestFiles | None]: # Prepare reference images in the correct Imagen format if image is None: raise ValueError("Vertex AI Imagen image edit requires at least one reference image.") @@ -182,14 +183,14 @@ class VertexAIImagenImageEditConfig(BaseImageEditConfig, VertexLLM): parameters["guidanceScale"] = 7.5 # Default guidance scale parameters["seed"] = None # Let Vertex AI choose random seed - request_body: Final[dict[str, Any]] = { + request_body: Final[dict[str, object]] = { "instances": instances, "parameters": parameters, } - payload: Final[Any] = json.dumps(request_body) + payload: Final = json.dumps(request_body) empty_files: Final = cast(RequestFiles, []) - return cast(tuple[dict[str, Any], RequestFiles | None], (payload, empty_files)) + return cast(tuple[dict[str, object], RequestFiles | None], (payload, empty_files)) def transform_image_edit_response( self, @@ -237,8 +238,8 @@ class VertexAIImagenImageEditConfig(BaseImageEditConfig, VertexLLM): def _prepare_reference_images( self, image: FileTypes | list[FileTypes], - image_edit_optional_request_params: dict[str, Any], - ) -> list[dict[str, Any]]: + image_edit_optional_request_params: Mapping[str, object], + ) -> list[dict[str, object]]: """ Prepare reference images in the correct Imagen API format """ @@ -248,7 +249,7 @@ class VertexAIImagenImageEditConfig(BaseImageEditConfig, VertexLLM): else: images = [image] - reference_images: Final[list[dict[str, Any]]] = [] + reference_images: Final[list[dict[str, object]]] = [] for idx, img in enumerate(images): if img is None: @@ -258,7 +259,7 @@ class VertexAIImagenImageEditConfig(BaseImageEditConfig, VertexLLM): base64_data = base64.b64encode(image_bytes).decode("utf-8") # Create reference image structure - reference_image = { + reference_image: dict[str, object] = { "referenceType": "REFERENCE_TYPE_RAW", "referenceId": idx + 1, "referenceImage": {"bytesBase64Encoded": base64_data}, @@ -272,7 +273,7 @@ class VertexAIImagenImageEditConfig(BaseImageEditConfig, VertexLLM): mask_bytes: Final = self._read_all_bytes(mask_image) mask_base64: Final = base64.b64encode(mask_bytes).decode("utf-8") - mask_reference: Final = { + mask_reference: Final[dict[str, object]] = { "referenceType": "REFERENCE_TYPE_MASK", "referenceId": len(reference_images) + 1, "referenceImage": {"bytesBase64Encoded": mask_base64}, diff --git a/litellm/llms/vertex_ai/image_generation/vertex_gemini_transformation.py b/litellm/llms/vertex_ai/image_generation/vertex_gemini_transformation.py index d7a2491c04a..b2c52c53580 100644 --- a/litellm/llms/vertex_ai/image_generation/vertex_gemini_transformation.py +++ b/litellm/llms/vertex_ai/image_generation/vertex_gemini_transformation.py @@ -218,10 +218,10 @@ class VertexAIGeminiImageGenerationConfig(BaseImageGenerationConfig, VertexLLM): contents: Final = [{"role": "user", "parts": [{"text": prompt}]}] # Prepare generation config - generation_config: Final[dict[str, Any]] = {"responseModalities": ["IMAGE"]} + generation_config: Final[dict[str, object]] = {"responseModalities": ["IMAGE"]} # Seed from user-supplied imageConfig dict; flat params are overlaid for backward compat. - image_config: Final[dict[str, Any]] = dict(optional_params.get("imageConfig") or {}) + image_config: Final[dict[str, object]] = dict(optional_params.get("imageConfig") or {}) if "aspectRatio" in optional_params: image_config["aspectRatio"] = optional_params["aspectRatio"] @@ -242,7 +242,7 @@ class VertexAIGeminiImageGenerationConfig(BaseImageGenerationConfig, VertexLLM): elif "n" in optional_params: generation_config["candidateCount"] = optional_params["n"] - request_body: Final[dict[str, Any]] = { + request_body: Final[dict[str, object]] = { "contents": contents, "generationConfig": generation_config, } diff --git a/litellm/proxy/_experimental/mcp_server/db.py b/litellm/proxy/_experimental/mcp_server/db.py index 082a90fdcfb..980a67b30d6 100644 --- a/litellm/proxy/_experimental/mcp_server/db.py +++ b/litellm/proxy/_experimental/mcp_server/db.py @@ -37,6 +37,7 @@ from litellm.repositories.table_repositories import ( MCPServerOAuthClientRepository, MCPServerRepository, MCPUserCredentialsRepository, + PrismaTableRepository, ) from litellm.repositories.team_repository import TeamRepository from litellm.repositories.verification_token_repository import ( @@ -522,11 +523,14 @@ def _user_credential_actions( return table +class _MCPUserEnvVarsRepository(PrismaTableRepository["prisma_db_models.LiteLLM_MCPUserEnvVars"]): + table_name = "litellm_mcpuserenvvars" + + def _user_env_var_actions( prisma_client: PrismaClient, ) -> "TableActions[prisma_db_models.LiteLLM_MCPUserEnvVars]": - table: Final[TableActions[prisma_db_models.LiteLLM_MCPUserEnvVars]] = prisma_client.db.litellm_mcpuserenvvars - return table + return _MCPUserEnvVarsRepository(prisma_client).table async def _db_find_user_credential_row( diff --git a/litellm/proxy/common_utils/reset_budget_job.py b/litellm/proxy/common_utils/reset_budget_job.py index f2648c8466e..dc29a8377c9 100644 --- a/litellm/proxy/common_utils/reset_budget_job.py +++ b/litellm/proxy/common_utils/reset_budget_job.py @@ -48,7 +48,7 @@ from litellm.proxy.db.db_transaction_queue.pod_lock_manager import PodLockManage from litellm.proxy.db.exception_handler import call_with_db_reconnect_retry from litellm.proxy.utils import PrismaClient, ProxyLogging from litellm.repositories.organization_repository import OrganizationRepository -from litellm.repositories.prisma_protocols import SpendLinkedTable +from litellm.repositories.prisma_protocols import PrismaBatch, SpendLinkedTable from litellm.repositories.table_repositories import ( EndUserRepository, ModelAccessGroupBudgetRepository, @@ -435,6 +435,11 @@ class ResetBudgetJob: self.reset_settings: BudgetResetSettings = reset_settings or get_budget_reset_settings() self.pod_lock_manager: PodLockManager | None = pod_lock_manager + @property + def _new_batch(self) -> Callable[[], PrismaBatch]: + new_batch: Final[Callable[[], PrismaBatch]] = self.prisma_client.db.batch_ + return new_batch + async def _lease_is_held(self, lock_manager: PodLockManager) -> bool: """True only when the lease is readable and someone holds it. @@ -721,7 +726,7 @@ class ResetBudgetJob: ) async def _commit_budget_cascade_once(self, cascade: _BudgetCascade) -> None: - async with budget_cascade_unit_of_work(self.prisma_client.db.batch_) as uow: + async with budget_cascade_unit_of_work(self._new_batch) as uow: _queue_budget_linked_resets(uow.team_memberships, cascade) _queue_budget_linked_resets(uow.keys, cascade, extra=_LINKED_KEYS_WHERE) _queue_budget_linked_resets(uow.organizations, cascade, extra=_SPENT_ROWS_WHERE) @@ -861,7 +866,7 @@ class ResetBudgetJob: ) async def _write_key_reset_updates_once(self, updated_keys: list[LiteLLM_VerificationToken]) -> None: - async with spend_reset_unit_of_work(self.prisma_client.db.batch_) as uow: + async with spend_reset_unit_of_work(self._new_batch) as uow: for k in updated_keys: if k.token is None: continue @@ -885,7 +890,7 @@ class ResetBudgetJob: ) async def _write_user_reset_updates_once(self, updated_users: list[LiteLLM_UserTable]) -> None: - async with spend_reset_unit_of_work(self.prisma_client.db.batch_) as uow: + async with spend_reset_unit_of_work(self._new_batch) as uow: for u in updated_users: uow.users.queue_spend_reset( user_id=u.user_id, @@ -907,7 +912,7 @@ class ResetBudgetJob: ) async def _write_team_reset_updates_once(self, updated_teams: list[LiteLLM_TeamTable]) -> None: - async with spend_reset_unit_of_work(self.prisma_client.db.batch_) as uow: + async with spend_reset_unit_of_work(self._new_batch) as uow: for t in updated_teams: uow.teams.queue_spend_reset( team_id=t.team_id, diff --git a/litellm/proxy/guardrails/guardrail_hooks/cisco_ai_defense/cisco_ai_defense.py b/litellm/proxy/guardrails/guardrail_hooks/cisco_ai_defense/cisco_ai_defense.py index facb822d00d..017ef6e09f6 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/cisco_ai_defense/cisco_ai_defense.py +++ b/litellm/proxy/guardrails/guardrail_hooks/cisco_ai_defense/cisco_ai_defense.py @@ -26,6 +26,7 @@ from typing import TYPE_CHECKING, Any, Final, Literal import httpx from fastapi import HTTPException +from typing_extensions import TypedDict, Unpack from litellm import DualCache from litellm._logging import verbose_proxy_logger @@ -111,6 +112,10 @@ class CiscoAIDefenseGuardrailAPIError(Exception): """Raised when there is an error talking to the Cisco AI Defense API.""" +class _CustomGuardrailOptions(TypedDict, total=False, extra_items=object): + """Base-class constructor options this guardrail forwards untouched to CustomGuardrail.""" + + class CiscoAIDefenseGuardrail(_CiscoAIDefenseMcpMixin, CustomGuardrail): """ Cisco AI Defense guardrail integration. @@ -144,7 +149,7 @@ class CiscoAIDefenseGuardrail(_CiscoAIDefenseMcpMixin, CustomGuardrail): on_flagged_action: str | None = None, fallback_on_error: str | None = None, timeout: float | None = None, - **kwargs: Any, + **kwargs: Unpack[_CustomGuardrailOptions], ) -> None: resolved_api_key: Final = api_key or os.environ.get("CISCO_AI_DEFENSE_API_KEY") if not resolved_api_key: diff --git a/litellm/proxy/guardrails/guardrail_hooks/deepkeep/deepkeep.py b/litellm/proxy/guardrails/guardrail_hooks/deepkeep/deepkeep.py index 214d4b486d4..539dc1ea1e9 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/deepkeep/deepkeep.py +++ b/litellm/proxy/guardrails/guardrail_hooks/deepkeep/deepkeep.py @@ -7,10 +7,10 @@ import os from collections.abc import Mapping -from typing import TYPE_CHECKING, Any, Final, Literal, Optional, Protocol +from typing import TYPE_CHECKING, Final, Literal, Optional, Protocol import httpx -from typing_extensions import NotRequired, ReadOnly, TypedDict +from typing_extensions import NotRequired, ReadOnly, TypedDict, Unpack from litellm._logging import verbose_proxy_logger from litellm._version import version as litellm_version @@ -56,7 +56,13 @@ class DeepKeepFirewallResponse(TypedDict): class _DeepKeepInitKwargsView(TypedDict): """Typed read of the guardrail name carried in the untyped base-guardrail kwargs.""" - guardrail_name: ReadOnly[str] + guardrail_name: ReadOnly[str | None] + + +class _CustomGuardrailOptions(TypedDict, total=False, extra_items=object): + """Base-class constructor options this guardrail forwards untouched to CustomGuardrail.""" + + guardrail_name: ReadOnly[str | None] class _DeepKeepMetadataSource(TypedDict, total=False): @@ -110,7 +116,7 @@ class DeepKeepGuardrail(CustomGuardrail): firewall_id: str | None = None, unreachable_fallback: Literal["fail_closed", "fail_open"] = "fail_closed", extra_headers: Mapping[str, str] | list[str] | None = None, - **kwargs: Any, + **kwargs: Unpack[_CustomGuardrailOptions], ): self.async_handler = get_async_httpx_client(llm_provider=httpxSpecialProvider.GuardrailCallback) diff --git a/litellm/proxy/guardrails/guardrail_hooks/headroom/headroom.py b/litellm/proxy/guardrails/guardrail_hooks/headroom/headroom.py index fa113aa4d33..eb62b896784 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/headroom/headroom.py +++ b/litellm/proxy/guardrails/guardrail_hooks/headroom/headroom.py @@ -7,7 +7,7 @@ import time import uuid from collections.abc import Mapping, Sequence from dataclasses import dataclass -from typing import TYPE_CHECKING, Any, ClassVar, Final, Literal, TypeGuard +from typing import TYPE_CHECKING, ClassVar, Final, Literal, TypeGuard import httpx from fastapi import HTTPException @@ -50,6 +50,7 @@ from litellm.types.utils import CallTypes, GenericGuardrailAPIInputs if TYPE_CHECKING: from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj + from litellm.llms.base_llm.anthropic_messages.transformation import BaseAnthropicMessagesConfig from litellm.types.guardrails import LitellmParams from litellm.types.proxy.guardrails.guardrail_hooks.base import GuardrailConfigModel @@ -878,9 +879,9 @@ class HeadroomGuardrail(CustomGuardrail): async def async_pre_call_deployment_hook( self, - kwargs: dict[str, Any], + kwargs: dict[str, object], call_type: CallTypes | None, - ) -> dict[str, Any] | None: # mutable-ok: overrides CustomLogger hook whose contract is a plain dict + ) -> dict[str, object] | None: # mutable-ok: overrides CustomLogger hook whose contract is a plain dict base_result: Final = await super().async_pre_call_deployment_hook(kwargs, call_type) effective: Final = base_result if base_result is not None else kwargs if call_type not in _STREAM_CONVERTIBLE_CALL_TYPES: @@ -897,7 +898,7 @@ class HeadroomGuardrail(CustomGuardrail): async def async_should_run_agentic_loop( self, - response: Any, + response: object, model: str, messages: list[dict], tools: list[dict] | None, @@ -919,8 +920,8 @@ class HeadroomGuardrail(CustomGuardrail): tools: dict, model: str, messages: list[dict], - response: Any, - anthropic_messages_provider_config: Any, + response: object, + anthropic_messages_provider_config: BaseAnthropicMessagesConfig | None, anthropic_messages_optional_request_params: dict, logging_obj: LiteLLMLoggingObj | None, stream: bool, diff --git a/litellm/proxy/guardrails/guardrail_hooks/model_armor/model_armor.py b/litellm/proxy/guardrails/guardrail_hooks/model_armor/model_armor.py index fde40111d49..88c7f21ad72 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/model_armor/model_armor.py +++ b/litellm/proxy/guardrails/guardrail_hooks/model_armor/model_armor.py @@ -360,7 +360,7 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase): else: return {"modelResponseData": {"byteItem": {"byteDataType": file_type, "byteData": base64_data}}} - def _should_block_content(self, armor_response: Mapping[str, Any], allow_sanitization: bool = False) -> bool: + def _should_block_content(self, armor_response: Mapping[str, object], allow_sanitization: bool = False) -> bool: """Check if Model Armor response indicates content should be blocked, including both inspectResult and deidentifyResult.""" for filt in self._filter_result_items(armor_response): # Check RAI, PI/Jailbreak, Malicious URI, CSAM, Virus scan as before @@ -429,7 +429,7 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase): return filter_results return [] - def _has_deidentify_match(self, armor_response: Mapping[str, Any]) -> bool: + def _has_deidentify_match(self, armor_response: Mapping[str, object]) -> bool: """Whether an SDP de-identify filter matched, i.e. Model Armor owes this response a redaction.""" for filter_entry in self._filter_result_items(armor_response): sdp = filter_entry.get("sdpFilterResult") @@ -439,7 +439,7 @@ class ModelArmorGuardrail(CustomGuardrail, VertexBase): def _resolve_streaming_outcome( self, - armor_response: Mapping[str, Any], + armor_response: Mapping[str, object], assembled_response: object, content: str, ) -> tuple[bool, str | None]: diff --git a/litellm/proxy/management_endpoints/auto_router_endpoints.py b/litellm/proxy/management_endpoints/auto_router_endpoints.py index bbc914a772a..12d23550cb9 100644 --- a/litellm/proxy/management_endpoints/auto_router_endpoints.py +++ b/litellm/proxy/management_endpoints/auto_router_endpoints.py @@ -8,7 +8,6 @@ POST /auto_router/validate_complexity_router_config - Dry-run the complexity-rou from collections.abc import Mapping, Sequence from datetime import datetime, timedelta, timezone from itertools import chain, groupby -from operator import attrgetter from types import MappingProxyType from typing import TYPE_CHECKING, Annotated, Final, Protocol from uuid import uuid4 @@ -1094,6 +1093,10 @@ def _slices(rows: Sequence[_AttemptAggRow]) -> tuple[ShadowEvalSlice, ...]: ) +def _leg_group_id(leg: "_LegRow") -> str: + return leg.group_id + + class _LegRow(BaseModel): """One LiteLLM_ShadowEvalJob row, validated off the untyped prisma record. A row is one target's leg of a job; the legs of a job share group_id and identical config, @@ -1598,10 +1601,7 @@ async def list_shadow_eval_jobs( or () ) by_group: Final[Mapping[str, tuple[_LegRow, ...]]] = MappingProxyType( - { - group_id: tuple(group) - for group_id, group in groupby(sorted(legs, key=attrgetter("group_id")), key=attrgetter("group_id")) - } + {group_id: tuple(group) for group_id, group in groupby(sorted(legs, key=_leg_group_id), key=_leg_group_id)} ) newest_first: Final = sorted( by_group, key=lambda group_id: max(leg.created_at for leg in by_group[group_id]), reverse=True diff --git a/litellm/proxy/spend_tracking/ptu_flat_cost_rollup.py b/litellm/proxy/spend_tracking/ptu_flat_cost_rollup.py index 0e6412a2c64..b8af432029f 100644 --- a/litellm/proxy/spend_tracking/ptu_flat_cost_rollup.py +++ b/litellm/proxy/spend_tracking/ptu_flat_cost_rollup.py @@ -31,11 +31,31 @@ from litellm.constants import ( ) from litellm.litellm_core_utils.ptu_pricing import ptu_terms from litellm.proxy.spend_tracking.ptu_feature_flag import is_ptu_cost_attribution_enabled +from litellm.repositories.model_repository import ModelRepository +from litellm.repositories.prisma_protocols import TableActions +from litellm.repositories.table_repositories import PrismaTableRepository if TYPE_CHECKING: + from prisma import models as prisma_models + from litellm.proxy.db.db_transaction_queue.pod_lock_manager import PodLockManager from litellm.proxy.utils import PrismaClient + +class _DailyTeamSpendRepository(PrismaTableRepository["prisma_models.LiteLLM_DailyTeamSpend"]): + table_name = "litellm_dailyteamspend" + + +def _daily_team_spend_table(prisma_client: "PrismaClient") -> "TableActions[prisma_models.LiteLLM_DailyTeamSpend]": + """The sentinel rows this rollup writes, reads back and prunes.""" + return _DailyTeamSpendRepository(prisma_client).table + + +def _proxy_model_table(prisma_client: "PrismaClient") -> "TableActions[prisma_models.LiteLLM_ProxyModelTable]": + """The stored deployments the rollup scans for PTU config.""" + return ModelRepository(prisma_client).table + + _HOURS_PER_DAY: Final = 24 _PRUNE_ID_CHUNK_SIZE: Final = 5_000 _UPSERT_ATTEMPTS: Final = 3 @@ -97,7 +117,7 @@ def _decode_model_info(raw: object) -> "Mapping[str, object] | None": """ if isinstance(raw, str): try: - decoded: Final = json.loads(raw) + decoded: Final[object] = json.loads(raw) except (TypeError, ValueError): return None return decoded if isinstance(decoded, dict) else None @@ -240,7 +260,7 @@ async def _upsert_ptu_daily_row( } } now: Final = datetime.now(timezone.utc) - await prisma_client.db.litellm_dailyteamspend.upsert( + await _daily_team_spend_table(prisma_client).upsert( where=where, data={ # mutable-ok: prisma upsert data payload "create": { # mutable-ok: prisma create payload @@ -353,7 +373,7 @@ async def _load_ptu_models(prisma_client: "PrismaClient", *, router: object | No The router is handed in rather than read off the proxy module, so a run prices exactly the deployments its caller declares and nothing a co-resident process left behind. """ - rows: Final = await prisma_client.db.litellm_proxymodeltable.find_many() + rows: Final = await _proxy_model_table(prisma_client).find_many() db_ids: Final = frozenset(model_id for row in rows if (model_id := str(getattr(row, "model_id", "") or ""))) config_records: Final = _config_deployments(router, owned_by_db=db_ids) models: Final = tuple( @@ -503,7 +523,7 @@ async def _existing_sentinel_keys( survives a rename. Nothing here reads the display name. """ date_range: Final = {"gte": start.isoformat(), "lte": end.isoformat()} # mutable-ok: prisma range filter - rows: Final = await prisma_client.db.litellm_dailyteamspend.find_many( + rows: Final = await _daily_team_spend_table(prisma_client).find_many( where={"api_key": PTU_SENTINEL_API_KEY, "date": date_range} # mutable-ok: prisma find filter ) return frozenset( @@ -771,7 +791,7 @@ async def _prune_unrefreshed_sentinel_rows( ) filters: Final = tuple(_prune_filter(date_str=date_str, cutoff=cutoff, chunk=chunk) for chunk in chunks) deletions: Final = tuple( - [await prisma_client.db.litellm_dailyteamspend.delete_many(where=where) for where in filters] + [await _daily_team_spend_table(prisma_client).delete_many(where=where) for where in filters] ) deleted: Final = sum(deletions) if deleted: diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index ddf31cb1d8a..21328a4962a 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -3401,7 +3401,7 @@ class _ConfigRow: __slots__ = ("param_name", "param_value") - def __init__(self, param_name: str, param_value: Any) -> None: + def __init__(self, param_name: str, param_value: object) -> None: self.param_name = param_name self.param_value = param_value @@ -3414,7 +3414,7 @@ def _pack_config_row(row: Any) -> dict[str, object]: return {"param_name": row.param_name, "param_value": row.param_value} -def _unpack_config_row(cached: Any) -> _ConfigRow | None: +def _unpack_config_row(cached: object) -> _ConfigRow | None: if cached is None or cached == _CONFIG_CACHE_MISS: return None if isinstance(cached, dict): @@ -3557,6 +3557,7 @@ class PrismaClient: verbose_proxy_logger.debug("Creating Prisma Client..") try: from prisma import Prisma + from prisma.types import DatasourceOverride except Exception as e: verbose_proxy_logger.error("Failed to import Prisma client: %s", e) verbose_proxy_logger.error("This usually means 'prisma generate' hasn't been run yet.") @@ -3607,11 +3608,11 @@ class PrismaClient: reader_token: Final = mint_database_token(token_auth, reader_iam_endpoint) read_replica_url = reader_iam_endpoint.build_url(reader_token) os.environ["DATABASE_URL_READ_REPLICA"] = read_replica_url - reader_kwargs: Final[dict[str, Any]] = {"datasource": {"url": read_replica_url}} + reader_datasource: Final = DatasourceOverride(url=read_replica_url) if http_client is not None: - reader_prisma = Prisma(http=http_client, **reader_kwargs) + reader_prisma = Prisma(http=http_client, datasource=reader_datasource) else: - reader_prisma = Prisma(**reader_kwargs) + reader_prisma = Prisma(datasource=reader_datasource) reader_wrapper: Final = PrismaWrapper( original_prisma=reader_prisma, token_auth=token_auth, diff --git a/litellm/realtime_api/main.py b/litellm/realtime_api/main.py index b824a5928c6..9d81c725c70 100644 --- a/litellm/realtime_api/main.py +++ b/litellm/realtime_api/main.py @@ -53,7 +53,7 @@ bedrock_realtime: Final = BedrockRealtime() xai_realtime: Final = XAIRealtime() vertex_llm_base: Final = VertexBase() base_llm_http_handler = BaseLLMHTTPHandler() -_EMPTY_MODEL_PARAMS: Final[Mapping[str, Any]] = MappingProxyType({}) +_EMPTY_MODEL_PARAMS: Final[Mapping[str, object]] = MappingProxyType({}) def _with_resolved_session_model(session: dict[str, object], model_name: str) -> dict[str, object]: diff --git a/litellm/repositories/model_repository.py b/litellm/repositories/model_repository.py index d24eb8ffc62..8ee76b93923 100644 --- a/litellm/repositories/model_repository.py +++ b/litellm/repositories/model_repository.py @@ -4,29 +4,23 @@ Model repository for database operations on LiteLLM_ProxyModelTable. import json from collections.abc import Mapping, Sequence -from typing import TYPE_CHECKING, Any, Final, Protocol +from typing import TYPE_CHECKING, Any, Final from litellm.models.model import LiteLLM_ProxyModelTable -from litellm.proxy.common_utils.config_sync_pubsub import wrap_table_actions_for_config_sync from litellm.proxy.common_utils.encrypt_decrypt_utils import ( decrypt_value_helper, encrypt_value_helper, ) from litellm.repositories.base_repository import BaseRepository from litellm.repositories.prisma_protocols import TableActions +from litellm.repositories.table_repositories import PrismaTableRepository if TYPE_CHECKING: from prisma import models as prisma_models -class _PrismaModelDb(Protocol): - @property - def litellm_proxymodeltable(self) -> TableActions["prisma_models.LiteLLM_ProxyModelTable"]: ... - - -class _PrismaClientView(Protocol): - @property - def db(self) -> _PrismaModelDb: ... +class _ProxyModelTableRepository(PrismaTableRepository["prisma_models.LiteLLM_ProxyModelTable"]): + table_name = "litellm_proxymodeltable" class ModelRepository(BaseRepository[LiteLLM_ProxyModelTable]): @@ -38,11 +32,7 @@ class ModelRepository(BaseRepository[LiteLLM_ProxyModelTable]): @property def table(self) -> TableActions["prisma_models.LiteLLM_ProxyModelTable"]: - client: Final[_PrismaClientView] = self.prisma_client - return wrap_table_actions_for_config_sync( - actions=client.db.litellm_proxymodeltable, - table_name="litellm_proxymodeltable", - ) + return _ProxyModelTableRepository(self._prisma_client).table @property def model_class(self) -> type[LiteLLM_ProxyModelTable]: diff --git a/litellm/repositories/team_repository.py b/litellm/repositories/team_repository.py index 5ff07d76b5d..cbe263699c9 100644 --- a/litellm/repositories/team_repository.py +++ b/litellm/repositories/team_repository.py @@ -5,6 +5,7 @@ Team repository for database operations on LiteLLM_TeamTable. import json from collections.abc import Mapping, Sequence from datetime import datetime +from types import TracebackType from typing import TYPE_CHECKING, Final, Protocol from pydantic import TypeAdapter @@ -40,6 +41,36 @@ def _team_arrays(team: LiteLLM_TeamTable) -> _TeamArrays: return team +class _TeamTables(Protocol): + """The two team tables this repository reads and writes.""" + + @property + def litellm_teamtable(self) -> TableActions["prisma_models.LiteLLM_TeamTable"]: ... + + @property + def litellm_deletedteamtable(self) -> TableActions["prisma_models.LiteLLM_DeletedTeamTable"]: ... + + +class _TeamTransactionManager(Protocol): + async def __aenter__(self) -> _TeamTables: ... + + async def __aexit__( + self, + exc_type: type[BaseException] | None, + exc_value: BaseException | None, + traceback: TracebackType | None, + ) -> bool | None: ... + + +class _PrismaTeamDb(_TeamTables, Protocol): + def tx(self) -> _TeamTransactionManager: ... + + +class _PrismaClientView(Protocol): + @property + def db(self) -> _PrismaTeamDb: ... + + _MEMBERS_WITH_ROLES_ADAPTER: Final = TypeAdapter(list[Member]) _JSON_ENCODED_TEAM_FIELDS: Final = ( "metadata", @@ -54,13 +85,18 @@ _JSON_ENCODED_TEAM_FIELDS: Final = ( class TeamRepository(BaseRepository[LiteLLM_TeamTable]): """Repository for team database operations.""" + @property + def _db(self) -> _PrismaTeamDb: + client: Final[_PrismaClientView] = self.prisma_client + return client.db + @property def table(self) -> TableActions["prisma_models.LiteLLM_TeamTable"]: - return self.prisma_client.db.litellm_teamtable + return self._db.litellm_teamtable @property def deleted_table(self) -> TableActions["prisma_models.LiteLLM_DeletedTeamTable"]: - return self.prisma_client.db.litellm_deletedteamtable + return self._db.litellm_deletedteamtable @property def model_class(self) -> type[LiteLLM_TeamTable]: @@ -256,7 +292,7 @@ class TeamRepository(BaseRepository[LiteLLM_TeamTable]): archive_data["litellm_changed_by"] = litellm_changed_by archive_data["deleted_at"] = datetime.utcnow() - async with self.prisma_client.db.tx() as tx: + async with self._db.tx() as tx: await tx.litellm_deletedteamtable.create(data=archive_data) await tx.litellm_teamtable.delete(where={"team_id": team_id}) diff --git a/litellm/router_utils/fallback_event_handlers.py b/litellm/router_utils/fallback_event_handlers.py index 7fda5d96fb0..bf1838e50f7 100644 --- a/litellm/router_utils/fallback_event_handlers.py +++ b/litellm/router_utils/fallback_event_handlers.py @@ -266,7 +266,7 @@ def get_pre_routing_selection(kwargs: Mapping[str, object]) -> str | None: DISABLE_FALLBACKS_METADATA_KEY: Final = "_disable_fallbacks" -def record_disable_fallbacks(request_kwargs: Mapping[str, Any] | None, disabled: bool) -> None: +def record_disable_fallbacks(request_kwargs: Mapping[str, object] | None, disabled: bool) -> None: """ Write-or-clear the request's disable_fallbacks verdict into the router-internal metadata bucket. The wrapper pops the raw kwarg before any downstream frame runs, so the refusal @@ -286,7 +286,7 @@ def record_disable_fallbacks(request_kwargs: Mapping[str, Any] | None, disabled: bucket.pop(DISABLE_FALLBACKS_METADATA_KEY, None) -def fallbacks_disabled_for_request(kwargs: Mapping[str, Any]) -> bool: +def fallbacks_disabled_for_request(kwargs: Mapping[str, object]) -> bool: """True when this request opted out of fallbacks, read from the raw kwarg (pre-pop snapshots keep it) or the router-internal bucket the wrapper stamps after popping it.""" if kwargs.get("disable_fallbacks") is True: @@ -639,7 +639,7 @@ async def log_failure_fallback_event(original_model_group: str, kwargs: dict, or verbose_router_logger.error("Error in log_failure_fallback_event: %s", e) -def _check_non_standard_fallback_format(fallbacks: list[Any] | None) -> bool: +def _check_non_standard_fallback_format(fallbacks: Sequence[object] | None) -> bool: """ Checks if the fallbacks list is a list of strings or a list of dictionaries. @@ -653,8 +653,9 @@ def _check_non_standard_fallback_format(fallbacks: list[Any] | None) -> bool: return False if all(isinstance(item, str) for item in fallbacks): return True - elif all(isinstance(item, dict) for item in fallbacks): - for item in fallbacks: + dict_entries: Final = tuple(item for item in fallbacks if isinstance(item, dict)) + if len(dict_entries) == len(fallbacks): + for item in dict_entries: for key in LiteLLMParamsTypedDict.__annotations__: if key in item: # If the value is a list, it's likely a standard fallback model group mapping From f747bc67048f2b38e6e99bd57ca284fe6f77bdb5 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Tue, 8 Sep 2026 10:14:07 +0000 Subject: [PATCH 05/44] refactor(types): replace Any with real types across 29 more backend files Second batch of the fifth basedpyright Any reduction round. Every change is typing-only and leaves runtime behavior identical. Guardrail hooks and file, vector store and usage endpoints move their payload, header and response annotations from Any to object, Mapping[str, object] or the concrete response model the call site already produces. Two private aggregation helpers in the guardrail usage endpoints take a key accessor function instead of an attribute name string, so the key they read is checked against the row type. The verification token repository reaches its two tables through Protocols that name the handles it calls, rather than reading them off an untyped prisma client, and the Azure AD credential wrapper describes the azure-identity credential it wraps the same way. --- .../litellm_core_utils/llm_cost_calc/utils.py | 15 +++-- .../azure/text_to_speech/transformation.py | 9 +-- .../image_edit/stability_transformation.py | 4 +- .../responses/transformation.py | 14 ++--- .../llms/cohere/embed/v1_transformation.py | 15 +++-- litellm/llms/gdc/chat/transformation.py | 22 ++++++- .../llama3/transformation.py | 4 +- litellm/llms/voyage/rerank/transformation.py | 6 +- litellm/proxy/client/users.py | 3 +- .../proxy/common_utils/http_parsing_utils.py | 14 +++-- litellm/proxy/db/prisma_client.py | 9 ++- .../block_code_execution.py | 9 ++- .../cato_networks/cato_networks.py | 10 ++-- .../guardrail_hooks/compresr/compresr.py | 11 ++-- .../llm_as_a_judge/__init__.py | 23 +++++++- .../guardrails/guardrail_hooks/noma/noma.py | 15 +++-- .../panw_prisma_airs/panw_prisma_airs.py | 4 +- litellm/proxy/guardrails/usage_endpoints.py | 25 ++++---- .../proxy/hooks/proxy_track_cost_callback.py | 28 +++++++-- ...model_access_group_management_endpoints.py | 25 ++++---- .../usage_endpoints/ai_usage_chat.py | 11 ++-- .../openai_files_endpoints/files_endpoints.py | 4 +- .../storage_backend_service.py | 9 ++- .../proxy/vector_store_endpoints/endpoints.py | 2 +- .../vector_store_files_endpoints/endpoints.py | 10 ++-- litellm/proxy_auth/credentials.py | 58 ++++++++++++++----- litellm/rag/rag_query.py | 4 +- .../verification_token_repository.py | 46 ++++++++++++--- litellm/router_utils/search_api_router.py | 4 +- 29 files changed, 274 insertions(+), 139 deletions(-) diff --git a/litellm/litellm_core_utils/llm_cost_calc/utils.py b/litellm/litellm_core_utils/llm_cost_calc/utils.py index 68dc27ec25e..52dac92ee22 100644 --- a/litellm/litellm_core_utils/llm_cost_calc/utils.py +++ b/litellm/litellm_core_utils/llm_cost_calc/utils.py @@ -6,7 +6,7 @@ from collections.abc import Mapping, Sequence from dataclasses import dataclass from datetime import datetime, timezone, tzinfo from types import MappingProxyType -from typing import Any, Final, Literal, TypedDict, cast +from typing import Final, Literal, TypedDict, cast from zoneinfo import ZoneInfo, ZoneInfoNotFoundError import litellm @@ -89,7 +89,7 @@ def _requested_image_size(optional_params: Mapping[str, object] | None) -> str | return value if value is not None and _IMAGE_SIZE_PATTERN.fullmatch(value) else None -def get_web_search_requests(server_tool_use: Any) -> int | None: +def get_web_search_requests(server_tool_use: object) -> int | None: """ Tolerantly read ``web_search_requests`` from a ``server_tool_use`` value that may be ``None``, a ``dict``, a ``ServerToolUse`` pydantic instance, @@ -1494,7 +1494,7 @@ def calculate_image_response_cost_from_usage( if prompt_tokens == 0 and completion_tokens == 0 and total_tokens == 0: return None - input_tokens_details: Final = getattr(usage, "input_tokens_details", None) + input_tokens_details: Final[object] = getattr(usage, "input_tokens_details", None) prompt_tokens_details: PromptTokensDetailsWrapper | None = None if input_tokens_details is not None: # input_tokens_details may be a dict (e.g. OpenAI image edit responses) @@ -1507,9 +1507,12 @@ def calculate_image_response_cost_from_usage( cached_tokens=0, ) - output_tokens_details = getattr(usage, "completion_tokens_details", None) - if output_tokens_details is None: - output_tokens_details = getattr(usage, "output_tokens_details", None) + completion_tokens_details_attr: Final[object] = getattr(usage, "completion_tokens_details", None) + output_tokens_details: Final[object] = ( + getattr(usage, "output_tokens_details", None) + if completion_tokens_details_attr is None + else completion_tokens_details_attr + ) if output_tokens_details is None: completion_tokens_details = CompletionTokensDetailsWrapper( diff --git a/litellm/llms/azure/text_to_speech/transformation.py b/litellm/llms/azure/text_to_speech/transformation.py index d8ccf26ce60..eed7a3178ca 100644 --- a/litellm/llms/azure/text_to_speech/transformation.py +++ b/litellm/llms/azure/text_to_speech/transformation.py @@ -19,6 +19,7 @@ from litellm.secret_managers.main import get_secret_str if TYPE_CHECKING: from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj + from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler from litellm.types.llms.openai import HttpxBinaryResponseContent else: LiteLLMLoggingObj = Any @@ -67,15 +68,15 @@ class AzureAVATextToSpeechConfig(BaseTextToSpeechConfig): litellm_params_dict: dict, logging_obj: "LiteLLMLoggingObj", timeout: float | httpx.Timeout, - extra_headers: dict[str, Any] | None, - base_llm_http_handler: Any, + extra_headers: dict[str, object] | None, + base_llm_http_handler: "BaseLLMHTTPHandler", aspeech: bool, api_base: str | None, api_key: str | None, - **kwargs: Any, + **kwargs: object, ) -> Union[ "HttpxBinaryResponseContent", - Coroutine[Any, Any, "HttpxBinaryResponseContent"], + Coroutine[object, object, "HttpxBinaryResponseContent"], ]: """ Dispatch method to handle Azure AVA TTS requests diff --git a/litellm/llms/bedrock/image_edit/stability_transformation.py b/litellm/llms/bedrock/image_edit/stability_transformation.py index bc9a64f587a..01e25f4671e 100644 --- a/litellm/llms/bedrock/image_edit/stability_transformation.py +++ b/litellm/llms/bedrock/image_edit/stability_transformation.py @@ -125,7 +125,7 @@ class BedrockStabilityImageEditConfig(BaseImageEditConfig): } # Create a copy to not mutate original - convert TypedDict to regular dict - mapped_params: Final[dict[str, Any]] = dict(image_edit_optional_params) + mapped_params: Final[dict[str, object]] = dict(image_edit_optional_params) for k, v in image_edit_optional_params.items(): if k in param_mapping: @@ -172,7 +172,7 @@ class BedrockStabilityImageEditConfig(BaseImageEditConfig): Returns the request body dict that will be JSON-encoded by the handler. """ # Build Bedrock Stability request - data: Final[dict[str, Any]] = { + data: Final[dict[str, object]] = { "output_format": "png", # Default to PNG } diff --git a/litellm/llms/bedrock_mantle/responses/transformation.py b/litellm/llms/bedrock_mantle/responses/transformation.py index bbbda4d14b6..3d2eab8fcee 100644 --- a/litellm/llms/bedrock_mantle/responses/transformation.py +++ b/litellm/llms/bedrock_mantle/responses/transformation.py @@ -16,8 +16,8 @@ BaseAWSLLM._sign_request after the request body is finalized. """ import json -from collections.abc import Mapping -from typing import Any, Final +from collections.abc import Mapping, Sequence +from typing import Final import httpx from typing_extensions import ReadOnly, TypedDict @@ -142,9 +142,9 @@ class BedrockMantleResponsesAPIConfig(BedrockMantleAuthMixin, OpenAIResponsesAPI return False @staticmethod - def _filter_unsupported_tools(tools: list[Any]) -> list[Any]: + def _filter_unsupported_tools(tools: "Sequence[object]") -> "list[object]": """Keep only tool types Mantle's Responses API accepts.""" - kept: Final[list[Any]] = [] + kept: Final[list[object]] = [] dropped_types: Final[list[str]] = [] for tool in tools: if not isinstance(tool, dict): @@ -217,11 +217,11 @@ class BedrockMantleResponsesAPIConfig(BedrockMantleAuthMixin, OpenAIResponsesAPI ) @staticmethod - def _is_codex_additional_tools_item(item: Any) -> bool: + def _is_codex_additional_tools_item(item: object) -> bool: return isinstance(item, dict) and item.get("type") == _CODEX_ADDITIONAL_TOOLS_INPUT_ITEM_TYPE @staticmethod - def _tools_of_additional_tools_item(item: "dict[str, Any]") -> "list[Any]": + def _tools_of_additional_tools_item(item: "Mapping[str, object]") -> "list[object]": tools: Final = item.get("tools") return tools if isinstance(tools, list) else [] @@ -229,7 +229,7 @@ class BedrockMantleResponsesAPIConfig(BedrockMantleAuthMixin, OpenAIResponsesAPI def _hoist_codex_additional_tools( cls, input: "str | ResponseInputParam", - ) -> "tuple[str | ResponseInputParam, list[Any]]": + ) -> "tuple[str | ResponseInputParam, list[object]]": """Codex's "responses lite" wire mode ships tool definitions inside `input` as {"type": "additional_tools", "role": "developer", "tools": [...]} items. api.openai.com accepts that item type; Mantle diff --git a/litellm/llms/cohere/embed/v1_transformation.py b/litellm/llms/cohere/embed/v1_transformation.py index ee40464362d..b35fae5a1ac 100644 --- a/litellm/llms/cohere/embed/v1_transformation.py +++ b/litellm/llms/cohere/embed/v1_transformation.py @@ -2,7 +2,8 @@ Legacy /v1/embedding transformation logic for Bedrock Cohere. """ -from typing import Any, Final +from collections.abc import Sized +from typing import Final, Protocol import httpx @@ -16,6 +17,12 @@ from litellm.types.utils import EmbeddingResponse, PromptTokensDetailsWrapper, U from litellm.utils import is_base64_encoded +class _SupportsEncode(Protocol): + """Tokenizer handle: the embedding usage path only encodes text to measure its token length.""" + + def encode(self, text: str, /) -> Sized: ... + + class CohereEmbeddingConfig: """ Reference: https://docs.cohere.com/v2/reference/embed @@ -61,7 +68,7 @@ class CohereEmbeddingConfig: return transformed_request - def _calculate_usage(self, input: list[str], encoding: Any, meta: dict) -> Usage: + def _calculate_usage(self, input: list[str], encoding: _SupportsEncode, meta: dict) -> Usage: input_tokens = 0 text_tokens: Final[int | None] = meta.get("billed_units", {}).get("input_tokens") @@ -97,7 +104,7 @@ class CohereEmbeddingConfig: data: dict | CohereEmbeddingRequest, model_response: EmbeddingResponse, model: str, - encoding: Any, + encoding: _SupportsEncode, input: list, ) -> EmbeddingResponse: response_json: Final = response.json() @@ -121,7 +128,7 @@ class CohereEmbeddingConfig: response_json: dict, model_response: EmbeddingResponse, model: str, - encoding: Any, + encoding: _SupportsEncode, input: list, ) -> EmbeddingResponse: """ diff --git a/litellm/llms/gdc/chat/transformation.py b/litellm/llms/gdc/chat/transformation.py index 03037512551..6eac3ac79cd 100644 --- a/litellm/llms/gdc/chat/transformation.py +++ b/litellm/llms/gdc/chat/transformation.py @@ -7,14 +7,32 @@ import os import re import threading from collections.abc import Callable -from typing import Any, Final, Protocol +from typing import Final, Protocol from urllib.parse import urlsplit +from typing_extensions import ReadOnly, TypedDict, Unpack + import litellm from litellm.llms.openai_like.chat.transformation import OpenAILikeChatConfig from litellm.types.llms.openai import AllMessageValues +class _OpenAIGPTConfigOptions(TypedDict, total=False): + """The sampling defaults ``OpenAIGPTConfig.__init__`` accepts and stashes on the class.""" + + frequency_penalty: ReadOnly[int | None] + function_call: ReadOnly[str | dict[str, object] | None] + functions: ReadOnly[list[object] | None] + logit_bias: ReadOnly[dict[str, object] | None] + max_tokens: ReadOnly[int | None] + n: ReadOnly[int | None] + presence_penalty: ReadOnly[int | None] + stop: ReadOnly[str | list[object] | None] + temperature: ReadOnly[int | None] + top_p: ReadOnly[int | None] + response_format: ReadOnly[dict[str, object] | None] + + class _GDCHAudienceCredentials(Protocol): """A GDCH service account credential already bound to an audience, ready to mint a bearer token.""" @@ -32,7 +50,7 @@ class GDCGeminiConfig(OpenAILikeChatConfig): _GDCH_CREDENTIAL_TYPE: Final[str] = "gdch_service_account" _PATH_ID_PATTERN: Final[re.Pattern[str]] = re.compile(r"^[a-zA-Z0-9_-]+$") - def __init__(self, **kwargs: Any) -> None: + def __init__(self, **kwargs: Unpack[_OpenAIGPTConfigOptions]) -> None: super().__init__(**kwargs) self._creds_lock = threading.Lock() self._gdch_creds_cache: dict[tuple[str, str], _GDCHAudienceCredentials] = {} diff --git a/litellm/llms/vertex_ai/vertex_ai_partner_models/llama3/transformation.py b/litellm/llms/vertex_ai/vertex_ai_partner_models/llama3/transformation.py index 279035c455d..89a5b8a570e 100644 --- a/litellm/llms/vertex_ai/vertex_ai_partner_models/llama3/transformation.py +++ b/litellm/llms/vertex_ai/vertex_ai_partner_models/llama3/transformation.py @@ -1,6 +1,6 @@ import types from collections.abc import AsyncIterator, Iterator -from typing import TYPE_CHECKING, Any, Final +from typing import TYPE_CHECKING, Final import httpx @@ -95,7 +95,7 @@ class VertexAILlama3Config(OpenAIGPTConfig): streaming_response: Iterator[str] | AsyncIterator[str] | ModelResponse, sync_stream: bool, json_mode: bool | None = False, - ) -> Any: + ) -> "VertexAILlama3StreamingHandler": return VertexAILlama3StreamingHandler( streaming_response=streaming_response, sync_stream=sync_stream, diff --git a/litellm/llms/voyage/rerank/transformation.py b/litellm/llms/voyage/rerank/transformation.py index fea8452d934..6f8d024f0b1 100644 --- a/litellm/llms/voyage/rerank/transformation.py +++ b/litellm/llms/voyage/rerank/transformation.py @@ -4,8 +4,8 @@ Transformation logic for Voyage AI's /v1/rerank endpoint. Docs - https://docs.voyageai.com/docs/reranker """ -from collections.abc import Mapping -from typing import Any, Final +from collections.abc import Mapping, Sequence +from typing import Final import httpx @@ -33,7 +33,7 @@ class VoyageRerankConfig(BaseRerankConfig): model: str, drop_params: bool, query: str, - documents: list[str | dict[str, Any]], + documents: Sequence[str | Mapping[str, object]], custom_llm_provider: str | None = None, top_n: int | None = None, rank_fields: list[str] | None = None, diff --git a/litellm/proxy/client/users.py b/litellm/proxy/client/users.py index 3f11fe94043..503c92228a8 100644 --- a/litellm/proxy/client/users.py +++ b/litellm/proxy/client/users.py @@ -1,3 +1,4 @@ +from collections.abc import Mapping from typing import Any, Final import requests @@ -50,7 +51,7 @@ class UsersManagementClient: response.raise_for_status() return response.json() - def create_user(self, user_data: dict[str, Any]) -> dict[str, Any]: + def create_user(self, user_data: Mapping[str, object]) -> dict[str, Any]: """Create a new user (POST /user/new)""" url: Final = f"{self.base_url}/user/new" response: Final = requests.post(url, headers=self._get_headers(), json=user_data, timeout=self.timeout) diff --git a/litellm/proxy/common_utils/http_parsing_utils.py b/litellm/proxy/common_utils/http_parsing_utils.py index 54a0f18fd63..2bae3e946f7 100644 --- a/litellm/proxy/common_utils/http_parsing_utils.py +++ b/litellm/proxy/common_utils/http_parsing_utils.py @@ -52,14 +52,18 @@ def _unqualified(annotation: object) -> object: return _unqualified(qualified[0]) +def _union_members(annotation: object) -> tuple[object, ...]: + """The non-``None`` members of a union annotation, or the annotation itself when it is not a union.""" + if get_origin(annotation) not in (Union, UnionType): + return (annotation,) + members: Final[tuple[object, ...]] = get_args(annotation) + return tuple(arg for arg in members if arg is not type(None)) + + def _numeric_form_type(annotation: object) -> type[int] | type[float] | None: """The scalar to parse an ``int``/``float``-typed field as, else ``None``.""" unwrapped: Final = _unqualified(annotation) - candidates: Final = ( - tuple(arg for arg in get_args(unwrapped) if arg is not type(None)) - if get_origin(unwrapped) in (Union, UnionType) - else (unwrapped,) - ) + candidates: Final = _union_members(unwrapped) if len(candidates) != 1: return None if candidates[0] is int: diff --git a/litellm/proxy/db/prisma_client.py b/litellm/proxy/db/prisma_client.py index 2190ae55fd2..2f2ebfdf2bb 100644 --- a/litellm/proxy/db/prisma_client.py +++ b/litellm/proxy/db/prisma_client.py @@ -13,7 +13,7 @@ import urllib import urllib.parse from collections.abc import Callable from datetime import datetime, timedelta -from typing import Any, Final, Protocol +from typing import TYPE_CHECKING, Any, Final, Protocol from litellm._logging import verbose_proxy_logger from litellm.proxy.db.token_auth import ( @@ -27,6 +27,9 @@ from litellm.proxy.db.token_auth import ( ) from litellm.secret_managers.main import str_to_bool +if TYPE_CHECKING: + from prisma import Prisma + __all__ = ( "IAMEndpoint", "PrismaManager", @@ -242,7 +245,7 @@ class PrismaWrapper: def _write_engine(prisma_client: _PrismaClient, engine: _PrismaEngine) -> None: prisma_client._Prisma__engine = engine - def _instrument_prisma_client(self, prisma_client: _PrismaClient) -> _PrismaDrainTracker | None: + def _instrument_prisma_client(self, prisma_client: "Prisma | _PrismaClient") -> _PrismaDrainTracker | None: from prisma.errors import ClientNotConnectedError try: @@ -255,7 +258,7 @@ class PrismaWrapper: self._write_engine(prisma_client, _TrackedPrismaEngine(engine, tracker)) return tracker - def _get_engine_pid(self, prisma_client: _PrismaClient | None = None) -> int: + def _get_engine_pid(self, prisma_client: "Prisma | _PrismaClient | None" = None) -> int: """Get the PID of the current Prisma engine subprocess, or 0 if unavailable. Must never raise: it runs inside the reconnect path, where the client diff --git a/litellm/proxy/guardrails/guardrail_hooks/block_code_execution/block_code_execution.py b/litellm/proxy/guardrails/guardrail_hooks/block_code_execution/block_code_execution.py index bf2aa1f76e0..2b697671eda 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/block_code_execution/block_code_execution.py +++ b/litellm/proxy/guardrails/guardrail_hooks/block_code_execution/block_code_execution.py @@ -8,9 +8,10 @@ confidence scoring and a tunable threshold (only block when confidence >= thresh import re from datetime import datetime -from typing import TYPE_CHECKING, Any, Final, Literal, Optional, cast +from typing import TYPE_CHECKING, Final, Literal, Optional, cast from fastapi import HTTPException +from typing_extensions import TypedDict, Unpack from litellm.integrations.custom_guardrail import ( CustomGuardrail, @@ -314,6 +315,10 @@ def _confidence_for_block( return 0.0 +class _CustomGuardrailOptions(TypedDict, total=False, extra_items=object): + """Base-class constructor options this guardrail forwards untouched to CustomGuardrail.""" + + class BlockCodeExecutionGuardrail(CustomGuardrail): """ Guardrail that detects fenced code blocks (markdown ```) and blocks or masks them @@ -332,7 +337,7 @@ class BlockCodeExecutionGuardrail(CustomGuardrail): detect_execution_intent: bool = True, event_hook: Literal["pre_call", "post_call", "during_call"] | list[str] | None = None, default_on: bool = False, - **kwargs: Any, + **kwargs: Unpack[_CustomGuardrailOptions], ) -> None: # Normalize to type expected by CustomGuardrail _event_hook: GuardrailEventHooks | list[GuardrailEventHooks] | None = None diff --git a/litellm/proxy/guardrails/guardrail_hooks/cato_networks/cato_networks.py b/litellm/proxy/guardrails/guardrail_hooks/cato_networks/cato_networks.py index 176c308eda6..2d203c31974 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/cato_networks/cato_networks.py +++ b/litellm/proxy/guardrails/guardrail_hooks/cato_networks/cato_networks.py @@ -264,7 +264,7 @@ class CatoNetworksGuardrail(CustomGuardrail): stack.extend(reversed(node)) @classmethod - def _extra_inspection_sources(cls, data: Mapping[str, Any]) -> Sequence[tuple[str, Sequence[Mapping[str, str]]]]: + def _extra_inspection_sources(cls, data: Mapping[str, object]) -> Sequence[tuple[str, Sequence[Mapping[str, str]]]]: """Text the proxy forwards to the model outside chat ``messages``: Responses-API ``input`` and ``instructions``, legacy completion ``prompt`` and tool/function/``response_format`` schema strings. Returned @@ -336,7 +336,7 @@ class CatoNetworksGuardrail(CustomGuardrail): ) raise HTTPException(status_code=400, detail=detection_message) - def _anonymize_request(self, res: Any, data: dict) -> dict: + def _anonymize_request(self, res: _CatoAnalyzeResponse, data: dict) -> dict: verbose_proxy_logger.info("Cato: anonymize action") redacted_chat: Final = res.get("redacted_chat") if not redacted_chat: @@ -379,7 +379,7 @@ class CatoNetworksGuardrail(CustomGuardrail): return data @classmethod - def _apply_extra_redaction(cls, data: dict, field: str, redacted: list) -> bool: + def _apply_extra_redaction(cls, data: dict, field: str, redacted: Sequence[Mapping[str, object]]) -> bool: if field == "input": input_only: Final = {"input": data["input"]} if not redacted: @@ -400,7 +400,7 @@ class CatoNetworksGuardrail(CustomGuardrail): return True @classmethod - def _apply_schema_string_redaction(cls, data: dict, redacted: list) -> None: + def _apply_schema_string_redaction(cls, data: dict, redacted: Sequence[Mapping[str, object]]) -> None: redactions: Final = iter(redacted) for container, key in cls._iter_schema_string_refs(data): replacement = next(redactions, None) @@ -408,7 +408,7 @@ class CatoNetworksGuardrail(CustomGuardrail): container[key] = replacement["content"] @staticmethod - def _apply_prompt_redaction(data: dict, redacted: list) -> None: + def _apply_prompt_redaction(data: dict, redacted: Sequence[Mapping[str, object]]) -> None: contents: Final = [m.get("content") for m in redacted if isinstance(m, dict)] prompt: Final = data.get("prompt") if isinstance(prompt, str): diff --git a/litellm/proxy/guardrails/guardrail_hooks/compresr/compresr.py b/litellm/proxy/guardrails/guardrail_hooks/compresr/compresr.py index 93d859066b0..1ecdb1b0f63 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/compresr/compresr.py +++ b/litellm/proxy/guardrails/guardrail_hooks/compresr/compresr.py @@ -22,7 +22,7 @@ import json import time from collections import Counter, OrderedDict from dataclasses import dataclass, field -from typing import TYPE_CHECKING, Any, Final, Literal, TypeGuard +from typing import TYPE_CHECKING, Final, Literal, TypeGuard from urllib.parse import urlparse import httpx @@ -64,6 +64,9 @@ if TYPE_CHECKING: from litellm.litellm_core_utils.litellm_logging import ( Logging as LiteLLMLoggingObj, ) + from litellm.llms.base_llm.anthropic_messages.transformation import ( + BaseAnthropicMessagesConfig, + ) from litellm.types.proxy.guardrails.guardrail_hooks.base import ( GuardrailConfigModel, ) @@ -1049,7 +1052,7 @@ class CompresrGuardrail(CustomGuardrail): async def async_should_run_agentic_loop( self, - response: Any, + response: object, model: str, messages: list[dict], tools: list[dict] | None, @@ -1069,8 +1072,8 @@ class CompresrGuardrail(CustomGuardrail): tools: dict, model: str, messages: list[dict], - response: Any, - anthropic_messages_provider_config: Any, + response: object, + anthropic_messages_provider_config: BaseAnthropicMessagesConfig | None, anthropic_messages_optional_request_params: dict, logging_obj: LiteLLMLoggingObj | None, stream: bool, diff --git a/litellm/proxy/guardrails/guardrail_hooks/llm_as_a_judge/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/llm_as_a_judge/__init__.py index 172b1440ca3..806ab5161e8 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/llm_as_a_judge/__init__.py +++ b/litellm/proxy/guardrails/guardrail_hooks/llm_as_a_judge/__init__.py @@ -2,10 +2,10 @@ from collections.abc import Callable, Sequence from datetime import datetime -from typing import TYPE_CHECKING, Any, Final, Generic, Literal, Optional, TypeVar +from typing import TYPE_CHECKING, Final, Generic, Literal, Optional, TypeVar from fastapi import HTTPException -from typing_extensions import NotRequired, ReadOnly, TypedDict +from typing_extensions import NotRequired, ReadOnly, TypedDict, Unpack import litellm from litellm._logging import verbose_logger @@ -106,6 +106,23 @@ def _build_judge_prompt( ) +class _CustomGuardrailOptions(TypedDict, total=False): + """The ``CustomGuardrail`` options this guardrail accepts and forwards untouched.""" + + mask_request_content: ReadOnly[bool] + mask_response_content: ReadOnly[bool] + violation_message_template: ReadOnly[str | None] + end_session_after_n_fails: ReadOnly[int | None] + on_violation: ReadOnly[str | None] + realtime_violation_message: ReadOnly[str | None] + on_sensitive_data: ReadOnly[str | None] + sensitive_data_route_to_model: ReadOnly[str | None] + sticky_session_routing: ReadOnly[bool] + run_in_parallel: ReadOnly[bool] + scan_raw_request: ReadOnly[bool] + only_scan_new_messages: ReadOnly[bool] + + class LLMAsAJudgeGuardrail(CustomGuardrail): """Post-call guardrail that judges response quality via an LLM.""" @@ -119,7 +136,7 @@ class LLMAsAJudgeGuardrail(CustomGuardrail): event_hook: GuardrailEventHooks | list[GuardrailEventHooks] | None = None, default_on: bool = False, router_provider: "Callable[[], Router | None] | None" = None, - **kwargs: Any, + **kwargs: Unpack[_CustomGuardrailOptions], ) -> None: _event_hook: GuardrailEventHooks | list[GuardrailEventHooks] | None = None if event_hook is not None: diff --git a/litellm/proxy/guardrails/guardrail_hooks/noma/noma.py b/litellm/proxy/guardrails/guardrail_hooks/noma/noma.py index 7ef0a9f73f3..edd78e0bbc6 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/noma/noma.py +++ b/litellm/proxy/guardrails/guardrail_hooks/noma/noma.py @@ -9,13 +9,13 @@ import asyncio import json import os import warnings -from collections.abc import AsyncGenerator +from collections.abc import AsyncGenerator, AsyncIterable from datetime import datetime from typing import ( TYPE_CHECKING, - Any, Final, Literal, + TypeVar, ) from urllib.parse import urljoin @@ -39,9 +39,7 @@ from litellm.types.guardrails import GuardrailEventHooks from litellm.types.utils import ( CallTypes, CallTypesLiteral, - EmbeddingResponse, GuardrailStatus, - ImageResponse, ModelResponseStream, TextCompletionResponse, ) @@ -53,7 +51,8 @@ SENSITIVE_DATA_DETECTOR_KEYS: Final[list[str]] = ["sensitiveData", "dataDetector # Type aliases MessageRole = Literal["user", "assistant"] -LLMResponse = Any | ModelResponse | EmbeddingResponse | ImageResponse +LLMResponse = object +_LLMResponseT: Final = TypeVar("_LLMResponseT") _LEGACY_NOMA_DEPRECATION_WARNED = False if TYPE_CHECKING: @@ -709,10 +708,10 @@ class NomaGuardrail(CustomGuardrail): async def _check_llm_response( self, request_data: dict, - response: LLMResponse, + response: _LLMResponseT, user_auth: UserAPIKeyAuth, event_type: GuardrailEventHooks | None = None, - ) -> Any: + ) -> _LLMResponseT: """Check LLM response for policy violations""" content: Final = await self._process_llm_response_check(request_data, response, user_auth, event_type) if not content: @@ -798,7 +797,7 @@ class NomaGuardrail(CustomGuardrail): async def async_post_call_streaming_iterator_hook( self, user_api_key_dict: UserAPIKeyAuth, - response: Any, + response: AsyncIterable[ModelResponseStream], request_data: dict, ) -> AsyncGenerator[ModelResponseStream, None]: """Process streaming response chunks with Noma guardrail.""" diff --git a/litellm/proxy/guardrails/guardrail_hooks/panw_prisma_airs/panw_prisma_airs.py b/litellm/proxy/guardrails/guardrail_hooks/panw_prisma_airs/panw_prisma_airs.py index b73d3adb99e..86ad2f9db5f 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/panw_prisma_airs/panw_prisma_airs.py +++ b/litellm/proxy/guardrails/guardrail_hooks/panw_prisma_airs/panw_prisma_airs.py @@ -793,7 +793,7 @@ class PanwPrismaAirsHandler(CustomGuardrail): }, ) - def _prepare_metadata_from_request(self, data: dict[str, Any]) -> dict[str, Any]: + def _prepare_metadata_from_request(self, data: dict[str, Any]) -> dict[str, object]: """ Extract and prepare metadata from request data for PANW API call. @@ -809,7 +809,7 @@ class PanwPrismaAirsHandler(CustomGuardrail): """ user_metadata: Final = data.get("metadata", {}) or {} requester_meta: Final = user_metadata.get("requester_metadata", {}) or {} - metadata: Final = { + metadata: Final[dict[str, object]] = { "user": data.get("user") or "litellm_user", "model": data.get("model") or "unknown", } diff --git a/litellm/proxy/guardrails/usage_endpoints.py b/litellm/proxy/guardrails/usage_endpoints.py index 6259efb6654..3f7c36bbbf2 100644 --- a/litellm/proxy/guardrails/usage_endpoints.py +++ b/litellm/proxy/guardrails/usage_endpoints.py @@ -45,6 +45,7 @@ _EMPTY_UNITS: Final[Mapping[str, int]] = MappingProxyType({}) _ACTION_SEVERITY: Final[Mapping[str, int]] = MappingProxyType({"passed": 0, "flagged": 1, "blocked": 2}) _T = TypeVar("_T") +_MetricsRowT = TypeVar("_MetricsRowT", bound="_DailyMetricsRow") _USAGE_MAX_RANGE_DAYS: Final = 366 @@ -360,10 +361,12 @@ def _trend_from_comparison(current_fail: float, previous_fail: float) -> str: return "stable" -def _aggregate_daily_metrics(metrics: "Sequence[_DailyMetricsRow]", id_attr: str) -> Mapping[str, _MetricTotals]: +def _aggregate_daily_metrics( + metrics: "Sequence[_MetricsRowT]", id_of: "Callable[[_MetricsRowT], str]" +) -> Mapping[str, _MetricTotals]: agg: Final[dict[str, _MetricTotals]] = {} for m in metrics: - gid: str = getattr(m, id_attr) + gid: str = id_of(m) if gid not in agg: agg[gid] = {"requests": 0, "passed": 0, "blocked": 0, "flagged": 0} agg[gid]["requests"] += int(m.requests_evaluated or 0) @@ -373,10 +376,12 @@ def _aggregate_daily_metrics(metrics: "Sequence[_DailyMetricsRow]", id_attr: str return agg -def _prev_fail_rates(metrics_prev: "Sequence[_DailyMetricsRow]", id_attr: str) -> Mapping[str, float]: +def _prev_fail_rates( + metrics_prev: "Sequence[_MetricsRowT]", id_of: "Callable[[_MetricsRowT], str]" +) -> Mapping[str, float]: prev_agg_raw: Final[dict[str, _PrevPeriodCounts]] = {} for m in metrics_prev: - gid: str = getattr(m, id_attr) + gid: str = id_of(m) r, b = int(m.requests_evaluated or 0), int(m.blocked_count or 0) if gid not in prev_agg_raw: prev_agg_raw[gid] = {"req": 0, "blocked": 0} @@ -429,7 +434,7 @@ def _field_str(mapping: Mapping[str, object], key: str, default: str) -> str: return str(mapping.get(key, default)) -def _get_guardrail_attrs(g: "_DbOrConfigGuardrail") -> tuple[Any, str]: +def _get_guardrail_attrs(g: "_DbOrConfigGuardrail") -> tuple[str | None, str]: """Get (guardrail_id, display_name) from guardrail - handles Prisma model or dict.""" gid: Final = _get_guardrail_field(g, "guardrail_id") name: Final = _get_guardrail_field(g, "guardrail_name") @@ -592,8 +597,8 @@ async def guardrails_usage_overview( Sequence[prisma_models.LiteLLM_DailyGuardrailUsageUnits] ] = await _find_daily_guardrail_usage_units(prisma_client, where=units_where) - agg: Final = _aggregate_daily_metrics(metrics, "guardrail_id") - prev_agg: Final = _prev_fail_rates(metrics_prev, "guardrail_id") + agg: Final = _aggregate_daily_metrics(metrics, lambda m: m.guardrail_id) + prev_agg: Final = _prev_fail_rates(metrics_prev, lambda m: m.guardrail_id) units_agg: Final = _by(units_rows, lambda r: r.guardrail_id, _sum_counter_units) cost_agg: Final = _by(units_rows, lambda r: r.guardrail_id, _sum_tracked_cost) untracked_agg: Final = _by(units_rows, lambda r: r.guardrail_id, _sum_untracked_units) @@ -811,7 +816,7 @@ def _usage_log_entry_from_row( ) -def _snippet(text: Any, max_len: int = 200) -> str | None: +def _snippet(text: object, max_len: int = 200) -> str | None: if text is None: return None if isinstance(text, str): @@ -964,8 +969,8 @@ async def policies_usage_overview( } }, ) - agg: Final = _aggregate_daily_metrics(metrics, "policy_id") - prev_agg: Final = _prev_fail_rates(metrics_prev, "policy_id") + agg: Final = _aggregate_daily_metrics(metrics, lambda m: m.policy_id) + prev_agg: Final = _prev_fail_rates(metrics_prev, lambda m: m.policy_id) chart: Final = _chart_from_metrics(metrics) total_requests: Final = sum(a["requests"] for a in agg.values()) total_blocked: Final = sum(a["blocked"] for a in agg.values()) diff --git a/litellm/proxy/hooks/proxy_track_cost_callback.py b/litellm/proxy/hooks/proxy_track_cost_callback.py index c4fba8ecf9e..e6ada41f062 100644 --- a/litellm/proxy/hooks/proxy_track_cost_callback.py +++ b/litellm/proxy/hooks/proxy_track_cost_callback.py @@ -2,7 +2,7 @@ import asyncio import traceback from collections.abc import Sequence from datetime import datetime -from typing import TYPE_CHECKING, Any, Final, cast +from typing import TYPE_CHECKING, Any, Final, Protocol, cast import litellm from litellm._logging import verbose_proxy_logger @@ -578,18 +578,36 @@ def _get_request_tags_for_cost_tracking( return None +class _IncrementSpendCounters(Protocol): + """The ``increment_spend_counters`` coroutine :func:`_update_database_and_spend_counters` awaits.""" + + async def __call__( + self, + token: str | None, + team_id: str | None, + user_id: str | None, + response_cost: float | None, + org_id: str | None = None, + budget_reservation: dict[str, object] | None = None, + end_user_id: str | None = None, + tags: list[str] | None = None, + request_started_at: datetime | None = None, + model_access_groups: Sequence[str] | None = None, + ) -> None: ... + + async def _update_database_and_spend_counters( proxy_logging_obj: "ProxyLogging", - increment_spend_counters: Any, + increment_spend_counters: _IncrementSpendCounters, user_api_key: str | None, user_id: str | None, end_user_id: str | None, team_id: str | None, org_id: str | None, kwargs: dict, - completion_response: litellm.ModelResponse | Any | None, - start_time: Any, - end_time: Any, + completion_response: object, + start_time: datetime | None, + end_time: datetime | None, response_cost: float, budget_reservation: dict | None, request_tags: list[str] | None = None, diff --git a/litellm/proxy/management_endpoints/model_access_group_management_endpoints.py b/litellm/proxy/management_endpoints/model_access_group_management_endpoints.py index a48130a4f22..e960bdfe337 100644 --- a/litellm/proxy/management_endpoints/model_access_group_management_endpoints.py +++ b/litellm/proxy/management_endpoints/model_access_group_management_endpoints.py @@ -540,7 +540,7 @@ async def get_all_access_groups_from_db( deployments: Final = await ModelRepository(prisma_client).table.find_many() # Build access group map - access_group_map: Final[dict[str, dict[str, Any]]] = {} + model_names_by_group: Final[dict[str, list[str]]] = {} for deployment in deployments: model_info = deployment.model_info or {} @@ -550,25 +550,20 @@ async def get_all_access_groups_from_db( model_name = deployment.model_name for access_group in access_groups: - if access_group not in access_group_map: - access_group_map[access_group] = { - "model_names": set(), - "deployment_count": 0, - } + if access_group not in model_names_by_group: + model_names_by_group[access_group] = [] - access_group_map[access_group]["model_names"].add(model_name) - access_group_map[access_group]["deployment_count"] += 1 + model_names_by_group[access_group].append(model_name) # Convert to AccessGroupInfo objects - result: Final = {} - for access_group, data in access_group_map.items(): - result[access_group] = AccessGroupInfo( + return { + access_group: AccessGroupInfo( access_group=access_group, - model_names=sorted(list(data["model_names"])), - deployment_count=data["deployment_count"], + model_names=sorted(frozenset(model_names)), + deployment_count=len(model_names), ) - - return result + for access_group, model_names in model_names_by_group.items() + } @router.post( diff --git a/litellm/proxy/management_endpoints/usage_endpoints/ai_usage_chat.py b/litellm/proxy/management_endpoints/usage_endpoints/ai_usage_chat.py index da4ddbd0aac..1265da99d89 100644 --- a/litellm/proxy/management_endpoints/usage_endpoints/ai_usage_chat.py +++ b/litellm/proxy/management_endpoints/usage_endpoints/ai_usage_chat.py @@ -6,7 +6,7 @@ usage/spend data by querying the aggregated daily activity endpoints. import json from collections.abc import AsyncGenerator, AsyncIterator, Awaitable, Callable, Mapping, Sequence from datetime import date -from typing import Any, Final, Literal, NamedTuple, Protocol, cast, overload +from typing import Final, Literal, NamedTuple, Protocol, cast, overload from typing_extensions import ReadOnly, TypedDict @@ -16,6 +16,7 @@ from litellm.constants import DEFAULT_COMPETITOR_DISCOVERY_MODEL from litellm.types.proxy.management_endpoints.common_daily_activity import ( SpendAnalyticsPaginatedResponse, ) +from litellm.types.utils import ChatCompletionMessageToolCall # --------------------------------------------------------------------------- # Constants @@ -489,19 +490,19 @@ async def _execute_tool_call( async def _process_tool_call( - tc: Any, + tc: ChatCompletionMessageToolCall, chat_messages: list[Mapping[str, object]], user_id: str | None, is_admin: bool, ) -> AsyncIterator[str]: """Execute a single tool call, yielding SSE events for status.""" - fn_name: Final[str] = tc.function.name + fn_name: Final = tc.function.name fn_args: Final[Mapping[str, str]] = json.loads(tc.function.arguments) allowed_names: Final = {t["function"]["name"] for t in get_tools_for_role(is_admin)} - handler: Final = TOOL_HANDLERS.get(fn_name) + handler: Final = TOOL_HANDLERS.get(fn_name) if fn_name is not None else None - if fn_name not in allowed_names or not handler: + if fn_name is None or fn_name not in allowed_names or not handler: chat_messages.append( { "role": "tool", diff --git a/litellm/proxy/openai_files_endpoints/files_endpoints.py b/litellm/proxy/openai_files_endpoints/files_endpoints.py index bf07f4748ef..4d3a397e519 100644 --- a/litellm/proxy/openai_files_endpoints/files_endpoints.py +++ b/litellm/proxy/openai_files_endpoints/files_endpoints.py @@ -485,8 +485,8 @@ async def create_file( # Parse expires_after if provided expires_after: FileExpiresAfter | None = None form_data_raw: Final = await request.form() - form_data_dict: Final[dict[str, Any]] = dict(form_data_raw) - extracted_litellm_metadata: Final[dict[str, Any] | None] = extract_nested_form_metadata( + form_data_dict: Final[Mapping[str, object]] = dict(form_data_raw) + extracted_litellm_metadata: Final[Mapping[str, object] | None] = extract_nested_form_metadata( form_data=form_data_dict, prefix="litellm_metadata[" ) expires_after_anchor: Final = form_data_raw.get("expires_after[anchor]") diff --git a/litellm/proxy/openai_files_endpoints/storage_backend_service.py b/litellm/proxy/openai_files_endpoints/storage_backend_service.py index e766f335071..0407499cbcc 100644 --- a/litellm/proxy/openai_files_endpoints/storage_backend_service.py +++ b/litellm/proxy/openai_files_endpoints/storage_backend_service.py @@ -7,7 +7,6 @@ storage backends (e.g., Azure Blob Storage) and managing associated metadata. import base64 import time -from collections.abc import Mapping from typing import Any, Final, cast from litellm._logging import verbose_proxy_logger @@ -17,7 +16,7 @@ from litellm.llms.base_llm.files.transformation import BaseFileEndpoints from litellm.proxy._types import ProxyException, UserAPIKeyAuth from litellm.proxy.utils import ProxyLogging from litellm.types.llms.openai import OpenAIFileObject, OpenAIFilesPurpose -from litellm.types.utils import SpecialEnums +from litellm.types.utils import ExtractedFileData, SpecialEnums class StorageBackendFileService: @@ -33,7 +32,7 @@ class StorageBackendFileService: @staticmethod async def upload_file_to_storage_backend( - file_data: Mapping[str, Any], + file_data: ExtractedFileData, target_storage: str, target_model_names: list[str], purpose: OpenAIFilesPurpose, @@ -163,7 +162,7 @@ class StorageBackendFileService: @staticmethod def _create_unified_file_id( - file_type: str, + file_type: str | None, target_model_names: list[str], file_id: str, ) -> str: @@ -193,7 +192,7 @@ class StorageBackendFileService: @staticmethod async def _store_in_managed_files( file_object: OpenAIFileObject, - file_data: Mapping[str, Any], + file_data: ExtractedFileData, target_model_names: list[str], target_storage: str, storage_url: str, diff --git a/litellm/proxy/vector_store_endpoints/endpoints.py b/litellm/proxy/vector_store_endpoints/endpoints.py index 1feda0b0bb5..f21c294e5a2 100644 --- a/litellm/proxy/vector_store_endpoints/endpoints.py +++ b/litellm/proxy/vector_store_endpoints/endpoints.py @@ -244,7 +244,7 @@ async def vector_store_create( ) # Create vector store across multiple models - response: Final = await managed_vector_stores.acreate_vector_store( + response: Final[object] = await managed_vector_stores.acreate_vector_store( create_request=data, llm_router=llm_router, target_model_names_list=target_model_names_list, diff --git a/litellm/proxy/vector_store_files_endpoints/endpoints.py b/litellm/proxy/vector_store_files_endpoints/endpoints.py index 957ed9fd0b9..3ddea288ab8 100644 --- a/litellm/proxy/vector_store_files_endpoints/endpoints.py +++ b/litellm/proxy/vector_store_files_endpoints/endpoints.py @@ -553,7 +553,7 @@ async def vector_store_file_create( processor: Final = ProxyBaseLLMRequestProcessing(data=data) try: - response = await processor.base_process_llm_request( + response: object = await processor.base_process_llm_request( request=request, fastapi_response=fastapi_response, user_api_key_dict=user_api_key_dict, @@ -756,7 +756,7 @@ async def vector_store_file_retrieve( processor: Final = ProxyBaseLLMRequestProcessing(data=data) try: - response = await processor.base_process_llm_request( + response: object = await processor.base_process_llm_request( request=request, fastapi_response=fastapi_response, user_api_key_dict=user_api_key_dict, @@ -863,7 +863,7 @@ async def vector_store_file_content( processor: Final = ProxyBaseLLMRequestProcessing(data=data) try: - response = await processor.base_process_llm_request( + response: object = await processor.base_process_llm_request( request=request, fastapi_response=fastapi_response, user_api_key_dict=user_api_key_dict, @@ -973,7 +973,7 @@ async def vector_store_file_update( processor: Final = ProxyBaseLLMRequestProcessing(data=data) try: - response = await processor.base_process_llm_request( + response: object = await processor.base_process_llm_request( request=request, fastapi_response=fastapi_response, user_api_key_dict=user_api_key_dict, @@ -1080,7 +1080,7 @@ async def vector_store_file_delete( processor: Final = ProxyBaseLLMRequestProcessing(data=data) try: - response = await processor.base_process_llm_request( + response: object = await processor.base_process_llm_request( request=request, fastapi_response=fastapi_response, user_api_key_dict=user_api_key_dict, diff --git a/litellm/proxy_auth/credentials.py b/litellm/proxy_auth/credentials.py index a4e29241959..f8814954a7a 100644 --- a/litellm/proxy_auth/credentials.py +++ b/litellm/proxy_auth/credentials.py @@ -7,7 +7,7 @@ It follows the same TokenCredential protocol used by Azure SDK. import time from dataclasses import dataclass -from typing import Any, Final, Protocol, runtime_checkable +from typing import Final, Protocol, runtime_checkable @dataclass @@ -50,6 +50,22 @@ class TokenCredential(Protocol): ... +class _AzureAccessToken(Protocol): + """The two attributes :class:`AzureADCredential` reads off an azure-identity token.""" + + @property + def token(self) -> str: ... + + @property + def expires_on(self) -> int: ... + + +class _AzureTokenCredential(Protocol): + """The single method :class:`AzureADCredential` calls on the credential it wraps.""" + + def get_token(self, *scopes: str) -> _AzureAccessToken: ... + + class AzureADCredential: """ Wrapper for Azure Identity credentials. @@ -71,7 +87,7 @@ class AzureADCredential: cred = AzureADCredential(credential=azure_cred) """ - def __init__(self, credential: Any | None = None): + def __init__(self, credential: _AzureTokenCredential | None = None): """ Initialize with an optional Azure credential. @@ -79,7 +95,7 @@ class AzureADCredential: credential: An azure-identity credential object. If None, DefaultAzureCredential will be used on first token request. """ - self._credential: Any = credential + self._credential: _AzureTokenCredential | None = credential self._initialized = credential is not None def get_token(self, scope: str) -> AccessToken: @@ -95,20 +111,30 @@ class AzureADCredential: Raises: ImportError: If azure-identity is not installed. """ - if not self._initialized: - try: - from azure.identity import DefaultAzureCredential - - self._credential = DefaultAzureCredential() - self._initialized = True - except ImportError: - raise ImportError( - "azure-identity is required for AzureADCredential. Install it with: pip install azure-identity" - ) - - result: Final = self._credential.get_token(scope) + result: Final = self._resolve_credential().get_token(scope) return AccessToken(token=result.token, expires_on=result.expires_on) + def _resolve_credential(self) -> _AzureTokenCredential: + """Return the wrapped credential, building the Azure default chain on first use. + + Raises: + ImportError: If azure-identity is not installed. + """ + existing: Final = self._credential + if existing is not None: + return existing + try: + from azure.identity import DefaultAzureCredential + + created: Final = DefaultAzureCredential() + except ImportError: + raise ImportError( + "azure-identity is required for AzureADCredential. Install it with: pip install azure-identity" + ) + self._credential = created + self._initialized = True + return created + class GenericOAuth2Credential: """ @@ -228,7 +254,7 @@ class ProxyAuthHandler: self._cached_token = self.credential.get_token(self.scope) return self._cached_token - def get_auth_headers(self) -> dict: + def get_auth_headers(self) -> dict[str, str]: """ Get HTTP headers for authentication. diff --git a/litellm/rag/rag_query.py b/litellm/rag/rag_query.py index 255faf94402..9325547c17d 100644 --- a/litellm/rag/rag_query.py +++ b/litellm/rag/rag_query.py @@ -124,9 +124,9 @@ class RAGQuery: @staticmethod def extract_documents_from_search( search_response: Any, - ) -> list[str | dict[str, Any]]: + ) -> list[str | dict[str, object]]: """Extract text documents from vector store search response.""" - documents: Final[list[str | dict[str, Any]]] = [] + documents: Final[list[str | dict[str, object]]] = [] search_data: Final[_SearchDataView] = {"results": search_response.get("data", [])} for result in search_data["results"]: content_list = result.get("content", []) diff --git a/litellm/repositories/verification_token_repository.py b/litellm/repositories/verification_token_repository.py index c0e59f9b975..d02c2114136 100644 --- a/litellm/repositories/verification_token_repository.py +++ b/litellm/repositories/verification_token_repository.py @@ -5,7 +5,8 @@ VerificationToken repository for database operations on LiteLLM_VerificationToke import json from collections.abc import Mapping, Sequence from datetime import datetime -from typing import TYPE_CHECKING, Final +from types import TracebackType +from typing import TYPE_CHECKING, Final, Protocol from litellm.models.verification_token import ( LiteLLM_VerificationToken, @@ -25,7 +26,36 @@ if TYPE_CHECKING: LiteLLM_VerificationToken as PrismaVerificationToken, ) - from litellm.proxy.utils import PrismaClient + +class _VerificationTokenTables(Protocol): + """The two verification token tables this repository reads and writes.""" + + @property + def litellm_verificationtoken(self) -> TableActions["PrismaVerificationToken"]: ... + + @property + def litellm_deletedverificationtoken(self) -> TableActions["PrismaDeletedVerificationToken"]: ... + + +class _VerificationTokenTransactionManager(Protocol): + async def __aenter__(self) -> _VerificationTokenTables: ... + + async def __aexit__( + self, + exc_type: type[BaseException] | None, + exc_value: BaseException | None, + traceback: TracebackType | None, + ) -> bool | None: ... + + +class _PrismaVerificationTokenDb(_VerificationTokenTables, Protocol): + def tx(self) -> _VerificationTokenTransactionManager: ... + + +class _PrismaClientView(Protocol): + @property + def db(self) -> _PrismaVerificationTokenDb: ... + _JSON_ENCODED_TOKEN_FIELDS: Final = ( "aliases", @@ -44,17 +74,17 @@ class VerificationTokenRepository(BaseRepository[LiteLLM_VerificationToken]): """Repository for verification token (API key) database operations.""" @property - def prisma_client(self) -> "PrismaClient": - prisma_client: Final[PrismaClient] = super().prisma_client - return prisma_client + def _db(self) -> _PrismaVerificationTokenDb: + client: Final[_PrismaClientView] = self.prisma_client + return client.db @property def table(self) -> TableActions["PrismaVerificationToken"]: - return self.prisma_client.db.litellm_verificationtoken + return self._db.litellm_verificationtoken @property def deleted_table(self) -> TableActions["PrismaDeletedVerificationToken"]: - return self.prisma_client.db.litellm_deletedverificationtoken + return self._db.litellm_deletedverificationtoken @property def model_class(self) -> type[LiteLLM_VerificationToken]: @@ -325,7 +355,7 @@ class VerificationTokenRepository(BaseRepository[LiteLLM_VerificationToken]): archive_data["litellm_changed_by"] = litellm_changed_by archive_data["deleted_at"] = datetime.utcnow() - async with self.prisma_client.db.tx() as tx: + async with self._db.tx() as tx: await tx.litellm_deletedverificationtoken.create(data=archive_data) await tx.litellm_verificationtoken.delete(where={"token": token}) diff --git a/litellm/router_utils/search_api_router.py b/litellm/router_utils/search_api_router.py index 309894957ea..76e833563ba 100644 --- a/litellm/router_utils/search_api_router.py +++ b/litellm/router_utils/search_api_router.py @@ -15,7 +15,7 @@ from typing import TYPE_CHECKING, Any, Final, Protocol from litellm._logging import verbose_router_logger if TYPE_CHECKING: - from litellm.types.router import SearchToolTypedDict + from litellm.types.router import SearchToolLiteLLMParams, SearchToolTypedDict class _SearchToolsRouter(Protocol): @@ -34,7 +34,7 @@ class SearchAPIRouter: @staticmethod def _resolve_search_provider_credentials( *, - tool_litellm_params: dict[str, Any], + tool_litellm_params: "SearchToolLiteLLMParams", ) -> tuple[str | None, str | None]: """ Resolve search provider credentials from tool configuration ONLY. From 1ec5083ab4d94c758227c22be052f69e2ae2cf49 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Tue, 8 Sep 2026 12:41:07 +0000 Subject: [PATCH 06/44] refactor(types): replace Any with real types across 54 more backend files Third batch of the fifth basedpyright Any reduction round. Every change is typing-only and leaves runtime behavior identical. Provider transformation configs, video and rerank base classes, OTel metadata and the guardrail and realtime type modules move their payload, header and optional-parameter annotations from Any to object, Mapping[str, object] or the concrete model the call site already produces. Repositories and endpoints that reached Prisma through an untyped handle now name the actions they call with the repo's own TableActions protocol. The pydantic field retypes were checked against pydantic to confirm object and Any validate, serialize and generate JSON schema identically. --- .../providers/pydantic_ai_agents/config.py | 4 +-- litellm/integrations/argilla.py | 7 +++-- litellm/integrations/dynamodb.py | 19 +++++++++++-- litellm/integrations/focus/database.py | 6 ++-- .../focus/destinations/vantage_destination.py | 5 ++-- .../langfuse/langfuse_otel_attributes.py | 3 +- litellm/integrations/langsmith.py | 8 +++--- litellm/integrations/otel/model/metadata.py | 14 +++++----- litellm/integrations/otel/mount.py | 19 +++++++++++-- .../litellm_core_utils/coroutine_checker.py | 6 ++-- .../exception_mapping_utils.py | 6 ++-- .../huggingface_template_handler.py | 26 +++++++++++++---- .../anthropic/count_tokens/handler.py | 2 +- .../llms/base_llm/videos/transformation.py | 10 +++---- .../count_tokens/bedrock_token_counter.py | 11 ++++---- litellm/llms/bedrock/files/transformation.py | 6 ++-- .../guardrail_translation/handler.py | 5 +++- litellm/llms/chatgpt/common_utils.py | 4 +-- .../responses/transformation.py | 7 +++-- .../llms/hosted_vllm/chat/transformation.py | 21 ++++++-------- litellm/llms/openai/common_utils.py | 6 ++-- litellm/llms/snowflake/chat/transformation.py | 4 +-- .../stability/image_edit/transformations.py | 4 +-- .../llms/triton/completion/transformation.py | 8 +++--- .../llms/vertex_ai/rerank/transformation.py | 4 +-- .../volcengine/embedding/transformation.py | 9 +++--- litellm/llms/watsonx/rerank/transformation.py | 8 +++--- litellm/llms/xai/realtime/transformation.py | 18 ++++++------ .../analytics_endpoints/cache_activity.py | 20 +++++++++---- .../proxy/client/cli/commands/model_groups.py | 18 ++++++++---- litellm/proxy/client/cli/commands/up.py | 3 +- .../common_utils/cache_pydantic_utils.py | 2 +- .../common_utils/proxy_rate_limit_error.py | 4 +-- .../proxy/container_endpoints/endpoints.py | 10 ++++--- litellm/proxy/db/exception_handler.py | 16 +++++++++-- litellm/proxy/db/spend_log_tool_index.py | 2 +- .../crowdstrike_aidr/crowdstrike_aidr.py | 16 +++++------ .../guardrail_hooks/enkryptai/enkryptai.py | 4 +-- .../guardrail_hooks/qualifire/qualifire.py | 2 +- .../semantic_guard/route_loader.py | 3 +- .../proxy/hooks/parallel_request_limiter.py | 6 ++-- .../common_daily_activity.py | 7 +++-- litellm/proxy/ocr_endpoints/endpoints.py | 14 +++++----- litellm/repositories/budget_repository.py | 28 +++++++++++++++---- .../repositories/organization_repository.py | 28 +++++++++++++++---- litellm/repositories/project_repository.py | 11 ++++---- .../adaptive_router/adaptive_router.py | 9 +++--- .../adaptive_router/signals.py | 8 +++--- .../adaptive_router/update_queue.py | 11 ++++---- litellm/types/containers/main.py | 2 +- litellm/types/guardrails.py | 16 +++++------ litellm/types/integrations/prometheus.py | 8 +++--- litellm/types/realtime.py | 16 +++++------ litellm/vector_store_files/utils.py | 11 ++++---- 54 files changed, 323 insertions(+), 202 deletions(-) diff --git a/litellm/a2a_protocol/providers/pydantic_ai_agents/config.py b/litellm/a2a_protocol/providers/pydantic_ai_agents/config.py index b7546e1a2a1..20404e3702b 100644 --- a/litellm/a2a_protocol/providers/pydantic_ai_agents/config.py +++ b/litellm/a2a_protocol/providers/pydantic_ai_agents/config.py @@ -23,7 +23,7 @@ class PydanticAIProviderConfig(BaseA2AProviderConfig): params: dict[str, Any], api_base: str | None = None, **kwargs: Any, - ) -> dict[str, Any]: + ) -> dict[str, object]: """Handle non-streaming request to Pydantic AI agent.""" if api_base is None: raise ValueError("api_base is required for PydanticAIProviderConfig") @@ -41,7 +41,7 @@ class PydanticAIProviderConfig(BaseA2AProviderConfig): params: dict[str, Any], api_base: str | None = None, **kwargs, - ) -> AsyncIterator[dict[str, Any]]: + ) -> AsyncIterator[dict[str, object]]: """Handle streaming request with fake streaming.""" if not api_base: raise ValueError("api_base is required for Pydantic AI agents") diff --git a/litellm/integrations/argilla.py b/litellm/integrations/argilla.py index 9a87a94cf0b..664ef8efda1 100644 --- a/litellm/integrations/argilla.py +++ b/litellm/integrations/argilla.py @@ -7,7 +7,8 @@ import json import os import random import types -from typing import Any, Final +from collections.abc import Mapping +from typing import Final import httpx from pydantic import BaseModel @@ -69,7 +70,7 @@ class ArgillaLogger(CustomBatchLogger): self.flush_lock = asyncio.Lock() super().__init__(**kwargs, flush_lock=self.flush_lock) - def validate_argilla_transformation_object(self, argilla_transformation_object: dict[str, Any]): + def validate_argilla_transformation_object(self, argilla_transformation_object: Mapping[str, object]): if not isinstance(argilla_transformation_object, dict): raise Exception("'argilla_transformation_object' must be a dictionary, to log your payload to Argilla.") @@ -115,7 +116,7 @@ class ArgillaLogger(CustomBatchLogger): ARGILLA_DATASET_NAME=_credentials_dataset_name, ) - def get_chat_messages(self, payload: StandardLoggingPayload) -> list[dict[str, Any]]: + def get_chat_messages(self, payload: StandardLoggingPayload) -> list[dict[str, object]]: payload_messages: Final = payload.get("messages", None) if payload_messages is None: diff --git a/litellm/integrations/dynamodb.py b/litellm/integrations/dynamodb.py index 38f5924a233..3401ced4efb 100644 --- a/litellm/integrations/dynamodb.py +++ b/litellm/integrations/dynamodb.py @@ -3,12 +3,25 @@ import os import traceback -from typing import Any, Final +from collections.abc import Mapping +from typing import Final, Protocol import litellm from litellm._uuid import uuid +class _DynamoTable(Protocol): + """The one boto3 DynamoDB table call this logger makes.""" + + def put_item(self, *, Item: Mapping[str, object]) -> object: ... + + +class _DynamoResource(Protocol): + """The one boto3 DynamoDB resource call this logger makes.""" + + def Table(self, name: str) -> _DynamoTable: ... + + class DyanmoDBLogger: # Class variables or attributes @@ -16,7 +29,7 @@ class DyanmoDBLogger: # Instance variables import boto3 - self.dynamodb: Any = boto3.resource("dynamodb", region_name=os.environ["AWS_REGION_NAME"]) + self.dynamodb: Final[_DynamoResource] = boto3.resource("dynamodb", region_name=os.environ["AWS_REGION_NAME"]) if litellm.dynamodb_table_name is None: raise ValueError( "LiteLLM Error, trying to use DynamoDB but not table name passed. Create a table and set `litellm.dynamodb_table_name=`" @@ -41,7 +54,7 @@ class DyanmoDBLogger: id: Final = response_obj.get("id", str(uuid.uuid4())) # Build the initial payload - payload: Final = { + payload: Final[dict[str, object]] = { "id": id, "call_type": call_type, "startTime": start_time, diff --git a/litellm/integrations/focus/database.py b/litellm/integrations/focus/database.py index 657c7e0d264..891318f1c54 100644 --- a/litellm/integrations/focus/database.py +++ b/litellm/integrations/focus/database.py @@ -3,7 +3,7 @@ from __future__ import annotations from datetime import datetime -from typing import Any, Final +from typing import Final import polars as pl @@ -32,7 +32,7 @@ class FocusLiteLLMDatabase: client: Final = self._ensure_prisma_client() where_clauses: Final[list[str]] = [] - query_params: Final[list[Any]] = [] + query_params: Final[list[datetime | int]] = [] placeholder_index = 1 if start_time_utc: where_clauses.append(f"dus.updated_at >= ${placeholder_index}::timestamptz") @@ -112,7 +112,7 @@ class FocusLiteLLMDatabase: except Exception as exc: raise RuntimeError(f"Error retrieving usage data: {exc}") from exc - async def get_table_info(self) -> dict[str, Any]: + async def get_table_info(self) -> dict[str, object]: """Return metadata about the spend table for diagnostics.""" client: Final = self._ensure_prisma_client() diff --git a/litellm/integrations/focus/destinations/vantage_destination.py b/litellm/integrations/focus/destinations/vantage_destination.py index 132f27779c2..68b0d399975 100644 --- a/litellm/integrations/focus/destinations/vantage_destination.py +++ b/litellm/integrations/focus/destinations/vantage_destination.py @@ -4,7 +4,8 @@ from __future__ import annotations import csv import io -from typing import Any, Final +from collections.abc import Mapping +from typing import Final import httpx # noqa: F401 - used at runtime (AsyncClient, HTTPStatusError) @@ -94,7 +95,7 @@ class FocusVantageDestination(FocusDestination): self, *, prefix: str, - config: dict[str, Any] | None = None, + config: Mapping[str, object] | None = None, ) -> None: config = config or {} api_key: Final = config.get("api_key") diff --git a/litellm/integrations/langfuse/langfuse_otel_attributes.py b/litellm/integrations/langfuse/langfuse_otel_attributes.py index 70fea1abb3b..1eb3ce1a9c2 100644 --- a/litellm/integrations/langfuse/langfuse_otel_attributes.py +++ b/litellm/integrations/langfuse/langfuse_otel_attributes.py @@ -5,6 +5,7 @@ Relevant Issue: https://github.com/BerriAI/litellm/issues/13764 """ import json +from collections.abc import Mapping from typing import TYPE_CHECKING, Any, Final from pydantic import BaseModel @@ -40,7 +41,7 @@ def get_output_content_by_type( | HttpxBinaryResponseContent | ResponsesAPIResponse | list, - kwargs: dict[str, Any] | None = None, + kwargs: Mapping[str, object] | None = None, ) -> str: """ Extract output content from response objects based on their type. diff --git a/litellm/integrations/langsmith.py b/litellm/integrations/langsmith.py index 9607eccef52..2c3837b12ec 100644 --- a/litellm/integrations/langsmith.py +++ b/litellm/integrations/langsmith.py @@ -75,9 +75,9 @@ class LangsmithLogger(CustomBatchLogger): if _batch_size: self.batch_size = int(_batch_size) self.log_queue: list[LangsmithQueueObject] = [] - self._flush_task: asyncio.Task[Any] | None = self._start_periodic_flush_task() + self._flush_task: asyncio.Task[None] | None = self._start_periodic_flush_task() - def _start_periodic_flush_task(self) -> asyncio.Task[Any] | None: + def _start_periodic_flush_task(self) -> asyncio.Task[None] | None: """Start the periodic flush task only when an event loop is already running.""" try: loop: Final = asyncio.get_running_loop() @@ -152,9 +152,9 @@ class LangsmithLogger(CustomBatchLogger): return self._redact_metadata(extra_metadata) - def _build_outputs_with_usage(self, payload: StandardLoggingPayload) -> dict[str, Any]: + def _build_outputs_with_usage(self, payload: StandardLoggingPayload) -> dict[str, object]: response: Final = payload["response"] - outputs: dict[str, Any] + outputs: dict[str, object] if isinstance(response, dict): outputs = {**response} else: diff --git a/litellm/integrations/otel/model/metadata.py b/litellm/integrations/otel/model/metadata.py index ee116aca46b..9eab3021d58 100644 --- a/litellm/integrations/otel/model/metadata.py +++ b/litellm/integrations/otel/model/metadata.py @@ -36,7 +36,7 @@ model. They coincide on the SDK path, which is correct. from __future__ import annotations -from collections.abc import Iterator, Mapping +from collections.abc import Callable, Iterator, Mapping from dataclasses import dataclass, field from types import MappingProxyType from typing import TYPE_CHECKING, Any, Final, cast @@ -57,7 +57,7 @@ class RequestIdentity: # The team's free-form metadata, carried raw (empty/missing -> None) and # filtered to an operator allowlist only at Baggage-promotion time, so an # unconfigured deployment never promotes any of it. - team_metadata: Mapping[str, Any] | None = None + team_metadata: Mapping[str, object] | None = None key_hash: str | None = None end_user: str | None = None # The model litellm dispatched to the provider. Only known once the call @@ -103,7 +103,7 @@ class RequestIdentity: ``user_api_key_*`` names that ``baggage.DEFAULT_BAGGAGE_METADATA_KEYS`` promotes. """ - get: Final = lambda name: getattr(auth, name, None) # noqa: E731 + get: Final[Callable[[str], object]] = lambda name: getattr(auth, name, None) # noqa: E731 metadata: Final = { meta_key: str(value) for meta_key, attr in ( @@ -217,7 +217,7 @@ class LLMCallEvent: time_to_first_chunk_seconds: float | None @classmethod - def from_dict(cls, kwargs: Mapping[str, Any]) -> LLMCallEvent: + def from_dict(cls, kwargs: Mapping[str, object]) -> LLMCallEvent: raw_payload: Final = kwargs.get("standard_logging_object") payload: Final = cast("StandardLoggingPayload", raw_payload) if raw_payload else None operation: Final = resolve_operation(as_str(kwargs.get("call_type"))) @@ -239,7 +239,7 @@ def time_to_first_chunk_seconds(kwargs: Mapping[str, Any]) -> float | None: to the first streamed chunk (``completion_start_time``); ``None`` for non-streaming calls, where ``completion_start_time`` is backfilled with the end time and would not measure first-chunk latency.""" - optional_params: Final = cast(Mapping[str, Any], kwargs.get("optional_params") or {}) + optional_params: Final = cast(Mapping[str, object], kwargs.get("optional_params") or {}) if not optional_params.get("stream"): return None api_call_start: Final = to_seconds(kwargs.get("api_call_start_time")) @@ -307,7 +307,7 @@ def _metadata_dicts( ) -def _call_id(payload: StandardLoggingPayload | None, kwargs: Mapping[str, Any]) -> str | None: +def _call_id(payload: StandardLoggingPayload | None, kwargs: Mapping[str, object]) -> str | None: """The call id from the payload (when closed) or the bare kwargs (at pre_call).""" if payload is not None: call_id: Final = as_str(payload.get("litellm_call_id")) or as_str(payload.get("id")) @@ -351,7 +351,7 @@ def _model_info_id(model_info: object) -> str | None: return None -def _team_metadata_dict(value: object) -> Mapping[str, Any] | None: +def _team_metadata_dict(value: object) -> Mapping[str, object] | None: """The team's free-form metadata as a raw mapping, or ``None`` when missing or empty. diff --git a/litellm/integrations/otel/mount.py b/litellm/integrations/otel/mount.py index ac647c2c4f6..776d8722d14 100644 --- a/litellm/integrations/otel/mount.py +++ b/litellm/integrations/otel/mount.py @@ -12,11 +12,14 @@ when the feature gate is off. """ import os -from typing import Any, Final +from typing import TYPE_CHECKING, Final, Protocol from litellm._logging import verbose_logger from litellm.integrations.otel.model.config import is_otel_v2_enabled +if TYPE_CHECKING: + from fastapi import FastAPI + # Routes excluded from server-span tracing by default: high-frequency pollers and # static UI/docs assets, none of which are LLM traffic. Entries are substring-matched # against the request path (unanchored, so they survive a ``server_root_path`` prefix @@ -65,7 +68,17 @@ PASSTHROUGH_PREFIXES: Final = frozenset( ) -def _passthrough_span_name_hook(span: Any, scope: dict) -> None: +class _RenameableSpan(Protocol): + """The span surface the passthrough naming hook drives.""" + + def is_recording(self) -> bool: ... + + def update_name(self, name: str) -> None: ... + + def set_attribute(self, key: str, value: str) -> None: ... + + +def _passthrough_span_name_hook(span: "_RenameableSpan | None", scope: dict) -> None: """FastAPI ``server_request_hook``: give passthrough server spans a useful name. The instrumentation matches the route at span creation, so both the span name @@ -88,7 +101,7 @@ def _passthrough_span_name_hook(span: Any, scope: dict) -> None: pass -def instrument_fastapi_app(app: Any) -> None: +def instrument_fastapi_app(app: "FastAPI") -> None: """Attach OTel server-span instrumentation to the proxy FastAPI app. Safe no-op when the V2 gate is off or ``opentelemetry-instrumentation-fastapi`` diff --git a/litellm/litellm_core_utils/coroutine_checker.py b/litellm/litellm_core_utils/coroutine_checker.py index 52fc44ba8dc..a5d4f8b968f 100644 --- a/litellm/litellm_core_utils/coroutine_checker.py +++ b/litellm/litellm_core_utils/coroutine_checker.py @@ -16,7 +16,7 @@ class CoroutineChecker: """ def __init__(self): - self._cache = WeakKeyDictionary() + self._cache: WeakKeyDictionary[object, bool] = WeakKeyDictionary() self._max_size = COROUTINE_CHECKER_MAX_SIZE_IN_MEMORY def is_async_callable(self, callback: Any) -> bool: @@ -33,10 +33,10 @@ class CoroutineChecker: pass # Determine target - optimized path for common cases - target = callback + target: object = callback if not inspect.isfunction(target) and not inspect.ismethod(target): try: - call_attr: Final = getattr(target, "__call__", None) + call_attr: Final[object] = getattr(target, "__call__", None) if call_attr is not None: target = call_attr except Exception: diff --git a/litellm/litellm_core_utils/exception_mapping_utils.py b/litellm/litellm_core_utils/exception_mapping_utils.py index 82708d412c9..92f07057da5 100644 --- a/litellm/litellm_core_utils/exception_mapping_utils.py +++ b/litellm/litellm_core_utils/exception_mapping_utils.py @@ -1,7 +1,7 @@ import json import re import traceback -from typing import Any, Final, Protocol, cast +from typing import Final, Protocol, cast import httpx @@ -191,7 +191,7 @@ def _get_response_headers(original_exception: Exception) -> httpx.Headers | None _response_headers: httpx.Headers | None = None try: _response_headers = getattr(original_exception, "headers", None) - error_response: Final = getattr(original_exception, "response", None) + error_response: Final[object] = getattr(original_exception, "response", None) if not _response_headers and error_response: _response_headers = getattr(error_response, "headers", None) if not _response_headers: @@ -203,7 +203,7 @@ def _get_response_headers(original_exception: Exception) -> httpx.Headers | None def extract_and_raise_litellm_exception( - response: Any | None, + response: object | None, error_str: str, model: str, custom_llm_provider: str, diff --git a/litellm/litellm_core_utils/prompt_templates/huggingface_template_handler.py b/litellm/litellm_core_utils/prompt_templates/huggingface_template_handler.py index 3878b36cd91..8f8228d6dfd 100644 --- a/litellm/litellm_core_utils/prompt_templates/huggingface_template_handler.py +++ b/litellm/litellm_core_utils/prompt_templates/huggingface_template_handler.py @@ -1,6 +1,8 @@ import json from datetime import datetime -from typing import Any, Final +from typing import Any, Final, Literal + +from typing_extensions import NotRequired, ReadOnly, TypedDict from litellm.llms.custom_httpx.http_handler import ( _get_httpx_client, @@ -9,6 +11,20 @@ from litellm.llms.custom_httpx.http_handler import ( from litellm.types.llms.custom_http import httpxSpecialProvider +class _TokenizerConfigResult(TypedDict): + """Outcome of a tokenizer_config.json fetch, carrying the parsed document when the fetch succeeded.""" + + status: ReadOnly[Literal["success", "failure"]] + tokenizer: NotRequired[ReadOnly[object]] + + +class _ChatTemplateFileResult(TypedDict): + """Outcome of a chat template file fetch, carrying the template body when the fetch succeeded.""" + + status: ReadOnly[Literal["success", "failure"]] + chat_template: NotRequired[ReadOnly[str]] + + def strftime_now(fmt: str) -> str: """ Custom function for templates that need current date/time formatting (e.g., gpt-oss) @@ -22,7 +38,7 @@ def strftime_now(fmt: str) -> str: return datetime.now().strftime(fmt) -def _get_tokenizer_config(hf_model_name: str) -> dict[str, Any]: +def _get_tokenizer_config(hf_model_name: str) -> _TokenizerConfigResult: """ Fetch tokenizer_config.json from HuggingFace (sync) @@ -45,7 +61,7 @@ def _get_tokenizer_config(hf_model_name: str) -> dict[str, Any]: return {"status": "failure"} -async def _aget_tokenizer_config(hf_model_name: str) -> dict[str, Any]: +async def _aget_tokenizer_config(hf_model_name: str) -> _TokenizerConfigResult: """ Fetch tokenizer_config.json from HuggingFace (async) @@ -70,7 +86,7 @@ async def _aget_tokenizer_config(hf_model_name: str) -> dict[str, Any]: return {"status": "failure"} -def _get_chat_template_file(hf_model_name: str) -> dict[str, Any]: +def _get_chat_template_file(hf_model_name: str) -> _ChatTemplateFileResult: """ Fetch chat template from separate .jinja file (sync) @@ -98,7 +114,7 @@ def _get_chat_template_file(hf_model_name: str) -> dict[str, Any]: return {"status": "failure"} -async def _aget_chat_template_file(hf_model_name: str) -> dict[str, Any]: +async def _aget_chat_template_file(hf_model_name: str) -> _ChatTemplateFileResult: """ Fetch chat template from separate .jinja file (async) diff --git a/litellm/llms/azure_ai/anthropic/count_tokens/handler.py b/litellm/llms/azure_ai/anthropic/count_tokens/handler.py index 3cc90823af9..36d5a56db0d 100644 --- a/litellm/llms/azure_ai/anthropic/count_tokens/handler.py +++ b/litellm/llms/azure_ai/anthropic/count_tokens/handler.py @@ -33,7 +33,7 @@ class AzureAIAnthropicCountTokensHandler(AzureAIAnthropicCountTokensConfig): litellm_params: dict[str, Any] | None = None, timeout: float | httpx.Timeout | None = None, tools: list[dict[str, Any]] | None = None, - system: Any | None = None, + system: object = None, ) -> dict[str, Any]: """ Handle a CountTokens request using httpx with Azure authentication. diff --git a/litellm/llms/base_llm/videos/transformation.py b/litellm/llms/base_llm/videos/transformation.py index f725b295d0f..88cdcd61dd6 100644 --- a/litellm/llms/base_llm/videos/transformation.py +++ b/litellm/llms/base_llm/videos/transformation.py @@ -180,7 +180,7 @@ class BaseVideoConfig(ABC): api_base: str, litellm_params: GenericLiteLLMParams, headers: dict, - extra_body: dict[str, Any] | None = None, + extra_body: dict[str, object] | None = None, ) -> tuple[str, dict]: """ Transform the video remix request into a URL and data @@ -207,7 +207,7 @@ class BaseVideoConfig(ABC): after: str | None = None, limit: int | None = None, order: str | None = None, - extra_query: dict[str, Any] | None = None, + extra_query: dict[str, object] | None = None, ) -> tuple[str, dict]: """ Transform the video list request into a URL and params @@ -342,8 +342,8 @@ class BaseVideoConfig(ABC): litellm_params: GenericLiteLLMParams, headers: dict, video_file: FileContent | None = None, - extra_body: dict[str, Any] | None = None, - prefetched_source_data: dict[str, Any] | None = None, + extra_body: dict[str, object] | None = None, + prefetched_source_data: dict[str, object] | None = None, ) -> tuple[str, Mapping[str, object], RequestFiles | None]: """ Transform the video edit request into a URL plus either JSON data or @@ -373,7 +373,7 @@ class BaseVideoConfig(ABC): api_base: str, litellm_params: GenericLiteLLMParams, headers: dict, - extra_body: dict[str, Any] | None = None, + extra_body: dict[str, object] | None = None, ) -> tuple[str, dict]: """ Transform the video extension request into a URL and JSON data. diff --git a/litellm/llms/bedrock/count_tokens/bedrock_token_counter.py b/litellm/llms/bedrock/count_tokens/bedrock_token_counter.py index 9c5211ed072..2d02b152c61 100644 --- a/litellm/llms/bedrock/count_tokens/bedrock_token_counter.py +++ b/litellm/llms/bedrock/count_tokens/bedrock_token_counter.py @@ -2,6 +2,7 @@ Bedrock Token Counter implementation using the CountTokens API. """ +from collections.abc import Mapping, Sequence from typing import Any, Final from litellm._logging import verbose_logger @@ -26,12 +27,12 @@ class BedrockTokenCounter(BaseTokenCounter): async def count_tokens( self, model_to_use: str, - messages: list[dict[str, Any]] | None, - contents: list[dict[str, Any]] | None, + messages: Sequence[Mapping[str, object]] | None, + contents: Sequence[Mapping[str, object]] | None, deployment: dict[str, Any] | None = None, request_model: str = "", - tools: list[dict[str, Any]] | None = None, - system: Any | None = None, + tools: Sequence[Mapping[str, object]] | None = None, + system: object | None = None, ) -> TokenCountResponse | None: """ Count tokens using AWS Bedrock's CountTokens API. @@ -56,7 +57,7 @@ class BedrockTokenCounter(BaseTokenCounter): litellm_params: Final = deployment.get("litellm_params", {}) # Build request data in the format expected by BedrockCountTokensHandler - request_data: Final[dict[str, Any]] = { + request_data: Final[dict[str, object]] = { "model": model_to_use, "messages": messages, } diff --git a/litellm/llms/bedrock/files/transformation.py b/litellm/llms/bedrock/files/transformation.py index 33b27943ad8..f50d63c9f6f 100644 --- a/litellm/llms/bedrock/files/transformation.py +++ b/litellm/llms/bedrock/files/transformation.py @@ -247,7 +247,7 @@ def _validate_file_id_against_configured_buckets( return validate_against(configured_bucket_names[-1]) -def _uploaded_object_size(litellm_params: Mapping[str, object], raw_response: Response) -> int: +def _uploaded_object_size(litellm_params: Mapping[str, object], response_headers: Mapping[str, str]) -> int: """ S3 answers PutObject with an empty body, so the stored object size comes from the signed request recorded by `transform_create_file_request`, not the response headers. @@ -255,7 +255,7 @@ def _uploaded_object_size(litellm_params: Mapping[str, object], raw_response: Re uploaded_size: Final = litellm_params.get(UPLOAD_CONTENT_LENGTH_PARAM) if isinstance(uploaded_size, int): return uploaded_size - response_content_length: Final = raw_response.headers.get("Content-Length", "0") + response_content_length: Final = response_headers.get("Content-Length", "0") return int(response_content_length) if response_content_length.isdigit() else 0 @@ -1161,7 +1161,7 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig): filename=filename, created_at=int(time.time()), # Current timestamp status="uploaded", - bytes=_uploaded_object_size(litellm_params=litellm_params, raw_response=raw_response), + bytes=_uploaded_object_size(litellm_params=litellm_params, response_headers=raw_response.headers), object="file", ) diff --git a/litellm/llms/bedrock/passthrough/guardrail_translation/handler.py b/litellm/llms/bedrock/passthrough/guardrail_translation/handler.py index b87f6196e51..eac8afd767c 100644 --- a/litellm/llms/bedrock/passthrough/guardrail_translation/handler.py +++ b/litellm/llms/bedrock/passthrough/guardrail_translation/handler.py @@ -14,6 +14,9 @@ from litellm.types.utils import GenericGuardrailAPIInputs if TYPE_CHECKING: from litellm.integrations.custom_guardrail import CustomGuardrail from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj + from litellm.llms.pass_through.guardrail_translation.handler import ( + PassThroughEndpointHandler, + ) from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.utils import ProxyLogging @@ -27,7 +30,7 @@ def _is_converse_endpoint(endpoint: str) -> bool: return bool(parts) and parts[-1] in _CONVERSE_ACTIONS -def _generic_passthrough_handler() -> BaseTranslation: +def _generic_passthrough_handler() -> "PassThroughEndpointHandler": """ Fallback for non-Converse Bedrock routes (e.g. invoke). The generic handler scans the full request/response payload so blocking guardrails diff --git a/litellm/llms/chatgpt/common_utils.py b/litellm/llms/chatgpt/common_utils.py index 35e32e4172f..fe33219f110 100644 --- a/litellm/llms/chatgpt/common_utils.py +++ b/litellm/llms/chatgpt/common_utils.py @@ -268,7 +268,7 @@ def _normalize_litellm_params(litellm_params: Any | None) -> dict: return {} -def get_chatgpt_session_id(litellm_params: Any | None) -> str | None: +def get_chatgpt_session_id(litellm_params: object) -> str | None: params: Final = _normalize_litellm_params(litellm_params) for key in ("litellm_session_id", "session_id"): value = params.get(key) @@ -286,5 +286,5 @@ def get_chatgpt_session_id(litellm_params: Any | None) -> str | None: return None -def ensure_chatgpt_session_id(litellm_params: Any | None) -> str: +def ensure_chatgpt_session_id(litellm_params: object) -> str: return get_chatgpt_session_id(litellm_params) or str(uuid4()) diff --git a/litellm/llms/github_copilot/responses/transformation.py b/litellm/llms/github_copilot/responses/transformation.py index 5a4bb798851..8b85b668cba 100644 --- a/litellm/llms/github_copilot/responses/transformation.py +++ b/litellm/llms/github_copilot/responses/transformation.py @@ -19,6 +19,7 @@ from litellm.llms.openai.responses.transformation import OpenAIResponsesAPIConfi from litellm.types.llms.openai import ( ResponseInputParam, ResponsesAPIOptionalRequestParams, + ResponsesAPIStreamingResponse, ) from litellm.types.router import GenericLiteLLMParams from litellm.types.utils import LlmProviders @@ -129,7 +130,7 @@ class GithubCopilotResponsesAPIConfig(OpenAIResponsesAPIConfig): model: str, parsed_chunk: dict, logging_obj: LiteLLMLoggingObj, - ) -> Any: + ) -> ResponsesAPIStreamingResponse: parsed_chunk = self._normalize_stream_item_id(parsed_chunk) return super().transform_streaming_response( model=model, @@ -262,7 +263,7 @@ class GithubCopilotResponsesAPIConfig(OpenAIResponsesAPIConfig): # Return the responses endpoint return f"{effective_api_base}/responses" - def _handle_reasoning_item(self, item: dict[str, Any]) -> dict[str, Any]: + def _handle_reasoning_item(self, item: dict[str, object]) -> dict[str, object]: """ Handle reasoning items for GitHub Copilot, preserving encrypted_content. @@ -280,7 +281,7 @@ class GithubCopilotResponsesAPIConfig(OpenAIResponsesAPIConfig): # Filter out None values for known problematic fields, # but preserve encrypted_content even if it exists - filtered_item: Final[dict[str, Any]] = {} + filtered_item: Final[dict[str, object]] = {} for k, v in item.items(): # Always include encrypted_content if present (even if None) if k == "encrypted_content": diff --git a/litellm/llms/hosted_vllm/chat/transformation.py b/litellm/llms/hosted_vllm/chat/transformation.py index 29dc485732f..92bc857e385 100644 --- a/litellm/llms/hosted_vllm/chat/transformation.py +++ b/litellm/llms/hosted_vllm/chat/transformation.py @@ -28,12 +28,12 @@ from ...openai.chat.gpt_transformation import OpenAIGPTConfig class HostedVLLMChatConfig(OpenAIGPTConfig): - def _convert_custom_tools_to_function_tools(self, tools: list[dict[str, Any]]) -> list[dict[str, Any]]: + def _convert_custom_tools_to_function_tools(self, tools: list[dict[str, Any]]) -> list[dict[str, object]]: """ vLLM chat completions currently accepts only OpenAI function tools. Convert custom tools into function tools so request validation does not fail. """ - converted_tools: Final[list[dict[str, Any]]] = [] + converted_tools: Final[list[dict[str, object]]] = [] for idx, tool in enumerate(tools): if not isinstance(tool, dict): converted_tools.append(tool) @@ -63,17 +63,14 @@ class HostedVLLMChatConfig(OpenAIGPTConfig): "required": ["input"], } - function_tool: dict[str, Any] = { - "type": "function", - "function": { - "name": str(tool_name), - "parameters": tool_parameters, - }, + function_definition: dict[str, object] = { + "name": str(tool_name), + "parameters": tool_parameters, } if isinstance(tool_description, str): - function_tool["function"]["description"] = tool_description + function_definition["description"] = tool_description - converted_tools.append(function_tool) + converted_tools.append({"type": "function", "function": function_definition}) return converted_tools @@ -148,7 +145,7 @@ class HostedVLLMChatConfig(OpenAIGPTConfig): @overload def _transform_messages( self, messages: list[AllMessageValues], model: str, is_async: Literal[True] - ) -> Coroutine[Any, Any, list[AllMessageValues]]: ... + ) -> Coroutine[object, object, list[AllMessageValues]]: ... @overload def _transform_messages( @@ -160,7 +157,7 @@ class HostedVLLMChatConfig(OpenAIGPTConfig): def _transform_messages( self, messages: list[AllMessageValues], model: str, is_async: bool = False - ) -> list[AllMessageValues] | Coroutine[Any, Any, list[AllMessageValues]]: + ) -> list[AllMessageValues] | Coroutine[object, object, list[AllMessageValues]]: """ Support translating: - video files from file_id or file_data to video_url diff --git a/litellm/llms/openai/common_utils.py b/litellm/llms/openai/common_utils.py index 2db6d78a218..e60b6145b44 100644 --- a/litellm/llms/openai/common_utils.py +++ b/litellm/llms/openai/common_utils.py @@ -10,7 +10,7 @@ import ssl import time import uuid from collections.abc import AsyncIterator, Iterator, Mapping -from typing import TYPE_CHECKING, Any, Final, Literal, NamedTuple, Optional +from typing import TYPE_CHECKING, Final, Literal, NamedTuple, Optional from urllib.parse import urlsplit import httpx @@ -87,8 +87,8 @@ class OpenAIError(BaseLLMException): ################################################################### def drop_params_from_unprocessable_entity_error( e: openai.UnprocessableEntityError | httpx.HTTPStatusError, - data: dict[str, Any], -) -> dict[str, Any]: + data: Mapping[str, object], +) -> dict[str, object]: """ Helper function to read OpenAI UnprocessableEntityError and drop the params that raised an error from the error message. diff --git a/litellm/llms/snowflake/chat/transformation.py b/litellm/llms/snowflake/chat/transformation.py index f65b0876202..aeff902f655 100644 --- a/litellm/llms/snowflake/chat/transformation.py +++ b/litellm/llms/snowflake/chat/transformation.py @@ -315,7 +315,7 @@ class SnowflakeConfig(SnowflakeBaseConfig, OpenAIGPTConfig): for msg in messages: if isinstance(msg, dict): role = msg.get("role", "") - content: Any = msg.get("content", "") + content: object = msg.get("content", "") msg_cache_control: object = msg.get("cache_control") else: role = getattr(msg, "role", "") @@ -463,7 +463,7 @@ class SnowflakeConfig(SnowflakeBaseConfig, OpenAIGPTConfig): return body - def _transform_tool_choice_to_anthropic(self, tool_choice: Any) -> dict[str, Any]: + def _transform_tool_choice_to_anthropic(self, tool_choice: object) -> Mapping[str, object]: """ Convert tool_choice from OpenAI format to Anthropic format. diff --git a/litellm/llms/stability/image_edit/transformations.py b/litellm/llms/stability/image_edit/transformations.py index 0b6052ad593..94711d21b50 100644 --- a/litellm/llms/stability/image_edit/transformations.py +++ b/litellm/llms/stability/image_edit/transformations.py @@ -74,7 +74,7 @@ class StabilityImageEditConfig(BaseImageEditConfig): } # Create a copy to not mutate original - convert TypedDict to regular dict - mapped_params: Final[dict[str, Any]] = dict(image_edit_optional_params) + mapped_params: Final[dict[str, object]] = dict(image_edit_optional_params) for k, v in image_edit_optional_params.items(): if k in param_mapping: @@ -182,7 +182,7 @@ class StabilityImageEditConfig(BaseImageEditConfig): # Build Stability request # Populate multipart form-data as separate text fields (data) and files. # Stability expects prompt/output_format/etc. as normal form fields, not file parts. - data: Final[dict[str, Any]] = { + data: Final[dict[str, object]] = { "output_format": "png", # Default to PNG } diff --git a/litellm/llms/triton/completion/transformation.py b/litellm/llms/triton/completion/transformation.py index 98a68ba2c36..3c868b3a96f 100644 --- a/litellm/llms/triton/completion/transformation.py +++ b/litellm/llms/triton/completion/transformation.py @@ -4,7 +4,7 @@ Translates from OpenAI's `/v1/chat/completions` endpoint to Triton's `/generate` import json from collections.abc import AsyncIterator, Iterator -from typing import TYPE_CHECKING, Any, Final, Literal +from typing import TYPE_CHECKING, Final, Literal from httpx import Headers, Response @@ -172,7 +172,7 @@ class TritonConfig(BaseConfig): streaming_response: Iterator[str] | AsyncIterator[str] | ModelResponse, sync_stream: bool, json_mode: bool | None = False, - ) -> Any: + ) -> "TritonResponseIterator": return TritonResponseIterator( streaming_response=streaming_response, sync_stream=sync_stream, @@ -195,14 +195,14 @@ class TritonGenerateConfig(TritonConfig): ) -> dict: inference_params: Final = optional_params.copy() stream: Final = inference_params.pop("stream", False) - data_for_triton: Final[dict[str, Any]] = { + data_for_triton: Final[dict[str, object]] = { "text_input": prompt_factory(model=model, messages=messages), "parameters": { "max_tokens": int(optional_params.get("max_tokens", DEFAULT_MAX_TOKENS_FOR_TRITON)), + **inference_params, }, "stream": bool(stream), } - data_for_triton["parameters"].update(inference_params) return data_for_triton def transform_response( diff --git a/litellm/llms/vertex_ai/rerank/transformation.py b/litellm/llms/vertex_ai/rerank/transformation.py index 2ec4f2da79b..2b1ed3a29a8 100644 --- a/litellm/llms/vertex_ai/rerank/transformation.py +++ b/litellm/llms/vertex_ai/rerank/transformation.py @@ -5,7 +5,7 @@ Why separate file? Make it easy to see how transformation works """ from collections.abc import Mapping -from typing import Any, Final +from typing import Final import httpx @@ -227,7 +227,7 @@ class VertexAIRerankConfig(BaseRerankConfig, VertexBase): model: str, drop_params: bool, query: str, - documents: list[str | dict[str, Any]], + documents: list[str | dict[str, object]], custom_llm_provider: str | None = None, top_n: int | None = None, rank_fields: list[str] | None = None, diff --git a/litellm/llms/volcengine/embedding/transformation.py b/litellm/llms/volcengine/embedding/transformation.py index 091b9dfd334..7c626c66e9d 100644 --- a/litellm/llms/volcengine/embedding/transformation.py +++ b/litellm/llms/volcengine/embedding/transformation.py @@ -3,7 +3,8 @@ Volcengine Embedding Transformation Transforms OpenAI embedding requests to Volcengine format """ -from typing import Any, Final +from collections.abc import Mapping +from typing import Final import httpx @@ -83,11 +84,11 @@ class VolcEngineEmbeddingConfig(BaseEmbeddingConfig): def map_openai_params( self, - non_default_params: dict[str, Any], - optional_params: dict[str, Any], + non_default_params: Mapping[str, object], + optional_params: dict[str, object], model: str, drop_params: bool, - ) -> dict[str, Any]: + ) -> dict[str, object]: """ Map OpenAI embedding parameters to Volcengine format. diff --git a/litellm/llms/watsonx/rerank/transformation.py b/litellm/llms/watsonx/rerank/transformation.py index 293880b188d..0c81d50a5fe 100644 --- a/litellm/llms/watsonx/rerank/transformation.py +++ b/litellm/llms/watsonx/rerank/transformation.py @@ -5,8 +5,8 @@ Docs - https://cloud.ibm.com/apidocs/watsonx-ai#text-rerank """ import uuid -from collections.abc import Mapping -from typing import Any, Final, cast +from collections.abc import Mapping, Sequence +from typing import Final, cast import httpx @@ -96,7 +96,7 @@ class IBMWatsonXRerankConfig(IBMWatsonXMixin, BaseRerankConfig): model: str, drop_params: bool, query: str, - documents: list[str | dict[str, Any]], + documents: Sequence[str | Mapping[str, object]], custom_llm_provider: str | None = None, top_n: int | None = None, rank_fields: list[str] | None = None, @@ -178,7 +178,7 @@ class IBMWatsonXRerankConfig(IBMWatsonXMixin, BaseRerankConfig): transformed_results: Final = [] for result in _results: - transformed_result: dict[str, Any] = { + transformed_result: dict[str, object] = { "index": result["index"], "relevance_score": result["score"], } diff --git a/litellm/llms/xai/realtime/transformation.py b/litellm/llms/xai/realtime/transformation.py index e9d16daad7c..5efe125ee60 100644 --- a/litellm/llms/xai/realtime/transformation.py +++ b/litellm/llms/xai/realtime/transformation.py @@ -16,7 +16,7 @@ construction time (see ``handler.py``) so all normalization is isolated here and ``RealTimeStreaming`` stays provider-agnostic. """ -from typing import Any, Final +from typing import Final class XAIRealtimeNormalizer: @@ -58,7 +58,7 @@ class XAIRealtimeNormalizer: # Cache content-part objects keyed by (response_id, item_id, content_index) # so that ``response.content_part.done`` events missing ``part`` can be # back-filled from earlier ``content_part.added`` / delta-done events. - self._content_part_by_key: dict[tuple, dict[str, Any]] = {} + self._content_part_by_key: dict[tuple, dict[str, object]] = {} # --------------------------------------------------------------------------- # Public interface consumed by RealTimeStreaming @@ -140,7 +140,7 @@ class XAIRealtimeNormalizer: } self._content_part_by_key[key] = updated - def _resolve_content_part(self, event: dict) -> dict[str, Any]: + def _resolve_content_part(self, event: dict) -> dict[str, object]: part: Final = event.get("part") if isinstance(part, dict): return part @@ -214,7 +214,7 @@ class XAIRealtimeNormalizer: needs_content: Final = event_type in self._EVENTS_NEEDING_CONTENT_INDEX if not needs_output and not needs_content: return event - patch: Final[dict[str, Any]] = {} + patch: Final[dict[str, object]] = {} if needs_output and "output_index" not in event: patch["output_index"] = 0 if needs_content and "content_index" not in event: @@ -228,8 +228,8 @@ class XAIRealtimeNormalizer: # --------------------------------------------------------------------------- @staticmethod - def _default_ga_usage() -> dict[str, Any]: - default_details: Final[dict[str, Any]] = { + def _default_ga_usage() -> dict[str, object]: + default_details: Final[dict[str, int]] = { "cached_tokens": 0, "text_tokens": 0, "audio_tokens": 0, @@ -243,7 +243,7 @@ class XAIRealtimeNormalizer: } @staticmethod - def _normalize_usage(usage: object, *, empty_as_null: bool) -> dict[str, Any] | None: + def _normalize_usage(usage: object, *, empty_as_null: bool) -> dict[str, object] | None: """Coerce a usage object into the full OpenAI GA shape. ``empty_as_null=True`` for ``response.created`` (usage optional). @@ -253,12 +253,12 @@ class XAIRealtimeNormalizer: return None if not usage: return None if empty_as_null else XAIRealtimeNormalizer._default_ga_usage() - default_details: Final[dict[str, Any]] = { + default_details: Final[dict[str, int]] = { "cached_tokens": 0, "text_tokens": 0, "audio_tokens": 0, } - normalized: Final[dict[str, Any]] = { + normalized: Final[dict[str, object]] = { "total_tokens": usage.get("total_tokens", 0), "input_tokens": usage.get("input_tokens", 0), "output_tokens": usage.get("output_tokens", 0), diff --git a/litellm/proxy/analytics_endpoints/cache_activity.py b/litellm/proxy/analytics_endpoints/cache_activity.py index b87b8eac3ef..952ef5bdcfc 100644 --- a/litellm/proxy/analytics_endpoints/cache_activity.py +++ b/litellm/proxy/analytics_endpoints/cache_activity.py @@ -2,16 +2,26 @@ import asyncio import json from collections.abc import Sequence from datetime import datetime -from typing import TYPE_CHECKING, Final +from typing import Final, Protocol from pydantic import BaseModel, TypeAdapter -if TYPE_CHECKING: - from litellm.proxy.utils import PrismaClient - UNKNOWN_CALL_TYPE: Final = "Unknown" +class _SupportsQueryRaw(Protocol): + """The single database operation the cache-activity queries issue.""" + + async def query_raw(self, query: str, *args: object) -> Sequence[object]: ... + + +class _SupportsRawQueryDb(Protocol): + """A prisma client handle, narrowed to the raw-query surface used here.""" + + @property + def db(self) -> _SupportsQueryRaw: ... + + class CacheActivityGroup(BaseModel): call_type: str api_requests: int @@ -143,7 +153,7 @@ def compute_totals(groups: Sequence[CacheActivityGroup]) -> CacheActivityTotals: async def get_cache_activity( - prisma_client: "PrismaClient", + prisma_client: _SupportsRawQueryDb, start_date: datetime, end_date: datetime, key_aliases: Sequence[str], diff --git a/litellm/proxy/client/cli/commands/model_groups.py b/litellm/proxy/client/cli/commands/model_groups.py index c904e5bed49..367c2063b6b 100644 --- a/litellm/proxy/client/cli/commands/model_groups.py +++ b/litellm/proxy/client/cli/commands/model_groups.py @@ -1,3 +1,4 @@ +from collections.abc import Mapping from typing import Final, Literal import click @@ -5,10 +6,17 @@ import rich import rich.table from ... import Client +from ._cli_context import cli_context_values def create_client(ctx: click.Context) -> Client: - return Client(base_url=ctx.obj["base_url"], api_key=ctx.obj["api_key"]) + context: Final = cli_context_values(ctx) + return Client(base_url=context["base_url"], api_key=context["api_key"]) + + +def _rendered_field(group: Mapping[str, object], key: str, default: str) -> str: + """The rendered value of one model group field, or ``default`` when the group omits it.""" + return str(group.get(key, default)) @click.group(name="model-groups") @@ -46,10 +54,10 @@ def list_model_groups(ctx: click.Context, output_format: Literal["table", "json" for group in groups: table.add_row( - str(group.get("model_group", "")), - str(group.get("mode", "chat")), - str(group.get("input_cost_per_token", "")), - str(group.get("output_cost_per_token", "")), + _rendered_field(group, "model_group", ""), + _rendered_field(group, "mode", "chat"), + _rendered_field(group, "input_cost_per_token", ""), + _rendered_field(group, "output_cost_per_token", ""), ) rich.print(table) diff --git a/litellm/proxy/client/cli/commands/up.py b/litellm/proxy/client/cli/commands/up.py index b7c02866d6f..45d21b0c0b5 100644 --- a/litellm/proxy/client/cli/commands/up.py +++ b/litellm/proxy/client/cli/commands/up.py @@ -166,7 +166,8 @@ def up(ctx: click.Context) -> None: is already running (this does not start one for you). Cursor is not supported: it has no equivalent file-based config to patch. """ - base_url: Final = ctx.obj["base_url"] + ctx_obj: Final[CliContextObj] = ctx.obj + base_url: Final = ctx_obj["base_url"] try: _ensure_fresh_login(ctx) diff --git a/litellm/proxy/common_utils/cache_pydantic_utils.py b/litellm/proxy/common_utils/cache_pydantic_utils.py index 725c2b61145..3703cf7c916 100644 --- a/litellm/proxy/common_utils/cache_pydantic_utils.py +++ b/litellm/proxy/common_utils/cache_pydantic_utils.py @@ -37,7 +37,7 @@ class CacheCodec: """ @staticmethod - def serialize(value: Any, model_type: type[T] | None = None) -> Any: + def serialize(value: object, model_type: type[T] | None = None) -> object: """ Encode a value for DualCache / Redis (``json.dumps``-safe). diff --git a/litellm/proxy/common_utils/proxy_rate_limit_error.py b/litellm/proxy/common_utils/proxy_rate_limit_error.py index c109da6f571..8d3a587a4cd 100644 --- a/litellm/proxy/common_utils/proxy_rate_limit_error.py +++ b/litellm/proxy/common_utils/proxy_rate_limit_error.py @@ -66,7 +66,7 @@ def map_v3_rate_limit_type( return None -def _coerce_message(detail: Any) -> str: +def _coerce_message(detail: object) -> str: """Best-effort, JSON-friendly stringification of an HTTPException-style detail.""" if detail is None: return "" @@ -144,7 +144,7 @@ class ProxyRateLimitError(HTTPException, RateLimitError): def __init__( self, detail: Any, - headers: Mapping[str, Any] | None = None, + headers: Mapping[str, object] | None = None, category: str | RateLimitErrorCategory = RateLimitErrorCategory.LITELLM_RATE_LIMIT, rate_limit_type: str | RateLimitType | None = None, model: str | None = None, diff --git a/litellm/proxy/container_endpoints/endpoints.py b/litellm/proxy/container_endpoints/endpoints.py index e852eb5d6f9..c1407979f29 100644 --- a/litellm/proxy/container_endpoints/endpoints.py +++ b/litellm/proxy/container_endpoints/endpoints.py @@ -106,7 +106,7 @@ async def create_container( # Process request using ProxyBaseLLMRequestProcessing processor: Final = ProxyBaseLLMRequestProcessing(data=data) try: - response: Final = await processor.base_process_llm_request( + response: Final[object] = await processor.base_process_llm_request( request=request, fastapi_response=fastapi_response, user_api_key_dict=user_api_key_dict, @@ -216,7 +216,7 @@ async def list_containers( or get_custom_llm_provider_from_request_query(request=request) or "openai" ) - data: Final[dict[str, Any]] = { + data: Final[dict[str, object]] = { "query_params": query_params, "model": query_params.get("model"), "order": order, @@ -341,7 +341,7 @@ async def retrieve_container( # Process request using ProxyBaseLLMRequestProcessing processor: Final = ProxyBaseLLMRequestProcessing(data=data) try: - return await processor.base_process_llm_request( + container: Final[object] = await processor.base_process_llm_request( request=request, fastapi_response=fastapi_response, user_api_key_dict=user_api_key_dict, @@ -366,6 +366,7 @@ async def retrieve_container( proxy_logging_obj=proxy_logging_obj, version=version, ) + return container @router.delete( @@ -446,7 +447,7 @@ async def delete_container( # Process request using ProxyBaseLLMRequestProcessing processor: Final = ProxyBaseLLMRequestProcessing(data=data) try: - return await processor.base_process_llm_request( + deleted_container: Final[object] = await processor.base_process_llm_request( request=request, fastapi_response=fastapi_response, user_api_key_dict=user_api_key_dict, @@ -471,6 +472,7 @@ async def delete_container( proxy_logging_obj=proxy_logging_obj, version=version, ) + return deleted_container # Register JSON-configured container file endpoints diff --git a/litellm/proxy/db/exception_handler.py b/litellm/proxy/db/exception_handler.py index f469587ab8e..de20416ad9c 100644 --- a/litellm/proxy/db/exception_handler.py +++ b/litellm/proxy/db/exception_handler.py @@ -1,5 +1,5 @@ from collections.abc import Awaitable, Callable, Iterator -from typing import Any, Final, TypeVar +from typing import Final, Protocol, TypeVar from litellm._logging import verbose_proxy_logger from litellm.proxy._types import ( @@ -407,8 +407,20 @@ def _coerce_timeout(value: object, fallback: float) -> float: _ReadResultT: Final = TypeVar("_ReadResultT") +class _DBReconnectClient(Protocol): + """The one method `call_with_db_reconnect_retry` needs from a Prisma client.""" + + async def attempt_db_reconnect( + self, + *, + reason: str, + timeout_seconds: float | None = None, + lock_timeout_seconds: float | None = None, + ) -> bool: ... + + async def call_with_db_reconnect_retry( - prisma_client: Any, + prisma_client: _DBReconnectClient, coro_factory: Callable[[], Awaitable[_ReadResultT]], *, reason: str, diff --git a/litellm/proxy/db/spend_log_tool_index.py b/litellm/proxy/db/spend_log_tool_index.py index 93bbc567430..7a7982481fb 100644 --- a/litellm/proxy/db/spend_log_tool_index.py +++ b/litellm/proxy/db/spend_log_tool_index.py @@ -148,4 +148,4 @@ async def flush_tool_usage_transactions( except DB_RETRY_SAFE_ERROR_TYPES: if attempt >= n_retry_times: raise - await asyncio.sleep(2**attempt + random.uniform(0, 1)) + await asyncio.sleep(2.0**attempt + random.uniform(0, 1)) diff --git a/litellm/proxy/guardrails/guardrail_hooks/crowdstrike_aidr/crowdstrike_aidr.py b/litellm/proxy/guardrails/guardrail_hooks/crowdstrike_aidr/crowdstrike_aidr.py index f7f500b1adc..db6fc238b9b 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/crowdstrike_aidr/crowdstrike_aidr.py +++ b/litellm/proxy/guardrails/guardrail_hooks/crowdstrike_aidr/crowdstrike_aidr.py @@ -6,7 +6,7 @@ from typing import TYPE_CHECKING, Annotated, Final, Literal, NamedTuple, Optiona from fastapi import HTTPException from pydantic import BaseModel, ConfigDict, Field, ValidationError -from typing_extensions import Any, override +from typing_extensions import override from litellm._logging import verbose_proxy_logger from litellm.integrations.custom_guardrail import ( @@ -78,7 +78,7 @@ class _GuardChatCompletionsResult(BaseModel): """Whether or not the prompt triggered a block detection.""" transformed: bool | None = None """Whether or not the original input was transformed.""" - detectors: dict[str, Any] | None = None + detectors: dict[str, object] | None = None """Result of the policy analyzing and input prompt.""" @@ -146,8 +146,8 @@ def _extract_text_from_message(message: _Message) -> str: return "\n".join(part.text for part in content if isinstance(part, _TextContentPart)) -def _merge_metadata_bags(request_data: Mapping[str, Any]) -> Mapping[str, Any] | None: - merged: Final[dict[str, Any]] = {} +def _merge_metadata_bags(request_data: Mapping[str, object]) -> Mapping[str, object] | None: + merged: Final[dict[str, object]] = {} present = False for bag in (request_data.get("metadata"), request_data.get("litellm_metadata")): if isinstance(bag, Mapping): @@ -313,7 +313,7 @@ class CrowdStrikeAIDRHandler(CustomGuardrail): self._set_streaming_params(streaming_params_from_litellm_params(litellm_params)) async def _call_crowdstrike_aidr_guard( - self, payload: dict[str, Any], hook_name: str + self, payload: dict[str, object], hook_name: str ) -> _GuardChatCompletionsResult: """ Makes the API call to the CrowdStrike AIDR AI Guard endpoint. @@ -423,7 +423,7 @@ class CrowdStrikeAIDRHandler(CustomGuardrail): return [_extract_text_from_message(msg) for msg in tail] async def _call_or_fail_open( - self, payload: dict[str, Any], hook_name: str, request_data: dict[str, object] + self, payload: dict[str, object], hook_name: str, request_data: dict[str, object] ) -> _GuardChatCompletionsResult: start_time: Final = time.time() try: @@ -506,7 +506,7 @@ class CrowdStrikeAIDRHandler(CustomGuardrail): event_type = "output" hook_name = "apply_guardrail (response)" - ai_guard_payload: Final[dict[str, Any]] = { + ai_guard_payload: Final[dict[str, object]] = { "guard_input": guard_input.model_dump(mode="json"), "event_type": event_type, } @@ -521,7 +521,7 @@ class CrowdStrikeAIDRHandler(CustomGuardrail): if user_id: ai_guard_payload["user_id"] = user_id - extra_info: Final[dict[str, str]] = {} + extra_info: Final[dict[str, object]] = {} user_email: Final = metadata.get("user_api_key_user_email") if user_email: extra_info["user_name"] = user_email diff --git a/litellm/proxy/guardrails/guardrail_hooks/enkryptai/enkryptai.py b/litellm/proxy/guardrails/guardrail_hooks/enkryptai/enkryptai.py index 89afecafb0f..efe959bd186 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/enkryptai/enkryptai.py +++ b/litellm/proxy/guardrails/guardrail_hooks/enkryptai/enkryptai.py @@ -6,7 +6,7 @@ # +-------------------------------------------------------------+ import os -from collections.abc import AsyncGenerator +from collections.abc import AsyncGenerator, AsyncIterable from datetime import datetime from typing import TYPE_CHECKING, Any, Final, Literal, Optional @@ -465,7 +465,7 @@ class EnkryptAIGuardrails(CustomGuardrail): async def async_post_call_streaming_iterator_hook( self, user_api_key_dict: UserAPIKeyAuth, - response: Any, + response: AsyncIterable[ModelResponseStream], request_data: dict, ) -> AsyncGenerator[ModelResponseStream, None]: """ diff --git a/litellm/proxy/guardrails/guardrail_hooks/qualifire/qualifire.py b/litellm/proxy/guardrails/guardrail_hooks/qualifire/qualifire.py index d82944c44ed..eceb54681f6 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/qualifire/qualifire.py +++ b/litellm/proxy/guardrails/guardrail_hooks/qualifire/qualifire.py @@ -383,7 +383,7 @@ class QualifireGuardrail(CustomGuardrail): result: Final = response.json() # Extract response info for logging - qualifire_response: Final = { + qualifire_response: Final[dict[str, object]] = { "score": result.get("score"), "status": result.get("status"), } diff --git a/litellm/proxy/guardrails/guardrail_hooks/semantic_guard/route_loader.py b/litellm/proxy/guardrails/guardrail_hooks/semantic_guard/route_loader.py index 8d5923d1302..b2139779925 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/semantic_guard/route_loader.py +++ b/litellm/proxy/guardrails/guardrail_hooks/semantic_guard/route_loader.py @@ -6,6 +6,7 @@ then builds a SemanticRouter for prompt matching. """ import os +from collections.abc import Mapping, Sequence from typing import TYPE_CHECKING, Any, Final import yaml @@ -66,7 +67,7 @@ class SemanticGuardRouteLoader: cls, route_templates: list[str] | None, custom_routes_file: str | None, - custom_routes: list[dict[str, Any]] | None, + custom_routes: Sequence[Mapping[str, object]] | None, global_threshold: float = DEFAULT_SEMANTIC_GUARD_SIMILARITY_THRESHOLD, ) -> list["Route"]: """Build semantic-router Route objects from templates + custom config.""" diff --git a/litellm/proxy/hooks/parallel_request_limiter.py b/litellm/proxy/hooks/parallel_request_limiter.py index b313cb64c3f..d41acadc4dd 100644 --- a/litellm/proxy/hooks/parallel_request_limiter.py +++ b/litellm/proxy/hooks/parallel_request_limiter.py @@ -27,7 +27,7 @@ if TYPE_CHECKING: from litellm.proxy.utils import InternalUsageCache as _InternalUsageCache - Span = _Span | Any + Span = _Span InternalUsageCache = _InternalUsageCache else: Span = Any @@ -75,7 +75,7 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger): current: dict | None, request_count_api_key: str, rate_limit_type: Literal["key", "model_per_key", "user", "customer", "team"], - values_to_update_in_cache: list[tuple[Any, Any]], + values_to_update_in_cache: list[tuple[str, object]], ) -> dict: verbose_proxy_logger.info("Current Usage of %s in this minute: %s", rate_limit_type, current) if current is None: @@ -266,7 +266,7 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger): rpm_limit = sys.maxsize values_to_update_in_cache: list[ - tuple[Any, Any] + tuple[str, object] ] = [] # values that need to get updated in cache, will run a batch_set_cache after this function # ------------ diff --git a/litellm/proxy/management_endpoints/common_daily_activity.py b/litellm/proxy/management_endpoints/common_daily_activity.py index ce6a97708ab..6841cadf972 100644 --- a/litellm/proxy/management_endpoints/common_daily_activity.py +++ b/litellm/proxy/management_endpoints/common_daily_activity.py @@ -17,6 +17,7 @@ from litellm.proxy.spend_tracking.key_metadata_recovery import ( ) from litellm.proxy.spend_tracking.ptu_feature_flag import is_ptu_cost_attribution_enabled from litellm.proxy.utils import PrismaClient +from litellm.repositories.prisma_protocols import TableActions from litellm.repositories.table_repositories import DeletedVerificationTokenRepository from litellm.repositories.verification_token_repository import ( VerificationTokenRepository, @@ -1155,8 +1156,10 @@ async def get_daily_activity( include_current_utc_day=include_current_utc_day, ) + spend_table: Final[TableActions[DailySpendRecord]] = getattr(prisma_client.db, table_name) + # Get total count for pagination - total_count: Final[int] = await getattr(prisma_client.db, table_name).count(where=where_conditions) + total_count: Final[int] = await spend_table.count(where=where_conditions) # Fetch paginated results. # ``date`` alone is not a unique sort key -- a busy tenant has many @@ -1168,7 +1171,7 @@ async def get_daily_activity( # total. Adding ``id`` (the row's UUID primary key, present on both # LiteLLM_DailyUserSpend and LiteLLM_DailyTeamSpend) as a tiebreaker # gives every page a stable cursor (#30164). - daily_spend_data: Final[Sequence[DailySpendRecord]] = await getattr(prisma_client.db, table_name).find_many( + daily_spend_data: Final[Sequence[DailySpendRecord]] = await spend_table.find_many( where=where_conditions, order=[ {"date": "desc"}, diff --git a/litellm/proxy/ocr_endpoints/endpoints.py b/litellm/proxy/ocr_endpoints/endpoints.py index ebf4d988fdd..1d0108cc2ba 100644 --- a/litellm/proxy/ocr_endpoints/endpoints.py +++ b/litellm/proxy/ocr_endpoints/endpoints.py @@ -2,7 +2,7 @@ import json from collections.abc import Mapping -from typing import Any, Final, cast +from typing import Final, cast import orjson from fastapi import APIRouter, Depends, HTTPException, Request, Response, UploadFile @@ -48,7 +48,7 @@ def _build_document_from_upload( ) -def _with_request_format(data: Mapping[str, Any], request: Request) -> Mapping[str, Any]: +def _with_request_format(data: Mapping[str, object], request: Request) -> Mapping[str, object]: """ Resolve the requested response format from the body or the `x-req-format` header. @@ -90,7 +90,7 @@ def _native_response(response: object, fastapi_response: Response) -> Response | ) -async def _parse_multipart_form(request: Request) -> dict[str, Any]: +async def _parse_multipart_form(request: Request) -> dict[str, object]: """ Extract OCR data from a multipart form request. @@ -130,7 +130,7 @@ async def _parse_multipart_form(request: Request) -> dict[str, Any]: content_type=uploaded_file.content_type, ) - data: Final[dict[str, Any]] = {"document": document} + data: Final[dict[str, object]] = {"document": document} for field_name, field_value in form.items(): if field_name in ("file", "document"): @@ -154,12 +154,12 @@ async def _parse_multipart_form(request: Request) -> dict[str, Any]: return data -async def _parse_ocr_request(request: Request) -> Mapping[str, Any]: +async def _parse_ocr_request(request: Request) -> Mapping[str, object]: """Parse an OCR request and apply the `x-req-format` header, if any.""" return _with_request_format(await _parse_ocr_request_body(request), request) -async def _parse_ocr_request_body(request: Request) -> dict[str, Any]: +async def _parse_ocr_request_body(request: Request) -> dict[str, object]: """ Parse an OCR request, supporting both JSON and multipart form data. @@ -320,7 +320,7 @@ async def ocr( # Process request using ProxyBaseLLMRequestProcessing processor = ProxyBaseLLMRequestProcessing(data=data) - response: Final = await processor.base_process_llm_request( + response: Final[object] = await processor.base_process_llm_request( request=request, fastapi_response=fastapi_response, user_api_key_dict=user_api_key_dict, diff --git a/litellm/repositories/budget_repository.py b/litellm/repositories/budget_repository.py index 62632ffb5f6..205646c8393 100644 --- a/litellm/repositories/budget_repository.py +++ b/litellm/repositories/budget_repository.py @@ -2,7 +2,8 @@ Budget repository for database operations on LiteLLM_BudgetTable. """ -from typing import TYPE_CHECKING, Any, Final +from collections.abc import Mapping +from typing import TYPE_CHECKING, Final, Protocol from litellm.models.budget import LiteLLM_BudgetTable from litellm.repositories.base_repository import BaseRepository @@ -12,12 +13,27 @@ if TYPE_CHECKING: from prisma import models as prisma_models +class _BudgetDb(Protocol): + """The single Prisma table this repository reaches for on ``prisma_client.db``.""" + + @property + def litellm_budgettable(self) -> TableActions["prisma_models.LiteLLM_BudgetTable"]: ... + + +class _PrismaClientView(Protocol): + """The one attribute this repository reads off the untyped Prisma client wrapper.""" + + @property + def db(self) -> _BudgetDb: ... + + class BudgetRepository(BaseRepository[LiteLLM_BudgetTable]): """Repository for budget database operations.""" @property def table(self) -> TableActions["prisma_models.LiteLLM_BudgetTable"]: - return self.prisma_client.db.litellm_budgettable + client: Final[_PrismaClientView] = self.prisma_client + return client.db.litellm_budgettable @property def model_class(self) -> type[LiteLLM_BudgetTable]: @@ -34,12 +50,12 @@ class BudgetRepository(BaseRepository[LiteLLM_BudgetTable]): max_parallel_requests: int | None = None, tpm_limit: int | None = None, rpm_limit: int | None = None, - model_max_budget: dict[str, Any] | None = None, + model_max_budget: Mapping[str, object] | None = None, budget_duration: str | None = None, allowed_models: list[str] | None = None, ) -> LiteLLM_BudgetTable: """Create a new budget record.""" - data: Final[dict[str, Any]] = { + data: Final[dict[str, object]] = { "created_by": created_by, "updated_by": created_by, } @@ -71,12 +87,12 @@ class BudgetRepository(BaseRepository[LiteLLM_BudgetTable]): max_parallel_requests: int | None = None, tpm_limit: int | None = None, rpm_limit: int | None = None, - model_max_budget: dict[str, Any] | None = None, + model_max_budget: Mapping[str, object] | None = None, budget_duration: str | None = None, allowed_models: list[str] | None = None, ) -> LiteLLM_BudgetTable | None: """Update an existing budget record.""" - data: Final[dict[str, Any]] = {"updated_by": updated_by} + data: Final[dict[str, object]] = {"updated_by": updated_by} if max_budget is not None: data["max_budget"] = max_budget if soft_budget is not None: diff --git a/litellm/repositories/organization_repository.py b/litellm/repositories/organization_repository.py index 5a9bd3724e0..47eb8f4a609 100644 --- a/litellm/repositories/organization_repository.py +++ b/litellm/repositories/organization_repository.py @@ -2,7 +2,8 @@ Organization repository for database operations on LiteLLM_OrganizationTable. """ -from typing import TYPE_CHECKING, Any, Final +from collections.abc import Mapping +from typing import TYPE_CHECKING, Final, Protocol from litellm.models.organization import LiteLLM_OrganizationTable from litellm.repositories.base_repository import BaseRepository @@ -12,12 +13,27 @@ if TYPE_CHECKING: from prisma import models as prisma_models +class _OrganizationDb(Protocol): + """The single Prisma table this repository reaches for on ``prisma_client.db``.""" + + @property + def litellm_organizationtable(self) -> TableActions["prisma_models.LiteLLM_OrganizationTable"]: ... + + +class _PrismaClientView(Protocol): + """The one attribute this repository reads off the untyped Prisma client wrapper.""" + + @property + def db(self) -> _OrganizationDb: ... + + class OrganizationRepository(BaseRepository[LiteLLM_OrganizationTable]): """Repository for organization database operations.""" @property def table(self) -> TableActions["prisma_models.LiteLLM_OrganizationTable"]: - return self.prisma_client.db.litellm_organizationtable + client: Final[_PrismaClientView] = self.prisma_client + return client.db.litellm_organizationtable @property def model_class(self) -> type[LiteLLM_OrganizationTable]: @@ -39,12 +55,12 @@ class OrganizationRepository(BaseRepository[LiteLLM_OrganizationTable]): budget_id: str, created_by: str, organization_id: str | None = None, - metadata: dict[str, Any] | None = None, + metadata: Mapping[str, object] | None = None, models: list[str] | None = None, object_permission_id: str | None = None, ) -> LiteLLM_OrganizationTable: """Create a new organization.""" - data: Final[dict[str, Any]] = { + data: Final[dict[str, object]] = { "organization_alias": organization_alias, "budget_id": budget_id, "created_by": created_by, @@ -67,12 +83,12 @@ class OrganizationRepository(BaseRepository[LiteLLM_OrganizationTable]): updated_by: str, organization_alias: str | None = None, budget_id: str | None = None, - metadata: dict[str, Any] | None = None, + metadata: Mapping[str, object] | None = None, models: list[str] | None = None, object_permission_id: str | None = None, ) -> LiteLLM_OrganizationTable | None: """Update an organization.""" - data: Final[dict[str, Any]] = {"updated_by": updated_by} + data: Final[dict[str, object]] = {"updated_by": updated_by} if organization_alias is not None: data["organization_alias"] = organization_alias if budget_id is not None: diff --git a/litellm/repositories/project_repository.py b/litellm/repositories/project_repository.py index 48e55efd258..905e813f35e 100644 --- a/litellm/repositories/project_repository.py +++ b/litellm/repositories/project_repository.py @@ -2,7 +2,8 @@ Project repository for database operations on LiteLLM_ProjectTable. """ -from typing import TYPE_CHECKING, Any, Final +from collections.abc import Mapping +from typing import TYPE_CHECKING, Final from litellm.models.project import LiteLLM_ProjectTable from litellm.repositories.base_repository import BaseRepository @@ -43,14 +44,14 @@ class ProjectRepository(BaseRepository[LiteLLM_ProjectTable]): description: str | None = None, team_id: str | None = None, budget_id: str | None = None, - metadata: dict[str, Any] | None = None, + metadata: Mapping[str, object] | None = None, models: list[str] | None = None, model_rpm_limit: dict[str, int] | None = None, model_tpm_limit: dict[str, int] | None = None, object_permission_id: str | None = None, ) -> LiteLLM_ProjectTable: """Create a new project.""" - data: Final[dict[str, Any]] = { + data: Final[dict[str, object]] = { "created_by": created_by, "updated_by": created_by, } @@ -85,7 +86,7 @@ class ProjectRepository(BaseRepository[LiteLLM_ProjectTable]): description: str | None = None, team_id: str | None = None, budget_id: str | None = None, - metadata: dict[str, Any] | None = None, + metadata: Mapping[str, object] | None = None, models: list[str] | None = None, model_rpm_limit: dict[str, int] | None = None, model_tpm_limit: dict[str, int] | None = None, @@ -93,7 +94,7 @@ class ProjectRepository(BaseRepository[LiteLLM_ProjectTable]): object_permission_id: str | None = None, ) -> LiteLLM_ProjectTable | None: """Update a project.""" - data: Final[dict[str, Any]] = {"updated_by": updated_by} + data: Final[dict[str, object]] = {"updated_by": updated_by} if project_alias is not None: data["project_alias"] = project_alias if description is not None: diff --git a/litellm/router_strategy/adaptive_router/adaptive_router.py b/litellm/router_strategy/adaptive_router/adaptive_router.py index 12ccacbbc1d..8727c3a69f0 100644 --- a/litellm/router_strategy/adaptive_router/adaptive_router.py +++ b/litellm/router_strategy/adaptive_router/adaptive_router.py @@ -16,6 +16,7 @@ from __future__ import annotations import asyncio import time from collections import OrderedDict +from collections.abc import Mapping from dataclasses import asdict, dataclass from typing import Any, Final, cast @@ -122,7 +123,7 @@ class AdaptiveRouter: prefs = self.model_to_prefs.get(model) or _default_prefs() self._cells[(rt, model)] = initial_cell(prefs, rt) - async def load_state_from_db(self, prisma_client: Any) -> None: + async def load_state_from_db(self, prisma_client: object) -> None: """Add each row's persisted delta to a freshly computed cold-start prior. A row holds an accumulated delta, not a full posterior, and can be one-sided @@ -237,7 +238,7 @@ class AdaptiveRouter: cost_weight=self.config.weights.cost, ) - async def get_state_snapshot(self) -> dict[str, Any]: + async def get_state_snapshot(self) -> dict[str, object]: """In-memory snapshot for the introspection endpoint. Cheap; no DB hit.""" cells: Final = [] for (rt, model), cell in sorted(self._cells.items(), key=lambda kv: (kv[0][0].value, kv[0][1])): @@ -278,7 +279,7 @@ class AdaptiveRouter: @staticmethod def _extract_min_quality_tier( - request_kwargs: dict[str, Any], + request_kwargs: Mapping[str, object], ) -> int | None: """Pull `min_quality_tier` from request headers or metadata. @@ -484,7 +485,7 @@ class AdaptiveRouter: return combined_delta @staticmethod - def _persistable_session_snapshot(state: SessionState) -> dict[str, Any]: + def _persistable_session_snapshot(state: SessionState) -> dict[str, object]: snapshot: Final = asdict(state) for sensitive in ( "last_user_content", diff --git a/litellm/router_strategy/adaptive_router/signals.py b/litellm/router_strategy/adaptive_router/signals.py index c28613b54eb..72e8d27d2bf 100644 --- a/litellm/router_strategy/adaptive_router/signals.py +++ b/litellm/router_strategy/adaptive_router/signals.py @@ -92,7 +92,7 @@ class Turn: user_content: str | None = None assistant_content: str | None = None - tool_calls: list[dict[str, Any]] = field(default_factory=list) + tool_calls: Sequence[Mapping[str, object]] = field(default_factory=list[Mapping[str, object]]) tool_results: Sequence[Mapping[str, object]] = field(default_factory=list) response_status: int | None = None @@ -174,7 +174,7 @@ def _detect_failure(tool_results: Sequence[Mapping[str, object]]) -> bool: return False -def _signature(call: dict[str, Any]) -> str: +def _signature(call: Mapping[str, Any]) -> str: """Stable signature for loop detection: name + sorted JSON-ish args.""" name: Final = call.get("name") or call.get("function", {}).get("name", "") call_args = call.get("arguments") @@ -185,7 +185,7 @@ def _signature(call: dict[str, Any]) -> str: return f"{name}({call_args})" -def _detect_loop(history: list[str], new_calls: list[dict[str, Any]]) -> bool: +def _detect_loop(history: list[str], new_calls: Sequence[Mapping[str, object]]) -> bool: """Fires if any new call's signature appears >= LOOP_REPEAT_THRESHOLD-1 times in recent history (so this call would be the Nth).""" if not new_calls: @@ -238,7 +238,7 @@ def detect_response_signals( previous_assistant_content: str | None, current_assistant_content: str | None, tool_call_history: list[str], - tool_calls: list[dict[str, Any]], + tool_calls: Sequence[Mapping[str, object]], tool_results: Sequence[Mapping[str, object]], response_status: int | None, ) -> SignalDelta: diff --git a/litellm/router_strategy/adaptive_router/update_queue.py b/litellm/router_strategy/adaptive_router/update_queue.py index 1b9fce284ac..e28f2379f9c 100644 --- a/litellm/router_strategy/adaptive_router/update_queue.py +++ b/litellm/router_strategy/adaptive_router/update_queue.py @@ -19,7 +19,8 @@ to the in-memory aggregator). Flush is async and batched. from __future__ import annotations import asyncio -from typing import Any, Final +from collections.abc import Mapping +from typing import Final from litellm._logging import verbose_router_logger from litellm.repositories.table_repositories import ( @@ -39,7 +40,7 @@ class AdaptiveRouterUpdateQueue: def __init__(self) -> None: self._state_agg: dict[StateKey, dict[str, float]] = {} - self._session_agg: dict[SessionKey, dict[str, Any]] = {} + self._session_agg: dict[SessionKey, Mapping[str, object]] = {} self._lock = asyncio.Lock() self._max_state_size_seen = 0 self._max_session_size_seen = 0 @@ -77,7 +78,7 @@ class AdaptiveRouterUpdateQueue: session_id: str, router_name: str, model_name: str, - state_dict: dict[str, Any], + state_dict: Mapping[str, object], ) -> None: """ Last-write-wins per session row. The state_dict is a snapshot of the @@ -91,7 +92,7 @@ class AdaptiveRouterUpdateQueue: # ---- Flushers (called by background task) ---------------------------- - async def flush_state_to_db(self, prisma_client: Any) -> int: + async def flush_state_to_db(self, prisma_client: object) -> int: """ Drain state aggregator and apply to LiteLLM_AdaptiveRouterState. Returns number of cells flushed. @@ -147,7 +148,7 @@ class AdaptiveRouterUpdateQueue: return len(batch) - async def flush_session_to_db(self, prisma_client: Any) -> int: + async def flush_session_to_db(self, prisma_client: object) -> int: """ Drain session aggregator and upsert into LiteLLM_AdaptiveRouterSession. Returns number of session rows flushed. diff --git a/litellm/types/containers/main.py b/litellm/types/containers/main.py index 6a339fd2eac..62ef524a435 100644 --- a/litellm/types/containers/main.py +++ b/litellm/types/containers/main.py @@ -140,7 +140,7 @@ class ContainerFileObject(BaseModel): created_at: int path: str source: str - _hidden_params: dict[str, Any] = {} + _hidden_params: dict[str, builtins.object] = {} def __contains__(self, key: str) -> bool: return hasattr(self, key) diff --git a/litellm/types/guardrails.py b/litellm/types/guardrails.py index 02dee40f2a3..175226f8d4b 100644 --- a/litellm/types/guardrails.py +++ b/litellm/types/guardrails.py @@ -2,10 +2,10 @@ from collections.abc import Mapping from datetime import datetime from enum import Enum from types import MappingProxyType -from typing import Any, Final, Literal +from typing import Final, Literal from pydantic import BaseModel, ConfigDict, Field, field_validator, model_validator -from typing_extensions import Required, TypedDict +from typing_extensions import ReadOnly, Required, TypedDict from litellm.constants import BEDROCK_APPLY_GUARDRAIL_CHUNK_BUDGET_CHARS from litellm.types.proxy.guardrails.guardrail_hooks.akto import ( @@ -935,7 +935,7 @@ class BaseLitellmParams(ContentFilterConfigModel): # works for new and patch up ), ) - additional_provider_specific_params: dict[str, Any] | None = Field( + additional_provider_specific_params: dict[str, object] | None = Field( default=None, description="Additional provider-specific parameters for generic guardrail APIs", ) @@ -1157,7 +1157,7 @@ class GuardrailEventHooks(str, Enum): class DynamicGuardrailParams(TypedDict): - extra_body: dict[str, Any] + extra_body: ReadOnly[dict[str, object]] class GUARDRAIL_DEFINITION_LOCATION(str, Enum): @@ -1188,7 +1188,7 @@ class GuardrailUIAddGuardrailSettings(BaseModel): supported_modes: list[str] supported_modes_by_provider: dict[str, list[str]] pii_entity_categories: list[PiiEntityCategoryMap] - content_filter_settings: dict[str, Any] | None = None + content_filter_settings: dict[str, object] | None = None class PresidioPerRequestConfig(BaseModel): @@ -1206,8 +1206,8 @@ class ApplyGuardrailRequest(BaseModel): language: str | None = None entities: list[PiiEntityType] | None = None input_type: str = "request" - messages: list[dict[str, Any]] | None = None - metadata: dict[str, Any] | None = None + messages: list[dict[str, object]] | None = None + metadata: dict[str, object] | None = None class ApplyGuardrailResponse(BaseModel): @@ -1217,4 +1217,4 @@ class ApplyGuardrailResponse(BaseModel): class PatchGuardrailRequest(BaseModel): guardrail_name: str | None = None litellm_params: BaseLitellmParams | None = None - guardrail_info: dict[str, Any] | None = None + guardrail_info: dict[str, object] | None = None diff --git a/litellm/types/integrations/prometheus.py b/litellm/types/integrations/prometheus.py index 8498b6f6d00..323f43f4b01 100644 --- a/litellm/types/integrations/prometheus.py +++ b/litellm/types/integrations/prometheus.py @@ -44,7 +44,7 @@ def _sanitize_prometheus_label_name(label: str) -> str: _PROMETHEUS_LABEL_VALUE_TRANSLATE_V1: Final = str.maketrans("\n", " ", "\r\u2028\u2029") -def _sanitize_prometheus_label_value(value: Any | None) -> str | None: +def _sanitize_prometheus_label_value(value: object | None) -> str | None: """ Same semantics as :func:`_sanitize_prometheus_label_value`, implemented with ``str.translate`` plus a single escape pass instead of chained ``replace``. @@ -1023,7 +1023,7 @@ class UserAPIKeyLabelValues: ``hashed_api_key``. This supports ``**standard_logging_payload`` in tests. """ field_names: Final = {f.name for f in fields(self)} - merged: Final[dict[str, Any]] = {} + merged: Final[dict[str, object]] = {} for f in fields(self): if f.default_factory is not MISSING: merged[f.name] = f.default_factory() @@ -1060,9 +1060,9 @@ class UserAPIKeyLabelValues: # stays cheap. (Dataclass default `str()` delegates to `__repr__`.) return "" - def model_dump(self) -> dict[str, Any]: + def model_dump(self) -> dict[str, object]: """Same shape as the former Pydantic ``model_dump()`` (plain dict, list tags).""" - d: Final[dict[str, Any]] = {f.name: getattr(self, f.name) for f in fields(self)} + d: Final[dict[str, object]] = {f.name: getattr(self, f.name) for f in fields(self)} d["tags"] = list(self.tags) d["custom_metadata_labels"] = dict(self.custom_metadata_labels) return d diff --git a/litellm/types/realtime.py b/litellm/types/realtime.py index 17dc70126f3..bda3865c46a 100644 --- a/litellm/types/realtime.py +++ b/litellm/types/realtime.py @@ -78,15 +78,15 @@ class RealtimeSessionConfig(BaseModel): type: str | None = None model: str | None = None instructions: str | None = None - audio: dict[str, Any] | None = None + audio: dict[str, object] | None = None include: list[str] | None = None max_output_tokens: int | str | None = None output_modalities: list[str] | None = None - tool_choice: Any | None = None - tools: list[dict[str, Any]] | None = None - tracing: Any | None = None - truncation: Any | None = None - prompt: dict[str, Any] | None = None + tool_choice: object | None = None + tools: list[dict[str, object]] | None = None + tracing: object | None = None + truncation: object | None = None + prompt: dict[str, object] | None = None class RealtimeClientSecretRequest(BaseModel): @@ -114,7 +114,7 @@ class RealtimeClientSecretResponse(BaseModel): expires_at: int | None = None value: str - session: dict[str, Any] | None = None + session: dict[str, object] | None = None class RealtimeTranscriptionSessionRequest(BaseModel): @@ -151,7 +151,7 @@ class RealtimeTranscriptionSessionResponse(BaseModel): model_config = {"extra": "allow"} - client_secret: dict[str, Any] | None = None + client_secret: dict[str, object] | None = None class RealtimeErrorDetail(TypedDict): diff --git a/litellm/vector_store_files/utils.py b/litellm/vector_store_files/utils.py index 94ad5c0ecdf..8b4bff921f8 100644 --- a/litellm/vector_store_files/utils.py +++ b/litellm/vector_store_files/utils.py @@ -1,4 +1,5 @@ -from typing import Any, Final, cast, get_type_hints +from collections.abc import Mapping +from typing import Final, cast, get_type_hints from litellm.types.vector_store_files import ( VectorStoreFileCreateRequest, @@ -11,25 +12,25 @@ class VectorStoreFileRequestUtils: """Helper utilities for constructing vector store file requests.""" @staticmethod - def _filter_params(params: dict[str, Any], model: Any) -> dict[str, Any]: + def _filter_params(params: Mapping[str, object], model: type[object]) -> dict[str, object]: valid_keys: Final = get_type_hints(model).keys() return {key: value for key, value in params.items() if key in valid_keys and value is not None} @staticmethod def get_create_request_params( - params: dict[str, Any], + params: Mapping[str, object], ) -> VectorStoreFileCreateRequest: filtered: Final = VectorStoreFileRequestUtils._filter_params(params=params, model=VectorStoreFileCreateRequest) return cast(VectorStoreFileCreateRequest, filtered) @staticmethod - def get_list_query_params(params: dict[str, Any]) -> VectorStoreFileListQueryParams: + def get_list_query_params(params: Mapping[str, object]) -> VectorStoreFileListQueryParams: filtered = VectorStoreFileRequestUtils._filter_params(params=params, model=VectorStoreFileListQueryParams) return cast(VectorStoreFileListQueryParams, filtered) @staticmethod def get_update_request_params( - params: dict[str, Any], + params: Mapping[str, object], ) -> VectorStoreFileUpdateRequest: filtered: Final = VectorStoreFileRequestUtils._filter_params(params=params, model=VectorStoreFileUpdateRequest) return cast(VectorStoreFileUpdateRequest, filtered) From bf9437780b7fd0aa74827239cfe9ea2e19bbacc7 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Tue, 8 Sep 2026 14:09:22 +0000 Subject: [PATCH 07/44] refactor(types): replace Any with real types across 11 more backend files Final batch of the fifth basedpyright Any reduction round. Every change is typing-only and leaves runtime behavior identical. These are the densest remaining files, so the yield per file is small and most of the batch was left alone deliberately. The Riva transcription handler describes the SDK module attributes it reads with Protocols instead of a bare ModuleType, the AWS secret manager stops hiding a botocore header object behind Any, and the sensitive data masker, MCP SSO assertion store and Ovalix guardrail move payload and option annotations to object and Mapping[str, object]. --- .../sensitive_data_masker.py | 10 ++--- .../audio_transcription/handler.py | 37 +++++++++++++--- .../vertex_gemma_models/transformation.py | 7 +-- .../mcp_server/openapi_to_mcp_generator.py | 2 +- .../sso_assertion_store.py | 43 ++++++++++++++++--- litellm/proxy/_lazy_features.py | 4 +- .../guardrail_hooks/ovalix/ovalix.py | 17 +++++--- .../team_callback_endpoints.py | 8 ++-- .../management_endpoints.py | 5 ++- litellm/rag/ingestion/gemini_ingestion.py | 2 +- .../secret_managers/aws_secret_manager_v2.py | 7 ++- 11 files changed, 104 insertions(+), 38 deletions(-) diff --git a/litellm/litellm_core_utils/sensitive_data_masker.py b/litellm/litellm_core_utils/sensitive_data_masker.py index 22dd4170963..e78c4bd24c3 100644 --- a/litellm/litellm_core_utils/sensitive_data_masker.py +++ b/litellm/litellm_core_utils/sensitive_data_masker.py @@ -91,13 +91,13 @@ class SensitiveDataMasker: def _mask_sequence( self, - values: list[Any], + values: Sequence[object], depth: int, max_depth: int, excluded_keys: set[str] | None, key_is_sensitive: bool, - ) -> list[Any]: - masked_items: Final[list[Any]] = [] + ) -> Sequence[object]: + masked_items: Final[list[object]] = [] if depth >= max_depth: return values @@ -197,7 +197,7 @@ def _walk_payload(node: object, key_is_sensitive: bool, depth: int) -> object: return node -def mask_sensitive_keys(data: dict[str, Any], sensitive_fields: set[str]) -> dict[str, Any]: +def mask_sensitive_keys(data: Mapping[str, object], sensitive_fields: set[str]) -> dict[str, object]: """Return a new dict with values masked for keys listed in ``sensitive_fields``. Unlike :meth:`SensitiveDataMasker.mask_dict`, this does exact key-name @@ -209,7 +209,7 @@ def mask_sensitive_keys(data: dict[str, Any], sensitive_fields: set[str]) -> dic range and are replaced with a fixed-length all-mask string, so a short credential is never returned verbatim. """ - masked: Final[dict[str, Any]] = {} + masked: Final[dict[str, object]] = {} mask_char: Final = _default_masker.mask_char min_visible: Final = _default_masker.visible_prefix + _default_masker.visible_suffix for key, value in data.items(): diff --git a/litellm/llms/nvidia_riva/audio_transcription/handler.py b/litellm/llms/nvidia_riva/audio_transcription/handler.py index d188fac8704..bea77a6761c 100644 --- a/litellm/llms/nvidia_riva/audio_transcription/handler.py +++ b/litellm/llms/nvidia_riva/audio_transcription/handler.py @@ -27,7 +27,6 @@ without the optional STT extras installed. import asyncio import inspect from collections.abc import Callable, Iterable -from types import ModuleType from typing import TYPE_CHECKING, Any, Final, Protocol from litellm.litellm_core_utils.audio_utils.utils import ( @@ -95,11 +94,37 @@ class _AudioEncoding(Protocol): def LINEAR_PCM(self) -> object: ... -def _auth_factory(riva_module: ModuleType) -> Callable[..., _RivaAuth]: +class _RivaClientModule(Protocol): + """The ``riva.client`` entry points this handler calls.""" + + @property + def Auth(self) -> Callable[..., _RivaAuth]: ... + + @property + def ASRService(self) -> Callable[[_RivaAuth], _AsrService]: ... + + +class _RivaAsrModule(Protocol): + """The protobuf constructors this handler calls, from whichever module exposes them.""" + + @property + def AudioEncoding(self) -> _AudioEncoding: ... + + @property + def RecognitionConfig(self) -> Callable[..., _RecognitionConfig]: ... + + @property + def StreamingRecognitionConfig(self) -> Callable[..., _StreamingRecognitionConfig]: ... + + @property + def EndpointingConfig(self) -> Callable[..., _EndpointingConfig]: ... + + +def _auth_factory(riva_module: _RivaClientModule) -> Callable[..., _RivaAuth]: return riva_module.Auth -def _audio_encoding(riva_asr_module: ModuleType) -> _AudioEncoding: +def _audio_encoding(riva_asr_module: _RivaAsrModule) -> _AudioEncoding: return riva_asr_module.AudioEncoding @@ -317,7 +342,7 @@ class NvidiaRivaAudioTranscription: def _construct_auth( self, - riva_module: ModuleType, + riva_module: _RivaClientModule, api_base: str, api_key: str | None, optional_params: dict, @@ -349,7 +374,7 @@ class NvidiaRivaAudioTranscription: return _auth_factory(riva_module)(None, use_ssl, api_base, metadata) def _build_recognition_config_proto( - self, riva_asr_module: ModuleType, recognition_config_dict: dict[str, Any] + self, riva_asr_module: _RivaAsrModule, recognition_config_dict: dict[str, Any] ) -> _RecognitionConfig: encoding_name: Final = (recognition_config_dict.get("encoding") or "LINEAR_PCM").upper() encoding_enum: Final[object] = getattr( @@ -436,7 +461,7 @@ class NvidiaRivaAudioTranscription: return final_results -def _import_riva() -> tuple[ModuleType, ModuleType]: +def _import_riva() -> tuple[_RivaClientModule, _RivaAsrModule]: """ Lazy import of ``riva.client`` and ``riva.client.proto.riva_asr_pb2``. diff --git a/litellm/llms/vertex_ai/vertex_gemma_models/transformation.py b/litellm/llms/vertex_ai/vertex_gemma_models/transformation.py index 58cf7c7e702..67b01c2dc43 100644 --- a/litellm/llms/vertex_ai/vertex_gemma_models/transformation.py +++ b/litellm/llms/vertex_ai/vertex_gemma_models/transformation.py @@ -27,6 +27,7 @@ if TYPE_CHECKING: import tiktoken from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj + from litellm.llms.base_llm.base_model_iterator import MockResponseIterator class VertexGemmaConfig(OpenAIGPTConfig): @@ -56,7 +57,7 @@ class VertexGemmaConfig(OpenAIGPTConfig): self, model_response: ModelResponse, stream: bool, - ) -> ModelResponse | Any: + ) -> "ModelResponse | MockResponseIterator": """ Helper method to return fake stream iterator if streaming is requested. @@ -138,7 +139,7 @@ class VertexGemmaConfig(OpenAIGPTConfig): client: HTTPHandler | httpx.Client | None, api_base: str, headers: dict[str, str], # mutable-ok: forwarded to post(headers: dict | None) - request_data: dict[str, Any], # mutable-ok: forwarded to post(json: dict | ...) + request_data: dict[str, object], # mutable-ok: forwarded to post(json: dict | ...) timeout: float | httpx.Timeout | None, ) -> httpx.Response: if isinstance(client, HTTPHandler): @@ -173,7 +174,7 @@ class VertexGemmaConfig(OpenAIGPTConfig): client: AsyncHTTPHandler | httpx.AsyncClient | None, api_base: str, headers: dict[str, str], # mutable-ok: forwarded to post(headers: dict | None) - request_data: dict[str, Any], # mutable-ok: forwarded to post(json: dict | ...) + request_data: dict[str, object], # mutable-ok: forwarded to post(json: dict | ...) timeout: float | httpx.Timeout | None, ) -> httpx.Response: from litellm.llms.custom_httpx.http_handler import get_async_httpx_client diff --git a/litellm/proxy/_experimental/mcp_server/openapi_to_mcp_generator.py b/litellm/proxy/_experimental/mcp_server/openapi_to_mcp_generator.py index 16f58ef5b76..4cf63b1fceb 100644 --- a/litellm/proxy/_experimental/mcp_server/openapi_to_mcp_generator.py +++ b/litellm/proxy/_experimental/mcp_server/openapi_to_mcp_generator.py @@ -453,7 +453,7 @@ def _raise_for_upstream_failure( if response.status_code == 401 and relays_upstream_auth: raise MCPUpstreamAuthError( status_code=response.status_code, - www_authenticate=response.headers.get("www-authenticate"), + www_authenticate=dict(response.headers).get("www-authenticate"), server_name=upstream, ) raise MCPOpenApiUpstreamError(response.status_code, upstream) diff --git a/litellm/proxy/_experimental/mcp_server/outbound_credentials/sso_assertion_store.py b/litellm/proxy/_experimental/mcp_server/outbound_credentials/sso_assertion_store.py index f7b92df5ba3..5503d19211b 100644 --- a/litellm/proxy/_experimental/mcp_server/outbound_credentials/sso_assertion_store.py +++ b/litellm/proxy/_experimental/mcp_server/outbound_credentials/sso_assertion_store.py @@ -18,6 +18,7 @@ TTL ``MCP_SSO_ASSERTION_CACHE_TTL_SECONDS``; invalidation also guards against st from __future__ import annotations import json +from collections.abc import Mapping, Sequence from datetime import datetime, timezone from typing import TYPE_CHECKING, Final, Protocol @@ -29,6 +30,8 @@ from litellm.caching.in_memory_cache import InMemoryCache from litellm.constants import MCP_OAUTH2_TOKEN_CACHE_MAX_SIZE, MCP_SSO_ASSERTION_CACHE_TTL_SECONDS if TYPE_CHECKING: + from prisma.models import LiteLLM_SSOIdentityAssertion + from litellm.proxy.utils import PrismaClient _ASSERTION_DECRYPT_LOG_KEY: Final = "sso_identity_assertion" @@ -36,6 +39,34 @@ _STR_ADAPTER: Final[TypeAdapter[str]] = TypeAdapter(str) _MAYBE_STR_ADAPTER: Final[TypeAdapter[str | None]] = TypeAdapter(str | None) +class _SSOAssertionTable(Protocol): + """The ``LiteLLM_SSOIdentityAssertion`` table operations this store calls.""" + + async def find_unique(self, *, where: Mapping[str, str]) -> LiteLLM_SSOIdentityAssertion | None: ... + + async def find_many(self) -> Sequence[LiteLLM_SSOIdentityAssertion]: ... + + async def upsert(self, *, where: Mapping[str, str], data: Mapping[str, Mapping[str, str]]) -> object: ... + + async def update(self, *, where: Mapping[str, str], data: Mapping[str, str]) -> object: ... + + +class _MCPServerTable(Protocol): + """The ``LiteLLM_MCPServerTable`` lookup the retention gate calls.""" + + async def find_first(self, *, where: Mapping[str, str]) -> object | None: ... + + +def _assertion_table(prisma_client: PrismaClient) -> _SSOAssertionTable: + """The SSO assertion table, typed so the untyped prisma client surface stops here.""" + return prisma_client.db.litellm_ssoidentityassertion + + +def _mcp_server_table(prisma_client: PrismaClient) -> _MCPServerTable: + """The MCP server table, typed so the untyped prisma client surface stops here.""" + return prisma_client.db.litellm_mcpservertable + + class SSOIdentityAssertion(BaseModel): """The IdP material an EMA exchange needs: ``id_token`` is the RFC 8693 subject token, ``expires_at`` bounds its usefulness, and the refresh token renews it without re-login.""" @@ -163,9 +194,7 @@ async def ema_assertion_retention_enabled() -> bool: return True if prisma_client is None: return False - row: Final = await prisma_client.db.litellm_mcpservertable.find_first( - where={"auth_type": MCPAuth.oauth2_id_jag.value} - ) + row: Final = await _mcp_server_table(prisma_client).find_first(where={"auth_type": MCPAuth.oauth2_id_jag.value}) return row is not None @@ -184,7 +213,7 @@ async def persist_sso_identity_assertion( **({"expires_at": assertion.expires_at.isoformat()} if assertion.expires_at else {}), } encoded: Final = _STR_ADAPTER.validate_python(encrypt_value_helper(json.dumps(payload))) - await prisma_client.db.litellm_ssoidentityassertion.upsert( + await _assertion_table(prisma_client).upsert( where={"user_id": user_id}, data={ "create": {"user_id": user_id, "assertion_b64": encoded}, @@ -200,7 +229,7 @@ async def _read_assertion_from_db(user_id: str) -> SSOIdentityAssertion | None: if prisma_client is None: return None - row: Final = await prisma_client.db.litellm_ssoidentityassertion.find_unique(where={"user_id": user_id}) + row: Final = await _assertion_table(prisma_client).find_unique(where={"user_id": user_id}) if row is None: return None raw: Final = _MAYBE_STR_ADAPTER.validate_python( @@ -310,13 +339,13 @@ async def rotate_sso_identity_assertions_master_key(prisma_client: PrismaClient, re_encrypted: Final = _STR_ADAPTER.validate_python( encrypt_value_helper(plaintext, new_encryption_key=new_master_key) ) - await prisma_client.db.litellm_ssoidentityassertion.update( + await _assertion_table(prisma_client).update( where={"user_id": row.user_id}, data={"assertion_b64": re_encrypted}, ) return True - rows: Final = await prisma_client.db.litellm_ssoidentityassertion.find_many() + rows: Final = await _assertion_table(prisma_client).find_many() outcomes: Final = [await _rotate_row(row) for row in rows] verbose_proxy_logger.info( "rotate_sso_identity_assertions_master_key: rotated %d row(s), skipped %d", diff --git a/litellm/proxy/_lazy_features.py b/litellm/proxy/_lazy_features.py index 50e0a961a49..17743e153f2 100644 --- a/litellm/proxy/_lazy_features.py +++ b/litellm/proxy/_lazy_features.py @@ -13,7 +13,7 @@ from collections.abc import Set as AbstractSet from dataclasses import dataclass, field from typing import TYPE_CHECKING, Final -from starlette.types import Receive, Scope, Send +from starlette.types import ASGIApp, Receive, Scope, Send from litellm._logging import verbose_proxy_logger @@ -266,7 +266,7 @@ class LazyFeatureMiddleware: def __init__( self, - app, + app: ASGIApp, fastapi_app: "FastAPI", features: tuple[LazyFeature, ...] = LAZY_FEATURES, ): diff --git a/litellm/proxy/guardrails/guardrail_hooks/ovalix/ovalix.py b/litellm/proxy/guardrails/guardrail_hooks/ovalix/ovalix.py index b31ed4b0f4a..c69b24c0553 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/ovalix/ovalix.py +++ b/litellm/proxy/guardrails/guardrail_hooks/ovalix/ovalix.py @@ -10,6 +10,7 @@ import os from typing import TYPE_CHECKING, Any, Final, Literal import httpx +from typing_extensions import ReadOnly, TypedDict, Unpack from litellm._logging import verbose_proxy_logger from litellm.exceptions import GuardrailRaisedException @@ -33,6 +34,12 @@ BLOCKED_BY_OVALIX_FALLBACK_MESSAGE: Final = "This message was blocked by Ovalix" BLOCKED_ACTION_TYPE: Final = "block" +class _CustomGuardrailOptions(TypedDict, total=False, extra_items=object): + """Base-class constructor options this guardrail forwards untouched to CustomGuardrail.""" + + supported_event_hooks: ReadOnly[list[GuardrailEventHooks]] + + class OvalixGuardrailMissingSecrets(Exception): """Raised when required Ovalix config (API base, key, application/checkpoint IDs) is missing.""" @@ -80,7 +87,7 @@ class OvalixGuardrail(CustomGuardrail): application_id: str | None = None, pre_checkpoint_id: str | None = None, post_checkpoint_id: str | None = None, - **kwargs: Any, + **kwargs: Unpack[_CustomGuardrailOptions], ): self._tracker_api_base = tracker_api_base or os.environ.get("OVALIX_TRACKER_API_BASE") self._tracker_api_key = tracker_api_key or os.environ.get("OVALIX_TRACKER_API_KEY") @@ -88,10 +95,9 @@ class OvalixGuardrail(CustomGuardrail): self._pre_checkpoint_id = pre_checkpoint_id or os.environ.get("OVALIX_PRE_CHECKPOINT_ID") self._post_checkpoint_id = post_checkpoint_id or os.environ.get("OVALIX_POST_CHECKPOINT_ID") - if "supported_event_hooks" not in kwargs: - kwargs["supported_event_hooks"] = [] + supported_event_hooks: Final = kwargs.get("supported_event_hooks", []) - self._validate_config(kwargs["supported_event_hooks"]) + self._validate_config(supported_event_hooks) self._tracker_headers = httpx.Headers( { @@ -103,7 +109,8 @@ class OvalixGuardrail(CustomGuardrail): self._async_handler = get_async_httpx_client(llm_provider=httpxSpecialProvider.GuardrailCallback) - super().__init__(**kwargs) + forwarded: Final[_CustomGuardrailOptions] = {**kwargs, "supported_event_hooks": supported_event_hooks} + super().__init__(**forwarded) verbose_proxy_logger.debug( "Ovalix Guardrail initialized: tracker=%s, application_id=%s, pre_checkpoint_id=%s, post_checkpoint_id=%s", self._tracker_api_base, diff --git a/litellm/proxy/management_endpoints/team_callback_endpoints.py b/litellm/proxy/management_endpoints/team_callback_endpoints.py index fe658a13c24..b7fcd4ac7a9 100644 --- a/litellm/proxy/management_endpoints/team_callback_endpoints.py +++ b/litellm/proxy/management_endpoints/team_callback_endpoints.py @@ -473,7 +473,7 @@ async def delete_team_callback( raise _callback_error(404, f"callback_name = {callback_name} is not registered for team_id = {team_id}.") updated_metadata: Final = {**team_metadata, "logging": remaining_callbacks} # mutable-ok: persisted as JSON - encrypted_metadata: Final = encrypt_callback_vars(updated_metadata) + encrypted_metadata: Final[object] = encrypt_callback_vars(updated_metadata) team_metadata_json: Final = json.dumps(encrypted_metadata) updated_team: Final = await TeamRepository(prisma_client).table.update( @@ -610,8 +610,8 @@ async def disable_team_logging( # _get_dynamic_logging_metadata stops at metadata["logging"], where the API # and Admin UI register callbacks, without ever reading callback_settings. team_metadata["logging"] = [] # mutable-ok: the disabled state is persisted as an empty JSON array - team_metadata = encrypt_callback_vars(team_metadata) - team_metadata_json: Final = json.dumps(team_metadata) + encrypted_metadata: Final[object] = encrypt_callback_vars(team_metadata) + team_metadata_json: Final = json.dumps(encrypted_metadata) # Update team in database updated_team: Final = await TeamRepository(prisma_client).table.update( @@ -643,7 +643,7 @@ async def disable_team_logging( await _emit_team_callback_audit_log( team_id=team_id, before_metadata=before_metadata, - after_metadata=team_metadata, + after_metadata=encrypted_metadata, user_api_key_dict=user_api_key_dict, litellm_changed_by=litellm_changed_by, ) diff --git a/litellm/proxy/vector_store_endpoints/management_endpoints.py b/litellm/proxy/vector_store_endpoints/management_endpoints.py index 0ca2c4c8865..2fb6813a471 100644 --- a/litellm/proxy/vector_store_endpoints/management_endpoints.py +++ b/litellm/proxy/vector_store_endpoints/management_endpoints.py @@ -88,7 +88,7 @@ def _redact_sensitive_litellm_params(litellm_params: object, _depth: int = 0) -> return None if isinstance(litellm_params, str): try: - parsed: Final = json.loads(litellm_params) + parsed: Final[object] = json.loads(litellm_params) except (TypeError, ValueError): return REDACTED_BY_LITELM_STRING return json.dumps(_redact_sensitive_litellm_params(parsed, _depth + 1)) @@ -589,7 +589,8 @@ async def update_vector_store( try: update_data: Final = data.model_dump(exclude_unset=True) - vector_store_id: Final[str] = update_data.pop("vector_store_id") + vector_store_id: Final[str] = data.vector_store_id + update_data.pop("vector_store_id") # Per-store access control: anyone authenticated who passes the # premium-feature gate could otherwise update *any* vector store — diff --git a/litellm/rag/ingestion/gemini_ingestion.py b/litellm/rag/ingestion/gemini_ingestion.py index 73a0159fc9f..b81c2cc0ebe 100644 --- a/litellm/rag/ingestion/gemini_ingestion.py +++ b/litellm/rag/ingestion/gemini_ingestion.py @@ -277,7 +277,7 @@ class GeminiRAGIngestion(BaseRAGIngestion): raise Exception(error_msg) verbose_logger.debug("Initiate resumable upload response: %s", response.headers) # Extract upload URL from response headers - upload_url: Final = response.headers.get("x-goog-upload-url") + upload_url: Final = dict(response.headers).get("x-goog-upload-url") if not upload_url: raise Exception("No upload URL returned in response headers") diff --git a/litellm/secret_managers/aws_secret_manager_v2.py b/litellm/secret_managers/aws_secret_manager_v2.py index e86c8e7c919..9f5bf783958 100644 --- a/litellm/secret_managers/aws_secret_manager_v2.py +++ b/litellm/secret_managers/aws_secret_manager_v2.py @@ -16,7 +16,7 @@ Requires: import json import os -from typing import Any, Final +from typing import TYPE_CHECKING, Final import httpx @@ -35,6 +35,9 @@ from litellm.types.secret_managers.main import KeyManagementSettings from .base_secret_manager import BaseSecretManager +if TYPE_CHECKING: + from botocore.awsrequest import HTTPHeaders + class AWSSecretsManagerV2(BaseAWSLLM, BaseSecretManager): def __init__( @@ -530,7 +533,7 @@ class AWSSecretsManagerV2(BaseAWSLLM, BaseSecretManager): secret_value: str | None = None, optional_params: dict | None = None, request_data: dict | None = None, - ) -> tuple[str, Any, bytes]: + ) -> tuple[str, "HTTPHeaders", bytes]: """Prepare the AWS Secrets Manager request""" try: from botocore.auth import SigV4Auth From e46c816ec7b42bdab4ac6cdbb7df10022629e693 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Mon, 14 Sep 2026 07:31:41 +0000 Subject: [PATCH 08/44] fix: address review feedback on Any reduction - drop redundant Protocol docstrings in dynamodb and otel mount - widen hosted_vllm custom-tool conversion signature from Any to object --- litellm/integrations/dynamodb.py | 4 ---- litellm/integrations/otel/mount.py | 2 -- litellm/llms/hosted_vllm/chat/transformation.py | 4 ++-- 3 files changed, 2 insertions(+), 8 deletions(-) diff --git a/litellm/integrations/dynamodb.py b/litellm/integrations/dynamodb.py index 3401ced4efb..3fbbfe91ddf 100644 --- a/litellm/integrations/dynamodb.py +++ b/litellm/integrations/dynamodb.py @@ -11,14 +11,10 @@ from litellm._uuid import uuid class _DynamoTable(Protocol): - """The one boto3 DynamoDB table call this logger makes.""" - def put_item(self, *, Item: Mapping[str, object]) -> object: ... class _DynamoResource(Protocol): - """The one boto3 DynamoDB resource call this logger makes.""" - def Table(self, name: str) -> _DynamoTable: ... diff --git a/litellm/integrations/otel/mount.py b/litellm/integrations/otel/mount.py index 776d8722d14..9340f6e9e15 100644 --- a/litellm/integrations/otel/mount.py +++ b/litellm/integrations/otel/mount.py @@ -69,8 +69,6 @@ PASSTHROUGH_PREFIXES: Final = frozenset( class _RenameableSpan(Protocol): - """The span surface the passthrough naming hook drives.""" - def is_recording(self) -> bool: ... def update_name(self, name: str) -> None: ... diff --git a/litellm/llms/hosted_vllm/chat/transformation.py b/litellm/llms/hosted_vllm/chat/transformation.py index 92bc857e385..32c60bd01b5 100644 --- a/litellm/llms/hosted_vllm/chat/transformation.py +++ b/litellm/llms/hosted_vllm/chat/transformation.py @@ -4,7 +4,7 @@ Translate from OpenAI's `/v1/chat/completions` to VLLM's `/v1/chat/completions` import json from collections.abc import Coroutine -from typing import Any, Final, Literal, cast, overload +from typing import Final, Literal, cast, overload from litellm.litellm_core_utils.prompt_templates.common_utils import ( _get_image_mime_type_from_url, @@ -28,7 +28,7 @@ from ...openai.chat.gpt_transformation import OpenAIGPTConfig class HostedVLLMChatConfig(OpenAIGPTConfig): - def _convert_custom_tools_to_function_tools(self, tools: list[dict[str, Any]]) -> list[dict[str, object]]: + def _convert_custom_tools_to_function_tools(self, tools: list[dict[str, object]]) -> list[dict[str, object]]: """ vLLM chat completions currently accepts only OpenAI function tools. Convert custom tools into function tools so request validation does not fail. From da7d5fe1284da59502cb6738d853acc8c431af40 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Mon, 14 Sep 2026 11:13:39 +0000 Subject: [PATCH 09/44] refactor: replace Any with precise types across 54 modules Narrow or remove reportAny / reportExplicitAny sites in provider transformations, caching, guardrails, proxy endpoints and enterprise batch-cost polling. Public parameters widen to Mapping/Sequence rather than dict/list so no caller signature breaks, and runtime behavior is unchanged. --- .../proxy/common_utils/check_batch_cost.py | 127 +++++++++++------- .../common_utils/check_responses_cost.py | 31 ++++- .../proxy/hooks/managed_files.py | 22 ++- litellm/caching/caching.py | 2 +- litellm/caching/caching_handler.py | 10 +- litellm/caching/redis_cache.py | 4 +- .../handler.py | 6 +- .../transformation.py | 30 +++-- litellm/cost_calculator.py | 6 +- litellm/integrations/braintrust_logging.py | 8 +- litellm/integrations/custom_guardrail.py | 2 +- .../integrations/datadog/datadog_llm_obs.py | 22 +-- litellm/integrations/galileo.py | 11 +- litellm/integrations/langfuse/langfuse.py | 8 +- .../llm_response_utils/response_metadata.py | 4 +- .../prompt_templates/common_utils.py | 39 ++---- .../a2a/chat/guardrail_translation/handler.py | 2 +- litellm/llms/anthropic/chat/transformation.py | 24 ++-- litellm/llms/anthropic/common_utils.py | 10 +- .../messages/agentic_streaming_iterator.py | 16 ++- litellm/llms/anthropic/files/handler.py | 6 +- .../bedrock/chat/converse_transformation.py | 2 +- .../llms/chatgpt/responses/transformation.py | 22 +-- litellm/llms/cohere/chat/transformation.py | 2 +- .../llms/databricks/chat/transformation.py | 10 +- .../llms/fireworks_ai/chat/transformation.py | 4 +- litellm/llms/gemini/count_tokens/handler.py | 2 +- .../llms/gemini/image_edit/transformation.py | 13 +- litellm/llms/gigachat/chat/transformation.py | 14 +- .../huggingface/embedding/transformation.py | 10 +- .../chat/guardrail_translation/handler.py | 20 +-- litellm/llms/openai/videos/transformation.py | 6 +- .../openrouter/image_edit/transformation.py | 11 +- .../perplexity/embedding/transformation.py | 2 +- .../llms/vertex_ai/gemini/transformation.py | 4 +- .../mcp_server/mcp_server_manager.py | 2 +- litellm/proxy/auth/auth_utils.py | 26 ++-- .../proxy/client/cli/commands/configure.py | 60 +++++---- litellm/proxy/common_utils/callback_utils.py | 6 +- litellm/proxy/db/db_spend_update_writer.py | 20 ++- .../guardrails/guardrail_hooks/akto/akto.py | 36 ++++- .../custom_code/custom_code_guardrail.py | 9 +- .../guardrails/guardrail_hooks/lasso/lasso.py | 12 +- .../vigil_guard/vigil_guard.py | 23 +++- .../guardrail_hooks/xecguard/xecguard.py | 12 +- .../model_management_endpoints.py | 5 +- .../vertex_passthrough_logging_handler.py | 3 +- .../proxy/response_api_endpoints/endpoints.py | 2 +- litellm/proxy/video_endpoints/endpoints.py | 28 ++-- .../mcp/litellm_proxy_mcp_handler.py | 10 +- litellm/responses/streaming_iterator.py | 2 +- .../complexity_router/complexity_router.py | 32 ++--- litellm/types/router.py | 16 +-- .../vector_stores/vector_store_registry.py | 5 +- 54 files changed, 485 insertions(+), 336 deletions(-) diff --git a/enterprise/litellm_enterprise/proxy/common_utils/check_batch_cost.py b/enterprise/litellm_enterprise/proxy/common_utils/check_batch_cost.py index 13e9e5093a8..3ea9b7d9bfd 100644 --- a/enterprise/litellm_enterprise/proxy/common_utils/check_batch_cost.py +++ b/enterprise/litellm_enterprise/proxy/common_utils/check_batch_cost.py @@ -2,10 +2,11 @@ Polls LiteLLM_ManagedObjectTable to check if the batch job is complete, and if the cost has been tracked. """ +from collections.abc import Sequence from dataclasses import replace as dataclasses_replace from datetime import datetime, timedelta, timezone from types import MappingProxyType -from typing import TYPE_CHECKING, Final, List, Literal, Optional, Tuple, cast +from typing import TYPE_CHECKING, Final, List, Literal, Optional, Protocol, Tuple, cast from litellm._logging import verbose_proxy_logger from litellm._uuid import uuid @@ -18,8 +19,8 @@ if TYPE_CHECKING: from prisma import models as prisma_models from litellm.integrations.prometheus import PrometheusLogger - from litellm.proxy._types import LiteLLM_ManagedObjectTable from litellm.proxy.utils import PrismaClient, ProxyLogging + from litellm.repositories.prisma_protocols import TableActions from litellm.router import Router from litellm.types.router import Deployment from litellm.types.utils import LiteLLMBatch @@ -41,6 +42,48 @@ TERMINAL_MANAGED_OBJECT_STATUSES: Final[Tuple[str, ...]] = ( ) +class _ManagedObjectRow(Protocol): + """The managed-object row fields this poller reads off whatever the DB hands back.""" + + @property + def id(self) -> str: ... + + @property + def unified_object_id(self) -> str: ... + + @property + def created_by(self) -> str | None: ... + + @property + def file_object(self) -> object: ... + + +def _managed_object_table(prisma_client: "PrismaClient") -> "TableActions[_ManagedObjectRow]": + """The managed-object table's prisma actions, typed to the row fields this module reads.""" + table: Final[TableActions[_ManagedObjectRow]] = prisma_client.db.litellm_managedobjecttable + return table + + +def _user_table(prisma_client: "PrismaClient") -> "TableActions[prisma_models.LiteLLM_UserTable]": + """The user table's prisma actions.""" + table: Final[TableActions[prisma_models.LiteLLM_UserTable]] = prisma_client.db.litellm_usertable + return table + + +def _token_table(prisma_client: "PrismaClient") -> "TableActions[prisma_models.LiteLLM_VerificationToken]": + """The virtual-key table's prisma actions.""" + table: Final[TableActions[prisma_models.LiteLLM_VerificationToken]] = ( + prisma_client.db.litellm_verificationtoken + ) + return table + + +def _team_table(prisma_client: "PrismaClient") -> "TableActions[prisma_models.LiteLLM_TeamTable]": + """The team table's prisma actions.""" + table: Final[TableActions[prisma_models.LiteLLM_TeamTable]] = prisma_client.db.litellm_teamtable + return table + + class CheckBatchCost: def __init__( self, @@ -73,7 +116,7 @@ class CheckBatchCost: inline for a batch the first poll cycle then accounts again. """ try: - await self.prisma_client.db.litellm_managedobjecttable.find_first( + await _managed_object_table(self.prisma_client).find_first( where={"file_purpose": "batch", "batch_processed": False} ) except Exception as probe_err: @@ -97,10 +140,8 @@ class CheckBatchCost: if not user_id: return {} try: - user_row: prisma_models.LiteLLM_UserTable | None = ( - await self.prisma_client.db.litellm_usertable.find_unique( - where={"user_id": user_id} - ) + user_row: prisma_models.LiteLLM_UserTable | None = await _user_table(self.prisma_client).find_unique( + where={"user_id": user_id} ) if user_row is None: return {} @@ -117,11 +158,9 @@ class CheckBatchCost: if not api_key: return None try: - key_row: prisma_models.LiteLLM_VerificationToken | None = ( - await self.prisma_client.db.litellm_verificationtoken.find_unique( - where={"token": api_key} - ) - ) + key_row: prisma_models.LiteLLM_VerificationToken | None = await _token_table( + self.prisma_client + ).find_unique(where={"token": api_key}) return getattr(key_row, "key_alias", None) if key_row is not None else None except Exception as e: verbose_proxy_logger.error(f"CheckBatchCost: could not look up key alias for batch {batch_id}: {e}") @@ -132,17 +171,15 @@ class CheckBatchCost: if not team_id: return None try: - team_row: prisma_models.LiteLLM_TeamTable | None = ( - await self.prisma_client.db.litellm_teamtable.find_unique( - where={"team_id": team_id} - ) + team_row: prisma_models.LiteLLM_TeamTable | None = await _team_table(self.prisma_client).find_unique( + where={"team_id": team_id} ) return getattr(team_row, "team_alias", None) if team_row is not None else None except Exception as e: verbose_proxy_logger.error(f"CheckBatchCost: could not look up team alias for team {team_id}: {e}") return None - async def _get_org_id(self, job: "LiteLLM_ManagedObjectTable", batch_id: str) -> str | None: + async def _get_org_id(self, job: "_ManagedObjectRow", batch_id: str) -> str | None: org_id = getattr(job, "org_id", None) if org_id: return org_id @@ -150,11 +187,9 @@ class CheckBatchCost: team_id = getattr(job, "team_id", None) if api_key: try: - key_row: prisma_models.LiteLLM_VerificationToken | None = ( - await self.prisma_client.db.litellm_verificationtoken.find_unique( - where={"token": api_key} - ) - ) + key_row: prisma_models.LiteLLM_VerificationToken | None = await _token_table( + self.prisma_client + ).find_unique(where={"token": api_key}) key_org_id = getattr(key_row, "organization_id", None) if key_row is not None else None if key_org_id: return key_org_id @@ -166,10 +201,8 @@ class CheckBatchCost: if not team_id: return None try: - team_row: prisma_models.LiteLLM_TeamTable | None = ( - await self.prisma_client.db.litellm_teamtable.find_unique( - where={"team_id": team_id} - ) + team_row: prisma_models.LiteLLM_TeamTable | None = await _team_table(self.prisma_client).find_unique( + where={"team_id": team_id} ) return getattr(team_row, "organization_id", None) if team_row is not None else None except Exception as e: @@ -177,7 +210,7 @@ class CheckBatchCost: return None async def _build_creator_attribution_metadata( - self, job: "LiteLLM_ManagedObjectTable", batch_id: str + self, job: "_ManagedObjectRow", batch_id: str ) -> dict[str, object]: """ Rebuild the spend-tracking metadata for the key, team, and tags that created the @@ -225,7 +258,7 @@ class CheckBatchCost: should not be polled. """ cutoff: Final = datetime.now(timezone.utc) - timedelta(days=MANAGED_OBJECT_STALENESS_CUTOFF_DAYS) - result: Final = await self.prisma_client.db.litellm_managedobjecttable.update_many( + result: Final = await _managed_object_table(self.prisma_client).update_many( where={ "file_purpose": "batch", "status": {"not_in": list(TERMINAL_MANAGED_OBJECT_STATUSES)}, @@ -244,7 +277,7 @@ class CheckBatchCost: # A row already in a terminal status is never rewritten by the sweep above, so # without this it keeps a poll-page slot forever and starves newer batches. - retired: Final = await self.prisma_client.db.litellm_managedobjecttable.update_many( + retired: Final = await _managed_object_table(self.prisma_client).update_many( where={ "file_purpose": "batch", "batch_processed": False, @@ -259,9 +292,9 @@ class CheckBatchCost: f"{MANAGED_OBJECT_STALENESS_CUTOFF_DAYS} days that were never costed" ) - async def _fallback_find_jobs(self) -> list: + async def _fallback_find_jobs(self) -> "Sequence[_ManagedObjectRow]": """Query batch jobs without the batch_processed filter (for older schemas).""" - return await self.prisma_client.db.litellm_managedobjecttable.find_many( + return await _managed_object_table(self.prisma_client).find_many( where={ "file_purpose": "batch", "status": { @@ -279,7 +312,7 @@ class CheckBatchCost: order={"created_at": "asc"}, ) - async def _retire_job(self, job: "LiteLLM_ManagedObjectTable", reason: str) -> None: + async def _retire_job(self, job: "_ManagedObjectRow", reason: str) -> None: """ Take a row that can never be costed out of the poll page. Leaving it selectable would burn one of the MAX_OBJECTS_PER_POLL_CYCLE slots on every future cycle, and @@ -292,7 +325,7 @@ class CheckBatchCost: else {"status": "stale_expired"} ) try: - await self.prisma_client.db.litellm_managedobjecttable.update( + await _managed_object_table(self.prisma_client).update( where={"id": job.id}, data=data, ) @@ -306,7 +339,7 @@ class CheckBatchCost: "so it will no longer be polled" ) - async def _claim_job_for_costing(self, job: "LiteLLM_ManagedObjectTable") -> bool: + async def _claim_job_for_costing(self, job: "_ManagedObjectRow") -> bool: """ Atomically flip batch_processed from false to true, returning whether this pod won the row. Every pod and uvicorn worker schedules its own poller against the shared @@ -321,7 +354,7 @@ class CheckBatchCost: if not self._has_batch_processed_column: return True try: - claimed: Final = await self.prisma_client.db.litellm_managedobjecttable.update_many( + claimed: Final = await _managed_object_table(self.prisma_client).update_many( where={"id": job.id, "batch_processed": False}, data={"batch_processed": True}, ) @@ -332,7 +365,7 @@ class CheckBatchCost: return False return claimed > 0 - async def _release_job_claim(self, job: "LiteLLM_ManagedObjectTable") -> None: + async def _release_job_claim(self, job: "_ManagedObjectRow") -> None: """Give a claimed row back once billing it failed, so a later poll cycle retries it. Safe to match on batch_processed=True: while this poller is active the retrieve @@ -342,7 +375,7 @@ class CheckBatchCost: if not self._has_batch_processed_column: return try: - await self.prisma_client.db.litellm_managedobjecttable.update_many( + await _managed_object_table(self.prisma_client).update_many( where={"id": job.id, "batch_processed": True}, data={"batch_processed": False}, ) @@ -353,7 +386,7 @@ class CheckBatchCost: ) @staticmethod - def _has_unified_id_without_model(job: "LiteLLM_ManagedObjectTable") -> bool: + def _has_unified_id_without_model(job: "_ManagedObjectRow") -> bool: """A unified id that decodes but carries no model_id can never be routed.""" from litellm.proxy.openai_files_endpoints.common_utils import ( convert_b64_uid_to_unified_uid, @@ -402,7 +435,7 @@ class CheckBatchCost: return isinstance(error, (NotFoundError, openai.NotFoundError)) and output_file_id in str(error) async def _finalize_unbilled_terminal_job( - self, job: "prisma_models.LiteLLM_ManagedObjectTable", response: "LiteLLMBatch" + self, job: "_ManagedObjectRow", response: "LiteLLMBatch" ) -> None: """Persist a terminal batch that has nothing billable, converting any raw provider file ids to managed ids, and take it out of the poll page.""" @@ -426,7 +459,7 @@ class CheckBatchCost: "file_object": response.model_dump_json(), **({"batch_processed": True} if self._has_batch_processed_column else {}), } - await self.prisma_client.db.litellm_managedobjecttable.update( + await _managed_object_table(self.prisma_client).update( where={"id": job.id}, data=update_data, ) @@ -447,7 +480,7 @@ class CheckBatchCost: def _resolve_job_routing( self, - job: "LiteLLM_ManagedObjectTable", + job: "_ManagedObjectRow", prom_logger: Optional["PrometheusLogger"], ) -> Optional[Tuple[str, str]]: """ @@ -524,7 +557,7 @@ class CheckBatchCost: def _resolve_unmanaged_provider_routing( self, - job: "LiteLLM_ManagedObjectTable", + job: "_ManagedObjectRow", prom_logger: Optional["PrometheusLogger"], llm_provider: str, bare_model_name: str, @@ -620,7 +653,7 @@ class CheckBatchCost: @classmethod def _get_managed_file_model_name( cls, - job: "LiteLLM_ManagedObjectTable", + job: "_ManagedObjectRow", deployment_info: "Deployment", ) -> Optional[str]: """ @@ -640,7 +673,7 @@ class CheckBatchCost: ) @staticmethod - def _get_input_file_id(job: "LiteLLM_ManagedObjectTable") -> Optional[str]: + def _get_input_file_id(job: "_ManagedObjectRow") -> Optional[str]: import json from litellm.types.utils import LiteLLMBatch @@ -660,7 +693,7 @@ class CheckBatchCost: async def _track_completed_batch_cost( self, - job: "LiteLLM_ManagedObjectTable", + job: "_ManagedObjectRow", response: "LiteLLMBatch", model_id: str, batch_id: str, @@ -936,7 +969,7 @@ class CheckBatchCost: # endpoint may transition a batch to "complete" before # CheckBatchCost runs. The batch_processed=False filter # already prevents reprocessing finished batches. - jobs = await self.prisma_client.db.litellm_managedobjecttable.find_many( + jobs = await _managed_object_table(self.prisma_client).find_many( where={ "file_purpose": "batch", "batch_processed": False, @@ -1038,7 +1071,7 @@ class CheckBatchCost: } if self._has_batch_processed_column: update_data["batch_processed"] = True - await self.prisma_client.db.litellm_managedobjecttable.update( + await _managed_object_table(self.prisma_client).update( where={"id": job.id}, data=update_data, ) diff --git a/enterprise/litellm_enterprise/proxy/common_utils/check_responses_cost.py b/enterprise/litellm_enterprise/proxy/common_utils/check_responses_cost.py index 06cf5fcf82f..1bc41f2aa5b 100644 --- a/enterprise/litellm_enterprise/proxy/common_utils/check_responses_cost.py +++ b/enterprise/litellm_enterprise/proxy/common_utils/check_responses_cost.py @@ -6,7 +6,7 @@ same route are non-inference and free. """ from datetime import datetime, timedelta, timezone -from typing import TYPE_CHECKING, Dict, Optional, cast +from typing import TYPE_CHECKING, Dict, Final, Optional, Protocol, cast import litellm from litellm._logging import verbose_proxy_logger @@ -22,11 +22,34 @@ from litellm.types.utils import BACKGROUND_RESPONSE_COST_POLL_CALL_ORIGIN if TYPE_CHECKING: from litellm.proxy.utils import PrismaClient, ProxyLogging + from litellm.repositories.prisma_protocols import TableActions from litellm.router import Router TERMINAL_RESPONSE_STATUSES = frozenset({"completed", "failed", "cancelled", "incomplete"}) +class _ManagedObjectRow(Protocol): + """The managed-object row fields this poller reads off whatever the DB hands back.""" + + @property + def id(self) -> str: ... + + @property + def unified_object_id(self) -> str: ... + + @property + def created_by(self) -> str | None: ... + + @property + def file_object(self) -> object: ... + + +def _managed_object_table(prisma_client: "PrismaClient") -> "TableActions[_ManagedObjectRow]": + """The managed-object table's prisma actions, typed to the row fields this poller reads.""" + table: Final[TableActions[_ManagedObjectRow]] = prisma_client.db.litellm_managedobjecttable + return table + + class CheckResponsesCost: def __init__( self, @@ -128,7 +151,7 @@ class CheckResponsesCost: f"CheckResponsesCost: stale cleanup failed (poll will continue): {cleanup_err}" ) - jobs = await self.prisma_client.db.litellm_managedobjecttable.find_many( + jobs = await _managed_object_table(self.prisma_client).find_many( where={ "status": {"in": ["queued", "in_progress"]}, "file_purpose": "response", @@ -138,7 +161,7 @@ class CheckResponsesCost: ) verbose_proxy_logger.debug(f"Found {len(jobs)} response jobs to check") - completed_jobs = [] + completed_jobs: Final[list[_ManagedObjectRow]] = [] for job in jobs: unified_object_id = job.unified_object_id @@ -189,7 +212,7 @@ class CheckResponsesCost: # Mark completed jobs in the database if len(completed_jobs) > 0: - await self.prisma_client.db.litellm_managedobjecttable.update_many( + await _managed_object_table(self.prisma_client).update_many( where={"id": {"in": [job.id for job in completed_jobs]}}, data={"status": "completed"}, ) diff --git a/enterprise/litellm_enterprise/proxy/hooks/managed_files.py b/enterprise/litellm_enterprise/proxy/hooks/managed_files.py index 4899b87da7a..5204894bee6 100644 --- a/enterprise/litellm_enterprise/proxy/hooks/managed_files.py +++ b/enterprise/litellm_enterprise/proxy/hooks/managed_files.py @@ -465,10 +465,8 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): """ if self.prisma_client is None: return - managed_object = ( - await self.prisma_client.db.litellm_managedobjecttable.find_first( - where={"OR": [{"unified_object_id": object_id}, {"model_object_id": object_id}]} - ) + managed_object = await _managed_object_table(self.prisma_client).find_first( + where={"OR": [{"unified_object_id": object_id}, {"model_object_id": object_id}]} ) if managed_object is None: return @@ -493,10 +491,8 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): """ if self.prisma_client is None: return - managed_file = ( - await self.prisma_client.db.litellm_managedfiletable.find_first( - where={"OR": [{"unified_file_id": file_id}, {"flat_model_file_ids": {"has": file_id}}]} - ) + managed_file = await _managed_file_table(self.prisma_client).find_first( + where={"OR": [{"unified_file_id": file_id}, {"flat_model_file_ids": {"has": file_id}}]} ) if managed_file is None: return @@ -519,8 +515,8 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): provider_file_ids = tuple( file_id for file_id in ( - getattr(response, "output_file_id", None), - getattr(response, "error_file_id", None), + response.output_file_id, + response.error_file_id, ) if file_id and not _is_base64_encoded_unified_file_id(file_id) ) @@ -528,10 +524,8 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): return if self.prisma_client is None: return - batch_row = ( - await self.prisma_client.db.litellm_managedobjecttable.find_first( - where={"unified_object_id": response.id} - ) + batch_row = await _managed_object_table(self.prisma_client).find_first( + where={"unified_object_id": response.id} ) if batch_row is None or ( batch_row.created_by is None and batch_row.team_id is None diff --git a/litellm/caching/caching.py b/litellm/caching/caching.py index d6dd2a073af..813ee655a1d 100644 --- a/litellm/caching/caching.py +++ b/litellm/caching/caching.py @@ -81,7 +81,7 @@ class Cache: s3_aws_access_key_id: str | None = None, s3_aws_secret_access_key: str | None = None, s3_aws_session_token: str | None = None, - s3_config: Any | None = None, + s3_config: object | None = None, s3_path: str | None = None, gcs_bucket_name: str | None = None, gcs_path_service_account: str | None = None, diff --git a/litellm/caching/caching_handler.py b/litellm/caching/caching_handler.py index 139dcf058d2..2c4f80b8708 100644 --- a/litellm/caching/caching_handler.py +++ b/litellm/caching/caching_handler.py @@ -73,7 +73,7 @@ class CachingHandlerResponse(BaseModel): For embeddings there can be a cache hit for some of the inputs in the list and a cache miss for others """ - cached_result: Any | None = None + cached_result: object | None = None final_embedding_cached_response: EmbeddingResponse | None = None embedding_all_elements_cache_hit: bool = False # this is set to True when all elements in the list have a cache hit in the embedding cache, if true return the final_embedding_cached_response no need to make an API call @@ -707,7 +707,7 @@ class LLMCachingHandler: async def _retrieve_from_cache( self, call_type: str, kwargs: dict[str, object], args: tuple[object, ...] - ) -> Any | None: + ) -> object | None: """ Internal method to - get cache key @@ -953,7 +953,7 @@ class LLMCachingHandler: def _convert_cached_stream_response( self, - cached_result: Any, + cached_result: dict[str, object], call_type: str, logging_obj: LiteLLMLoggingObj, model: str, @@ -982,7 +982,7 @@ class LLMCachingHandler: async def async_set_cache( self, - result: Any, + result: object, original_function: Callable, kwargs: dict[str, Any], args: tuple[object, ...] | None = None, @@ -1050,7 +1050,7 @@ class LLMCachingHandler: def sync_set_cache( self, - result: Any, + result: object, kwargs: dict[str, object], args: tuple[object, ...] | None = None, ): diff --git a/litellm/caching/redis_cache.py b/litellm/caching/redis_cache.py index d0cefcb6086..e2b48159e01 100644 --- a/litellm/caching/redis_cache.py +++ b/litellm/caching/redis_cache.py @@ -934,7 +934,7 @@ class RedisCache(BaseCache): client: object = None, ) -> object: async def execute() -> object: - executor: Callable[..., Awaitable[Any]] | None = litellm.in_memory_llm_clients_cache.get_cache( + executor: Callable[..., Awaitable[object]] | None = litellm.in_memory_llm_clients_cache.get_cache( key=script_cache_key ) if executor is None: @@ -946,7 +946,7 @@ class RedisCache(BaseCache): return run_script - def _register_script_for_current_loop(self, script: str) -> Callable[..., Awaitable[Any]]: + def _register_script_for_current_loop(self, script: str) -> Callable[..., Awaitable[object]]: """ Register the script against the current event loop's Redis client. diff --git a/litellm/completion_extras/litellm_responses_transformation/handler.py b/litellm/completion_extras/litellm_responses_transformation/handler.py index f494d6610a1..642a78789b2 100644 --- a/litellm/completion_extras/litellm_responses_transformation/handler.py +++ b/litellm/completion_extras/litellm_responses_transformation/handler.py @@ -2,7 +2,7 @@ Handler for transforming /chat/completions api requests to litellm.responses requests """ -from collections.abc import Coroutine +from collections.abc import AsyncIterable, Coroutine, Iterable from typing import TYPE_CHECKING, Any, Final, Union from typing_extensions import TypedDict @@ -74,7 +74,7 @@ class ResponsesToCompletionBridgeHandler: existing.setdefault(key, value) return response - def _collect_response_from_stream(self, stream_iter: Any) -> "ResponsesAPIResponse": + def _collect_response_from_stream(self, stream_iter: Iterable[object]) -> "ResponsesAPIResponse": for _ in stream_iter: pass @@ -89,7 +89,7 @@ class ResponsesToCompletionBridgeHandler: raise ValueError("Stream completed response is invalid") return response - async def _collect_response_from_stream_async(self, stream_iter: Any) -> "ResponsesAPIResponse": + async def _collect_response_from_stream_async(self, stream_iter: AsyncIterable[object]) -> "ResponsesAPIResponse": async for _ in stream_iter: pass diff --git a/litellm/completion_extras/litellm_responses_transformation/transformation.py b/litellm/completion_extras/litellm_responses_transformation/transformation.py index 5a6debc4af5..a5d67273768 100644 --- a/litellm/completion_extras/litellm_responses_transformation/transformation.py +++ b/litellm/completion_extras/litellm_responses_transformation/transformation.py @@ -6,7 +6,7 @@ import json import os from collections.abc import AsyncIterator, Callable, Iterable, Iterator, Mapping, Sequence from types import MappingProxyType -from typing import TYPE_CHECKING, Any, Final, Literal, TypedDict, Union, cast, get_args +from typing import TYPE_CHECKING, Any, Final, Literal, TypedDict, TypeVar, Union, cast, get_args from openai.types.chat import ChatCompletion from openai.types.responses import Response @@ -52,7 +52,7 @@ from litellm.types.llms.openai import ( from litellm.types.utils import GenericStreamingChunk, ModelResponseStream if TYPE_CHECKING: - from openai.types.responses import ResponseInputImageParam + from openai.types.responses import ResponseInputImageParam, ResponseOutputItem from openai.types.responses.response_text_config_param import ( ResponseTextConfigParam as ResponseText, ) @@ -197,6 +197,9 @@ def _as_chat_reasoning_items( return cast(list[ChatCompletionReasoningItem], list(reasoning_items)) +_ToolChoiceT = TypeVar("_ToolChoiceT") + + def _map_incomplete_reason_to_finish_reason(incomplete_reason: str | None) -> Literal["length", "content_filter"]: if incomplete_reason == "content_filter": return "content_filter" @@ -291,7 +294,9 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): def __init__(self): pass - def _normalize_tool_choice_for_responses_api(self, tool_choice: Any) -> Any: + def _normalize_tool_choice_for_responses_api( + self, tool_choice: _ToolChoiceT + ) -> _ToolChoiceT | ToolChoiceFunctionParam | ToolChoiceCustomParam | Literal["auto", "none", "required"]: """Chat tool_choice nests the name under function/custom; Responses API expects top-level name.""" if not isinstance(tool_choice, dict): return tool_choice @@ -497,7 +502,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): responses_api_request["max_output_tokens"] = value elif key == "tools" and value is not None: responses_api_request["tools"] = self._convert_tools_to_responses_format( - cast(list[dict[str, Any]], value) + cast(list[dict[str, object]], value) ) elif key == "response_format": text_format = self._transform_response_format_to_text_format(value) @@ -810,7 +815,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): response_output: Final = response_payload.get("output") if not isinstance(response_output, list) or len(response_output) == 0: return None - return cast(list[dict[str, Any]], response_output) + return cast(list[dict[str, object]], response_output) @classmethod def _recover_output_items_from_raw_sse(cls, raw_sse: str | None) -> list[dict[str, object]]: @@ -893,10 +898,12 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): output_items = raw_response.output if len(output_items) == 0: - recovered_output_items: Final = self._recover_output_items_from_logging(logging_obj) + recovered_output_items: Final[list[ResponseOutputItem | dict[str, object]]] = [ + *self._recover_output_items_from_logging(logging_obj) + ] if recovered_output_items: - output_items = cast(Any, recovered_output_items) - raw_response.output = cast(Any, recovered_output_items) + output_items = recovered_output_items + raw_response.output = recovered_output_items verbose_logger.warning( "Recovered empty Responses API output from raw SSE for model=%s", model, @@ -1092,7 +1099,9 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): verbose_logger.debug("Chat provider: Other content type -> %s", result) return result - def _convert_tools_to_responses_format(self, tools: list[dict[str, Any]]) -> list["ALL_RESPONSES_API_TOOL_PARAMS"]: + def _convert_tools_to_responses_format( + self, tools: list[dict[str, object]] + ) -> list["ALL_RESPONSES_API_TOOL_PARAMS"]: """Convert chat completion tools to responses API tools format""" responses_tools: Final[list[ALL_RESPONSES_API_TOOL_PARAMS]] = [] for tool in tools: @@ -1108,12 +1117,11 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): description=function_tool.get("description"), ) ) - elif tool.get("type") == "custom" and isinstance(tool.get("custom"), dict): + elif tool.get("type") == "custom" and isinstance(custom_payload := tool.get("custom"), dict): from litellm.litellm_core_utils.prompt_templates.common_utils import ( convert_custom_tool_format_to_responses_shape, ) - custom_payload = tool["custom"] flat_custom = CustomToolParam(type="custom", name=custom_payload.get("name", "")) if custom_payload.get("description") is not None: flat_custom["description"] = custom_payload["description"] diff --git a/litellm/cost_calculator.py b/litellm/cost_calculator.py index 440d97d13be..d93bf1e8769 100644 --- a/litellm/cost_calculator.py +++ b/litellm/cost_calculator.py @@ -351,7 +351,7 @@ def cost_per_token( data_residency: str | None = None, # for OpenAI regional-processing uplift (e.g. "eu", "us") ### VERTEX LOCATION ### vertex_location: str | None = None, # for Vertex AI regional-endpoint uplift (e.g. "us-east5", "global") - response: Any | None = None, + response: object | None = None, ### REQUEST MODEL ### request_model: str | None = None, # original request model for router detection custom_model_info: OCRPricing | None = None, @@ -607,7 +607,7 @@ def cost_per_token( model=model, custom_llm_provider=custom_llm_provider, number_of_queries=number_of_queries or 1, - optional_params=(response._hidden_params if response and hasattr(response, "_hidden_params") else None), + optional_params=(getattr(response, "_hidden_params", None) if response else None), ) elif custom_llm_provider == "vertex_ai": cost_router: Final = google_cost_router( @@ -996,7 +996,7 @@ def _is_known_usage_objects(usage_obj): ) -def _infer_call_type(call_type: CallTypesLiteral | None, completion_response: Any) -> CallTypesLiteral | None: +def _infer_call_type(call_type: CallTypesLiteral | None, completion_response: object) -> CallTypesLiteral | None: if call_type is not None: return call_type diff --git a/litellm/integrations/braintrust_logging.py b/litellm/integrations/braintrust_logging.py index aaf72a0bc4e..501f5749ea4 100644 --- a/litellm/integrations/braintrust_logging.py +++ b/litellm/integrations/braintrust_logging.py @@ -139,13 +139,13 @@ class BraintrustLogger(CustomLogger): ): output = None elif response_obj is not None and isinstance(response_obj, litellm.ModelResponse): - output = response_obj["choices"][0]["message"].json() + output = response_obj.choices[0].message.json() choices = response_obj["choices"] elif response_obj is not None and isinstance(response_obj, litellm.TextCompletionResponse): output = response_obj.choices[0].text choices = response_obj.choices elif response_obj is not None and isinstance(response_obj, litellm.ImageResponse): - output = response_obj["data"] + output = response_obj.data litellm_params: Final = kwargs.get("litellm_params", {}) or {} dynamic_metadata: Final = litellm_params.get("metadata", {}) or {} @@ -264,13 +264,13 @@ class BraintrustLogger(CustomLogger): ): output = None elif response_obj is not None and isinstance(response_obj, litellm.ModelResponse): - output = response_obj["choices"][0]["message"].json() + output = response_obj.choices[0].message.json() choices = response_obj["choices"] elif response_obj is not None and isinstance(response_obj, litellm.TextCompletionResponse): output = response_obj.choices[0].text choices = response_obj.choices elif response_obj is not None and isinstance(response_obj, litellm.ImageResponse): - output = response_obj["data"] + output = response_obj.data litellm_params: Final = kwargs.get("litellm_params", {}) dynamic_metadata: Final = litellm_params.get("metadata", {}) or {} diff --git a/litellm/integrations/custom_guardrail.py b/litellm/integrations/custom_guardrail.py index 39adea30828..eaee84fae93 100644 --- a/litellm/integrations/custom_guardrail.py +++ b/litellm/integrations/custom_guardrail.py @@ -155,7 +155,7 @@ class CustomGuardrail(CustomLogger): def __init_subclass__(cls, **kwargs: object) -> None: # kwargs-ok: forwarded to cooperative __init_subclass__ hooks super().__init_subclass__(**kwargs) - own_apply_guardrail: Final = cls.__dict__.get("apply_guardrail") + own_apply_guardrail: Final[object] = cls.__dict__.get("apply_guardrail") if own_apply_guardrail is None or LOGS_GUARDRAIL_INFORMATION_MARKER in vars(own_apply_guardrail): return cls.apply_guardrail = log_guardrail_information(own_apply_guardrail) diff --git a/litellm/integrations/datadog/datadog_llm_obs.py b/litellm/integrations/datadog/datadog_llm_obs.py index c64a12c6d75..3afe38ea075 100644 --- a/litellm/integrations/datadog/datadog_llm_obs.py +++ b/litellm/integrations/datadog/datadog_llm_obs.py @@ -54,7 +54,7 @@ from litellm.types.utils import ( StandardLoggingPayloadErrorInformation, ) -_EMPTY_MAPPING: Final[Mapping[str, Any]] = MappingProxyType({}) +_EMPTY_MAPPING: Final[Mapping[str, object]] = MappingProxyType({}) _EMPTY_MESSAGE: Final[Message] = {"role": "", "content": ""} _MAX_PARSED_TOOL_ARGUMENT_CHARS: Final = 256 * 1024 _SAFE_REDACTED_MESSAGE_ROLES: Final = frozenset( @@ -154,7 +154,7 @@ def _guardrail_information_without_prompt_carriers( return tuple(_guardrail_entry_without_prompt_carriers(entry) for entry in _guardrail_entries(guardrail_information)) -def _metadata_without_prompt_carriers(standard_logging_metadata: Mapping[str, Any]) -> Mapping[str, Any]: +def _metadata_without_prompt_carriers(standard_logging_metadata: Mapping[str, object]) -> Mapping[str, object]: """The metadata minus the records that quote prompts, tool arguments, tool results, or retrieved text.""" return MappingProxyType( { @@ -237,7 +237,7 @@ def _declared_cost_tags(span_tags: Sequence[str]) -> tuple[str, ...]: return tuple(dimension for dimension in _COST_DIMENSIONS if dimension in present) -def _reasoning_output_tokens(usage_object: Mapping[str, Any] | None) -> float: +def _reasoning_output_tokens(usage_object: Mapping[str, object] | None) -> float: """The provider's reasoning-token count, from either the chat or the responses spelling.""" if usage_object is None: return 0.0 @@ -254,20 +254,20 @@ def _reasoning_output_tokens(usage_object: Mapping[str, Any] | None) -> float: ) -def _mapping_field(source: Mapping[str, Any], key: str) -> Mapping[str, Any]: +def _mapping_field(source: Mapping[str, object], key: str) -> Mapping[str, Any]: """The value at `key` when it is a mapping, else an empty one.""" value: Final = source.get(key) return value if isinstance(value, dict) else _EMPTY_MAPPING -def _content_blocks(message: Mapping[str, Any]) -> tuple[Mapping[str, Any], ...]: +def _content_blocks(message: Mapping[str, object]) -> tuple[Mapping[str, Any], ...]: content: Final = message.get("content") if not isinstance(content, list): return () return tuple(block for block in content if isinstance(block, dict)) -def _to_dd_arguments(raw_arguments: object) -> dict[str, Any] | str: +def _to_dd_arguments(raw_arguments: object) -> dict[str, object] | str: """ Arguments as the object LLM Obs types them as, or the raw string when they are not one. @@ -282,7 +282,7 @@ def _to_dd_arguments(raw_arguments: object) -> dict[str, Any] | str: return parsed if isinstance(parsed, dict) else raw_arguments -def _to_dd_tool_calls(message: Mapping[str, Any]) -> tuple[ToolCall, ...]: +def _to_dd_tool_calls(message: Mapping[str, object]) -> tuple[ToolCall, ...]: """ The tool calls a message carries, in LLM Obs' ToolCall schema, from either dialect. @@ -315,7 +315,7 @@ def _to_dd_tool_calls(message: Mapping[str, Any]) -> tuple[ToolCall, ...]: return openai_calls + anthropic_calls -def _to_dd_tool_results(message: Mapping[str, Any], tool_call_names: Mapping[str, str]) -> tuple[ToolResult, ...]: +def _to_dd_tool_results(message: Mapping[str, object], tool_call_names: Mapping[str, str]) -> tuple[ToolResult, ...]: """ The tool results a message carries, linked back to the call each answers. @@ -400,7 +400,7 @@ def _to_dd_messages(messages: object) -> tuple[Message, ...]: return tuple(_to_dd_message(message, tool_call_names) for message in messages) -def _to_dd_tool_definition(entry: Mapping[str, Any]) -> ToolDefinition | None: +def _to_dd_tool_definition(entry: Mapping[str, object]) -> ToolDefinition | None: function: Final = entry.get("function") declared: Final[Mapping[str, Any]] = function if isinstance(function, dict) else entry name: Final = declared.get("name") @@ -683,7 +683,7 @@ class DataDogLLMObsLogger(CustomBatchLogger): if callable(current_span_fn): current_span: Final = current_span_fn() if current_span is not None: - trace_id: Final = getattr(current_span, "trace_id", None) + trace_id: Final[object] = getattr(current_span, "trace_id", None) if trace_id is not None: return str(trace_id) except Exception: @@ -716,7 +716,7 @@ class DataDogLLMObsLogger(CustomBatchLogger): def redacts_messages_itself(self) -> bool: return True - def _payload_logging_is_off(self, kwargs: Mapping[str, Any]) -> bool: + def _payload_logging_is_off(self, kwargs: Mapping[str, object]) -> bool: return ( bool(self.turn_off_message_logging) or self.message_logging is not True diff --git a/litellm/integrations/galileo.py b/litellm/integrations/galileo.py index b27618993a3..010f8ad8ef2 100644 --- a/litellm/integrations/galileo.py +++ b/litellm/integrations/galileo.py @@ -396,12 +396,13 @@ class GalileoObserve(CustomLogger): ) @staticmethod - def _log_v2_payload_validation(payload: dict[str, Any]) -> None: + def _log_v2_payload_validation(payload: dict[str, object]) -> None: missing_fields: Final[list[str]] = [] - traces: Final[Sequence[object]] = payload.get("traces", []) - if not traces: + traces_value: Final = payload.get("traces", []) + if not traces_value: missing_fields.append("traces") + traces: Final[Sequence[object]] = traces_value if isinstance(traces_value, list) else [] for trace_index, trace in enumerate(traces): if not isinstance(trace, dict): continue @@ -425,8 +426,8 @@ class GalileoObserve(CustomLogger): missing_fields, ) - def _log_flush_payload(self, url: str, payload: dict[str, Any]) -> None: - traces: Final[Sequence[object]] = payload.get("traces", []) + def _log_flush_payload(self, url: str, payload: dict[str, object]) -> None: + traces: Final = payload.get("traces") verbose_logger.debug( "Galileo Logger flush URL: %s trace_count=%s", url, diff --git a/litellm/integrations/langfuse/langfuse.py b/litellm/integrations/langfuse/langfuse.py index b75369965de..c89506facbd 100644 --- a/litellm/integrations/langfuse/langfuse.py +++ b/litellm/integrations/langfuse/langfuse.py @@ -4,7 +4,7 @@ import inspect import os import re import traceback -from collections.abc import Callable, Iterable, Mapping +from collections.abc import Callable, Iterable, Mapping, Sequence from datetime import datetime from functools import lru_cache from types import MappingProxyType @@ -447,7 +447,7 @@ class LangFuseLogger: prompt: dict, level: str, status_message: str | None, - ) -> tuple[dict | None, str | dict | list | None]: + ) -> tuple[dict | None, str | dict | Sequence[object] | None]: """ Get the input and output content for Langfuse logging @@ -463,7 +463,7 @@ class LangFuseLogger: output: The output content for Langfuse logging """ input = None - output: str | dict | list[Any] | None = None + output: str | dict | Sequence[object] | None = None if level == "ERROR" and status_message is not None and isinstance(status_message, str): input = prompt output = status_message @@ -575,7 +575,7 @@ class LangFuseLogger: user_id: str | None, metadata: dict[str, object], litellm_params: dict, - output: str | dict | list | None, + output: str | dict | Sequence[object] | None, start_time: datetime | None, end_time: datetime | None, kwargs: dict, diff --git a/litellm/litellm_core_utils/llm_response_utils/response_metadata.py b/litellm/litellm_core_utils/llm_response_utils/response_metadata.py index c83c266a17e..d279fb9e259 100644 --- a/litellm/litellm_core_utils/llm_response_utils/response_metadata.py +++ b/litellm/litellm_core_utils/llm_response_utils/response_metadata.py @@ -1,6 +1,6 @@ import datetime from collections.abc import Mapping -from typing import Any, Final +from typing import Final import httpx @@ -54,7 +54,7 @@ class ResponseMetadata: Handles setting and managing `_hidden_params`, `response_time_ms`, and `litellm_overhead_time_ms` for LiteLLM responses """ - def __init__(self, result: Any): + def __init__(self, result: object): self.result = result self._hidden_params: HiddenParams | dict = getattr(result, "_hidden_params", {}) or {} diff --git a/litellm/litellm_core_utils/prompt_templates/common_utils.py b/litellm/litellm_core_utils/prompt_templates/common_utils.py index 2485896184e..f2b03e72497 100644 --- a/litellm/litellm_core_utils/prompt_templates/common_utils.py +++ b/litellm/litellm_core_utils/prompt_templates/common_utils.py @@ -13,14 +13,6 @@ from pathlib import Path from types import MappingProxyType from typing import TYPE_CHECKING, Any, Final, Literal, TypeVar, cast -from openai.types.chat.chat_completion_custom_tool_param import ( - CustomFormatGrammar, - CustomFormatGrammarGrammar, -) -from openai.types.shared_params.custom_tool_input_format import ( - Grammar as ResponsesGrammarFormat, -) - import litellm from litellm import verbose_logger from litellm.router_utils.batch_utils import InMemoryFile @@ -59,7 +51,7 @@ if TYPE_CHECKING: def handle_any_messages_to_chat_completion_str_messages_conversion( - messages: Any, + messages: object, ) -> list[dict[str, str]]: """ Handles any messages to chat completion str messages conversion @@ -804,7 +796,7 @@ def extract_file_metadata(file_data: FileTypes) -> tuple[str | None, str | None] """ filename: str | None = None content_type: str | None = None - file_content: Any = None + file_content: object = None if isinstance(file_data, tuple): if len(file_data) == 2: @@ -1002,7 +994,7 @@ def unpack_defs( # Use iterative approach with queue to avoid recursion # Each item in queue is (node, parent_container, key/index, active_defs, ref_chain) - queue: Final[deque[tuple[Any, dict | list | None, str | int | None, dict, set]]] = deque( + queue: Final[deque[tuple[object, dict | list | None, str | int | None, dict, set]]] = deque( [(schema, None, None, root_defs, set())] ) inlined_bytes = 0 @@ -1624,7 +1616,10 @@ def is_function_call(optional_params: dict) -> bool: return False -def convert_custom_tool_format_to_chat_shape(format_obj: Mapping[str, Any]) -> Mapping[str, Any]: +_CUSTOM_GRAMMAR_FIELDS: Final = ("definition", "syntax") + + +def convert_custom_tool_format_to_chat_shape(format_obj: Mapping[str, object]) -> Mapping[str, object]: """ Responses API grammar formats are flat ({"type": "grammar", "definition", "syntax"}); Chat Completions wraps the same fields in a "grammar" object. Text formats are @@ -1632,15 +1627,11 @@ def convert_custom_tool_format_to_chat_shape(format_obj: Mapping[str, Any]) -> M """ if format_obj.get("type") != "grammar" or "grammar" in format_obj: return format_obj - grammar: Final = CustomFormatGrammarGrammar() - if "definition" in format_obj: - grammar["definition"] = format_obj["definition"] - if "syntax" in format_obj: - grammar["syntax"] = format_obj["syntax"] - return CustomFormatGrammar(type="grammar", grammar=grammar) + grammar: Final[Mapping[str, object]] = {key: format_obj[key] for key in _CUSTOM_GRAMMAR_FIELDS if key in format_obj} + return {"type": "grammar", "grammar": grammar} -def convert_custom_tool_format_to_responses_shape(format_obj: Mapping[str, Any]) -> Mapping[str, Any]: +def convert_custom_tool_format_to_responses_shape(format_obj: Mapping[str, object]) -> Mapping[str, object]: """ Inverse of convert_custom_tool_format_to_chat_shape: unwrap the Chat Completions "grammar" object into the flat Responses API grammar shape. @@ -1648,12 +1639,10 @@ def convert_custom_tool_format_to_responses_shape(format_obj: Mapping[str, Any]) grammar: Final = format_obj.get("grammar") if format_obj.get("type") != "grammar" or not isinstance(grammar, dict): return format_obj - flat: Final = ResponsesGrammarFormat(type="grammar") - if "definition" in grammar: - flat["definition"] = grammar["definition"] - if "syntax" in grammar: - flat["syntax"] = grammar["syntax"] - return flat + return { + "type": "grammar", + **{key: grammar[key] for key in _CUSTOM_GRAMMAR_FIELDS if key in grammar}, + } def get_file_ids_from_messages(messages: list[AllMessageValues]) -> list[str]: diff --git a/litellm/llms/a2a/chat/guardrail_translation/handler.py b/litellm/llms/a2a/chat/guardrail_translation/handler.py index 5c30ff4747a..92dc49ea9c1 100644 --- a/litellm/llms/a2a/chat/guardrail_translation/handler.py +++ b/litellm/llms/a2a/chat/guardrail_translation/handler.py @@ -125,7 +125,7 @@ class A2AGuardrailHandler(BaseTranslation): litellm_logging_obj: Optional["LiteLLMLoggingObj"] = None, user_api_key_dict: Optional["UserAPIKeyAuth"] = None, request_data: dict | None = None, - ) -> Any: + ) -> object: """ Process A2A output response by applying guardrails to text content. diff --git a/litellm/llms/anthropic/chat/transformation.py b/litellm/llms/anthropic/chat/transformation.py index 0f99441a115..6899c334618 100644 --- a/litellm/llms/anthropic/chat/transformation.py +++ b/litellm/llms/anthropic/chat/transformation.py @@ -6,7 +6,7 @@ from types import MappingProxyType from typing import TYPE_CHECKING, Any, Final, NoReturn, cast import httpx -from pydantic import ValidationError +from pydantic import BaseModel, ValidationError from typing_extensions import ReadOnly, TypedDict import litellm @@ -150,7 +150,7 @@ class _AnthropicToolResultBlock(TypedDict, total=False): content: ReadOnly[object] -_ENUM_TYPE_CHECKS: Final[Mapping[str, Callable[[object], bool]]] = MappingProxyType( +_ENUM_TYPE_CHECKS: Final[Mapping[object, Callable[[object], bool]]] = MappingProxyType( { "null": lambda v: v is None, "boolean": lambda v: isinstance(v, bool), @@ -163,7 +163,7 @@ _ENUM_TYPE_CHECKS: Final[Mapping[str, Callable[[object], bool]]] = MappingProxyT ) -def _enum_conflicts_with_declared_type(schema: Mapping[str, Any]) -> bool: +def _enum_conflicts_with_declared_type(schema: Mapping[str, object]) -> bool: """Whether ``schema``'s ``enum`` cannot match its declared ``type``.""" enum_values: Final = schema.get("enum") declared_type: Final = schema.get("type") @@ -658,7 +658,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): return result - def get_json_schema_from_pydantic_object(self, response_format: Any | dict | None) -> dict | None: + def get_json_schema_from_pydantic_object(self, response_format: type[BaseModel] | dict | None) -> dict | None: return type_to_response_format_param( response_format, ref_template="/$defs/{model}" ) # Relevant issue: https://github.com/BerriAI/litellm/issues/7755 @@ -1061,7 +1061,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): @staticmethod def _sanitize_tool_names_in_request( - optional_params: dict[str, Any], + optional_params: dict[str, object], ) -> tuple[dict[str, str], dict[str, str]]: """Sanitize ``optional_params['tools']`` and ``optional_params['tool_choice']`` in place so every name matches Anthropic's ``^[a-zA-Z0-9_-]{1,128}$``. @@ -1108,7 +1108,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): # so a caller reusing the same tool list/dicts across requests # doesn't see its inputs permanently rewritten (which would also # drop the original key from `forward` on the next request). - new_tools: Final[list[Any]] = [] + new_tools: Final[list[object]] = [] for t in tools: if ( isinstance(t, dict) @@ -1431,7 +1431,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): entry_type = entry.get("type") if entry_type == "compaction": - anthropic_edit: dict[str, Any] = {"type": "compact_20260112"} + anthropic_edit: dict[str, object] = {"type": "compact_20260112"} compact_threshold = entry.get("compact_threshold") # Rewrite to 'trigger' with correct nesting if threshold exists if compact_threshold is not None and isinstance(compact_threshold, (int, float)): @@ -2431,9 +2431,11 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): code_by_id: Final[dict[str, str]] = {} for tc in tool_calls: try: - args = json.loads(tc.get("function", {}).get("arguments", "{}")) + args: object = json.loads(tc.get("function", {}).get("arguments", "{}")) + if not isinstance(args, Mapping): + continue call_id = tc.get("id") - command = args.get("command", "") + command: object = args.get("command", "") if isinstance(call_id, str): code_by_id[call_id] = command if isinstance(command, str) else "" except Exception: @@ -2503,8 +2505,8 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): tool_results: Sequence[_AnthropicToolResultBlock] | None, compaction_blocks: Sequence[object] | None, tool_calls: list[ChatCompletionToolCallChunk], - ) -> dict[str, Any]: - provider_specific_fields: Final[dict[str, Any]] = { + ) -> dict[str, object]: + provider_specific_fields: Final[dict[str, object]] = { "citations": citations, "thinking_blocks": thinking_blocks, } diff --git a/litellm/llms/anthropic/common_utils.py b/litellm/llms/anthropic/common_utils.py index 87c4ec8938e..7ec15ebd1f4 100644 --- a/litellm/llms/anthropic/common_utils.py +++ b/litellm/llms/anthropic/common_utils.py @@ -7,7 +7,7 @@ import re from collections.abc import Mapping, MutableMapping, Sequence from datetime import datetime, timezone from types import MappingProxyType -from typing import Any, Final, Literal +from typing import Any, Final, Literal, TypeVar import httpx from pydantic import BaseModel, ConfigDict, TypeAdapter, ValidationError @@ -39,6 +39,8 @@ from litellm.types.llms.anthropic import ( from litellm.types.llms.openai import AllMessageValues from litellm.types.proxy.model_listing import ModelInfoResponse +_MessageT = TypeVar("_MessageT") + DROP_FORCED_TOOL_CHOICE_WARNING: Final = ( "Downgrading forced tool_choice to 'auto' for model=%s (drop_params=True): this model rejects tool_choice type " "'any'/'tool' with a 400 because thinking is always on and a forced call would skip it." @@ -1074,7 +1076,7 @@ class AnthropicModelInfo(BaseLLMModelInfo): return AnthropicTokenCounter() -def strip_advisor_blocks_from_messages(messages: list[Any], replace_with_text: bool = False) -> list[Any]: +def strip_advisor_blocks_from_messages(messages: list[_MessageT], replace_with_text: bool = False) -> list[_MessageT]: """ Remove (or replace) server_tool_use (name='advisor') and advisor_tool_result blocks from assistant message content. @@ -1181,7 +1183,7 @@ def is_anthropic_invalid_thinking_block_error(error_text: str) -> bool: return "must contain thinking" in lower -def strip_thinking_blocks_from_anthropic_messages(messages: list[Any]) -> list[Any]: +def strip_thinking_blocks_from_anthropic_messages(messages: Sequence[object]) -> list[object]: """ Return a new message list with thinking / redacted_thinking content blocks removed from each message. Used to recover from invalid thinking signatures on retry. @@ -1189,7 +1191,7 @@ def strip_thinking_blocks_from_anthropic_messages(messages: list[Any]) -> list[A Messages whose content is a list and becomes empty after stripping are omitted, since Anthropic rejects empty content arrays. """ - out: Final[list[Any]] = [] + out: Final[list[object]] = [] for m in messages: if not isinstance(m, dict): out.append(m) diff --git a/litellm/llms/anthropic/experimental_pass_through/messages/agentic_streaming_iterator.py b/litellm/llms/anthropic/experimental_pass_through/messages/agentic_streaming_iterator.py index 171f5156594..306041d9949 100644 --- a/litellm/llms/anthropic/experimental_pass_through/messages/agentic_streaming_iterator.py +++ b/litellm/llms/anthropic/experimental_pass_through/messages/agentic_streaming_iterator.py @@ -25,6 +25,9 @@ from litellm.constants import STREAM_SSE_KEEPALIVE_PING_BYTES if TYPE_CHECKING: from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj + from litellm.llms.base_llm.anthropic_messages.transformation import ( + BaseAnthropicMessagesConfig, + ) HOLD_BACK_PING_INTERVAL_SECONDS: Final = 15.0 SERVER_FULFILLED_TOOL_LEAK_ERROR_SSE_BYTES: Final = ( @@ -182,7 +185,7 @@ class AgenticAnthropicStreamingIterator: http_handler: Any, model: str, messages: list[dict], - anthropic_messages_provider_config: Any, + anthropic_messages_provider_config: "BaseAnthropicMessagesConfig", anthropic_messages_optional_request_params: dict, logging_obj: "LiteLLMLoggingObj", custom_llm_provider: str, @@ -402,7 +405,7 @@ class AgenticAnthropicStreamingIterator: @staticmethod def _rebuild_anthropic_response_from_sse( raw_bytes: list[bytes], - ) -> dict[str, Any] | None: + ) -> dict[str, object] | None: """ Parse collected SSE bytes into an Anthropic Messages response dict. @@ -416,17 +419,18 @@ class AgenticAnthropicStreamingIterator: """ events: Final = _parse_sse_events(b"".join(raw_bytes)) - response: Final[dict[str, Any]] = { + content: Final[list[dict[str, object]]] = [] + response: Final[dict[str, object]] = { "id": "", "type": "message", "role": "assistant", "model": "", - "content": [], + "content": content, "stop_reason": None, "stop_sequence": None, "usage": {"input_tokens": 0, "output_tokens": 0}, } - content_blocks: Final[dict[int, dict[str, Any]]] = {} + content_blocks: Final[dict[int, dict[str, object]]] = {} saw_message_start = False for event_type, data in events: @@ -448,6 +452,6 @@ class AgenticAnthropicStreamingIterator: for idx in sorted(content_blocks.keys()): block = content_blocks[idx] block.pop("_partial_json", None) - response["content"].append(block) + content.append(block) return response diff --git a/litellm/llms/anthropic/files/handler.py b/litellm/llms/anthropic/files/handler.py index dfd62ca575b..e4c75a704ec 100644 --- a/litellm/llms/anthropic/files/handler.py +++ b/litellm/llms/anthropic/files/handler.py @@ -185,7 +185,11 @@ class AnthropicFilesHandler: if not line.strip(): continue - anthropic_result = json.loads(line) + anthropic_result: object = json.loads(line) + if not isinstance(anthropic_result, dict): + raise TypeError( + f"Anthropic batch result line is not a JSON object: {type(anthropic_result).__name__}" + ) custom_id = anthropic_result.get("custom_id", "") result = anthropic_result.get("result", {}) result_type = result.get("type", "") diff --git a/litellm/llms/bedrock/chat/converse_transformation.py b/litellm/llms/bedrock/chat/converse_transformation.py index fa18361e44c..aabd4f1afdc 100644 --- a/litellm/llms/bedrock/chat/converse_transformation.py +++ b/litellm/llms/bedrock/chat/converse_transformation.py @@ -1037,7 +1037,7 @@ class AmazonConverseConfig(BaseConfig): return optional_params - def _map_request_metadata_param(self, value: Any, optional_params: dict) -> None: + def _map_request_metadata_param(self, value: object, optional_params: dict) -> None: if value is not None and isinstance(value, dict): self._validate_request_metadata(value) optional_params["requestMetadata"] = value diff --git a/litellm/llms/chatgpt/responses/transformation.py b/litellm/llms/chatgpt/responses/transformation.py index b96e06be3d8..9774b762396 100644 --- a/litellm/llms/chatgpt/responses/transformation.py +++ b/litellm/llms/chatgpt/responses/transformation.py @@ -1,5 +1,8 @@ +from collections.abc import Mapping from typing import TYPE_CHECKING, Any, Final +import httpx + from litellm.exceptions import AuthenticationError from litellm.litellm_core_utils.core_helpers import process_response_headers from litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response import ( @@ -13,6 +16,7 @@ from litellm.responses.sse_output_recovery import ( record_output_text_chunk, ) from litellm.types.llms.openai import ( + ResponseInputParam, ResponsesAPIResponse, ResponsesAPIStreamEvents, ) @@ -64,7 +68,7 @@ class ChatGPTResponsesAPIConfig(OpenAIResponsesAPIConfig): def transform_responses_api_request( self, model: str, - input: Any, + input: str | ResponseInputParam, response_api_optional_request_params: dict, litellm_params: GenericLiteLLMParams, headers: dict, @@ -109,9 +113,9 @@ class ChatGPTResponsesAPIConfig(OpenAIResponsesAPIConfig): def transform_response_api_response( self, model: str, - raw_response: Any, + raw_response: httpx.Response, logging_obj: "LiteLLMLoggingObj", - ): + ) -> ResponsesAPIResponse: body_text: Final = raw_response.text or "" if not self._should_parse_as_sse(raw_response=raw_response, body_text=body_text): return super().transform_response_api_response( @@ -135,7 +139,7 @@ class ChatGPTResponsesAPIConfig(OpenAIResponsesAPIConfig): self._attach_response_headers(completed_response=completed_response, raw_response=raw_response) return completed_response - def _should_parse_as_sse(self, raw_response: Any, body_text: str) -> bool: + def _should_parse_as_sse(self, raw_response: httpx.Response, body_text: str) -> bool: content_type: Final = (raw_response.headers or {}).get("content-type", "") if "text/event-stream" in content_type.lower(): return True @@ -150,8 +154,8 @@ class ChatGPTResponsesAPIConfig(OpenAIResponsesAPIConfig): def _extract_completed_response_from_sse(self, body_text: str) -> tuple[ResponsesAPIResponse | None, str | None]: completed_response = None error_message = None - streamed_output_items: Final[dict[int, dict]] = {} - text_only_output_items: Final[dict[int, dict]] = {} + streamed_output_items: Final[dict[int, dict[str, object]]] = {} + text_only_output_items: Final[dict[int, dict[str, object]]] = {} for chunk in body_text.splitlines(): parsed_chunk = parse_sse_json_chunk(chunk) if parsed_chunk is None: @@ -178,7 +182,7 @@ class ChatGPTResponsesAPIConfig(OpenAIResponsesAPIConfig): # output_index, but text-only items at indices without a # matching OUTPUT_ITEM_DONE must still be preserved (e.g. # providers that emit only OUTPUT_TEXT_DONE for some indices). - merged_items: dict[int, dict] = {**text_only_output_items} + merged_items: dict[int, dict[str, object]] = {**text_only_output_items} merged_items.update(streamed_output_items) completed_response = self._build_completed_response_from_chunk( parsed_chunk=parsed_chunk, @@ -197,7 +201,7 @@ class ChatGPTResponsesAPIConfig(OpenAIResponsesAPIConfig): return completed_response, error_message def _build_completed_response_from_chunk( - self, parsed_chunk: dict[str, Any], streamed_output_items: dict[int, dict] + self, parsed_chunk: Mapping[str, object], streamed_output_items: Mapping[int, dict[str, object]] ) -> ResponsesAPIResponse | None: response_payload = parsed_chunk.get("response") if not isinstance(response_payload, dict): @@ -223,7 +227,7 @@ class ChatGPTResponsesAPIConfig(OpenAIResponsesAPIConfig): def _attach_response_headers( self, completed_response: ResponsesAPIResponse, - raw_response: Any, + raw_response: httpx.Response, ) -> None: raw_headers: Final = dict(raw_response.headers) processed_headers: Final = process_response_headers(raw_headers) diff --git a/litellm/llms/cohere/chat/transformation.py b/litellm/llms/cohere/chat/transformation.py index 319603b0dad..fa46bd7f6cf 100644 --- a/litellm/llms/cohere/chat/transformation.py +++ b/litellm/llms/cohere/chat/transformation.py @@ -110,7 +110,7 @@ class CohereChatConfig(BaseConfig): tool_results: list | None = None, seed: int | None = None, ) -> None: - locals_: Final = locals().copy() + locals_: Final[dict[str, object]] = locals().copy() for key, value in locals_.items(): if key != "self" and value is not None: setattr(self.__class__, key, value) diff --git a/litellm/llms/databricks/chat/transformation.py b/litellm/llms/databricks/chat/transformation.py index 82c3b5d91d3..dd257cd68b0 100644 --- a/litellm/llms/databricks/chat/transformation.py +++ b/litellm/llms/databricks/chat/transformation.py @@ -3,7 +3,7 @@ Translates from OpenAI's `/v1/chat/completions` to Databricks' `/chat/completion """ import os -from collections.abc import AsyncIterator, Coroutine, Iterator, Mapping +from collections.abc import AsyncIterator, Coroutine, Iterator, Mapping, Sequence from typing import TYPE_CHECKING, Any, Final, Literal, cast, overload import httpx @@ -67,7 +67,7 @@ def _is_bare_assistant_message(message_dict: Mapping[str, object]) -> bool: ) -def _sanitize_empty_content(message_dict: dict[str, Any]) -> None: +def _sanitize_empty_content(message_dict: dict[str, object]) -> None: """ Remove or filter content so empty text blocks are not sent. Databricks Model Serving uses Anthropic Messages API spec and rejects empty text blocks. @@ -430,7 +430,7 @@ class DatabricksConfig(DatabricksBase, OpenAILikeChatConfig, AnthropicConfig): @overload def _transform_messages( self, messages: list[AllMessageValues], model: str, is_async: Literal[True] - ) -> Coroutine[Any, Any, list[AllMessageValues]]: ... + ) -> Coroutine[object, object, list[AllMessageValues]]: ... @overload def _transform_messages( @@ -442,7 +442,7 @@ class DatabricksConfig(DatabricksBase, OpenAILikeChatConfig, AnthropicConfig): def _transform_messages( self, messages: list[AllMessageValues], model: str, is_async: bool = False - ) -> list[AllMessageValues] | Coroutine[Any, Any, list[AllMessageValues]]: + ) -> list[AllMessageValues] | Coroutine[object, object, list[AllMessageValues]]: """ Databricks does not support: - 'name' in user message. @@ -564,7 +564,7 @@ class DatabricksConfig(DatabricksBase, OpenAILikeChatConfig, AnthropicConfig): @staticmethod def extract_citations( content: AllDatabricksContentValues | None, - ) -> list[Any] | None: + ) -> Sequence[Sequence[Mapping[str, object]]] | None: if content is None: return None citations: Final = [] diff --git a/litellm/llms/fireworks_ai/chat/transformation.py b/litellm/llms/fireworks_ai/chat/transformation.py index 05160d83c12..1e9082a1ef1 100644 --- a/litellm/llms/fireworks_ai/chat/transformation.py +++ b/litellm/llms/fireworks_ai/chat/transformation.py @@ -1,6 +1,6 @@ import json from collections.abc import AsyncIterator, Iterator, Mapping -from typing import TYPE_CHECKING, Any, Final, Literal, cast +from typing import TYPE_CHECKING, Final, Literal, cast import httpx @@ -751,7 +751,7 @@ class FireworksAIConfig(FireworksAIMixin, OpenAIGPTConfig): streaming_response: Iterator[str] | AsyncIterator[str] | ModelResponse, sync_stream: bool, json_mode: bool | None = False, - ) -> Any: + ) -> "FireworksAIChatCompletionStreamingHandler": return FireworksAIChatCompletionStreamingHandler( streaming_response=streaming_response, sync_stream=sync_stream, diff --git a/litellm/llms/gemini/count_tokens/handler.py b/litellm/llms/gemini/count_tokens/handler.py index cb2be2c860e..c2f0ef473ae 100644 --- a/litellm/llms/gemini/count_tokens/handler.py +++ b/litellm/llms/gemini/count_tokens/handler.py @@ -84,7 +84,7 @@ class GoogleAIStudioTokenCounter: api_key: str | None = None, api_base: str | None = None, timeout: float | httpx.Timeout | None = None, - **kwargs, + **kwargs: object, ) -> dict[str, Any]: """ Count tokens using Google Gen AI Studio countTokens endpoint. diff --git a/litellm/llms/gemini/image_edit/transformation.py b/litellm/llms/gemini/image_edit/transformation.py index e6c22dc60b4..31d3963c70c 100644 --- a/litellm/llms/gemini/image_edit/transformation.py +++ b/litellm/llms/gemini/image_edit/transformation.py @@ -1,4 +1,5 @@ import base64 +from collections.abc import Mapping from io import BufferedReader, BytesIO from typing import TYPE_CHECKING, Any, Final, cast @@ -44,7 +45,7 @@ class GeminiImageEditConfig(BaseImageEditConfig): image_edit_optional_params: ImageEditOptionalRequestParams, model: str, drop_params: bool, - ) -> dict[str, Any]: + ) -> dict[str, object]: return map_openai_image_params_to_gemini( params=image_edit_optional_params, model=model, @@ -87,10 +88,10 @@ class GeminiImageEditConfig(BaseImageEditConfig): model: str, prompt: str | None, image: FileTypes | None, - image_edit_optional_request_params: dict[str, Any], + image_edit_optional_request_params: Mapping[str, object], litellm_params: GenericLiteLLMParams, headers: dict, - ) -> tuple[dict[str, Any], RequestFiles | None]: + ) -> tuple[dict[str, object], RequestFiles | None]: inline_parts: Final = self._prepare_inline_image_parts(image) if image else [] if not inline_parts: raise ValueError("Gemini image edit requires at least one image.") @@ -106,7 +107,7 @@ class GeminiImageEditConfig(BaseImageEditConfig): } ] - request_body: Final[dict[str, Any]] = {"contents": contents} + request_body: Final[dict[str, object]] = {"contents": contents} request_body["generationConfig"] = get_gemini_image_generation_config( model=model, @@ -153,14 +154,14 @@ class GeminiImageEditConfig(BaseImageEditConfig): model_response.usage = transform_gemini_image_usage(response_json["usageMetadata"]) return model_response - def _prepare_inline_image_parts(self, image: FileTypes | list[FileTypes]) -> list[dict[str, Any]]: + def _prepare_inline_image_parts(self, image: FileTypes | list[FileTypes]) -> list[dict[str, object]]: images: list[FileTypes] if isinstance(image, list): images = image else: images = [image] - inline_parts: Final[list[dict[str, Any]]] = [] + inline_parts: Final[list[dict[str, object]]] = [] for img in images: if img is None: continue diff --git a/litellm/llms/gigachat/chat/transformation.py b/litellm/llms/gigachat/chat/transformation.py index 89920ebd27b..d9250ea8836 100644 --- a/litellm/llms/gigachat/chat/transformation.py +++ b/litellm/llms/gigachat/chat/transformation.py @@ -81,9 +81,17 @@ class GigaChatConfig(BaseConfig): repetition_penalty: float | None = None, profanity_check: bool | None = None, ) -> None: - locals_: Final = locals().copy() - for key, value in locals_.items(): - if key != "self" and value is not None: + config_params: Final[Mapping[str, float | int | bool | None]] = MappingProxyType( + { + "temperature": temperature, + "top_p": top_p, + "max_tokens": max_tokens, + "repetition_penalty": repetition_penalty, + "profanity_check": profanity_check, + } + ) + for key, value in config_params.items(): + if value is not None: setattr(self.__class__, key, value) # Instance variables for current request context self._current_credentials: str | None = None diff --git a/litellm/llms/huggingface/embedding/transformation.py b/litellm/llms/huggingface/embedding/transformation.py index f6fe7f2fa10..33b0e21e326 100644 --- a/litellm/llms/huggingface/embedding/transformation.py +++ b/litellm/llms/huggingface/embedding/transformation.py @@ -84,13 +84,13 @@ class HuggingFaceEmbeddingConfig(BaseConfig): typical_p: float | None = None, watermark: bool | None = None, ) -> None: - locals_: Final = locals().copy() + locals_: Final[dict[str, object]] = locals().copy() for key, value in locals_.items(): if key != "self" and value is not None: setattr(self.__class__, key, value) @classmethod - def get_config(cls): + def get_config(cls) -> dict[str, object]: return super().get_config() def get_special_options_params(self): @@ -352,17 +352,17 @@ class HuggingFaceEmbeddingConfig(BaseConfig): model: str, data: dict, api_key: str | None = None, - ) -> list[dict[str, Any]]: + ) -> list[dict[str, str]]: streamed_response: Final = CustomStreamWrapper( completion_stream=response.iter_lines(), model=model, custom_llm_provider="huggingface", logging_obj=logging_obj, ) - content = "" + content: str = "" for chunk in streamed_response: content += chunk["choices"][0]["delta"]["content"] - completion_response: Final[list[dict[str, Any]]] = [{"generated_text": content}] + completion_response: Final[list[dict[str, str]]] = [{"generated_text": content}] ## LOGGING logging_obj.post_call( input=data, diff --git a/litellm/llms/openai/chat/guardrail_translation/handler.py b/litellm/llms/openai/chat/guardrail_translation/handler.py index 58ff03e6a0d..665215303d9 100644 --- a/litellm/llms/openai/chat/guardrail_translation/handler.py +++ b/litellm/llms/openai/chat/guardrail_translation/handler.py @@ -17,7 +17,7 @@ This pattern can be replicated for other message formats (e.g., Anthropic). import json import time import uuid -from collections.abc import Sequence +from collections.abc import Mapping, Sequence from typing import TYPE_CHECKING, Any, Final, Union, cast from typing_extensions import NotRequired, ReadOnly, TypedDict @@ -232,7 +232,7 @@ class OpenAIChatCompletionsHandler(BaseTranslation): def _extract_inputs( self, - message: dict[str, Any], + message: Mapping[str, object], msg_idx: int, texts_to_check: list[str], images_to_check: list[str], @@ -293,7 +293,7 @@ class OpenAIChatCompletionsHandler(BaseTranslation): async def _apply_guardrail_responses_to_input_texts( self, - messages: list[dict[str, Any]], + messages: list[dict[str, object]], responses: list[str], task_mappings: list[tuple[int, int | None]], ) -> None: @@ -318,12 +318,12 @@ class OpenAIChatCompletionsHandler(BaseTranslation): elif isinstance(content, list) and content_idx_optional is not None: # Replace specific text item in list content - messages[msg_idx]["content"][content_idx_optional]["text"] = guardrail_response + content[content_idx_optional]["text"] = guardrail_response async def _apply_guardrail_responses_to_input_tool_calls( self, - messages: list[dict[str, Any]], - tool_calls: list[dict[str, Any]], + messages: Sequence[Mapping[str, object]], + tool_calls: Sequence[Mapping[str, object]], task_mappings: list[tuple[int, int]], ) -> None: """ @@ -375,7 +375,7 @@ class OpenAIChatCompletionsHandler(BaseTranslation): texts_to_check: Final[list[str]] = [] images_to_check: Final[list[str]] = [] - tool_calls_to_check: Final[list[dict[str, Any]]] = [] + tool_calls_to_check: Final[list[dict[str, object]]] = [] text_task_mappings: Final[list[tuple[int, int | None]]] = [] tool_call_task_mappings: Final[list[tuple[int, int]]] = [] # text_task_mappings: Track (choice_index, content_index) for each text @@ -424,8 +424,8 @@ class OpenAIChatCompletionsHandler(BaseTranslation): guardrailed_texts: Final = guardrailed_inputs.get("texts", []) returned_tool_calls: Final = guardrailed_inputs.get("tool_calls") - guardrailed_tool_calls: Final[list[dict[str, Any]]] = ( - cast(list[dict[str, Any]], returned_tool_calls) + guardrailed_tool_calls: Final[list[dict[str, object]]] = ( + cast(list[dict[str, object]], returned_tool_calls) if isinstance(returned_tool_calls, list) and len(returned_tool_calls) == len(tool_calls_to_check) else tool_calls_to_check ) @@ -864,7 +864,7 @@ class OpenAIChatCompletionsHandler(BaseTranslation): choice_idx: int, texts_to_check: list[str], images_to_check: list[str], - tool_calls_to_check: list[dict[str, Any]], + tool_calls_to_check: list[dict[str, object]], text_task_mappings: list[tuple[int, int | None]], tool_call_task_mappings: list[tuple[int, int]], ) -> None: diff --git a/litellm/llms/openai/videos/transformation.py b/litellm/llms/openai/videos/transformation.py index 94dc30f41e5..9a4b030993f 100644 --- a/litellm/llms/openai/videos/transformation.py +++ b/litellm/llms/openai/videos/transformation.py @@ -237,7 +237,7 @@ class OpenAIVideoConfig(BaseVideoConfig): api_base: str, litellm_params: GenericLiteLLMParams, headers: dict, - extra_body: dict[str, Any] | None = None, + extra_body: dict[str, object] | None = None, ) -> tuple[str, dict]: """ Transform the video remix request for OpenAI API. @@ -252,7 +252,7 @@ class OpenAIVideoConfig(BaseVideoConfig): url: Final = f"{api_base.rstrip('/')}/{encoded_video_id}/remix" # Prepare the request data - data: Final = {"prompt": prompt} + data: Final[dict[str, object]] = {"prompt": prompt} # Add any extra body parameters if extra_body: @@ -305,7 +305,7 @@ class OpenAIVideoConfig(BaseVideoConfig): after: str | None = None, limit: int | None = None, order: str | None = None, - extra_query: dict[str, Any] | None = None, + extra_query: dict[str, object] | None = None, ) -> tuple[str, dict]: """ Transform the video list request for OpenAI API. diff --git a/litellm/llms/openrouter/image_edit/transformation.py b/litellm/llms/openrouter/image_edit/transformation.py index b01c25aad0c..3d46277a69e 100644 --- a/litellm/llms/openrouter/image_edit/transformation.py +++ b/litellm/llms/openrouter/image_edit/transformation.py @@ -90,20 +90,21 @@ class OpenRouterImageEditConfig(BaseImageEditConfig): drop_params: bool, ) -> dict: supported_params: Final = self.get_supported_openai_params(model) - mapped_params: Final[dict[str, Any]] = {} + mapped_params: Final[dict[str, object]] = {} + image_config: Final[dict[str, str]] = {} for key, value in image_edit_optional_params.items(): if key in supported_params: if key == "size": if "image_config" not in mapped_params: - mapped_params["image_config"] = {} - mapped_params["image_config"]["aspect_ratio"] = self._map_size_to_aspect_ratio(cast(str, value)) + mapped_params["image_config"] = image_config + image_config["aspect_ratio"] = self._map_size_to_aspect_ratio(cast(str, value)) elif key == "quality": image_size = self._map_quality_to_image_size(cast(str, value)) if image_size: if "image_config" not in mapped_params: - mapped_params["image_config"] = {} - mapped_params["image_config"]["image_size"] = image_size + mapped_params["image_config"] = image_config + image_config["image_size"] = image_size else: mapped_params[key] = value diff --git a/litellm/llms/perplexity/embedding/transformation.py b/litellm/llms/perplexity/embedding/transformation.py index a911fa62719..c93206db2bb 100644 --- a/litellm/llms/perplexity/embedding/transformation.py +++ b/litellm/llms/perplexity/embedding/transformation.py @@ -130,7 +130,7 @@ class PerplexityEmbeddingConfig(BaseEmbeddingConfig): if isinstance(embedding_value, str): raw_bytes: Final = base64.b64decode(embedding_value) count: Final = len(raw_bytes) - int8_values: Final = struct.unpack(f"{count}b", raw_bytes) + int8_values: Final[tuple[int, ...]] = struct.unpack(f"{count}b", raw_bytes) return [float(v) / 127.0 for v in int8_values] return embedding_value diff --git a/litellm/llms/vertex_ai/gemini/transformation.py b/litellm/llms/vertex_ai/gemini/transformation.py index 13e2238fdf6..e3cc3bbb2dc 100644 --- a/litellm/llms/vertex_ai/gemini/transformation.py +++ b/litellm/llms/vertex_ai/gemini/transformation.py @@ -179,7 +179,7 @@ def _apply_gemini_metadata( part: PartType, model: str | None, media_resolution_enum: dict[str, str] | None, - video_metadata: dict[str, Any] | None, + video_metadata: Mapping[str, object] | None, ) -> PartType: """ Apply media_resolution and video_metadata parameters to a Gemini part. @@ -480,7 +480,7 @@ def _process_gemini_media( format: str | None = None, media_resolution_enum: dict[str, str] | None = None, model: str | None = None, - video_metadata: dict[str, Any] | None = None, + video_metadata: Mapping[str, object] | None = None, vertex_project: str | None = None, vertex_credentials: object = None, ) -> PartType: diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index fb0c623473a..f640db7e5dc 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -5574,7 +5574,7 @@ class MCPServerManager: async def pre_call_tool_check( self, name: str, - arguments: dict[str, Any], + arguments: _ToolArguments, server_name: str, user_api_key_auth: UserAPIKeyAuth | None, proxy_logging_obj: ProxyLogging | None, diff --git a/litellm/proxy/auth/auth_utils.py b/litellm/proxy/auth/auth_utils.py index be65c3b39ec..bd248da52ef 100644 --- a/litellm/proxy/auth/auth_utils.py +++ b/litellm/proxy/auth/auth_utils.py @@ -166,7 +166,7 @@ def check_regex_or_str_match(request_body_value: Any, regex_str: str) -> bool: def _is_param_allowed( param: str, - request_body_value: Any, + request_body_value: object, configurable_clientside_auth_params: CONFIGURABLE_CLIENTSIDE_AUTH_PARAMS, ) -> bool: """ @@ -189,7 +189,7 @@ def _is_param_allowed( def _allow_model_level_clientside_configurable_parameters( - model: str, param: str, request_body_value: Any, llm_router: Router | None + model: str, param: str, request_body_value: object, llm_router: Router | None ) -> bool: """ Check if model is allowed to use configurable client-side params @@ -532,7 +532,7 @@ def is_request_body_safe(request_body: dict, general_settings: dict, llm_router: return True -def _coerce_metadata_to_dict(value: Any) -> dict[str, Any] | None: +def _coerce_metadata_to_dict(value: object) -> dict[str, object] | None: """Return ``value`` as a dict, parsing it from JSON if delivered as a string. Multipart/form-data and ``extra_body`` callers send ``litellm_metadata`` @@ -891,7 +891,7 @@ async def check_if_request_size_is_safe(request: Request) -> bool: return True -async def check_response_size_is_safe(response: Any) -> bool: +async def check_response_size_is_safe(response: object) -> bool: """ Enterprise Only: - Checks if the response size is within the limit @@ -1526,7 +1526,7 @@ def get_customer_user_header_from_mapping(user_id_mapping) -> list | None: def _get_customer_id_from_standard_headers( - request_headers: dict | None, + request_headers: Mapping[str, object] | None, ) -> str | None: """ Check standard customer ID headers for a customer/end-user ID. @@ -1552,7 +1552,7 @@ def _get_customer_id_from_standard_headers( return None -def _coerce_user_id_to_str(value: Any) -> str | None: +def _coerce_user_id_to_str(value: object) -> str | None: """Return a usable end-user identifier string, or None if the value isn't one. Always drops non-string structured values (dict/list/tuple/set) because @@ -1579,7 +1579,7 @@ def _coerce_user_id_to_str(value: Any) -> str | None: # behind the flag preserves backwards compatibility for deployments # that intentionally pass JSON-encoded user identifiers. if litellm.validate_end_user_id_in_db and stripped[:1] in ("{", "["): - parsed: Final = safe_json_loads(stripped) + parsed: Final[object] = safe_json_loads(stripped) if isinstance(parsed, (dict, list)): return None return stripped @@ -1587,7 +1587,9 @@ def _coerce_user_id_to_str(value: Any) -> str | None: return None -def get_end_user_id_from_request_body(request_body: dict, request_headers: dict | None = None) -> str | None: +def get_end_user_id_from_request_body( + request_body: Mapping[str, object], request_headers: Mapping[str, object] | None = None +) -> str | None: # Import general_settings here to avoid potential circular import issues at module level # and to ensure it's fetched at runtime. from litellm.proxy.proxy_server import general_settings @@ -1636,7 +1638,7 @@ def get_end_user_id_from_request_body(request_body: dict, request_headers: dict if user_id_str: return user_id_str - def _as_dict(value: Any) -> dict: + def _as_dict(value: object) -> dict: # metadata / litellm_metadata can arrive as JSON strings from # multipart/form-data or extra_body; coerce so string-encoded # payloads can't evade end-user attribution. @@ -1721,11 +1723,11 @@ _MODEL_ROUTING_ID_FIELDS: Final = ( ) -def _append_model_candidates(candidates: list[str], value: Any) -> None: +def _append_model_candidates(candidates: list[str], value: object) -> None: if value is None: return - values: Final = value if isinstance(value, (list, tuple, set)) else [value] + values: Final[tuple[object, ...]] = tuple(value) if isinstance(value, (list, tuple, set)) else (value,) for item in values: if item is None: continue @@ -1766,7 +1768,7 @@ def _route_uses_model_routing_sources(route: str) -> bool: def _extract_models_from_managed_resource_id( - resource_id: Any, + resource_id: object, resource_id_field: str | None = None, llm_router: Router | None = None, ) -> list[str]: diff --git a/litellm/proxy/client/cli/commands/configure.py b/litellm/proxy/client/cli/commands/configure.py index 7988f8aef3c..539aa2581f5 100644 --- a/litellm/proxy/client/cli/commands/configure.py +++ b/litellm/proxy/client/cli/commands/configure.py @@ -98,13 +98,15 @@ def _preflight(target: str) -> None: raise click.ClickException(str(e)) from e -def _start(ctx: click.Context, api_key: str | None, target: str = _CLAUDE_TARGET) -> tuple[StaticToken, _Listing]: +def _start( + ctx: click.Context, base_url: str, api_key: str | None, target: str = _CLAUDE_TARGET +) -> tuple[StaticToken, _Listing]: _preflight(target) try: credential: Final = resolve_credential(ctx, api_key) except ClaudeSettingsError as e: raise click.ClickException(str(e)) - return credential, _listed_models(ctx.obj["base_url"], credential.token, target) + return credential, _listed_models(base_url, credential.token, target) def _listing_error(base_url: str, error: PiSyncError, target: str) -> str: @@ -147,9 +149,7 @@ def _validated_model(model: str | None, listing: _Listing, base_url: str) -> str return starting -def _apply_claude(ctx: click.Context, credential: StaticToken, listing: _Listing, model: str | None) -> None: - ctx_obj: Final[CliContextObj] = ctx.obj - base_url: Final = ctx_obj["base_url"] +def _apply_claude(base_url: str, credential: StaticToken, listing: _Listing, model: str | None) -> None: listed: Final = listing.ids starting: Final = _validated_model(model, listing, base_url) settings_path: Final = claude_settings_path(os.environ) @@ -214,8 +214,7 @@ def _pick_codex_model(listed: Sequence[str]) -> str: return str(inquirer.fuzzy(message="Model Codex starts on (type to filter):", choices=choices).execute()) -def _apply_codex(ctx: click.Context, credential: StaticToken, listing: _Listing, model: str) -> None: - base_url: Final[str] = ctx.obj["base_url"] +def _apply_codex(base_url: str, credential: StaticToken, listing: _Listing, model: str) -> None: _validated_model(model, listing, base_url) settings_path: Final = codex_config_path(os.environ) try: @@ -237,13 +236,12 @@ class _Setup: def _choose_setup( - ctx: click.Context, + base_url: str, target: str, credential: StaticToken, pick_model: Callable[[Sequence[str]], str | None], pick_codex_model: Callable[[Sequence[str]], str], ) -> _Setup: - base_url: Final[str] = ctx.obj["base_url"] listing: Final = _listed_models(base_url, credential.token, target) model: Final = ( pick_model(tuple(item.source_model or item.id for item in listing.models)) @@ -270,12 +268,15 @@ def interactive_configure( credential: Final = resolve_credential(ctx, None) except ClaudeSettingsError as e: raise click.ClickException(str(e)) from e - setups: Final = tuple(_choose_setup(ctx, target, credential, pick_model, pick_codex_model) for target in targets) + base_url: Final[str] = ctx.obj["base_url"] + setups: Final = tuple( + _choose_setup(base_url, target, credential, pick_model, pick_codex_model) for target in targets + ) for setup in setups: if setup.target == _CLAUDE_TARGET: - _apply_claude(ctx, credential, setup.listing, setup.model) + _apply_claude(base_url, credential, setup.listing, setup.model) elif setup.model is not None: - _apply_codex(ctx, credential, setup.listing, setup.model) + _apply_codex(base_url, credential, setup.listing, setup.model) class _ConnectionOptions(BaseModel): @@ -283,7 +284,8 @@ class _ConnectionOptions(BaseModel): gateway_url: str | None = None -def _connection_context(ctx: click.Context, api_key: str | None, gateway_url: str | None) -> click.Context: +def _connection_settings(ctx: click.Context, api_key: str | None, gateway_url: str | None) -> CliContextObj: + """The context object a subcommand runs with: its own --api-key / --gateway-url over the group's, over `lite`'s.""" ctx_obj: Final[CliContextObj] = ctx.obj group: Final = ( _ConnectionOptions.model_validate(ctx.parent.params) @@ -300,7 +302,11 @@ def _connection_context(ctx: click.Context, api_key: str | None, gateway_url: st "api_key": key if key is not None else ctx_obj.get("api_key"), "api_key_from_token_file": False if key is not None else ctx_obj.get("api_key_from_token_file", False), } - return click.Context(ctx.command, parent=ctx.parent, obj=connection) + return connection + + +def _connection_context(ctx: click.Context, settings: CliContextObj) -> click.Context: + return click.Context(ctx.command, parent=ctx.parent, obj=settings) @click.group(name="configure", invoke_without_command=True) @@ -316,19 +322,19 @@ def configure_group(ctx: click.Context, api_key: str | None, gateway_url: str | """ if ctx.invoked_subcommand is not None: return - connection: Final = _connection_context(ctx, api_key, gateway_url) + settings: Final = _connection_settings(ctx, api_key, gateway_url) + connection: Final = _connection_context(ctx, settings) if not sys.stdin.isatty(): raise click.ClickException( "`lite configure` asks questions, so it needs a terminal. Non-interactively, run " "`lite configure claude --api-key --model ` or " "`lite configure codex --api-key --model `." ) - prompted: Final = ( - connection - if connection.obj.get("base_url_explicit") - else _connection_context(connection, None, click.prompt("Gateway URL", default=connection.obj["base_url"])) - ) - interactive_configure(prompted) + if settings.get("base_url_explicit"): + interactive_configure(connection) + return + prompted: Final = _connection_settings(connection, None, click.prompt("Gateway URL", default=settings["base_url"])) + interactive_configure(_connection_context(connection, prompted)) @click.group(name="unconfigure") @@ -356,9 +362,9 @@ def configure_claude(ctx: click.Context, api_key: str | None, model: str | None, setting is kept, and what changed is recorded so `lite unconfigure claude` can put it back. Assumes the proxy is already running. """ - connection: Final = _connection_context(ctx, api_key, gateway_url) - credential, listing = _start(connection, api_key) - _apply_claude(connection, credential, listing, model) + settings: Final = _connection_settings(ctx, api_key, gateway_url) + credential, listing = _start(_connection_context(ctx, settings), settings["base_url"], api_key) + _apply_claude(settings["base_url"], credential, listing, model) @configure_group.command(name="codex") @@ -368,9 +374,9 @@ def configure_claude(ctx: click.Context, api_key: str | None, model: str | None, @click.pass_context def configure_codex(ctx: click.Context, api_key: str | None, gateway_url: str | None, model: str) -> None: """Route plain `codex` through the gateway until `lite unconfigure codex`.""" - connection: Final = _connection_context(ctx, api_key, gateway_url) - credential, listing = _start(connection, api_key, _CODEX_TARGET) - _apply_codex(connection, credential, listing, model) + settings: Final = _connection_settings(ctx, api_key, gateway_url) + credential, listing = _start(_connection_context(ctx, settings), settings["base_url"], api_key, _CODEX_TARGET) + _apply_codex(settings["base_url"], credential, listing, model) @unconfigure_group.command(name="codex") diff --git a/litellm/proxy/common_utils/callback_utils.py b/litellm/proxy/common_utils/callback_utils.py index 561a53409f4..eafbb1ef95f 100644 --- a/litellm/proxy/common_utils/callback_utils.py +++ b/litellm/proxy/common_utils/callback_utils.py @@ -713,7 +713,7 @@ def strip_callback_config(metadata: dict[str, object] | None) -> dict[str, objec return {k: v for k, v in metadata.items() if k not in _CALLBACK_CONFIG_SLOTS} -def encrypt_callback_vars(metadata: Any) -> Any: +def encrypt_callback_vars(metadata: object) -> Any: """Return a deep copy of metadata with callback_vars values encrypted at rest. Idempotent: a value that already decrypts cleanly is left unchanged so @@ -722,7 +722,7 @@ def encrypt_callback_vars(metadata: Any) -> Any: return _transform_callback_vars(metadata, _encrypt_if_plaintext) -def decrypt_callback_vars(metadata: Any) -> Any: +def decrypt_callback_vars(metadata: object) -> Any: """Return a deep copy of metadata with callback_vars values decrypted. Legacy plaintext rows pass through unchanged (decrypt failure → original). @@ -730,7 +730,7 @@ def decrypt_callback_vars(metadata: Any) -> Any: return _transform_callback_vars(metadata, _decrypt_or_passthrough) -def _transform_callback_vars(metadata: object, transform: Callable[[str, Any], Any]) -> object: +def _transform_callback_vars(metadata: object, transform: Callable[[str, object], object]) -> object: if not isinstance(metadata, dict): return metadata out: Final = copy.deepcopy(metadata) diff --git a/litellm/proxy/db/db_spend_update_writer.py b/litellm/proxy/db/db_spend_update_writer.py index eaa03c5d7f7..19af995932a 100644 --- a/litellm/proxy/db/db_spend_update_writer.py +++ b/litellm/proxy/db/db_spend_update_writer.py @@ -15,7 +15,7 @@ import traceback from collections.abc import Mapping, Sequence from datetime import datetime, timedelta, timezone from types import MappingProxyType -from typing import TYPE_CHECKING, Any, Final, Literal, Protocol, cast, overload +from typing import TYPE_CHECKING, Any, Final, Literal, Protocol, TypeAlias, cast, overload import litellm from litellm._logging import verbose_proxy_logger @@ -115,6 +115,20 @@ class _SpendBatch(Protocol): litellm_modelaccessgroupbudgettable: BatchTable +_EntitySpendTable: TypeAlias = Literal["litellm_tagtable", "litellm_agentstable", "litellm_modelaccessgroupbudgettable"] + + +def _entity_spend_table(batcher: _SpendBatch, table_accessor: _EntitySpendTable) -> BatchTable: + """The batch table an entity type's spend increments are written to.""" + match table_accessor: + case "litellm_tagtable": + return batcher.litellm_tagtable + case "litellm_agentstable": + return batcher.litellm_agentstable + case "litellm_modelaccessgroupbudgettable": + return batcher.litellm_modelaccessgroupbudgettable + + class _SpendBatchManager(Protocol): async def __aenter__(self) -> _SpendBatch: ... @@ -1750,7 +1764,7 @@ class DBSpendUpdateWriter: async def _update_entity_spend_in_db( entity_name: str, transactions: dict[str, float] | None, - table_accessor: Literal["litellm_tagtable", "litellm_agentstable", "litellm_modelaccessgroupbudgettable"], + table_accessor: _EntitySpendTable, where_field: str, n_retry_times: int, prisma_client: PrismaClient, @@ -1784,7 +1798,7 @@ class DBSpendUpdateWriter: entity_id, response_cost, ) - getattr(batcher, table_accessor).update_many( + _entity_spend_table(batcher, table_accessor).update_many( where={where_field: entity_id}, data={"spend": {"increment": response_cost}}, ) diff --git a/litellm/proxy/guardrails/guardrail_hooks/akto/akto.py b/litellm/proxy/guardrails/guardrail_hooks/akto/akto.py index 2c27531cea1..a7f45a37ae6 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/akto/akto.py +++ b/litellm/proxy/guardrails/guardrail_hooks/akto/akto.py @@ -11,10 +11,11 @@ import asyncio import json import os from datetime import datetime -from typing import TYPE_CHECKING, Any, Final, Literal +from typing import TYPE_CHECKING, Final, Literal import httpx from fastapi import HTTPException +from typing_extensions import NotRequired, ReadOnly, TypedDict, Unpack from litellm._logging import verbose_proxy_logger from litellm.integrations.custom_guardrail import ( @@ -25,12 +26,34 @@ from litellm.llms.custom_httpx.http_handler import ( get_async_httpx_client, httpxSpecialProvider, ) -from litellm.types.guardrails import GuardrailEventHooks +from litellm.types.guardrails import GuardrailEventHooks, Mode from litellm.types.utils import GenericGuardrailAPIInputs if TYPE_CHECKING: from litellm.types.proxy.guardrails.guardrail_hooks.base import GuardrailConfigModel + +class _CustomGuardrailKwargs(TypedDict): + """Keyword arguments forwarded verbatim to CustomGuardrail.__init__.""" + + guardrail_name: NotRequired[ReadOnly[str | None]] + event_hook: NotRequired[ReadOnly[GuardrailEventHooks | list[GuardrailEventHooks] | Mode | None]] + default_on: NotRequired[ReadOnly[bool]] + mask_request_content: NotRequired[ReadOnly[bool]] + mask_response_content: NotRequired[ReadOnly[bool]] + violation_message_template: NotRequired[ReadOnly[str | None]] + end_session_after_n_fails: NotRequired[ReadOnly[int | None]] + on_violation: NotRequired[ReadOnly[str | None]] + realtime_violation_message: NotRequired[ReadOnly[str | None]] + on_sensitive_data: NotRequired[ReadOnly[str | None]] + sensitive_data_route_to_model: NotRequired[ReadOnly[str | None]] + sticky_session_routing: NotRequired[ReadOnly[bool]] + run_in_parallel: NotRequired[ReadOnly[bool]] + scan_raw_request: NotRequired[ReadOnly[bool]] + only_scan_new_messages: NotRequired[ReadOnly[bool]] + supported_event_hooks: NotRequired[ReadOnly[list[GuardrailEventHooks]]] + + HTTP_PROXY_PATH: Final = "/api/http-proxy" AKTO_CONNECTOR_NAME: Final = "litellm" DEFAULT_GUARDRAIL_TIMEOUT: Final = 5 @@ -66,7 +89,7 @@ class AktoGuardrail(CustomGuardrail): akto_vxlan_id: str | None = None, unreachable_fallback: Literal["fail_closed", "fail_open"] = "fail_closed", guardrail_timeout: int | None = None, - **kwargs: Any, + **kwargs: Unpack[_CustomGuardrailKwargs], ) -> None: """Initialize the Akto guardrail. @@ -96,8 +119,11 @@ class AktoGuardrail(CustomGuardrail): self.akto_account_id = akto_account_id or os.environ.get("AKTO_ACCOUNT_ID", "1000000") self.akto_vxlan_id = akto_vxlan_id or os.environ.get("AKTO_VXLAN_ID", "0") - kwargs["supported_event_hooks"] = list(self.get_supported_event_hooks()) - super().__init__(**kwargs) + init_kwargs: Final[_CustomGuardrailKwargs] = { + **kwargs, + "supported_event_hooks": list(self.get_supported_event_hooks()), + } + super().__init__(**init_kwargs) verbose_proxy_logger.debug( "Akto guardrail initialized: base_url=%s fallback=%s", diff --git a/litellm/proxy/guardrails/guardrail_hooks/custom_code/custom_code_guardrail.py b/litellm/proxy/guardrails/guardrail_hooks/custom_code/custom_code_guardrail.py index d5ef1e949b8..252b94b76c8 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/custom_code/custom_code_guardrail.py +++ b/litellm/proxy/guardrails/guardrail_hooks/custom_code/custom_code_guardrail.py @@ -38,9 +38,10 @@ import asyncio import threading import time from collections.abc import Callable, Mapping -from typing import TYPE_CHECKING, Any, Final, Literal, Optional, cast +from typing import TYPE_CHECKING, Final, Literal, Optional, cast from fastapi import HTTPException +from typing_extensions import TypedDict, Unpack from litellm._logging import verbose_proxy_logger from litellm.exceptions import ModifyResponseException @@ -74,6 +75,10 @@ class CustomCodeExecutionError(CustomCodeGuardrailError): """Raised when custom code fails during execution.""" +class _CustomGuardrailOptions(TypedDict, total=False, extra_items=object): + """Base-class constructor options this guardrail forwards untouched to CustomGuardrail.""" + + class CustomCodeGuardrailConfigModel(GuardrailConfigModel): """Configuration parameters for the custom code guardrail.""" @@ -109,7 +114,7 @@ class CustomCodeGuardrail(CustomGuardrail): self, custom_code: str, guardrail_name: str | None = "custom_code", - **kwargs: Any, + **kwargs: Unpack[_CustomGuardrailOptions], ) -> None: """ Initialize the custom code guardrail. diff --git a/litellm/proxy/guardrails/guardrail_hooks/lasso/lasso.py b/litellm/proxy/guardrails/guardrail_hooks/lasso/lasso.py index cf5da27e9ca..63821428c62 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/lasso/lasso.py +++ b/litellm/proxy/guardrails/guardrail_hooks/lasso/lasso.py @@ -121,7 +121,7 @@ class LassoGuardrail(CustomGuardrail): super().__init__(**kwargs) @staticmethod - def _get_field(obj: Any, field: str, default: object = None) -> Any: + def _get_field(obj: object, field: str, default: object = None) -> object: """Get a field from either a dict or a Pydantic object.""" if isinstance(obj, dict): return obj.get(field, default) @@ -130,7 +130,7 @@ class LassoGuardrail(CustomGuardrail): @staticmethod def _extract_tool_call_fields( call: object, - ) -> tuple[str | None, str | None, dict[str, object] | None]: + ) -> tuple[object, object, dict[str, object] | None]: """Extract (call_id, name, parsed_input) from a tool call. Handles both dict-style and Pydantic object-style tool_calls. @@ -146,7 +146,7 @@ class LassoGuardrail(CustomGuardrail): input_data: dict[str, object] | None = None if args_str: try: - parsed = json.loads(args_str) + parsed = json.loads(args_str) if isinstance(args_str, (str, bytes, bytearray)) else None except (json.JSONDecodeError, TypeError): parsed = None if isinstance(parsed, dict): @@ -488,7 +488,7 @@ class LassoGuardrail(CustomGuardrail): while preserving the original structure. """ # Index masked content by type so we can look up by id without caring about order. - masked_tool_use: Final[dict[str, dict[str, object]]] = {} + masked_tool_use: Final[dict[object, dict[str, object]]] = {} masked_tool_result: Final[dict[str, str]] = {} masked_text: Final[list[str]] = [] @@ -565,7 +565,7 @@ class LassoGuardrail(CustomGuardrail): def _update_tool_calls_from_masked( self, tool_calls: list[object], - masked_tool_use: dict[str, dict[str, object]], + masked_tool_use: Mapping[object, Mapping[str, object]], ) -> list[object]: """Replace tool_call arguments with masked values returned by Lasso.""" updated: Final = [] @@ -922,7 +922,7 @@ class LassoGuardrail(CustomGuardrail): ) -> None: """Apply masking to the actual model response when mask=True and masked content is available.""" # Index masked tool_use blocks by id for O(1) lookup. - masked_tool_use: Final[dict[str, dict[str, object]]] = {} + masked_tool_use: Final[dict[object, dict[str, object]]] = {} masked_text: Final[list[str]] = [] for masked_msg in masked_messages: content = masked_msg.get("content") diff --git a/litellm/proxy/guardrails/guardrail_hooks/vigil_guard/vigil_guard.py b/litellm/proxy/guardrails/guardrail_hooks/vigil_guard/vigil_guard.py index a5945a39589..e807da7079e 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/vigil_guard/vigil_guard.py +++ b/litellm/proxy/guardrails/guardrail_hooks/vigil_guard/vigil_guard.py @@ -3,7 +3,7 @@ from json import JSONDecodeError from typing import TYPE_CHECKING, Any, Final, Literal, Optional, Protocol, TypeAlias, cast import httpx -from typing_extensions import ReadOnly, TypedDict +from typing_extensions import ReadOnly, TypedDict, Unpack from litellm._logging import verbose_proxy_logger from litellm.exceptions import GuardrailRaisedException @@ -18,7 +18,8 @@ from litellm.llms.custom_httpx.http_handler import ( ) from litellm.secret_managers.main import get_secret_str from litellm.types.guardrails import GuardrailEventHooks -from litellm.types.utils import GenericGuardrailAPIInputs +from litellm.types.llms.openai import ChatCompletionToolCallChunk +from litellm.types.utils import ChatCompletionMessageToolCall, GenericGuardrailAPIInputs if TYPE_CHECKING: from litellm.litellm_core_utils.litellm_logging import ( @@ -53,6 +54,7 @@ _METADATA_ALLOWLIST: Final = ( _FallbackMode: TypeAlias = Literal["fail_closed", "fail_open"] _MetadataValue: TypeAlias = str | int | float | Sequence[str | int | float] +_ToolCalls: TypeAlias = list[ChatCompletionToolCallChunk] | list[ChatCompletionMessageToolCall] class _AnalyzePayload(TypedDict): @@ -70,6 +72,12 @@ class _AnalysisView(TypedDict): analysis: ReadOnly[Mapping[str, object]] +class _CustomGuardrailOptions(TypedDict, total=False, extra_items=object): + """Base-class constructor options this guardrail forwards untouched to CustomGuardrail.""" + + supported_event_hooks: ReadOnly[list[GuardrailEventHooks]] + + class _AsyncPostHandler(Protocol): def post( self, @@ -93,7 +101,7 @@ class VigilGuardGuardrail(CustomGuardrail): unreachable_fallback: str | None = None, timeout: float | None = None, async_handler: _AsyncPostHandler | None = None, - **kwargs: Any, + **kwargs: Unpack[_CustomGuardrailOptions], ) -> None: resolved_base: Final = api_base or get_secret_str("VIGIL_GUARD_URL") if not resolved_base: @@ -122,9 +130,12 @@ class VigilGuardGuardrail(CustomGuardrail): llm_provider=httpxSpecialProvider.GuardrailCallback, ) - kwargs.setdefault("supported_event_hooks", list(self.get_supported_event_hooks())) + forwarded: Final[_CustomGuardrailOptions] = { + "supported_event_hooks": list(self.get_supported_event_hooks()), + **kwargs, + } - super().__init__(**kwargs) + super().__init__(**forwarded) @staticmethod def get_config_model() -> type["GuardrailConfigModel"] | None: @@ -264,7 +275,7 @@ class VigilGuardGuardrail(CustomGuardrail): inputs: GenericGuardrailAPIInputs, source: str, final_texts: list[str], - final_tool_calls: Any, + final_tool_calls: _ToolCalls | None, ) -> GenericGuardrailAPIInputs: if self.unreachable_fallback == "fail_open": verbose_proxy_logger.error( diff --git a/litellm/proxy/guardrails/guardrail_hooks/xecguard/xecguard.py b/litellm/proxy/guardrails/guardrail_hooks/xecguard/xecguard.py index 831df43692b..f4330ad6aa9 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/xecguard/xecguard.py +++ b/litellm/proxy/guardrails/guardrail_hooks/xecguard/xecguard.py @@ -196,9 +196,9 @@ class XecGuardGuardrail(CustomGuardrail): async def async_logging_hook( self, kwargs: dict, - result: Any, + result: object, call_type: str, - ) -> tuple[dict, Any]: + ) -> tuple[dict, object]: """Observe-only scan for logging_only mode. Never blocks, never raises - all errors are swallowed. Records a @@ -275,9 +275,9 @@ class XecGuardGuardrail(CustomGuardrail): def logging_hook( self, kwargs: dict, - result: Any, + result: object, call_type: str, - ) -> tuple[dict, Any]: + ) -> tuple[dict, object]: """Sync counterpart to ``async_logging_hook``. Runs the async version on an available loop, swallowing every @@ -433,7 +433,7 @@ class XecGuardGuardrail(CustomGuardrail): return {"role": role, "content": ""} @staticmethod - def _synthesize_user_from_inputs(inputs: Any) -> dict | None: + def _synthesize_user_from_inputs(inputs: object) -> dict | None: if not isinstance(inputs, dict): return None texts: Final = inputs.get("texts") @@ -490,7 +490,7 @@ class XecGuardGuardrail(CustomGuardrail): return None @staticmethod - def _content_to_text(content: Any) -> str | None: + def _content_to_text(content: object) -> str | None: if isinstance(content, str) and content: return content if isinstance(content, list): diff --git a/litellm/proxy/management_endpoints/model_management_endpoints.py b/litellm/proxy/management_endpoints/model_management_endpoints.py index 2234e825090..50a498cf9bf 100644 --- a/litellm/proxy/management_endpoints/model_management_endpoints.py +++ b/litellm/proxy/management_endpoints/model_management_endpoints.py @@ -455,7 +455,6 @@ async def _auto_router_capability_slot( ENFORCE_RPM_TPM_ON_MODEL_ADD_SETTING: Final = "enforce_rpm_tpm_on_model_add" -_REQUIRED_RATE_LIMIT_FIELDS: Final = ("rpm", "tpm") def _raise_if_rate_limits_required_but_missing(*, litellm_params: GenericLiteLLMParams, enforced: bool) -> None: @@ -470,8 +469,8 @@ def _raise_if_rate_limits_required_but_missing(*, litellm_params: GenericLiteLLM return missing: Final = tuple( field - for field in _REQUIRED_RATE_LIMIT_FIELDS - if (value := getattr(litellm_params, field)) is None or value <= 0 + for field, value in (("rpm", litellm_params.rpm), ("tpm", litellm_params.tpm)) + if value is None or value <= 0 ) if not missing: return diff --git a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/vertex_passthrough_logging_handler.py b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/vertex_passthrough_logging_handler.py index 119a53c2411..ce44596c19c 100644 --- a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/vertex_passthrough_logging_handler.py +++ b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/vertex_passthrough_logging_handler.py @@ -28,6 +28,7 @@ from litellm.proxy.pass_through_endpoints.llm_provider_handlers.batch_attributio optional_str, request_tags_from_metadata, ) +from litellm.types.passthrough_endpoints.pass_through_endpoints import EndpointType from litellm.types.utils import ( Choices, EmbeddingResponse, @@ -47,8 +48,6 @@ else: PassThroughEndpointLogging = Any LiteLLMBatch = Any -EndpointType = Any - class VertexPassthroughLoggingHandler: @staticmethod diff --git a/litellm/proxy/response_api_endpoints/endpoints.py b/litellm/proxy/response_api_endpoints/endpoints.py index 5907ffc64eb..7f618526f11 100644 --- a/litellm/proxy/response_api_endpoints/endpoints.py +++ b/litellm/proxy/response_api_endpoints/endpoints.py @@ -1357,7 +1357,7 @@ async def _read_ws_model_from_first_frame( return model, first_message -def _extract_model_from_first_ws_event(first_event: Any) -> str | None: +def _extract_model_from_first_ws_event(first_event: object) -> str | None: """Extract model from a response.create WS event, handling flat and nested formats. Flat: {"type": "response.create", "model": "gpt-4o", ...} diff --git a/litellm/proxy/video_endpoints/endpoints.py b/litellm/proxy/video_endpoints/endpoints.py index 66071c05b4f..fe966c2e31a 100644 --- a/litellm/proxy/video_endpoints/endpoints.py +++ b/litellm/proxy/video_endpoints/endpoints.py @@ -89,7 +89,7 @@ async def video_generation( # Process request using ProxyBaseLLMRequestProcessing processor: Final = ProxyBaseLLMRequestProcessing(data=data) try: - return await processor.base_process_llm_request( + generated: Final[object] = await processor.base_process_llm_request( request=request, fastapi_response=fastapi_response, user_api_key_dict=user_api_key_dict, @@ -114,6 +114,8 @@ async def video_generation( proxy_logging_obj=proxy_logging_obj, version=version, ) + else: + return generated @router.get( @@ -174,7 +176,7 @@ async def video_list( # Process request using ProxyBaseLLMRequestProcessing processor: Final = ProxyBaseLLMRequestProcessing(data=data) try: - return await processor.base_process_llm_request( + listed: Final[object] = await processor.base_process_llm_request( request=request, fastapi_response=fastapi_response, user_api_key_dict=user_api_key_dict, @@ -199,6 +201,8 @@ async def video_list( proxy_logging_obj=proxy_logging_obj, version=version, ) + else: + return listed @router.get( @@ -272,7 +276,7 @@ async def video_status( # Process request using ProxyBaseLLMRequestProcessing processor: Final = ProxyBaseLLMRequestProcessing(data=data) try: - return await processor.base_process_llm_request( + status: Final[object] = await processor.base_process_llm_request( request=request, fastapi_response=fastapi_response, user_api_key_dict=user_api_key_dict, @@ -297,6 +301,8 @@ async def video_status( proxy_logging_obj=proxy_logging_obj, version=version, ) + else: + return status @router.get( @@ -478,7 +484,7 @@ async def video_remix( # Process request using ProxyBaseLLMRequestProcessing processor: Final = ProxyBaseLLMRequestProcessing(data=data) try: - return await processor.base_process_llm_request( + remixed: Final[object] = await processor.base_process_llm_request( request=request, fastapi_response=fastapi_response, user_api_key_dict=user_api_key_dict, @@ -503,6 +509,8 @@ async def video_remix( proxy_logging_obj=proxy_logging_obj, version=version, ) + else: + return remixed @router.post( @@ -571,7 +579,7 @@ async def video_create_character( processor: Final = ProxyBaseLLMRequestProcessing(data=data) try: - response = await processor.base_process_llm_request( + response: object = await processor.base_process_llm_request( request=request, fastapi_response=fastapi_response, user_api_key_dict=user_api_key_dict, @@ -678,7 +686,7 @@ async def video_get_character( processor: Final = ProxyBaseLLMRequestProcessing(data=data) try: - response = await processor.base_process_llm_request( + response: object = await processor.base_process_llm_request( request=request, fastapi_response=fastapi_response, user_api_key_dict=user_api_key_dict, @@ -789,7 +797,7 @@ async def video_edit( processor: Final = ProxyBaseLLMRequestProcessing(data=data) try: - return await processor.base_process_llm_request( + edited: Final[object] = await processor.base_process_llm_request( request=request, fastapi_response=fastapi_response, user_api_key_dict=user_api_key_dict, @@ -814,6 +822,8 @@ async def video_edit( proxy_logging_obj=proxy_logging_obj, version=version, ) + else: + return edited @router.post( @@ -884,7 +894,7 @@ async def video_extension( processor: Final = ProxyBaseLLMRequestProcessing(data=data) try: - return await processor.base_process_llm_request( + extended: Final[object] = await processor.base_process_llm_request( request=request, fastapi_response=fastapi_response, user_api_key_dict=user_api_key_dict, @@ -909,3 +919,5 @@ async def video_extension( proxy_logging_obj=proxy_logging_obj, version=version, ) + else: + return extended diff --git a/litellm/responses/mcp/litellm_proxy_mcp_handler.py b/litellm/responses/mcp/litellm_proxy_mcp_handler.py index a5021e2f777..c4bbf8c8b25 100644 --- a/litellm/responses/mcp/litellm_proxy_mcp_handler.py +++ b/litellm/responses/mcp/litellm_proxy_mcp_handler.py @@ -1266,14 +1266,14 @@ class LiteLLM_Proxy_MCP_Handler: return tool_execution_events @staticmethod - def _prepare_initial_call_params(call_params: dict[str, Any], should_auto_execute: bool) -> dict[str, Any]: + def _prepare_initial_call_params(call_params: Mapping[str, object], should_auto_execute: bool) -> dict[str, Any]: """ Prepare call parameters for the initial LLM call. For auto-execute scenarios, we need to disable streaming for the initial call so we can process the tool calls before streaming the final response. """ - initial_params: Final = call_params.copy() + initial_params: Final = dict(call_params) if should_auto_execute: # Disable streaming for initial call when auto-executing tools @@ -1282,14 +1282,16 @@ class LiteLLM_Proxy_MCP_Handler: return initial_params @staticmethod - def _prepare_follow_up_call_params(call_params: dict[str, Any], original_stream_setting: bool) -> dict[str, Any]: + def _prepare_follow_up_call_params( + call_params: Mapping[str, object], original_stream_setting: bool + ) -> dict[str, Any]: """ Prepare call parameters for the follow-up LLM call after tool execution. Restores the original streaming setting and removes tool_choice since we're now providing tool results, not requesting tool calls. """ - follow_up_params: Final = call_params.copy() + follow_up_params: Final = dict(call_params) # Restore original streaming setting for follow-up call follow_up_params["stream"] = original_stream_setting diff --git a/litellm/responses/streaming_iterator.py b/litellm/responses/streaming_iterator.py index 40ff88fc557..e1cf7847972 100644 --- a/litellm/responses/streaming_iterator.py +++ b/litellm/responses/streaming_iterator.py @@ -2353,7 +2353,7 @@ class ManagedResponsesWebSocketHandler: await self.websocket.send_text(serialized) @staticmethod - def _build_base_call_kwargs(msg_obj: _MutableJsonObject) -> dict[str, Any]: + def _build_base_call_kwargs(msg_obj: _MutableJsonObject) -> dict[str, object]: """ Extract Responses API params from the event, handling both wire formats: Nested: {"type": "response.create", "response": {"input": [...], ...}} diff --git a/litellm/router_strategy/complexity_router/complexity_router.py b/litellm/router_strategy/complexity_router/complexity_router.py index 9deccc9a468..eea0b2ec564 100644 --- a/litellm/router_strategy/complexity_router/complexity_router.py +++ b/litellm/router_strategy/complexity_router/complexity_router.py @@ -363,7 +363,7 @@ def _append_custom_keywords(base_keywords: list[str], custom_keywords: list[str] return [*base_keywords, *deduped_custom.values()] -def _parent_session_kwargs(request_kwargs: Mapping[str, Any] | None) -> Mapping[str, Any]: +def _parent_session_kwargs(request_kwargs: Mapping[str, object] | None) -> Mapping[str, Any]: kwargs: Final = request_kwargs or {} return {k: kwargs[k] for k in ("litellm_session_id", "litellm_trace_id") if kwargs.get(k) is not None} @@ -1165,7 +1165,7 @@ class ComplexityRouter(CustomLogger): self, model_name: str, litellm_router_instance: Router, - complexity_router_config: dict[str, Any] | None = None, + complexity_router_config: Mapping[str, object] | None = None, default_model: str | None = None, derive_savings_baseline: bool = True, ): @@ -1736,7 +1736,7 @@ class ComplexityRouter(CustomLogger): self, prompt: str, system_prompt: str | None, - request_kwargs: dict[str, Any] | None, # mutable-ok: handed to _classify_with_llm as-is + request_kwargs: dict[str, object] | None, # mutable-ok: handed to _classify_with_llm as-is messages: Sequence[Mapping[str, object]] | None, ) -> ClassificationOutcome: """Score locally, and only pay for the classifier call when the scorer did not confidently @@ -1769,7 +1769,7 @@ class ComplexityRouter(CustomLogger): self, prompt: str, system_prompt: str | None, - request_kwargs: dict[str, Any] | None, # mutable-ok: handed to _classify_with_llm as-is + request_kwargs: dict[str, object] | None, # mutable-ok: handed to _classify_with_llm as-is messages: Sequence[Mapping[str, object]] | None, ) -> ClassificationOutcome: """Score locally, and only pay for the classifier when the score sits near a tier boundary. @@ -1824,7 +1824,7 @@ class ComplexityRouter(CustomLogger): self, prompt: str, system_prompt: str | None, - request_kwargs: dict[str, Any] | None, # mutable-ok: handed to _classify_with_llm as-is + request_kwargs: dict[str, object] | None, # mutable-ok: handed to _classify_with_llm as-is messages: Sequence[Mapping[str, object]] | None, scored: ClassificationOutcome | None = None, ) -> ClassificationOutcome: @@ -1902,8 +1902,8 @@ class ComplexityRouter(CustomLogger): self, prompt: str, system_prompt: str | None, - request_kwargs: dict[str, Any] | None, # mutable-ok: handed to resolve_structured_messages as-is - raw_messages: list[dict[str, Any]] | None, # mutable-ok: same shape _run_routing_plugins receives + request_kwargs: dict[str, object] | None, # mutable-ok: handed to resolve_structured_messages as-is + raw_messages: list[dict[str, object]] | None, # mutable-ok: same shape _run_routing_plugins receives ) -> ClassificationOutcome: from litellm.litellm_core_utils.prompt_templates.factory import resolve_structured_messages from litellm.types.router import RoutingContext @@ -2262,8 +2262,8 @@ class ComplexityRouter(CustomLogger): async def _pick_model_for_tier( self, tier: ComplexityTier | str, - raw_messages: list[dict[str, Any]] | None, - resolved_messages: list[dict[str, Any]] | None, + raw_messages: list[dict[str, object]] | None, + resolved_messages: list[dict[str, object]] | None, request_kwargs: dict, allowed_models: tuple[str, ...] | None = None, ) -> str: @@ -2373,7 +2373,7 @@ class ComplexityRouter(CustomLogger): self, classified_tier: ComplexityTier | str, user_message: str, - request_kwargs: dict[str, Any] | None = None, + request_kwargs: dict[str, object] | None = None, hard_floor: ComplexityTier | str | None = None, hard_ceiling: ComplexityTier | str | None = None, fit_filter: frozenset[str] | None = None, @@ -2903,7 +2903,7 @@ class ComplexityRouter(CustomLogger): async def _gate_response_modality( self, response: PreRoutingHookResponse, - messages: list[dict[str, Any]] | None, # mutable-ok: forwarded verbatim to the list-typed re-pick + messages: list[dict[str, object]] | None, # mutable-ok: forwarded verbatim to the list-typed re-pick resolved_messages: Sequence[Mapping[str, object]] | None, request_kwargs: dict, # mutable-ok: same shape the hook receives context_fit: _RequestContextFit | None = None, @@ -3093,7 +3093,7 @@ class ComplexityRouter(CustomLogger): async def _gate_response_health( self, response: PreRoutingHookResponse, - messages: list[dict[str, Any]] | None, # mutable-ok: forwarded verbatim to the list-typed re-pick + messages: list[dict[str, object]] | None, # mutable-ok: forwarded verbatim to the list-typed re-pick input: str | list | None, # mutable-ok: mirrors the owner's own input parameter, which this forwards verbatim resolved_messages: Sequence[Mapping[str, object]] | None, request_kwargs: dict, # mutable-ok: same shape the hook receives @@ -3405,9 +3405,9 @@ class ComplexityRouter(CustomLogger): def _resolve_messages( self, - messages: list[dict[str, Any]] | None, + messages: list[dict[str, object]] | None, request_kwargs: dict, - ) -> list[dict[str, Any]] | None: + ) -> list[dict[str, object]] | None: """ Resolve messages from the request, converting from other formats if needed. @@ -3422,7 +3422,7 @@ class ComplexityRouter(CustomLogger): @staticmethod def _extract_user_message_and_system_prompt( - messages: list[dict[str, Any]], + messages: Sequence[Mapping[str, object]], ) -> tuple[str | None, str | None]: """ Deprecated: use _extract_current_ask_and_system_prompt instead. @@ -3729,7 +3729,7 @@ class ComplexityRouter(CustomLogger): self, model: str, request_kwargs: dict, - messages: list[dict[str, Any]] | None = None, + messages: list[dict[str, object]] | None = None, input: str | list | None = None, specific_deployment: bool | None = False, conversation_continuing: bool = True, diff --git a/litellm/types/router.py b/litellm/types/router.py index c7363502017..f0d2405c88f 100644 --- a/litellm/types/router.py +++ b/litellm/types/router.py @@ -248,7 +248,7 @@ class ModelInfo(MirroredPricingParams): # Custom .get() method to access attributes with a default value if the attribute doesn't exist return getattr(self, key, default) - def __getitem__(self, key): + def __getitem__(self, key) -> object: # Allow dictionary-style access to attributes return getattr(self, key) @@ -358,7 +358,7 @@ class GenericLiteLLMParams(CredentialLiteLLMParams, CustomPricingLiteLLMParams): model_config = ConfigDict(extra="allow", arbitrary_types_allowed=True) merge_reasoning_content_in_choices: bool | None = False model_info: dict | None = None - mock_response: str | ModelResponse | Exception | Any | None = None + mock_response: str | ModelResponse | Exception | object | None = None # tag-based routing tags: list[str] | None = None @@ -435,7 +435,7 @@ class GenericLiteLLMParams(CredentialLiteLLMParams, CustomPricingLiteLLMParams): # Custom .get() method to access attributes with a default value if the attribute doesn't exist return getattr(self, key, default) - def __getitem__(self, key): + def __getitem__(self, key) -> object: # Allow dictionary-style access to attributes return getattr(self, key) @@ -460,7 +460,7 @@ class LiteLLM_Params(GenericLiteLLMParams): # Custom .get() method to access attributes with a default value if the attribute doesn't exist return getattr(self, key, default) - def __getitem__(self, key): + def __getitem__(self, key) -> object: # Allow dictionary-style access to attributes return getattr(self, key) @@ -1043,11 +1043,11 @@ class RoutingContext(BaseModel): plugins that need the exact original payload can read `raw_messages`. """ - raw_messages: list[dict[str, Any]] - structured_messages: list[dict[str, Any]] + raw_messages: list[dict[str, object]] + structured_messages: list[dict[str, object]] candidate_models: list[str] - metadata: dict[str, Any] = Field(default_factory=dict) - signals: dict[str, Any] = Field(default_factory=dict) + metadata: dict[str, object] = Field(default_factory=dict) + signals: dict[str, object] = Field(default_factory=dict) @runtime_checkable diff --git a/litellm/vector_stores/vector_store_registry.py b/litellm/vector_stores/vector_store_registry.py index b71d6784873..c7aed77286c 100644 --- a/litellm/vector_stores/vector_store_registry.py +++ b/litellm/vector_stores/vector_store_registry.py @@ -112,9 +112,8 @@ class VectorStoreRegistry: Dynamically extracts all parameters defined in VECTOR_STORE_OPENAI_PARAMS. """ # Get the list of supported param names from the Literal type - supported_params: Final = tuple( - param for param in get_args(VECTOR_STORE_OPENAI_PARAMS) if isinstance(param, str) - ) + declared_params: Final[tuple[object, ...]] = get_args(VECTOR_STORE_OPENAI_PARAMS) + supported_params: Final = tuple(param for param in declared_params if isinstance(param, str)) # Extract only the params that exist in the tool kwargs: Final = {param: tool.get(param) for param in supported_params if param in tool} From 6e2c088bc265d17ae31a8924e0f1494d291f6cf5 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Mon, 14 Sep 2026 11:13:42 +0000 Subject: [PATCH 10/44] chore(ui): regenerate dashboard API types Picks up the user-endpoint docstring removal already on main. --- ui/litellm-dashboard/src/lib/http/schema.d.ts | 2 -- 1 file changed, 2 deletions(-) diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index 7eadaa6c991..839aa52fa84 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -16781,7 +16781,6 @@ export interface paths { * - permissions: Optional[dict] - [Not Implemented Yet] User-specific permissions, eg. turning off pii masking. * - metadata: Optional[dict] - Metadata for user, store information for user. Example metadata = {"team": "core-infra", "app": "app2", "email": "ishaan@berri.ai" } * - max_parallel_requests: Optional[int] - Rate limit a user based on the number of parallel requests. Raises 429 error, if user's parallel requests > x. - * - 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) * - budget_fallbacks: Optional[Dict[str, List[str]]] - Per-model fallback chain tried in order when that model's own `model_max_budget` is exceeded, e.g. {"gpt-4o": ["gpt-4o-mini"]}. * - 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) @@ -16887,7 +16886,6 @@ export interface paths { * - permissions: Optional[dict] - [Not Implemented Yet] User-specific permissions, eg. turning off pii masking. * - metadata: Optional[dict] - Metadata for user, store information for user. Example metadata = {"team": "core-infra", "app": "app2", "email": "ishaan@berri.ai" } * - max_parallel_requests: Optional[int] - Rate limit a user based on the number of parallel requests. Raises 429 error, if user's parallel requests > x. - * - 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) * - budget_fallbacks: Optional[Dict[str, List[str]]] - Per-model fallback chain tried in order when that model's own `model_max_budget` is exceeded, e.g. {"gpt-4o": ["gpt-4o-mini"]}. * - 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) From bbee8692a80bf2be2a917c88fb3fb2ba76b4e5a3 Mon Sep 17 00:00:00 2001 From: Tejas Chopra Date: Wed, 16 Sep 2026 21:27:45 -0700 Subject: [PATCH 11/44] fix(responses): stop agentic follow-up from passing request params twice --- litellm/llms/custom_httpx/llm_http_handler.py | 22 +++++----- .../custom_httpx/test_llm_http_handler.py | 44 +++++++++++++++++++ 2 files changed, 56 insertions(+), 10 deletions(-) diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index 2fe4130a310..b4f2a74b3e6 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -5573,17 +5573,19 @@ class BaseLLMHTTPHandler: internal_keys: Final = {"litellm_logging_obj"} kwargs_for_followup: Final = { - k: v - for k, v in kwargs.items() - if not is_interception_internal_key(k, prefixes=NON_CODE_INTERPRETER_INTERCEPTION_INTERNAL_PREFIXES) - and k != "_code_interpreter_interception_converted_stream" - and k not in internal_keys - and k not in optional_params + **{ + k: v + for k, v in kwargs.items() + if not is_interception_internal_key(k, prefixes=NON_CODE_INTERPRETER_INTERCEPTION_INTERNAL_PREFIXES) + and k != "_code_interpreter_interception_converted_stream" + and k not in internal_keys + and k not in optional_params + }, + **{k: v for k, v in patch.kwargs.items() if k not in optional_params}, + "_agentic_loop_depth": depth + 1, + "max_agentic_loops": max_loops, + "_agentic_loop_fingerprints": fingerprints + [fingerprint], } - kwargs_for_followup.update(patch.kwargs) - kwargs_for_followup["_agentic_loop_depth"] = depth + 1 - kwargs_for_followup["max_agentic_loops"] = max_loops - kwargs_for_followup["_agentic_loop_fingerprints"] = fingerprints + [fingerprint] try: response: ResponsesAPIResponse | BaseResponsesAPIStreamingIterator = await litellm.aresponses( diff --git a/tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py b/tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py index c39779972c0..5d6060c7666 100644 --- a/tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py +++ b/tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py @@ -3713,3 +3713,47 @@ def test_image_edit_handler_keeps_the_sync_transform(): assert config.transform_calls == ["sync"] assert captured["body"] == {"transformed_by": "sync"} assert response.data[0].b64_json == "sync" + + +@pytest.mark.asyncio +async def test_responses_agentic_followup_does_not_repeat_request_params_from_plan_kwargs(monkeypatch): + """A plan whose kwargs repeat a request param must not crash the Responses follow-up with a duplicate keyword""" + from litellm.integrations.custom_logger import CustomLogger + from litellm.types.integrations.custom_logger import AgenticLoopPlan, AgenticLoopRequestPatch + + followup_calls: list[dict[str, object]] = [] + + async def fake_aresponses(**kwargs: object) -> str: + followup_calls.append(kwargs) + return "followup-response" + + monkeypatch.setattr(litellm, "aresponses", fake_aresponses) + request_kwargs: Final = {"prompt_cache_key": "thread-1", "metadata": {"user": "u1"}} + plan: Final = AgenticLoopPlan( + run_agentic_loop=True, + request_patch=AgenticLoopRequestPatch( + model="gpt-5", + messages=[{"role": "user", "content": "x"}], + optional_params={"prompt_cache_key": "thread-1"}, + kwargs=dict(request_kwargs), + ), + ) + + response: Final = await BaseLLMHTTPHandler()._execute_responses_agentic_plan( + plan=plan, + model="gpt-5", + response_api_optional_request_params={"prompt_cache_key": "thread-1"}, + logging_obj=Mock(litellm_call_id="call-1"), + kwargs=dict(request_kwargs), + depth=0, + max_loops=3, + fingerprints=[], + fingerprint="fp", + callback=CustomLogger(), + ) + + assert response == "followup-response" + assert len(followup_calls) == 1 + assert followup_calls[0]["prompt_cache_key"] == "thread-1" + assert followup_calls[0]["metadata"] == {"user": "u1"} + assert followup_calls[0]["_agentic_loop_depth"] == 1 From 3298b416878ec8b2d409b47799f76e4f9bcbc0f7 Mon Sep 17 00:00:00 2001 From: Tejas Chopra Date: Wed, 16 Sep 2026 21:44:10 -0700 Subject: [PATCH 12/44] refactor(responses): build agentic follow-up kwargs as one frozen mapping --- litellm/llms/custom_httpx/llm_http_handler.py | 38 ++++++++++++------- 1 file changed, 24 insertions(+), 14 deletions(-) diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index b4f2a74b3e6..50420ab3301 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -4,6 +4,7 @@ import ssl from collections.abc import AsyncIterator, Coroutine, Iterator, Mapping, Sequence from contextlib import asynccontextmanager from functools import lru_cache +from itertools import chain from types import MappingProxyType, ModuleType from typing import TYPE_CHECKING, Any, Final, Literal, Optional, TypedDict, TypeVar, Union, cast, get_type_hints from urllib.parse import parse_qs, urlencode, urlparse, urlunparse @@ -5572,20 +5573,29 @@ class BaseLLMHTTPHandler: } internal_keys: Final = {"litellm_logging_obj"} - kwargs_for_followup: Final = { - **{ - k: v - for k, v in kwargs.items() - if not is_interception_internal_key(k, prefixes=NON_CODE_INTERPRETER_INTERCEPTION_INTERNAL_PREFIXES) - and k != "_code_interpreter_interception_converted_stream" - and k not in internal_keys - and k not in optional_params - }, - **{k: v for k, v in patch.kwargs.items() if k not in optional_params}, - "_agentic_loop_depth": depth + 1, - "max_agentic_loops": max_loops, - "_agentic_loop_fingerprints": fingerprints + [fingerprint], - } + kwargs_for_followup: Final = MappingProxyType( + { + key: value + for key, value in chain( + ( + (k, v) + for k, v in kwargs.items() + if not is_interception_internal_key( + k, prefixes=NON_CODE_INTERPRETER_INTERCEPTION_INTERNAL_PREFIXES + ) + and k != "_code_interpreter_interception_converted_stream" + and k not in internal_keys + and k not in optional_params + ), + ((k, v) for k, v in patch.kwargs.items() if k not in optional_params), + ( + ("_agentic_loop_depth", depth + 1), + ("max_agentic_loops", max_loops), + ("_agentic_loop_fingerprints", fingerprints + [fingerprint]), + ), + ) + } + ) try: response: ResponsesAPIResponse | BaseResponsesAPIStreamingIterator = await litellm.aresponses( From 484524b70b9cd28108547b201b2fa88e1fce27fb Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Sat, 19 Sep 2026 16:04:13 -0700 Subject: [PATCH 13/44] fix(auth): fail closed when the team membership lookup hits a db outage --- litellm/proxy/auth/auth_checks.py | 28 ++++---- .../proxy/auth/test_auth_checks.py | 68 ++++++++++++++----- .../proxy/auth/test_resolvers_grants.py | 28 ++++++++ 3 files changed, 90 insertions(+), 34 deletions(-) diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 61d2fa572a1..32ae9bc81df 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -2296,23 +2296,19 @@ async def _load_team_membership_on_cache_miss( parent_otel_span: Span | None, proxy_logging_obj: ProxyLogging | None, ) -> LiteLLM_TeamMembership | None: - try: - redis_cached: Final[object] = await user_api_key_cache.async_get_cache(key=cache_key) - redis_membership: Final = _membership_from_cached_payload(redis_cached) - if not isinstance(redis_membership, _TeamMembershipCacheMiss): - return redis_membership + redis_cached: Final[object] = await user_api_key_cache.async_get_cache(key=cache_key) + redis_membership: Final = _membership_from_cached_payload(redis_cached) + if not isinstance(redis_membership, _TeamMembershipCacheMiss): + return redis_membership - return await _fetch_team_membership_from_db( - user_id=user_id, - team_id=team_id, - prisma_client=prisma_client, - user_api_key_cache=user_api_key_cache, - parent_otel_span=parent_otel_span, - proxy_logging_obj=proxy_logging_obj, - ) - except Exception: - verbose_proxy_logger.exception("Error getting team membership") - return None + return await _fetch_team_membership_from_db( + user_id=user_id, + team_id=team_id, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + parent_otel_span=parent_otel_span, + proxy_logging_obj=proxy_logging_obj, + ) async def get_team_membership( diff --git a/tests/test_litellm/proxy/auth/test_auth_checks.py b/tests/test_litellm/proxy/auth/test_auth_checks.py index 1ae986db23b..233660b634a 100644 --- a/tests/test_litellm/proxy/auth/test_auth_checks.py +++ b/tests/test_litellm/proxy/auth/test_auth_checks.py @@ -7198,7 +7198,7 @@ async def test_common_checks_skips_membership_load_when_no_check_reads_it(): @pytest.mark.asyncio -async def test_get_team_membership_db_error_returns_none_and_retries_next_call(): +async def test_get_team_membership_db_error_surfaces_and_retries_next_call(): from litellm.proxy.auth.auth_checks import get_team_membership from litellm.proxy.common_utils.user_api_key_cache import team_membership_reservation_cache_key @@ -7210,12 +7210,13 @@ async def test_get_team_membership_db_error_returns_none_and_retries_next_call() ) cache = UserApiKeyCache() - failed = await get_team_membership( - user_id="u-fail", - team_id="t-fail", - prisma_client=mock_prisma_client, - user_api_key_cache=cache, - ) + with pytest.raises(RuntimeError, match="db down"): + await get_team_membership( + user_id="u-fail", + team_id="t-fail", + prisma_client=mock_prisma_client, + user_api_key_cache=cache, + ) cached_after_failure = await cache.async_get_cache( key=team_membership_reservation_cache_key(user_id="u-fail", team_id="t-fail") ) @@ -7226,24 +7227,55 @@ async def test_get_team_membership_db_error_returns_none_and_retries_next_call() user_api_key_cache=cache, ) - assert failed is None assert cached_after_failure is None assert recovered is not None assert recovered.user_id == "u-fail" assert mock_prisma_client.db.litellm_teammembership.find_unique.await_count == 2 -@pytest.mark.asyncio -async def test_get_team_membership_string_prisma_client_returns_none(): - from litellm.proxy.auth.auth_checks import get_team_membership +class _UnreachableMembershipPrisma: + class db: + class litellm_teammembership: + @staticmethod + async def find_unique(where: dict[str, dict[str, str]], include: dict[str, bool]) -> None: + raise httpx.ConnectError("All connection attempts failed") - result = await get_team_membership( - user_id="u-str", - team_id="t-str", - prisma_client="hello-world", - user_api_key_cache=UserApiKeyCache(), - ) - assert result is None + +def _restricted_member_check_deps() -> dict[str, object]: + from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache + from litellm.proxy.utils import ProxyLogging + + cache = UserApiKeyCache() + return { + "team_object": LiteLLM_TeamTable(team_id="team-outage", models=["claude-sonnet-5"]), + "valid_token": UserAPIKeyAuth(token="hashed-fake", user_id="bob", team_id="team-outage"), + "prisma_client": _UnreachableMembershipPrisma(), + "user_api_key_cache": cache, + "proxy_logging_obj": ProxyLogging(user_api_key_cache=cache), + } + + +@pytest.mark.asyncio +async def test_check_team_member_model_access_fails_closed_when_the_membership_read_hits_a_db_outage(): + """Regression: with the member's row uncached and the database unreachable, the loader used to swallow the + transport error and return None, which every check reads as "no per-member restriction", so a member + limited to other models got a 200. The outage must surface as the 503 the rest of auth answers with.""" + from litellm.proxy.auth.auth_checks import _check_team_member_model_access + from litellm.proxy.auth.auth_exception_handler import _as_proxy_exception + + with pytest.raises(httpx.ConnectError) as raised: + await _check_team_member_model_access( + model="claude-sonnet-5", llm_router=None, **_restricted_member_check_deps() + ) + + surfaced = _as_proxy_exception(raised.value) + assert (surfaced.code, surfaced.type) == ("503", ProxyErrorTypes.no_db_connection) + + +@pytest.mark.asyncio +async def test_check_team_member_budget_fails_closed_when_the_membership_read_hits_a_db_outage(): + with pytest.raises(httpx.ConnectError): + await _check_team_member_budget(user_object=None, **_restricted_member_check_deps()) @pytest.mark.asyncio diff --git a/tests/test_litellm/proxy/auth/test_resolvers_grants.py b/tests/test_litellm/proxy/auth/test_resolvers_grants.py index 3f6d943bf98..415d1191b5d 100644 --- a/tests/test_litellm/proxy/auth/test_resolvers_grants.py +++ b/tests/test_litellm/proxy/auth/test_resolvers_grants.py @@ -1,4 +1,5 @@ from fastapi import HTTPException +import httpx import pytest from litellm.proxy._types import ( @@ -8,6 +9,7 @@ from litellm.proxy._types import ( ProxyException, ) from litellm.proxy.auth.auth_checks import TeamNotFoundError, UserNotFoundError +from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache from litellm.proxy.auth.resolvers.grants import ( GrantResolver, LookupDegraded, @@ -172,6 +174,32 @@ async def test_resolve_identity_lets_loader_errors_surface(): await loaders.resolver().resolve_identity(UserLookup(user_id=USER_ID), team_id=None) +class _UnreachableMembershipPrisma: + class db: + class litellm_teammembership: + @staticmethod + async def find_unique(where: dict[str, dict[str, str]], include: dict[str, bool]) -> None: + raise httpx.ConnectError("All connection attempts failed") + + +async def test_resolve_marks_a_membership_read_that_hits_a_db_outage_as_degraded(): + """Regression: the real membership loader swallowed a database transport error into None, so this outcome + was ResolvedGrants with no membership, never LookupDegraded, and a member's own model or budget limits + silently dropped for the request.""" + loaders = _Loaders(user=_user(), team=_team()) + resolver = GrantResolver( + _UnreachableMembershipPrisma(), + UserApiKeyCache(), + load_user=loaders.load_user, + load_team=loaders.load_team, + ) + + outcome = await resolver.resolve(UserLookup(user_id=USER_ID), team_id=TEAM_ID) + + assert isinstance(outcome, LookupDegraded) + assert isinstance(outcome.error, httpx.ConnectError) + + def test_raise_public_maps_a_deleted_user_to_401(): with pytest.raises(ProxyException) as exc_info: raise_public(UserGone(user_id=USER_ID)) From b3b280d46381552b0624261b57e6911d08469fb9 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Sat, 19 Sep 2026 16:32:22 -0700 Subject: [PATCH 14/44] test(auth): model the membership row read in the fakes the loader now reaches --- tests/proxy_unit_tests/test_user_api_key_auth.py | 10 +++++++++- .../mcp_server/test_discoverable_endpoints.py | 4 +++- tests/test_litellm/proxy/auth/test_auth_checks.py | 3 --- tests/test_litellm/proxy/auth/test_resolvers_grants.py | 3 --- 4 files changed, 12 insertions(+), 8 deletions(-) diff --git a/tests/proxy_unit_tests/test_user_api_key_auth.py b/tests/proxy_unit_tests/test_user_api_key_auth.py index a8fce58c60b..f5e8d861d79 100644 --- a/tests/proxy_unit_tests/test_user_api_key_auth.py +++ b/tests/proxy_unit_tests/test_user_api_key_auth.py @@ -201,6 +201,14 @@ async def test_returned_user_api_key_auth(user_role, expected_role): assert new_obj.user_role == expected_role +class _NoMembershipRowPrisma: + class db: + class litellm_teammembership: + @staticmethod + async def find_unique(where: dict[str, dict[str, str]], include: dict[str, bool]) -> None: + return None + + @pytest.mark.parametrize("key_ownership", ["user_key", "team_key"]) @pytest.mark.asyncio async def test_aaauser_personal_budgets(key_ownership): @@ -253,7 +261,7 @@ async def test_aaauser_personal_budgets(key_ownership): setattr(litellm.proxy.proxy_server, "user_api_key_cache", user_api_key_cache) setattr(litellm.proxy.proxy_server, "master_key", "sk-1234") - setattr(litellm.proxy.proxy_server, "prisma_client", "hello-world") + setattr(litellm.proxy.proxy_server, "prisma_client", _NoMembershipRowPrisma()) request = Request(scope={"type": "http"}) request._url = URL(url="/chat/completions") 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 d7666f5e694..bed12892665 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 @@ -11515,7 +11515,9 @@ def jwt_oauth_identity(monkeypatch: pytest.MonkeyPatch) -> tuple["JWTHandler", " monkeypatch.setattr(proxy_server, "general_settings", {"enable_jwt_auth": True}) monkeypatch.setattr(proxy_server, "premium_user", True) monkeypatch.setattr(proxy_server, "user_api_key_cache", cache) - monkeypatch.setattr(proxy_server, "prisma_client", MagicMock()) + prisma: Final = MagicMock() + prisma.db.litellm_teammembership.find_unique = AsyncMock(return_value=None) + monkeypatch.setattr(proxy_server, "prisma_client", prisma) return handler, signing_key diff --git a/tests/test_litellm/proxy/auth/test_auth_checks.py b/tests/test_litellm/proxy/auth/test_auth_checks.py index 233660b634a..4f3e31fd70d 100644 --- a/tests/test_litellm/proxy/auth/test_auth_checks.py +++ b/tests/test_litellm/proxy/auth/test_auth_checks.py @@ -7257,9 +7257,6 @@ def _restricted_member_check_deps() -> dict[str, object]: @pytest.mark.asyncio async def test_check_team_member_model_access_fails_closed_when_the_membership_read_hits_a_db_outage(): - """Regression: with the member's row uncached and the database unreachable, the loader used to swallow the - transport error and return None, which every check reads as "no per-member restriction", so a member - limited to other models got a 200. The outage must surface as the 503 the rest of auth answers with.""" from litellm.proxy.auth.auth_checks import _check_team_member_model_access from litellm.proxy.auth.auth_exception_handler import _as_proxy_exception diff --git a/tests/test_litellm/proxy/auth/test_resolvers_grants.py b/tests/test_litellm/proxy/auth/test_resolvers_grants.py index 415d1191b5d..e61269ec1c7 100644 --- a/tests/test_litellm/proxy/auth/test_resolvers_grants.py +++ b/tests/test_litellm/proxy/auth/test_resolvers_grants.py @@ -183,9 +183,6 @@ class _UnreachableMembershipPrisma: async def test_resolve_marks_a_membership_read_that_hits_a_db_outage_as_degraded(): - """Regression: the real membership loader swallowed a database transport error into None, so this outcome - was ResolvedGrants with no membership, never LookupDegraded, and a member's own model or budget limits - silently dropped for the request.""" loaders = _Loaders(user=_user(), team=_team()) resolver = GrantResolver( _UnreachableMembershipPrisma(), From 86960fb127a6d7fe508e6632f4f189a047c46a65 Mon Sep 17 00:00:00 2001 From: shivam Date: Mon, 21 Sep 2026 19:11:38 +0000 Subject: [PATCH 15/44] fix(router): walk every entry of a fallback list after a mid-stream failure A fallback hop that dies before its first chunk surfaces inside the streaming iterator, where the chain lookup is keyed by the hop's own group. That group has no chain of its own, so the remaining entries of the original list were never tried. Resume the original group's chain as the last lookup key; attempted_targets already skips the entries that were tried. Resolves LIT-7400 Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../router_utils/fallback_event_handlers.py | 12 ++- .../test_fallback_event_handlers.py | 8 ++ tests/test_litellm/test_router.py | 75 +++++++++++++++++++ 3 files changed, 92 insertions(+), 3 deletions(-) diff --git a/litellm/router_utils/fallback_event_handlers.py b/litellm/router_utils/fallback_event_handlers.py index d0abaed4d3a..156195f0587 100644 --- a/litellm/router_utils/fallback_event_handlers.py +++ b/litellm/router_utils/fallback_event_handlers.py @@ -307,13 +307,19 @@ def fallbacks_disabled_for_request(kwargs: Mapping[str, Any]) -> bool: def fallback_lookup_groups(kwargs: Mapping[str, object], model_group: str | None) -> tuple[str, ...]: """ Ordered keys for resolving a fallback chain: the tier a pre-routing hook selected wins, - then the routed group, then the requested group. The routed group differs when Claude Code - session affinity remaps a subagent's concrete model to its bound router. + then the routed group, then the requested group, then the group the request was + originally for. The routed group differs when Claude Code session affinity remaps a + subagent's concrete model to its bound router. The original group differs on a fallback + hop that fails after `run_async_fallback` already returned its stream: the hop has no + chain of its own, so it resumes the original group's chain, and `attempted_targets` keeps + the entries already tried from being repeated. """ metadata: Final = kwargs.get(get_metadata_variable_name_from_kwargs(kwargs)) routed_group_value: Final = metadata.get("model_group") if isinstance(metadata, Mapping) else None routed_group: Final = routed_group_value if isinstance(routed_group_value, str) else None - ordered: Final = (get_pre_routing_selection(kwargs), routed_group, model_group) + original_group_value: Final = metadata.get("original_model_group") if isinstance(metadata, Mapping) else None + original_group: Final = original_group_value if isinstance(original_group_value, str) else None + ordered: Final = (get_pre_routing_selection(kwargs), routed_group, model_group, original_group) return tuple(dict.fromkeys(group for group in ordered if group)) diff --git a/tests/test_litellm/router_utils/test_fallback_event_handlers.py b/tests/test_litellm/router_utils/test_fallback_event_handlers.py index dfe06bffd09..69772561172 100644 --- a/tests/test_litellm/router_utils/test_fallback_event_handlers.py +++ b/tests/test_litellm/router_utils/test_fallback_event_handlers.py @@ -1305,6 +1305,14 @@ class TestOrderedFallbackLookupGroups: "requested-model", ) + def test_fallback_hop_resumes_the_original_groups_chain_last(self): + from litellm.router_utils.fallback_event_handlers import fallback_lookup_groups + + kwargs = {"metadata": {"model_group": "fb1", "original_model_group": "primary"}} + + assert fallback_lookup_groups(kwargs, "fb1") == ("fb1", "primary") + assert fallback_lookup_groups({"metadata": {"original_model_group": 42}}, "fb1") == ("fb1",) + def test_first_resolving_group_wins_and_generic_idx_survives_a_miss(self): from litellm.router_utils.fallback_event_handlers import ( get_fallback_model_group_for_lookup_groups, diff --git a/tests/test_litellm/test_router.py b/tests/test_litellm/test_router.py index 8310d30d90e..cf143744067 100644 --- a/tests/test_litellm/test_router.py +++ b/tests/test_litellm/test_router.py @@ -3447,6 +3447,81 @@ def test_completion_streaming_iterator_adopts_the_deployment_that_served_a_neste assert result._hidden_params["model_id"] == "served-deployment" +@pytest.mark.asyncio +async def test_acompletion_mid_stream_fallback_walks_every_entry_of_the_configured_list(): + """LIT-7400: fallbacks=[{primary: [fb1, fb2]}] must reach fb2 when fb1 dies before its first chunk. + + run_async_fallback returns as soon as fb1's stream wrapper exists, so fb1's failure surfaces + inside the streaming iterator, where the lookup is keyed by fb1. That key has no chain of its + own, so the iterator has to resume the chain of the group the request was originally for. + """ + from unittest.mock import MagicMock, patch + + from litellm.exceptions import MidStreamFallbackError + from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper + + attempted_model_groups: list[str] = [] + + class FailingStream(CustomStreamWrapper): + def __init__(self, model: str): + super().__init__( + completion_stream=object(), model=model, custom_llm_provider="openai", logging_obj=MagicMock() + ) + + def __aiter__(self): + return self + + async def __anext__(self): + raise MidStreamFallbackError( + message=f"provider 500 from {self.model}", + model=self.model, + llm_provider="openai", + generated_content="", + is_pre_first_chunk=True, + original_exception=litellm.InternalServerError( + message=f"provider 500 from {self.model}", model=self.model, llm_provider="openai" + ), + ) + + class OkStream(FailingStream): + def __init__(self, model: str): + super().__init__(model) + self._chunks = iter( + [litellm.ModelResponseStream(choices=[{"index": 0, "delta": {"content": f"ok-from-{model}"}}])] + ) + + async def __anext__(self): + try: + return next(self._chunks) + except StopIteration: + raise StopAsyncIteration from None + + async def fake_acompletion(**kwargs): + attempted_model_groups.append(kwargs["metadata"]["model_group"]) + if "fb2" in kwargs["model"]: + return OkStream(kwargs["model"]) + return FailingStream(kwargs["model"]) + + router = litellm.Router( + model_list=[ + {"model_name": "primary", "litellm_params": {"model": "openai/primary-model", "api_key": "fake-key"}}, + {"model_name": "fb1", "litellm_params": {"model": "openai/fb1-model", "api_key": "fake-key"}}, + {"model_name": "fb2", "litellm_params": {"model": "openai/fb2-model", "api_key": "fake-key"}}, + ], + fallbacks=[{"primary": ["fb1", "fb2"]}], + num_retries=0, + ) + + with patch("litellm.acompletion", side_effect=fake_acompletion): + response = await router.acompletion(model="primary", messages=[{"role": "user", "content": "hi"}], stream=True) + content: Final = "".join( + [chunk.choices[0].delta.content or "" async for chunk in response if chunk is not None] + ) + + assert content == "ok-from-openai/fb2-model" + assert attempted_model_groups == ["primary", "fb1", "fb2"] + + def test_completion_streaming_iterator_adopts_fallback_response_headers(): """LIT-6767, sync counterpart of the fallback-adoption test.""" from unittest.mock import MagicMock, patch From 1098604ed661ee54b66ec7a378a6163f589b10df Mon Sep 17 00:00:00 2001 From: Joshua Valluru <326636767+joshua-berri@users.noreply.github.com> Date: Mon, 21 Sep 2026 12:24:15 -0700 Subject: [PATCH 16/44] refactor(mcp): extract explicit operation context and dispatch --- .github/workflows/test-linting.yml | 4 + Makefile | 1 + .../_experimental/mcp_server/contracts.py | 95 + .../mcp_server/legacy_callbacks.py | 83 + .../mcp_server/mcp_server_manager.py | 129 +- .../_experimental/mcp_server/operations.py | 3102 +++++++++++++++++ .../mcp_server/rest_endpoints.py | 13 +- .../proxy/_experimental/mcp_server/server.py | 3059 ++-------------- .../_experimental/mcp_server/tool_search.py | 11 +- litellm/proxy/_types.py | 2 + scripts/check_mcp_operation_boundary.py | 65 + scripts/pre_commit_lint.sh | 3 + .../mcp_server/test_byok_oauth_endpoints.py | 19 +- .../mcp_server/test_contracts.py | 60 + .../mcp_server/test_mcp_block_recording.py | 5 +- .../mcp_server/test_mcp_hook_extra_headers.py | 6 +- .../test_mcp_oauth_passthrough_tools.py | 61 +- .../mcp_server/test_mcp_proxy_mode.py | 8 +- .../mcp_server/test_mcp_server.py | 385 +- .../mcp_server/test_mcp_server_manager.py | 174 +- .../mcp_server/test_mcp_stale_session.py | 40 +- .../mcp_server/test_mcp_tool_search.py | 40 +- .../mcp_server/test_mcp_toolset_scope.py | 2 +- .../mcp_server/test_openapi_tool_auth.py | 75 +- .../mcp_server/test_operations.py | 343 ++ .../mcp_server/test_rest_endpoints.py | 13 +- .../test_check_mcp_operation_boundary.py | 52 + 27 files changed, 4624 insertions(+), 3226 deletions(-) create mode 100644 litellm/proxy/_experimental/mcp_server/contracts.py create mode 100644 litellm/proxy/_experimental/mcp_server/legacy_callbacks.py create mode 100644 litellm/proxy/_experimental/mcp_server/operations.py create mode 100644 scripts/check_mcp_operation_boundary.py create mode 100644 tests/test_litellm/proxy/_experimental/mcp_server/test_contracts.py create mode 100644 tests/test_litellm/proxy/_experimental/mcp_server/test_operations.py create mode 100644 tests/test_litellm/test_check_mcp_operation_boundary.py diff --git a/.github/workflows/test-linting.yml b/.github/workflows/test-linting.yml index 06d369eabcd..592d8edf6b8 100644 --- a/.github/workflows/test-linting.yml +++ b/.github/workflows/test-linting.yml @@ -130,6 +130,10 @@ jobs: echo "File content around line 43:" head -50 litellm/litellm_core_utils/custom_logger_registry.py | tail -10 + - name: Check MCP operation boundary + if: steps.changes.outputs.decision != 'skip' + run: uv run --no-sync python scripts/check_mcp_operation_boundary.py + - name: Run Ruff linting if: steps.changes.outputs.decision != 'skip' run: | diff --git a/Makefile b/Makefile index 0e9d2bbf82c..ab7fab6aa99 100644 --- a/Makefile +++ b/Makefile @@ -164,6 +164,7 @@ lint-format-check-changed: $(LINT_DEP_INSTALL) $(LINT_DEP_BASE) # Linting targets lint-ruff: $(LINT_DEP_INSTALL) + $(UV_RUN) python scripts/check_mcp_operation_boundary.py cd litellm && $(UV_RUN) ruff check . && cd .. $(UV_RUN) ruff check --config ruff-tests.toml tests diff --git a/litellm/proxy/_experimental/mcp_server/contracts.py b/litellm/proxy/_experimental/mcp_server/contracts.py new file mode 100644 index 00000000000..c3129d171ad --- /dev/null +++ b/litellm/proxy/_experimental/mcp_server/contracts.py @@ -0,0 +1,95 @@ +from collections.abc import Mapping +from copy import deepcopy +from dataclasses import dataclass, field +from datetime import datetime +from types import MappingProxyType +from typing import Final, Protocol + +from litellm.proxy._types import UserAPIKeyAuth +from litellm.types.mcp_server.mcp_server_manager import MCPServer + + +def copy_caller(auth: UserAPIKeyAuth | None) -> UserAPIKeyAuth | None: + if auth is None: + return None + span: Final = auth.parent_otel_span + return deepcopy(auth, {id(span): span} if span is not None else None) # mutable-ok: deepcopy mutates its memo + + +@dataclass(frozen=True, slots=True) +class OperationContext: + _caller: UserAPIKeyAuth | None = field(repr=False) + mcp_auth_header: str | None = field(default=None, repr=False) + mcp_servers: tuple[str, ...] | None = None + mcp_server_auth_headers: Mapping[str, Mapping[str, str]] | None = field(default=None, repr=False) + oauth2_headers: Mapping[str, str] | None = field(default=None, repr=False) + raw_headers: Mapping[str, str] | None = field(default=None, repr=False) + client_ip: str | None = None + mcp_proxy_mode: bool = False + + def __post_init__(self) -> None: + object.__setattr__(self, "_caller", copy_caller(self._caller)) + object.__setattr__(self, "mcp_servers", tuple(self.mcp_servers) if self.mcp_servers is not None else None) + object.__setattr__( + self, + "oauth2_headers", + MappingProxyType(dict(self.oauth2_headers)) if self.oauth2_headers is not None else None, + ) + object.__setattr__( + self, "raw_headers", MappingProxyType(dict(self.raw_headers)) if self.raw_headers is not None else None + ) + object.__setattr__( + self, + "mcp_server_auth_headers", + MappingProxyType( + {key: MappingProxyType(dict(value)) for key, value in self.mcp_server_auth_headers.items()} + ) + if self.mcp_server_auth_headers is not None + else None, + ) + + @property + def user_api_key_auth(self) -> UserAPIKeyAuth | None: + return copy_caller(self._caller) + + def legacy_auth( + self, + ) -> tuple[ + UserAPIKeyAuth | None, + str | None, + list[str] | None, # mutable-ok: detached legacy server-list payload + dict[str, dict[str, str]] | None, # mutable-ok: legacy auth dispatch requires concrete dict headers + dict[str, str] | None, # mutable-ok: detached legacy header payload + dict[str, str] | None, # mutable-ok: detached legacy header payload + str | None, + ]: + return ( + self.user_api_key_auth, + self.mcp_auth_header, + list(self.mcp_servers) if self.mcp_servers is not None else None, # mutable-ok: legacy policy list input + { + key: dict(value) for key, value in self.mcp_server_auth_headers.items() + } # mutable-ok: legacy auth dispatch checks concrete dict headers + if self.mcp_server_auth_headers is not None + else None, + dict(self.oauth2_headers) + if self.oauth2_headers is not None + else None, # mutable-ok: legacy OAuth header input + dict(self.raw_headers) if self.raw_headers is not None else None, # mutable-ok: legacy request header input + self.client_ip, + ) + + +class ProgressCallback(Protocol): + async def __call__(self, progress: float, total: float | None, /) -> None: ... + + +@dataclass(frozen=True, slots=True) +class AuthorizedToolCall: + name: str + arguments: Mapping[str, object] + allowed_mcp_servers: tuple[MCPServer, ...] + start_time: datetime + host_progress_callback: ProgressCallback | None + guardrail_context: Mapping[str, object] | None + logging_data: Mapping[str, object] diff --git a/litellm/proxy/_experimental/mcp_server/legacy_callbacks.py b/litellm/proxy/_experimental/mcp_server/legacy_callbacks.py new file mode 100644 index 00000000000..4424907c28a --- /dev/null +++ b/litellm/proxy/_experimental/mcp_server/legacy_callbacks.py @@ -0,0 +1,83 @@ +from collections.abc import Mapping +from typing import Final, Protocol + +from mcp.client.session import ClientRequestContext +from mcp.types import ( + CreateMessageRequestParams, + CreateMessageResult, + CreateMessageResultWithTools, + ElicitRequestParams, + ElicitResult, + ErrorData, +) + +from litellm.proxy._experimental.mcp_server.contracts import OperationContext +from litellm.proxy._types import UserAPIKeyAuth + + +class SamplingCallback(Protocol): + async def __call__( + self, context: ClientRequestContext, params: CreateMessageRequestParams, / + ) -> CreateMessageResult | CreateMessageResultWithTools | ErrorData: ... + + +class ElicitationCallback(Protocol): + async def __call__(self, context: object, params: ElicitRequestParams, /) -> ElicitResult | ErrorData: ... + + +def create_sampling_callback( + user_api_key_auth: UserAPIKeyAuth | None = None, + raw_headers: Mapping[str, str] | None = None, + client_ip: str | None = None, + operation_context: OperationContext | None = None, +) -> SamplingCallback: + from litellm.proxy._experimental.mcp_server.server import get_active_auth_context + + auth: Final = get_active_auth_context() if operation_context is None else None + captured: Final = ( + operation_context + if operation_context is not None + else OperationContext( + _caller=user_api_key_auth if user_api_key_auth is not None else (auth.user_api_key_auth if auth else None), + raw_headers=raw_headers if raw_headers is not None else (auth.raw_headers if auth else None), + client_ip=client_ip if client_ip is not None else (auth.client_ip if auth else None), + ) + ) + + async def callback( + context: ClientRequestContext, params: CreateMessageRequestParams + ) -> CreateMessageResult | CreateMessageResultWithTools | ErrorData: + import litellm + from litellm.proxy._experimental.mcp_server.sampling_handler import handle_sampling_create_message + + return await handle_sampling_create_message( + context=context, + params=params, + default_model=getattr(litellm, "default_mcp_sampling_model", None), + user_api_key_auth=captured.user_api_key_auth, + raw_headers=dict(captured.raw_headers) + if captured.raw_headers is not None + else None, # mutable-ok: handler consumes an owned request header dict + client_ip=captured.client_ip, + ) + + return callback + + +def create_elicitation_callback() -> ElicitationCallback: + from litellm.proxy._experimental.mcp_server.server import get_active_mcp_session + + downstream_session: Final = get_active_mcp_session() + downstream_capabilities: Final = getattr(downstream_session, "capabilities", None) + + async def callback(context: object, params: ElicitRequestParams) -> ElicitResult | ErrorData: + from litellm.proxy._experimental.mcp_server.elicitation_handler import handle_elicitation_request + + return await handle_elicitation_request( + context=context, + params=params, + downstream_session=downstream_session, + downstream_capabilities=downstream_capabilities, + ) + + return callback diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index b293ab5a206..264d7cf19aa 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -73,6 +73,7 @@ from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import ( MCPServerAccess, _is_mcp_admitted_user_subject, ) +from litellm.proxy._experimental.mcp_server.contracts import OperationContext from litellm.proxy._experimental.mcp_server.elicitation_handler import ( MCP_ELICITATION_AVAILABLE, ) @@ -195,9 +196,6 @@ from litellm.types.mcp_server.mcp_server_manager import ( from litellm.types.utils import CallTypes if TYPE_CHECKING: - from mcp.client.session import ClientRequestContext - from mcp.types import CreateMessageRequestParams - from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj from litellm.types.mcp_server.mcp_toolset import MCPToolset @@ -1218,7 +1216,7 @@ async def _resolve_byok_mcp_auth_header( if not mcp_server.is_byok: return mcp_auth_header - from litellm.proxy._experimental.mcp_server.server import ( + from litellm.proxy._experimental.mcp_server.operations import ( _check_byok_credential, _get_byok_credential, ) @@ -1577,77 +1575,25 @@ def _normalize_mcp_server_cost_info(mcp_info: MCPInfo) -> None: mcp_info["mcp_server_cost_info"] = normalized -def _create_sampling_callback(user_api_key_auth: UserAPIKeyAuth | None = None): - """ - Create a sampling callback for MCP ClientSession. - Returns a callable that handles sampling/createMessage requests from - upstream MCP servers by routing them through litellm.acompletion(). - """ +def _create_sampling_callback( + user_api_key_auth: UserAPIKeyAuth | None = None, + raw_headers: Mapping[str, str] | None = None, + client_ip: str | None = None, + operation_context: OperationContext | None = None, +): if not MCP_SAMPLING_AVAILABLE: return None + from litellm.proxy._experimental.mcp_server.legacy_callbacks import create_sampling_callback - async def _sampling_callback( - context: "ClientRequestContext", - params: "CreateMessageRequestParams", - ): - import litellm - from litellm.proxy._experimental.mcp_server.sampling_handler import ( - handle_sampling_create_message, - ) - from litellm.proxy._experimental.mcp_server.server import ( - get_active_auth_context, - ) - - auth_context: Final = get_active_auth_context() - resolved_auth: Final = user_api_key_auth or (auth_context.user_api_key_auth if auth_context else None) - # Forward original HTTP headers and client IP so that - # header-dependent guardrails, tag-based routing, trace - # correlation, and forward_llm_provider_auth_headers work - # correctly for sampling sub-calls. - _raw_headers: Final = getattr(auth_context, "raw_headers", None) - _client_ip: Final = getattr(auth_context, "client_ip", None) - - return await handle_sampling_create_message( - context=context, - params=params, - default_model=getattr(litellm, "default_mcp_sampling_model", None), - user_api_key_auth=resolved_auth, - raw_headers=_raw_headers, - client_ip=_client_ip, - ) - - return _sampling_callback + return create_sampling_callback(user_api_key_auth, raw_headers, client_ip, operation_context) def _create_elicitation_callback(): - """ - Create an elicitation callback for MCP ClientSession. - Returns a callable that handles elicitation/create requests from - upstream MCP servers. In gateway mode, this relays to the downstream - client; in tool bridge mode, it returns a decline response. - """ if not MCP_ELICITATION_AVAILABLE: return None + from litellm.proxy._experimental.mcp_server.legacy_callbacks import create_elicitation_callback - async def _elicitation_callback(context, params): - from litellm.proxy._experimental.mcp_server.elicitation_handler import ( - handle_elicitation_request, - ) - from litellm.proxy._experimental.mcp_server.server import get_active_mcp_session - - # In Gateway mode, we relay the elicitation request to the downstream client - # that triggered the current operation. - downstream_session: Final = get_active_mcp_session() - downstream_capabilities = getattr(downstream_session, "capabilities", None) if downstream_session else None - - return await handle_elicitation_request( - context=context, - params=params, - downstream_session=downstream_session, - downstream_capabilities=downstream_capabilities, - ) - - return _elicitation_callback + return create_elicitation_callback() def _record_mcp_guardrail_evaluations( @@ -3373,17 +3319,13 @@ class MCPServerManager: listable but uninvokable. Empty inside a toolset scope: toolset_mcp_route / dynamic_mcp_route set - ``_mcp_active_toolset_id`` before calling the handler, pinning the request to the toolset's + the caller's server-only ``mcp_toolset_id`` before calling the handler, pinning the request to the toolset's own servers (checking op.mcp_toolsets==[] instead would false-positive on DB-default rows where Postgres initialises the column to ARRAY[]::TEXT[]). ``allow_all_server_ids`` / ``submitted_server_ids`` are injectable so the server union, which precomputes both for its fallback path, does not compute them twice.""" - from litellm.proxy._experimental.mcp_server.mcp_context import ( # noqa: PLC0415 - _mcp_active_toolset_id, - ) - - if _mcp_active_toolset_id.get() is not None: + if user_api_key_auth is not None and user_api_key_auth.mcp_toolset_id is not None: return set() if allow_all_server_ids is None: allow_all_server_ids = self.get_allow_all_keys_server_ids() @@ -4151,6 +4093,8 @@ class MCPServerManager: subject_token: str | None = None, user_api_key_auth: UserAPIKeyAuth | None = None, cred_provider: UpstreamCredentialProvider | None = None, + raw_headers: Mapping[str, str] | None = None, + client_ip: str | None = None, ) -> MCPClient: """ Create an MCPClient instance for the given server. @@ -4199,7 +4143,13 @@ class MCPServerManager: # Create sampling and elicitation callbacks for this client sampling_cb = ( - _create_sampling_callback(user_api_key_auth=user_api_key_auth) if resolved_server.allow_sampling else None + _create_sampling_callback( + operation_context=OperationContext( + _caller=user_api_key_auth, raw_headers=raw_headers, client_ip=client_ip + ) + ) + if resolved_server.allow_sampling + else None ) elicitation_cb: Final = _create_elicitation_callback() if resolved_server.allow_elicitation else None @@ -4344,6 +4294,7 @@ class MCPServerManager: raw_headers: dict[str, str] | None = None, user_api_key_auth: UserAPIKeyAuth | None = None, oauth2_headers: dict[str, str] | None = None, + client_ip: str | None = None, ) -> list[MCPTool]: """ Helper method to get tools from a single MCP server with prefixed names. @@ -4433,6 +4384,8 @@ class MCPServerManager: stdio_env=stdio_env, subject_token=subject_token, user_api_key_auth=user_api_key_auth, + raw_headers=raw_headers, + client_ip=client_ip, ) ## HANDLE OPENAPI TOOLS @@ -4543,6 +4496,7 @@ class MCPServerManager: extra_headers: dict[str, str] | None = None, add_prefix: bool = True, raw_headers: dict[str, str] | None = None, + client_ip: str | None = None, ) -> list[Prompt]: try: headers: Final = ( @@ -4563,6 +4517,8 @@ class MCPServerManager: stdio_env=stdio_env, subject_token=subject_token, user_api_key_auth=user_api_key_auth, + raw_headers=raw_headers, + client_ip=client_ip, ) credential_fingerprint: Final = await client.discovery_auth_fingerprint() key: Final = self._discovery_key( @@ -4586,6 +4542,7 @@ class MCPServerManager: extra_headers: dict[str, str] | None = None, add_prefix: bool = True, raw_headers: dict[str, str] | None = None, + client_ip: str | None = None, ) -> list[Resource]: try: headers: Final = ( @@ -4606,6 +4563,8 @@ class MCPServerManager: stdio_env=stdio_env, subject_token=subject_token, user_api_key_auth=user_api_key_auth, + raw_headers=raw_headers, + client_ip=client_ip, ) credential_fingerprint: Final = await client.discovery_auth_fingerprint() key: Final = self._discovery_key( @@ -4629,6 +4588,7 @@ class MCPServerManager: extra_headers: dict[str, str] | None = None, add_prefix: bool = True, raw_headers: dict[str, str] | None = None, + client_ip: str | None = None, ) -> list[ResourceTemplate]: try: headers: Final = ( @@ -4649,6 +4609,8 @@ class MCPServerManager: stdio_env=stdio_env, subject_token=subject_token, user_api_key_auth=user_api_key_auth, + raw_headers=raw_headers, + client_ip=client_ip, ) credential_fingerprint: Final = await client.discovery_auth_fingerprint() key: Final = self._discovery_key( @@ -4672,6 +4634,7 @@ class MCPServerManager: mcp_auth_header: str | dict[str, str] | None = None, extra_headers: dict[str, str] | None = None, raw_headers: dict[str, str] | None = None, + client_ip: str | None = None, ) -> ReadResourceResult: """Read resource contents from a specific MCP server.""" @@ -4692,6 +4655,9 @@ class MCPServerManager: extra_headers=extra_headers, stdio_env=stdio_env, subject_token=subject_token, + raw_headers=raw_headers, + client_ip=client_ip, + user_api_key_auth=user_api_key_auth, ) return await client.read_resource(url) @@ -4705,6 +4671,7 @@ class MCPServerManager: mcp_auth_header: str | dict[str, str] | None = None, extra_headers: dict[str, str] | None = None, raw_headers: dict[str, str] | None = None, + client_ip: str | None = None, ) -> GetPromptResult: """Fetch a specific prompt definition from a single MCP server.""" @@ -4725,6 +4692,9 @@ class MCPServerManager: extra_headers=extra_headers, stdio_env=stdio_env, subject_token=subject_token, + raw_headers=raw_headers, + client_ip=client_ip, + user_api_key_auth=user_api_key_auth, ) get_prompt_request_params: Final = GetPromptRequestParams( @@ -5805,6 +5775,8 @@ class MCPServerManager: stdio_env: dict[str, str] | None, subject_token: str | None, user_api_key_auth: UserAPIKeyAuth | None, + raw_headers: Mapping[str, str] | None = None, + client_ip: str | None = None, ) -> CallToolResult: """Call a token_exchange (OBO) tool; on an upstream 401/403 re-mint the token once and retry. @@ -5830,6 +5802,8 @@ class MCPServerManager: stdio_env=stdio_env, subject_token=subject_token, user_api_key_auth=user_api_key_auth, + raw_headers=raw_headers, + client_ip=client_ip, ) return await retry_client.call_tool(call_tool_params, host_progress_callback=host_progress_callback) @@ -5847,6 +5821,7 @@ class MCPServerManager: host_progress_callback: Callable | None = None, hook_extra_headers: dict[str, str] | None = None, user_api_key_auth: UserAPIKeyAuth | None = None, + client_ip: str | None = None, ) -> CallToolResult: """ Call a regular MCP tool using the MCP client. @@ -5991,6 +5966,8 @@ class MCPServerManager: stdio_env=stdio_env, subject_token=subject_token, user_api_key_auth=user_api_key_auth, + raw_headers=raw_headers, + client_ip=client_ip, ) call_tool_params: Final = MCPCallToolRequestParams( @@ -6014,6 +5991,8 @@ class MCPServerManager: stdio_env=stdio_env, subject_token=subject_token, user_api_key_auth=user_api_key_auth, + raw_headers=raw_headers, + client_ip=client_ip, ) tool_call_coro = _obo_call_tool_limited() @@ -6189,7 +6168,7 @@ class MCPServerManager: return oauth2_headers try: - from litellm.proxy._experimental.mcp_server.server import ( # noqa: PLC0415 + from litellm.proxy._experimental.mcp_server.operations import ( # noqa: PLC0415 _get_user_oauth_extra_headers_from_db, ) @@ -6295,6 +6274,7 @@ class MCPServerManager: host_progress_callback: Callable | None = None, litellm_logging_obj: "LiteLLMLoggingObj | None" = None, guardrail_context: Mapping[str, object] | None = None, + client_ip: str | None = None, ) -> CallToolResult: """ Call a tool with the given name and arguments @@ -6421,6 +6401,7 @@ class MCPServerManager: mcp_server_auth_headers=mcp_server_auth_headers, oauth2_headers=oauth2_headers, raw_headers=raw_headers, + client_ip=client_ip, proxy_logging_obj=proxy_logging_obj, host_progress_callback=host_progress_callback, hook_extra_headers=hook_result.get("extra_headers"), diff --git a/litellm/proxy/_experimental/mcp_server/operations.py b/litellm/proxy/_experimental/mcp_server/operations.py new file mode 100644 index 00000000000..3ee4f37add7 --- /dev/null +++ b/litellm/proxy/_experimental/mcp_server/operations.py @@ -0,0 +1,3102 @@ +"""Shared MCP operation policy and dispatch.""" + +import asyncio +import traceback +import types +import uuid +from collections.abc import Mapping, Sequence +from datetime import datetime +from typing import Any, Final, NoReturn, TypeAlias, assert_never, overload + +from fastapi import HTTPException +from mcp import ReadResourceResult, Resource +from mcp.types import ( + CallToolRequest, + CallToolRequestParams, + CallToolResult, + GetPromptRequest, + GetPromptRequestParams, + GetPromptResult, + ListPromptsRequest, + ListPromptsResult, + ListResourcesRequest, + ListResourcesResult, + ListResourceTemplatesRequest, + ListResourceTemplatesResult, + ListToolsRequest, + ListToolsResult, + PaginatedRequestParams, + Prompt, + ReadResourceRequest, + ReadResourceRequestParams, + ResourceTemplate, + TextContent, +) +from mcp.types import Tool as MCPTool +from pydantic import AnyUrl, ConfigDict, Field, TypeAdapter +from typing_extensions import ReadOnly, TypedDict + +from litellm._logging import verbose_logger +from litellm.constants import ( + MAXIMUM_TRACEBACK_LINES_TO_LOG, +) +from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj +from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import ( + MCPRequestHandler, +) +from litellm.proxy._experimental.mcp_server.byok_credential_cache import ( + byok_credential_cache, + byok_credential_cache_key, + cache_byok_credential, + get_cached_byok_credential, +) +from litellm.proxy._experimental.mcp_server.contracts import ( + AuthorizedToolCall, + OperationContext, + ProgressCallback, +) +from litellm.proxy._experimental.mcp_server.db import OAuthCredentialPayload +from litellm.proxy._experimental.mcp_server.exceptions import ( + MCPToolResultError, + MCPUpstreamAuthError, +) +from litellm.proxy._experimental.mcp_server.faults.list_outcomes import ( + SERVER_OUTCOMES_META_KEY, + AggregateToolListing, + ServerListOk, + ServerOutcome, + classify_list_exception, + outcome_wire_value, +) +from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + MCPServerManager, + _caller_authorization_fans_out, + _client_forwarded_authorization_headers, + _resolve_openapi_tool_auth, + _should_strip_caller_authorization, + global_mcp_server_manager, +) +from litellm.proxy._experimental.mcp_server.oauth_utils import ( + _redact_mcp_resource_url, + get_byok_www_authenticate, +) +from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import ( + _request_auth_header, + _request_extra_headers, + _request_resolved_auth_headers, +) +from litellm.proxy._experimental.mcp_server.tool_registry import ( + global_mcp_tool_registry, +) +from litellm.proxy._experimental.mcp_server.utils import ( + MCP_TOOL_PREFIX_SEPARATOR, + MCPMissingUserEnvVarsError, + add_server_prefix_to_name, + build_synthetic_mcp_request, + extract_mcp_tool_result_error_message, + get_server_prefix, + is_tool_name_prefixed, + iter_known_server_prefixes, + logging_safe_mcp_headers, + match_known_tool_name, + normalize_server_name, + split_server_prefix_from_name, + strip_known_server_prefix, +) +from litellm.proxy._types import ( + UserAPIKeyAuth, +) +from litellm.proxy.common_utils.auth_cache_invalidation_pubsub import ( + publish_auth_cache_invalidation, +) +from litellm.proxy.litellm_pre_call_utils import ( + LiteLLMProxyRequestSetup, + get_chain_id_from_headers, +) +from litellm.types.mcp import ( + DEFAULT_CREDENTIAL_HEADER, + MCPAuth, + without_header, +) +from litellm.types.mcp_server.mcp_server_manager import MCPInfo, MCPServer +from litellm.types.utils import CallTypes, StandardLoggingMCPToolCall +from litellm.utils import Rules, client, function_setup + +__all__ = ( + "_MCP_CREDENTIAL_REQUEST_FIELDS", + "ListMCPToolsRestAPIResponseObject", + "MCPInfo", + "MCPServer", + "_McpDeniedDetail", + "_aggregate_server_key", + "_build_virtual_call_logging_obj", + "_check_byok_credential", + "_client_has_passthrough_authorization", + "_client_has_per_server_auth_header", + "_dispatch_virtual_mcp_tool", + "_fire_mcp_tool_call_logging", + "_get_allowed_mcp_servers", + "_get_allowed_mcp_servers_from_mcp_server_names", + "_get_byok_credential", + "_get_prompts_from_mcp_servers", + "_get_resource_templates_from_mcp_servers", + "_get_resources_from_mcp_servers", + "_get_standard_logging_mcp_tool_call", + "_get_tools_from_mcp_servers", + "_get_user_oauth_extra_headers_from_db", + "_handle_local_mcp_tool", + "_handle_managed_mcp_tool", + "_http_detail_message", + "_invalidate_byok_cred_cache", + "_list_mcp_prompts", + "_list_mcp_resource_templates", + "_list_mcp_resources", + "_list_mcp_tools", + "_list_tools_before_first_call", + "_mcp_session_id_from_headers", + "_merge_gateway_initialize_instructions", + "_prefetch_oauth_creds_for_user", + "_prepare_mcp_server_headers", + "_raise_if_initialize_grants_no_mcp_servers", + "_resolve_display_name_to_original", + "_run_post_mcp_call_guardrails", + "_server_answers_to", + "_tool_name_matches", + "apply_tool_overrides", + "call_mcp_tool", + "execute_mcp_tool", + "filter_tools_by_allowed_tools", + "filter_tools_by_key_team_permissions", + "fire_mcp_tool_call_failure_logging", + "mcp_get_prompt", + "mcp_read_resource", + "raise_denied_scoped_mcp_access", +) + + +async def _invalidate_byok_cred_cache(user_id: str, server_id: str) -> None: + """Drop a stored-or-deleted BYOK credential from this worker's cache and from every peer worker's.""" + cache_key: Final = byok_credential_cache_key(user_id, server_id) + byok_credential_cache.delete_cache(cache_key) + await publish_auth_cache_invalidation(cache_key=cache_key) + + +def _mcp_session_id_from_headers( + raw_headers: dict[str, str] | None, +) -> str | None: + """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 + + +class ListMCPToolsRestAPIResponseObject(MCPTool): + """ + Object returned by the /tools/list REST API route. + """ + + mcp_info: MCPInfo | None = Field(default=None, alias="mcp_info") + model_config = ConfigDict(arbitrary_types_allowed=True) + + +async def _build_virtual_call_logging_obj( + name: str, + arguments: dict[str, object], + user_api_key_auth: UserAPIKeyAuth, + raw_headers: Mapping[str, str] | None = None, + client_ip: str | None = None, +) -> LiteLLMLoggingObj | None: + """Run the pre-call pipeline (guardrails + logging setup) for a virtual + mcp_tool_call so the SSE path spend-logs like the REST path.""" + from litellm.proxy.common_request_processing import ( + ProxyBaseLLMRequestProcessing, + ) + from litellm.proxy.proxy_server import ( + general_settings, + proxy_config, + proxy_logging_obj, + ) + + request: Final = build_synthetic_mcp_request( + path="/mcp/tools/call", + raw_headers=raw_headers, + client_ip=client_ip, + ) + _, virtual_logging_obj = await ProxyBaseLLMRequestProcessing( + data={"name": name, "arguments": arguments} + ).common_processing_pre_call_logic( + request=request, + user_api_key_dict=user_api_key_auth, + proxy_config=proxy_config, + route_type=CallTypes.call_mcp_tool.value, + proxy_logging_obj=proxy_logging_obj, + general_settings=general_settings, + ) + return virtual_logging_obj + + +async def _dispatch_virtual_mcp_tool( + name: str, + arguments: dict[str, object] | None, + user_api_key_auth: UserAPIKeyAuth | None, + client_ip: str | None, + mcp_servers: list[str] | None = None, + mcp_auth_header: str | None = None, + mcp_server_auth_headers: dict[str, dict[str, str]] | None = None, + oauth2_headers: dict[str, str] | None = None, + raw_headers: dict[str, str] | None = None, + mcp_proxy_mode: bool = False, +) -> CallToolResult | None: + """Handle the mcp_tool_search / mcp_tool_call virtual tools. + + Returns a CallToolResult when ``name`` is a virtual tool, else ``None`` so + the caller falls through to normal tool routing. + """ + from litellm.llms.litellm_proxy.skills.skill_search import DEFAULT_SKILL_SEARCH_TOP_K + from litellm.proxy._experimental.mcp_server.tool_search import ( + AGENT_SEARCH_TOOL_NAME, + DEFAULT_AGENT_SEARCH_TOP_K, + MCP_PROXY_CALL_TOOL_NAME, + MCP_PROXY_TOOL_NAMES, + MCP_TOOL_SEARCH_TOOL_NAME, + SKILL_SEARCH_TOOL_NAME, + VIRTUAL_TOOL_NAMES, + coerce_top_k, + handle_agent_search, + handle_mcp_proxy_tool, + handle_mcp_tool_call, + handle_mcp_tool_search, + handle_skill_search, + ) + + if mcp_proxy_mode and name not in MCP_PROXY_TOOL_NAMES: + return CallToolResult( + content=[ # mutable-ok: MCP result content + TextContent(type="text", text=f"Tool {name} is unavailable on /mcp/proxy") + ], + is_error=True, + ) + + if mcp_proxy_mode and name in MCP_PROXY_TOOL_NAMES: + assert user_api_key_auth is not None + proxy_call_start: Final = datetime.now() # noqa: DTZ005 # logging pipeline uses naive datetimes + proxy_logging_obj: Final = ( + await _build_virtual_call_logging_obj( + name=name, + arguments=arguments or {}, # mutable-ok: logging pipeline payload + user_api_key_auth=user_api_key_auth, + raw_headers=raw_headers, + client_ip=client_ip, + ) + if name == MCP_PROXY_CALL_TOOL_NAME + else None + ) + try: + proxy_result: Final = await handle_mcp_proxy_tool( + name=name, + arguments=arguments or {}, # mutable-ok: proxy handler payload + user_api_key_dict=user_api_key_auth, + client_ip=client_ip, + mcp_servers=mcp_servers, + mcp_auth_header=mcp_auth_header, + mcp_server_auth_headers=mcp_server_auth_headers, + oauth2_headers=oauth2_headers, + raw_headers=raw_headers, + litellm_logging_obj=proxy_logging_obj, + ) + except Exception as exc: + if proxy_logging_obj is not None: + from litellm.proxy.proxy_server import proxy_logging_obj as request_logging_obj + + failure_end: Final = datetime.now() # noqa: DTZ005 # matches the logging pipeline start time + failure_traceback: Final = traceback.format_exc(limit=MAXIMUM_TRACEBACK_LINES_TO_LOG) + try: + proxy_logging_obj.failure_handler(exc, failure_traceback, proxy_call_start, failure_end) + await proxy_logging_obj.async_failure_handler(exc, failure_traceback, proxy_call_start, failure_end) + if not isinstance(exc, MCPUpstreamAuthError): + await request_logging_obj.post_call_failure_hook( + request_data={ # mutable-ok: failure hook mutates its request payload + "name": name, + "arguments": arguments, + "litellm_logging_obj": proxy_logging_obj, + }, + original_exception=exc, + user_api_key_dict=user_api_key_auth, + route="/mcp/call_tool", + traceback_str=failure_traceback, + ) + except Exception: # noqa: BLE001 # a failing failure hook must not mask the tool call's own error + verbose_logger.exception("Error logging failed MCP proxy tool call") + raise + if proxy_logging_obj is not None: + return await _fire_mcp_tool_call_logging( + logging_obj=proxy_logging_obj, + result=proxy_result, + start_time=proxy_call_start, + end_time=datetime.now(), # noqa: DTZ005 # matches the logging pipeline start time + user_api_key_auth=user_api_key_auth, + request_data=types.MappingProxyType({"name": name, "arguments": arguments}), + ) + return proxy_result + + if name not in VIRTUAL_TOOL_NAMES: + return None + + if not getattr( + getattr(user_api_key_auth, "object_permission", None), + "mcp_tool_search_enabled", + False, + ): + return CallToolResult( + content=[ + TextContent( + type="text", + text=f"Tool {name} requires mcp_tool_search_enabled on the key", + ) + ], + is_error=True, + ) + + args: Final = arguments or {} + if name == MCP_TOOL_SEARCH_TOOL_NAME: + return await handle_mcp_tool_search( + query=TypeAdapter(str).validate_python(args.get("query", "")), + top_k=coerce_top_k(args.get("top_k", 5)), + user_api_key_dict=user_api_key_auth, + client_ip=client_ip, + mcp_servers=mcp_servers, + mcp_auth_header=mcp_auth_header, + mcp_server_auth_headers=mcp_server_auth_headers, + oauth2_headers=oauth2_headers, + raw_headers=raw_headers, + ) + + assert user_api_key_auth is not None # guaranteed by the flag check above + if name == AGENT_SEARCH_TOOL_NAME: + return await handle_agent_search( + query=str(args.get("query", "")), + top_k=coerce_top_k(args.get("top_k", DEFAULT_AGENT_SEARCH_TOP_K), default=DEFAULT_AGENT_SEARCH_TOP_K), + user_api_key_dict=user_api_key_auth, + ) + if name == SKILL_SEARCH_TOOL_NAME: + return await handle_skill_search( + query=str(args.get("query", "")), + top_k=coerce_top_k(args.get("top_k", DEFAULT_SKILL_SEARCH_TOP_K), default=DEFAULT_SKILL_SEARCH_TOP_K), + user_api_key_dict=user_api_key_auth, + ) + virtual_logging_obj: Final = await _build_virtual_call_logging_obj( + name=name, + arguments=args, + user_api_key_auth=user_api_key_auth, + raw_headers=raw_headers, + client_ip=client_ip, + ) + tool_request: Final = CallToolRequestParams.model_validate( + types.MappingProxyType({"name": args.get("tool_name", ""), "arguments": args.get("arguments") or {}}) + ) + return await handle_mcp_tool_call( + tool_name=tool_request.name, + arguments=tool_request.arguments or {}, + user_api_key_dict=user_api_key_auth, + client_ip=client_ip, + mcp_servers=mcp_servers, + mcp_auth_header=mcp_auth_header, + mcp_server_auth_headers=mcp_server_auth_headers, + oauth2_headers=oauth2_headers, + raw_headers=raw_headers, + litellm_logging_obj=virtual_logging_obj, + ) + + +async def _get_allowed_mcp_servers_from_mcp_server_names( + mcp_servers: Sequence[str] | None, + allowed_mcp_servers: list[MCPServer], +) -> list[MCPServer]: + """ + Get the filtered MCP servers from the MCP server names. + + Fails closed when ``mcp_servers`` is explicitly provided (path- or + header-derived) but none of the names resolve to a server alias or + access group the caller can access. The previous behavior returned + the full ``allowed_mcp_servers`` set, which silently widened scope + when a client targeted ``/mcp//`` and made URL/header + namespacing appear to work when it did not. + """ + + filtered_server: Final[dict[str, MCPServer]] = {} + # Filter servers based on mcp_servers parameter if provided + if mcp_servers is not None: + for server_or_group in mcp_servers: + server_name_matched = False + + for server in allowed_mcp_servers: + if server and _server_answers_to(server, server_or_group): + filtered_server[server.server_id] = server + server_name_matched = True + break + + if not server_name_matched: + try: + access_group_server_ids = await MCPRequestHandler._get_mcp_servers_from_access_groups( + [server_or_group] + ) + # Only include servers that the user has access to + for server_id in access_group_server_ids: + for server in allowed_mcp_servers: + if server_id == server.server_id: + filtered_server[server.server_id] = server + except Exception as e: + verbose_logger.debug("Could not resolve '%s' as access group: %s", server_or_group, e) + + if filtered_server: + return list(filtered_server.values()) + + if mcp_servers is not None: + # Caller asked for a specific scope but nothing resolved. Fail + # closed so URL/header namespacing cannot silently fall back to + # the caller's full allowed-server set. + verbose_logger.debug( + "MCP scope filter resolved to no servers for requested names %s; returning empty list (fail-closed).", + mcp_servers, + ) + return [] + + return allowed_mcp_servers + + +def _http_detail_message(detail: object) -> str: + return str(detail.get("error")) if isinstance(detail, dict) and detail.get("error") else str(detail) + + +def _server_answers_to(server: MCPServer, name: str) -> bool: + requested: Final = name.lower() + return any(requested == known.lower() for known in iter_known_server_prefixes(server) if known) + + +async def raise_denied_scoped_mcp_access( + requested_names: Sequence[str], + user_api_key_auth: UserAPIKeyAuth | None, + client_ip: str | None = None, +) -> None: + """A scoped request (``/mcp/`` path or ``x-mcp-servers`` header) resolved to zero + allowed servers, so the denial must be loud: a silent 200 with no tools reads as a healthy + server with no tools. Unknown, unauthorized, and access-group names all share one generic + error so scoping cannot probe which servers exist; the agent variant fires only when the + same request resolves once the agent binding is stripped, proving the binding caused the veto.""" + agent_id: Final = user_api_key_auth.agent_id if user_api_key_auth else None + if user_api_key_auth is not None and agent_id: + resolved_without_agent: Final = await _get_allowed_mcp_servers( + user_api_key_auth=user_api_key_auth.model_copy(update=types.MappingProxyType({"agent_id": None})), + mcp_servers=requested_names, + client_ip=client_ip, + ) + + def _resolved_to_server(name: str) -> bool: + return any(_server_answers_to(server, name) for server in resolved_without_agent) + + vetoed_server: Final = next((name for name in requested_names if _resolved_to_server(name)), None) + if vetoed_server is not None: + agent_denial: Final[_McpDeniedDetail] = { + "error": ( + f"MCP server '{vetoed_server}' is not available to this key: the key is bound to " + f"agent '{agent_id}', whose MCP grants do not include this server. Add the server " + f"to the agent's object_permission.mcp_servers (edit the agent in the Admin UI or " + f"PATCH /v1/agents/{agent_id}), or use a key that is not bound to the agent." + ) + } + raise HTTPException(status_code=403, detail=agent_denial) + vetoed_group: Final = next( + ( + name + for name in requested_names + if not _resolved_to_server(name) + and any(name in (server.access_groups or ()) for server in resolved_without_agent) + ), + None, + ) + if vetoed_group is not None: + group_denial: Final[_McpDeniedDetail] = { + "error": ( + f"MCP access group '{vetoed_group}' is not available to this key: the key is bound to " + f"agent '{agent_id}', whose MCP grants do not include it. Add the group to the " + f"agent's object_permission.mcp_access_groups (edit the agent in the Admin UI or " + f"PATCH /v1/agents/{agent_id}), or use a key that is not bound to the agent." + ) + } + raise HTTPException(status_code=403, detail=group_denial) + generic_denial: Final[_McpDeniedDetail] = { + "error": f"The key is not allowed to access the requested MCP servers: {', '.join(requested_names)}" + } + raise HTTPException(status_code=403, detail=generic_denial) + + +def _tool_name_matches(tool_name: str, filter_list: list[str], mcp_server: MCPServer) -> bool: + """ + Check if a tool name matches any name in the filter list. + + Reads the same owner the server-level permission checks use, so discovery hides + exactly what dispatch refuses. ``mcp_server`` is required: guessing the boundary + at the first separator mismatches every tool on a server whose prefix contains + the separator. + """ + bare_name: Final = strip_known_server_prefix(tool_name, mcp_server) + return match_known_tool_name(bare_name, mcp_server, filter_list) is not None + + +def filter_tools_by_allowed_tools( + tools: list[MCPTool], + mcp_server: MCPServer, +) -> list[MCPTool]: + """ + Filter tools by allowed/disallowed tools configuration. + + If allowed_tools is set, only tools in that list are returned. + If disallowed_tools is set, tools in that list are excluded. + Tool names are matched with and without server prefixes for flexibility. + + Args: + tools: List of tools to filter + mcp_server: Server configuration with allowed_tools/disallowed_tools + + Returns: + Filtered list of tools + """ + from litellm.proxy._experimental.mcp_server.utils import ( + server_applies_tool_allowlist, + ) + + tools_to_return = tools + + # Filter by allowed_tools (whitelist) + if server_applies_tool_allowlist(mcp_server): + if not mcp_server.allowed_tools: + return [] + tools_to_return = [ + tool for tool in tools if _tool_name_matches(tool.name, mcp_server.allowed_tools, mcp_server) + ] + + # Filter by disallowed_tools (blacklist) + if mcp_server.disallowed_tools: + tools_to_return = [ + tool + for tool in tools_to_return + if not _tool_name_matches(tool.name, mcp_server.disallowed_tools, mcp_server) + ] + + return tools_to_return + + +def apply_tool_overrides( + tools: list[MCPTool], + mcp_server: MCPServer, +) -> list[MCPTool]: + """Apply admin-configured display name/description overrides to tools. + + Overrides are keyed by the unprefixed tool name, same convention as + allowed_tools configuration. + """ + display_name_map: Final = mcp_server.tool_name_to_display_name or {} + description_map: Final = mcp_server.tool_name_to_description or {} + if not display_name_map and not description_map: + return tools + + for tool in tools: + unprefixed = strip_known_server_prefix(tool.name, mcp_server) + lookup_key = unprefixed or tool.name + if lookup_key in display_name_map: + tool.name = display_name_map[lookup_key] + if lookup_key in description_map: + tool.description = description_map[lookup_key] + return tools + + +async def _get_allowed_mcp_servers( + user_api_key_auth: UserAPIKeyAuth | None, + mcp_servers: Sequence[str] | None, + client_ip: str | None = None, +) -> list[MCPServer]: + """Return allowed MCP servers for a request after applying filters. + + Args: + user_api_key_auth: The authenticated user's API key info. + mcp_servers: Optional list of server names to filter to. + client_ip: Client IP for IP-based access control. If None, falls back to + auth context. Pass explicitly from request handlers for safety. + Note: If client_ip is None and auth context is not set, IP filtering is skipped. + This is intentional for internal callers but may indicate a bug if called + from a request handler without proper context setup. + """ + allowed_mcp_server_ids = await global_mcp_server_manager.get_allowed_mcp_servers(user_api_key_auth) + ( + allowed_mcp_server_ids, + _ip_blocked, + ) = global_mcp_server_manager.filter_server_ids_by_ip_with_info(allowed_mcp_server_ids, client_ip) + verbose_logger.debug( + "MCP IP filter: client_ip=%s, allowed_server_ids=%s", + client_ip, + allowed_mcp_server_ids, + ) + if _ip_blocked > 0: + verbose_logger.debug( + "MCP IP filtering: %d server(s) are not accessible from client IP %s " + "because they are restricted to internal networks. " + "No tools from those servers will be returned. " + "To expose a server externally, set 'available_on_public_internet: true' " + "in its configuration.", + _ip_blocked, + client_ip, + ) + allowed_mcp_servers: list[MCPServer] = [] + for allowed_mcp_server_id in allowed_mcp_server_ids: + mcp_server = global_mcp_server_manager.get_mcp_server_by_id(allowed_mcp_server_id) + if mcp_server is not None: + # Apply the request-time oauth2_flow backstop for legacy null rows. + mcp_server = MCPServerManager.resolve_oauth2_flow_for_request(mcp_server) + allowed_mcp_servers.append(mcp_server) + + if mcp_servers is not None: + allowed_mcp_servers = await _get_allowed_mcp_servers_from_mcp_server_names( + mcp_servers=mcp_servers, + allowed_mcp_servers=allowed_mcp_servers, + ) + + return allowed_mcp_servers + + +def _client_has_per_server_auth_header( + server: MCPServer, + mcp_server_auth_headers: dict[str, dict[str, str]] | None, +) -> bool: + """True if the request carries a per-server ``x-mcp-{alias}-authorization`` + header for this server. This is the multi-server binding: it names one + upstream, so it is unambiguously the caller's upstream token regardless of + auth mode (never the LiteLLM admission credential). + + Resolves through the same ``lookup_mcp_server_auth_in_headers`` egress uses, so + the connect gate and egress agree on which per-server header names match: a + dashboard client sends ``x-mcp-{sanitize_mcp_alias_for_header(alias)}-authorization``, + and matching only the raw alias here would 401 a token egress would forward. + """ + if not mcp_server_auth_headers: + return False + from litellm.proxy._experimental.mcp_server.utils import ( + lookup_mcp_server_auth_in_headers, + ) + + server_headers: Final = lookup_mcp_server_auth_in_headers( + mcp_server_auth_headers, + alias=server.alias, + server_name=server.server_name, + access_groups=server.access_groups, + ) + if isinstance(server_headers, str): + return bool(server_headers.strip()) + if isinstance(server_headers, dict): + return any(isinstance(hk, str) and hk.lower() == "authorization" for hk in server_headers) + return False + + +def _client_has_passthrough_authorization( + server: MCPServer, + oauth2_headers: dict[str, str] | None, + mcp_server_auth_headers: dict[str, dict[str, str]] | None, +) -> 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: + if k.lower() == "authorization": + return True + return _client_has_per_server_auth_header(server, mcp_server_auth_headers) + + +async def _get_user_oauth_extra_headers_from_db( + server: MCPServer, + user_api_key_auth: UserAPIKeyAuth | None, + prefetched_creds: 'Mapping[str, "OAuthCredentialPayload"] | None' = None, +) -> dict[str, str] | None: + """Stored OAuth2 token for (user, server) as an ``Authorization: Bearer`` header, or None. + + Thin wrapper over ``resolve_user_oauth_access_token`` (Redis cache, else DB + refresh); + ``prefetched_creds`` skips the per-server Redis/DB lookups for the batch path. + """ + if server.auth_type != MCPAuth.oauth2 or user_api_key_auth is None: + return None + from litellm.proxy._experimental.mcp_server.db import ( # noqa: PLC0415 + resolve_user_oauth_access_token, + ) + + token: Final = await resolve_user_oauth_access_token( + getattr(user_api_key_auth, "user_id", None), server, prefetched_creds + ) + return {"Authorization": f"Bearer {token}"} if token else None + + +async def _prefetch_oauth_creds_for_user( + user_api_key_auth: UserAPIKeyAuth | None, +) -> dict[str, "OAuthCredentialPayload"]: + """Fetch all OAuth2 credentials for the user in one DB query. + + Returns a dict keyed by server_id to avoid N+1 queries in asyncio.gather loops. + """ + user_id: Final[str | None] = getattr(user_api_key_auth, "user_id", None) if user_api_key_auth else None + if not user_id: + return {} + try: + from litellm.proxy._experimental.mcp_server.db import ( # noqa: PLC0415 + list_user_oauth_credentials, + ) + from litellm.proxy.utils import get_prisma_client_or_throw # noqa: PLC0415 + + prisma_client: Final = get_prisma_client_or_throw( + "Database not connected. Connect a database to use OAuth2 MCP tools." + ) + creds: Final = await list_user_oauth_credentials(prisma_client, user_id) + return {c["server_id"]: c for c in creds if "server_id" in c} + except Exception as e: + verbose_logger.warning("_prefetch_oauth_creds_for_user: failed to prefetch for user=%s: %s", user_id, e) + return {} + + +def _prepare_mcp_server_headers( + server: MCPServer, + mcp_server_auth_headers: dict[str, dict[str, str]] | None, + mcp_auth_header: str | None, + oauth2_headers: dict[str, str] | None, + raw_headers: dict[str, str] | None, + user_api_key_auth: UserAPIKeyAuth | None = None, + scope_servers: list[MCPServer] | None = None, +) -> tuple[dict[str, str] | str | None, dict[str, str] | None]: + """Build auth and extra headers for a server. + + ``scope_servers`` is the full server list a fan-out handler iterates. Passing it lets the + client-forwarded token modes withhold the caller's request-wide ``Authorization`` when + another server in the scope would also receive it (``_caller_authorization_fans_out``); + explicitly-addressed operations leave it None. Per-server ``x-mcp-{alias}-authorization`` + headers are unaffected — they bind one token to one server and are the multi-server shape. + """ + server_auth_header: dict[str, str] | str | None = None + if mcp_server_auth_headers: + from litellm.proxy._experimental.mcp_server.utils import ( + lookup_mcp_server_auth_in_headers, + ) + + server_auth_header = lookup_mcp_server_auth_in_headers( + mcp_server_auth_headers, + alias=server.alias, + server_name=server.server_name, + access_groups=server.access_groups, + ) + + extra_headers: dict[str, str] | None = None + is_client_forwarded_mode: Final = server.is_client_forwarded_token + # In a multi-server listing scope the request-wide Authorization can only carry one token, + # so it is withheld from a client-forwarded server when another server in scope also consumes + # it (RFC 9700 cross-resource replay); such scopes must bind per-server via + # x-mcp-{alias}-authorization. The decision is computed once so BOTH the forwarding branch and + # the extra_headers copy loop below honor it — otherwise a server that lists Authorization in + # extra_headers would re-copy the withheld bearer from raw_headers and replay it anyway. + withhold_forwarded_authorization: Final = is_client_forwarded_mode and _caller_authorization_fans_out( + server, scope_servers + ) + if server.auth_type == MCPAuth.oauth2: + # For OAuth2 M2M servers, upstream Authorization must come from + # client_credentials token fetch, never from caller headers. + if server.has_client_credentials: + extra_headers = None + else: + # Copy to avoid mutating the original dict (important for parallel fetching) + extra_headers = oauth2_headers.copy() if oauth2_headers else None + # Migrated authorization_code: the v2 resolver injects the stored per-user + # token, so drop the caller-forwarded Authorization (apply-if-absent would + # otherwise let it shadow the resolved token). Delegate keeps it. Centralized + # via _should_strip_caller_authorization to match _call_regular_mcp_tool. + if extra_headers and _should_strip_caller_authorization( + mcp_server=server, + raw_headers=raw_headers, + user_api_key_auth=user_api_key_auth, + ): + extra_headers = without_header(extra_headers, DEFAULT_CREDENTIAL_HEADER) + elif is_client_forwarded_mode: + if not withhold_forwarded_authorization: + extra_headers = _client_forwarded_authorization_headers( + mcp_server=server, + oauth2_headers=oauth2_headers, + raw_headers=raw_headers, + user_api_key_auth=user_api_key_auth, + ) + + if server.extra_headers and raw_headers: + if extra_headers is None: + extra_headers = {} + + normalized_raw_headers: Final = {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: Final = _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 header.lower() == "authorization" and (strip_caller_authorization or withhold_forwarded_authorization): + continue + header_value = normalized_raw_headers.get(header.lower()) + if header_value is None: + continue + extra_headers[header] = header_value + + # Reset to None if no headers were actually added + if extra_headers is not None and len(extra_headers) == 0: + extra_headers = None + + if server_auth_header is None: + server_auth_header = mcp_auth_header + + return server_auth_header, extra_headers + + +def _merge_gateway_initialize_instructions( + allowed_mcp_servers: list[MCPServer], +) -> str | None: + """YAML/DB override, else upstream text (prefetch on init, or list_tools / health_check / call_tool cache).""" + if not allowed_mcp_servers: + return None + + texts: Final[list[tuple[str, str]]] = [] + for server in allowed_mcp_servers: + label = server.alias or server.server_name or server.name or server.server_id or "mcp" + if server.instructions and server.instructions.strip(): + texts.append((label, server.instructions.strip())) + continue + if server.spec_path: + continue + cached = global_mcp_server_manager._upstream_initialize_instructions_by_server_id.get(server.server_id) + if cached and cached.strip(): + texts.append((label, cached.strip())) + + if not texts: + return None + if len(texts) == 1: + return texts[0][1] + return "\n\n---\n\n".join(f"[{lbl}]\n{txt}" for lbl, txt in texts) + + +async def _raise_if_initialize_grants_no_mcp_servers( + allowed: Sequence[MCPServer], + user_api_key_auth: UserAPIKeyAuth | None, + mcp_servers: Sequence[str] | None, + client_ip: str | None, +) -> None: + if allowed or user_api_key_auth is None or not user_api_key_auth.api_key: + return + if mcp_servers: + await raise_denied_scoped_mcp_access( + requested_names=mcp_servers, + user_api_key_auth=user_api_key_auth, + client_ip=client_ip, + ) + no_servers_denial: Final[_McpDeniedDetail] = { + "error": ( + "The key has no MCP servers granted, or none of its granted servers is loaded and allowed for " + "this client IP. Grant servers or access groups to the key, its team, or its organization " + "(object_permission.mcp_servers), check the server's allowed IPs, and reconnect." + ) + } + raise HTTPException(status_code=403, detail=no_servers_denial) + + +def _aggregate_server_key(server: MCPServer) -> str: + """The client-visible key for a server in listing outcomes and spend metadata: the same + display prefix (alias, or the short prefix when that mode is enabled) the caller already + sees on the tool names. Canonical internal server names never key a caller-readable + surface; when the display naming deliberately hides them, the outcome keys must too.""" + return get_server_prefix(server) or "unknown" + + +async def _get_tools_from_mcp_servers( + user_api_key_auth: UserAPIKeyAuth | None, + mcp_auth_header: str | None, + mcp_servers: list[str] | None, + mcp_server_auth_headers: dict[str, dict[str, str]] | None = None, + oauth2_headers: dict[str, str] | None = None, + raw_headers: dict[str, str] | None = None, + log_list_tools_to_spendlogs: bool = False, + list_tools_log_source: str | None = None, + litellm_trace_id: str | None = None, + request_tags: list[str] | None = None, + client_ip: str | None = None, + mcp_proxy_mode: bool = False, +) -> AggregateToolListing: + """ + Helper method to fetch tools from MCP servers based on server filtering criteria. + + Args: + user_api_key_auth: User authentication info for access control + mcp_auth_header: Optional auth header for MCP server (deprecated) + mcp_servers: Optional list of server names/aliases to filter by + mcp_server_auth_headers: Optional dict of server-specific auth headers + oauth2_headers: Optional dict of oauth2 headers + + Returns: + AggregateToolListing: Combined tools from filtered servers plus each server's + classified listing outcome + """ + + list_tools_start_time: Final = datetime.now() + litellm_logging_obj: LiteLLMLoggingObj | None = None + list_tools_request_data: dict[str, object] = {} + + if log_list_tools_to_spendlogs: + # This is intentionally minimal: only async_success_handler / post_call_failure_hook + rules_obj: Final = Rules() + list_tools_call_id: Final = str(uuid.uuid4()) + # Derive trace_id from raw_headers when not explicitly passed (same as A2A / MCP call_tool) + effective_litellm_trace_id: Final = litellm_trace_id or get_chain_id_from_headers(raw_headers) + spend_logs_metadata: Final[dict[str, object]] = { + "mcp_operation": "list_tools", + } + if isinstance(list_tools_log_source, str): + spend_logs_metadata["source"] = list_tools_log_source + if isinstance(mcp_servers, list): + spend_logs_metadata["requested_mcp_servers"] = mcp_servers + + list_tools_request_data = { + "model": "MCP: list_tools", + "call_type": CallTypes.list_mcp_tools.value, + "litellm_call_id": list_tools_call_id, + "litellm_trace_id": effective_litellm_trace_id, + "metadata": { + "spend_logs_metadata": spend_logs_metadata, + "headers": logging_safe_mcp_headers(raw_headers), + **({"tags": request_tags} if request_tags else {}), + }, + # Provide a small input payload for standard logging + "input": [ + { + "role": "system", + "content": { + "mcp_operation": "list_tools", + "requested_mcp_servers": mcp_servers, + }, + } + ], + } + + # Attach user identifiers using the standard helper + if user_api_key_auth is not None: + LiteLLMProxyRequestSetup.add_user_api_key_auth_to_request_metadata( + data=list_tools_request_data, + user_api_key_dict=user_api_key_auth, + _metadata_variable_name="metadata", + ) + + user_identifier: Final = getattr(user_api_key_auth, "end_user_id", None) or getattr( + user_api_key_auth, "user_id", None + ) + if user_identifier: + list_tools_request_data["user"] = user_identifier + + try: + litellm_logging_obj, _ = function_setup( + original_function="list_mcp_tools", + is_async_call=False, + rules_obj=rules_obj, + start_time=list_tools_start_time, + **list_tools_request_data, + ) + if litellm_logging_obj: + litellm_logging_obj.call_type = CallTypes.list_mcp_tools.value + litellm_logging_obj.model = "MCP: list_tools" + except Exception as logging_error: + verbose_logger.debug("Failed to initialize logging for MCP list_tools: %s", logging_error) + litellm_logging_obj = None + + try: + allowed_mcp_servers: Final = await _get_allowed_mcp_servers( + user_api_key_auth=user_api_key_auth, + mcp_servers=mcp_servers, + client_ip=client_ip, + ) + if mcp_servers and not allowed_mcp_servers: + await raise_denied_scoped_mcp_access( + requested_names=mcp_servers, + user_api_key_auth=user_api_key_auth, + client_ip=client_ip, + ) + + # Pre-fetch OAuth credentials only when at least one server uses OAuth2, + # to avoid an unnecessary DB round-trip on requests with no OAuth2 MCP servers. + _has_oauth2_server = any(getattr(s, "auth_type", None) == MCPAuth.oauth2 for s in allowed_mcp_servers) + _prefetched_oauth_creds: Final = ( + await _prefetch_oauth_creds_for_user(user_api_key_auth) if _has_oauth2_server else {} + ) + + async def _fetch_and_filter_server_tools( + server: MCPServer, + ) -> "tuple[list[MCPTool], ServerOutcome]": + """Fetch and filter tools from a single server, classifying any failure into that + server's outcome so the aggregate can keep serving the healthy subset without a + broken server masquerading as an empty one.""" + if server is None: + return [], ServerListOk(tool_count=0) + + server_auth_header, extra_headers = _prepare_mcp_server_headers( + server=server, + mcp_server_auth_headers=mcp_server_auth_headers, + mcp_auth_header=mcp_auth_header, + oauth2_headers=oauth2_headers, + raw_headers=raw_headers, + user_api_key_auth=user_api_key_auth, + scope_servers=allowed_mcp_servers, + ) + + # Prefer server-stored per-user OAuth when configured, so a stale + # Authorization header from the MCP client cannot override Redis/DB + # (same issue as call_tool in mcp_server_manager: VS Code caches tokens). + from litellm.proxy._experimental.mcp_server.outbound_credentials.adapter import ( # noqa: PLC0415 + to_server_spec, + ) + + # A server migrated to the v2 resolver gets its token from the resolver at connect + # time; building it here would double-resolve and be shadowed by the v2 graft. The + # preemptive 401 already challenged a missing token, so one exists for the connect. + migrated_to_v2: Final = to_server_spec(server) is not None + if ( + not migrated_to_v2 + and server.auth_type == MCPAuth.oauth2 + and getattr(server, "needs_user_oauth_token", False) + and user_api_key_auth is not None + ): + db_headers: Final = await _get_user_oauth_extra_headers_from_db( + server, + user_api_key_auth, + prefetched_creds=_prefetched_oauth_creds, + ) + if db_headers: + extra_headers = db_headers + + # If still no OAuth2 token, fall back to pre-fetched creds (non-stale-client path) + elif not migrated_to_v2 and extra_headers is None and server.auth_type == MCPAuth.oauth2: + extra_headers = await _get_user_oauth_extra_headers_from_db( + server, + user_api_key_auth, + prefetched_creds=_prefetched_oauth_creds, + ) + + if server.is_byok and server.auth_type != MCPAuth.oauth2 and server_auth_header is None: + server_auth_header = await _get_byok_credential(server, user_api_key_auth) + + try: + tools: Final = await global_mcp_server_manager._get_tools_from_server( + server=server, + mcp_auth_header=server_auth_header, + extra_headers=extra_headers, + add_prefix=True, # Always add server prefix + raw_headers=raw_headers, + client_ip=client_ip, + user_api_key_auth=user_api_key_auth, + oauth2_headers=oauth2_headers, + ) + filtered_tools = filter_tools_by_allowed_tools(tools, server) + + filtered_tools = await filter_tools_by_key_team_permissions( + tools=filtered_tools, + server_id=server.server_id, + user_api_key_auth=user_api_key_auth, + ) + + if mcp_proxy_mode: + from litellm.proxy._experimental.mcp_server.tool_search import with_mcp_proxy_identity + + filtered_tools = [ # mutable-ok: MCP tool pipeline + with_mcp_proxy_identity(tool, server.server_id) for tool in filtered_tools + ] + else: + filtered_tools = apply_tool_overrides(filtered_tools, server) + + verbose_logger.debug( + "Successfully fetched %s tools from server %s, %s after filtering", + len(tools), + server.name, + len(filtered_tools), + ) + return filtered_tools, ServerListOk(tool_count=len(filtered_tools)) + except MCPUpstreamAuthError as e: + # Absorb so one unauthenticated server does not empty every other server's + # tools. Surfacing the upstream 401 to the client as a re-auth challenge is + # intentionally not done here: raising from this list handler cannot produce a + # 401 + WWW-Authenticate (the MCP session manager serializes it as a JSON-RPC + # error). Single-server routes surface it via the request-scope preemptive + # check in _raise_preemptive_401_for_unauthenticated_servers instead. + verbose_logger.debug("MCP list_tools: omitting %s; it needs upstream auth", server.name) + return [], classify_list_exception(e) + except Exception as e: + verbose_logger.exception("Error getting tools from server %s: %s", server.name, e) + return [], classify_list_exception(e) + + # Fetch tools from all servers in parallel + tasks: Final = [_fetch_and_filter_server_tools(server) for server in allowed_mcp_servers] + results: Final = await asyncio.gather(*tasks) + + # Flatten results into single list + all_tools: Final[list[MCPTool]] = [tool for tools, _ in results for tool in tools] + server_outcomes: Final[dict[str, ServerOutcome]] = { + _aggregate_server_key(server): outcome + for server, (_, outcome) in zip(allowed_mcp_servers, results) + if server is not None + } + + # If logging is enabled, enrich spend_logs_metadata with counts + if litellm_logging_obj: + per_server_tool_counts: Final[dict[str, int]] = { + _aggregate_server_key(server): len(server_tools) + for server, (server_tools, _) in zip(allowed_mcp_servers, results) + if server is not None + } + + metadata_dict: Final = litellm_logging_obj.model_call_details.get("metadata") + if isinstance(metadata_dict, dict): + spend_meta = metadata_dict.get("spend_logs_metadata") + if not isinstance(spend_meta, dict): + spend_meta = {} + metadata_dict["spend_logs_metadata"] = spend_meta + spend_meta["allowed_server_count"] = len(allowed_mcp_servers) + spend_meta["tool_count_total"] = len(all_tools) + spend_meta["per_server_tool_counts"] = per_server_tool_counts + spend_meta["per_server_list_outcomes"] = { + key: outcome_wire_value(outcome) for key, outcome in server_outcomes.items() + } + + end_time: Final = datetime.now() + try: + await litellm_logging_obj.async_success_handler( + result=[tool.model_dump(mode="json") if isinstance(tool, MCPTool) else tool for tool in all_tools], + start_time=list_tools_start_time, + end_time=end_time, + ) + except Exception as log_exc: + # list_tools responses must not be dropped due to non-blocking + # observability/serialization failures. + verbose_logger.warning( + "MCP list_tools success logging failed (continuing): %s", + log_exc, + ) + + verbose_logger.info("Successfully fetched %s tools total from all MCP servers", len(all_tools)) + + return AggregateToolListing(tools=all_tools, outcomes=server_outcomes) + except Exception as e: + # Only fire failure hook if logging was requested for this list-tools execution + if log_list_tools_to_spendlogs and user_api_key_auth is not None: + try: + from litellm.proxy.proxy_server import proxy_logging_obj + + if proxy_logging_obj: + traceback_str: Final = traceback.format_exc(limit=MAXIMUM_TRACEBACK_LINES_TO_LOG) + await proxy_logging_obj.post_call_failure_hook( + request_data=list_tools_request_data or {}, + original_exception=e, + user_api_key_dict=user_api_key_auth, + route="/mcp/list_tools", + traceback_str=traceback_str, + ) + except Exception: + verbose_logger.debug("Failed to log MCP list_tools failure via post_call_failure_hook") + raise + + +async def _get_prompts_from_mcp_servers( + user_api_key_auth: UserAPIKeyAuth | None, + mcp_auth_header: str | None, + mcp_servers: list[str] | None, + mcp_server_auth_headers: dict[str, dict[str, str]] | None = None, + oauth2_headers: dict[str, str] | None = None, + raw_headers: dict[str, str] | None = None, + client_ip: str | None = None, +) -> list[Prompt]: + """ + Helper method to fetch prompt from MCP servers based on server filtering criteria. + + Args: + user_api_key_auth: User authentication info for access control + mcp_auth_header: Optional auth header for MCP server (deprecated) + mcp_servers: Optional list of server names/aliases to filter by + mcp_server_auth_headers: Optional dict of server-specific auth headers + oauth2_headers: Optional dict of oauth2 headers + + Returns: + List[Prompt]: Combined list of prompts from filtered servers + """ + + allowed_mcp_servers: Final = await _get_allowed_mcp_servers( + user_api_key_auth=user_api_key_auth, + mcp_servers=mcp_servers, + client_ip=client_ip, + ) + + # Get prompts from each allowed server + all_prompts: Final = [] + for server in allowed_mcp_servers: + if server is None: + continue + + server_auth_header, extra_headers = _prepare_mcp_server_headers( + server=server, + mcp_server_auth_headers=mcp_server_auth_headers, + mcp_auth_header=mcp_auth_header, + oauth2_headers=oauth2_headers, + raw_headers=raw_headers, + user_api_key_auth=user_api_key_auth, + scope_servers=allowed_mcp_servers, + ) + + try: + prompts = await global_mcp_server_manager.get_prompts_from_server( + server=server, + user_api_key_auth=user_api_key_auth, + mcp_auth_header=server_auth_header, + extra_headers=extra_headers, + add_prefix=True, # Always add server prefix + raw_headers=raw_headers, + client_ip=client_ip, + ) + + all_prompts.extend(prompts) + + verbose_logger.debug("Successfully fetched %s prompts from server %s", len(prompts), server.name) + except Exception as e: + verbose_logger.exception("Error getting prompts from server %s: %s", server.name, e) + # Continue with other servers instead of failing completely + + verbose_logger.info("Successfully fetched %s prompts total from all MCP servers", len(all_prompts)) + + return all_prompts + + +async def _get_resources_from_mcp_servers( + user_api_key_auth: UserAPIKeyAuth | None, + mcp_auth_header: str | None, + mcp_servers: list[str] | None, + mcp_server_auth_headers: dict[str, dict[str, str]] | None = None, + oauth2_headers: dict[str, str] | None = None, + raw_headers: dict[str, str] | None = None, + client_ip: str | None = None, +) -> list[Resource]: + """Fetch resources from allowed MCP servers.""" + + allowed_mcp_servers: Final = await _get_allowed_mcp_servers( + user_api_key_auth=user_api_key_auth, + mcp_servers=mcp_servers, + client_ip=client_ip, + ) + + all_resources: Final[list[Resource]] = [] + for server in allowed_mcp_servers: + if server is None: + continue + + server_auth_header, extra_headers = _prepare_mcp_server_headers( + server=server, + mcp_server_auth_headers=mcp_server_auth_headers, + mcp_auth_header=mcp_auth_header, + oauth2_headers=oauth2_headers, + raw_headers=raw_headers, + user_api_key_auth=user_api_key_auth, + scope_servers=allowed_mcp_servers, + ) + + try: + resources = await global_mcp_server_manager.get_resources_from_server( + server=server, + user_api_key_auth=user_api_key_auth, + mcp_auth_header=server_auth_header, + extra_headers=extra_headers, + add_prefix=True, # Always add server prefix + raw_headers=raw_headers, + client_ip=client_ip, + ) + all_resources.extend(resources) + + verbose_logger.debug("Successfully fetched %s resources from server %s", len(resources), server.name) + except Exception as e: + verbose_logger.exception("Error getting resources from server %s: %s", server.name, e) + + verbose_logger.info("Successfully fetched %s resources total from all MCP servers", len(all_resources)) + + return all_resources + + +async def _get_resource_templates_from_mcp_servers( + user_api_key_auth: UserAPIKeyAuth | None, + mcp_auth_header: str | None, + mcp_servers: list[str] | None, + mcp_server_auth_headers: dict[str, dict[str, str]] | None = None, + oauth2_headers: dict[str, str] | None = None, + raw_headers: dict[str, str] | None = None, + client_ip: str | None = None, +) -> list[ResourceTemplate]: + """Fetch resource templates from allowed MCP servers.""" + + allowed_mcp_servers: Final = await _get_allowed_mcp_servers( + user_api_key_auth=user_api_key_auth, + mcp_servers=mcp_servers, + client_ip=client_ip, + ) + + all_resource_templates: Final[list[ResourceTemplate]] = [] + for server in allowed_mcp_servers: + if server is None: + continue + + server_auth_header, extra_headers = _prepare_mcp_server_headers( + server=server, + mcp_server_auth_headers=mcp_server_auth_headers, + mcp_auth_header=mcp_auth_header, + oauth2_headers=oauth2_headers, + raw_headers=raw_headers, + user_api_key_auth=user_api_key_auth, + scope_servers=allowed_mcp_servers, + ) + + try: + resource_templates = await global_mcp_server_manager.get_resource_templates_from_server( + server=server, + user_api_key_auth=user_api_key_auth, + mcp_auth_header=server_auth_header, + extra_headers=extra_headers, + add_prefix=True, # Always add server prefix + raw_headers=raw_headers, + client_ip=client_ip, + ) + all_resource_templates.extend(resource_templates) + verbose_logger.debug( + "Successfully fetched %s resource templates from server %s", + len(resource_templates), + server.name, + ) + except Exception as e: + verbose_logger.exception( + "Error getting resource templates from server %s: %s", + server.name, + str(e), + ) + + verbose_logger.info( + "Successfully fetched %s resource templates total from all MCP servers", + len(all_resource_templates), + ) + + return all_resource_templates + + +async def filter_tools_by_key_team_permissions( + tools: list[MCPTool], + server_id: str, + user_api_key_auth: UserAPIKeyAuth | None, +) -> list[MCPTool]: + """ + Filter tools based on key/team mcp_tool_permissions. + + Note: Tool names in the DB are stored without server prefixes, + but tool names from MCP servers are prefixed. We need to strip + the prefix before comparing. + """ + # Filter by key/team tool-level permissions + allowed_tool_names: Final = await MCPRequestHandler.get_allowed_tools_for_server( + server_id=server_id, + user_api_key_auth=user_api_key_auth, + ) + + # Tools arrive prefixed with the server's own prefix; strip exactly that + # prefix (resolved from the server) rather than the first separator, so a + # prefix containing the separator still reduces to the stored bare name. + server: Final = global_mcp_server_manager.get_mcp_server_by_id(server_id) + return [ + t + for t in tools + if MCPRequestHandler.tool_is_granted(strip_known_server_prefix(t.name, server), allowed_tool_names) + ] + + +async def _list_mcp_tools( + user_api_key_auth: UserAPIKeyAuth | None = None, + mcp_auth_header: str | None = None, + mcp_servers: list[str] | None = None, + mcp_server_auth_headers: dict[str, dict[str, str]] | None = None, + oauth2_headers: dict[str, str] | None = None, + raw_headers: dict[str, str] | None = None, + log_list_tools_to_spendlogs: bool = False, + list_tools_log_source: str | None = None, + client_ip: str | None = None, + mcp_proxy_mode: bool = False, +) -> AggregateToolListing: + """ + List all available MCP tools. + + Args: + user_api_key_auth: User authentication info for access control + mcp_auth_header: Optional auth header for MCP server (deprecated) + mcp_servers: Optional list of server names/aliases to filter by + mcp_server_auth_headers: Optional dict of server-specific auth headers {server_alias: auth_value} + client_ip: Client IP for IP-based server access control + + Returns: + AggregateToolListing: Combined tools from all accessible servers plus each server's + classified listing outcome + """ + + try: + listing: Final = await _get_tools_from_mcp_servers( + user_api_key_auth=user_api_key_auth, + mcp_auth_header=mcp_auth_header, + mcp_servers=mcp_servers, + mcp_server_auth_headers=mcp_server_auth_headers, + oauth2_headers=oauth2_headers, + raw_headers=raw_headers, + log_list_tools_to_spendlogs=log_list_tools_to_spendlogs, + list_tools_log_source=list_tools_log_source, + client_ip=client_ip, + mcp_proxy_mode=mcp_proxy_mode, + ) + verbose_logger.debug("Successfully fetched %s tools from managed MCP servers", len(listing.tools)) + return listing + except HTTPException: + raise + except Exception as e: + verbose_logger.exception("Error getting tools from managed MCP servers: %s", e) + # Continue with an empty listing instead of failing completely + return AggregateToolListing(tools=[], outcomes={}) + + +async def _list_mcp_prompts( + user_api_key_auth: UserAPIKeyAuth | None = None, + mcp_auth_header: str | None = None, + mcp_servers: list[str] | None = None, + mcp_server_auth_headers: dict[str, dict[str, str]] | None = None, + oauth2_headers: dict[str, str] | None = None, + raw_headers: dict[str, str] | None = None, + client_ip: str | None = None, +) -> list[Prompt]: + """ + List all available MCP prompts. + + Args: + user_api_key_auth: User authentication info for access control + mcp_auth_header: Optional auth header for MCP server (deprecated) + mcp_servers: Optional list of server names/aliases to filter by + mcp_server_auth_headers: Optional dict of server-specific auth headers {server_alias: auth_value} + + Returns: + List[Prompt]: Combined list of tools from all accessible servers + """ + # Get tools from managed MCP servers with error handling + managed_prompts = [] + try: + managed_prompts = await _get_prompts_from_mcp_servers( + user_api_key_auth=user_api_key_auth, + mcp_auth_header=mcp_auth_header, + mcp_servers=mcp_servers, + mcp_server_auth_headers=mcp_server_auth_headers, + oauth2_headers=oauth2_headers, + raw_headers=raw_headers, + client_ip=client_ip, + ) + verbose_logger.debug("Successfully fetched %s prompts from managed MCP servers", len(managed_prompts)) + except Exception as e: + verbose_logger.exception("Error getting tools from managed MCP servers: %s", e) + # Continue with empty managed tools list instead of failing completely + + return managed_prompts + + +async def _list_mcp_resources( + user_api_key_auth: UserAPIKeyAuth | None = None, + mcp_auth_header: str | None = None, + mcp_servers: list[str] | None = None, + mcp_server_auth_headers: dict[str, dict[str, str]] | None = None, + oauth2_headers: dict[str, str] | None = None, + raw_headers: dict[str, str] | None = None, + client_ip: str | None = None, +) -> list[Resource]: + """List all available MCP resources.""" + + managed_resources: list[Resource] = [] + try: + managed_resources = await _get_resources_from_mcp_servers( + user_api_key_auth=user_api_key_auth, + mcp_auth_header=mcp_auth_header, + mcp_servers=mcp_servers, + mcp_server_auth_headers=mcp_server_auth_headers, + oauth2_headers=oauth2_headers, + raw_headers=raw_headers, + client_ip=client_ip, + ) + verbose_logger.debug("Successfully fetched %s resources from managed MCP servers", len(managed_resources)) + except Exception as e: + verbose_logger.exception("Error getting resources from managed MCP servers: %s", e) + + return managed_resources + + +async def _list_mcp_resource_templates( + user_api_key_auth: UserAPIKeyAuth | None = None, + mcp_auth_header: str | None = None, + mcp_servers: list[str] | None = None, + mcp_server_auth_headers: dict[str, dict[str, str]] | None = None, + oauth2_headers: dict[str, str] | None = None, + raw_headers: dict[str, str] | None = None, + client_ip: str | None = None, +) -> list[ResourceTemplate]: + """List all available MCP resource templates.""" + + managed_resource_templates: list[ResourceTemplate] = [] + try: + managed_resource_templates = await _get_resource_templates_from_mcp_servers( + user_api_key_auth=user_api_key_auth, + mcp_auth_header=mcp_auth_header, + mcp_servers=mcp_servers, + mcp_server_auth_headers=mcp_server_auth_headers, + oauth2_headers=oauth2_headers, + raw_headers=raw_headers, + client_ip=client_ip, + ) + verbose_logger.debug( + "Successfully fetched %s resource templates from managed MCP servers", + len(managed_resource_templates), + ) + except Exception as e: + verbose_logger.exception( + "Error getting resource templates from managed MCP servers: %s", + str(e), + ) + + return managed_resource_templates + + +def _resolve_display_name_to_original( + name: str, + allowed_mcp_servers: list[MCPServer], +) -> str: + """Translate a display-name override back to the original prefixed tool name. + + When a client received a customised display name from tools/list (e.g. + "Get Pet") it will call tools/call with that same string. We need to + reverse-map it to the original prefixed name (e.g. + "petstore_mcp-getPetById") before any routing or permission logic runs. + """ + for server in allowed_mcp_servers: + display_map = server.tool_name_to_display_name or {} + for unprefixed_name, display_name in display_map.items(): + if display_name == name: + return add_server_prefix_to_name(unprefixed_name, get_server_prefix(server)) + return name + + +async def _get_byok_credential( + mcp_server: MCPServer, + user_api_key_auth: UserAPIKeyAuth | None, +) -> str | None: + """Retrieve the stored BYOK credential for a user+server pair, served from the worker cache within its TTL.""" + if not mcp_server.is_byok: + return None + user_id: Final = (user_api_key_auth.user_id if user_api_key_auth else None) or "" + if not user_id: + return None + + cached: Final = get_cached_byok_credential(user_id, mcp_server.server_id) + if cached is not None: + return cached.credential + + from litellm.proxy._experimental.mcp_server.db import get_user_credential + from litellm.proxy.proxy_server import prisma_client + + if prisma_client is None: + return None + credential: Final = await get_user_credential( + prisma_client=prisma_client, + user_id=user_id, + server_id=mcp_server.server_id, + ) + cache_byok_credential(user_id, mcp_server.server_id, credential) + return credential + + +async def _check_byok_credential( + mcp_server: MCPServer, + user_api_key_auth: UserAPIKeyAuth | None, +) -> None: + """ + If the MCP server is BYOK-enabled, verify that the requesting user has a + stored credential. When no credential is found, raise an HTTP 401 with a + WWW-Authenticate header that points the MCP client to our OAuth metadata + endpoint so it can drive the authorization flow. + """ + if not mcp_server.is_byok: + return + + user_id: Final = (user_api_key_auth.user_id if user_api_key_auth else None) or "" + if not user_id: + raise HTTPException( + status_code=401, + detail={ + "error": "byok_auth_required", + "server_id": mcp_server.server_id, + "server_name": mcp_server.server_name or mcp_server.name, + "message": "User identity is required for BYOK servers", + }, + headers={"WWW-Authenticate": get_byok_www_authenticate()}, + ) + + cached: Final = get_cached_byok_credential(user_id, mcp_server.server_id) + if cached is not None: + if cached.credential is None: + raise HTTPException( + status_code=401, + detail={ + "error": "byok_auth_required", + "server_id": mcp_server.server_id, + "server_name": mcp_server.server_name or mcp_server.name, + "message": ( + "No stored credential found for this BYOK server. " + "Complete the OAuth authorization flow to provide your API key." + ), + }, + headers={"WWW-Authenticate": get_byok_www_authenticate()}, + ) + return + + from litellm.proxy._experimental.mcp_server.db import get_user_credential + from litellm.proxy.proxy_server import prisma_client + + if prisma_client is None: + # Fail closed on DB unavailability: returning here previously + # bypassed the ownership check and let any proxy-authenticated + # caller invoke BYOK tools during outage windows. + raise HTTPException( + status_code=503, + detail={ + "error": "byok_auth_unavailable", + "server_id": mcp_server.server_id, + "server_name": mcp_server.server_name or mcp_server.name, + "message": "BYOK credential check requires a database connection.", + }, + ) + + credential: Final = await get_user_credential( + prisma_client=prisma_client, + user_id=user_id, + server_id=mcp_server.server_id, + ) + cache_byok_credential(user_id, mcp_server.server_id, credential) + if credential is None: + raise HTTPException( + status_code=401, + detail={ + "error": "byok_auth_required", + "server_id": mcp_server.server_id, + "server_name": mcp_server.server_name or mcp_server.name, + "message": ( + "No stored credential found for this BYOK server. " + "Complete the OAuth authorization flow to provide your API key." + ), + }, + headers={"WWW-Authenticate": get_byok_www_authenticate()}, + ) + + +async def _list_tools_before_first_call( + server: MCPServer | None, + tool_name: str, + allowed_mcp_servers: list[MCPServer], + user_api_key_auth: UserAPIKeyAuth | None, + mcp_auth_header: str | None, + mcp_server_auth_headers: dict[str, dict[str, str]] | None, + oauth2_headers: dict[str, str] | None, + raw_headers: dict[str, str] | None, + client_ip: str | None = None, +) -> None: + """List ``server`` with the caller's own credentials when it does not yet expose ``tool_name`` here. + + The startup fill skips a server whose upstream wants the caller's token, and mcp 2 no + longer lists before an uncached tools/call, so a worker that has not served tools/list + for this caller would otherwise answer 404 for a tool the caller can see. Gating on the + requested tool, not on any prior listing, keeps callers with different upstream catalogs + from masking each other. + """ + if server is None or global_mcp_server_manager.server_exposes_tool(server, tool_name): + return + if all(allowed.server_id != server.server_id for allowed in allowed_mcp_servers): + return + try: + await _get_tools_from_mcp_servers( + user_api_key_auth=user_api_key_auth, + mcp_auth_header=mcp_auth_header, + mcp_servers=[server.server_id], + mcp_server_auth_headers=mcp_server_auth_headers, + oauth2_headers=oauth2_headers, + raw_headers=raw_headers, + client_ip=client_ip, + ) + except Exception as e: # noqa: BLE001 # best effort: resolution below answers as it did before + verbose_logger.debug("MCP tools/call: listing %s before its first call failed: %s", server.name, e) + + +async def execute_mcp_tool( + name: str, + arguments: dict[str, object], + allowed_mcp_servers: list[MCPServer], + start_time: datetime, + user_api_key_auth: UserAPIKeyAuth | None = None, + mcp_auth_header: str | None = None, + mcp_server_auth_headers: dict[str, dict[str, str]] | None = None, + oauth2_headers: dict[str, str] | None = None, + raw_headers: dict[str, str] | None = None, + host_progress_callback: ProgressCallback | None = None, + guardrail_context: Mapping[str, object] | None = None, + client_ip: str | None = None, + **kwargs: object, # kwargs-ok: preserves the existing REST and decorated logging call contract +) -> CallToolResult: + context: Final = prepare_context( + user_api_key_auth=user_api_key_auth, + mcp_auth_header=mcp_auth_header, + mcp_server_auth_headers=mcp_server_auth_headers, + oauth2_headers=oauth2_headers, + raw_headers=raw_headers, + client_ip=client_ip, + ) + operation: Final = AuthorizedToolCall( + name=name, + arguments=arguments, + allowed_mcp_servers=tuple(allowed_mcp_servers), + start_time=start_time, + host_progress_callback=host_progress_callback, + guardrail_context=guardrail_context, + logging_data=types.MappingProxyType(kwargs), + ) + return await GatewayOperations().execute(operation, context) + + +async def _execute_mcp_tool( + name: str, + arguments: dict[str, object], + allowed_mcp_servers: list[MCPServer], + start_time: datetime, + user_api_key_auth: UserAPIKeyAuth | None = None, + mcp_auth_header: str | None = None, + mcp_server_auth_headers: dict[str, dict[str, str]] | None = None, + oauth2_headers: dict[str, str] | None = None, + raw_headers: dict[str, str] | None = None, + host_progress_callback: ProgressCallback | None = None, + guardrail_context: Mapping[str, object] | None = None, + client_ip: str | None = None, + **kwargs: Any, +) -> CallToolResult: + """ + Execute MCP tool. + + This function assumes permission checks have already been performed. + + Args: + name: Tool name (may include server prefix) + arguments: Tool arguments + allowed_mcp_servers: Pre-validated list of servers the user can access + start_time: Start time for logging + user_api_key_auth: Optional user API key auth for logging + mcp_auth_header: Optional MCP auth header + mcp_server_auth_headers: Optional server-specific auth headers + oauth2_headers: Optional OAuth2 headers + raw_headers: Optional raw HTTP headers + **kwargs: Additional arguments (e.g., litellm_logging_obj) + + Returns: + CallToolResult: Tool execution result + """ + # Track resolved MCP server for both permission checks and dispatch + mcp_server: MCPServer | None = None + requested_server_id: Final[str | None] = kwargs.get("requested_server_id") + + # If the client called with a display-name override (e.g. "Get Pet"), + # translate it back to the original prefixed name before any routing. + name = _resolve_display_name_to_original(name, allowed_mcp_servers) + + # Remove prefix from tool name for logging and processing + original_tool_name, server_name = split_server_prefix_from_name(name) + + requested_server: MCPServer | None = None + if requested_server_id: + requested_server = next( + (s for s in allowed_mcp_servers if s.server_id == requested_server_id), + None, + ) + + name_is_prefixed = False + if requested_server is not None and MCP_TOOL_PREFIX_SEPARATOR in name: + all_registry_prefixes: Final[set[str]] = set() + for registry_server in global_mcp_server_manager.get_registry().values(): + for known_prefix in iter_known_server_prefixes(registry_server): + all_registry_prefixes.add(normalize_server_name(known_prefix)) + name_is_prefixed = is_tool_name_prefixed(name, known_server_prefixes=all_registry_prefixes) + + first_call_target: Final = ( + requested_server + if requested_server is not None and not name_is_prefixed + else global_mcp_server_manager.server_owning_tool_name_prefix(name) + ) + first_call_tool_name: Final = ( + name + if first_call_target is None or (requested_server is not None and not name_is_prefixed) + else strip_known_server_prefix(name, first_call_target) + ) + await _list_tools_before_first_call( + server=first_call_target, + tool_name=first_call_tool_name, + allowed_mcp_servers=allowed_mcp_servers, + user_api_key_auth=user_api_key_auth, + mcp_auth_header=mcp_auth_header, + mcp_server_auth_headers=mcp_server_auth_headers, + oauth2_headers=oauth2_headers, + raw_headers=raw_headers, + client_ip=client_ip, + ) + + if requested_server is not None and not name_is_prefixed: + # REST callers may pass server_id with the upstream tool name (no + # LiteLLM prefix). The first segment is not a registered server + # prefix, so the whole string is the upstream tool name and may + # legitimately contain the separator (e.g. "text-to-speech"). + # server_id is authoritative for routing and auth. + mcp_server = requested_server + server_name = requested_server.name + original_tool_name = name + else: + # Resolve from tool name (MCP JSON-RPC or prefixed REST tool names). + mcp_server = global_mcp_server_manager._get_mcp_server_from_tool_name(name) + if mcp_server is None and requested_server is not None: + for known_prefix in iter_known_server_prefixes(requested_server): + candidate = global_mcp_server_manager._get_mcp_server_from_tool_name( + add_server_prefix_to_name(name, known_prefix) + ) + if candidate is not None: + mcp_server = candidate + break + if mcp_server is not None: + server_name = mcp_server.name + original_tool_name = strip_known_server_prefix(name, mcp_server) + + if requested_server is not None: + if mcp_server is not None and mcp_server.server_id != requested_server.server_id: + raise HTTPException( + status_code=403, + detail={ + "error": "tool_server_mismatch", + "message": ( + f"Tool '{name}' belongs to MCP server " + f"'{mcp_server.name}' but request specified " + f"server_id for '{requested_server.name}'." + ), + }, + ) + if mcp_server is None: + mcp_server = requested_server + server_name = requested_server.name + original_tool_name = strip_known_server_prefix(name, requested_server) + + # Only enforce server-level permissions when we can resolve a server + if server_name: + if not MCPRequestHandler.is_tool_allowed( + allowed_mcp_servers=[server.name for server in allowed_mcp_servers], + server_name=server_name, + ): + raise HTTPException( + status_code=403, + detail="User not allowed to call this tool.", + ) + + standard_logging_mcp_tool_call: Final[StandardLoggingMCPToolCall] = _get_standard_logging_mcp_tool_call( + 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: Final[LiteLLMLoggingObj | None] = kwargs.get("litellm_logging_obj", None) + if litellm_logging_obj: + litellm_logging_obj.model_call_details["mcp_tool_call_metadata"] = standard_logging_mcp_tool_call + litellm_logging_obj.model = f"MCP: {name}" + litellm_logging_obj.model_call_details["model"] = f"MCP: {name}" + # Resolve the MCP server early so BYOK checks and credential injection + # apply to ALL dispatch paths (local tool registry AND managed MCP server). + if mcp_server is None: + mcp_server = global_mcp_server_manager._get_mcp_server_from_tool_name(name) + + if mcp_server: + standard_logging_mcp_tool_call["mcp_server_cost_info"] = (mcp_server.mcp_info or {}).get("mcp_server_cost_info") + if litellm_logging_obj: + litellm_logging_obj.model_call_details["mcp_tool_call_metadata"] = standard_logging_mcp_tool_call + + # BYOK: retrieve the stored per-user credential. A single DB call + # both checks existence and fetches the value, avoiding a double query. + if mcp_server.is_byok and not mcp_auth_header: + byok_cred: Final = await _get_byok_credential(mcp_server, user_api_key_auth) + if byok_cred is None: + raise HTTPException( + status_code=401, + detail={ + "error": "byok_auth_required", + "server_id": mcp_server.server_id, + "server_name": mcp_server.server_name or mcp_server.name, + "message": ( + "No stored credential found for this BYOK server. " + "Complete the OAuth authorization flow to provide your API key." + ), + }, + headers={"WWW-Authenticate": get_byok_www_authenticate()}, + ) + mcp_auth_header = byok_cred + elif mcp_server.is_byok: + # External auth header supplied; still enforce user-identity check. + await _check_byok_credential(mcp_server, user_api_key_auth) + + # Check if tool exists in local registry first (for OpenAPI-based tools) + # These tools are registered with their prefixed names + ######################################################### + local_tool: Final = global_mcp_tool_registry.get_tool(name) + if local_tool: + # OpenAPI-backed tools used to bypass `pre_call_tool_check` — + # only the managed path ran allowed/banned-tool checks, key/team + # tool permissions, and parameter validation. Run the same checks + # before dispatching to the local registry. Refuse the call if + # we cannot resolve a server: tools registered via + # openapi_to_mcp_generator are always tied to a server, so a + # missing mcp_server here means the tool->server mapping has + # not finished initializing or the registry entry is orphaned. + # Skipping the check would re-open the same authorization gap. + if mcp_server is None: + raise HTTPException( + status_code=503, + detail=( + f"MCP server for tool '{name}' is not available; " + "refusing to dispatch without authorization checks. " + "Retry once the server is registered." + ), + ) + + # `pre_call_tool_check` calls into `proxy_logging_obj` for the + # pre-call guardrail hooks, so source it from the canonical + # `proxy_server` module the same way `_handle_managed_mcp_tool` + # does. `kwargs.get("proxy_logging_obj")` is None on the MCP + # entry path and would crash with AttributeError after the + # security checks pass. + from litellm.proxy.proxy_server import proxy_logging_obj + + hook_result = await global_mcp_server_manager.pre_call_tool_check( + name=original_tool_name, + arguments=arguments or {}, + server_name=server_name or mcp_server.name, + user_api_key_auth=user_api_key_auth, + proxy_logging_obj=proxy_logging_obj, + server=mcp_server, + raw_headers=raw_headers, + litellm_logging_obj=litellm_logging_obj, + guardrail_context=guardrail_context, + ) + # `pre_call_tool_check` may return guardrail-modified + # arguments; honor them on the local path too. + if isinstance(hook_result, dict) and "arguments" in hook_result: + arguments = hook_result["arguments"] + + verbose_logger.debug("Executing local registry tool: %s", name) + # The credential rides ContextVars because the tool function has its + # headers baked into the closure at registration time. + auth_header_value, openapi_forwarded_headers, upstream_credential = _resolve_openapi_tool_auth( + mcp_server=mcp_server, + mcp_auth_header=mcp_auth_header, + mcp_server_auth_headers=mcp_server_auth_headers, + raw_headers=raw_headers, + user_api_key_auth=user_api_key_auth, + ) + ( + resolved_auth_headers, + forwarded_headers, + ) = await global_mcp_server_manager.resolve_openapi_upstream_auth( + mcp_server=mcp_server, + oauth2_headers=oauth2_headers, + raw_headers=raw_headers, + mcp_auth_header=upstream_credential, + user_api_key_auth=user_api_key_auth, + forwarded_headers=openapi_forwarded_headers, + ) + + _auth_token: Final = _request_auth_header.set(auth_header_value) + _extra_token: Final = _request_extra_headers.set(forwarded_headers) + _resolved_token: Final = _request_resolved_auth_headers.set(resolved_auth_headers) + try: + response = await _handle_local_mcp_tool(name, arguments) + finally: + _request_auth_header.reset(_auth_token) + _request_extra_headers.reset(_extra_token) + _request_resolved_auth_headers.reset(_resolved_token) + + # Try managed MCP server tool (the name is bare; the prefix boundary was + # already resolved above against this server's registered prefixes) + # Primary and recommended way to use external MCP servers + ######################################################### + elif mcp_server: + response = await _handle_managed_mcp_tool( + server_name=server_name, + name=original_tool_name, + arguments=arguments, + user_api_key_auth=user_api_key_auth, + mcp_auth_header=mcp_auth_header, + mcp_server_auth_headers=mcp_server_auth_headers, + oauth2_headers=oauth2_headers, + raw_headers=raw_headers, + client_ip=client_ip, + litellm_logging_obj=litellm_logging_obj, + guardrail_context=guardrail_context, + host_progress_callback=host_progress_callback, + ) + + # Fall back to local tool registry with original name (legacy support) + ######################################################### + # Deprecated: Local MCP Server Tool + ######################################################### + else: + # Gate only what can actually dispatch. When the unprefixed name is + # not in the registry either, `_handle_local_mcp_tool` below reports + # 404 and nothing runs, so demanding a server here would turn every + # unknown tool name into a misleading 503. + if global_mcp_tool_registry.get_tool(original_tool_name) is not None: + # `mcp_server` is None here because the tool name is not in the + # tool -> server mapping, but the name still carries a prefix + # that the server-level check above compared against the + # caller's `allowed_mcp_servers` by exact `name`. So the named + # server is in that list and can carry the tool-level checks, + # even with the mapping cold. Resolve it from + # `allowed_mcp_servers` rather than the registry: the registry + # would happily return a server the caller holds no grant for, + # and matching anything other than `name` would accept a server + # the check never validated. + prefix_server: Final = next( + (candidate for candidate in allowed_mcp_servers if candidate.name == server_name), + None, + ) + if prefix_server is None: + # A non-empty prefix that passed the server-level check + # always matches here, so this arm only fires when the + # prefix was empty, which is exactly the case that check + # skips. Fail closed rather than dispatch with no server to + # evaluate a tool ceiling against. + raise HTTPException( + status_code=503, + detail=( + f"MCP server for tool '{original_tool_name}' is not available; " + "refusing to dispatch without authorization checks. " + "Retry once the server is registered." + ), + ) + + from litellm.proxy.proxy_server import proxy_logging_obj + + hook_result = await global_mcp_server_manager.pre_call_tool_check( + name=original_tool_name, + arguments=arguments, + server_name=server_name, + user_api_key_auth=user_api_key_auth, + proxy_logging_obj=proxy_logging_obj, + server=prefix_server, + raw_headers=raw_headers, + litellm_logging_obj=litellm_logging_obj, + guardrail_context=guardrail_context, + ) + if "arguments" in hook_result: + arguments = hook_result["arguments"] # pyright: ignore[reportAny] # hook returns untyped args + + response = await _handle_local_mcp_tool(original_tool_name, arguments) + + return await _run_post_mcp_call_guardrails( + result=response, + litellm_logging_obj=litellm_logging_obj, + user_api_key_auth=user_api_key_auth, + request_data=kwargs, + ) + + +async def _run_post_mcp_call_guardrails( + result: CallToolResult, + litellm_logging_obj: LiteLLMLoggingObj | None, + user_api_key_auth: UserAPIKeyAuth | None, + request_data: Mapping[str, object], +) -> CallToolResult: + """Run ``post_mcp_call`` guardrails over an executed tool result. + + Lives on ``execute_mcp_tool``'s return path rather than inside + ``_fire_mcp_tool_call_logging`` so enforcement never depends on logging + being configured, and so every dispatch route gets it: the MCP protocol + handler, the REST endpoint, and tool search all funnel through here. + A guardrail that rejects the result raises, matching ``pre_mcp_call``. + """ + from litellm.proxy.proxy_server import proxy_logging_obj + + if proxy_logging_obj is None: + return result + return await proxy_logging_obj.post_mcp_call_hook( + response=result, + request_data=( + litellm_logging_obj.model_call_details if litellm_logging_obj is not None else dict(request_data) + ), + user_api_key_dict=user_api_key_auth, + ) + + +async def _fire_mcp_tool_call_logging( + logging_obj: LiteLLMLoggingObj, + result: CallToolResult, + start_time: datetime, + end_time: datetime, + user_api_key_auth: UserAPIKeyAuth | None = None, + request_data: Mapping[str, object] | None = None, +) -> CallToolResult: + """Fire post-call logging for an executed MCP tool call, returning the result to send. + + The returned result is what the caller must forward to the client: a + ``post_mcp_call`` guardrail may rewrite the tool output (e.g. mask + sensitive values) or reject it, in which case its exception propagates. + Guardrails run before the success/failure logging so the masked text, not + the raw one, is what gets logged. + + A result with ``is_error=True`` is logged as a failure (``status="failure"`` + payload, so OTel marks the span ERROR) while the HTTP wire behavior stays + 200 + ``isError: true`` per the MCP spec. The error check runs after + ``async_post_mcp_tool_call_hook`` because guardrails may flip the result + to ``is_error=True`` in that hook. Raised exceptions never reach here (the + ``@client`` wrapper and ``call_mcp_tool``'s except path log those), so + this cannot double-log a failure. + + ``request_data`` may carry credential-bearing fields (the REST path puts + ``raw_headers``, ``mcp_auth_header``, ``mcp_server_auth_headers``, and + ``oauth2_headers`` at the top level of its data dict), so those are + stripped before the dict is handed to ``post_call_failure_hook`` + callbacks. + """ + from litellm.proxy.proxy_server import proxy_logging_obj + + logging_obj.post_call(original_response=result) + await logging_obj.async_post_mcp_tool_call_hook( + kwargs=logging_obj.model_call_details, + response_obj=result, + start_time=start_time, + end_time=end_time, + ) + logging_obj.call_type = CallTypes.call_mcp_tool.value + error_message: Final = extract_mcp_tool_result_error_message(result) + if error_message is None: + await logging_obj.async_success_handler(result=result, start_time=start_time, end_time=end_time) + return result + + logging_obj.has_run_logging(event_type="sync_success") + logging_obj.has_run_logging(event_type="async_success") + tool_error: Final = MCPToolResultError(error_message) + logging_obj.failure_handler(tool_error, "", start_time, end_time) + await logging_obj.async_failure_handler(tool_error, "", start_time, end_time) + + if user_api_key_auth is None: + return result + + if proxy_logging_obj: + sanitized_request_data: Final = { + key: value for key, value in (request_data or {}).items() if key not in _MCP_CREDENTIAL_REQUEST_FIELDS + } + await proxy_logging_obj.post_call_failure_hook( + request_data=sanitized_request_data, + original_exception=tool_error, + user_api_key_dict=user_api_key_auth, + route="/mcp/call_tool", + ) + return result + + +async def fire_mcp_tool_call_failure_logging( + logging_obj: LiteLLMLoggingObj | None, + exception: Exception, + start_time: datetime, + user_api_key_auth: UserAPIKeyAuth | None, + request_data: Mapping[str, object], +) -> None: + """Failure logging shared by the ``/mcp`` path and the REST endpoint. Call from + inside the ``except`` block so the traceback is still available. + + The failure handlers run first because ``_ProxyDBLogger.async_post_call_failure_hook`` + builds the failure spend-log row from the ``standard_logging_object`` they produce; + both gate on ``should_run_logging``, so the ``@client`` wrapper does not log twice. + A relayed upstream 401 (``MCPUpstreamAuthError``) is an expected caller-must-reauth + signal and skips ``post_call_failure_hook``, which fires the ``llm_exceptions`` alert. + """ + from litellm.proxy.proxy_server import proxy_logging_obj + + traceback_str: Final = traceback.format_exc(limit=MAXIMUM_TRACEBACK_LINES_TO_LOG) + if logging_obj is not None: + end_time: Final = datetime.now() # noqa: DTZ005 # naive to match `start_time`, which it is subtracted from + logging_obj.failure_handler(exception, traceback_str, start_time, end_time) + await logging_obj.async_failure_handler(exception, traceback_str, start_time, end_time) + + if isinstance(exception, MCPUpstreamAuthError) or not proxy_logging_obj or user_api_key_auth is None: + return + sanitized_request_data: Final = { + key: value for key, value in request_data.items() if key not in _MCP_CREDENTIAL_REQUEST_FIELDS + } + await proxy_logging_obj.post_call_failure_hook( + request_data=sanitized_request_data, + original_exception=exception, + user_api_key_dict=user_api_key_auth, + route="/mcp/call_tool", + traceback_str=traceback_str, + ) + + +@client +async def call_mcp_tool( + name: str, + arguments: dict[str, object] | None = None, + user_api_key_auth: UserAPIKeyAuth | None = None, + mcp_auth_header: str | None = None, + mcp_servers: list[str] | None = None, + mcp_server_auth_headers: dict[str, dict[str, str]] | None = None, + oauth2_headers: dict[str, str] | None = None, + raw_headers: dict[str, str] | None = None, + client_ip: str | None = None, + **kwargs: Any, +) -> CallToolResult: + """ + Call a specific tool with the provided arguments (handles prefixed tool names). + """ + start_time: Final = datetime.now() + litellm_logging_obj: Final[LiteLLMLoggingObj | None] = kwargs.get("litellm_logging_obj", None) + + try: + if arguments is None: + raise HTTPException(status_code=400, detail="Request arguments are required") + + ## CHECK IF USER IS ALLOWED TO CALL THIS TOOL + allowed_mcp_server_ids: Final = await global_mcp_server_manager.get_allowed_mcp_servers( + user_api_key_auth=user_api_key_auth, + ) + + allowed_mcp_servers: list[MCPServer] = [] + for allowed_mcp_server_id in allowed_mcp_server_ids: + allowed_server = global_mcp_server_manager.get_mcp_server_by_id(allowed_mcp_server_id) + if allowed_server is not None: + # Same request-time oauth2_flow backstop the listing path applies, + # so a null-flow M2M-shape row is treated as M2M on tool calls too. + allowed_server = MCPServerManager.resolve_oauth2_flow_for_request(allowed_server) + allowed_mcp_servers.append(allowed_server) + + allowed_mcp_servers = await _get_allowed_mcp_servers_from_mcp_server_names( + mcp_servers=mcp_servers, + allowed_mcp_servers=allowed_mcp_servers, + ) + if mcp_servers and not allowed_mcp_servers: + await raise_denied_scoped_mcp_access( + requested_names=mcp_servers, + user_api_key_auth=user_api_key_auth, + client_ip=client_ip, + ) + if not allowed_mcp_servers: + raise HTTPException( + status_code=403, + detail="User not allowed to call this tool.", + ) + + # Delegate to execute_mcp_tool for execution + response = await execute_mcp_tool( + name=name, + arguments=arguments, + allowed_mcp_servers=allowed_mcp_servers, + start_time=start_time, + user_api_key_auth=user_api_key_auth, + mcp_auth_header=mcp_auth_header, + mcp_server_auth_headers=mcp_server_auth_headers, + oauth2_headers=oauth2_headers, + raw_headers=raw_headers, + client_ip=client_ip, + **kwargs, + ) + except Exception as e: + await fire_mcp_tool_call_failure_logging(litellm_logging_obj, e, start_time, user_api_key_auth, kwargs) + raise + + if litellm_logging_obj: + response = await _fire_mcp_tool_call_logging( + logging_obj=litellm_logging_obj, + result=response, + start_time=start_time, + end_time=datetime.now(), + user_api_key_auth=user_api_key_auth, + request_data=kwargs, + ) + return response + + +async def mcp_get_prompt( + name: str, + arguments: dict[str, str] | None = None, + user_api_key_auth: UserAPIKeyAuth | None = None, + mcp_auth_header: str | None = None, + mcp_servers: list[str] | None = None, + mcp_server_auth_headers: dict[str, dict[str, str]] | None = None, + oauth2_headers: dict[str, str] | None = None, + raw_headers: dict[str, str] | None = None, + client_ip: str | None = None, +) -> GetPromptResult: + """ + Fetch a specific MCP prompt, handling both prefixed and unprefixed names. + """ + allowed_mcp_servers: Final = await _get_allowed_mcp_servers( + user_api_key_auth=user_api_key_auth, + mcp_servers=mcp_servers, + client_ip=client_ip, + ) + + if not allowed_mcp_servers: + raise HTTPException( + status_code=403, + detail="User not allowed to get this prompt.", + ) + + # Extract server name from prefixed prompt name + original_prompt_name, server_name = split_server_prefix_from_name(name) + + server: Final = next((s for s in allowed_mcp_servers if s.name == server_name), None) + if server is None: + raise HTTPException( + status_code=403, + detail="User not allowed to get this prompt.", + ) + + server_auth_header, extra_headers = _prepare_mcp_server_headers( + server=server, + mcp_server_auth_headers=mcp_server_auth_headers, + 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( + server=server, + user_api_key_auth=user_api_key_auth, + prompt_name=original_prompt_name, + arguments=arguments, + mcp_auth_header=server_auth_header, + extra_headers=extra_headers, + raw_headers=raw_headers, + client_ip=client_ip, + ) + + +async def mcp_read_resource( + url: AnyUrl, + user_api_key_auth: UserAPIKeyAuth | None = None, + mcp_auth_header: str | None = None, + mcp_servers: list[str] | None = None, + mcp_server_auth_headers: dict[str, dict[str, str]] | None = None, + oauth2_headers: dict[str, str] | None = None, + raw_headers: dict[str, str] | None = None, + client_ip: str | None = None, +) -> ReadResourceResult: + """Read resource contents from upstream MCP servers.""" + + allowed_mcp_servers: Final = await _get_allowed_mcp_servers( + user_api_key_auth=user_api_key_auth, + mcp_servers=mcp_servers, + client_ip=client_ip, + ) + + if not allowed_mcp_servers: + raise HTTPException( + status_code=403, + detail="User not allowed to read this resource.", + ) + + if len(allowed_mcp_servers) != 1: + raise HTTPException( + status_code=400, + detail=("Multiple MCP servers configured; read_resource currently supports exactly one allowed server."), + ) + + server: Final = allowed_mcp_servers[0] + + server_auth_header, extra_headers = _prepare_mcp_server_headers( + server=server, + mcp_server_auth_headers=mcp_server_auth_headers, + 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( + server=server, + user_api_key_auth=user_api_key_auth, + url=url, + mcp_auth_header=server_auth_header, + extra_headers=extra_headers, + raw_headers=raw_headers, + client_ip=client_ip, + ) + + +def _get_standard_logging_mcp_tool_call( + name: str, + arguments: dict[str, object], + server_name: str | None, + session_id: str | None = None, +) -> StandardLoggingMCPToolCall: + mcp_server: Final = global_mcp_server_manager._get_mcp_server_from_tool_name( + add_server_prefix_to_name(name, server_name) if server_name else name + ) + namespaced_tool_name: Final = f"{server_name}/{name}" if server_name else name + if mcp_server: + mcp_info: Final = mcp_server.mcp_info or {} + return StandardLoggingMCPToolCall( + name=name, + arguments=arguments, + mcp_server_name=mcp_info.get("server_name"), + mcp_server_logo_url=mcp_info.get("logo_url"), + namespaced_tool_name=namespaced_tool_name, + mcp_session_id=session_id, + mcp_auth_mode=mcp_server.auth_type, + mcp_server_resource=_redact_mcp_resource_url(mcp_server.url), + ) + else: + return StandardLoggingMCPToolCall( + name=name, + arguments=arguments, + namespaced_tool_name=namespaced_tool_name, + mcp_session_id=session_id, + ) + + +async def _handle_managed_mcp_tool( + server_name: str, + name: str, + arguments: dict[str, object], + user_api_key_auth: UserAPIKeyAuth | None = None, + mcp_auth_header: str | None = None, + mcp_server_auth_headers: dict[str, dict[str, str]] | None = None, + oauth2_headers: dict[str, str] | None = None, + raw_headers: dict[str, str] | None = None, + litellm_logging_obj: LiteLLMLoggingObj | None = None, + host_progress_callback: ProgressCallback | None = None, + guardrail_context: Mapping[str, object] | None = None, + client_ip: str | None = None, +) -> CallToolResult: + """Handle tool execution for managed server tools""" + # Import here to avoid circular import + from litellm.proxy.proxy_server import proxy_logging_obj + + call_tool_result: Final = await global_mcp_server_manager.call_tool( + server_name=server_name, + name=name, + arguments=arguments, + user_api_key_auth=user_api_key_auth, + mcp_auth_header=mcp_auth_header, + mcp_server_auth_headers=mcp_server_auth_headers, + oauth2_headers=oauth2_headers, + raw_headers=raw_headers, + client_ip=client_ip, + proxy_logging_obj=proxy_logging_obj, + host_progress_callback=host_progress_callback, + litellm_logging_obj=litellm_logging_obj, + guardrail_context=guardrail_context, + ) + verbose_logger.debug("CALL TOOL RESULT: %s", call_tool_result) + return call_tool_result + + +async def _handle_local_mcp_tool(name: str, arguments: dict[str, object]) -> CallToolResult: + """Execute a local-registry tool and report whether it succeeded. + + Returns the result rather than bare content because the verdict is part of it: the content + alone cannot say whether the handler failed, so callers used to stamp is_error=False on every + outcome and an upstream rejection was served as tool output. + + A failure is reported as ``is_error=True`` here rather than raised, because the REST surface + turns an unrecognized exception into a 500 and an upstream 403 or 429 is not a gateway crash. + ``MCPUpstreamAuthError`` is the exception: it propagates so the caller is told to + re-authenticate, which both renderers already know how to say. + + Note: Local tools don't use prefixes, so we use the original name + """ + import inspect + + tool: Final = global_mcp_tool_registry.get_tool(name) + if not tool: + raise HTTPException(status_code=404, detail=f"Tool '{name}' not found") + + try: + if inspect.iscoroutinefunction(tool.handler): + result = await tool.handler(**arguments) + else: + result = tool.handler(**arguments) + except MCPUpstreamAuthError: + raise + except Exception as e: + verbose_logger.exception("Error executing local tool %s: %s", name, e) + return CallToolResult( + content=[TextContent(text=f"Error: {e}", type="text")], # mutable-ok: MCP result content + is_error=True, + ) + return CallToolResult( + content=[TextContent(text=str(result), type="text")], # mutable-ok: MCP result content + is_error=False, + ) + + +_MCP_CREDENTIAL_REQUEST_FIELDS: Final = frozenset( + { + "raw_headers", + "mcp_auth_header", + "mcp_server_auth_headers", + "oauth2_headers", + "user_api_key_auth", + } +) + + +class _McpDeniedDetail(TypedDict): + error: ReadOnly[str] + + +async def _execute_handle_list_tools( + context: OperationContext, params: PaginatedRequestParams, host_progress_callback: ProgressCallback | None = None +) -> ListToolsResult: + try: + ( + user_api_key_auth, + mcp_auth_header, + mcp_servers, + mcp_server_auth_headers, + oauth2_headers, + raw_headers, + _client_ip, + ) = context.legacy_auth() + verbose_logger.debug("MCP list_tools - User API Key Auth from context: %s", user_api_key_auth) + verbose_logger.debug("MCP list_tools - MCP servers from context: %s", mcp_servers) + verbose_logger.debug( + "MCP list_tools - MCP server auth headers: %s", + list(mcp_server_auth_headers.keys()) if mcp_server_auth_headers else None, + ) + from mcp.types import Tool + + from litellm.proxy._experimental.mcp_server.tool_search import ( + get_mcp_proxy_tool_definitions, + get_virtual_tool_definitions, + ) + + if context.mcp_proxy_mode: + return ListToolsResult(tools=[Tool.model_validate(d) for d in get_mcp_proxy_tool_definitions()]) + if getattr( + getattr(user_api_key_auth, "object_permission", None), + "mcp_tool_search_enabled", + False, + ): + return ListToolsResult(tools=[Tool.model_validate(d) for d in get_virtual_tool_definitions()]) + + # Get mcp_servers from context variable + verbose_logger.debug("MCP list_tools - Calling _list_mcp_tools") + listing: Final = await _list_mcp_tools( + user_api_key_auth=user_api_key_auth, + mcp_auth_header=mcp_auth_header, + mcp_servers=mcp_servers, + mcp_server_auth_headers=mcp_server_auth_headers, + oauth2_headers=oauth2_headers, + raw_headers=raw_headers, + log_list_tools_to_spendlogs=True, + list_tools_log_source="mcp_protocol", + client_ip=_client_ip, + ) + verbose_logger.info("MCP list_tools - Successfully returned %s tools", len(listing.tools)) + if not listing.outcomes: + return ListToolsResult(tools=listing.tools) + outcome_meta: Final = { + SERVER_OUTCOMES_META_KEY: {key: outcome_wire_value(outcome) for key, outcome in listing.outcomes.items()} + } + return ListToolsResult.model_validate({"tools": listing.tools, "_meta": outcome_meta}) + except HTTPException as e: + from mcp.shared.exceptions import MCPError + from mcp.types import INVALID_REQUEST + + raise MCPError(code=INVALID_REQUEST, message=_http_detail_message(e.detail)) from e + except Exception as e: + verbose_logger.exception("Error in list_tools endpoint: %s", e) + # Return empty list instead of failing completely + # This prevents the HTTP stream from failing and allows the client to get a response + return ListToolsResult(tools=[]) # mutable-ok: MCP result payload + + +async def _execute_mcp_server_tool_call( + context: OperationContext, params: CallToolRequestParams, host_progress_callback: ProgressCallback | None = None +) -> CallToolResult: + from mcp.types import CallToolResult + + from litellm.exceptions import BlockedPiiEntityError, GuardrailRaisedException + from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request + from litellm.proxy.proxy_server import proxy_config + + ( + user_api_key_auth, + mcp_auth_header, + mcp_servers, + mcp_server_auth_headers, + oauth2_headers, + raw_headers, + _client_ip, + ) = context.legacy_auth() + verbose_logger.debug( + "MCP mcp_server_tool_call - user_api_key_auth=%s, user_role=%s", + user_api_key_auth, + getattr(user_api_key_auth, "user_role", "N/A"), + ) + + verbose_logger.debug("MCP mcp_server_tool_call - User API Key Auth from context: %s", user_api_key_auth) + + try: + # Inside this try so virtual-tool errors convert to isError + # CallToolResult instead of raising out of the protocol handler. + virtual_tool_result: Final = await _dispatch_virtual_mcp_tool( + name=params.name, + arguments=params.arguments, + user_api_key_auth=user_api_key_auth, + client_ip=_client_ip, + mcp_servers=mcp_servers, + mcp_auth_header=mcp_auth_header, + mcp_server_auth_headers=mcp_server_auth_headers, + oauth2_headers=oauth2_headers, + raw_headers=raw_headers, + mcp_proxy_mode=context.mcp_proxy_mode, + ) + if virtual_tool_result is not None: + return virtual_tool_result + + # Create a body date for logging + body_data: Final = {"name": params.name, "arguments": params.arguments} # mutable-ok: logging payload + # Set trace/session id from raw_headers so spend logs and logging_obj stay consistent (same as A2A) + chain_id: Final = get_chain_id_from_headers(raw_headers) + if chain_id: + body_data["litellm_trace_id"] = chain_id + body_data["litellm_session_id"] = chain_id + + request: Final = build_synthetic_mcp_request( + path="/mcp/tools/call", + raw_headers=raw_headers, + client_ip=_client_ip, + ) + if user_api_key_auth is not None: + data = await add_litellm_data_to_request( + data=body_data, + request=request, + # Bill a team-derived call to the team that granted it. A keyless admitted + # subject carries no team_id, so spend skipped team updates entirely and + # charged the user's PRIMARY org — the granting team's budget never + # accumulated (so it could never begin to block) and, cross-org, the wrong + # organization was charged. This is the ACCOUNTING half; the enforcement + # half (an already-over-budget team stops granting) lives in the source gate. + # Authorization is unaffected: it ran before this, and the union is resolved + # from the untouched auth object passed to call_mcp_tool below. + user_api_key_dict=await MCPRequestHandler.billing_auth_for_tool_call( + user_api_key_auth, tool_name=params.name + ), + proxy_config=proxy_config, + ) + else: + data = body_data + + response: Final = await call_mcp_tool( + user_api_key_auth=user_api_key_auth, + mcp_auth_header=mcp_auth_header, + mcp_servers=mcp_servers, + mcp_server_auth_headers=mcp_server_auth_headers, + oauth2_headers=oauth2_headers, + raw_headers=raw_headers, + client_ip=_client_ip, + host_progress_callback=host_progress_callback, + **data, # for logging + ) + except MCPMissingUserEnvVarsError as e: + verbose_logger.info( + "MCP mcp_server_tool_call missing per-user env vars: server_id=%s missing=%s", + e.server_id, + e.missing, + ) + return CallToolResult( + content=[TextContent(text=str(e), type="text")], + is_error=True, + ) + except BlockedPiiEntityError as e: + verbose_logger.error("BlockedPiiEntityError in MCP tool call: %s", e) + return CallToolResult( + content=[ + TextContent( + text=f"Error: Blocked PII entity detected - {e}", + type="text", + ) + ], + is_error=True, + ) + except GuardrailRaisedException as e: + verbose_logger.error("GuardrailRaisedException in MCP tool call: %s", e) + return CallToolResult( + content=[TextContent(text=f"Error: Guardrail violation - {e}", type="text")], + is_error=True, + ) + except HTTPException as e: + verbose_logger.error("HTTPException in MCP tool call: %s", e) + return CallToolResult( + content=[TextContent(text=f"Error: {_http_detail_message(e.detail)}", type="text")], + is_error=True, + ) + except MCPUpstreamAuthError as e: + # The MCP session manager serializes handler exceptions as JSON-RPC errors, so a + # mid-session tool call cannot emit a raw 401 + WWW-Authenticate the way the REST + # call path and the connect-time preemptive check do. Return an explicit isError + # naming the upstream status (at info level, not a traceback) so the client still + # learns it must re-authenticate upstream and expected pass-through 401s don't spam. + verbose_logger.info("Upstream auth failure calling MCP tool: HTTP %s", e.status_code) + return CallToolResult( + content=[ + TextContent( + text=f"Error: upstream authentication required (HTTP {e.status_code})", + type="text", + ) + ], + is_error=True, + ) + except Exception as e: + verbose_logger.exception("MCP mcp_server_tool_call - error: %s", e) + return CallToolResult( + content=[TextContent(text=f"Error: {e}", type="text")], + is_error=True, + ) + + return response + + +async def _execute_list_prompts( + context: OperationContext, params: PaginatedRequestParams, host_progress_callback: ProgressCallback | None = None +) -> ListPromptsResult: + if context.mcp_proxy_mode: + _reject_mcp_proxy_operation() + try: + ( + user_api_key_auth, + mcp_auth_header, + mcp_servers, + mcp_server_auth_headers, + oauth2_headers, + raw_headers, + _client_ip, + ) = context.legacy_auth() + verbose_logger.debug("MCP list_prompts - User API Key Auth from context: %s", user_api_key_auth) + verbose_logger.debug("MCP list_prompts - MCP servers from context: %s", mcp_servers) + verbose_logger.debug( + "MCP list_prompts - MCP server auth headers: %s", + list(mcp_server_auth_headers.keys()) if mcp_server_auth_headers else None, + ) + # Get mcp_servers from context variable + verbose_logger.debug("MCP list_prompts - Calling _list_prompts") + prompts: Final = await _list_mcp_prompts( + user_api_key_auth=user_api_key_auth, + mcp_auth_header=mcp_auth_header, + mcp_servers=mcp_servers, + mcp_server_auth_headers=mcp_server_auth_headers, + oauth2_headers=oauth2_headers, + raw_headers=raw_headers, + client_ip=_client_ip, + ) + verbose_logger.info("MCP list_prompts - Successfully returned %s prompts", len(prompts)) + return ListPromptsResult(prompts=prompts) + except Exception as e: + verbose_logger.exception("Error in list_prompts endpoint: %s", e) + # Return empty list instead of failing completely + # This prevents the HTTP stream from failing and allows the client to get a response + return ListPromptsResult(prompts=[]) # mutable-ok: MCP result payload + + +async def _execute_get_prompt( + context: OperationContext, params: GetPromptRequestParams, host_progress_callback: ProgressCallback | None = None +) -> GetPromptResult: + if context.mcp_proxy_mode: + _reject_mcp_proxy_operation() + ( + user_api_key_auth, + mcp_auth_header, + mcp_servers, + mcp_server_auth_headers, + oauth2_headers, + raw_headers, + _client_ip, + ) = context.legacy_auth() + + verbose_logger.debug("MCP mcp_server_tool_call - User API Key Auth from context: %s", user_api_key_auth) + return await mcp_get_prompt( + name=params.name, + arguments=params.arguments, + user_api_key_auth=user_api_key_auth, + mcp_auth_header=mcp_auth_header, + mcp_servers=mcp_servers, + mcp_server_auth_headers=mcp_server_auth_headers, + oauth2_headers=oauth2_headers, + raw_headers=raw_headers, + client_ip=_client_ip, + ) + + +async def _execute_list_resources( + context: OperationContext, params: PaginatedRequestParams, host_progress_callback: ProgressCallback | None = None +) -> ListResourcesResult: + if context.mcp_proxy_mode: + _reject_mcp_proxy_operation() + try: + ( + user_api_key_auth, + mcp_auth_header, + mcp_servers, + mcp_server_auth_headers, + oauth2_headers, + raw_headers, + _client_ip, + ) = context.legacy_auth() + verbose_logger.debug("MCP list_resources - User API Key Auth from context: %s", user_api_key_auth) + verbose_logger.debug("MCP list_resources - MCP servers from context: %s", mcp_servers) + verbose_logger.debug( + "MCP list_resources - MCP server auth headers: %s", + list(mcp_server_auth_headers.keys()) if mcp_server_auth_headers else None, + ) + + resources: Final = await _list_mcp_resources( + user_api_key_auth=user_api_key_auth, + mcp_auth_header=mcp_auth_header, + mcp_servers=mcp_servers, + mcp_server_auth_headers=mcp_server_auth_headers, + oauth2_headers=oauth2_headers, + raw_headers=raw_headers, + client_ip=_client_ip, + ) + verbose_logger.info("MCP list_resources - Successfully returned %s resources", len(resources)) + return ListResourcesResult(resources=resources) + except Exception as e: + verbose_logger.exception("Error in list_resources endpoint: %s", e) + return ListResourcesResult(resources=[]) # mutable-ok: MCP result payload + + +async def _execute_list_resource_templates( + context: OperationContext, params: PaginatedRequestParams, host_progress_callback: ProgressCallback | None = None +) -> ListResourceTemplatesResult: + if context.mcp_proxy_mode: + _reject_mcp_proxy_operation() + try: + ( + user_api_key_auth, + mcp_auth_header, + mcp_servers, + mcp_server_auth_headers, + oauth2_headers, + raw_headers, + _client_ip, + ) = context.legacy_auth() + verbose_logger.debug("MCP list_resource_templates - User API Key Auth from context: %s", user_api_key_auth) + verbose_logger.debug("MCP list_resource_templates - MCP servers from context: %s", mcp_servers) + verbose_logger.debug( + "MCP list_resource_templates - MCP server auth headers: %s", + list(mcp_server_auth_headers.keys()) if mcp_server_auth_headers else None, + ) + + resource_templates: Final = await _list_mcp_resource_templates( + user_api_key_auth=user_api_key_auth, + mcp_auth_header=mcp_auth_header, + mcp_servers=mcp_servers, + mcp_server_auth_headers=mcp_server_auth_headers, + oauth2_headers=oauth2_headers, + raw_headers=raw_headers, + client_ip=_client_ip, + ) + verbose_logger.info( + "MCP list_resource_templates - Successfully returned %s resource templates", len(resource_templates) + ) + return ListResourceTemplatesResult(resource_templates=resource_templates) + except Exception as e: + verbose_logger.exception("Error in list_resource_templates endpoint: %s", e) + return ListResourceTemplatesResult(resource_templates=[]) # mutable-ok: MCP result payload + + +async def _execute_read_resource( + context: OperationContext, params: ReadResourceRequestParams, host_progress_callback: ProgressCallback | None = None +) -> ReadResourceResult: + if context.mcp_proxy_mode: + _reject_mcp_proxy_operation() + ( + user_api_key_auth, + mcp_auth_header, + mcp_servers, + mcp_server_auth_headers, + oauth2_headers, + raw_headers, + _client_ip, + ) = context.legacy_auth() + + read_resource_result: Final = await mcp_read_resource( + url=params.uri, + user_api_key_auth=user_api_key_auth, + mcp_auth_header=mcp_auth_header, + mcp_servers=mcp_servers, + mcp_server_auth_headers=mcp_server_auth_headers, + oauth2_headers=oauth2_headers, + raw_headers=raw_headers, + client_ip=_client_ip, + ) + + return read_resource_result + + +def _reject_mcp_proxy_operation() -> NoReturn: + from mcp.shared.exceptions import MCPError + from mcp.types import METHOD_NOT_FOUND + + raise MCPError(code=METHOD_NOT_FOUND, message="Operation unavailable on /mcp/proxy") + + +def prepare_context( + user_api_key_auth: UserAPIKeyAuth | None = None, + mcp_auth_header: str | None = None, + mcp_servers: Sequence[str] | None = None, + mcp_server_auth_headers: Mapping[str, Mapping[str, str]] | None = None, + oauth2_headers: Mapping[str, str] | None = None, + raw_headers: Mapping[str, str] | None = None, + client_ip: str | None = None, + mcp_proxy_mode: bool = False, +) -> OperationContext: + return OperationContext( + _caller=user_api_key_auth, + mcp_auth_header=mcp_auth_header, + mcp_servers=tuple(mcp_servers) if mcp_servers is not None else None, + mcp_server_auth_headers=mcp_server_auth_headers, + oauth2_headers=oauth2_headers, + raw_headers=raw_headers, + client_ip=client_ip, + mcp_proxy_mode=mcp_proxy_mode, + ) + + +GatewayOperation: TypeAlias = ( + AuthorizedToolCall + | ListToolsRequest + | CallToolRequest + | ListPromptsRequest + | GetPromptRequest + | ListResourcesRequest + | ListResourceTemplatesRequest + | ReadResourceRequest +) +GatewayResult: TypeAlias = ( + ListToolsResult + | CallToolResult + | ListPromptsResult + | GetPromptResult + | ListResourcesResult + | ListResourceTemplatesResult + | ReadResourceResult +) + + +class GatewayOperations: + def __init__(self, host_progress_callback: ProgressCallback | None = None) -> None: + self._host_progress_callback = host_progress_callback + + @overload + async def execute(self, operation: AuthorizedToolCall, context: OperationContext) -> CallToolResult: ... + + @overload + async def execute(self, operation: ListToolsRequest, context: OperationContext) -> ListToolsResult: ... + + @overload + async def execute(self, operation: CallToolRequest, context: OperationContext) -> CallToolResult: ... + + @overload + async def execute(self, operation: ListPromptsRequest, context: OperationContext) -> ListPromptsResult: ... + + @overload + async def execute(self, operation: GetPromptRequest, context: OperationContext) -> GetPromptResult: ... + + @overload + async def execute(self, operation: ListResourcesRequest, context: OperationContext) -> ListResourcesResult: ... + + @overload + async def execute( + self, operation: ListResourceTemplatesRequest, context: OperationContext + ) -> ListResourceTemplatesResult: ... + + @overload + async def execute(self, operation: ReadResourceRequest, context: OperationContext) -> ReadResourceResult: ... + + async def execute(self, operation: GatewayOperation, context: OperationContext) -> GatewayResult: + match operation: + case AuthorizedToolCall(): + auth, token, _servers, server_headers, oauth_headers, headers, _client_ip = context.legacy_auth() + return await _execute_mcp_tool( + name=operation.name, + arguments=dict(operation.arguments), # mutable-ok: existing tool hooks own mutable argument data + allowed_mcp_servers=list( + operation.allowed_mcp_servers + ), # mutable-ok: legacy dispatch list contract + start_time=operation.start_time, + user_api_key_auth=auth, + mcp_auth_header=token, + mcp_server_auth_headers=server_headers, + oauth2_headers=oauth_headers, + raw_headers=headers, + client_ip=_client_ip, + host_progress_callback=operation.host_progress_callback, + guardrail_context=operation.guardrail_context, + **operation.logging_data, + ) + case ListToolsRequest(params=params): + return await _execute_handle_list_tools( + context, params or PaginatedRequestParams(), self._host_progress_callback + ) + case CallToolRequest(params=params): + return await _execute_mcp_server_tool_call(context, params, self._host_progress_callback) + case ListPromptsRequest(params=params): + return await _execute_list_prompts( + context, params or PaginatedRequestParams(), self._host_progress_callback + ) + case GetPromptRequest(params=params): + return await _execute_get_prompt(context, params, self._host_progress_callback) + case ListResourcesRequest(params=params): + return await _execute_list_resources( + context, params or PaginatedRequestParams(), self._host_progress_callback + ) + case ListResourceTemplatesRequest(params=params): + return await _execute_list_resource_templates( + context, params or PaginatedRequestParams(), self._host_progress_callback + ) + case ReadResourceRequest(params=params): + return await _execute_read_resource(context, params, self._host_progress_callback) + case _: + assert_never(operation) diff --git a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py index 15f97a15b73..c2f7bf7d531 100644 --- a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py @@ -203,17 +203,19 @@ if MCP_AVAILABLE: from litellm.proxy._experimental.mcp_server.oauth_utils import ( get_request_base_url, ) - from litellm.proxy._experimental.mcp_server.server import ( + from litellm.proxy._experimental.mcp_server.operations import ( ListMCPToolsRestAPIResponseObject, MCPInfo, MCPServer, - _aggregate_server_key, # pyright: ignore[reportPrivateUsage] # same per-server key as the tools/list _meta outcomes - _apply_toolset_scope, + _aggregate_server_key, _fire_mcp_tool_call_logging, execute_mcp_tool, filter_tools_by_allowed_tools, filter_tools_by_key_team_permissions, fire_mcp_tool_call_failure_logging, + ) + from litellm.proxy._experimental.mcp_server.server import ( + _apply_toolset_scope, reject_disallowed_mcp_client, ) @@ -670,6 +672,7 @@ if MCP_AVAILABLE: user_api_key_auth: UserAPIKeyAuth | None = None, extra_headers: dict[str, str] | None = None, apply_tool_filters: bool = True, + client_ip: str | None = None, ): """Helper function to get tools for a single server. @@ -684,6 +687,7 @@ if MCP_AVAILABLE: extra_headers=extra_headers, add_prefix=False, raw_headers=raw_headers, + client_ip=client_ip, user_api_key_auth=user_api_key_auth, ) @@ -797,6 +801,7 @@ if MCP_AVAILABLE: user_api_key_dict, extra_headers=user_oauth_extra_headers, apply_tool_filters=apply_tool_filters, + client_ip=rest_client_ip, ) except MCPUpstreamAuthError: # Surface the upstream 401/403 to the caller so it can emit the @@ -1016,6 +1021,7 @@ if MCP_AVAILABLE: user_api_key_dict, extra_headers=user_oauth_extra_headers, apply_tool_filters=apply_tool_filters, + client_ip=_rest_client_ip, ) except Exception as e: verbose_logger.warning( @@ -1193,6 +1199,7 @@ if MCP_AVAILABLE: mcp_server_auth_headers=data.get("mcp_server_auth_headers"), oauth2_headers=user_oauth_extra_headers or data.get("oauth2_headers"), raw_headers=data.get("raw_headers"), + client_ip=IPAddressUtils.get_mcp_client_ip(request), litellm_logging_obj=data.get("litellm_logging_obj"), guardrail_context=MCPRequestContext.resolve_guardrail_context(data), requested_server_id=canonical_server_id, diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 3a9bca926b0..33978bb9182 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -11,28 +11,22 @@ import hashlib import json import os import time -import traceback import types -import uuid from collections import Counter -from collections.abc import AsyncIterator, Callable, Iterable, Mapping, Sequence -from datetime import datetime -from typing import TYPE_CHECKING, Any, Final, NoReturn, Protocol +from collections.abc import AsyncGenerator, AsyncIterator, Callable, Iterable, Mapping, Sequence +from typing import TYPE_CHECKING, Final, NoReturn, Protocol import httpx from fastapi import FastAPI, HTTPException -from pydantic import AnyUrl, ConfigDict, Field, TypeAdapter, ValidationError +from pydantic import ConfigDict, TypeAdapter, ValidationError from starlette.requests import Request as StarletteRequest from starlette.responses import JSONResponse from starlette.types import Message, Receive, Scope, Send -from typing_extensions import ReadOnly, TypedDict from litellm._logging import verbose_logger from litellm.constants import ( - MAXIMUM_TRACEBACK_LINES_TO_LOG, MCP_GATEWAY_SESSION_ID_PREFIX_LENGTH, ) -from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj from litellm.llms.custom_httpx.http_handler import ( get_async_httpx_client, httpxSpecialProvider, @@ -41,12 +35,6 @@ from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import ( MCPRequestHandler, _is_mcp_admitted_user_subject, ) -from litellm.proxy._experimental.mcp_server.byok_credential_cache import ( - byok_credential_cache, - byok_credential_cache_key, - cache_byok_credential, - get_cached_byok_credential, -) from litellm.proxy._experimental.mcp_server.client_allowlist import ( MCPClientAllowlist, check_mcp_client_allowed, @@ -56,7 +44,6 @@ from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( get_request_base_url, ) from litellm.proxy._experimental.mcp_server.exceptions import ( - MCPToolResultError, MCPUpstreamAuthError, ) from litellm.proxy._experimental.mcp_server.mcp_context import ( @@ -74,7 +61,6 @@ from litellm.proxy._experimental.mcp_server.mcp_debug import ( ) from litellm.proxy._experimental.mcp_server.oauth_utils import ( _redact_mcp_resource_url, - get_byok_www_authenticate, get_passthrough_www_authenticate, get_route_relative_request_path, well_known_root_suffix, @@ -84,14 +70,6 @@ from litellm.proxy._experimental.mcp_server.utils import ( LITELLM_MCP_SERVER_DESCRIPTION, LITELLM_MCP_SERVER_NAME, LITELLM_MCP_SERVER_VERSION, - MCPMissingUserEnvVarsError, - add_server_prefix_to_name, - build_synthetic_mcp_request, - extract_mcp_tool_result_error_message, - get_server_prefix, - iter_known_server_prefixes, - logging_safe_mcp_headers, - match_known_tool_name, ) from litellm.proxy._types import ( ProxyException, @@ -99,13 +77,6 @@ from litellm.proxy._types import ( UserAPIKeyAuth, ) from litellm.proxy.auth.ip_address_utils import IPAddressUtils -from litellm.proxy.common_utils.auth_cache_invalidation_pubsub import ( - publish_auth_cache_invalidation, -) -from litellm.proxy.litellm_pre_call_utils import ( - LiteLLMProxyRequestSetup, - get_chain_id_from_headers, -) from litellm.types.mcp import ( MCPAuth, MCPGatewaySession, @@ -114,14 +85,11 @@ from litellm.types.mcp import ( MCPGatewaySessionsTerminateResponse, MCPSpecVersion, ) -from litellm.types.mcp_server.mcp_server_manager import MCPInfo, MCPServer -from litellm.types.utils import CallTypes, StandardLoggingMCPToolCall -from litellm.utils import Rules, client, function_setup +from litellm.types.mcp_server.mcp_server_manager import MCPServer if TYPE_CHECKING: from mcp.server.session import ServerSession as _McpServerSession - from litellm.proxy._experimental.mcp_server.db import OAuthCredentialPayload _STATEFUL_SESSION_IDLE_TIMEOUT_SECONDS: Final = 30 * 60 # Upper bound on concurrent stateful sessions a single caller may hold. Each @@ -159,13 +127,6 @@ def unsupported_protocol_version(scope: Scope) -> str | None: return None -async def _invalidate_byok_cred_cache(user_id: str, server_id: str) -> None: - """Drop a stored-or-deleted BYOK credential from this worker's cache and from every peer worker's.""" - cache_key: Final = byok_credential_cache_key(user_id, server_id) - byok_credential_cache.delete_cache(cache_key) - await publish_auth_cache_invalidation(cache_key=cache_key) - - # Check if MCP is available # "mcp" requires python 3.10 or higher, but several litellm users use python 3.8 # We're making this conditional import to avoid breaking users who use python 3.8. @@ -210,19 +171,6 @@ _SESSION_MANAGERS_INITIALIZED = False _INITIALIZATION_LOCK: Final = asyncio.Lock() -def _mcp_session_id_from_headers( - raw_headers: dict[str, str] | None, -) -> str | None: - """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 - - def _jsonrpc_text_has_top_level_method(text: str) -> bool: """Whether a (possibly truncated) JSON-RPC envelope has a ``method`` key at the root object's top level. @@ -466,6 +414,56 @@ def _proxy_exception_to_http_exception(exc: ProxyException) -> HTTPException: if MCP_AVAILABLE: + __all__ = ( + "_MCP_CREDENTIAL_REQUEST_FIELDS", + "ListMCPToolsRestAPIResponseObject", + "_McpDeniedDetail", + "_aggregate_server_key", + "_build_virtual_call_logging_obj", + "_check_byok_credential", + "_client_has_passthrough_authorization", + "_client_has_per_server_auth_header", + "_dispatch_virtual_mcp_tool", + "_fire_mcp_tool_call_logging", + "_get_allowed_mcp_servers", + "_get_allowed_mcp_servers_from_mcp_server_names", + "_get_byok_credential", + "_get_prompts_from_mcp_servers", + "_get_resource_templates_from_mcp_servers", + "_get_resources_from_mcp_servers", + "_get_standard_logging_mcp_tool_call", + "_get_tools_from_mcp_servers", + "_get_user_oauth_extra_headers_from_db", + "_handle_local_mcp_tool", + "_handle_managed_mcp_tool", + "_http_detail_message", + "_invalidate_byok_cred_cache", + "_list_mcp_prompts", + "_list_mcp_resource_templates", + "_list_mcp_resources", + "_list_mcp_tools", + "_list_tools_before_first_call", + "_mcp_session_id_from_headers", + "_merge_gateway_initialize_instructions", + "_prefetch_oauth_creds_for_user", + "_prepare_mcp_server_headers", + "_raise_if_initialize_grants_no_mcp_servers", + "_redact_mcp_resource_url", + "_resolve_display_name_to_original", + "_run_post_mcp_call_guardrails", + "_server_answers_to", + "_tool_name_matches", + "apply_tool_overrides", + "call_mcp_tool", + "execute_mcp_tool", + "filter_tools_by_allowed_tools", + "filter_tools_by_key_team_permissions", + "fire_mcp_tool_call_failure_logging", + "global_mcp_server_manager", + "mcp_get_prompt", + "mcp_read_resource", + "raise_denied_scoped_mcp_access", + ) from mcp.server import Server # Import auth context variables and middleware @@ -476,6 +474,23 @@ if MCP_AVAILABLE: from mcp.server.context import ServerRequestContext from mcp.server.lowlevel.server import NotificationOptions from mcp.server.models import InitializationOptions + from mcp.shared.exceptions import MCPError + from mcp.types import ( + CallToolRequest, + GetPromptRequest, + ListPromptsRequest, + ListResourcesRequest, + ListResourceTemplatesRequest, + ListToolsRequest, + ReadResourceRequest, + ) + + from litellm.proxy._experimental.mcp_server import operations + from litellm.proxy._experimental.mcp_server.contracts import OperationContext + from litellm.proxy._experimental.mcp_server.operations import ( + _invalidate_byok_cred_cache, + _mcp_session_id_from_headers, + ) try: from mcp.server.streamable_http_manager import StreamableHTTPSessionManager @@ -493,62 +508,27 @@ if MCP_AVAILABLE: ListResourceTemplatesResult, ListToolsResult, PaginatedRequestParams, - Prompt, ReadResourceRequestParams, - TextContent, ) - from mcp.types import Tool as MCPTool from mcp_types.version import HANDSHAKE_PROTOCOL_VERSIONS from litellm.proxy._experimental.mcp_server.auth.litellm_auth_handler import ( MCPAuthenticatedUser, ) - from litellm.proxy._experimental.mcp_server.faults.list_outcomes import ( - SERVER_OUTCOMES_META_KEY, - AggregateToolListing, - ServerListOk, - ServerOutcome, - classify_list_exception, - outcome_wire_value, - ) from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( MCPServerManager, - _caller_authorization_fans_out, - _client_forwarded_authorization_headers, - _resolve_openapi_tool_auth, - _should_strip_caller_authorization, global_mcp_server_manager, ) - from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import ( - _request_auth_header, - _request_extra_headers, - _request_resolved_auth_headers, - ) - from litellm.proxy._experimental.mcp_server.sse_transport import SseServerTransport - from litellm.proxy._experimental.mcp_server.tool_registry import ( - global_mcp_tool_registry, - ) - from litellm.proxy._experimental.mcp_server.utils import ( - MCP_TOOL_PREFIX_SEPARATOR, - is_tool_name_prefixed, - normalize_server_name, - split_server_prefix_from_name, - strip_known_server_prefix, - ) - from litellm.types.mcp import DEFAULT_CREDENTIAL_HEADER, without_header ###################################################### ############ MCP Tools List REST API Response Object # # Defined here because we don't want to add `mcp` as a # required dependency for `litellm` pip package ###################################################### - class ListMCPToolsRestAPIResponseObject(MCPTool): - """ - Object returned by the /tools/list REST API route. - """ - - mcp_info: MCPInfo | None = Field(default=None, alias="mcp_info") - model_config = ConfigDict(arbitrary_types_allowed=True) + from litellm.proxy._experimental.mcp_server.operations import ( + ListMCPToolsRestAPIResponseObject, + ) + from litellm.proxy._experimental.mcp_server.sse_transport import SseServerTransport def _gateway_create_initialization_options( self, @@ -818,94 +798,45 @@ if MCP_AVAILABLE: ############### MCP Server Routes ####################### ######################################################## - async def handle_list_tools(ctx: ServerRequestContext, params: PaginatedRequestParams) -> ListToolsResult: - """ - List all available tools, with each server's listing outcome attached to the result's - ``_meta`` (SERVER_OUTCOMES_META_KEY) so a broken upstream is distinguishable from a healthy - server with no tools. Returning a ListToolsResult (rather than a bare list) makes the MCP SDK - pass the result through unwrapped, which is what lets the ``_meta`` survive to the client. - Also captures the active session for propagation to callbacks. - """ - req_ctx: Final = ctx - _ctx_reset_token: Final = active_mcp_request_ctx_var.set(ctx) - _session_reset_token: Final = active_mcp_session_var.set(ctx.session) - _trace_token = None - _transport_token = None - _destinations_token = None - - try: - _trace_token = _otel_set_mcp_trace_carrier(_mcp_meta_trace_carrier(req_ctx)) - _transport_token = _otel_set_mcp_transport_span(_otel_transport_span_from_message(req_ctx)) - _destinations_token = _otel_set_mcp_request_destinations(req_ctx) - # Get user authentication from context variable + @contextlib.asynccontextmanager + async def _legacy_operation_context(ctx: ServerRequestContext, *, trace: bool) -> AsyncGenerator[OperationContext]: + with contextlib.ExitStack() as cleanup: + cleanup.callback(active_mcp_request_ctx_var.reset, active_mcp_request_ctx_var.set(ctx)) + cleanup.callback(active_mcp_session_var.reset, active_mcp_session_var.set(ctx.session)) + if trace: + cleanup.callback( + _otel_reset_mcp_trace_carrier, _otel_set_mcp_trace_carrier(_mcp_meta_trace_carrier(ctx)) + ) + cleanup.callback( + _otel_reset_mcp_transport_span, _otel_set_mcp_transport_span(_otel_transport_span_from_message(ctx)) + ) + cleanup.callback(_otel_reset_mcp_request_destinations, _otel_set_mcp_request_destinations(ctx)) ( - user_api_key_auth, - mcp_auth_header, - mcp_servers, - mcp_server_auth_headers, - oauth2_headers, - raw_headers, - _client_ip, + auth, + token, + servers, + server_headers, + oauth_headers, + headers, + client_ip, ) = await get_or_extract_auth_context() - verbose_logger.debug("MCP list_tools - User API Key Auth from context: %s", user_api_key_auth) - verbose_logger.debug("MCP list_tools - MCP servers from context: %s", mcp_servers) - verbose_logger.debug( - "MCP list_tools - MCP server auth headers: %s", - list(mcp_server_auth_headers.keys()) if mcp_server_auth_headers else None, - ) - from mcp.types import Tool - - from litellm.proxy._experimental.mcp_server.tool_search import ( - get_mcp_proxy_tool_definitions, - get_virtual_tool_definitions, + yield operations.prepare_context( + auth, token, servers, server_headers, oauth_headers, headers, client_ip, _mcp_proxy_mode.get() ) - if _mcp_proxy_mode.get(): - return ListToolsResult(tools=[Tool.model_validate(d) for d in get_mcp_proxy_tool_definitions()]) - if getattr( - getattr(user_api_key_auth, "object_permission", None), - "mcp_tool_search_enabled", - False, - ): - return ListToolsResult(tools=[Tool.model_validate(d) for d in get_virtual_tool_definitions()]) - - # Get mcp_servers from context variable - verbose_logger.debug("MCP list_tools - Calling _list_mcp_tools") - listing: Final = await _list_mcp_tools( - user_api_key_auth=user_api_key_auth, - mcp_auth_header=mcp_auth_header, - mcp_servers=mcp_servers, - mcp_server_auth_headers=mcp_server_auth_headers, - oauth2_headers=oauth2_headers, - raw_headers=raw_headers, - log_list_tools_to_spendlogs=True, - list_tools_log_source="mcp_protocol", - ) - verbose_logger.info("MCP list_tools - Successfully returned %s tools", len(listing.tools)) - if not listing.outcomes: - return ListToolsResult(tools=listing.tools) - outcome_meta: Final = { - SERVER_OUTCOMES_META_KEY: { - key: outcome_wire_value(outcome) for key, outcome in listing.outcomes.items() - } - } - return ListToolsResult.model_validate({"tools": listing.tools, "_meta": outcome_meta}) - except HTTPException as e: - from mcp.shared.exceptions import MCPError - from mcp.types import INVALID_REQUEST - - raise MCPError(code=INVALID_REQUEST, message=_http_detail_message(e.detail)) from e - except Exception as e: - verbose_logger.exception("Error in list_tools endpoint: %s", e) - # Return empty list instead of failing completely - # This prevents the HTTP stream from failing and allows the client to get a response - return ListToolsResult(tools=[]) # mutable-ok: MCP result payload - finally: - _otel_reset_mcp_request_destinations(_destinations_token) - _otel_reset_mcp_transport_span(_transport_token) - _otel_reset_mcp_trace_carrier(_trace_token) - active_mcp_session_var.reset(_session_reset_token) - active_mcp_request_ctx_var.reset(_ctx_reset_token) + async def handle_list_tools(ctx: ServerRequestContext, params: PaginatedRequestParams) -> ListToolsResult: + try: + async with _legacy_operation_context(ctx, trace=True) as context: + return await operations.GatewayOperations(_capture_host_progress_callback(ctx)).execute( + ListToolsRequest(params=params), context + ) + except MCPError: + raise + except HTTPException as exc: + raise MCPError(code=INVALID_REQUEST, message=operations._http_detail_message(exc.detail)) from exc + except Exception as exc: # noqa: BLE001 # preserve native listing fallback for ingress failures + verbose_logger.exception("Error in list_tools endpoint: %s", exc) + return ListToolsResult(tools=[]) def _capture_host_progress_callback(ctx: ServerRequestContext) -> Callable | None: """Return a progress-forwarding callback bound to the host MCP session. @@ -942,581 +873,71 @@ if MCP_AVAILABLE: raise MCPError(code=METHOD_NOT_FOUND, message="Operation unavailable on /mcp/proxy") - async def _build_virtual_call_logging_obj( - name: str, - arguments: dict[str, object], - user_api_key_auth: UserAPIKeyAuth, - raw_headers: Mapping[str, str] | None = None, - client_ip: str | None = None, - ) -> LiteLLMLoggingObj | None: - """Run the pre-call pipeline (guardrails + logging setup) for a virtual - mcp_tool_call so the SSE path spend-logs like the REST path.""" - from litellm.proxy.common_request_processing import ( - ProxyBaseLLMRequestProcessing, - ) - from litellm.proxy.proxy_server import ( - general_settings, - proxy_config, - proxy_logging_obj, - ) - - request: Final = build_synthetic_mcp_request( - path="/mcp/tools/call", - raw_headers=raw_headers, - client_ip=client_ip, - ) - _, virtual_logging_obj = await ProxyBaseLLMRequestProcessing( - data={"name": name, "arguments": arguments} - ).common_processing_pre_call_logic( - request=request, - user_api_key_dict=user_api_key_auth, - proxy_config=proxy_config, - route_type=CallTypes.call_mcp_tool.value, - proxy_logging_obj=proxy_logging_obj, - general_settings=general_settings, - ) - return virtual_logging_obj - - async def _dispatch_virtual_mcp_tool( - name: str, - arguments: dict[str, object] | None, - user_api_key_auth: UserAPIKeyAuth | None, - client_ip: str | None, - mcp_servers: list[str] | None = None, - mcp_auth_header: str | None = None, - mcp_server_auth_headers: dict[str, dict[str, str]] | None = None, - oauth2_headers: dict[str, str] | None = None, - raw_headers: dict[str, str] | None = None, - ) -> CallToolResult | None: - """Handle the mcp_tool_search / mcp_tool_call virtual tools. - - Returns a CallToolResult when ``name`` is a virtual tool, else ``None`` so - the caller falls through to normal tool routing. - """ - from litellm.llms.litellm_proxy.skills.skill_search import DEFAULT_SKILL_SEARCH_TOP_K - from litellm.proxy._experimental.mcp_server.tool_search import ( - AGENT_SEARCH_TOOL_NAME, - DEFAULT_AGENT_SEARCH_TOP_K, - MCP_PROXY_CALL_TOOL_NAME, - MCP_PROXY_TOOL_NAMES, - MCP_TOOL_SEARCH_TOOL_NAME, - SKILL_SEARCH_TOOL_NAME, - VIRTUAL_TOOL_NAMES, - coerce_top_k, - handle_agent_search, - handle_mcp_proxy_tool, - handle_mcp_tool_call, - handle_mcp_tool_search, - handle_skill_search, - ) - - if _mcp_proxy_mode.get() and name not in MCP_PROXY_TOOL_NAMES: - return CallToolResult( - content=[ # mutable-ok: MCP result content - TextContent(type="text", text=f"Tool {name} is unavailable on /mcp/proxy") - ], - is_error=True, - ) - - if _mcp_proxy_mode.get() and name in MCP_PROXY_TOOL_NAMES: - assert user_api_key_auth is not None - proxy_call_start: Final = datetime.now() # noqa: DTZ005 # logging pipeline uses naive datetimes - proxy_logging_obj: Final = ( - await _build_virtual_call_logging_obj( - name=name, - arguments=arguments or {}, # mutable-ok: logging pipeline payload - user_api_key_auth=user_api_key_auth, - raw_headers=raw_headers, - client_ip=client_ip, - ) - if name == MCP_PROXY_CALL_TOOL_NAME - else None - ) - try: - proxy_result: Final = await handle_mcp_proxy_tool( - name=name, - arguments=arguments or {}, # mutable-ok: proxy handler payload - user_api_key_dict=user_api_key_auth, - client_ip=client_ip, - mcp_servers=mcp_servers, - mcp_auth_header=mcp_auth_header, - mcp_server_auth_headers=mcp_server_auth_headers, - oauth2_headers=oauth2_headers, - raw_headers=raw_headers, - litellm_logging_obj=proxy_logging_obj, - ) - except Exception as exc: - if proxy_logging_obj is not None: - from litellm.proxy.proxy_server import proxy_logging_obj as request_logging_obj - - failure_end: Final = datetime.now() # noqa: DTZ005 # matches the logging pipeline start time - failure_traceback: Final = traceback.format_exc(limit=MAXIMUM_TRACEBACK_LINES_TO_LOG) - try: - proxy_logging_obj.failure_handler(exc, failure_traceback, proxy_call_start, failure_end) - await proxy_logging_obj.async_failure_handler( - exc, failure_traceback, proxy_call_start, failure_end - ) - if not isinstance(exc, MCPUpstreamAuthError): - await request_logging_obj.post_call_failure_hook( - request_data={ # mutable-ok: failure hook mutates its request payload - "name": name, - "arguments": arguments, - "litellm_logging_obj": proxy_logging_obj, - }, - original_exception=exc, - user_api_key_dict=user_api_key_auth, - route="/mcp/call_tool", - traceback_str=failure_traceback, - ) - except Exception: # noqa: BLE001 # a failing failure hook must not mask the tool call's own error - verbose_logger.exception("Error logging failed MCP proxy tool call") - raise - if proxy_logging_obj is not None: - return await _fire_mcp_tool_call_logging( - logging_obj=proxy_logging_obj, - result=proxy_result, - start_time=proxy_call_start, - end_time=datetime.now(), # noqa: DTZ005 # matches the logging pipeline start time - user_api_key_auth=user_api_key_auth, - request_data=types.MappingProxyType({"name": name, "arguments": arguments}), - ) - return proxy_result - - if name not in VIRTUAL_TOOL_NAMES: - return None - - if not getattr( - getattr(user_api_key_auth, "object_permission", None), - "mcp_tool_search_enabled", - False, - ): - return CallToolResult( - content=[ - TextContent( - type="text", - text=f"Tool {name} requires mcp_tool_search_enabled on the key", - ) - ], - is_error=True, - ) - - args: Final = arguments or {} - if name == MCP_TOOL_SEARCH_TOOL_NAME: - return await handle_mcp_tool_search( - query=args.get("query", ""), - top_k=coerce_top_k(args.get("top_k", 5)), - user_api_key_dict=user_api_key_auth, - client_ip=client_ip, - mcp_servers=mcp_servers, - mcp_auth_header=mcp_auth_header, - mcp_server_auth_headers=mcp_server_auth_headers, - oauth2_headers=oauth2_headers, - raw_headers=raw_headers, - ) - - assert user_api_key_auth is not None # guaranteed by the flag check above - if name == AGENT_SEARCH_TOOL_NAME: - return await handle_agent_search( - query=str(args.get("query", "")), - top_k=coerce_top_k(args.get("top_k", DEFAULT_AGENT_SEARCH_TOP_K), default=DEFAULT_AGENT_SEARCH_TOP_K), - user_api_key_dict=user_api_key_auth, - ) - if name == SKILL_SEARCH_TOOL_NAME: - return await handle_skill_search( - query=str(args.get("query", "")), - top_k=coerce_top_k(args.get("top_k", DEFAULT_SKILL_SEARCH_TOP_K), default=DEFAULT_SKILL_SEARCH_TOP_K), - user_api_key_dict=user_api_key_auth, - ) - virtual_logging_obj: Final = await _build_virtual_call_logging_obj( - name=name, - arguments=args, - user_api_key_auth=user_api_key_auth, - raw_headers=raw_headers, - client_ip=client_ip, - ) - return await handle_mcp_tool_call( - tool_name=args.get("tool_name", ""), - arguments=args.get("arguments") or {}, - user_api_key_dict=user_api_key_auth, - client_ip=client_ip, - mcp_servers=mcp_servers, - mcp_auth_header=mcp_auth_header, - mcp_server_auth_headers=mcp_server_auth_headers, - oauth2_headers=oauth2_headers, - raw_headers=raw_headers, - litellm_logging_obj=virtual_logging_obj, - ) + from litellm.proxy._experimental.mcp_server.operations import ( + _build_virtual_call_logging_obj, + _dispatch_virtual_mcp_tool, + ) async def mcp_server_tool_call(ctx: ServerRequestContext, params: CallToolRequestParams) -> CallToolResult: - """ - Call a specific tool with the provided arguments - Args: - ctx: SDK request context carrying the client session and HTTP request - params (CallToolRequestParams): Tool name and arguments - Returns: - CallToolResult: Tool execution results - """ - from mcp.types import CallToolResult - - from litellm.exceptions import BlockedPiiEntityError, GuardrailRaisedException - from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request - from litellm.proxy.proxy_server import proxy_config - - req_ctx: Final = ctx - _ctx_reset_token: Final = active_mcp_request_ctx_var.set(ctx) - _session_reset_token: Final = active_mcp_session_var.set(ctx.session) - _trace_token = None - _transport_token = None - _destinations_token = None - - try: - _trace_token = _otel_set_mcp_trace_carrier(_mcp_meta_trace_carrier(req_ctx)) - _transport_token = _otel_set_mcp_transport_span(_otel_transport_span_from_message(req_ctx)) - _destinations_token = _otel_set_mcp_request_destinations(req_ctx) - # Validate arguments - ( - user_api_key_auth, - mcp_auth_header, - mcp_servers, - mcp_server_auth_headers, - oauth2_headers, - raw_headers, - _client_ip, - ) = await get_or_extract_auth_context() - verbose_logger.debug( - "MCP mcp_server_tool_call - user_api_key_auth=%s, user_role=%s", - user_api_key_auth, - getattr(user_api_key_auth, "user_role", "N/A"), + async with _legacy_operation_context(ctx, trace=True) as context: + return await operations.GatewayOperations(_capture_host_progress_callback(ctx)).execute( + CallToolRequest(params=params), context ) - verbose_logger.debug("MCP mcp_server_tool_call - User API Key Auth from context: %s", user_api_key_auth) - - try: - # Inside this try so virtual-tool errors convert to isError - # CallToolResult instead of raising out of the protocol handler. - virtual_tool_result: Final = await _dispatch_virtual_mcp_tool( - name=params.name, - arguments=params.arguments, - user_api_key_auth=user_api_key_auth, - client_ip=_client_ip, - mcp_servers=mcp_servers, - mcp_auth_header=mcp_auth_header, - mcp_server_auth_headers=mcp_server_auth_headers, - oauth2_headers=oauth2_headers, - raw_headers=raw_headers, - ) - if virtual_tool_result is not None: - return virtual_tool_result - - host_progress_callback: Final = _capture_host_progress_callback(ctx) - # Create a body date for logging - body_data: Final = {"name": params.name, "arguments": params.arguments} # mutable-ok: logging payload - # Set trace/session id from raw_headers so spend logs and logging_obj stay consistent (same as A2A) - chain_id: Final = get_chain_id_from_headers(raw_headers) - if chain_id: - body_data["litellm_trace_id"] = chain_id - body_data["litellm_session_id"] = chain_id - - request: Final = build_synthetic_mcp_request( - path="/mcp/tools/call", - raw_headers=raw_headers, - client_ip=_client_ip, - ) - if user_api_key_auth is not None: - data = await add_litellm_data_to_request( - data=body_data, - request=request, - # Bill a team-derived call to the team that granted it. A keyless admitted - # subject carries no team_id, so spend skipped team updates entirely and - # charged the user's PRIMARY org — the granting team's budget never - # accumulated (so it could never begin to block) and, cross-org, the wrong - # organization was charged. This is the ACCOUNTING half; the enforcement - # half (an already-over-budget team stops granting) lives in the source gate. - # Authorization is unaffected: it ran before this, and the union is resolved - # from the untouched auth object passed to call_mcp_tool below. - user_api_key_dict=await MCPRequestHandler.billing_auth_for_tool_call( - user_api_key_auth, tool_name=params.name - ), - proxy_config=proxy_config, - ) - else: - data = body_data - - response: Final = await call_mcp_tool( - user_api_key_auth=user_api_key_auth, - mcp_auth_header=mcp_auth_header, - mcp_servers=mcp_servers, - mcp_server_auth_headers=mcp_server_auth_headers, - oauth2_headers=oauth2_headers, - raw_headers=raw_headers, - client_ip=_client_ip, - host_progress_callback=host_progress_callback, - **data, # for logging - ) - except MCPMissingUserEnvVarsError as e: - verbose_logger.info( - "MCP mcp_server_tool_call missing per-user env vars: server_id=%s missing=%s", - e.server_id, - e.missing, - ) - return CallToolResult( - content=[TextContent(text=str(e), type="text")], - is_error=True, - ) - except BlockedPiiEntityError as e: - verbose_logger.error("BlockedPiiEntityError in MCP tool call: %s", e) - return CallToolResult( - content=[ - TextContent( - text=f"Error: Blocked PII entity detected - {e}", - type="text", - ) - ], - is_error=True, - ) - except GuardrailRaisedException as e: - verbose_logger.error("GuardrailRaisedException in MCP tool call: %s", e) - return CallToolResult( - content=[TextContent(text=f"Error: Guardrail violation - {e}", type="text")], - is_error=True, - ) - except HTTPException as e: - verbose_logger.error("HTTPException in MCP tool call: %s", e) - return CallToolResult( - content=[TextContent(text=f"Error: {_http_detail_message(e.detail)}", type="text")], - is_error=True, - ) - except MCPUpstreamAuthError as e: - # The MCP session manager serializes handler exceptions as JSON-RPC errors, so a - # mid-session tool call cannot emit a raw 401 + WWW-Authenticate the way the REST - # call path and the connect-time preemptive check do. Return an explicit isError - # naming the upstream status (at info level, not a traceback) so the client still - # learns it must re-authenticate upstream and expected pass-through 401s don't spam. - verbose_logger.info("Upstream auth failure calling MCP tool: HTTP %s", e.status_code) - return CallToolResult( - content=[ - TextContent( - text=f"Error: upstream authentication required (HTTP {e.status_code})", - type="text", - ) - ], - is_error=True, - ) - except Exception as e: - verbose_logger.exception("MCP mcp_server_tool_call - error: %s", e) - return CallToolResult( - content=[TextContent(text=f"Error: {e}", type="text")], - is_error=True, - ) - - return response - finally: - _otel_reset_mcp_request_destinations(_destinations_token) - _otel_reset_mcp_transport_span(_transport_token) - _otel_reset_mcp_trace_carrier(_trace_token) - active_mcp_session_var.reset(_session_reset_token) - active_mcp_request_ctx_var.reset(_ctx_reset_token) - async def list_prompts(ctx: ServerRequestContext, params: PaginatedRequestParams) -> ListPromptsResult: - """ - List all available prompts - """ if _mcp_proxy_mode.get(): _reject_mcp_proxy_operation() - _ctx_reset_token: Final = active_mcp_request_ctx_var.set(ctx) - _session_reset_token: Final = active_mcp_session_var.set(ctx.session) - try: - # Get user authentication from context variable - ( - user_api_key_auth, - mcp_auth_header, - mcp_servers, - mcp_server_auth_headers, - oauth2_headers, - raw_headers, - _client_ip, - ) = await get_or_extract_auth_context() - verbose_logger.debug("MCP list_prompts - User API Key Auth from context: %s", user_api_key_auth) - verbose_logger.debug("MCP list_prompts - MCP servers from context: %s", mcp_servers) - verbose_logger.debug( - "MCP list_prompts - MCP server auth headers: %s", - list(mcp_server_auth_headers.keys()) if mcp_server_auth_headers else None, - ) - # Get mcp_servers from context variable - verbose_logger.debug("MCP list_prompts - Calling _list_prompts") - prompts: Final = await _list_mcp_prompts( - user_api_key_auth=user_api_key_auth, - mcp_auth_header=mcp_auth_header, - mcp_servers=mcp_servers, - mcp_server_auth_headers=mcp_server_auth_headers, - oauth2_headers=oauth2_headers, - raw_headers=raw_headers, - ) - verbose_logger.info("MCP list_prompts - Successfully returned %s prompts", len(prompts)) - return ListPromptsResult(prompts=prompts) - except Exception as e: - verbose_logger.exception("Error in list_prompts endpoint: %s", e) - # Return empty list instead of failing completely - # This prevents the HTTP stream from failing and allows the client to get a response - return ListPromptsResult(prompts=[]) # mutable-ok: MCP result payload - finally: - active_mcp_session_var.reset(_session_reset_token) - active_mcp_request_ctx_var.reset(_ctx_reset_token) + async with _legacy_operation_context(ctx, trace=False) as context: + return await operations.GatewayOperations(_capture_host_progress_callback(ctx)).execute( + ListPromptsRequest(params=params), context + ) + except Exception as exc: # noqa: BLE001 # preserve native listing fallback for ingress failures + verbose_logger.exception("Error in list_prompts endpoint: %s", exc) + return ListPromptsResult(prompts=[]) async def get_prompt(ctx: ServerRequestContext, params: GetPromptRequestParams) -> GetPromptResult: - """ - Get a specific prompt with the provided arguments - """ if _mcp_proxy_mode.get(): _reject_mcp_proxy_operation() - _ctx_reset_token: Final = active_mcp_request_ctx_var.set(ctx) - _session_reset_token: Final = active_mcp_session_var.set(ctx.session) - - try: - ( - user_api_key_auth, - mcp_auth_header, - mcp_servers, - mcp_server_auth_headers, - oauth2_headers, - raw_headers, - _client_ip, - ) = await get_or_extract_auth_context() - - verbose_logger.debug("MCP mcp_server_tool_call - User API Key Auth from context: %s", user_api_key_auth) - return await mcp_get_prompt( - name=params.name, - arguments=params.arguments, - user_api_key_auth=user_api_key_auth, - mcp_auth_header=mcp_auth_header, - mcp_servers=mcp_servers, - mcp_server_auth_headers=mcp_server_auth_headers, - oauth2_headers=oauth2_headers, - raw_headers=raw_headers, + async with _legacy_operation_context(ctx, trace=False) as context: + return await operations.GatewayOperations(_capture_host_progress_callback(ctx)).execute( + GetPromptRequest(params=params), context ) - finally: - active_mcp_session_var.reset(_session_reset_token) - active_mcp_request_ctx_var.reset(_ctx_reset_token) async def list_resources(ctx: ServerRequestContext, params: PaginatedRequestParams) -> ListResourcesResult: - """List all available resources.""" if _mcp_proxy_mode.get(): _reject_mcp_proxy_operation() - _ctx_reset_token: Final = active_mcp_request_ctx_var.set(ctx) - _session_reset_token: Final = active_mcp_session_var.set(ctx.session) - try: - ( - user_api_key_auth, - mcp_auth_header, - mcp_servers, - mcp_server_auth_headers, - oauth2_headers, - raw_headers, - _client_ip, - ) = await get_or_extract_auth_context() - verbose_logger.debug("MCP list_resources - User API Key Auth from context: %s", user_api_key_auth) - verbose_logger.debug("MCP list_resources - MCP servers from context: %s", mcp_servers) - verbose_logger.debug( - "MCP list_resources - MCP server auth headers: %s", - list(mcp_server_auth_headers.keys()) if mcp_server_auth_headers else None, - ) - - resources: Final = await _list_mcp_resources( - user_api_key_auth=user_api_key_auth, - mcp_auth_header=mcp_auth_header, - mcp_servers=mcp_servers, - mcp_server_auth_headers=mcp_server_auth_headers, - oauth2_headers=oauth2_headers, - raw_headers=raw_headers, - ) - verbose_logger.info("MCP list_resources - Successfully returned %s resources", len(resources)) - return ListResourcesResult(resources=resources) - except Exception as e: - verbose_logger.exception("Error in list_resources endpoint: %s", e) - return ListResourcesResult(resources=[]) # mutable-ok: MCP result payload - finally: - active_mcp_session_var.reset(_session_reset_token) - active_mcp_request_ctx_var.reset(_ctx_reset_token) + async with _legacy_operation_context(ctx, trace=False) as context: + return await operations.GatewayOperations(_capture_host_progress_callback(ctx)).execute( + ListResourcesRequest(params=params), context + ) + except Exception as exc: # noqa: BLE001 # preserve native listing fallback for ingress failures + verbose_logger.exception("Error in list_resources endpoint: %s", exc) + return ListResourcesResult(resources=[]) async def list_resource_templates( ctx: ServerRequestContext, params: PaginatedRequestParams ) -> ListResourceTemplatesResult: - """List all available resource templates.""" if _mcp_proxy_mode.get(): _reject_mcp_proxy_operation() - _ctx_reset_token: Final = active_mcp_request_ctx_var.set(ctx) - _session_reset_token: Final = active_mcp_session_var.set(ctx.session) - try: - ( - user_api_key_auth, - mcp_auth_header, - mcp_servers, - mcp_server_auth_headers, - oauth2_headers, - raw_headers, - _client_ip, - ) = await get_or_extract_auth_context() - verbose_logger.debug("MCP list_resource_templates - User API Key Auth from context: %s", user_api_key_auth) - verbose_logger.debug("MCP list_resource_templates - MCP servers from context: %s", mcp_servers) - verbose_logger.debug( - "MCP list_resource_templates - MCP server auth headers: %s", - list(mcp_server_auth_headers.keys()) if mcp_server_auth_headers else None, - ) - - resource_templates: Final = await _list_mcp_resource_templates( - user_api_key_auth=user_api_key_auth, - mcp_auth_header=mcp_auth_header, - mcp_servers=mcp_servers, - mcp_server_auth_headers=mcp_server_auth_headers, - oauth2_headers=oauth2_headers, - raw_headers=raw_headers, - ) - verbose_logger.info( - "MCP list_resource_templates - Successfully returned %s resource templates", len(resource_templates) - ) - return ListResourceTemplatesResult(resource_templates=resource_templates) - except Exception as e: - verbose_logger.exception("Error in list_resource_templates endpoint: %s", e) - return ListResourceTemplatesResult(resource_templates=[]) # mutable-ok: MCP result payload - finally: - active_mcp_session_var.reset(_session_reset_token) - active_mcp_request_ctx_var.reset(_ctx_reset_token) + async with _legacy_operation_context(ctx, trace=False) as context: + return await operations.GatewayOperations(_capture_host_progress_callback(ctx)).execute( + ListResourceTemplatesRequest(params=params), context + ) + except Exception as exc: # noqa: BLE001 # preserve native listing fallback for ingress failures + verbose_logger.exception("Error in list_resource_templates endpoint: %s", exc) + return ListResourceTemplatesResult(resource_templates=[]) async def read_resource(ctx: ServerRequestContext, params: ReadResourceRequestParams) -> ReadResourceResult: if _mcp_proxy_mode.get(): _reject_mcp_proxy_operation() - _ctx_reset_token: Final = active_mcp_request_ctx_var.set(ctx) - _session_reset_token: Final = active_mcp_session_var.set(ctx.session) - - try: - ( - user_api_key_auth, - mcp_auth_header, - mcp_servers, - mcp_server_auth_headers, - oauth2_headers, - raw_headers, - _client_ip, - ) = await get_or_extract_auth_context() - - read_resource_result: Final = await mcp_read_resource( - url=params.uri, - user_api_key_auth=user_api_key_auth, - mcp_auth_header=mcp_auth_header, - mcp_servers=mcp_servers, - mcp_server_auth_headers=mcp_server_auth_headers, - oauth2_headers=oauth2_headers, - raw_headers=raw_headers, + async with _legacy_operation_context(ctx, trace=False) as context: + return await operations.GatewayOperations(_capture_host_progress_callback(ctx)).execute( + ReadResourceRequest(params=params), context ) - return read_resource_result - finally: - active_mcp_session_var.reset(_session_reset_token) - active_mcp_request_ctx_var.reset(_ctx_reset_token) - server.add_request_handler("tools/list", PaginatedRequestParams, handle_list_tools) server.add_request_handler("tools/call", CallToolRequestParams, mcp_server_tool_call) server.add_request_handler("prompts/list", PaginatedRequestParams, list_prompts) @@ -1533,527 +954,24 @@ if MCP_AVAILABLE: ############ Helper Functions ########################## ######################################################## - async def _get_allowed_mcp_servers_from_mcp_server_names( - mcp_servers: Sequence[str] | None, - allowed_mcp_servers: list[MCPServer], - ) -> list[MCPServer]: - """ - Get the filtered MCP servers from the MCP server names. - - Fails closed when ``mcp_servers`` is explicitly provided (path- or - header-derived) but none of the names resolve to a server alias or - access group the caller can access. The previous behavior returned - the full ``allowed_mcp_servers`` set, which silently widened scope - when a client targeted ``/mcp//`` and made URL/header - namespacing appear to work when it did not. - """ - - filtered_server: Final[dict[str, MCPServer]] = {} - # Filter servers based on mcp_servers parameter if provided - if mcp_servers is not None: - for server_or_group in mcp_servers: - server_name_matched = False - - for server in allowed_mcp_servers: - if server and _server_answers_to(server, server_or_group): - filtered_server[server.server_id] = server - server_name_matched = True - break - - if not server_name_matched: - try: - access_group_server_ids = await MCPRequestHandler._get_mcp_servers_from_access_groups( - [server_or_group] - ) - # Only include servers that the user has access to - for server_id in access_group_server_ids: - for server in allowed_mcp_servers: - if server_id == server.server_id: - filtered_server[server.server_id] = server - except Exception as e: - verbose_logger.debug("Could not resolve '%s' as access group: %s", server_or_group, e) - - if filtered_server: - return list(filtered_server.values()) - - if mcp_servers is not None: - # Caller asked for a specific scope but nothing resolved. Fail - # closed so URL/header namespacing cannot silently fall back to - # the caller's full allowed-server set. - verbose_logger.debug( - "MCP scope filter resolved to no servers for requested names %s; returning empty list (fail-closed).", - mcp_servers, - ) - return [] - - return allowed_mcp_servers - - def _http_detail_message(detail: object) -> str: - return str(detail.get("error")) if isinstance(detail, dict) and detail.get("error") else str(detail) - - def _server_answers_to(server: MCPServer, name: str) -> bool: - requested: Final = name.lower() - return any(requested == known.lower() for known in iter_known_server_prefixes(server) if known) - - class _McpDeniedDetail(TypedDict): - error: ReadOnly[str] - - async def raise_denied_scoped_mcp_access( - requested_names: Sequence[str], - user_api_key_auth: UserAPIKeyAuth | None, - client_ip: str | None = None, - ) -> None: - """A scoped request (``/mcp/`` path or ``x-mcp-servers`` header) resolved to zero - allowed servers, so the denial must be loud: a silent 200 with no tools reads as a healthy - server with no tools. Unknown, unauthorized, and access-group names all share one generic - error so scoping cannot probe which servers exist; the agent variant fires only when the - same request resolves once the agent binding is stripped, proving the binding caused the veto.""" - agent_id: Final = user_api_key_auth.agent_id if user_api_key_auth else None - if user_api_key_auth is not None and agent_id: - resolved_without_agent: Final = await _get_allowed_mcp_servers( - user_api_key_auth=user_api_key_auth.model_copy(update=types.MappingProxyType({"agent_id": None})), - mcp_servers=requested_names, - client_ip=client_ip, - ) - - def _resolved_to_server(name: str) -> bool: - return any(_server_answers_to(server, name) for server in resolved_without_agent) - - vetoed_server: Final = next((name for name in requested_names if _resolved_to_server(name)), None) - if vetoed_server is not None: - agent_denial: Final[_McpDeniedDetail] = { - "error": ( - f"MCP server '{vetoed_server}' is not available to this key: the key is bound to " - f"agent '{agent_id}', whose MCP grants do not include this server. Add the server " - f"to the agent's object_permission.mcp_servers (edit the agent in the Admin UI or " - f"PATCH /v1/agents/{agent_id}), or use a key that is not bound to the agent." - ) - } - raise HTTPException(status_code=403, detail=agent_denial) - vetoed_group: Final = next( - ( - name - for name in requested_names - if not _resolved_to_server(name) - and any(name in (server.access_groups or ()) for server in resolved_without_agent) - ), - None, - ) - if vetoed_group is not None: - group_denial: Final[_McpDeniedDetail] = { - "error": ( - f"MCP access group '{vetoed_group}' is not available to this key: the key is bound to " - f"agent '{agent_id}', whose MCP grants do not include it. Add the group to the " - f"agent's object_permission.mcp_access_groups (edit the agent in the Admin UI or " - f"PATCH /v1/agents/{agent_id}), or use a key that is not bound to the agent." - ) - } - raise HTTPException(status_code=403, detail=group_denial) - generic_denial: Final[_McpDeniedDetail] = { - "error": f"The key is not allowed to access the requested MCP servers: {', '.join(requested_names)}" - } - raise HTTPException(status_code=403, detail=generic_denial) - - def _tool_name_matches(tool_name: str, filter_list: list[str], mcp_server: MCPServer) -> bool: - """ - Check if a tool name matches any name in the filter list. - - Reads the same owner the server-level permission checks use, so discovery hides - exactly what dispatch refuses. ``mcp_server`` is required: guessing the boundary - at the first separator mismatches every tool on a server whose prefix contains - the separator. - """ - bare_name: Final = strip_known_server_prefix(tool_name, mcp_server) - return match_known_tool_name(bare_name, mcp_server, filter_list) is not None - - def filter_tools_by_allowed_tools( - tools: list[MCPTool], - mcp_server: MCPServer, - ) -> list[MCPTool]: - """ - Filter tools by allowed/disallowed tools configuration. - - If allowed_tools is set, only tools in that list are returned. - If disallowed_tools is set, tools in that list are excluded. - Tool names are matched with and without server prefixes for flexibility. - - Args: - tools: List of tools to filter - mcp_server: Server configuration with allowed_tools/disallowed_tools - - Returns: - Filtered list of tools - """ - from litellm.proxy._experimental.mcp_server.utils import ( - server_applies_tool_allowlist, - ) - - tools_to_return = tools - - # Filter by allowed_tools (whitelist) - if server_applies_tool_allowlist(mcp_server): - if not mcp_server.allowed_tools: - return [] - tools_to_return = [ - tool for tool in tools if _tool_name_matches(tool.name, mcp_server.allowed_tools, mcp_server) - ] - - # Filter by disallowed_tools (blacklist) - if mcp_server.disallowed_tools: - tools_to_return = [ - tool - for tool in tools_to_return - if not _tool_name_matches(tool.name, mcp_server.disallowed_tools, mcp_server) - ] - - return tools_to_return - - def apply_tool_overrides( - tools: list[MCPTool], - mcp_server: MCPServer, - ) -> list[MCPTool]: - """Apply admin-configured display name/description overrides to tools. - - Overrides are keyed by the unprefixed tool name, same convention as - allowed_tools configuration. - """ - display_name_map: Final = mcp_server.tool_name_to_display_name or {} - description_map: Final = mcp_server.tool_name_to_description or {} - if not display_name_map and not description_map: - return tools - - for tool in tools: - unprefixed = strip_known_server_prefix(tool.name, mcp_server) - lookup_key = unprefixed or tool.name - if lookup_key in display_name_map: - tool.name = display_name_map[lookup_key] - if lookup_key in description_map: - tool.description = description_map[lookup_key] - return tools - - def _get_client_ip_from_context() -> str | None: - """ - Extract client_ip from auth context. - Returns None if context not set (caller should handle this as "no IP filtering"). - """ - try: - auth_user: Final = auth_context_var.get() - if auth_user and isinstance(auth_user, MCPAuthenticatedUser): - return auth_user.client_ip - except Exception: - pass - return None - - async def _get_allowed_mcp_servers( - user_api_key_auth: UserAPIKeyAuth | None, - mcp_servers: Sequence[str] | None, - client_ip: str | None = None, - ) -> list[MCPServer]: - """Return allowed MCP servers for a request after applying filters. - - Args: - user_api_key_auth: The authenticated user's API key info. - mcp_servers: Optional list of server names to filter to. - client_ip: Client IP for IP-based access control. If None, falls back to - auth context. Pass explicitly from request handlers for safety. - Note: If client_ip is None and auth context is not set, IP filtering is skipped. - This is intentional for internal callers but may indicate a bug if called - from a request handler without proper context setup. - """ - # Use explicit client_ip if provided, otherwise try auth context - if client_ip is None: - client_ip = _get_client_ip_from_context() - if client_ip is None: - verbose_logger.debug( - "MCP _get_allowed_mcp_servers called without client_ip and no auth context. " - "IP filtering will be skipped. This is expected for internal calls." - ) - - allowed_mcp_server_ids = await global_mcp_server_manager.get_allowed_mcp_servers(user_api_key_auth) - ( - allowed_mcp_server_ids, - _ip_blocked, - ) = global_mcp_server_manager.filter_server_ids_by_ip_with_info(allowed_mcp_server_ids, client_ip) - verbose_logger.debug( - "MCP IP filter: client_ip=%s, allowed_server_ids=%s", - client_ip, - allowed_mcp_server_ids, - ) - if _ip_blocked > 0: - verbose_logger.debug( - "MCP IP filtering: %d server(s) are not accessible from client IP %s " - "because they are restricted to internal networks. " - "No tools from those servers will be returned. " - "To expose a server externally, set 'available_on_public_internet: true' " - "in its configuration.", - _ip_blocked, - client_ip, - ) - allowed_mcp_servers: list[MCPServer] = [] - for allowed_mcp_server_id in allowed_mcp_server_ids: - mcp_server = global_mcp_server_manager.get_mcp_server_by_id(allowed_mcp_server_id) - if mcp_server is not None: - # Apply the request-time oauth2_flow backstop for legacy null rows. - mcp_server = MCPServerManager.resolve_oauth2_flow_for_request(mcp_server) - allowed_mcp_servers.append(mcp_server) - - if mcp_servers is not None: - allowed_mcp_servers = await _get_allowed_mcp_servers_from_mcp_server_names( - mcp_servers=mcp_servers, - allowed_mcp_servers=allowed_mcp_servers, - ) - - return allowed_mcp_servers - - def _client_has_per_server_auth_header( - server: MCPServer, - mcp_server_auth_headers: dict[str, dict[str, str]] | None, - ) -> bool: - """True if the request carries a per-server ``x-mcp-{alias}-authorization`` - header for this server. This is the multi-server binding: it names one - upstream, so it is unambiguously the caller's upstream token regardless of - auth mode (never the LiteLLM admission credential). - - Resolves through the same ``lookup_mcp_server_auth_in_headers`` egress uses, so - the connect gate and egress agree on which per-server header names match: a - dashboard client sends ``x-mcp-{sanitize_mcp_alias_for_header(alias)}-authorization``, - and matching only the raw alias here would 401 a token egress would forward. - """ - if not mcp_server_auth_headers: - return False - from litellm.proxy._experimental.mcp_server.utils import ( - lookup_mcp_server_auth_in_headers, - ) - - server_headers: Final = lookup_mcp_server_auth_in_headers( - mcp_server_auth_headers, - alias=server.alias, - server_name=server.server_name, - access_groups=server.access_groups, - ) - if isinstance(server_headers, str): - return bool(server_headers.strip()) - if isinstance(server_headers, dict): - return any(isinstance(hk, str) and hk.lower() == "authorization" for hk in server_headers) - return False - - def _client_has_passthrough_authorization( - server: MCPServer, - oauth2_headers: dict[str, str] | None, - mcp_server_auth_headers: dict[str, dict[str, str]] | None, - ) -> 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: - if k.lower() == "authorization": - return True - return _client_has_per_server_auth_header(server, mcp_server_auth_headers) - - async def _get_user_oauth_extra_headers_from_db( - server: MCPServer, - user_api_key_auth: UserAPIKeyAuth | None, - prefetched_creds: 'Mapping[str, "OAuthCredentialPayload"] | None' = None, - ) -> dict[str, str] | None: - """Stored OAuth2 token for (user, server) as an ``Authorization: Bearer`` header, or None. - - Thin wrapper over ``resolve_user_oauth_access_token`` (Redis cache, else DB + refresh); - ``prefetched_creds`` skips the per-server Redis/DB lookups for the batch path. - """ - if server.auth_type != MCPAuth.oauth2 or user_api_key_auth is None: - return None - from litellm.proxy._experimental.mcp_server.db import ( # noqa: PLC0415 - resolve_user_oauth_access_token, - ) - - token: Final = await resolve_user_oauth_access_token( - getattr(user_api_key_auth, "user_id", None), server, prefetched_creds - ) - return {"Authorization": f"Bearer {token}"} if token else None - - async def _prefetch_oauth_creds_for_user( - user_api_key_auth: UserAPIKeyAuth | None, - ) -> dict[str, "OAuthCredentialPayload"]: - """Fetch all OAuth2 credentials for the user in one DB query. - - Returns a dict keyed by server_id to avoid N+1 queries in asyncio.gather loops. - """ - user_id: Final[str | None] = getattr(user_api_key_auth, "user_id", None) if user_api_key_auth else None - if not user_id: - return {} - try: - from litellm.proxy._experimental.mcp_server.db import ( # noqa: PLC0415 - list_user_oauth_credentials, - ) - from litellm.proxy.utils import get_prisma_client_or_throw # noqa: PLC0415 - - prisma_client: Final = get_prisma_client_or_throw( - "Database not connected. Connect a database to use OAuth2 MCP tools." - ) - creds: Final = await list_user_oauth_credentials(prisma_client, user_id) - return {c["server_id"]: c for c in creds if "server_id" in c} - except Exception as e: - verbose_logger.warning("_prefetch_oauth_creds_for_user: failed to prefetch for user=%s: %s", user_id, e) - return {} - - def _prepare_mcp_server_headers( - server: MCPServer, - mcp_server_auth_headers: dict[str, dict[str, str]] | None, - mcp_auth_header: str | None, - oauth2_headers: dict[str, str] | None, - raw_headers: dict[str, str] | None, - user_api_key_auth: UserAPIKeyAuth | None = None, - scope_servers: list[MCPServer] | None = None, - ) -> tuple[dict[str, str] | str | None, dict[str, str] | None]: - """Build auth and extra headers for a server. - - ``scope_servers`` is the full server list a fan-out handler iterates. Passing it lets the - client-forwarded token modes withhold the caller's request-wide ``Authorization`` when - another server in the scope would also receive it (``_caller_authorization_fans_out``); - explicitly-addressed operations leave it None. Per-server ``x-mcp-{alias}-authorization`` - headers are unaffected — they bind one token to one server and are the multi-server shape. - """ - server_auth_header: dict[str, str] | str | None = None - if mcp_server_auth_headers: - from litellm.proxy._experimental.mcp_server.utils import ( - lookup_mcp_server_auth_in_headers, - ) - - server_auth_header = lookup_mcp_server_auth_in_headers( - mcp_server_auth_headers, - alias=server.alias, - server_name=server.server_name, - access_groups=server.access_groups, - ) - - extra_headers: dict[str, str] | None = None - is_client_forwarded_mode: Final = server.is_client_forwarded_token - # In a multi-server listing scope the request-wide Authorization can only carry one token, - # so it is withheld from a client-forwarded server when another server in scope also consumes - # it (RFC 9700 cross-resource replay); such scopes must bind per-server via - # x-mcp-{alias}-authorization. The decision is computed once so BOTH the forwarding branch and - # the extra_headers copy loop below honor it — otherwise a server that lists Authorization in - # extra_headers would re-copy the withheld bearer from raw_headers and replay it anyway. - withhold_forwarded_authorization: Final = is_client_forwarded_mode and _caller_authorization_fans_out( - server, scope_servers - ) - if server.auth_type == MCPAuth.oauth2: - # For OAuth2 M2M servers, upstream Authorization must come from - # client_credentials token fetch, never from caller headers. - if server.has_client_credentials: - extra_headers = None - else: - # Copy to avoid mutating the original dict (important for parallel fetching) - extra_headers = oauth2_headers.copy() if oauth2_headers else None - # Migrated authorization_code: the v2 resolver injects the stored per-user - # token, so drop the caller-forwarded Authorization (apply-if-absent would - # otherwise let it shadow the resolved token). Delegate keeps it. Centralized - # via _should_strip_caller_authorization to match _call_regular_mcp_tool. - if extra_headers and _should_strip_caller_authorization( - mcp_server=server, - raw_headers=raw_headers, - user_api_key_auth=user_api_key_auth, - ): - extra_headers = without_header(extra_headers, DEFAULT_CREDENTIAL_HEADER) - elif is_client_forwarded_mode: - if not withhold_forwarded_authorization: - extra_headers = _client_forwarded_authorization_headers( - mcp_server=server, - oauth2_headers=oauth2_headers, - raw_headers=raw_headers, - user_api_key_auth=user_api_key_auth, - ) - - if server.extra_headers and raw_headers: - if extra_headers is None: - extra_headers = {} - - normalized_raw_headers: Final = {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: Final = _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 header.lower() == "authorization" and ( - strip_caller_authorization or withhold_forwarded_authorization - ): - continue - header_value = normalized_raw_headers.get(header.lower()) - if header_value is None: - continue - extra_headers[header] = header_value - - # Reset to None if no headers were actually added - if extra_headers is not None and len(extra_headers) == 0: - extra_headers = None - - if server_auth_header is None: - server_auth_header = mcp_auth_header - - return server_auth_header, extra_headers - - def _merge_gateway_initialize_instructions( - allowed_mcp_servers: list[MCPServer], - ) -> str | None: - """YAML/DB override, else upstream text (prefetch on init, or list_tools / health_check / call_tool cache).""" - if not allowed_mcp_servers: - return None - - texts: Final[list[tuple[str, str]]] = [] - for server in allowed_mcp_servers: - label = server.alias or server.server_name or server.name or server.server_id or "mcp" - if server.instructions and server.instructions.strip(): - texts.append((label, server.instructions.strip())) - continue - if server.spec_path: - continue - cached = global_mcp_server_manager._upstream_initialize_instructions_by_server_id.get(server.server_id) - if cached and cached.strip(): - texts.append((label, cached.strip())) - - if not texts: - return None - if len(texts) == 1: - return texts[0][1] - return "\n\n---\n\n".join(f"[{lbl}]\n{txt}" for lbl, txt in texts) - - async def _raise_if_initialize_grants_no_mcp_servers( - allowed: Sequence[MCPServer], - user_api_key_auth: UserAPIKeyAuth | None, - mcp_servers: Sequence[str] | None, - client_ip: str | None, - ) -> None: - if allowed or user_api_key_auth is None or not user_api_key_auth.api_key: - return - if mcp_servers: - await raise_denied_scoped_mcp_access( - requested_names=mcp_servers, - user_api_key_auth=user_api_key_auth, - client_ip=client_ip, - ) - no_servers_denial: Final[_McpDeniedDetail] = { - "error": ( - "The key has no MCP servers granted, or none of its granted servers is loaded and allowed for " - "this client IP. Grant servers or access groups to the key, its team, or its organization " - "(object_permission.mcp_servers), check the server's allowed IPs, and reconnect." - ) - } - raise HTTPException(status_code=403, detail=no_servers_denial) + from litellm.proxy._experimental.mcp_server.operations import ( + _client_has_passthrough_authorization, + _client_has_per_server_auth_header, + _get_allowed_mcp_servers, + _get_allowed_mcp_servers_from_mcp_server_names, + _get_user_oauth_extra_headers_from_db, + _http_detail_message, + _McpDeniedDetail, + _merge_gateway_initialize_instructions, + _prefetch_oauth_creds_for_user, + _prepare_mcp_server_headers, + _raise_if_initialize_grants_no_mcp_servers, + _server_answers_to, + _tool_name_matches, + apply_tool_overrides, + filter_tools_by_allowed_tools, + raise_denied_scoped_mcp_access, + ) @contextlib.asynccontextmanager async def _gateway_initialize_instructions_request_scope( @@ -2063,26 +981,28 @@ if MCP_AVAILABLE: scoped_server_endpoint: bool = False, is_initialize: bool = False, ) -> AsyncIterator[None]: - allowed: Final = await _get_allowed_mcp_servers( + allowed: Final = await operations._get_allowed_mcp_servers( user_api_key_auth=user_api_key_auth, mcp_servers=mcp_servers, client_ip=client_ip, ) if is_initialize: - await _raise_if_initialize_grants_no_mcp_servers(allowed, user_api_key_auth, mcp_servers, client_ip) + await operations._raise_if_initialize_grants_no_mcp_servers( + allowed, user_api_key_auth, mcp_servers, client_ip + ) if allowed: # return_exceptions=True: a per-server probe failure (incl. CancelledError # bubbled from anyio task group teardown on connection refused) must not # cancel sibling probes or 500 the gateway initialize request. await asyncio.gather( *[ - global_mcp_server_manager._ensure_upstream_initialize_instructions_cached(s) + operations.global_mcp_server_manager._ensure_upstream_initialize_instructions_cached(s) for s in allowed if s is not None ], return_exceptions=True, ) - merged: Final = _merge_gateway_initialize_instructions(allowed_mcp_servers=allowed) + merged: Final = operations._merge_gateway_initialize_instructions(allowed_mcp_servers=allowed) scoped_server_name = None if scoped_server_endpoint and len(allowed) == 1: scoped_server: Final = allowed[0] @@ -2097,1599 +1017,34 @@ if MCP_AVAILABLE: _mcp_gateway_initialize_instructions.reset(instructions_token) _mcp_gateway_server_name.reset(server_name_token) - def _aggregate_server_key(server: MCPServer) -> str: - """The client-visible key for a server in listing outcomes and spend metadata: the same - display prefix (alias, or the short prefix when that mode is enabled) the caller already - sees on the tool names. Canonical internal server names never key a caller-readable - surface; when the display naming deliberately hides them, the outcome keys must too.""" - return get_server_prefix(server) or "unknown" - - async def _get_tools_from_mcp_servers( - user_api_key_auth: UserAPIKeyAuth | None, - mcp_auth_header: str | None, - mcp_servers: list[str] | None, - mcp_server_auth_headers: dict[str, dict[str, str]] | None = None, - oauth2_headers: dict[str, str] | None = None, - raw_headers: dict[str, str] | None = None, - log_list_tools_to_spendlogs: bool = False, - list_tools_log_source: str | None = None, - litellm_trace_id: str | None = None, - request_tags: list[str] | None = None, - client_ip: str | None = None, - mcp_proxy_mode: bool = False, - ) -> AggregateToolListing: - """ - Helper method to fetch tools from MCP servers based on server filtering criteria. - - Args: - user_api_key_auth: User authentication info for access control - mcp_auth_header: Optional auth header for MCP server (deprecated) - mcp_servers: Optional list of server names/aliases to filter by - mcp_server_auth_headers: Optional dict of server-specific auth headers - oauth2_headers: Optional dict of oauth2 headers - - Returns: - AggregateToolListing: Combined tools from filtered servers plus each server's - classified listing outcome - """ - if not MCP_AVAILABLE: - return AggregateToolListing(tools=[], outcomes={}) - - list_tools_start_time: Final = datetime.now() - litellm_logging_obj: LiteLLMLoggingObj | None = None - list_tools_request_data: dict[str, object] = {} - - if log_list_tools_to_spendlogs: - # This is intentionally minimal: only async_success_handler / post_call_failure_hook - rules_obj: Final = Rules() - list_tools_call_id: Final = str(uuid.uuid4()) - # Derive trace_id from raw_headers when not explicitly passed (same as A2A / MCP call_tool) - effective_litellm_trace_id: Final = litellm_trace_id or get_chain_id_from_headers(raw_headers) - spend_logs_metadata: Final[dict[str, object]] = { - "mcp_operation": "list_tools", - } - if isinstance(list_tools_log_source, str): - spend_logs_metadata["source"] = list_tools_log_source - if isinstance(mcp_servers, list): - spend_logs_metadata["requested_mcp_servers"] = mcp_servers - - list_tools_request_data = { - "model": "MCP: list_tools", - "call_type": CallTypes.list_mcp_tools.value, - "litellm_call_id": list_tools_call_id, - "litellm_trace_id": effective_litellm_trace_id, - "metadata": { - "spend_logs_metadata": spend_logs_metadata, - "headers": logging_safe_mcp_headers(raw_headers), - **({"tags": request_tags} if request_tags else {}), - }, - # Provide a small input payload for standard logging - "input": [ - { - "role": "system", - "content": { - "mcp_operation": "list_tools", - "requested_mcp_servers": mcp_servers, - }, - } - ], - } - - # Attach user identifiers using the standard helper - if user_api_key_auth is not None: - LiteLLMProxyRequestSetup.add_user_api_key_auth_to_request_metadata( - data=list_tools_request_data, - user_api_key_dict=user_api_key_auth, - _metadata_variable_name="metadata", - ) - - user_identifier: Final = getattr(user_api_key_auth, "end_user_id", None) or getattr( - user_api_key_auth, "user_id", None - ) - if user_identifier: - list_tools_request_data["user"] = user_identifier - - try: - litellm_logging_obj, _ = function_setup( - original_function="list_mcp_tools", - rules_obj=rules_obj, - start_time=list_tools_start_time, - **list_tools_request_data, - ) - if litellm_logging_obj: - litellm_logging_obj.call_type = CallTypes.list_mcp_tools.value - litellm_logging_obj.model = "MCP: list_tools" - except Exception as logging_error: - verbose_logger.debug("Failed to initialize logging for MCP list_tools: %s", logging_error) - litellm_logging_obj = None - - try: - allowed_mcp_servers: Final = await _get_allowed_mcp_servers( - user_api_key_auth=user_api_key_auth, - mcp_servers=mcp_servers, - client_ip=client_ip, - ) - if mcp_servers and not allowed_mcp_servers: - await raise_denied_scoped_mcp_access( - requested_names=mcp_servers, - user_api_key_auth=user_api_key_auth, - client_ip=client_ip, - ) - - # Pre-fetch OAuth credentials only when at least one server uses OAuth2, - # to avoid an unnecessary DB round-trip on requests with no OAuth2 MCP servers. - _has_oauth2_server = any(getattr(s, "auth_type", None) == MCPAuth.oauth2 for s in allowed_mcp_servers) - _prefetched_oauth_creds: Final = ( - await _prefetch_oauth_creds_for_user(user_api_key_auth) if _has_oauth2_server else {} - ) - - async def _fetch_and_filter_server_tools( - server: MCPServer, - ) -> "tuple[list[MCPTool], ServerOutcome]": - """Fetch and filter tools from a single server, classifying any failure into that - server's outcome so the aggregate can keep serving the healthy subset without a - broken server masquerading as an empty one.""" - if server is None: - return [], ServerListOk(tool_count=0) - - server_auth_header, extra_headers = _prepare_mcp_server_headers( - server=server, - mcp_server_auth_headers=mcp_server_auth_headers, - mcp_auth_header=mcp_auth_header, - oauth2_headers=oauth2_headers, - raw_headers=raw_headers, - user_api_key_auth=user_api_key_auth, - scope_servers=allowed_mcp_servers, - ) - - # Prefer server-stored per-user OAuth when configured, so a stale - # Authorization header from the MCP client cannot override Redis/DB - # (same issue as call_tool in mcp_server_manager: VS Code caches tokens). - from litellm.proxy._experimental.mcp_server.outbound_credentials.adapter import ( # noqa: PLC0415 - to_server_spec, - ) - - # A server migrated to the v2 resolver gets its token from the resolver at connect - # time; building it here would double-resolve and be shadowed by the v2 graft. The - # preemptive 401 already challenged a missing token, so one exists for the connect. - migrated_to_v2: Final = to_server_spec(server) is not None - if ( - not migrated_to_v2 - and server.auth_type == MCPAuth.oauth2 - and getattr(server, "needs_user_oauth_token", False) - and user_api_key_auth is not None - ): - db_headers: Final = await _get_user_oauth_extra_headers_from_db( - server, - user_api_key_auth, - prefetched_creds=_prefetched_oauth_creds, - ) - if db_headers: - extra_headers = db_headers - - # If still no OAuth2 token, fall back to pre-fetched creds (non-stale-client path) - elif not migrated_to_v2 and extra_headers is None and server.auth_type == MCPAuth.oauth2: - extra_headers = await _get_user_oauth_extra_headers_from_db( - server, - user_api_key_auth, - prefetched_creds=_prefetched_oauth_creds, - ) - - if server.is_byok and server.auth_type != MCPAuth.oauth2 and server_auth_header is None: - server_auth_header = await _get_byok_credential(server, user_api_key_auth) - - try: - tools: Final = await global_mcp_server_manager._get_tools_from_server( - server=server, - mcp_auth_header=server_auth_header, - extra_headers=extra_headers, - add_prefix=True, # Always add server prefix - raw_headers=raw_headers, - user_api_key_auth=user_api_key_auth, - oauth2_headers=oauth2_headers, - ) - filtered_tools = filter_tools_by_allowed_tools(tools, server) - - filtered_tools = await filter_tools_by_key_team_permissions( - tools=filtered_tools, - server_id=server.server_id, - user_api_key_auth=user_api_key_auth, - ) - - if mcp_proxy_mode: - from litellm.proxy._experimental.mcp_server.tool_search import with_mcp_proxy_identity - - filtered_tools = [ # mutable-ok: MCP tool pipeline - with_mcp_proxy_identity(tool, server.server_id) for tool in filtered_tools - ] - else: - filtered_tools = apply_tool_overrides(filtered_tools, server) - - verbose_logger.debug( - "Successfully fetched %s tools from server %s, %s after filtering", - len(tools), - server.name, - len(filtered_tools), - ) - return filtered_tools, ServerListOk(tool_count=len(filtered_tools)) - except MCPUpstreamAuthError as e: - # Absorb so one unauthenticated server does not empty every other server's - # tools. Surfacing the upstream 401 to the client as a re-auth challenge is - # intentionally not done here: raising from this list handler cannot produce a - # 401 + WWW-Authenticate (the MCP session manager serializes it as a JSON-RPC - # error). Single-server routes surface it via the request-scope preemptive - # check in _raise_preemptive_401_for_unauthenticated_servers instead. - verbose_logger.debug("MCP list_tools: omitting %s; it needs upstream auth", server.name) - return [], classify_list_exception(e) - except Exception as e: - verbose_logger.exception("Error getting tools from server %s: %s", server.name, e) - return [], classify_list_exception(e) - - # Fetch tools from all servers in parallel - tasks: Final = [_fetch_and_filter_server_tools(server) for server in allowed_mcp_servers] - results: Final = await asyncio.gather(*tasks) - - # Flatten results into single list - all_tools: Final[list[MCPTool]] = [tool for tools, _ in results for tool in tools] - server_outcomes: Final[dict[str, ServerOutcome]] = { - _aggregate_server_key(server): outcome - for server, (_, outcome) in zip(allowed_mcp_servers, results) - if server is not None - } - - # If logging is enabled, enrich spend_logs_metadata with counts - if litellm_logging_obj: - per_server_tool_counts: Final[dict[str, int]] = { - _aggregate_server_key(server): len(server_tools) - for server, (server_tools, _) in zip(allowed_mcp_servers, results) - if server is not None - } - - metadata_dict: Final = litellm_logging_obj.model_call_details.get("metadata") - if isinstance(metadata_dict, dict): - spend_meta = metadata_dict.get("spend_logs_metadata") - if not isinstance(spend_meta, dict): - spend_meta = {} - metadata_dict["spend_logs_metadata"] = spend_meta - spend_meta["allowed_server_count"] = len(allowed_mcp_servers) - spend_meta["tool_count_total"] = len(all_tools) - spend_meta["per_server_tool_counts"] = per_server_tool_counts - spend_meta["per_server_list_outcomes"] = { - key: outcome_wire_value(outcome) for key, outcome in server_outcomes.items() - } - - end_time: Final = datetime.now() - try: - await litellm_logging_obj.async_success_handler( - result=[ - tool.model_dump(mode="json") if isinstance(tool, MCPTool) else tool for tool in all_tools - ], - start_time=list_tools_start_time, - end_time=end_time, - ) - except Exception as log_exc: - # list_tools responses must not be dropped due to non-blocking - # observability/serialization failures. - verbose_logger.warning( - "MCP list_tools success logging failed (continuing): %s", - log_exc, - ) - - verbose_logger.info("Successfully fetched %s tools total from all MCP servers", len(all_tools)) - - return AggregateToolListing(tools=all_tools, outcomes=server_outcomes) - except Exception as e: - # Only fire failure hook if logging was requested for this list-tools execution - if log_list_tools_to_spendlogs and user_api_key_auth is not None: - try: - from litellm.proxy.proxy_server import proxy_logging_obj - - if proxy_logging_obj: - traceback_str: Final = traceback.format_exc(limit=MAXIMUM_TRACEBACK_LINES_TO_LOG) - await proxy_logging_obj.post_call_failure_hook( - request_data=list_tools_request_data or {}, - original_exception=e, - user_api_key_dict=user_api_key_auth, - route="/mcp/list_tools", - traceback_str=traceback_str, - ) - except Exception: - verbose_logger.debug("Failed to log MCP list_tools failure via post_call_failure_hook") - raise - - async def _get_prompts_from_mcp_servers( - user_api_key_auth: UserAPIKeyAuth | None, - mcp_auth_header: str | None, - mcp_servers: list[str] | None, - mcp_server_auth_headers: dict[str, dict[str, str]] | None = None, - oauth2_headers: dict[str, str] | None = None, - raw_headers: dict[str, str] | None = None, - ) -> list[Prompt]: - """ - Helper method to fetch prompt from MCP servers based on server filtering criteria. - - Args: - user_api_key_auth: User authentication info for access control - mcp_auth_header: Optional auth header for MCP server (deprecated) - mcp_servers: Optional list of server names/aliases to filter by - mcp_server_auth_headers: Optional dict of server-specific auth headers - oauth2_headers: Optional dict of oauth2 headers - - Returns: - List[Prompt]: Combined list of prompts from filtered servers - """ - if not MCP_AVAILABLE: - return [] - - allowed_mcp_servers: Final = await _get_allowed_mcp_servers( - user_api_key_auth=user_api_key_auth, - mcp_servers=mcp_servers, - ) - - # Get prompts from each allowed server - all_prompts: Final = [] - for server in allowed_mcp_servers: - if server is None: - continue - - server_auth_header, extra_headers = _prepare_mcp_server_headers( - server=server, - mcp_server_auth_headers=mcp_server_auth_headers, - mcp_auth_header=mcp_auth_header, - oauth2_headers=oauth2_headers, - raw_headers=raw_headers, - user_api_key_auth=user_api_key_auth, - scope_servers=allowed_mcp_servers, - ) - - try: - prompts = await global_mcp_server_manager.get_prompts_from_server( - server=server, - user_api_key_auth=user_api_key_auth, - mcp_auth_header=server_auth_header, - extra_headers=extra_headers, - add_prefix=True, # Always add server prefix - raw_headers=raw_headers, - ) - - all_prompts.extend(prompts) - - verbose_logger.debug("Successfully fetched %s prompts from server %s", len(prompts), server.name) - except Exception as e: - verbose_logger.exception("Error getting prompts from server %s: %s", server.name, e) - # Continue with other servers instead of failing completely - - verbose_logger.info("Successfully fetched %s prompts total from all MCP servers", len(all_prompts)) - - return all_prompts - - async def _get_resources_from_mcp_servers( - user_api_key_auth: UserAPIKeyAuth | None, - mcp_auth_header: str | None, - mcp_servers: list[str] | None, - mcp_server_auth_headers: dict[str, dict[str, str]] | None = None, - oauth2_headers: dict[str, str] | None = None, - raw_headers: dict[str, str] | None = None, - ) -> list[Resource]: - """Fetch resources from allowed MCP servers.""" - - if not MCP_AVAILABLE: - return [] - - allowed_mcp_servers: Final = await _get_allowed_mcp_servers( - user_api_key_auth=user_api_key_auth, - mcp_servers=mcp_servers, - ) - - all_resources: Final[list[Resource]] = [] - for server in allowed_mcp_servers: - if server is None: - continue - - server_auth_header, extra_headers = _prepare_mcp_server_headers( - server=server, - mcp_server_auth_headers=mcp_server_auth_headers, - mcp_auth_header=mcp_auth_header, - oauth2_headers=oauth2_headers, - raw_headers=raw_headers, - user_api_key_auth=user_api_key_auth, - scope_servers=allowed_mcp_servers, - ) - - try: - resources = await global_mcp_server_manager.get_resources_from_server( - server=server, - user_api_key_auth=user_api_key_auth, - mcp_auth_header=server_auth_header, - extra_headers=extra_headers, - add_prefix=True, # Always add server prefix - raw_headers=raw_headers, - ) - all_resources.extend(resources) - - verbose_logger.debug("Successfully fetched %s resources from server %s", len(resources), server.name) - except Exception as e: - verbose_logger.exception("Error getting resources from server %s: %s", server.name, e) - - verbose_logger.info("Successfully fetched %s resources total from all MCP servers", len(all_resources)) - - return all_resources - - async def _get_resource_templates_from_mcp_servers( - user_api_key_auth: UserAPIKeyAuth | None, - mcp_auth_header: str | None, - mcp_servers: list[str] | None, - mcp_server_auth_headers: dict[str, dict[str, str]] | None = None, - oauth2_headers: dict[str, str] | None = None, - raw_headers: dict[str, str] | None = None, - ) -> list[ResourceTemplate]: - """Fetch resource templates from allowed MCP servers.""" - - if not MCP_AVAILABLE: - return [] - - allowed_mcp_servers: Final = await _get_allowed_mcp_servers( - user_api_key_auth=user_api_key_auth, - mcp_servers=mcp_servers, - ) - - all_resource_templates: Final[list[ResourceTemplate]] = [] - for server in allowed_mcp_servers: - if server is None: - continue - - server_auth_header, extra_headers = _prepare_mcp_server_headers( - server=server, - mcp_server_auth_headers=mcp_server_auth_headers, - mcp_auth_header=mcp_auth_header, - oauth2_headers=oauth2_headers, - raw_headers=raw_headers, - user_api_key_auth=user_api_key_auth, - scope_servers=allowed_mcp_servers, - ) - - try: - resource_templates = await global_mcp_server_manager.get_resource_templates_from_server( - server=server, - user_api_key_auth=user_api_key_auth, - mcp_auth_header=server_auth_header, - extra_headers=extra_headers, - add_prefix=True, # Always add server prefix - raw_headers=raw_headers, - ) - all_resource_templates.extend(resource_templates) - verbose_logger.debug( - "Successfully fetched %s resource templates from server %s", - len(resource_templates), - server.name, - ) - except Exception as e: - verbose_logger.exception( - "Error getting resource templates from server %s: %s", - server.name, - str(e), - ) - - verbose_logger.info( - "Successfully fetched %s resource templates total from all MCP servers", - len(all_resource_templates), - ) - - return all_resource_templates - - async def filter_tools_by_key_team_permissions( - tools: list[MCPTool], - server_id: str, - user_api_key_auth: UserAPIKeyAuth | None, - ) -> list[MCPTool]: - """ - Filter tools based on key/team mcp_tool_permissions. - - Note: Tool names in the DB are stored without server prefixes, - but tool names from MCP servers are prefixed. We need to strip - the prefix before comparing. - """ - # Filter by key/team tool-level permissions - allowed_tool_names: Final = await MCPRequestHandler.get_allowed_tools_for_server( - server_id=server_id, - user_api_key_auth=user_api_key_auth, - ) - - # Tools arrive prefixed with the server's own prefix; strip exactly that - # prefix (resolved from the server) rather than the first separator, so a - # prefix containing the separator still reduces to the stored bare name. - server: Final = global_mcp_server_manager.get_mcp_server_by_id(server_id) - return [ - t - for t in tools - if MCPRequestHandler.tool_is_granted(strip_known_server_prefix(t.name, server), allowed_tool_names) - ] - - async def _list_mcp_tools( - user_api_key_auth: UserAPIKeyAuth | None = None, - mcp_auth_header: str | None = None, - mcp_servers: list[str] | None = None, - mcp_server_auth_headers: dict[str, dict[str, str]] | None = None, - oauth2_headers: dict[str, str] | None = None, - raw_headers: dict[str, str] | None = None, - log_list_tools_to_spendlogs: bool = False, - list_tools_log_source: str | None = None, - client_ip: str | None = None, - mcp_proxy_mode: bool = False, - ) -> AggregateToolListing: - """ - List all available MCP tools. - - Args: - user_api_key_auth: User authentication info for access control - mcp_auth_header: Optional auth header for MCP server (deprecated) - mcp_servers: Optional list of server names/aliases to filter by - mcp_server_auth_headers: Optional dict of server-specific auth headers {server_alias: auth_value} - client_ip: Client IP for IP-based server access control - - Returns: - AggregateToolListing: Combined tools from all accessible servers plus each server's - classified listing outcome - """ - if not MCP_AVAILABLE: - return AggregateToolListing(tools=[], outcomes={}) - - try: - listing: Final = await _get_tools_from_mcp_servers( - user_api_key_auth=user_api_key_auth, - mcp_auth_header=mcp_auth_header, - mcp_servers=mcp_servers, - mcp_server_auth_headers=mcp_server_auth_headers, - oauth2_headers=oauth2_headers, - raw_headers=raw_headers, - log_list_tools_to_spendlogs=log_list_tools_to_spendlogs, - list_tools_log_source=list_tools_log_source, - client_ip=client_ip, - mcp_proxy_mode=mcp_proxy_mode, - ) - verbose_logger.debug("Successfully fetched %s tools from managed MCP servers", len(listing.tools)) - return listing - except HTTPException: - raise - except Exception as e: - verbose_logger.exception("Error getting tools from managed MCP servers: %s", e) - # Continue with an empty listing instead of failing completely - return AggregateToolListing(tools=[], outcomes={}) - - async def _list_mcp_prompts( - user_api_key_auth: UserAPIKeyAuth | None = None, - mcp_auth_header: str | None = None, - mcp_servers: list[str] | None = None, - mcp_server_auth_headers: dict[str, dict[str, str]] | None = None, - oauth2_headers: dict[str, str] | None = None, - raw_headers: dict[str, str] | None = None, - ) -> list[Prompt]: - """ - List all available MCP prompts. - - Args: - user_api_key_auth: User authentication info for access control - mcp_auth_header: Optional auth header for MCP server (deprecated) - mcp_servers: Optional list of server names/aliases to filter by - mcp_server_auth_headers: Optional dict of server-specific auth headers {server_alias: auth_value} - - Returns: - List[Prompt]: Combined list of tools from all accessible servers - """ - if not MCP_AVAILABLE: - return [] - # Get tools from managed MCP servers with error handling - managed_prompts = [] - try: - managed_prompts = await _get_prompts_from_mcp_servers( - user_api_key_auth=user_api_key_auth, - mcp_auth_header=mcp_auth_header, - mcp_servers=mcp_servers, - mcp_server_auth_headers=mcp_server_auth_headers, - oauth2_headers=oauth2_headers, - raw_headers=raw_headers, - ) - verbose_logger.debug("Successfully fetched %s prompts from managed MCP servers", len(managed_prompts)) - except Exception as e: - verbose_logger.exception("Error getting tools from managed MCP servers: %s", e) - # Continue with empty managed tools list instead of failing completely - - return managed_prompts - - async def _list_mcp_resources( - user_api_key_auth: UserAPIKeyAuth | None = None, - mcp_auth_header: str | None = None, - mcp_servers: list[str] | None = None, - mcp_server_auth_headers: dict[str, dict[str, str]] | None = None, - oauth2_headers: dict[str, str] | None = None, - raw_headers: dict[str, str] | None = None, - ) -> list[Resource]: - """List all available MCP resources.""" - - if not MCP_AVAILABLE: - return [] - - managed_resources: list[Resource] = [] - try: - managed_resources = await _get_resources_from_mcp_servers( - user_api_key_auth=user_api_key_auth, - mcp_auth_header=mcp_auth_header, - mcp_servers=mcp_servers, - mcp_server_auth_headers=mcp_server_auth_headers, - oauth2_headers=oauth2_headers, - raw_headers=raw_headers, - ) - verbose_logger.debug("Successfully fetched %s resources from managed MCP servers", len(managed_resources)) - except Exception as e: - verbose_logger.exception("Error getting resources from managed MCP servers: %s", e) - - return managed_resources - - async def _list_mcp_resource_templates( - user_api_key_auth: UserAPIKeyAuth | None = None, - mcp_auth_header: str | None = None, - mcp_servers: list[str] | None = None, - mcp_server_auth_headers: dict[str, dict[str, str]] | None = None, - oauth2_headers: dict[str, str] | None = None, - raw_headers: dict[str, str] | None = None, - ) -> list[ResourceTemplate]: - """List all available MCP resource templates.""" - - if not MCP_AVAILABLE: - return [] - - managed_resource_templates: list[ResourceTemplate] = [] - try: - managed_resource_templates = await _get_resource_templates_from_mcp_servers( - user_api_key_auth=user_api_key_auth, - mcp_auth_header=mcp_auth_header, - mcp_servers=mcp_servers, - mcp_server_auth_headers=mcp_server_auth_headers, - oauth2_headers=oauth2_headers, - raw_headers=raw_headers, - ) - verbose_logger.debug( - "Successfully fetched %s resource templates from managed MCP servers", - len(managed_resource_templates), - ) - except Exception as e: - verbose_logger.exception( - "Error getting resource templates from managed MCP servers: %s", - str(e), - ) - - return managed_resource_templates - - def _resolve_display_name_to_original( - name: str, - allowed_mcp_servers: list[MCPServer], - ) -> str: - """Translate a display-name override back to the original prefixed tool name. - - When a client received a customised display name from tools/list (e.g. - "Get Pet") it will call tools/call with that same string. We need to - reverse-map it to the original prefixed name (e.g. - "petstore_mcp-getPetById") before any routing or permission logic runs. - """ - for server in allowed_mcp_servers: - display_map = server.tool_name_to_display_name or {} - for unprefixed_name, display_name in display_map.items(): - if display_name == name: - return add_server_prefix_to_name(unprefixed_name, get_server_prefix(server)) - return name - - async def _get_byok_credential( - mcp_server: MCPServer, - user_api_key_auth: UserAPIKeyAuth | None, - ) -> str | None: - """Retrieve the stored BYOK credential for a user+server pair, served from the worker cache within its TTL.""" - if not mcp_server.is_byok: - return None - user_id: Final = (user_api_key_auth.user_id if user_api_key_auth else None) or "" - if not user_id: - return None - - cached: Final = get_cached_byok_credential(user_id, mcp_server.server_id) - if cached is not None: - return cached.credential - - from litellm.proxy._experimental.mcp_server.db import get_user_credential - from litellm.proxy.proxy_server import prisma_client - - if prisma_client is None: - return None - credential: Final = await get_user_credential( - prisma_client=prisma_client, - user_id=user_id, - server_id=mcp_server.server_id, - ) - cache_byok_credential(user_id, mcp_server.server_id, credential) - return credential - - async def _check_byok_credential( - mcp_server: MCPServer, - user_api_key_auth: UserAPIKeyAuth | None, - ) -> None: - """ - If the MCP server is BYOK-enabled, verify that the requesting user has a - stored credential. When no credential is found, raise an HTTP 401 with a - WWW-Authenticate header that points the MCP client to our OAuth metadata - endpoint so it can drive the authorization flow. - """ - if not mcp_server.is_byok: - return - - user_id: Final = (user_api_key_auth.user_id if user_api_key_auth else None) or "" - if not user_id: - raise HTTPException( - status_code=401, - detail={ - "error": "byok_auth_required", - "server_id": mcp_server.server_id, - "server_name": mcp_server.server_name or mcp_server.name, - "message": "User identity is required for BYOK servers", - }, - headers={"WWW-Authenticate": get_byok_www_authenticate()}, - ) - - cached: Final = get_cached_byok_credential(user_id, mcp_server.server_id) - if cached is not None: - if cached.credential is None: - raise HTTPException( - status_code=401, - detail={ - "error": "byok_auth_required", - "server_id": mcp_server.server_id, - "server_name": mcp_server.server_name or mcp_server.name, - "message": ( - "No stored credential found for this BYOK server. " - "Complete the OAuth authorization flow to provide your API key." - ), - }, - headers={"WWW-Authenticate": get_byok_www_authenticate()}, - ) - return - - from litellm.proxy._experimental.mcp_server.db import get_user_credential - from litellm.proxy.proxy_server import prisma_client - - if prisma_client is None: - # Fail closed on DB unavailability: returning here previously - # bypassed the ownership check and let any proxy-authenticated - # caller invoke BYOK tools during outage windows. - raise HTTPException( - status_code=503, - detail={ - "error": "byok_auth_unavailable", - "server_id": mcp_server.server_id, - "server_name": mcp_server.server_name or mcp_server.name, - "message": "BYOK credential check requires a database connection.", - }, - ) - - credential: Final = await get_user_credential( - prisma_client=prisma_client, - user_id=user_id, - server_id=mcp_server.server_id, - ) - cache_byok_credential(user_id, mcp_server.server_id, credential) - if credential is None: - raise HTTPException( - status_code=401, - detail={ - "error": "byok_auth_required", - "server_id": mcp_server.server_id, - "server_name": mcp_server.server_name or mcp_server.name, - "message": ( - "No stored credential found for this BYOK server. " - "Complete the OAuth authorization flow to provide your API key." - ), - }, - headers={"WWW-Authenticate": get_byok_www_authenticate()}, - ) - - async def _list_tools_before_first_call( - server: MCPServer | None, - tool_name: str, - allowed_mcp_servers: list[MCPServer], - user_api_key_auth: UserAPIKeyAuth | None, - mcp_auth_header: str | None, - mcp_server_auth_headers: dict[str, dict[str, str]] | None, - oauth2_headers: dict[str, str] | None, - raw_headers: dict[str, str] | None, - ) -> None: - """List ``server`` with the caller's own credentials when it does not yet expose ``tool_name`` here. - - The startup fill skips a server whose upstream wants the caller's token, and mcp 2 no - longer lists before an uncached tools/call, so a worker that has not served tools/list - for this caller would otherwise answer 404 for a tool the caller can see. Gating on the - requested tool, not on any prior listing, keeps callers with different upstream catalogs - from masking each other. - """ - if server is None or global_mcp_server_manager.server_exposes_tool(server, tool_name): - return - if all(allowed.server_id != server.server_id for allowed in allowed_mcp_servers): - return - try: - await _get_tools_from_mcp_servers( - user_api_key_auth=user_api_key_auth, - mcp_auth_header=mcp_auth_header, - mcp_servers=[server.server_id], - mcp_server_auth_headers=mcp_server_auth_headers, - oauth2_headers=oauth2_headers, - raw_headers=raw_headers, - ) - except Exception as e: # noqa: BLE001 # best effort: resolution below answers as it did before - verbose_logger.debug("MCP tools/call: listing %s before its first call failed: %s", server.name, e) - - async def execute_mcp_tool( - name: str, - arguments: dict[str, object], - allowed_mcp_servers: list[MCPServer], - start_time: datetime, - user_api_key_auth: UserAPIKeyAuth | None = None, - mcp_auth_header: str | None = None, - mcp_server_auth_headers: dict[str, dict[str, str]] | None = None, - oauth2_headers: dict[str, str] | None = None, - raw_headers: dict[str, str] | None = None, - host_progress_callback: Callable | None = None, - guardrail_context: Mapping[str, object] | None = None, - **kwargs: Any, - ) -> CallToolResult: - """ - Execute MCP tool. - - This function assumes permission checks have already been performed. - - Args: - name: Tool name (may include server prefix) - arguments: Tool arguments - allowed_mcp_servers: Pre-validated list of servers the user can access - start_time: Start time for logging - user_api_key_auth: Optional user API key auth for logging - mcp_auth_header: Optional MCP auth header - mcp_server_auth_headers: Optional server-specific auth headers - oauth2_headers: Optional OAuth2 headers - raw_headers: Optional raw HTTP headers - **kwargs: Additional arguments (e.g., litellm_logging_obj) - - Returns: - CallToolResult: Tool execution result - """ - # Track resolved MCP server for both permission checks and dispatch - mcp_server: MCPServer | None = None - requested_server_id: Final[str | None] = kwargs.get("requested_server_id") - - # If the client called with a display-name override (e.g. "Get Pet"), - # translate it back to the original prefixed name before any routing. - name = _resolve_display_name_to_original(name, allowed_mcp_servers) - - # Remove prefix from tool name for logging and processing - original_tool_name, server_name = split_server_prefix_from_name(name) - - requested_server: MCPServer | None = None - if requested_server_id: - requested_server = next( - (s for s in allowed_mcp_servers if s.server_id == requested_server_id), - None, - ) - - name_is_prefixed = False - if requested_server is not None and MCP_TOOL_PREFIX_SEPARATOR in name: - all_registry_prefixes: Final[set[str]] = set() - for registry_server in global_mcp_server_manager.get_registry().values(): - for known_prefix in iter_known_server_prefixes(registry_server): - all_registry_prefixes.add(normalize_server_name(known_prefix)) - name_is_prefixed = is_tool_name_prefixed(name, known_server_prefixes=all_registry_prefixes) - - first_call_target: Final = ( - requested_server - if requested_server is not None and not name_is_prefixed - else global_mcp_server_manager.server_owning_tool_name_prefix(name) - ) - first_call_tool_name: Final = ( - name - if first_call_target is None or (requested_server is not None and not name_is_prefixed) - else strip_known_server_prefix(name, first_call_target) - ) - await _list_tools_before_first_call( - server=first_call_target, - tool_name=first_call_tool_name, - allowed_mcp_servers=allowed_mcp_servers, - user_api_key_auth=user_api_key_auth, - mcp_auth_header=mcp_auth_header, - mcp_server_auth_headers=mcp_server_auth_headers, - oauth2_headers=oauth2_headers, - raw_headers=raw_headers, - ) - - if requested_server is not None and not name_is_prefixed: - # REST callers may pass server_id with the upstream tool name (no - # LiteLLM prefix). The first segment is not a registered server - # prefix, so the whole string is the upstream tool name and may - # legitimately contain the separator (e.g. "text-to-speech"). - # server_id is authoritative for routing and auth. - mcp_server = requested_server - server_name = requested_server.name - original_tool_name = name - else: - # Resolve from tool name (MCP JSON-RPC or prefixed REST tool names). - mcp_server = global_mcp_server_manager._get_mcp_server_from_tool_name(name) - if mcp_server is None and requested_server is not None: - for known_prefix in iter_known_server_prefixes(requested_server): - candidate = global_mcp_server_manager._get_mcp_server_from_tool_name( - add_server_prefix_to_name(name, known_prefix) - ) - if candidate is not None: - mcp_server = candidate - break - if mcp_server is not None: - server_name = mcp_server.name - original_tool_name = strip_known_server_prefix(name, mcp_server) - - if requested_server is not None: - if mcp_server is not None and mcp_server.server_id != requested_server.server_id: - raise HTTPException( - status_code=403, - detail={ - "error": "tool_server_mismatch", - "message": ( - f"Tool '{name}' belongs to MCP server " - f"'{mcp_server.name}' but request specified " - f"server_id for '{requested_server.name}'." - ), - }, - ) - if mcp_server is None: - mcp_server = requested_server - server_name = requested_server.name - original_tool_name = strip_known_server_prefix(name, requested_server) - - # Only enforce server-level permissions when we can resolve a server - if server_name: - if not MCPRequestHandler.is_tool_allowed( - allowed_mcp_servers=[server.name for server in allowed_mcp_servers], - server_name=server_name, - ): - raise HTTPException( - status_code=403, - detail="User not allowed to call this tool.", - ) - - standard_logging_mcp_tool_call: Final[StandardLoggingMCPToolCall] = _get_standard_logging_mcp_tool_call( - 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: Final[LiteLLMLoggingObj | None] = kwargs.get("litellm_logging_obj", None) - if litellm_logging_obj: - litellm_logging_obj.model_call_details["mcp_tool_call_metadata"] = standard_logging_mcp_tool_call - litellm_logging_obj.model = f"MCP: {name}" - litellm_logging_obj.model_call_details["model"] = f"MCP: {name}" - # Resolve the MCP server early so BYOK checks and credential injection - # apply to ALL dispatch paths (local tool registry AND managed MCP server). - if mcp_server is None: - mcp_server = global_mcp_server_manager._get_mcp_server_from_tool_name(name) - - if mcp_server: - standard_logging_mcp_tool_call["mcp_server_cost_info"] = (mcp_server.mcp_info or {}).get( - "mcp_server_cost_info" - ) - if litellm_logging_obj: - litellm_logging_obj.model_call_details["mcp_tool_call_metadata"] = standard_logging_mcp_tool_call - - # BYOK: retrieve the stored per-user credential. A single DB call - # both checks existence and fetches the value, avoiding a double query. - if mcp_server.is_byok and not mcp_auth_header: - byok_cred: Final = await _get_byok_credential(mcp_server, user_api_key_auth) - if byok_cred is None: - raise HTTPException( - status_code=401, - detail={ - "error": "byok_auth_required", - "server_id": mcp_server.server_id, - "server_name": mcp_server.server_name or mcp_server.name, - "message": ( - "No stored credential found for this BYOK server. " - "Complete the OAuth authorization flow to provide your API key." - ), - }, - headers={"WWW-Authenticate": get_byok_www_authenticate()}, - ) - mcp_auth_header = byok_cred - elif mcp_server.is_byok: - # External auth header supplied; still enforce user-identity check. - await _check_byok_credential(mcp_server, user_api_key_auth) - - # Check if tool exists in local registry first (for OpenAPI-based tools) - # These tools are registered with their prefixed names - ######################################################### - local_tool: Final = global_mcp_tool_registry.get_tool(name) - if local_tool: - # OpenAPI-backed tools used to bypass `pre_call_tool_check` — - # only the managed path ran allowed/banned-tool checks, key/team - # tool permissions, and parameter validation. Run the same checks - # before dispatching to the local registry. Refuse the call if - # we cannot resolve a server: tools registered via - # openapi_to_mcp_generator are always tied to a server, so a - # missing mcp_server here means the tool->server mapping has - # not finished initializing or the registry entry is orphaned. - # Skipping the check would re-open the same authorization gap. - if mcp_server is None: - raise HTTPException( - status_code=503, - detail=( - f"MCP server for tool '{name}' is not available; " - "refusing to dispatch without authorization checks. " - "Retry once the server is registered." - ), - ) - - # `pre_call_tool_check` calls into `proxy_logging_obj` for the - # pre-call guardrail hooks, so source it from the canonical - # `proxy_server` module the same way `_handle_managed_mcp_tool` - # does. `kwargs.get("proxy_logging_obj")` is None on the MCP - # entry path and would crash with AttributeError after the - # security checks pass. - from litellm.proxy.proxy_server import proxy_logging_obj - - hook_result = await global_mcp_server_manager.pre_call_tool_check( - name=original_tool_name, - arguments=arguments or {}, - server_name=server_name or mcp_server.name, - user_api_key_auth=user_api_key_auth, - proxy_logging_obj=proxy_logging_obj, - server=mcp_server, - raw_headers=raw_headers, - litellm_logging_obj=litellm_logging_obj, - guardrail_context=guardrail_context, - ) - # `pre_call_tool_check` may return guardrail-modified - # arguments; honor them on the local path too. - if isinstance(hook_result, dict) and "arguments" in hook_result: - arguments = hook_result["arguments"] - - verbose_logger.debug("Executing local registry tool: %s", name) - # The credential rides ContextVars because the tool function has its - # headers baked into the closure at registration time. - auth_header_value, openapi_forwarded_headers, upstream_credential = _resolve_openapi_tool_auth( - mcp_server=mcp_server, - mcp_auth_header=mcp_auth_header, - mcp_server_auth_headers=mcp_server_auth_headers, - raw_headers=raw_headers, - user_api_key_auth=user_api_key_auth, - ) - ( - resolved_auth_headers, - forwarded_headers, - ) = await global_mcp_server_manager.resolve_openapi_upstream_auth( - mcp_server=mcp_server, - oauth2_headers=oauth2_headers, - raw_headers=raw_headers, - mcp_auth_header=upstream_credential, - user_api_key_auth=user_api_key_auth, - forwarded_headers=openapi_forwarded_headers, - ) - - _auth_token: Final = _request_auth_header.set(auth_header_value) - _extra_token: Final = _request_extra_headers.set(forwarded_headers) - _resolved_token: Final = _request_resolved_auth_headers.set(resolved_auth_headers) - try: - response = await _handle_local_mcp_tool(name, arguments) - finally: - _request_auth_header.reset(_auth_token) - _request_extra_headers.reset(_extra_token) - _request_resolved_auth_headers.reset(_resolved_token) - - # Try managed MCP server tool (the name is bare; the prefix boundary was - # already resolved above against this server's registered prefixes) - # Primary and recommended way to use external MCP servers - ######################################################### - elif mcp_server: - response = await _handle_managed_mcp_tool( - server_name=server_name, - name=original_tool_name, - arguments=arguments, - user_api_key_auth=user_api_key_auth, - mcp_auth_header=mcp_auth_header, - mcp_server_auth_headers=mcp_server_auth_headers, - oauth2_headers=oauth2_headers, - raw_headers=raw_headers, - litellm_logging_obj=litellm_logging_obj, - guardrail_context=guardrail_context, - host_progress_callback=host_progress_callback, - ) - - # Fall back to local tool registry with original name (legacy support) - ######################################################### - # Deprecated: Local MCP Server Tool - ######################################################### - else: - # Gate only what can actually dispatch. When the unprefixed name is - # not in the registry either, `_handle_local_mcp_tool` below reports - # 404 and nothing runs, so demanding a server here would turn every - # unknown tool name into a misleading 503. - if global_mcp_tool_registry.get_tool(original_tool_name) is not None: - # `mcp_server` is None here because the tool name is not in the - # tool -> server mapping, but the name still carries a prefix - # that the server-level check above compared against the - # caller's `allowed_mcp_servers` by exact `name`. So the named - # server is in that list and can carry the tool-level checks, - # even with the mapping cold. Resolve it from - # `allowed_mcp_servers` rather than the registry: the registry - # would happily return a server the caller holds no grant for, - # and matching anything other than `name` would accept a server - # the check never validated. - prefix_server: Final = next( - (candidate for candidate in allowed_mcp_servers if candidate.name == server_name), - None, - ) - if prefix_server is None: - # A non-empty prefix that passed the server-level check - # always matches here, so this arm only fires when the - # prefix was empty, which is exactly the case that check - # skips. Fail closed rather than dispatch with no server to - # evaluate a tool ceiling against. - raise HTTPException( - status_code=503, - detail=( - f"MCP server for tool '{original_tool_name}' is not available; " - "refusing to dispatch without authorization checks. " - "Retry once the server is registered." - ), - ) - - from litellm.proxy.proxy_server import proxy_logging_obj - - hook_result = await global_mcp_server_manager.pre_call_tool_check( - name=original_tool_name, - arguments=arguments, - server_name=server_name, - user_api_key_auth=user_api_key_auth, - proxy_logging_obj=proxy_logging_obj, - server=prefix_server, - raw_headers=raw_headers, - litellm_logging_obj=litellm_logging_obj, - guardrail_context=guardrail_context, - ) - if "arguments" in hook_result: - arguments = hook_result["arguments"] # pyright: ignore[reportAny] # hook returns untyped args - - response = await _handle_local_mcp_tool(original_tool_name, arguments) - - return await _run_post_mcp_call_guardrails( - result=response, - litellm_logging_obj=litellm_logging_obj, - user_api_key_auth=user_api_key_auth, - request_data=kwargs, - ) - - async def _run_post_mcp_call_guardrails( - result: CallToolResult, - litellm_logging_obj: LiteLLMLoggingObj | None, - user_api_key_auth: UserAPIKeyAuth | None, - request_data: Mapping[str, object], - ) -> CallToolResult: - """Run ``post_mcp_call`` guardrails over an executed tool result. - - Lives on ``execute_mcp_tool``'s return path rather than inside - ``_fire_mcp_tool_call_logging`` so enforcement never depends on logging - being configured, and so every dispatch route gets it: the MCP protocol - handler, the REST endpoint, and tool search all funnel through here. - A guardrail that rejects the result raises, matching ``pre_mcp_call``. - """ - from litellm.proxy.proxy_server import proxy_logging_obj - - if proxy_logging_obj is None: - return result - return await proxy_logging_obj.post_mcp_call_hook( - response=result, - request_data=( - litellm_logging_obj.model_call_details if litellm_logging_obj is not None else dict(request_data) - ), - user_api_key_dict=user_api_key_auth, - ) - - _MCP_CREDENTIAL_REQUEST_FIELDS: Final = frozenset( - { - "raw_headers", - "mcp_auth_header", - "mcp_server_auth_headers", - "oauth2_headers", - "user_api_key_auth", - } + from litellm.proxy._experimental.mcp_server.operations import ( + _MCP_CREDENTIAL_REQUEST_FIELDS, + _aggregate_server_key, + _check_byok_credential, + _fire_mcp_tool_call_logging, + _get_byok_credential, + _get_prompts_from_mcp_servers, + _get_resource_templates_from_mcp_servers, + _get_resources_from_mcp_servers, + _get_standard_logging_mcp_tool_call, + _get_tools_from_mcp_servers, + _handle_local_mcp_tool, + _handle_managed_mcp_tool, + _list_mcp_prompts, + _list_mcp_resource_templates, + _list_mcp_resources, + _list_mcp_tools, + _list_tools_before_first_call, + _resolve_display_name_to_original, + _run_post_mcp_call_guardrails, + call_mcp_tool, + execute_mcp_tool, + filter_tools_by_key_team_permissions, + fire_mcp_tool_call_failure_logging, + mcp_get_prompt, + mcp_read_resource, ) - async def _fire_mcp_tool_call_logging( - logging_obj: LiteLLMLoggingObj, - result: CallToolResult, - start_time: datetime, - end_time: datetime, - user_api_key_auth: UserAPIKeyAuth | None = None, - request_data: Mapping[str, object] | None = None, - ) -> CallToolResult: - """Fire post-call logging for an executed MCP tool call, returning the result to send. - - The returned result is what the caller must forward to the client: a - ``post_mcp_call`` guardrail may rewrite the tool output (e.g. mask - sensitive values) or reject it, in which case its exception propagates. - Guardrails run before the success/failure logging so the masked text, not - the raw one, is what gets logged. - - A result with ``is_error=True`` is logged as a failure (``status="failure"`` - payload, so OTel marks the span ERROR) while the HTTP wire behavior stays - 200 + ``isError: true`` per the MCP spec. The error check runs after - ``async_post_mcp_tool_call_hook`` because guardrails may flip the result - to ``is_error=True`` in that hook. Raised exceptions never reach here (the - ``@client`` wrapper and ``call_mcp_tool``'s except path log those), so - this cannot double-log a failure. - - ``request_data`` may carry credential-bearing fields (the REST path puts - ``raw_headers``, ``mcp_auth_header``, ``mcp_server_auth_headers``, and - ``oauth2_headers`` at the top level of its data dict), so those are - stripped before the dict is handed to ``post_call_failure_hook`` - callbacks. - """ - from litellm.proxy.proxy_server import proxy_logging_obj - - logging_obj.post_call(original_response=result) - await logging_obj.async_post_mcp_tool_call_hook( - kwargs=logging_obj.model_call_details, - response_obj=result, - start_time=start_time, - end_time=end_time, - ) - logging_obj.call_type = CallTypes.call_mcp_tool.value - error_message: Final = extract_mcp_tool_result_error_message(result) - if error_message is None: - await logging_obj.async_success_handler(result=result, start_time=start_time, end_time=end_time) - return result - - logging_obj.has_run_logging(event_type="sync_success") - logging_obj.has_run_logging(event_type="async_success") - tool_error: Final = MCPToolResultError(error_message) - logging_obj.failure_handler(tool_error, "", start_time, end_time) - await logging_obj.async_failure_handler(tool_error, "", start_time, end_time) - - if user_api_key_auth is None: - return result - - if proxy_logging_obj: - sanitized_request_data: Final = { - key: value for key, value in (request_data or {}).items() if key not in _MCP_CREDENTIAL_REQUEST_FIELDS - } - await proxy_logging_obj.post_call_failure_hook( - request_data=sanitized_request_data, - original_exception=tool_error, - user_api_key_dict=user_api_key_auth, - route="/mcp/call_tool", - ) - return result - - async def fire_mcp_tool_call_failure_logging( - logging_obj: LiteLLMLoggingObj | None, - exception: Exception, - start_time: datetime, - user_api_key_auth: UserAPIKeyAuth | None, - request_data: Mapping[str, object], - ) -> None: - """Failure logging shared by the ``/mcp`` path and the REST endpoint. Call from - inside the ``except`` block so the traceback is still available. - - The failure handlers run first because ``_ProxyDBLogger.async_post_call_failure_hook`` - builds the failure spend-log row from the ``standard_logging_object`` they produce; - both gate on ``should_run_logging``, so the ``@client`` wrapper does not log twice. - A relayed upstream 401 (``MCPUpstreamAuthError``) is an expected caller-must-reauth - signal and skips ``post_call_failure_hook``, which fires the ``llm_exceptions`` alert. - """ - from litellm.proxy.proxy_server import proxy_logging_obj - - traceback_str: Final = traceback.format_exc(limit=MAXIMUM_TRACEBACK_LINES_TO_LOG) - if logging_obj is not None: - end_time: Final = datetime.now() # noqa: DTZ005 # naive to match `start_time`, which it is subtracted from - logging_obj.failure_handler(exception, traceback_str, start_time, end_time) - await logging_obj.async_failure_handler(exception, traceback_str, start_time, end_time) - - if isinstance(exception, MCPUpstreamAuthError) or not proxy_logging_obj or user_api_key_auth is None: - return - sanitized_request_data: Final = { - key: value for key, value in request_data.items() if key not in _MCP_CREDENTIAL_REQUEST_FIELDS - } - await proxy_logging_obj.post_call_failure_hook( - request_data=sanitized_request_data, - original_exception=exception, - user_api_key_dict=user_api_key_auth, - route="/mcp/call_tool", - traceback_str=traceback_str, - ) - - @client - async def call_mcp_tool( - name: str, - arguments: dict[str, object] | None = None, - user_api_key_auth: UserAPIKeyAuth | None = None, - mcp_auth_header: str | None = None, - mcp_servers: list[str] | None = None, - mcp_server_auth_headers: dict[str, dict[str, str]] | None = None, - oauth2_headers: dict[str, str] | None = None, - raw_headers: dict[str, str] | None = None, - client_ip: str | None = None, - **kwargs: Any, - ) -> CallToolResult: - """ - Call a specific tool with the provided arguments (handles prefixed tool names). - """ - start_time: Final = datetime.now() - litellm_logging_obj: Final[LiteLLMLoggingObj | None] = kwargs.get("litellm_logging_obj", None) - - try: - if arguments is None: - raise HTTPException(status_code=400, detail="Request arguments are required") - - ## CHECK IF USER IS ALLOWED TO CALL THIS TOOL - allowed_mcp_server_ids: Final = await global_mcp_server_manager.get_allowed_mcp_servers( - user_api_key_auth=user_api_key_auth, - ) - - allowed_mcp_servers: list[MCPServer] = [] - for allowed_mcp_server_id in allowed_mcp_server_ids: - allowed_server = global_mcp_server_manager.get_mcp_server_by_id(allowed_mcp_server_id) - if allowed_server is not None: - # Same request-time oauth2_flow backstop the listing path applies, - # so a null-flow M2M-shape row is treated as M2M on tool calls too. - allowed_server = MCPServerManager.resolve_oauth2_flow_for_request(allowed_server) - allowed_mcp_servers.append(allowed_server) - - allowed_mcp_servers = await _get_allowed_mcp_servers_from_mcp_server_names( - mcp_servers=mcp_servers, - allowed_mcp_servers=allowed_mcp_servers, - ) - if mcp_servers and not allowed_mcp_servers: - await raise_denied_scoped_mcp_access( - requested_names=mcp_servers, - user_api_key_auth=user_api_key_auth, - client_ip=client_ip, - ) - if not allowed_mcp_servers: - raise HTTPException( - status_code=403, - detail="User not allowed to call this tool.", - ) - - # Delegate to execute_mcp_tool for execution - response = await execute_mcp_tool( - name=name, - arguments=arguments, - allowed_mcp_servers=allowed_mcp_servers, - start_time=start_time, - user_api_key_auth=user_api_key_auth, - mcp_auth_header=mcp_auth_header, - mcp_server_auth_headers=mcp_server_auth_headers, - oauth2_headers=oauth2_headers, - raw_headers=raw_headers, - **kwargs, - ) - except Exception as e: - await fire_mcp_tool_call_failure_logging(litellm_logging_obj, e, start_time, user_api_key_auth, kwargs) - raise - - if litellm_logging_obj: - response = await _fire_mcp_tool_call_logging( - logging_obj=litellm_logging_obj, - result=response, - start_time=start_time, - end_time=datetime.now(), - user_api_key_auth=user_api_key_auth, - request_data=kwargs, - ) - return response - - async def mcp_get_prompt( - name: str, - arguments: dict[str, object] | None = None, - user_api_key_auth: UserAPIKeyAuth | None = None, - mcp_auth_header: str | None = None, - mcp_servers: list[str] | None = None, - mcp_server_auth_headers: dict[str, dict[str, str]] | None = None, - oauth2_headers: dict[str, str] | None = None, - raw_headers: dict[str, str] | None = None, - ) -> GetPromptResult: - """ - Fetch a specific MCP prompt, handling both prefixed and unprefixed names. - """ - allowed_mcp_servers: Final = await _get_allowed_mcp_servers( - user_api_key_auth=user_api_key_auth, - mcp_servers=mcp_servers, - ) - - if not allowed_mcp_servers: - raise HTTPException( - status_code=403, - detail="User not allowed to get this prompt.", - ) - - # Extract server name from prefixed prompt name - original_prompt_name, server_name = split_server_prefix_from_name(name) - - server: Final = next((s for s in allowed_mcp_servers if s.name == server_name), None) - if server is None: - raise HTTPException( - status_code=403, - detail="User not allowed to get this prompt.", - ) - - server_auth_header, extra_headers = _prepare_mcp_server_headers( - server=server, - mcp_server_auth_headers=mcp_server_auth_headers, - 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( - server=server, - user_api_key_auth=user_api_key_auth, - prompt_name=original_prompt_name, - arguments=arguments, - mcp_auth_header=server_auth_header, - extra_headers=extra_headers, - raw_headers=raw_headers, - ) - - async def mcp_read_resource( - url: AnyUrl, - user_api_key_auth: UserAPIKeyAuth | None = None, - mcp_auth_header: str | None = None, - mcp_servers: list[str] | None = None, - mcp_server_auth_headers: dict[str, dict[str, str]] | None = None, - oauth2_headers: dict[str, str] | None = None, - raw_headers: dict[str, str] | None = None, - ) -> ReadResourceResult: - """Read resource contents from upstream MCP servers.""" - - allowed_mcp_servers: Final = await _get_allowed_mcp_servers( - user_api_key_auth=user_api_key_auth, - mcp_servers=mcp_servers, - ) - - if not allowed_mcp_servers: - raise HTTPException( - status_code=403, - detail="User not allowed to read this resource.", - ) - - if len(allowed_mcp_servers) != 1: - raise HTTPException( - status_code=400, - detail=( - "Multiple MCP servers configured; read_resource currently supports exactly one allowed server." - ), - ) - - server: Final = allowed_mcp_servers[0] - - server_auth_header, extra_headers = _prepare_mcp_server_headers( - server=server, - mcp_server_auth_headers=mcp_server_auth_headers, - 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( - server=server, - user_api_key_auth=user_api_key_auth, - url=url, - mcp_auth_header=server_auth_header, - extra_headers=extra_headers, - raw_headers=raw_headers, - ) - - def _get_standard_logging_mcp_tool_call( - name: str, - arguments: dict[str, object], - server_name: str | None, - session_id: str | None = None, - ) -> StandardLoggingMCPToolCall: - mcp_server: Final = global_mcp_server_manager._get_mcp_server_from_tool_name( - add_server_prefix_to_name(name, server_name) if server_name else name - ) - namespaced_tool_name: Final = f"{server_name}/{name}" if server_name else name - if mcp_server: - mcp_info: Final = mcp_server.mcp_info or {} - return StandardLoggingMCPToolCall( - name=name, - arguments=arguments, - mcp_server_name=mcp_info.get("server_name"), - mcp_server_logo_url=mcp_info.get("logo_url"), - namespaced_tool_name=namespaced_tool_name, - mcp_session_id=session_id, - mcp_auth_mode=mcp_server.auth_type, - mcp_server_resource=_redact_mcp_resource_url(mcp_server.url), - ) - else: - return StandardLoggingMCPToolCall( - name=name, - arguments=arguments, - namespaced_tool_name=namespaced_tool_name, - mcp_session_id=session_id, - ) - - async def _handle_managed_mcp_tool( - server_name: str, - name: str, - arguments: dict[str, object], - user_api_key_auth: UserAPIKeyAuth | None = None, - mcp_auth_header: str | None = None, - mcp_server_auth_headers: dict[str, dict[str, str]] | None = None, - oauth2_headers: dict[str, str] | None = None, - raw_headers: dict[str, str] | None = None, - litellm_logging_obj: LiteLLMLoggingObj | None = None, - host_progress_callback: Callable | None = None, - guardrail_context: Mapping[str, object] | None = None, - ) -> CallToolResult: - """Handle tool execution for managed server tools""" - # Import here to avoid circular import - from litellm.proxy.proxy_server import proxy_logging_obj - - call_tool_result: Final = await global_mcp_server_manager.call_tool( - server_name=server_name, - name=name, - arguments=arguments, - user_api_key_auth=user_api_key_auth, - mcp_auth_header=mcp_auth_header, - mcp_server_auth_headers=mcp_server_auth_headers, - oauth2_headers=oauth2_headers, - raw_headers=raw_headers, - proxy_logging_obj=proxy_logging_obj, - host_progress_callback=host_progress_callback, - litellm_logging_obj=litellm_logging_obj, - guardrail_context=guardrail_context, - ) - verbose_logger.debug("CALL TOOL RESULT: %s", call_tool_result) - return call_tool_result - - async def _handle_local_mcp_tool(name: str, arguments: dict[str, object]) -> CallToolResult: - """Execute a local-registry tool and report whether it succeeded. - - Returns the result rather than bare content because the verdict is part of it: the content - alone cannot say whether the handler failed, so callers used to stamp is_error=False on every - outcome and an upstream rejection was served as tool output. - - A failure is reported as ``is_error=True`` here rather than raised, because the REST surface - turns an unrecognized exception into a 500 and an upstream 403 or 429 is not a gateway crash. - ``MCPUpstreamAuthError`` is the exception: it propagates so the caller is told to - re-authenticate, which both renderers already know how to say. - - Note: Local tools don't use prefixes, so we use the original name - """ - import inspect - - tool: Final = global_mcp_tool_registry.get_tool(name) - if not tool: - raise HTTPException(status_code=404, detail=f"Tool '{name}' not found") - - try: - if inspect.iscoroutinefunction(tool.handler): - result = await tool.handler(**arguments) - else: - result = tool.handler(**arguments) - except MCPUpstreamAuthError: - raise - except Exception as e: - verbose_logger.exception("Error executing local tool %s: %s", name, e) - return CallToolResult( - content=[TextContent(text=f"Error: {e}", type="text")], # mutable-ok: MCP result content - is_error=True, - ) - return CallToolResult( - content=[TextContent(text=str(result), type="text")], # mutable-ok: MCP result content - is_error=False, - ) - def _get_mcp_servers_in_path(path: str) -> list[str] | None: """ Get the MCP servers from the path @@ -4178,7 +1533,9 @@ if MCP_AVAILABLE: detail=f"API key does not have access to toolset '{toolset_id}'.", ) - tool_permissions = await global_mcp_server_manager.resolve_toolset_tool_permissions(toolset_ids=[toolset_id]) + tool_permissions = await operations.global_mcp_server_manager.resolve_toolset_tool_permissions( + toolset_ids=[toolset_id] + ) server_ids: Final = list(tool_permissions.keys()) existing_op: Final = user_api_key_auth.object_permission if existing_op is not None: @@ -4197,7 +1554,7 @@ if MCP_AVAILABLE: mcp_servers=server_ids, mcp_tool_permissions=tool_permissions, ) - return user_api_key_auth.model_copy(update={"object_permission": updated_op}) + return user_api_key_auth.model_copy(update={"object_permission": updated_op, "mcp_toolset_id": toolset_id}) async def _raise_preemptive_401_for_unauthenticated_servers( scope: Scope, @@ -4221,7 +1578,7 @@ if MCP_AVAILABLE: 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) + server = operations.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 @@ -4234,7 +1591,7 @@ if MCP_AVAILABLE: # authorization_url/token_url can change their inferred flow. continue if server is not None: - server = await global_mcp_server_manager.ensure_oauth_metadata_discovered(server) + server = await operations.global_mcp_server_manager.ensure_oauth_metadata_discovered(server) if server and server.auth_type == MCPAuth.oauth2: # The challenge decision is per oauth2 sub-mode, not per header: # gateway-managed modes (M2M and interactive authorization_code) @@ -4262,7 +1619,7 @@ if MCP_AVAILABLE: # authorization server is the gateway itself, vaulting via the # authorize interlude); the per-server relay advertised below # cannot vault without a litellm key on its token request. - if await global_mcp_server_manager.has_user_oauth_token(server, user_api_key_auth): + if await operations.global_mcp_server_manager.has_user_oauth_token(server, user_api_key_auth): continue if _is_mcp_admitted_user_subject(user_api_key_auth): @@ -4345,12 +1702,12 @@ if MCP_AVAILABLE: and server.server_id in frozenset( allowed.server_id - for allowed in await _get_allowed_mcp_servers( + for allowed in await operations._get_allowed_mcp_servers( user_api_key_auth=user_api_key_auth, mcp_servers=mcp_servers, client_ip=client_ip ) ) ): - await global_mcp_server_manager.preflight_token_exchange( + await operations.global_mcp_server_manager.preflight_token_exchange( server=server, oauth2_headers=oauth2_headers, user_api_key_auth=user_api_key_auth, @@ -4366,7 +1723,9 @@ if MCP_AVAILABLE: if ( server and server.is_oauth_passthrough - and not _client_has_passthrough_authorization(server, oauth2_headers, mcp_server_auth_headers) + and not operations._client_has_passthrough_authorization( + server, oauth2_headers, mcp_server_auth_headers + ) ): www_authenticate = get_passthrough_www_authenticate( scope=scope, @@ -4383,7 +1742,7 @@ if MCP_AVAILABLE: and server.is_oauth_delegate and len(mcp_servers or []) == 1 and _get_forwarded_auth_from_scope(scope) is None - and not _client_has_per_server_auth_header(server, mcp_server_auth_headers) + and not operations._client_has_per_server_auth_header(server, mcp_server_auth_headers) ): www_authenticate = get_passthrough_www_authenticate( scope=scope, @@ -4400,7 +1759,7 @@ if MCP_AVAILABLE: and server.is_true_passthrough and len(mcp_servers or []) == 1 and not _scope_has_authorization_header(scope) - and not _client_has_per_server_auth_header(server, mcp_server_auth_headers) + and not operations._client_has_per_server_auth_header(server, mcp_server_auth_headers) ): if server.is_dcr_bridge: raise HTTPException( @@ -4528,7 +1887,7 @@ if MCP_AVAILABLE: # Use the authorized server set, not the raw user-supplied names, so that # a caller cannot force a probe to a server their key is not allowed to use. - allowed_servers: Final = await _get_allowed_mcp_servers( + allowed_servers: Final = await operations._get_allowed_mcp_servers( user_api_key_auth=user_api_key_auth, mcp_servers=mcp_servers, client_ip=client_ip, diff --git a/litellm/proxy/_experimental/mcp_server/tool_search.py b/litellm/proxy/_experimental/mcp_server/tool_search.py index a482d02c31d..3650c722103 100644 --- a/litellm/proxy/_experimental/mcp_server/tool_search.py +++ b/litellm/proxy/_experimental/mcp_server/tool_search.py @@ -463,8 +463,8 @@ async def handle_mcp_tool_search( oauth2_headers: dict[str, str] | None = None, raw_headers: dict[str, str] | None = None, ) -> CallToolResult: - from litellm.proxy._experimental.mcp_server.server import ( - _list_mcp_tools, # pyright: ignore[reportPrivateUsage] # shared catalog owner + from litellm.proxy._experimental.mcp_server.operations import ( + _list_mcp_tools, ) from litellm.proxy.proxy_server import llm_router, proxy_logging_obj @@ -519,8 +519,8 @@ async def handle_mcp_proxy_tool( from jsonschema import validate from litellm.proxy import proxy_server - from litellm.proxy._experimental.mcp_server.server import ( # pyright: ignore[reportPrivateUsage] # shared catalog owner - _list_mcp_tools, # pyright: ignore[reportPrivateUsage] # shared catalog owner + from litellm.proxy._experimental.mcp_server.operations import ( + _list_mcp_tools, ) listing: Final = await _list_mcp_tools( @@ -607,7 +607,7 @@ async def handle_mcp_tool_call( requested_server_id: str | None = None, guardrail_context: Mapping[str, object] | None = None, ) -> CallToolResult: - from litellm.proxy._experimental.mcp_server.server import ( + from litellm.proxy._experimental.mcp_server.operations import ( _get_allowed_mcp_servers, execute_mcp_tool, raise_denied_scoped_mcp_access, @@ -643,6 +643,7 @@ async def handle_mcp_tool_call( mcp_server_auth_headers=mcp_server_auth_headers, oauth2_headers=oauth2_headers, raw_headers=raw_headers, + client_ip=client_ip, litellm_logging_obj=litellm_logging_obj, requested_server_id=requested_server_id, guardrail_context=guardrail_context, diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 3b34440d0bf..07338bcff52 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -3238,6 +3238,7 @@ class UserAPIKeyAuth(LiteLLM_VerificationTokenView): # the expected response ob # above; a forged value could at most narrow, but the stripping keeps the field's provenance # single-owner so its meaning stays trustworthy. mcp_session_resource_server_id: str | None = Field(default=None, exclude=True) + mcp_toolset_id: str | None = Field(default=None, exclude=True) via_virtual_key: bool = Field( default=False, exclude=True, @@ -3279,6 +3280,7 @@ class UserAPIKeyAuth(LiteLLM_VerificationTokenView): # the expected response ob values.pop("mcp_admitted_user_subject", None) values.pop("mcp_source_team_rpm_limits", None) values.pop("mcp_session_resource_server_id", None) + values.pop("mcp_toolset_id", None) values.pop("via_virtual_key", None) if values.get("api_key") is not None: values.update({"token": cls._safe_hash_litellm_api_key(values.get("api_key"))}) diff --git a/scripts/check_mcp_operation_boundary.py b/scripts/check_mcp_operation_boundary.py new file mode 100644 index 00000000000..b6c9dcefabf --- /dev/null +++ b/scripts/check_mcp_operation_boundary.py @@ -0,0 +1,65 @@ +import ast +import sys +from pathlib import Path +from typing import Final + +PACKAGE: Final = Path("litellm/proxy/_experimental/mcp_server") +LEGACY_ADAPTERS: Final = frozenset({"server.py", "legacy_callbacks.py", "mcp_context.py", "mcp_debug.py"}) +CONFINED_NAMES: Final = frozenset( + { + "auth_context_var", + "active_mcp_session_var", + "active_mcp_request_ctx_var", + "get_active_auth_context", + "get_active_mcp_session", + "get_active_mcp_request_ctx", + "get_or_extract_auth_context", + "_session_obj_auth_storage", + "WeakKeyDictionary", + "_mcp_active_toolset_id", + "_mcp_gateway_initialize_instructions", + "_mcp_gateway_server_name", + "_mcp_proxy_mode", + } +) + + +def is_confined(name: str) -> bool: + return name in CONFINED_NAMES or name.startswith("_stateful_session_") + + +def violations(path: Path, source: str) -> tuple[str, ...]: + if path.name in LEGACY_ADAPTERS: + return () + tree: Final = ast.parse(source, filename=str(path)) + return tuple( + f"{path}:{node.lineno}: MCP request/session state belongs in a legacy adapter" + for node in ast.walk(tree) + if ( + isinstance(node, ast.ImportFrom) + and ( + (node.module or "").endswith(".mcp_context") + or any(is_confined(alias.name) for alias in node.names) + or (path.name in {"operations.py", "contracts.py"} and (node.module or "").endswith(".server")) + ) + or isinstance(node, ast.Name) + and is_confined(node.id) + or isinstance(node, ast.Attribute) + and is_confined(node.attr) + ) + ) + + +def main() -> int: + findings: Final = tuple( + finding for path in sorted(PACKAGE.rglob("*.py")) for finding in violations(path, path.read_text()) + ) + if findings: + print("\n".join(findings), file=sys.stderr) + return 1 + print("MCP operation boundary: passed") + return 0 + + +if __name__ == "__main__": + sys.exit(main()) diff --git a/scripts/pre_commit_lint.sh b/scripts/pre_commit_lint.sh index 1abd415d237..22cc38f841c 100755 --- a/scripts/pre_commit_lint.sh +++ b/scripts/pre_commit_lint.sh @@ -102,6 +102,9 @@ ui_prettier_pattern='^ui/litellm-dashboard/.*\.(js|jsx|ts|tsx|mjs|cjs|json|css|s ui_eslint_pattern='^ui/litellm-dashboard/.*\.(js|jsx|ts|tsx|mjs|cjs)$' litellm_py_files=$(scope_match "$litellm_py_pattern") +if [ -n "$(scope_match '^(litellm/proxy/_experimental/mcp_server/|scripts/check_mcp_operation_boundary\.py)')" ]; then + uv run --no-sync python scripts/check_mcp_operation_boundary.py || exit 1 +fi e2e_py_files=$(scope_match "$e2e_py_pattern") test_tree_files=$(scope_match "$test_tree_pattern") # ruff format (and CI's format step) skip enterprise; the rest of make lint covers it. diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_byok_oauth_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_byok_oauth_endpoints.py index 87e23893616..a77b4c8d565 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_byok_oauth_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_byok_oauth_endpoints.py @@ -1,3 +1,4 @@ +from litellm.proxy._experimental.mcp_server import operations as mcp_operations """ Unit tests for the BYOK OAuth 2.1 authorization server endpoints. @@ -592,7 +593,7 @@ async def test_check_byok_credential_missing_credential(monkeypatch): monkeypatch.delenv("PROXY_BASE_URL", raising=False) monkeypatch.delenv("SERVER_ROOT_PATH", raising=False) - server_module.byok_credential_cache.flush_cache() + mcp_operations.byok_credential_cache.flush_cache() mock_prisma = MagicMock() with ( @@ -628,13 +629,13 @@ async def test_execute_byok_tool_missing_credential_advertises_api_key_flow(monk from litellm.types.mcp_server.mcp_server_manager import MCPServer monkeypatch.setenv("PROXY_BASE_URL", "https://gateway.example.com/proxy") - mcp_module.byok_credential_cache.flush_cache() + mcp_operations.byok_credential_cache.flush_cache() server = MCPServer(server_id="byok-discovery", name="byok-discovery", transport=MCPTransport.http, is_byok=True) prisma = MagicMock() prisma.db.litellm_mcpusercredentials.find_unique = AsyncMock(return_value=None) monkeypatch.setattr(proxy_server, "prisma_client", prisma) with pytest.raises(HTTPException) as exc_info: - await mcp_module.execute_mcp_tool( + await mcp_operations.execute_mcp_tool( name="list_regions", arguments={}, allowed_mcp_servers=[server], @@ -687,7 +688,7 @@ async def test_invalidate_byok_cred_cache_evicts_locally_and_broadcasts_the_same server = MCPServer(server_id="byok-revoke", name="byok-server", transport=MCPTransport.http, is_byok=True) user_auth = UserAPIKeyAuth(user_id="mallory", api_key="sk-test") - server_module.byok_credential_cache.flush_cache() + mcp_operations.byok_credential_cache.flush_cache() db_lookup = AsyncMock(side_effect=["sk-before-revoke", None]) publish = AsyncMock() @@ -699,13 +700,13 @@ async def test_invalidate_byok_cred_cache_evicts_locally_and_broadcasts_the_same "litellm.proxy.proxy_server.prisma_client", MagicMock() ), patch.object( # test-quality-ok: the redis publisher is module-level; asserting the broadcast without a redis - server_module, "publish_auth_cache_invalidation", new=publish + mcp_operations, "publish_auth_cache_invalidation", new=publish ), ): - assert await server_module._get_byok_credential(server, user_auth) == "sk-before-revoke" - assert await server_module._get_byok_credential(server, user_auth) == "sk-before-revoke" - await server_module._invalidate_byok_cred_cache("mallory", "byok-revoke") - assert await server_module._get_byok_credential(server, user_auth) is None + assert await mcp_operations._get_byok_credential(server, user_auth) == "sk-before-revoke" + assert await mcp_operations._get_byok_credential(server, user_auth) == "sk-before-revoke" + await mcp_operations._invalidate_byok_cred_cache("mallory", "byok-revoke") + assert await mcp_operations._get_byok_credential(server, user_auth) is None assert db_lookup.await_count == 2 publish.assert_awaited_once_with(cache_key=byok_credential_cache_key("mallory", "byok-revoke")) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_contracts.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_contracts.py new file mode 100644 index 00000000000..e13ecdfcce9 --- /dev/null +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_contracts.py @@ -0,0 +1,60 @@ +from dataclasses import FrozenInstanceError + +import pytest + +from litellm.proxy._experimental.mcp_server.operations import prepare_context +from litellm.proxy._types import UserAPIKeyAuth + + +def test_operation_context_isolates_nested_headers_and_caller_permissions(): + caller = UserAPIKeyAuth(user_id="alpha", models=["allowed"]) + caller.mcp_admitted_user_subject = True + caller.mcp_session_resource_server_id = "alpha-server" + caller.mcp_toolset_id = "toolset-alpha" + caller.mcp_source_team_rpm_limits = {"team": {"alpha-server": 2}} + headers = {"x-caller": "alpha"} + server_headers = {"alpha-server": {"authorization": "alpha-token"}} + context = prepare_context(caller, raw_headers=headers, mcp_server_auth_headers=server_headers) + + caller.models.append("forbidden") + caller.mcp_source_team_rpm_limits["team"]["alpha-server"] = 999 + headers["x-caller"] = "bravo" + server_headers["alpha-server"]["authorization"] = "bravo-token" + captured = context.user_api_key_auth + assert captured is not None + assert captured.models == ["allowed"] + assert captured.mcp_admitted_user_subject is True + assert captured.mcp_session_resource_server_id == "alpha-server" + assert captured.mcp_toolset_id == "toolset-alpha" + assert captured.mcp_source_team_rpm_limits == {"team": {"alpha-server": 2}} + captured.models.append("also-forbidden") + assert context.user_api_key_auth.models == ["allowed"] + assert context.raw_headers == {"x-caller": "alpha"} + assert context.mcp_server_auth_headers == {"alpha-server": {"authorization": "alpha-token"}} + with pytest.raises(TypeError): + context.raw_headers["x-caller"] = "changed" + with pytest.raises(TypeError): + context.mcp_server_auth_headers["alpha-server"]["authorization"] = "changed" + with pytest.raises(FrozenInstanceError): + context.client_ip = "untrusted" + + +def test_operation_context_preserves_missing_and_empty_inputs(): + missing = prepare_context() + empty = prepare_context(mcp_servers=[], raw_headers={}, oauth2_headers={}, mcp_server_auth_headers={}) + assert missing.user_api_key_auth is None + assert missing.mcp_servers is None + assert missing.raw_headers is None + assert missing.oauth2_headers is None + assert missing.mcp_server_auth_headers is None + assert empty.mcp_servers == () + assert empty.raw_headers == {} + assert empty.oauth2_headers == {} + assert empty.mcp_server_auth_headers == {} + + +def test_toolset_request_marker_cannot_be_supplied_by_caller_or_serialized(): + auth = UserAPIKeyAuth.model_validate({"user_id": "alpha", "mcp_toolset_id": "forged"}) + assert auth.mcp_toolset_id is None + auth.mcp_toolset_id = "server-resolved" + assert "mcp_toolset_id" not in auth.model_dump() diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_block_recording.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_block_recording.py index 64d926bc5e3..b8aadef430f 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_block_recording.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_block_recording.py @@ -1,5 +1,6 @@ +from litellm.proxy._experimental.mcp_server import operations as mcp_operations """Tests for guardrail-block recording in -``litellm.proxy._experimental.mcp_server.server.call_mcp_tool``. +``litellm.proxy._experimental.mcp_server.operations.call_mcp_tool``. A pre-call MCP guardrail block *raises* into ``call_mcp_tool``'s ``except Exception``. The failure spend-log row that the Guardrails Monitor's @@ -70,7 +71,7 @@ async def _call_block(logging_obj, order: list, *, user_api_key_auth=mock.sentin with mock.patch.dict(sys.modules, {"litellm.proxy.proxy_server": fake_proxy_server}): with contextlib.suppress(HTTPException): - await server.call_mcp_tool.__wrapped__( + await mcp_operations.call_mcp_tool.__wrapped__( name="t", arguments=None, user_api_key_auth=user_api_key_auth, 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 28faf375ab8..9659eb1cbc2 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 @@ -1229,7 +1229,7 @@ class TestResolveByokMcpAuthHeader: user_auth = UserAPIKeyAuth(user_id="user-1", api_key="sk-dashboard") with patch( - "litellm.proxy._experimental.mcp_server.server._get_byok_credential", + "litellm.proxy._experimental.mcp_server.operations._get_byok_credential", new=AsyncMock(return_value="stored-cred"), ): result = await _resolve_byok_mcp_auth_header(server, user_auth, None) @@ -1249,7 +1249,7 @@ class TestResolveByokMcpAuthHeader: user_auth = UserAPIKeyAuth(user_id="user-1", api_key="sk-dashboard") with patch( - "litellm.proxy._experimental.mcp_server.server._get_byok_credential", + "litellm.proxy._experimental.mcp_server.operations._get_byok_credential", new=AsyncMock(return_value=None), ): with pytest.raises(HTTPException) as exc_info: @@ -1272,7 +1272,7 @@ class TestResolveByokMcpAuthHeader: check_mock = AsyncMock(return_value=None) with patch( - "litellm.proxy._experimental.mcp_server.server._check_byok_credential", + "litellm.proxy._experimental.mcp_server.operations._check_byok_credential", new=check_mock, ): result = await _resolve_byok_mcp_auth_header(server, user_auth, "caller-header") 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 index 3f5d4ad83ea..1909e3306a2 100644 --- 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 @@ -1,3 +1,4 @@ +from litellm.proxy._experimental.mcp_server import operations as mcp_operations """Unit tests for MCP OAuth passthrough tool-fetch behavior.""" import logging @@ -339,16 +340,16 @@ async def test_aggregate_list_tools_absorbs_one_unauthenticated_server(): raise MCPUpstreamAuthError(status_code=401, www_authenticate=None, server_name=server.name) return [good_tool] - with patch.object(mcp_server, "_get_allowed_mcp_servers", AsyncMock(return_value=[delegate, working])), patch.object( - mcp_server, "_prefetch_oauth_creds_for_user", AsyncMock(return_value={}) - ), patch.object(mcp_server, "_prepare_mcp_server_headers", MagicMock(return_value=(None, None))), patch.object( - mcp_server, "_get_user_oauth_extra_headers_from_db", AsyncMock(return_value=None) + with patch.object(mcp_operations, "_get_allowed_mcp_servers", AsyncMock(return_value=[delegate, working])), patch.object( + mcp_operations, "_prefetch_oauth_creds_for_user", AsyncMock(return_value={}) + ), patch.object(mcp_operations, "_prepare_mcp_server_headers", MagicMock(return_value=(None, None))), patch.object( + mcp_operations, "_get_user_oauth_extra_headers_from_db", AsyncMock(return_value=None) ), patch.object( - mcp_server, "filter_tools_by_key_team_permissions", AsyncMock(side_effect=lambda tools, **k: tools) + mcp_operations, "filter_tools_by_key_team_permissions", AsyncMock(side_effect=lambda tools, **k: tools) ), patch.object( - mcp_server.global_mcp_server_manager, "_get_tools_from_server", AsyncMock(side_effect=fake_get_tools) + mcp_operations.global_mcp_server_manager, "_get_tools_from_server", AsyncMock(side_effect=fake_get_tools) ): - listing = await mcp_server._get_tools_from_mcp_servers( + listing = await mcp_operations._get_tools_from_mcp_servers( user_api_key_auth=UserAPIKeyAuth(token="h", user_id="u1"), mcp_auth_header=None, mcp_servers=None, @@ -382,14 +383,14 @@ async def test_single_server_route_also_absorbs_upstream_auth_error(): # //mcp sets the path-derived single-server scope; absorption must hold even then. token = _mcp_gateway_server_name.set("delegate_docs") try: - with patch.object(mcp_server, "_get_allowed_mcp_servers", AsyncMock(return_value=[delegate])), patch.object( - mcp_server, "_prefetch_oauth_creds_for_user", AsyncMock(return_value={}) - ), patch.object(mcp_server, "_prepare_mcp_server_headers", MagicMock(return_value=(None, None))), patch.object( - mcp_server, "_get_user_oauth_extra_headers_from_db", AsyncMock(return_value=None) + with patch.object(mcp_operations, "_get_allowed_mcp_servers", AsyncMock(return_value=[delegate])), patch.object( + mcp_operations, "_prefetch_oauth_creds_for_user", AsyncMock(return_value={}) + ), patch.object(mcp_operations, "_prepare_mcp_server_headers", MagicMock(return_value=(None, None))), patch.object( + mcp_operations, "_get_user_oauth_extra_headers_from_db", AsyncMock(return_value=None) ), patch.object( - mcp_server.global_mcp_server_manager, "_get_tools_from_server", AsyncMock(side_effect=fake_get_tools) + mcp_operations.global_mcp_server_manager, "_get_tools_from_server", AsyncMock(side_effect=fake_get_tools) ): - listing = await mcp_server._get_tools_from_mcp_servers( + listing = await mcp_operations._get_tools_from_mcp_servers( user_api_key_auth=UserAPIKeyAuth(token="h", user_id="u1"), mcp_auth_header=None, mcp_servers=["delegate_docs"], @@ -419,15 +420,15 @@ async def test_aggregate_with_single_accessible_server_still_absorbs(): async def fake_get_tools(server, **kwargs): raise MCPUpstreamAuthError(status_code=401, www_authenticate=None, server_name=server.name) - with patch.object(mcp_server, "_get_allowed_mcp_servers", AsyncMock(return_value=[delegate])), patch.object( - mcp_server, "_prefetch_oauth_creds_for_user", AsyncMock(return_value={}) - ), patch.object(mcp_server, "_prepare_mcp_server_headers", MagicMock(return_value=(None, None))), patch.object( - mcp_server, "_get_user_oauth_extra_headers_from_db", AsyncMock(return_value=None) + with patch.object(mcp_operations, "_get_allowed_mcp_servers", AsyncMock(return_value=[delegate])), patch.object( + mcp_operations, "_prefetch_oauth_creds_for_user", AsyncMock(return_value={}) + ), patch.object(mcp_operations, "_prepare_mcp_server_headers", MagicMock(return_value=(None, None))), patch.object( + mcp_operations, "_get_user_oauth_extra_headers_from_db", AsyncMock(return_value=None) ), patch.object( - mcp_server.global_mcp_server_manager, "_get_tools_from_server", AsyncMock(side_effect=fake_get_tools) + mcp_operations.global_mcp_server_manager, "_get_tools_from_server", AsyncMock(side_effect=fake_get_tools) ): # Aggregate route: no explicit server filter, even though only one server is accessible. - listing = await mcp_server._get_tools_from_mcp_servers( + listing = await mcp_operations._get_tools_from_mcp_servers( user_api_key_auth=UserAPIKeyAuth(token="h", user_id="u1"), mcp_auth_header=None, mcp_servers=None, @@ -475,3 +476,25 @@ async def test_client_creation_failure_logs_sanitized_exchange(monkeypatch, capl await manager._get_tools_from_server(server) assert "POST https://upstream/ -> HTTP 500" in caplog.text assert "missing_scope" in caplog.text and "query-secret" not in caplog.text + + +@pytest.mark.parametrize( + "oauth_headers,server_headers,authorized", + [ + ({"Authorization": "Bearer upstream"}, None, True), + ({"AUTHORIZATION": "Bearer upstream"}, None, True), + ({"x-unrelated": "present"}, None, False), + (None, {"catalog": {"Authorization": "Bearer scoped"}}, True), + (None, {"other-server": {"Authorization": "Bearer unrelated"}}, False), + (None, {"catalog": {"x-unrelated": "present"}}, False), + (None, {"catalog": "Bearer legacy"}, True), + (None, {"catalog": " "}, False), + ], +) +def test_passthrough_admission_recognizes_only_matching_authorization(oauth_headers, server_headers, authorized): + from litellm.proxy._experimental.mcp_server.operations import _client_has_passthrough_authorization + from litellm.types.mcp import MCPTransport + from litellm.types.mcp_server.mcp_server_manager import MCPServer + + server = MCPServer(server_id="catalog", name="catalog", alias="catalog", transport=MCPTransport.http) + assert _client_has_passthrough_authorization(server, oauth_headers, server_headers) is authorized diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_proxy_mode.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_proxy_mode.py index 84d4f1fd083..ed5d67164bd 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_proxy_mode.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_proxy_mode.py @@ -1,3 +1,4 @@ +from litellm.proxy._experimental.mcp_server import operations as mcp_operations import json from datetime import datetime @@ -27,8 +28,8 @@ def proxy_mode(): @pytest.mark.asyncio @pytest.mark.usefixtures("proxy_mode") async def test_proxy_call_rejects_non_proxy_tool_names() -> None: - result = await server._dispatch_virtual_mcp_tool( - name="math_stdio-add", arguments={"a": 1, "b": 2}, user_api_key_auth=AUTH, client_ip=None + result = await mcp_operations._dispatch_virtual_mcp_tool( + name="math_stdio-add", arguments={"a": 1, "b": 2}, user_api_key_auth=AUTH, client_ip=None, mcp_proxy_mode=True ) assert result is not None @@ -105,12 +106,13 @@ async def test_proxy_scope_exception_emits_failure_log(monkeypatch: pytest.Monke arguments = {"tool_id": "denied-scope", "arguments": {}} with pytest.raises(HTTPException) as denied: - await server._dispatch_virtual_mcp_tool( + await mcp_operations._dispatch_virtual_mcp_tool( name="call_tool", arguments=arguments, user_api_key_auth=auth, client_ip=None, mcp_servers=["ungranted"], + mcp_proxy_mode=True, raw_headers={"authorization": "Bearer raw-scope-secret", "x-litellm-call-id": "scope-denial"}, ) 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 3668a06203c..169de9095cf 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 @@ -1,3 +1,4 @@ +from litellm.proxy._experimental.mcp_server import operations as mcp_operations import asyncio import contextlib import contextvars @@ -138,7 +139,7 @@ async def test_mcp_server_tool_call_body_contains_request_data(_mcp_request_ctx) mock_add_litellm_data_to_request, ): with patch( - "litellm.proxy._experimental.mcp_server.server.call_mcp_tool", + "litellm.proxy._experimental.mcp_server.operations.call_mcp_tool", mock_call_mcp_tool, ): with patch( @@ -194,7 +195,7 @@ async def test_mcp_server_tool_call_forwards_client_headers_to_logging(_mcp_requ mock_add_litellm_data_to_request, ): with patch( - "litellm.proxy._experimental.mcp_server.server.call_mcp_tool", + "litellm.proxy._experimental.mcp_server.operations.call_mcp_tool", mock_call_mcp_tool, ): with patch("litellm.proxy.proxy_server.proxy_config", MagicMock()): @@ -241,7 +242,7 @@ async def test_mcp_server_tool_call_strips_custom_litellm_key_header(_mcp_reques capturing_add_litellm_data_to_request, ): with patch( - "litellm.proxy._experimental.mcp_server.server.call_mcp_tool", + "litellm.proxy._experimental.mcp_server.operations.call_mcp_tool", mock_call_mcp_tool, ): with patch("litellm.proxy.proxy_server.proxy_config", MagicMock()): @@ -287,11 +288,11 @@ async def test_mcp_server_tool_call_relays_upstream_auth_error_as_iserror(_mcp_r mock_add_litellm_data_to_request, ): with patch( - "litellm.proxy._experimental.mcp_server.server.call_mcp_tool", + "litellm.proxy._experimental.mcp_server.operations.call_mcp_tool", mock_call_mcp_tool, ): with patch("litellm.proxy.proxy_server.proxy_config", MagicMock()): - with patch("litellm.proxy._experimental.mcp_server.server.verbose_logger", mock_logger): + with patch("litellm.proxy._experimental.mcp_server.operations.verbose_logger", mock_logger): result = await mcp_server_tool_call(_mcp_request_ctx(), _call_tool_params("test_tool", {"param": "value"})) assert result.is_error is True @@ -867,15 +868,15 @@ async def test_get_prompts_from_mcp_servers_success(): with ( patch( - "litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", + "litellm.proxy._experimental.mcp_server.operations._get_allowed_mcp_servers", AsyncMock(return_value=[server_a, server_b]), ) as mock_allowed, patch( - "litellm.proxy._experimental.mcp_server.server._prepare_mcp_server_headers", + "litellm.proxy._experimental.mcp_server.operations._prepare_mcp_server_headers", return_value=(None, None), ) as mock_headers, patch( - "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager", + "litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager", ) as mock_manager, ): mock_manager.get_prompts_from_server = AsyncMock( @@ -927,15 +928,15 @@ async def test_get_resources_from_mcp_servers_success(): with ( patch( - "litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", + "litellm.proxy._experimental.mcp_server.operations._get_allowed_mcp_servers", AsyncMock(return_value=[server_a, server_b]), ) as mock_allowed, patch( - "litellm.proxy._experimental.mcp_server.server._prepare_mcp_server_headers", + "litellm.proxy._experimental.mcp_server.operations._prepare_mcp_server_headers", return_value=(None, None), ) as mock_headers, patch( - "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager", + "litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager", ) as mock_manager, ): mock_manager.get_resources_from_server = AsyncMock( @@ -992,15 +993,15 @@ async def test_get_resource_templates_from_mcp_servers_success(): with ( patch( - "litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", + "litellm.proxy._experimental.mcp_server.operations._get_allowed_mcp_servers", AsyncMock(return_value=[server]), ) as mock_allowed, patch( - "litellm.proxy._experimental.mcp_server.server._prepare_mcp_server_headers", + "litellm.proxy._experimental.mcp_server.operations._prepare_mcp_server_headers", return_value=(None, None), ) as mock_headers, patch( - "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager", + "litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager", ) as mock_manager, ): mock_manager.get_resource_templates_from_server = AsyncMock( @@ -1042,15 +1043,15 @@ async def test_mcp_get_prompt_success(): with ( patch( - "litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", + "litellm.proxy._experimental.mcp_server.operations._get_allowed_mcp_servers", AsyncMock(return_value=[server]), ) as mock_allowed, patch( - "litellm.proxy._experimental.mcp_server.server._prepare_mcp_server_headers", + "litellm.proxy._experimental.mcp_server.operations._prepare_mcp_server_headers", return_value=({"Authorization": "token"}, {"X-Test": "1"}), ) as mock_headers, patch( - "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager", + "litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager", ) as mock_manager, ): mock_manager.get_prompt_from_server = AsyncMock(return_value=prompt_result) @@ -1078,6 +1079,7 @@ async def test_mcp_get_prompt_success(): mcp_auth_header={"Authorization": "token"}, extra_headers={"X-Test": "1"}, raw_headers=None, + client_ip=None, ) assert result is prompt_result @@ -1106,15 +1108,15 @@ async def test_mcp_read_resource_success(): with ( patch( - "litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", + "litellm.proxy._experimental.mcp_server.operations._get_allowed_mcp_servers", AsyncMock(return_value=[server]), ) as mock_allowed, patch( - "litellm.proxy._experimental.mcp_server.server._prepare_mcp_server_headers", + "litellm.proxy._experimental.mcp_server.operations._prepare_mcp_server_headers", return_value=({"Authorization": "token"}, {"X-Test": "1"}), ) as mock_headers, patch( - "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager", + "litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager", ) as mock_manager, ): mock_manager.read_resource_from_server = AsyncMock(return_value=read_result) @@ -1140,6 +1142,7 @@ async def test_mcp_read_resource_success(): mcp_auth_header={"Authorization": "token"}, extra_headers={"X-Test": "1"}, raw_headers=None, + client_ip=None, ) assert result is read_result @@ -1264,7 +1267,7 @@ async def test_mcp_read_resource_multiple_servers_error(): server_b.name = "server_b" with patch( - "litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", + "litellm.proxy._experimental.mcp_server.operations._get_allowed_mcp_servers", AsyncMock(return_value=[server_a, server_b]), ) as mock_allowed: with pytest.raises(HTTPException) as exc_info: @@ -1354,11 +1357,11 @@ async def test_get_tools_from_mcp_servers_continues_when_one_server_fails(): mock_manager._get_tools_from_server = mock_get_tools_from_server with patch( - "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager", + "litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager", mock_manager, ): with patch( - "litellm.proxy._experimental.mcp_server.server.verbose_logger", + "litellm.proxy._experimental.mcp_server.operations.verbose_logger", ) as mock_logger: # Test with server-specific auth headers mcp_server_auth_headers = { @@ -1450,11 +1453,11 @@ async def test_get_tools_from_mcp_servers_handles_all_servers_failing(): mock_manager._get_tools_from_server = mock_get_tools_from_server with patch( - "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager", + "litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager", mock_manager, ): with patch( - "litellm.proxy._experimental.mcp_server.server.verbose_logger", + "litellm.proxy._experimental.mcp_server.operations.verbose_logger", ) as mock_logger: # Test with server-specific auth headers mcp_server_auth_headers = { @@ -1524,11 +1527,11 @@ async def _denied_scoped_list( with ( patch( # test-quality-ok: the permission resolver is a module-level function; the suite's only seam - "litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", + "litellm.proxy._experimental.mcp_server.operations._get_allowed_mcp_servers", resolver, ), patch( # test-quality-ok: the server registry is a module-level singleton; the suite's only seam - "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager", + "litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager", mock_manager, ), ): @@ -1575,11 +1578,11 @@ async def test_empty_scope_lists_nothing_instead_of_raising_a_nameless_denial(): with ( patch( # test-quality-ok: the permission resolver is a module-level function; the suite's only seam - "litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", + "litellm.proxy._experimental.mcp_server.operations._get_allowed_mcp_servers", resolver, ), patch( # test-quality-ok: the server registry is a module-level singleton; the suite's only seam - "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager", + "litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager", _denied_scope_manager({"github": "srv-github"}), ), ): @@ -1721,7 +1724,8 @@ async def test_scoped_list_agent_veto_attributed_for_differently_cased_server_na @pytest.mark.asyncio -async def test_handle_list_tools_converts_permission_httpexception_to_mcp_error(_mcp_request_ctx): +@pytest.mark.parametrize("denial_at_auth", [False, True]) +async def test_handle_list_tools_converts_permission_httpexception_to_mcp_error(_mcp_request_ctx, denial_at_auth): """The MCP protocol handler surfaces a permission HTTPException as a clean JSON-RPC error (MCPError, INVALID_REQUEST) carrying the denial message, instead of a raw 500.""" try: @@ -1738,10 +1742,10 @@ async def test_handle_list_tools_converts_permission_httpexception_to_mcp_error( with ( patch( # test-quality-ok: the protocol handler reads auth from module context; no injection seam "litellm.proxy._experimental.mcp_server.server.get_or_extract_auth_context", - new=AsyncMock(return_value=(None, None, None, None, None, None, None)), + new=AsyncMock(return_value=(None, None, None, None, None, None, None), side_effect=denial if denial_at_auth else None), ), patch( # test-quality-ok: the listing helper is the handler's only collaborator; the suite's seam - "litellm.proxy._experimental.mcp_server.server._list_mcp_tools", + "litellm.proxy._experimental.mcp_server.operations._list_mcp_tools", new=AsyncMock(side_effect=denial), ), ): @@ -1768,7 +1772,7 @@ async def test_mcp_server_tool_call_renders_denial_message_not_detail_dict(_mcp_ new=AsyncMock(return_value=(None, None, None, None, None, None, None)), ), patch( # test-quality-ok: the tool-call helper is the handler's only collaborator; the suite's seam - "litellm.proxy._experimental.mcp_server.server.call_mcp_tool", + "litellm.proxy._experimental.mcp_server.operations.call_mcp_tool", new=AsyncMock(side_effect=denial), ), ): @@ -1819,7 +1823,7 @@ async def test_mcp_server_tool_call_body_with_none_arguments(_mcp_request_ctx): mock_add_litellm_data_to_request, ): with patch( - "litellm.proxy._experimental.mcp_server.server.call_mcp_tool", + "litellm.proxy._experimental.mcp_server.operations.call_mcp_tool", mock_call_mcp_tool, ): with patch( @@ -1893,7 +1897,7 @@ async def test_concurrent_initialize_session_managers(): "run", return_value=mock_cm_sse, ) as mock_sse_run, - patch("litellm.proxy._experimental.mcp_server.server.verbose_logger"), + patch("litellm.proxy._experimental.mcp_server.operations.verbose_logger"), ): # Create multiple concurrent tasks that call initialize_session_managers async def init_task(): @@ -2333,7 +2337,7 @@ async def test_mcp_routing_chunked_initialize_to_stateful(): "litellm.proxy._experimental.mcp_server.server.set_auth_context", ), patch( # test-quality-ok: registry is empty in unit tests; key owns one server - "litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", + "litellm.proxy._experimental.mcp_server.operations._get_allowed_mcp_servers", new_callable=AsyncMock, return_value=[MagicMock()], ), @@ -2823,7 +2827,7 @@ async def test_initialize_request_tracks_active_session_after_response_header(): return_value=(owner_auth, None, None, None, None, None), ), patch( # test-quality-ok: registry is empty in unit tests; key owns one server - "litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", + "litellm.proxy._experimental.mcp_server.operations._get_allowed_mcp_servers", new_callable=AsyncMock, return_value=[MagicMock()], ), @@ -2976,7 +2980,7 @@ async def test_initialize_request_records_client_name_in_gateway_sessions_report return_value=(owner_auth, None, None, None, None, None), ), patch( # test-quality-ok: registry is empty in unit tests; key owns one server - "litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", + "litellm.proxy._experimental.mcp_server.operations._get_allowed_mcp_servers", new_callable=AsyncMock, return_value=[MagicMock()], ), @@ -3451,7 +3455,7 @@ async def test_initialize_request_with_existing_session_tracks_new_session(): ), ), patch( # test-quality-ok: registry is empty in unit tests; key owns one server - "litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", + "litellm.proxy._experimental.mcp_server.operations._get_allowed_mcp_servers", new_callable=AsyncMock, return_value=[MagicMock()], ), @@ -4164,7 +4168,7 @@ async def test_mcp_routing_with_conflicting_alias_and_group_name(): with ( patch( - "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager.get_allowed_mcp_servers", + "litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager.get_allowed_mcp_servers", mock_get_allowed, ), patch( @@ -4172,7 +4176,7 @@ async def test_mcp_routing_with_conflicting_alias_and_group_name(): mock_db_lookup, ), patch( - "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager._get_tools_from_server", + "litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager._get_tools_from_server", mock_get_tools_spy, ), ): @@ -4281,16 +4285,16 @@ async def test_oauth2_caller_headers_not_forwarded_for_migrated_server(): side_effect=mock_fetch_tools_with_timeout, ), patch( - "litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", + "litellm.proxy._experimental.mcp_server.operations._get_allowed_mcp_servers", AsyncMock(return_value=[oauth2_server]), ), patch( - "litellm.proxy._experimental.mcp_server.server._prefetch_oauth_creds_for_user", + "litellm.proxy._experimental.mcp_server.operations._prefetch_oauth_creds_for_user", new_callable=AsyncMock, return_value={}, ), patch( - "litellm.proxy._experimental.mcp_server.server._get_user_oauth_extra_headers_from_db", + "litellm.proxy._experimental.mcp_server.operations._get_user_oauth_extra_headers_from_db", new_callable=AsyncMock, return_value=None, ), @@ -4372,7 +4376,7 @@ async def test_list_tools_single_server_unprefixed_names(): mock_manager._get_tools_from_server = mock_get_tools_from_server with patch( - "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager", + "litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager", mock_manager, ): listing = await _get_tools_from_mcp_servers( @@ -4451,7 +4455,7 @@ async def test_list_tools_multiple_servers_prefixed_names(): mock_manager._get_tools_from_server = mock_get_tools_from_server with patch( - "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager", + "litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager", mock_manager, ): listing = await _get_tools_from_mcp_servers( @@ -4631,7 +4635,7 @@ async def test_call_mcp_tool_user_unauthorized_access(): AsyncMock(return_value=["allowed_server", "another_server"]), ), patch( - "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager.get_mcp_server_by_id", + "litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager.get_mcp_server_by_id", side_effect=mock_get_server_by_id, ), ): @@ -4661,11 +4665,11 @@ async def test_call_mcp_tool_scoped_denial_names_the_binding_agent(): with ( patch( # test-quality-ok: the server registry is a module-level singleton; the suite's only seam - "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager.get_allowed_mcp_servers", + "litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager.get_allowed_mcp_servers", AsyncMock(return_value=[]), ), patch( # test-quality-ok: the permission resolver is a module-level function; the suite's only seam - "litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", + "litellm.proxy._experimental.mcp_server.operations._get_allowed_mcp_servers", _scope_resolver({"github": "srv-github"}), ), ): @@ -4737,7 +4741,7 @@ async def test_call_mcp_tool_unauthorized_403_does_not_leak_server_credentials() AsyncMock(return_value=["allowed_server"]), ), patch( - "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager.get_mcp_server_by_id", + "litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager.get_mcp_server_by_id", side_effect=mock_get_server_by_id, ), ): @@ -4880,7 +4884,7 @@ async def test_list_tools_filters_by_key_team_permissions(): mock_manager._get_tools_from_server = mock_get_tools_from_server with patch( - "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager", + "litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager", mock_manager, ): listing = await _get_tools_from_mcp_servers( @@ -4991,7 +4995,7 @@ async def test_list_tools_with_team_tool_permissions_inheritance(): mock_manager._get_tools_from_server = mock_get_tools_from_server with patch( - "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager", + "litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager", mock_manager, ): # Mock the team object permission retrieval @@ -5083,7 +5087,7 @@ async def test_list_tools_with_no_tool_permissions_shows_all(): mock_manager._get_tools_from_server = mock_get_tools_from_server with patch( - "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager", + "litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager", mock_manager, ): listing = await _get_tools_from_mcp_servers( @@ -5189,7 +5193,7 @@ async def test_list_tools_strips_prefix_when_matching_permissions(): mock_manager._get_tools_from_server = mock_get_tools_from_server with patch( - "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager", + "litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager", mock_manager, ): listing = await _get_tools_from_mcp_servers( @@ -5631,12 +5635,12 @@ async def test_call_mcp_tool_logs_failure_via_post_call_failure_hook(): return_value=mock_server, ), patch( - "litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers_from_mcp_server_names", + "litellm.proxy._experimental.mcp_server.operations._get_allowed_mcp_servers_from_mcp_server_names", new_callable=AsyncMock, return_value=[mock_server], ), patch( - "litellm.proxy._experimental.mcp_server.server.execute_mcp_tool", + "litellm.proxy._experimental.mcp_server.operations.execute_mcp_tool", new_callable=AsyncMock, side_effect=Exception("boom"), ), @@ -5700,26 +5704,26 @@ async def test_get_tools_from_mcp_servers_logs_list_tools_to_spendlogs_when_enab with ( patch( - "litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", + "litellm.proxy._experimental.mcp_server.operations._get_allowed_mcp_servers", new=AsyncMock(return_value=[server_a]), ), patch( - "litellm.proxy._experimental.mcp_server.server._prepare_mcp_server_headers", + "litellm.proxy._experimental.mcp_server.operations._prepare_mcp_server_headers", return_value=(None, None), ), patch( - "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager", + "litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager", ) as mock_manager, patch( - "litellm.proxy._experimental.mcp_server.server.filter_tools_by_allowed_tools", + "litellm.proxy._experimental.mcp_server.operations.filter_tools_by_allowed_tools", side_effect=lambda tools, _server: tools, ), patch( - "litellm.proxy._experimental.mcp_server.server.filter_tools_by_key_team_permissions", + "litellm.proxy._experimental.mcp_server.operations.filter_tools_by_key_team_permissions", new=AsyncMock(side_effect=lambda tools, **_: tools), ), patch( - "litellm.proxy._experimental.mcp_server.server.function_setup", + "litellm.proxy._experimental.mcp_server.operations.function_setup", side_effect=_capture_function_setup, ), ): @@ -5782,26 +5786,26 @@ async def test_get_tools_from_mcp_servers_returns_tools_when_success_logging_fai with ( patch( - "litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", + "litellm.proxy._experimental.mcp_server.operations._get_allowed_mcp_servers", new=AsyncMock(return_value=[server_a]), ), patch( - "litellm.proxy._experimental.mcp_server.server._prepare_mcp_server_headers", + "litellm.proxy._experimental.mcp_server.operations._prepare_mcp_server_headers", return_value=(None, None), ), patch( - "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager", + "litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager", ) as mock_manager, patch( - "litellm.proxy._experimental.mcp_server.server.filter_tools_by_allowed_tools", + "litellm.proxy._experimental.mcp_server.operations.filter_tools_by_allowed_tools", side_effect=lambda tools, _server: tools, ), patch( - "litellm.proxy._experimental.mcp_server.server.filter_tools_by_key_team_permissions", + "litellm.proxy._experimental.mcp_server.operations.filter_tools_by_key_team_permissions", new=AsyncMock(side_effect=lambda tools, **_: tools), ), patch( - "litellm.proxy._experimental.mcp_server.server.function_setup", + "litellm.proxy._experimental.mcp_server.operations.function_setup", return_value=(dummy_logging_obj, None), ), ): @@ -6102,23 +6106,23 @@ async def test_get_tools_from_mcp_servers_injects_stored_oauth2_token(): with ( patch( - "litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", + "litellm.proxy._experimental.mcp_server.operations._get_allowed_mcp_servers", new=AsyncMock(return_value=[oauth2_server]), ), patch( # Patch the bulk prefetch so no real DB connection is needed - "litellm.proxy._experimental.mcp_server.server._prefetch_oauth_creds_for_user", + "litellm.proxy._experimental.mcp_server.operations._prefetch_oauth_creds_for_user", new=AsyncMock(return_value=prefetched_creds), ) as mock_prefetch, patch( - "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager", + "litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager", ) as mock_manager, patch( - "litellm.proxy._experimental.mcp_server.server.filter_tools_by_allowed_tools", + "litellm.proxy._experimental.mcp_server.operations.filter_tools_by_allowed_tools", side_effect=lambda tools, _server: tools, ), patch( - "litellm.proxy._experimental.mcp_server.server.filter_tools_by_key_team_permissions", + "litellm.proxy._experimental.mcp_server.operations.filter_tools_by_key_team_permissions", new=AsyncMock(side_effect=lambda tools, **_: tools), ), ): @@ -6450,7 +6454,7 @@ class TestGatewayCreateInitializationOptions: with ( patch( - "litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", + "litellm.proxy._experimental.mcp_server.operations._get_allowed_mcp_servers", new_callable=AsyncMock, return_value=[scoped_server], ), @@ -6480,7 +6484,7 @@ class TestGatewayCreateInitializationOptions: from litellm.proxy._types import UserAPIKeyAuth with patch( # test-quality-ok: grant resolution is the input under test - "litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", + "litellm.proxy._experimental.mcp_server.operations._get_allowed_mcp_servers", new_callable=AsyncMock, return_value=[], ): @@ -6506,7 +6510,7 @@ class TestGatewayCreateInitializationOptions: from litellm.proxy._types import UserAPIKeyAuth with patch( # test-quality-ok: grant resolution is the input under test - "litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", + "litellm.proxy._experimental.mcp_server.operations._get_allowed_mcp_servers", new_callable=AsyncMock, return_value=[], ): @@ -6531,7 +6535,7 @@ class TestGatewayCreateInitializationOptions: from litellm.proxy._types import UserAPIKeyAuth with patch( # test-quality-ok: grant resolution is the input under test - "litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", + "litellm.proxy._experimental.mcp_server.operations._get_allowed_mcp_servers", new_callable=AsyncMock, return_value=[], ): @@ -6587,7 +6591,7 @@ class TestGatewayCreateInitializationOptions: ), ), patch( - "litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", + "litellm.proxy._experimental.mcp_server.operations._get_allowed_mcp_servers", new_callable=AsyncMock, return_value=[scoped_server], ), @@ -6722,14 +6726,14 @@ async def test_list_tools_with_legacy_db_m2m_server_resolves_oauth2_flow(): with ( patch( - "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager", + "litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager", ) as mock_manager, patch( - "litellm.proxy._experimental.mcp_server.server.filter_tools_by_allowed_tools", + "litellm.proxy._experimental.mcp_server.operations.filter_tools_by_allowed_tools", side_effect=lambda tools, _server: tools, ), patch( - "litellm.proxy._experimental.mcp_server.server.filter_tools_by_key_team_permissions", + "litellm.proxy._experimental.mcp_server.operations.filter_tools_by_key_team_permissions", new=AsyncMock(side_effect=lambda tools, **_: tools), ), ): @@ -6992,7 +6996,7 @@ def _patch_delegate_resolver(server: MCPServer, *resolvable_names: str): return server if name in resolvable_names else None return patch( - "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager.get_mcp_server_by_name", + "litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager.get_mcp_server_by_name", side_effect=_resolve, ) @@ -7011,7 +7015,7 @@ async def test_legacy_delegate_bare_token_is_not_probed_upstream(): # test-qual with ( _patch_delegate_resolver(server, "delegate_test"), patch( - "litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", + "litellm.proxy._experimental.mcp_server.operations._get_allowed_mcp_servers", new=AsyncMock(return_value=[server]), ), patch( @@ -7047,7 +7051,7 @@ async def test_legacy_delegate_dual_credentials_are_not_probed_upstream(): # te with ( patch( # test-quality-ok: isolate authorized-server resolution so this test targets the preflight boundary - "litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", + "litellm.proxy._experimental.mcp_server.operations._get_allowed_mcp_servers", new=AsyncMock(return_value=[server]), ), patch( # test-quality-ok: the removed probe call is the security regression under test @@ -7094,7 +7098,7 @@ async def test_oauth_passthrough_preflight_preserves_status_contract(probe_statu with ( patch( # test-quality-ok: isolate authorized-server resolution so this test exercises the preflight contract - "litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", + "litellm.proxy._experimental.mcp_server.operations._get_allowed_mcp_servers", new=AsyncMock(return_value=[server]), ), patch( # test-quality-ok: the upstream transport boundary is the behavior being mapped to an HTTP response @@ -7140,7 +7144,7 @@ async def test_delegate_tokenless_request_not_probed(): with ( _patch_delegate_resolver(server, "delegate_test"), patch( - "litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", + "litellm.proxy._experimental.mcp_server.operations._get_allowed_mcp_servers", new=AsyncMock(return_value=[server]), ), patch( @@ -7173,7 +7177,7 @@ async def test_delegate_preflight_skipped_on_multi_server_routes(): with ( _patch_delegate_resolver(servers[0], "delegate_test", "other_server"), patch( - "litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", + "litellm.proxy._experimental.mcp_server.operations._get_allowed_mcp_servers", new=AsyncMock(return_value=servers), ), patch( @@ -7216,7 +7220,7 @@ async def test_bare_authorization_never_probes_passthrough_servers(): with ( _patch_delegate_resolver(passthrough_server, "pt_server"), patch( - "litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", + "litellm.proxy._experimental.mcp_server.operations._get_allowed_mcp_servers", new=AsyncMock(return_value=[passthrough_server]), ), patch( @@ -7262,7 +7266,7 @@ async def test_delegate_not_probed_when_named_only_via_server_id(): with ( _patch_delegate_resolver(server, "delegate_test"), patch( - "litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", + "litellm.proxy._experimental.mcp_server.operations._get_allowed_mcp_servers", new=AsyncMock(return_value=[server]), ), patch( @@ -7295,7 +7299,7 @@ async def test_delegate_probe_not_fanned_out_to_access_group_members(): with ( _patch_delegate_resolver(group_member, "delegate_test"), patch( - "litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", + "litellm.proxy._experimental.mcp_server.operations._get_allowed_mcp_servers", new=AsyncMock(return_value=[group_member]), ), patch( @@ -7391,11 +7395,11 @@ async def test_execute_mcp_tool_rest_server_id_authoritative_for_unprefixed_tool with ( patch.dict( - mcp_module.global_mcp_server_manager.tool_name_to_mcp_server_name_mapping, + mcp_operations.global_mcp_server_manager.tool_name_to_mcp_server_name_mapping, {"echo": oauth_server.name}, ), patch.object( - mcp_module.global_mcp_server_manager, + mcp_operations.global_mcp_server_manager, "get_registry", return_value={ api_key_server.server_id: api_key_server, @@ -7403,13 +7407,12 @@ async def test_execute_mcp_tool_rest_server_id_authoritative_for_unprefixed_tool }, ), patch.object( - mcp_module.global_mcp_server_manager, + mcp_operations.global_mcp_server_manager, "_get_mcp_server_from_tool_name", return_value=oauth_server, ), patch.object( - mcp_module, - "_handle_managed_mcp_tool", + mcp_operations, "_handle_managed_mcp_tool", new=fake_handle_managed_mcp_tool, ), patch.object( @@ -7418,12 +7421,12 @@ async def test_execute_mcp_tool_rest_server_id_authoritative_for_unprefixed_tool return_value=True, ), patch.object( - mcp_module.global_mcp_tool_registry, + mcp_operations.global_mcp_tool_registry, "get_tool", return_value=None, ), ): - await mcp_module.execute_mcp_tool( + await mcp_operations.execute_mcp_tool( name="echo", arguments={"message": "hello"}, allowed_mcp_servers=[api_key_server, oauth_server], @@ -7456,7 +7459,7 @@ def _worker_that_never_listed(server: MCPServer, upstream_tools: tuple[str, ...] from litellm.proxy._experimental.mcp_server import server as mcp_module - mcp_module.global_mcp_server_manager.registry[server.server_id] = server + mcp_operations.global_mcp_server_manager.registry[server.server_id] = server dispatched: dict[str, object] = {} async def fake_handle_managed_mcp_tool(**kwargs): @@ -7468,17 +7471,17 @@ def _worker_that_never_listed(server: MCPServer, upstream_tools: tuple[str, ...] with ( patch.object( # test-quality-ok: the upstream MCP session is the boundary; a real one needs an initialize handshake over a live server - mcp_module.global_mcp_server_manager, + mcp_operations.global_mcp_server_manager, "_create_mcp_client", new=AsyncMock(return_value=MagicMock()), ) as create_client, patch.object( # test-quality-ok: same boundary, this is the tools/list answer the upstream would give - mcp_module.global_mcp_server_manager, + mcp_operations.global_mcp_server_manager, "_fetch_tools_with_timeout", side_effect=fake_fetch_tools, ) as fetch_tools, patch.object( # test-quality-ok: records the resolved server and bare name the managed call would forward upstream - mcp_module, "_handle_managed_mcp_tool", new=fake_handle_managed_mcp_tool + mcp_operations, "_handle_managed_mcp_tool", new=fake_handle_managed_mcp_tool ), ): yield SimpleNamespace(create_client=create_client, fetch_tools=fetch_tools, dispatched=dispatched) @@ -7492,7 +7495,7 @@ async def test_execute_mcp_tool_lists_never_listed_passthrough_server_with_calle server = _never_listed_passthrough_server() with _worker_that_never_listed(server, upstream_tools=("add",)) as worker: - await mcp_module.execute_mcp_tool( + await mcp_operations.execute_mcp_tool( name="lazy_map-add", arguments={"a": 1, "b": 2}, allowed_mcp_servers=[server], @@ -7513,7 +7516,7 @@ async def test_execute_mcp_tool_rest_server_id_lists_never_listed_server_first() server = _never_listed_passthrough_server() with _worker_that_never_listed(server, upstream_tools=("add",)) as worker: - await mcp_module.execute_mcp_tool( + await mcp_operations.execute_mcp_tool( name="add", arguments={"a": 1, "b": 2}, allowed_mcp_servers=[server], @@ -7536,7 +7539,7 @@ async def test_execute_mcp_tool_unknown_tool_on_never_listed_server_lists_once_t _worker_that_never_listed(server, upstream_tools=("add",)) as worker, pytest.raises(HTTPException) as exc_info, ): - await mcp_module.execute_mcp_tool( + await mcp_operations.execute_mcp_tool( name="lazy_map-nope", arguments={}, allowed_mcp_servers=[server], @@ -7557,8 +7560,8 @@ async def test_execute_mcp_tool_does_not_relist_a_server_this_worker_already_lis server = _never_listed_passthrough_server() with _worker_that_never_listed(server, upstream_tools=("add",)) as worker: - mcp_module.global_mcp_server_manager._create_prefixed_tools([MCPTool(name="add", inputSchema={})], server) - await mcp_module.execute_mcp_tool( + mcp_operations.global_mcp_server_manager._create_prefixed_tools([MCPTool(name="add", inputSchema={})], server) + await mcp_operations.execute_mcp_tool( name="lazy_map-add", arguments={"a": 1, "b": 2}, allowed_mcp_servers=[server], @@ -7580,8 +7583,8 @@ async def test_execute_mcp_tool_lists_a_tool_this_worker_has_not_yet_seen_on_a_l server = _never_listed_passthrough_server() with _worker_that_never_listed(server, upstream_tools=("add", "multiply")) as worker: - mcp_module.global_mcp_server_manager._create_prefixed_tools([MCPTool(name="add", inputSchema={})], server) - await mcp_module.execute_mcp_tool( + mcp_operations.global_mcp_server_manager._create_prefixed_tools([MCPTool(name="add", inputSchema={})], server) + await mcp_operations.execute_mcp_tool( name="lazy_map-multiply", arguments={"a": 1, "b": 2}, allowed_mcp_servers=[server], @@ -7604,7 +7607,7 @@ async def test_execute_mcp_tool_never_lists_a_server_the_caller_cannot_access(): _worker_that_never_listed(server, upstream_tools=("add",)) as worker, pytest.raises(HTTPException) as exc_info, ): - await mcp_module.execute_mcp_tool( + await mcp_operations.execute_mcp_tool( name="lazy_map-add", arguments={"a": 1, "b": 2}, allowed_mcp_servers=[other_server], @@ -7650,13 +7653,12 @@ async def test_execute_mcp_tool_strips_a_prefix_that_contains_the_separator(): with ( patch.object( - mcp_module.global_mcp_server_manager, + mcp_operations.global_mcp_server_manager, "_get_mcp_server_from_tool_name", return_value=alias_less_server, ), patch.object( - mcp_module, - "_handle_managed_mcp_tool", + mcp_operations, "_handle_managed_mcp_tool", new=fake_handle_managed_mcp_tool, ), patch.object( @@ -7665,12 +7667,12 @@ async def test_execute_mcp_tool_strips_a_prefix_that_contains_the_separator(): return_value=True, ), patch.object( - mcp_module.global_mcp_tool_registry, + mcp_operations.global_mcp_tool_registry, "get_tool", return_value=None, ), ): - await mcp_module.execute_mcp_tool( + await mcp_operations.execute_mcp_tool( name=f"{server_id}-read_wiki_contents", arguments={"repoName": "acme/wiki"}, allowed_mcp_servers=[alias_less_server], @@ -7724,11 +7726,11 @@ async def test_execute_mcp_tool_rest_server_id_injects_requested_server_credenti with ( patch.dict( - mcp_module.global_mcp_server_manager.tool_name_to_mcp_server_name_mapping, + mcp_operations.global_mcp_server_manager.tool_name_to_mcp_server_name_mapping, {"echo": collision_server.name, "echo_requested-echo": requested_server.name}, ), patch.object( - mcp_module.global_mcp_server_manager, + mcp_operations.global_mcp_server_manager, "get_registry", return_value={ requested_server.server_id: requested_server, @@ -7736,7 +7738,7 @@ async def test_execute_mcp_tool_rest_server_id_injects_requested_server_credenti }, ), patch.object( - mcp_module.global_mcp_server_manager, + mcp_operations.global_mcp_server_manager, "_create_mcp_client", new=fake_create_mcp_client, ), @@ -7746,13 +7748,13 @@ async def test_execute_mcp_tool_rest_server_id_injects_requested_server_credenti return_value=True, ), patch.object( - mcp_module.global_mcp_tool_registry, + mcp_operations.global_mcp_tool_registry, "get_tool", return_value=None, ), patch("litellm.proxy.proxy_server.proxy_logging_obj", None), ): - await mcp_module.execute_mcp_tool( + await mcp_operations.execute_mcp_tool( name="echo", arguments={"message": "hello"}, allowed_mcp_servers=[requested_server, collision_server], @@ -7795,7 +7797,7 @@ async def test_execute_mcp_tool_rest_prefixed_tool_still_validates_server_id(): with ( patch.object( - mcp_module.global_mcp_server_manager, + mcp_operations.global_mcp_server_manager, "get_registry", return_value={ api_key_server.server_id: api_key_server, @@ -7803,7 +7805,7 @@ async def test_execute_mcp_tool_rest_prefixed_tool_still_validates_server_id(): }, ), patch.object( - mcp_module.global_mcp_server_manager, + mcp_operations.global_mcp_server_manager, "_get_mcp_server_from_tool_name", return_value=oauth_server, ), @@ -7813,13 +7815,13 @@ async def test_execute_mcp_tool_rest_prefixed_tool_still_validates_server_id(): return_value=True, ), patch.object( - mcp_module.global_mcp_tool_registry, + mcp_operations.global_mcp_tool_registry, "get_tool", return_value=None, ), pytest.raises(HTTPException) as exc_info, ): - await mcp_module.execute_mcp_tool( + await mcp_operations.execute_mcp_tool( name="echo_oauth_m2m-echo", arguments={"message": "hello"}, allowed_mcp_servers=[api_key_server, oauth_server], @@ -7857,7 +7859,7 @@ async def test_execute_mcp_tool_rest_unauthorized_prefix_still_mismatches(): with ( patch.object( - mcp_module.global_mcp_server_manager, + mcp_operations.global_mcp_server_manager, "get_registry", return_value={ api_key_server.server_id: api_key_server, @@ -7865,7 +7867,7 @@ async def test_execute_mcp_tool_rest_unauthorized_prefix_still_mismatches(): }, ), patch.object( - mcp_module.global_mcp_server_manager, + mcp_operations.global_mcp_server_manager, "_get_mcp_server_from_tool_name", return_value=restricted_server, ), @@ -7875,13 +7877,13 @@ async def test_execute_mcp_tool_rest_unauthorized_prefix_still_mismatches(): return_value=True, ), patch.object( - mcp_module.global_mcp_tool_registry, + mcp_operations.global_mcp_tool_registry, "get_tool", return_value=None, ), pytest.raises(HTTPException) as exc_info, ): - await mcp_module.execute_mcp_tool( + await mcp_operations.execute_mcp_tool( name="restricted_server-echo", arguments={"message": "hello"}, allowed_mcp_servers=[api_key_server], @@ -7921,18 +7923,17 @@ async def test_execute_mcp_tool_rest_hyphenated_upstream_tool_name_routes_to_req with ( patch.object( - mcp_module.global_mcp_server_manager, + mcp_operations.global_mcp_server_manager, "get_registry", return_value={api_key_server.server_id: api_key_server}, ), patch.object( - mcp_module.global_mcp_server_manager, + mcp_operations.global_mcp_server_manager, "_get_mcp_server_from_tool_name", return_value=None, ), patch.object( - mcp_module, - "_handle_managed_mcp_tool", + mcp_operations, "_handle_managed_mcp_tool", new=fake_handle_managed_mcp_tool, ), patch.object( @@ -7941,12 +7942,12 @@ async def test_execute_mcp_tool_rest_hyphenated_upstream_tool_name_routes_to_req return_value=True, ), patch.object( - mcp_module.global_mcp_tool_registry, + mcp_operations.global_mcp_tool_registry, "get_tool", return_value=None, ), ): - await mcp_module.execute_mcp_tool( + await mcp_operations.execute_mcp_tool( name="text-to-speech", arguments={"message": "hello"}, allowed_mcp_servers=[api_key_server], @@ -8005,22 +8006,22 @@ async def test_execute_mcp_tool_sets_model_in_model_call_details(): with ( patch.object( - mcp_module.global_mcp_server_manager, + mcp_operations.global_mcp_server_manager, "_get_mcp_server_from_tool_name", return_value=fake_server, ), patch.object( - mcp_module.global_mcp_server_manager, + mcp_operations.global_mcp_server_manager, "pre_call_tool_check", new=AsyncMock(return_value={}), ), patch.object( - mcp_module.global_mcp_tool_registry, + mcp_operations.global_mcp_tool_registry, "get_tool", return_value=fake_tool, ), patch( - "litellm.proxy._experimental.mcp_server.server._handle_local_mcp_tool", + "litellm.proxy._experimental.mcp_server.operations._handle_local_mcp_tool", new=AsyncMock(return_value=[]), ), patch( @@ -8028,7 +8029,7 @@ async def test_execute_mcp_tool_sets_model_in_model_call_details(): return_value=True, ), ): - await mcp_module.execute_mcp_tool( + await mcp_operations.execute_mcp_tool( name="list_pets", arguments={"limit": 10}, allowed_mcp_servers=[fake_server], @@ -8084,7 +8085,7 @@ async def test_execute_mcp_tool_rest_unresolved_prefixed_name_routes_to_requeste with ( patch.object( - mcp_module.global_mcp_server_manager, + mcp_operations.global_mcp_server_manager, "get_registry", return_value={ requested_server.server_id: requested_server, @@ -8092,13 +8093,12 @@ async def test_execute_mcp_tool_rest_unresolved_prefixed_name_routes_to_requeste }, ), patch.object( - mcp_module.global_mcp_server_manager, + mcp_operations.global_mcp_server_manager, "_get_mcp_server_from_tool_name", return_value=None, ), patch.object( - mcp_module, - "_handle_managed_mcp_tool", + mcp_operations, "_handle_managed_mcp_tool", new=fake_handle_managed_mcp_tool, ), patch.object( @@ -8107,12 +8107,12 @@ async def test_execute_mcp_tool_rest_unresolved_prefixed_name_routes_to_requeste return_value=True, ), patch.object( - mcp_module.global_mcp_tool_registry, + mcp_operations.global_mcp_tool_registry, "get_tool", return_value=None, ), ): - await mcp_module.execute_mcp_tool( + await mcp_operations.execute_mcp_tool( name="known_prefix-list_things", arguments={"message": "hello"}, allowed_mcp_servers=[requested_server, prefix_owner], @@ -8164,7 +8164,7 @@ async def test_execute_mcp_tool_rest_prefix_retry_resolution_still_enforces_serv with ( patch.object( - mcp_module.global_mcp_server_manager, + mcp_operations.global_mcp_server_manager, "get_registry", return_value={ requested_server.server_id: requested_server, @@ -8172,7 +8172,7 @@ async def test_execute_mcp_tool_rest_prefix_retry_resolution_still_enforces_serv }, ), patch.object( - mcp_module.global_mcp_server_manager, + mcp_operations.global_mcp_server_manager, "_get_mcp_server_from_tool_name", side_effect=resolve_only_when_requested_prefix_added, ), @@ -8182,13 +8182,13 @@ async def test_execute_mcp_tool_rest_prefix_retry_resolution_still_enforces_serv return_value=True, ), patch.object( - mcp_module.global_mcp_tool_registry, + mcp_operations.global_mcp_tool_registry, "get_tool", return_value=None, ), pytest.raises(HTTPException) as exc_info, ): - await mcp_module.execute_mcp_tool( + await mcp_operations.execute_mcp_tool( name="known_prefix-echo", arguments={"message": "hello"}, allowed_mcp_servers=[requested_server, prefix_owner], @@ -9091,14 +9091,14 @@ async def test_call_tool_with_legacy_db_m2m_server_resolves_oauth2_flow(): with ( patch( - "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager", + "litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager", ) as mock_manager, patch( - "litellm.proxy._experimental.mcp_server.server.execute_mcp_tool", + "litellm.proxy._experimental.mcp_server.operations.execute_mcp_tool", side_effect=capture_execute, ), patch( - "litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers_from_mcp_server_names", + "litellm.proxy._experimental.mcp_server.operations._get_allowed_mcp_servers_from_mcp_server_names", new=AsyncMock(side_effect=lambda mcp_servers, allowed_mcp_servers: allowed_mcp_servers), ), ): @@ -9176,12 +9176,12 @@ async def test_call_mcp_tool_skips_failure_hook_for_upstream_auth_error(): ), patch.object(global_mcp_server_manager, "get_mcp_server_by_id", return_value=mock_server), patch( - "litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers_from_mcp_server_names", + "litellm.proxy._experimental.mcp_server.operations._get_allowed_mcp_servers_from_mcp_server_names", new_callable=AsyncMock, return_value=[mock_server], ), patch( - "litellm.proxy._experimental.mcp_server.server.execute_mcp_tool", + "litellm.proxy._experimental.mcp_server.operations.execute_mcp_tool", new_callable=AsyncMock, side_effect=MCPUpstreamAuthError(status_code=401, www_authenticate="Bearer", server_name="test_server"), ), @@ -9261,7 +9261,7 @@ async def test_aggregate_listing_reports_per_server_outcomes(): mock_manager._get_tools_from_server = mock_get_tools_from_server with patch( - "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager", + "litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager", mock_manager, ): listing = await _get_tools_from_mcp_servers( @@ -9335,7 +9335,7 @@ async def test_handle_list_tools_attaches_outcome_meta(_mcp_request_ctx): new=AsyncMock(return_value=(None, None, None, None, None, None, None)), ), patch( - "litellm.proxy._experimental.mcp_server.server._list_mcp_tools", + "litellm.proxy._experimental.mcp_server.operations._list_mcp_tools", new=AsyncMock(return_value=listing), ), ): @@ -9401,12 +9401,12 @@ class TestPreemptive401ModeAware: with ( patch.object( - server_module.global_mcp_server_manager, + mcp_operations.global_mcp_server_manager, "get_mcp_server_by_name", return_value=server, ), patch.object( - server_module.global_mcp_server_manager, + mcp_operations.global_mcp_server_manager, "has_user_oauth_token", new_callable=AsyncMock, return_value=has_stored_token, @@ -9425,7 +9425,7 @@ class TestPreemptive401ModeAware: async def test_deferred_discovery_runs_before_delegate_challenge(self): from litellm.proxy._experimental.mcp_server import server as server_module - manager = server_module.global_mcp_server_manager + manager = mcp_operations.global_mcp_server_manager server = _make_oauth2_server( "lazy_delegate", oauth2_flow="authorization_code", @@ -9457,7 +9457,7 @@ class TestPreemptive401ModeAware: async def test_stamped_m2m_challenge_skips_deferred_discovery(self): from litellm.proxy._experimental.mcp_server import server as server_module - manager = server_module.global_mcp_server_manager + manager = mcp_operations.global_mcp_server_manager server = _make_oauth2_server("stamped_m2m", oauth2_flow="client_credentials") with patch.object( @@ -9500,12 +9500,12 @@ class TestPreemptive401ModeAware: with ( patch.dict(os.environ, {"SERVER_ROOT_PATH": "/litellm"}), patch.object( - server_module.global_mcp_server_manager, + mcp_operations.global_mcp_server_manager, "get_mcp_server_by_name", return_value=server, ), patch.object( - server_module.global_mcp_server_manager, + mcp_operations.global_mcp_server_manager, "has_user_oauth_token", new_callable=AsyncMock, return_value=False, @@ -9607,17 +9607,17 @@ class TestSingleServerPreflightReachesIdJag: with ( patch.object( # test-quality-ok: route wiring must use the manager's configured server - server_module.global_mcp_server_manager, + mcp_operations.global_mcp_server_manager, "get_mcp_server_by_name", return_value=server, ), patch.object( # test-quality-ok: route wiring must invoke the manager preflight - server_module.global_mcp_server_manager, + mcp_operations.global_mcp_server_manager, "preflight_token_exchange", preflight, ), patch.object( # test-quality-ok: allowed-set resolution needs the DB; the test controls its answer - server_module, "_get_allowed_mcp_servers", AsyncMock(return_value=[server]) + mcp_operations, "_get_allowed_mcp_servers", AsyncMock(return_value=[server]) ), ): await server_module._raise_preemptive_401_for_unauthenticated_servers( @@ -9667,12 +9667,12 @@ class TestSingleServerPreflightReachesIdJag: with ( patch.object( # test-quality-ok: route wiring must use the manager's configured server - server_module.global_mcp_server_manager, + mcp_operations.global_mcp_server_manager, "get_mcp_server_by_name", return_value=token_exchange, ), patch.object( # test-quality-ok: route wiring must invoke the manager preflight - server_module.global_mcp_server_manager, + mcp_operations.global_mcp_server_manager, "preflight_token_exchange", preflight, ), @@ -9733,13 +9733,13 @@ class TestOboPreflightScopedToAllowedServers: preflight = AsyncMock() with ( patch.object( # test-quality-ok: route handler reads the module-level manager, no injection seam - server_module.global_mcp_server_manager, "get_mcp_server_by_name", return_value=requested + mcp_operations.global_mcp_server_manager, "get_mcp_server_by_name", return_value=requested ), patch.object( # test-quality-ok: the exchanger is the observable; a real one would call an IdP - server_module.global_mcp_server_manager, "preflight_token_exchange", preflight + mcp_operations.global_mcp_server_manager, "preflight_token_exchange", preflight ), patch.object( # test-quality-ok: allowed-set resolution needs the DB; the test controls its answer - server_module, "_get_allowed_mcp_servers", allowed_lookup + mcp_operations, "_get_allowed_mcp_servers", allowed_lookup ), ): await server_module._raise_preemptive_401_for_unauthenticated_servers( @@ -10048,7 +10048,7 @@ class TestListFiltersHonorThePrefixBoundary: with ( patch.object(MCPRequestHandler, "get_allowed_tools_for_server", AsyncMock(return_value=grants)), - patch("litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager") as mock_manager, + patch("litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager") as mock_manager, ): mock_manager.get_mcp_server_by_id.return_value = server @@ -10111,11 +10111,11 @@ async def test_list_tools_injects_byok_credential_for_non_oauth2_auth_types(auth with ( patch( - "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager", + "litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager", mock_manager, ), patch( - "litellm.proxy._experimental.mcp_server.server._get_byok_credential", + "litellm.proxy._experimental.mcp_server.operations._get_byok_credential", AsyncMock(return_value="personal-api-key"), ), ): @@ -10208,3 +10208,44 @@ async def test_streamable_http_rejects_modern_protocol_version(header_value: str assert header_value in body["error"]["message"] for version in body["error"]["message"].split("supported: ")[1].split(", "): assert version in HANDSHAKE_PROTOCOL_VERSIONS + + +@pytest.mark.asyncio +@pytest.mark.parametrize("handler_name,field", [ + ("handle_list_tools", "tools"), + ("list_prompts", "prompts"), + ("list_resources", "resources"), + ("list_resource_templates", "resource_templates"), +]) +async def test_native_listing_preserves_empty_result_on_auth_failure(_mcp_request_ctx, handler_name, field): + from litellm.proxy._experimental.mcp_server import server + + with patch.object(server, "get_or_extract_auth_context", AsyncMock(side_effect=RuntimeError("auth failure"))): + result = await getattr(server, handler_name)(_mcp_request_ctx(), _paged_params()) + assert getattr(result, field) == [] + + +@pytest.mark.asyncio +@pytest.mark.parametrize("failure_hook_raises", [False, True]) +async def test_tool_listing_preserves_permission_denial_when_failure_logging_fails(failure_hook_raises): + from litellm.proxy._experimental.mcp_server import operations + from litellm.proxy import proxy_server + + auth = UserAPIKeyAuth(user_id="denied-caller") + denial = HTTPException(status_code=403, detail="scope denied") + logger = MagicMock() + logger.post_call_failure_hook = AsyncMock(side_effect=RuntimeError("log unavailable") if failure_hook_raises else None) + upstream = AsyncMock() + with ( + patch.object(operations, "_get_allowed_mcp_servers", AsyncMock(side_effect=denial)), + patch.object(operations, "function_setup", return_value=(None, None)), + patch.object(proxy_server, "proxy_logging_obj", logger), + patch.object(operations.global_mcp_server_manager, "_get_tools_from_server", upstream), + ): + with pytest.raises(HTTPException) as rejected: + await operations._get_tools_from_mcp_servers(user_api_key_auth=auth, mcp_auth_header=None, mcp_servers=["catalog"], log_list_tools_to_spendlogs=True) + assert rejected.value is denial + upstream.assert_not_awaited() + logger.post_call_failure_hook.assert_awaited_once() + assert logger.post_call_failure_hook.await_args.kwargs["original_exception"] is denial + assert logger.post_call_failure_hook.await_args.kwargs["user_api_key_dict"] == auth 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 9140ac61f1a..7be79f9b514 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 @@ -71,6 +71,135 @@ from litellm.proxy.utils import ProxyLogging from litellm.types.guardrails import GuardrailEventHooks +@pytest.mark.asyncio +async def test_manager_sampling_preserves_explicit_headers_without_ambient_context(): + from litellm.proxy._experimental.mcp_server import server as legacy_server + + caller = UserAPIKeyAuth(user_id="sampling-caller") + upstream = MCPServer( + server_id="sampling-context", + name="sampling_context", + url="https://example.invalid/mcp", + transport=MCPTransport.http, + allow_sampling=True, + ) + sampling = AsyncMock() + client = MagicMock() + client.call_tool = AsyncMock(return_value=CallToolResult(content=[])) + assert legacy_server.get_active_auth_context() is None + with ( + patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPClient", return_value=client) as factory, + patch("litellm.proxy._experimental.mcp_server.sampling_handler.handle_sampling_create_message", sampling), + ): + await MCPServerManager()._call_regular_mcp_tool( + mcp_server=upstream, + original_tool_name="probe", + arguments={}, + tasks=[], + mcp_auth_header=None, + mcp_server_auth_headers=None, + oauth2_headers=None, + raw_headers={"x-test-caller": "sampling-caller"}, + proxy_logging_obj=None, + user_api_key_auth=caller, + ) + callback = factory.call_args.kwargs["sampling_callback"] + await callback(None, None) + assert sampling.await_args.kwargs["user_api_key_auth"].user_id == "sampling-caller" + assert sampling.await_args.kwargs["raw_headers"] == {"x-test-caller": "sampling-caller"} + + + +@pytest.mark.asyncio +async def test_sampling_callback_keeps_creation_context_after_caller_switch(): + from mcp.server.auth.middleware.auth_context import auth_context_var + + from litellm.proxy._experimental.mcp_server import server as legacy_server + from litellm.proxy._experimental.mcp_server.mcp_server_manager import _create_sampling_callback + + token = auth_context_var.set(None) + recorder = AsyncMock() + try: + original = UserAPIKeyAuth(user_id="alpha", models=["alpha-model"]) + original.mcp_admitted_user_subject = True + headers = {"x-caller": "alpha"} + legacy_server.set_auth_context(original, raw_headers=headers, client_ip="192.0.2.1") + callback = _create_sampling_callback() + original.models.append("bravo-model") + headers["x-caller"] = "bravo" + legacy_server.set_auth_context(UserAPIKeyAuth(user_id="bravo"), raw_headers={"x-caller": "bravo"}) + with patch("litellm.proxy._experimental.mcp_server.sampling_handler.handle_sampling_create_message", recorder): + await callback(None, None) + observed = recorder.await_args.kwargs + assert observed["user_api_key_auth"].user_id == "alpha" + assert observed["user_api_key_auth"].models == ["alpha-model"] + assert observed["user_api_key_auth"].mcp_admitted_user_subject is True + assert observed["raw_headers"] == {"x-caller": "alpha"} + assert observed["client_ip"] == "192.0.2.1" + finally: + auth_context_var.reset(token) + + +@pytest.mark.asyncio +async def test_elicitation_callback_keeps_initiating_session(): + from litellm.proxy._experimental.mcp_server import server as legacy_server + from litellm.proxy._experimental.mcp_server.mcp_server_manager import _create_elicitation_callback + + initiating = MagicMock() + replacement = MagicMock() + recorder = AsyncMock() + token = legacy_server.active_mcp_session_var.set(initiating) + try: + callback = _create_elicitation_callback() + legacy_server.active_mcp_session_var.set(replacement) + with patch("litellm.proxy._experimental.mcp_server.elicitation_handler.handle_elicitation_request", recorder): + await callback(None, None) + assert recorder.await_args.kwargs["downstream_session"] is initiating + assert recorder.await_args.kwargs["downstream_capabilities"] is initiating.capabilities + finally: + legacy_server.active_mcp_session_var.reset(token) + + +@pytest.mark.asyncio +async def test_sampling_callbacks_isolate_callers_and_cancellation(): + from mcp.types import ErrorData + + from litellm.proxy._experimental.mcp_server.mcp_server_manager import _create_sampling_callback + + started = asyncio.Event() + cancelled = asyncio.Event() + observed = {} + + async def record_sampling(*, user_api_key_auth, raw_headers, **kwargs): + label = user_api_key_auth.user_id + if label == "cancelled": + started.set() + try: + await asyncio.Event().wait() + except asyncio.CancelledError: + cancelled.set() + raise + await asyncio.sleep(0) + observed[label] = raw_headers["x-caller"] + return ErrorData(code=-1, message=label) + + callbacks = tuple( + _create_sampling_callback(UserAPIKeyAuth(user_id=label), raw_headers={"x-caller": label}) + for label in ("alpha", "bravo", "cancelled") + ) + with patch( + "litellm.proxy._experimental.mcp_server.sampling_handler.handle_sampling_create_message", record_sampling + ): + tasks = tuple(asyncio.create_task(callback(None, None)) for callback in callbacks) + await asyncio.wait_for(started.wait(), timeout=2) + tasks[2].cancel() + results = await asyncio.gather(*tasks, return_exceptions=True) + assert observed == {"alpha": "alpha", "bravo": "bravo"} + assert [result.message for result in results[:2]] == ["alpha", "bravo"] + assert isinstance(results[2], asyncio.CancelledError) + assert cancelled.is_set() + + def _reload_mcp_manager_module(): utils_module = sys.modules["litellm.proxy._experimental.mcp_server.utils"] manager_module = sys.modules["litellm.proxy._experimental.mcp_server.mcp_server_manager"] @@ -82,6 +211,9 @@ def _reload_mcp_manager_module(): server_module = sys.modules.get("litellm.proxy._experimental.mcp_server.server") if server_module is not None and hasattr(server_module, "global_mcp_server_manager"): server_module.global_mcp_server_manager = reloaded.global_mcp_server_manager + operations_module = sys.modules.get("litellm.proxy._experimental.mcp_server.operations") + if operations_module is not None: + operations_module.global_mcp_server_manager = reloaded.global_mcp_server_manager return reloaded @@ -3921,6 +4053,7 @@ class TestMCPServerManager: result = await manager.get_resource_templates_from_server( server=server, user_api_key_auth=None, + raw_headers=None, mcp_auth_header="auth", extra_headers=None, add_prefix=False, @@ -3933,6 +4066,8 @@ class TestMCPServerManager: stdio_env=None, subject_token=None, user_api_key_auth=None, + raw_headers=None, + client_ip=None, ) mock_client.list_resource_templates.assert_awaited_once() assert result == expected_templates @@ -5847,7 +5982,7 @@ class TestMCPServerManager: stored = {"Authorization": "Bearer stored-user-token"} with patch( - "litellm.proxy._experimental.mcp_server.server._get_user_oauth_extra_headers_from_db", + "litellm.proxy._experimental.mcp_server.operations._get_user_oauth_extra_headers_from_db", new=AsyncMock(return_value=stored), ) as mock_lookup: result = await manager._resolve_oauth2_headers_for_tool_call( @@ -5874,7 +6009,7 @@ class TestMCPServerManager: user_auth = UserAPIKeyAuth(api_key="sk-test", user_id="alice") with patch( - "litellm.proxy._experimental.mcp_server.server._get_user_oauth_extra_headers_from_db", + "litellm.proxy._experimental.mcp_server.operations._get_user_oauth_extra_headers_from_db", new=AsyncMock(return_value={"Authorization": "Bearer should-not-be-used"}), ) as mock_lookup: result = await manager._resolve_oauth2_headers_for_tool_call( @@ -5900,7 +6035,7 @@ class TestMCPServerManager: user_auth = UserAPIKeyAuth(api_key="sk-test", user_id="alice") with patch( - "litellm.proxy._experimental.mcp_server.server._get_user_oauth_extra_headers_from_db", + "litellm.proxy._experimental.mcp_server.operations._get_user_oauth_extra_headers_from_db", new=AsyncMock(side_effect=RuntimeError("redis down")), ): result = await manager._resolve_oauth2_headers_for_tool_call( @@ -6056,7 +6191,7 @@ class TestMCPServerManager: user_auth = UserAPIKeyAuth(api_key="sk-test") with patch( - "litellm.proxy._experimental.mcp_server.server._get_user_oauth_extra_headers_from_db", + "litellm.proxy._experimental.mcp_server.operations._get_user_oauth_extra_headers_from_db", new=AsyncMock(return_value={"Authorization": "Bearer x"}), ) as mock_lookup: result = await manager._resolve_oauth2_headers_for_tool_call( @@ -6860,7 +6995,8 @@ class TestMCPServerManager: } user_api_key_auth = UserAPIKeyAuth(api_key="sk-test", user_id="user-123") - token = _mcp_active_toolset_id.set("toolset-abc") + user_api_key_auth.mcp_toolset_id = "toolset-abc" + token = _mcp_active_toolset_id.set("unrelated-ambient-toolset") try: with ( patch.object(proxy_server_module, "user_api_key_cache", cache), @@ -14075,3 +14211,31 @@ async def test_request_selected_during_guardrail_runs_concurrently_with_tool(mon assert guardrail_started.is_set() is selected assert result.is_error is False assert result.content[0].text == "executed" + + +@pytest.mark.asyncio +@pytest.mark.parametrize("with_caller", [True, False]) +async def test_client_sampling_does_not_fill_explicit_context_from_another_ambient_caller(with_caller): + from mcp.server.auth.middleware.auth_context import auth_context_var + from litellm.proxy._experimental.mcp_server import server as legacy_server + + upstream = MCPServer(server_id="explicit-empty", name="explicit_empty", url="https://example.invalid/mcp", transport=MCPTransport.http, allow_sampling=True) + token = auth_context_var.set(None) + sampling = AsyncMock() + try: + legacy_server.set_auth_context(UserAPIKeyAuth(user_id="unrelated"), raw_headers={"authorization": "unrelated-credential"}, client_ip="192.0.2.99") + with ( + patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPClient") as factory, + patch("litellm.proxy._experimental.mcp_server.sampling_handler.handle_sampling_create_message", sampling), + ): + await MCPServerManager()._create_mcp_client(upstream, user_api_key_auth=UserAPIKeyAuth(user_id="explicit") if with_caller else None) + await factory.call_args.kwargs["sampling_callback"](None, None) + captured = sampling.await_args.kwargs + if with_caller: + assert captured["user_api_key_auth"].user_id == "explicit" + else: + assert captured["user_api_key_auth"] is None + assert captured["raw_headers"] is None + assert captured["client_ip"] is None + finally: + auth_context_var.reset(token) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_stale_session.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_stale_session.py index 9420eecd222..ec6fdef69ee 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_stale_session.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_stale_session.py @@ -639,12 +639,12 @@ async def test_per_user_oauth_missing_stored_token_returns_preemptive_401(): return_value=False, ), patch( - "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager.has_user_oauth_token", + "litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager.has_user_oauth_token", new_callable=AsyncMock, return_value=False, ) as mock_has_token, patch( - "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager.get_mcp_server_by_name", + "litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager.get_mcp_server_by_name", return_value=oauth_server, ), patch.object( @@ -727,12 +727,12 @@ async def test_admitted_subject_missing_stored_token_challenged_with_resource_me return_value=False, ), patch( - "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager.has_user_oauth_token", + "litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager.has_user_oauth_token", new_callable=AsyncMock, return_value=False, ) as mock_has_token, patch( - "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager.get_mcp_server_by_name", + "litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager.get_mcp_server_by_name", return_value=oauth_server, ), patch.object( @@ -833,11 +833,11 @@ async def test_client_credentials_server_is_not_preemptively_challenged(m2m_fiel return_value=False, ), patch( - "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager.has_user_oauth_token", + "litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager.has_user_oauth_token", new_callable=AsyncMock, ) as mock_has_token, patch( - "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager.get_mcp_server_by_name", + "litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager.get_mcp_server_by_name", return_value=m2m_server, ), patch.object(session_manager_stateless, "handle_request", new_callable=AsyncMock) as mock_handle_request, @@ -929,16 +929,16 @@ async def test_handle_streamable_http_mcp_delegated_server_surfaces_upstream_cha return_value=False, ), patch( - "litellm.proxy._experimental.mcp_server.server._get_user_oauth_extra_headers_from_db", + "litellm.proxy._experimental.mcp_server.operations._get_user_oauth_extra_headers_from_db", new_callable=AsyncMock, return_value=None, ), patch( - "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager.get_mcp_server_by_name", + "litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager.get_mcp_server_by_name", return_value=delegated_server, ), patch( # test-quality-ok: registry is empty in unit tests; key owns the delegated server - "litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", + "litellm.proxy._experimental.mcp_server.operations._get_allowed_mcp_servers", new_callable=AsyncMock, return_value=[delegated_server], ), @@ -1022,12 +1022,12 @@ async def test_per_user_oauth_with_stored_token_skips_preemptive_401(): return_value=False, ), patch( - "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager.has_user_oauth_token", + "litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager.has_user_oauth_token", new_callable=AsyncMock, return_value=True, ) as mock_has_token, patch( - "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager.get_mcp_server_by_name", + "litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager.get_mcp_server_by_name", return_value=oauth_server, ), patch.object( @@ -1126,11 +1126,11 @@ async def test_handle_streamable_http_mcp_delegated_server_without_token_returns return_value=False, ), patch( - "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager.has_user_oauth_token", + "litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager.has_user_oauth_token", new_callable=AsyncMock, ) as mock_has_token, patch( - "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager.get_mcp_server_by_name", + "litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager.get_mcp_server_by_name", return_value=delegated_server, ), patch.object( @@ -1218,7 +1218,7 @@ async def test_handle_streamable_http_mcp_token_exchange_without_subject_returns return_value=False, ), patch( - "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager.get_mcp_server_by_name", + "litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager.get_mcp_server_by_name", return_value=obo_server, ), patch.object( @@ -1317,7 +1317,7 @@ async def test_handle_streamable_http_mcp_oauth_delegate_without_token_returns_g True, ), patch( - "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager.get_mcp_server_by_name", + "litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager.get_mcp_server_by_name", return_value=od_server, ), patch.object( @@ -1391,7 +1391,7 @@ async def test_handle_streamable_http_mcp_oauth_delegate_with_forwarded_token_sk new_callable=AsyncMock, ), patch( - "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager.get_mcp_server_by_name", + "litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager.get_mcp_server_by_name", return_value=od_server, ), patch.object( @@ -1453,7 +1453,7 @@ async def _run_passthrough_connect( new_callable=AsyncMock, ), patch( - "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager.get_mcp_server_by_name", + "litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager.get_mcp_server_by_name", return_value=server, ), patch.object(session_manager_stateless, "handle_request", new_callable=AsyncMock) as mock_handle_request, @@ -1574,7 +1574,7 @@ async def test_handle_streamable_http_mcp_true_passthrough_without_token_surface return_value=probe_client, ), patch( - "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager.get_mcp_server_by_name", + "litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager.get_mcp_server_by_name", return_value=tp_server, ), patch.object( @@ -1642,7 +1642,7 @@ async def test_handle_streamable_http_mcp_true_passthrough_dcr_bridge_challenges return_value=probe_client, ), patch( - "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager.get_mcp_server_by_name", + "litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager.get_mcp_server_by_name", return_value=bridge_server, ), patch.object( @@ -1720,7 +1720,7 @@ async def test_handle_streamable_http_mcp_true_passthrough_with_token_skips_prob return_value=probe_client, ), patch( - "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager.get_mcp_server_by_name", + "litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager.get_mcp_server_by_name", return_value=tp_server, ), patch.object( diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_tool_search.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_tool_search.py index cb43d2c2592..4575741aa8b 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_tool_search.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_tool_search.py @@ -1,3 +1,4 @@ +from litellm.proxy._experimental.mcp_server import operations as mcp_operations """ Tests for MCP tool search feature. @@ -572,7 +573,7 @@ class TestCallToolRestApiVirtualTools: mock_tool.input_schema = {"type": "object", "properties": {}} with patch( - "litellm.proxy._experimental.mcp_server.server._list_mcp_tools", + "litellm.proxy._experimental.mcp_server.operations._list_mcp_tools", new_callable=AsyncMock, return_value=AggregateToolListing(tools=[mock_tool], outcomes={}), ): @@ -616,12 +617,12 @@ class TestCallToolRestApiVirtualTools: with ( patch( - "litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", + "litellm.proxy._experimental.mcp_server.operations._get_allowed_mcp_servers", new_callable=AsyncMock, return_value=[MagicMock()], ), patch( - "litellm.proxy._experimental.mcp_server.server.execute_mcp_tool", + "litellm.proxy._experimental.mcp_server.operations.execute_mcp_tool", new_callable=AsyncMock, return_value=fake_result, ) as mock_execute, @@ -669,12 +670,12 @@ class TestCallToolRestApiVirtualTools: return_value="203.0.113.7", ), patch( - "litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", + "litellm.proxy._experimental.mcp_server.operations._get_allowed_mcp_servers", new_callable=AsyncMock, return_value=[MagicMock()], ) as mock_allowed, patch( - "litellm.proxy._experimental.mcp_server.server.execute_mcp_tool", + "litellm.proxy._experimental.mcp_server.operations.execute_mcp_tool", new_callable=AsyncMock, return_value=fake_result, ), @@ -699,7 +700,7 @@ class TestCallToolRestApiVirtualTools: return_value="203.0.113.7", ), patch( - "litellm.proxy._experimental.mcp_server.server._list_mcp_tools", + "litellm.proxy._experimental.mcp_server.operations._list_mcp_tools", new_callable=AsyncMock, return_value=AggregateToolListing(tools=[], outcomes={}), ) as mock_list, @@ -832,7 +833,7 @@ class TestCallToolRestApiVirtualTools: "litellm.proxy.proxy_server.proxy_logging_obj", key_limits ), patch( # test-quality-ok: the authorized catalog is the seam every virtual tool shares; the ranking under test stays real - "litellm.proxy._experimental.mcp_server.server._list_mcp_tools", + "litellm.proxy._experimental.mcp_server.operations._list_mcp_tools", new_callable=AsyncMock, return_value=AggregateToolListing(tools=list(CATALOG), outcomes={}), ) as mock_list, @@ -939,7 +940,7 @@ class TestDispatchVirtualMcpTool: new_callable=AsyncMock, return_value="SEARCH_RESULT", ) as mock_search: - result = await srv._dispatch_virtual_mcp_tool( + result = await mcp_operations._dispatch_virtual_mcp_tool( name=MCP_TOOL_SEARCH_TOOL_NAME, arguments={"query": "q", "top_k": 3}, user_api_key_auth=uak, @@ -961,7 +962,7 @@ class TestDispatchVirtualMcpTool: new_callable=AsyncMock, return_value="AGENT_RESULT", ) as mock_agent_search: - result = await srv._dispatch_virtual_mcp_tool( + result = await mcp_operations._dispatch_virtual_mcp_tool( name=AGENT_SEARCH_TOOL_NAME, arguments={"query": "translate a document", "top_k": "2"}, user_api_key_auth=uak, @@ -996,7 +997,7 @@ class TestDispatchVirtualMcpTool: new_callable=AsyncMock, return_value="CALL_RESULT", ) as mock_call: - result = await srv._dispatch_virtual_mcp_tool( + result = await mcp_operations._dispatch_virtual_mcp_tool( name=MCP_TOOL_CALL_TOOL_NAME, arguments={"tool_name": "math-add", "arguments": {"a": 1, "b": 2}}, user_api_key_auth=uak, @@ -1027,8 +1028,7 @@ class TestDispatchVirtualMcpTool: sentinel_logging_obj = object() with ( patch.object( - srv, - "_build_virtual_call_logging_obj", + mcp_operations, "_build_virtual_call_logging_obj", new_callable=AsyncMock, return_value=sentinel_logging_obj, ) as mock_build, @@ -1038,7 +1038,7 @@ class TestDispatchVirtualMcpTool: return_value="CALL_RESULT", ) as mock_call, ): - await srv._dispatch_virtual_mcp_tool( + await mcp_operations._dispatch_virtual_mcp_tool( name=MCP_TOOL_CALL_TOOL_NAME, arguments={"tool_name": "math-add", "arguments": {"a": 1}}, user_api_key_auth=uak, @@ -1060,7 +1060,7 @@ class TestDispatchVirtualMcpTool: new_callable=AsyncMock, return_value="SEARCH_RESULT", ) as mock_search: - await srv._dispatch_virtual_mcp_tool( + await mcp_operations._dispatch_virtual_mcp_tool( name=MCP_TOOL_SEARCH_TOOL_NAME, arguments={"query": "issue", "top_k": "not-a-number"}, user_api_key_auth=uak, @@ -1083,12 +1083,12 @@ class TestDispatchVirtualMcpTool: fake = CallToolResult(content=[TextContent(type="text", text="ok")], isError=False) with ( patch( - "litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", + "litellm.proxy._experimental.mcp_server.operations._get_allowed_mcp_servers", new_callable=AsyncMock, return_value=[MagicMock()], ) as mock_allowed, patch( - "litellm.proxy._experimental.mcp_server.server.execute_mcp_tool", + "litellm.proxy._experimental.mcp_server.operations.execute_mcp_tool", new_callable=AsyncMock, return_value=fake, ) as mock_exec, @@ -1130,12 +1130,12 @@ class TestDispatchVirtualMcpTool: uak = UserAPIKeyAuth(api_key="k", object_permission=_make_perm(mcp_tool_search_enabled=True)) with ( patch( - "litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", + "litellm.proxy._experimental.mcp_server.operations._get_allowed_mcp_servers", new_callable=AsyncMock, return_value=[], ), patch( - "litellm.proxy._experimental.mcp_server.server.execute_mcp_tool", + "litellm.proxy._experimental.mcp_server.operations.execute_mcp_tool", new_callable=AsyncMock, ) as mock_exec, ): @@ -1217,7 +1217,7 @@ class TestMcpServerToolCallErrorHandling: return_value=(uak, None, None, None, None, None, None), ), patch( - "litellm.proxy._experimental.mcp_server.server._dispatch_virtual_mcp_tool", + "litellm.proxy._experimental.mcp_server.operations._dispatch_virtual_mcp_tool", new_callable=AsyncMock, side_effect=HTTPException(status_code=403, detail="User not allowed to call this tool"), ), @@ -1254,7 +1254,7 @@ async def test_handle_mcp_tool_call_scoped_denial_names_the_binding_agent() -> N ] with patch( # test-quality-ok: the permission resolver is a module-level function; the suite's only seam - "litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", + "litellm.proxy._experimental.mcp_server.operations._get_allowed_mcp_servers", new=AsyncMock(side_effect=resolve), ): with pytest.raises(HTTPException) as exc_info: diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_toolset_scope.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_toolset_scope.py index 519acc241c6..60098b1f656 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_toolset_scope.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_toolset_scope.py @@ -564,7 +564,7 @@ class TestMCPActiveToolsetContextVar: MagicMock(get_mcp_client_ip=MagicMock(return_value="127.0.0.1")), ), patch( - "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager", + "litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager", MagicMock(get_mcp_server_by_name=MagicMock(return_value=None)), ), patch( diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_openapi_tool_auth.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_openapi_tool_auth.py index ac716bace3c..bb70f38285c 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_openapi_tool_auth.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_openapi_tool_auth.py @@ -1,3 +1,4 @@ +from litellm.proxy._experimental.mcp_server import operations as mcp_operations """ VERIA-7 regression: OpenAPI-backed (local-registry) MCP tools must run through `pre_call_tool_check` before dispatch, the same as managed @@ -49,22 +50,22 @@ async def test_openapi_local_tool_runs_pre_call_tool_check(): with ( patch.object( - mcp_module.global_mcp_server_manager, + mcp_operations.global_mcp_server_manager, "_get_mcp_server_from_tool_name", return_value=fake_server, ), patch.object( - mcp_module.global_mcp_server_manager, + mcp_operations.global_mcp_server_manager, "pre_call_tool_check", new=pre_call, ), patch.object( - mcp_module.global_mcp_tool_registry, + mcp_operations.global_mcp_tool_registry, "get_tool", return_value=fake_tool, ), patch( - "litellm.proxy._experimental.mcp_server.server._handle_local_mcp_tool", + "litellm.proxy._experimental.mcp_server.operations._handle_local_mcp_tool", new=handle_local, ), patch( @@ -72,7 +73,7 @@ async def test_openapi_local_tool_runs_pre_call_tool_check(): return_value=True, ), ): - await mcp_module.execute_mcp_tool( + await mcp_operations.execute_mcp_tool( name="list_pets", arguments={"limit": 10}, allowed_mcp_servers=[fake_server], @@ -92,7 +93,7 @@ async def test_openapi_local_tool_runs_pre_call_tool_check(): assert pre_call_kwargs["guardrail_context"] == {"metadata": {"guardrails": ("block-all",)}} assert pre_call_kwargs["name"] == "list_pets" assert pre_call_kwargs["server"] is fake_server - assert pre_call_kwargs["user_api_key_auth"] is user + assert pre_call_kwargs["user_api_key_auth"] == user # `proxy_logging_obj` must be sourced from the canonical proxy_server # module (same as the managed path) — passing None would crash the # downstream `_create_mcp_request_object_from_kwargs` call with @@ -134,22 +135,22 @@ async def test_openapi_local_tool_blocked_when_pre_call_check_raises(): with ( patch.object( - mcp_module.global_mcp_server_manager, + mcp_operations.global_mcp_server_manager, "_get_mcp_server_from_tool_name", return_value=fake_server, ), patch.object( - mcp_module.global_mcp_server_manager, + mcp_operations.global_mcp_server_manager, "pre_call_tool_check", new=pre_call, ), patch.object( - mcp_module.global_mcp_tool_registry, + mcp_operations.global_mcp_tool_registry, "get_tool", return_value=fake_tool, ), patch( - "litellm.proxy._experimental.mcp_server.server._handle_local_mcp_tool", + "litellm.proxy._experimental.mcp_server.operations._handle_local_mcp_tool", new=handle_local, ), patch( @@ -158,7 +159,7 @@ async def test_openapi_local_tool_blocked_when_pre_call_check_raises(): ), ): with pytest.raises(HTTPException) as exc: - await mcp_module.execute_mcp_tool( + await mcp_operations.execute_mcp_tool( name="delete_pet", arguments={}, allowed_mcp_servers=[fake_server], @@ -195,24 +196,24 @@ async def test_openapi_local_tool_denied_when_server_not_resolvable(): # `_get_mcp_server_from_tool_name` returns None — no server context. with ( - patch.object(mcp_module, "_resolve_openapi_tool_auth", new=resolve_auth), + patch.object(mcp_operations, "_resolve_openapi_tool_auth", new=resolve_auth), patch.object( - mcp_module.global_mcp_server_manager, + mcp_operations.global_mcp_server_manager, "_get_mcp_server_from_tool_name", return_value=None, ), patch.object( - mcp_module.global_mcp_server_manager, + mcp_operations.global_mcp_server_manager, "pre_call_tool_check", new=pre_call, ), patch.object( - mcp_module.global_mcp_tool_registry, + mcp_operations.global_mcp_tool_registry, "get_tool", return_value=fake_tool, ), patch( - "litellm.proxy._experimental.mcp_server.server._handle_local_mcp_tool", + "litellm.proxy._experimental.mcp_server.operations._handle_local_mcp_tool", new=handle_local, ), patch( @@ -221,7 +222,7 @@ async def test_openapi_local_tool_denied_when_server_not_resolvable(): ), ): with pytest.raises(HTTPException) as exc: - await mcp_module.execute_mcp_tool( + await mcp_operations.execute_mcp_tool( name="list_pets", arguments={}, allowed_mcp_servers=[], @@ -280,27 +281,27 @@ async def test_openapi_local_tool_injects_resolved_oauth_token(): with ( patch.object( - mcp_module.global_mcp_server_manager, + mcp_operations.global_mcp_server_manager, "_get_mcp_server_from_tool_name", return_value=oauth_server, ), patch.object( - mcp_module.global_mcp_server_manager, + mcp_operations.global_mcp_server_manager, "pre_call_tool_check", new=AsyncMock(return_value={}), ), patch.object( - mcp_module.global_mcp_tool_registry, + mcp_operations.global_mcp_tool_registry, "get_tool", return_value=fake_tool, ), patch.object( - mcp_module.global_mcp_server_manager._cred_provider, + mcp_operations.global_mcp_server_manager._cred_provider, "resolve_credentials", new=AsyncMock(return_value=Ok(StaticHeaderAuth("Bearer stored-user-token"))), ), patch( - "litellm.proxy._experimental.mcp_server.server._handle_local_mcp_tool", + "litellm.proxy._experimental.mcp_server.operations._handle_local_mcp_tool", new=handle_local, ), patch( @@ -308,7 +309,7 @@ async def test_openapi_local_tool_injects_resolved_oauth_token(): return_value=True, ), ): - await mcp_module.execute_mcp_tool( + await mcp_operations.execute_mcp_tool( name="get_values", arguments={}, allowed_mcp_servers=[oauth_server], @@ -417,7 +418,7 @@ async def test_legacy_local_tool_fallback_refuses_unentitled_caller(legacy_local ) with pytest.raises(HTTPException) as exc: - await mcp_module.execute_mcp_tool( + await mcp_operations.execute_mcp_tool( name=f"{LEGACY_SERVER_NAME}-{LEGACY_TOOL}", arguments={}, allowed_mcp_servers=[server], @@ -451,7 +452,7 @@ async def test_legacy_local_tool_fallback_still_dispatches_entitled_caller( server, executed = legacy_local_tool user = _caller_entitled_to([LEGACY_TOOL]) - result = await mcp_module.execute_mcp_tool( + result = await mcp_operations.execute_mcp_tool( name=f"{LEGACY_SERVER_NAME}-{LEGACY_TOOL}", arguments={}, allowed_mcp_servers=[server], @@ -481,7 +482,7 @@ async def test_legacy_local_tool_fallback_fails_closed_on_empty_prefix( _server, executed = legacy_local_tool with pytest.raises(HTTPException) as exc: - await mcp_module.execute_mcp_tool( + await mcp_operations.execute_mcp_tool( name=f"-{LEGACY_TOOL}", arguments={}, allowed_mcp_servers=[], @@ -523,7 +524,7 @@ async def test_legacy_local_tool_fallback_fails_closed_when_prefix_names_no_serv return_value=True, ): with pytest.raises(HTTPException) as exc: - await mcp_module.execute_mcp_tool( + await mcp_operations.execute_mcp_tool( name=f"{LEGACY_SERVER_NAME}-{LEGACY_TOOL}", arguments={}, allowed_mcp_servers=[other_server], @@ -546,7 +547,7 @@ async def test_unknown_tool_name_still_reports_not_found(): from litellm.proxy._experimental.mcp_server import server as mcp_module with pytest.raises(HTTPException) as exc: - await mcp_module.execute_mcp_tool( + await mcp_operations.execute_mcp_tool( name="tool_no_registry_knows", arguments={}, allowed_mcp_servers=[], @@ -610,7 +611,7 @@ async def test_per_server_auth_header_reaches_both_openapi_dispatch_arms(dispatc captured["injected"] = _request_auth_header.get() return [] - manager = mcp_module.global_mcp_server_manager + manager = mcp_operations.global_mcp_server_manager with ( patch.object(manager, "resolve_openapi_upstream_auth", new=fake_resolver), patch.object(manager, "pre_call_tool_check", new=AsyncMock(return_value={})), @@ -620,9 +621,9 @@ async def test_per_server_auth_header_reaches_both_openapi_dispatch_arms(dispatc fake_tool.name = "list_reports" with ( patch.object(manager, "_get_mcp_server_from_tool_name", return_value=server), - patch.object(mcp_module.global_mcp_tool_registry, "get_tool", return_value=fake_tool), + patch.object(mcp_operations.global_mcp_tool_registry, "get_tool", return_value=fake_tool), patch( - "litellm.proxy._experimental.mcp_server.server._handle_local_mcp_tool", + "litellm.proxy._experimental.mcp_server.operations._handle_local_mcp_tool", new=capture_local, ), patch( @@ -630,7 +631,7 @@ async def test_per_server_auth_header_reaches_both_openapi_dispatch_arms(dispatc return_value=True, ), ): - await mcp_module.execute_mcp_tool( + await mcp_operations.execute_mcp_tool( name="list_reports", arguments={}, allowed_mcp_servers=[server], @@ -702,11 +703,11 @@ async def test_local_dispatch_reports_the_outcome_instead_of_success(failure: st user = UserAPIKeyAuth(api_key="sk-user", user_id="alice", user_role=LitellmUserRoles.INTERNAL_USER.value) with ( - patch.object(mcp_module.global_mcp_server_manager, "_get_mcp_server_from_tool_name", return_value=server), - patch.object(mcp_module.global_mcp_server_manager, "pre_call_tool_check", new=AsyncMock(return_value={})), - patch.object(mcp_module.global_mcp_tool_registry, "get_tool", return_value=fake_tool), + patch.object(mcp_operations.global_mcp_server_manager, "_get_mcp_server_from_tool_name", return_value=server), + patch.object(mcp_operations.global_mcp_server_manager, "pre_call_tool_check", new=AsyncMock(return_value={})), + patch.object(mcp_operations.global_mcp_tool_registry, "get_tool", return_value=fake_tool), patch.object( - mcp_module.global_mcp_server_manager, + mcp_operations.global_mcp_server_manager, "resolve_openapi_upstream_auth", new=AsyncMock(return_value=(None, None)), ), @@ -715,7 +716,7 @@ async def test_local_dispatch_reports_the_outcome_instead_of_success(failure: st return_value=True, ), ): - call = mcp_module.execute_mcp_tool( + call = mcp_operations.execute_mcp_tool( name="list_reports", arguments={}, allowed_mcp_servers=[server], diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_operations.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_operations.py new file mode 100644 index 00000000000..870f2cd5f8f --- /dev/null +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_operations.py @@ -0,0 +1,343 @@ +import asyncio +from unittest.mock import AsyncMock, patch + +import pytest +from mcp.types import GetPromptRequest, GetPromptRequestParams, GetPromptResult + +from litellm.proxy._experimental.mcp_server.operations import GatewayOperations, prepare_context +from litellm.proxy._types import UserAPIKeyAuth + + +@pytest.mark.asyncio +async def test_dispatch_uses_explicit_context_when_ambient_caller_differs(): + from mcp.server.auth.middleware.auth_context import auth_context_var + from litellm.proxy._experimental.mcp_server.server import set_auth_context + + context = prepare_context( + UserAPIKeyAuth(user_id="alpha"), + raw_headers={"x-caller": "alpha"}, + mcp_servers=["alpha-server"], + client_ip="192.0.2.1", + ) + token = auth_context_var.set(None) + handler = AsyncMock(return_value=GetPromptResult(messages=[])) + try: + set_auth_context(UserAPIKeyAuth(user_id="bravo"), raw_headers={"x-caller": "bravo"}) + with patch("litellm.proxy._experimental.mcp_server.operations.mcp_get_prompt", handler): + result = await GatewayOperations().execute( + GetPromptRequest(params=GetPromptRequestParams(name="alpha-prompt")), context + ) + assert result.messages == [] + assert handler.await_args.kwargs["name"] == "alpha-prompt" + assert handler.await_args.kwargs["user_api_key_auth"].user_id == "alpha" + assert handler.await_args.kwargs["raw_headers"] == {"x-caller": "alpha"} + assert handler.await_args.kwargs["mcp_servers"] == ["alpha-server"] + assert handler.await_args.kwargs["client_ip"] == "192.0.2.1" + finally: + auth_context_var.reset(token) + + +@pytest.mark.asyncio +async def test_legacy_adapter_cleans_context_after_cancelled_operation(): + from types import SimpleNamespace + from litellm.proxy._experimental.mcp_server import server + from litellm.proxy._experimental.mcp_server.mcp_context import active_mcp_request_ctx_var + + previous_session = server.active_mcp_session_var.get() + previous_request = active_mcp_request_ctx_var.get() + request = SimpleNamespace(session=object()) + auth = (None, None, None, None, None, None, None) + + async def cancelled_operation(): + async with server._legacy_operation_context(request, trace=False): + assert server.active_mcp_session_var.get() is request.session + assert active_mcp_request_ctx_var.get() is request + raise asyncio.CancelledError + + with patch( + "litellm.proxy._experimental.mcp_server.server.get_or_extract_auth_context", AsyncMock(return_value=auth) + ): + with pytest.raises(asyncio.CancelledError): + await cancelled_operation() + assert server.active_mcp_session_var.get() is previous_session + assert active_mcp_request_ctx_var.get() is previous_request + + +@pytest.mark.asyncio +async def test_legacy_adapter_cleans_context_when_trace_setup_fails(): + from types import SimpleNamespace + from litellm.proxy._experimental.mcp_server import server + from litellm.proxy._experimental.mcp_server.mcp_context import active_mcp_request_ctx_var + + previous_session = server.active_mcp_session_var.get() + previous_request = active_mcp_request_ctx_var.get() + request = SimpleNamespace(session=object()) + + async def enter_operation(): + async with server._legacy_operation_context(request, trace=True): + pytest.fail("Trace setup failure must prevent dispatch") + + with patch.object(server, "_otel_set_mcp_transport_span", side_effect=RuntimeError("trace failure")): + with pytest.raises(RuntimeError, match="trace failure"): + await enter_operation() + assert server.active_mcp_session_var.get() is previous_session + assert active_mcp_request_ctx_var.get() is previous_request + + +@pytest.mark.asyncio +async def test_prompt_sampling_receives_explicit_operation_caller_headers_and_ip(): + from unittest.mock import MagicMock + from litellm.proxy._experimental.mcp_server import operations + from litellm.types.mcp import MCPTransport + from litellm.types.mcp_server.mcp_server_manager import MCPServer + + upstream = MCPServer( + server_id="explicit-prompt", + name="explicit_prompt", + url="https://example.invalid/mcp", + transport=MCPTransport.http, + allow_sampling=True, + ) + context = prepare_context( + UserAPIKeyAuth(user_id="prompt-caller"), + raw_headers={"x-caller": "prompt-caller"}, + client_ip="192.0.2.41", + ) + client = MagicMock() + client.get_prompt = AsyncMock(return_value=GetPromptResult(messages=[])) + sampling = AsyncMock() + with ( + patch.object(operations, "_get_allowed_mcp_servers", AsyncMock(return_value=[upstream])), + patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPClient", return_value=client) as factory, + patch("litellm.proxy._experimental.mcp_server.sampling_handler.handle_sampling_create_message", sampling), + ): + result = await GatewayOperations().execute( + GetPromptRequest(params=GetPromptRequestParams(name="explicit_prompt-prompt")), context + ) + assert result.messages == [] + await factory.call_args.kwargs["sampling_callback"](None, None) + captured = sampling.await_args.kwargs + assert captured["user_api_key_auth"] is not None + assert captured["user_api_key_auth"].user_id == "prompt-caller" + assert captured["raw_headers"] == {"x-caller": "prompt-caller"} + assert captured["client_ip"] == "192.0.2.41" + + +def _catalog_case(method): + from mcp import types + + cases = { + "prompts/list": ( + types.ListPromptsRequest(), + "list_prompts", + "get_prompts_from_server", + [types.Prompt(name="catalog-prompt")], + "prompts", + ), + "prompts/get": ( + types.GetPromptRequest( + params=types.GetPromptRequestParams(name="catalog-prompt", arguments={"topic": "test"}) + ), + "get_prompt", + "get_prompt_from_server", + types.GetPromptResult(messages=[]), + None, + ), + "resources/list": ( + types.ListResourcesRequest(), + "list_resources", + "get_resources_from_server", + [types.Resource(name="document", uri="https://example.com/document")], + "resources", + ), + "resources/templates/list": ( + types.ListResourceTemplatesRequest(), + "list_resource_templates", + "get_resource_templates_from_server", + [types.ResourceTemplate(name="document", uri_template="https://example.com/{name}")], + "resource_templates", + ), + "resources/read": ( + types.ReadResourceRequest(params=types.ReadResourceRequestParams(uri="https://example.com/document")), + "read_resource", + "read_resource_from_server", + types.ReadResourceResult( + contents=[types.TextResourceContents(uri="https://example.com/document", text="document body")] + ), + None, + ), + } + return cases[method] + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "method", ["prompts/list", "prompts/get", "resources/list", "resources/templates/list", "resources/read"] +) +@pytest.mark.parametrize("state", ["success", "denied", "upstream_failure", "scope_failure"]) +async def test_native_catalog_operations_preserve_context_results_and_failure_policy(method, state): + from types import SimpleNamespace + from fastapi import HTTPException + from mcp.server.context import ServerRequestContext + from mcp.types import PaginatedRequestParams + from litellm.proxy._experimental.mcp_server import operations, server + from litellm.types.mcp import MCPTransport + from litellm.types.mcp_server.mcp_server_manager import MCPServer + + operation, handler_name, manager_method, payload, collection = _catalog_case(method) + caller = UserAPIKeyAuth(user_id="catalog-caller") + headers = {"x-caller": "catalog-caller"} + upstream_server = MCPServer(server_id="catalog", name="catalog", transport=MCPTransport.http) + allowed = AsyncMock( + return_value=[] if state == "denied" else [upstream_server], + side_effect=HTTPException(status_code=403, detail="scope denied") if state == "scope_failure" else None, + ) + upstream = AsyncMock( + return_value=payload, side_effect=RuntimeError("upstream unavailable") if state == "upstream_failure" else None + ) + ctx = ServerRequestContext( + session=SimpleNamespace(), lifespan_context={}, protocol_version="2025-06-18", method=method + ) + auth = (caller, None, ["catalog"], None, None, headers, "192.0.2.41") + with ( + patch.object(server, "get_or_extract_auth_context", AsyncMock(return_value=auth)), + patch.object(operations, "_get_allowed_mcp_servers", allowed), + patch.object(operations.global_mcp_server_manager, manager_method, upstream), + ): + if collection is None and state != "success": + expected_error = RuntimeError if state == "upstream_failure" else HTTPException + with pytest.raises(expected_error): + await getattr(server, handler_name)(ctx, operation.params) + else: + result = await getattr(server, handler_name)(ctx, operation.params or PaginatedRequestParams()) + if collection: + assert getattr(result, collection) == (payload if state == "success" else []) + else: + assert result == payload + assert allowed.await_args.kwargs == { + "user_api_key_auth": caller, + "mcp_servers": ["catalog"], + "client_ip": "192.0.2.41", + } + if state in ("denied", "scope_failure"): + upstream.assert_not_awaited() + else: + upstream.assert_awaited_once() + forwarded = upstream.await_args.kwargs + assert forwarded["user_api_key_auth"] == caller + assert forwarded["raw_headers"] == headers + assert forwarded["client_ip"] == "192.0.2.41" + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "method", ["prompts/list", "prompts/get", "resources/list", "resources/templates/list", "resources/read"] +) +async def test_explicit_proxy_context_rejects_catalog_operations_before_upstream_access(method): + from mcp.shared.exceptions import MCPError + from mcp.types import METHOD_NOT_FOUND + from litellm.proxy._experimental.mcp_server import operations + + operation, _, manager_method, _, _ = _catalog_case(method) + upstream = AsyncMock() + with patch.object(operations.global_mcp_server_manager, manager_method, upstream): + with pytest.raises(MCPError) as rejected: + await GatewayOperations().execute(operation, prepare_context(mcp_proxy_mode=True)) + assert rejected.value.error.code == METHOD_NOT_FOUND + assert rejected.value.error.message == "Operation unavailable on /mcp/proxy" + upstream.assert_not_awaited() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("failure", ["missing_env", "pii", "guardrail", "unexpected"]) +async def test_tool_operation_preserves_failure_messages_and_request_trace(failure): + from mcp.types import CallToolRequest, CallToolRequestParams + from litellm.exceptions import BlockedPiiEntityError, GuardrailRaisedException + from litellm.proxy._experimental.mcp_server import operations + from litellm.proxy._experimental.mcp_server.utils import MCPMissingUserEnvVarsError + + failures = { + "missing_env": ( + MCPMissingUserEnvVarsError( + server_id="server", server_name="server", missing=["TOKEN"], setup_url="https://example.com/setup" + ), + "https://example.com/setup", + ), + "pii": ( + BlockedPiiEntityError(entity_type="EMAIL_ADDRESS", guardrail_name="test"), + "Blocked PII entity detected", + ), + "guardrail": (GuardrailRaisedException(message="request denied"), "Guardrail violation"), + "unexpected": (RuntimeError("upstream unavailable"), "Error: upstream unavailable"), + } + error, expected = failures[failure] + dispatch = AsyncMock(side_effect=error) + context = prepare_context( + raw_headers={"x-litellm-trace-id": "operation-trace", "authorization": "private-test-header"} + ) + with patch.object(operations, "call_mcp_tool", dispatch): + result = await GatewayOperations().execute( + CallToolRequest(params=CallToolRequestParams(name="catalog-tool", arguments={})), context + ) + assert result.is_error is True + assert expected in result.content[0].text + assert "private-test-header" not in result.content[0].text + dispatch.assert_awaited_once() + assert dispatch.await_args.kwargs["litellm_trace_id"] == "operation-trace" + assert dispatch.await_args.kwargs["litellm_session_id"] == "operation-trace" + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "method,helper", + [ + ("prompts/list", "_list_mcp_prompts"), + ("resources/list", "_list_mcp_resources"), + ("resources/templates/list", "_list_mcp_resource_templates"), + ], +) +async def test_catalog_operation_preserves_empty_result_for_malformed_upstream_items(method, helper): + from litellm.proxy._experimental.mcp_server import operations + + operation, _, _, _, collection = _catalog_case(method) + with patch.object(operations, helper, AsyncMock(return_value=[{"unexpected": "item"}])): + result = await GatewayOperations().execute(operation, prepare_context()) + assert getattr(result, collection) == [] + + +@pytest.mark.asyncio +@pytest.mark.parametrize("catalog_unavailable", [False, True]) +async def test_tool_listing_returns_empty_result_without_dispatch_for_unavailable_catalog(catalog_unavailable): + from mcp.types import ListToolsRequest + from litellm.proxy._experimental.mcp_server import operations + + allowed = AsyncMock( + return_value=[], side_effect=RuntimeError("catalog unavailable") if catalog_unavailable else None + ) + upstream = AsyncMock() + with ( + patch.object(operations, "_get_allowed_mcp_servers", allowed), + patch.object(operations.global_mcp_server_manager, "_get_tools_from_server", upstream), + ): + result = await GatewayOperations().execute(ListToolsRequest(), prepare_context()) + assert result.tools == [] + allowed.assert_awaited_once() + upstream.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_explicit_proxy_context_lists_builtin_tools_and_blocks_direct_tool_dispatch(): + from mcp.types import CallToolRequest, CallToolRequestParams, ListToolsRequest + from litellm.proxy._experimental.mcp_server import operations + + context = prepare_context(mcp_proxy_mode=True) + allowed = AsyncMock() + with patch.object(operations, "_get_allowed_mcp_servers", allowed): + listing = await GatewayOperations().execute(ListToolsRequest(), context) + denied = await GatewayOperations().execute( + CallToolRequest(params=CallToolRequestParams(name="catalog-tool", arguments={})), context + ) + assert {tool.name for tool in listing.tools} == {"search_tools", "get_tool_schema", "call_tool"} + assert denied.is_error is True + assert "unavailable on /mcp/proxy" in denied.content[0].text + allowed.assert_not_awaited() 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 13af58c15c0..233a8cc96ba 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 @@ -1,3 +1,4 @@ +from litellm.proxy._experimental.mcp_server import operations as mcp_operations import asyncio import inspect import json @@ -1253,6 +1254,7 @@ class TestListToolsRestAPI: user_api_key_auth=None, extra_headers=None, apply_tool_filters=True, + client_ip=None, ): captured["called"] = True captured["server"] = server @@ -1338,6 +1340,7 @@ class TestListToolsRestAPI: user_api_key_auth=None, extra_headers=None, apply_tool_filters=True, + client_ip=None, ): captured["user_api_key_auth"] = user_api_key_auth return ["tool-1"] @@ -1891,6 +1894,7 @@ class TestListToolsRestAPI: user_api_key_auth=None, extra_headers=None, apply_tool_filters=True, + client_ip=None, ): captured["called"] = True captured["server_arg"] = server @@ -2027,6 +2031,7 @@ class TestListToolsRestAPI: user_api_key_auth=None, extra_headers=None, apply_tool_filters=True, + client_ip=None, ): captured["called"] = True captured["server_arg"] = server @@ -2112,6 +2117,7 @@ class TestListToolsRestAPI: user_api_key_auth=None, extra_headers=None, apply_tool_filters=True, + client_ip=None, ): return ["scoped-tool"] @@ -2319,6 +2325,7 @@ class TestListToolsRestAPI: user_api_key_auth=None, extra_headers=None, apply_tool_filters=True, + client_ip=None, ): captured["server"] = server captured["auth_header"] = server_auth_header @@ -3145,10 +3152,10 @@ async def test_request_selected_tool_specific_guardrail_applies_to_virtual_execu monkeypatch.setattr(litellm, "callbacks", [guardrail]) monkeypatch.setattr(tool_registry, "global_mcp_tool_registry", registry) monkeypatch.setattr(mcp_server_manager, "global_mcp_server_manager", manager) - monkeypatch.setattr(server, "global_mcp_tool_registry", registry) - monkeypatch.setattr(server, "global_mcp_server_manager", manager) + monkeypatch.setattr(mcp_operations, "global_mcp_tool_registry", registry) + monkeypatch.setattr(mcp_operations, "global_mcp_server_manager", manager) monkeypatch.setattr(rest_endpoints, "global_mcp_server_manager", manager) - monkeypatch.setattr(server, "_get_allowed_mcp_servers", AsyncMock(return_value=[managed_server])) + monkeypatch.setattr(mcp_operations, "_get_allowed_mcp_servers", AsyncMock(return_value=[managed_server])) monkeypatch.setattr(proxy_server, "proxy_logging_obj", ProxyLogging(user_api_key_cache=DualCache())) monkeypatch.setattr(proxy_server, "add_litellm_data_to_request", passthrough_request_data) monkeypatch.setattr(proxy_server, "proxy_config", {}) diff --git a/tests/test_litellm/test_check_mcp_operation_boundary.py b/tests/test_litellm/test_check_mcp_operation_boundary.py new file mode 100644 index 00000000000..d7ac72de9f0 --- /dev/null +++ b/tests/test_litellm/test_check_mcp_operation_boundary.py @@ -0,0 +1,52 @@ +from pathlib import Path + +import pytest + +from scripts.check_mcp_operation_boundary import main, violations + + +@pytest.mark.parametrize( + "source", + ( + "from mcp.server.auth.middleware.auth_context import auth_context_var as hidden", + "from litellm.proxy._experimental.mcp_server.mcp_context import _mcp_proxy_mode as mode", + "caller = legacy.get_active_auth_context()", + "owners = transport._stateful_session_owners", + "from weakref import WeakKeyDictionary", + "from litellm.proxy._experimental.mcp_server.server import get_auth_context", + ), +) +def test_shared_operation_boundary_rejects_ambient_state(source): + assert violations(Path("operations.py"), source) + + +def test_legacy_adapter_may_resolve_context_but_policy_must_receive_it(): + source = "from litellm.proxy._experimental.mcp_server.mcp_context import _mcp_proxy_mode" + assert violations(Path("server.py"), source) == () + assert violations(Path("legacy_callbacks.py"), source) == () + assert violations(Path("operations.py"), "def execute(context):\n return context.client_ip") == () + assert violations(Path("mcp_server_manager.py"), "def _mcp_registry_key(server):\n return server.name") == () + + +def test_boundary_command_rejects_shared_state_and_accepts_explicit_context(tmp_path, monkeypatch, capsys): + import subprocess + import sys + + package = tmp_path / "litellm/proxy/_experimental/mcp_server" + package.mkdir(parents=True) + module = package / "operations.py" + module.write_text("from mcp.server.auth.middleware.auth_context import auth_context_var as hidden\n") + command = [sys.executable, str(Path(__file__).resolve().parents[2] / "scripts/check_mcp_operation_boundary.py")] + monkeypatch.chdir(tmp_path) + assert main() == 1 + assert "operations.py:1:" in capsys.readouterr().err + rejected = subprocess.run(command, cwd=tmp_path, capture_output=True, text=True, check=False) + assert rejected.returncode == 1 + assert "operations.py:1: MCP request/session state belongs in a legacy adapter" in rejected.stderr + + module.write_text("def execute(context):\n return context.client_ip\n") + assert main() == 0 + assert "MCP operation boundary: passed" in capsys.readouterr().out + accepted = subprocess.run(command, cwd=tmp_path, capture_output=True, text=True, check=False) + assert accepted.returncode == 0 + assert "MCP operation boundary: passed" in accepted.stdout From d2f30a77fded3a39c32b5dfc2f7f464bcd8f0409 Mon Sep 17 00:00:00 2001 From: Joshua Valluru <326636767+joshua-berri@users.noreply.github.com> Date: Mon, 21 Sep 2026 12:31:48 -0700 Subject: [PATCH 17/44] test(mcp): preserve toolset scope across explicit context --- .../mcp_server/test_mcp_toolset_scope.py | 16 ++++++++++++++++ 1 file changed, 16 insertions(+) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_toolset_scope.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_toolset_scope.py index 60098b1f656..1398884783e 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_toolset_scope.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_toolset_scope.py @@ -58,6 +58,22 @@ class TestApplyToolsetScope: assert set(op.mcp_servers or []) == {"server-a", "server-b"} assert op.mcp_tool_permissions == toolset_perms + from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager + from litellm.proxy._experimental.mcp_server.operations import prepare_context + + manager = MCPServerManager() + unscoped_open = await manager.operator_open_server_ids( + auth, allow_all_server_ids=["operator-open-outside-toolset"], submitted_server_ids=[] + ) + scoped_open = await manager.operator_open_server_ids( + prepare_context(result).user_api_key_auth, + allow_all_server_ids=["operator-open-outside-toolset"], + submitted_server_ids=[], + ) + assert unscoped_open == {"operator-open-outside-toolset"} + assert scoped_open == set() + assert auth.mcp_toolset_id is None + @pytest.mark.asyncio async def test_admin_creates_object_permission_when_none(self): """Admin key with object_permission=None can access any toolset.""" From e208b4e89ed41c6e4239c638f8e51a1239560959 Mon Sep 17 00:00:00 2001 From: Joshua Valluru <326636767+joshua-berri@users.noreply.github.com> Date: Mon, 21 Sep 2026 12:40:12 -0700 Subject: [PATCH 18/44] fix(mcp): keep explicit legacy sampling callers isolated --- .../_experimental/mcp_server/legacy_callbacks.py | 2 +- .../mcp_server/test_mcp_server_manager.py | 13 +++++++++---- 2 files changed, 10 insertions(+), 5 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/legacy_callbacks.py b/litellm/proxy/_experimental/mcp_server/legacy_callbacks.py index 4424907c28a..9e321062643 100644 --- a/litellm/proxy/_experimental/mcp_server/legacy_callbacks.py +++ b/litellm/proxy/_experimental/mcp_server/legacy_callbacks.py @@ -33,7 +33,7 @@ def create_sampling_callback( ) -> SamplingCallback: from litellm.proxy._experimental.mcp_server.server import get_active_auth_context - auth: Final = get_active_auth_context() if operation_context is None else None + auth: Final = get_active_auth_context() if operation_context is None and user_api_key_auth is None else None captured: Final = ( operation_context if operation_context is not None 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 7be79f9b514..2dc7e01966a 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 @@ -14214,10 +14214,11 @@ async def test_request_selected_during_guardrail_runs_concurrently_with_tool(mon @pytest.mark.asyncio -@pytest.mark.parametrize("with_caller", [True, False]) -async def test_client_sampling_does_not_fill_explicit_context_from_another_ambient_caller(with_caller): +@pytest.mark.parametrize("with_caller,legacy_factory", [(True, False), (False, False), (True, True)]) +async def test_client_sampling_does_not_fill_explicit_context_from_another_ambient_caller(with_caller, legacy_factory): from mcp.server.auth.middleware.auth_context import auth_context_var from litellm.proxy._experimental.mcp_server import server as legacy_server + from litellm.proxy._experimental.mcp_server.mcp_server_manager import _create_sampling_callback upstream = MCPServer(server_id="explicit-empty", name="explicit_empty", url="https://example.invalid/mcp", transport=MCPTransport.http, allow_sampling=True) token = auth_context_var.set(None) @@ -14228,8 +14229,12 @@ async def test_client_sampling_does_not_fill_explicit_context_from_another_ambie patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.MCPClient") as factory, patch("litellm.proxy._experimental.mcp_server.sampling_handler.handle_sampling_create_message", sampling), ): - await MCPServerManager()._create_mcp_client(upstream, user_api_key_auth=UserAPIKeyAuth(user_id="explicit") if with_caller else None) - await factory.call_args.kwargs["sampling_callback"](None, None) + if legacy_factory: + callback = _create_sampling_callback(user_api_key_auth=UserAPIKeyAuth(user_id="explicit")) + else: + await MCPServerManager()._create_mcp_client(upstream, user_api_key_auth=UserAPIKeyAuth(user_id="explicit") if with_caller else None) + callback = factory.call_args.kwargs["sampling_callback"] + await callback(None, None) captured = sampling.await_args.kwargs if with_caller: assert captured["user_api_key_auth"].user_id == "explicit" From f0e87f2457011f01eecfbd3d07ddc94cd99bdf94 Mon Sep 17 00:00:00 2001 From: shivam Date: Mon, 21 Sep 2026 19:50:07 +0000 Subject: [PATCH 19/44] fix(router): keep refusal gates closed once every fallback entry was tried Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/router.py | 9 +++--- .../router_utils/fallback_event_handlers.py | 12 +++++++ .../test_fallback_event_handlers.py | 19 ++++++++++- tests/test_litellm/test_router.py | 32 +++++++++++++++++++ 4 files changed, 67 insertions(+), 5 deletions(-) diff --git a/litellm/router.py b/litellm/router.py index 98c7c319eaa..d9dd48a2c6a 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -191,6 +191,7 @@ from litellm.router_utils.fallback_event_handlers import ( fallbacks_disabled_for_request, get_fallback_model_group_for_lookup_groups, get_pre_routing_selection, + has_unattempted_fallback_target, record_disable_fallbacks, record_pre_routing_selection, run_async_fallback, @@ -8338,12 +8339,12 @@ class Router: """ content_policy_fallbacks: Final = kwargs.get("content_policy_fallbacks", self.content_policy_fallbacks) if content_policy_fallbacks is not None: - return ( + return has_unattempted_fallback_target( self._get_fallback_model_group_for_lookup_groups( fallbacks=content_policy_fallbacks, lookup_groups=fallback_lookup_groups(kwargs, model_group), - ) - is not None + ), + kwargs, ) if self._has_default_fallbacks(): return True @@ -8375,7 +8376,7 @@ class Router: fallbacks=fallbacks, lookup_groups=fallback_lookup_groups(kwargs, model_group), ) - return resolved is not None + return has_unattempted_fallback_target(resolved, kwargs) def _should_raise_content_policy_error(self, model: str, response: ModelResponse, kwargs: dict) -> bool: """ diff --git a/litellm/router_utils/fallback_event_handlers.py b/litellm/router_utils/fallback_event_handlers.py index 156195f0587..08b9246e562 100644 --- a/litellm/router_utils/fallback_event_handlers.py +++ b/litellm/router_utils/fallback_event_handlers.py @@ -199,6 +199,18 @@ class AttemptedFallbackTargets: self.keys = self.keys | frozenset((key,)) +def has_unattempted_fallback_target( + fallback_model_group: Sequence[object] | None, kwargs: Mapping[str, object] +) -> bool: + """Whether a resolved chain still holds an entry this request has not tried.""" + if fallback_model_group is None: + return False + attempted: Final = kwargs.get("attempted_targets") + if not isinstance(attempted, AttemptedFallbackTargets): + return True + return any((key := fallback_attempt_key(target)) is None or key not in attempted for target in fallback_model_group) + + def _check_stripped_model_group(model_group: str, fallback_key: str) -> bool: """ Handles wildcard routing scenario diff --git a/tests/test_litellm/router_utils/test_fallback_event_handlers.py b/tests/test_litellm/router_utils/test_fallback_event_handlers.py index 69772561172..b6ab21dfcef 100644 --- a/tests/test_litellm/router_utils/test_fallback_event_handlers.py +++ b/tests/test_litellm/router_utils/test_fallback_event_handlers.py @@ -1,6 +1,6 @@ import json from datetime import datetime, timedelta -from typing import NoReturn +from typing import Final, NoReturn from unittest.mock import MagicMock, patch import httpx @@ -1323,3 +1323,20 @@ class TestOrderedFallbackLookupGroups: assert get_fallback_model_group_for_lookup_groups(fallbacks, ("tier9", "smart-router")) == (["backup-b"], None) assert get_fallback_model_group_for_lookup_groups(fallbacks, ("tier9", "no-such")) == (["backup-c"], 2) assert get_fallback_model_group_for_lookup_groups([{"tier1": ["backup-a"]}], ("no", "nope")) == (None, None) + + +class TestHasUnattemptedFallbackTarget: + def test_exhausted_chain_is_not_recoverable_but_a_fresh_entry_is(self): + from litellm.router_utils.fallback_event_handlers import ( + has_unattempted_fallback_target, + ) + + attempted: Final = AttemptedFallbackTargets() + attempted.record("primary") + attempted.record("fb1") + attempted.record("fb2") + + assert has_unattempted_fallback_target(["fb1", "fb2"], {"attempted_targets": attempted}) is False + assert has_unattempted_fallback_target(["fb1", "fb3"], {"attempted_targets": attempted}) is True + assert has_unattempted_fallback_target(["fb1"], {}) is True + assert has_unattempted_fallback_target(None, {}) is False diff --git a/tests/test_litellm/test_router.py b/tests/test_litellm/test_router.py index cf143744067..e31316009fa 100644 --- a/tests/test_litellm/test_router.py +++ b/tests/test_litellm/test_router.py @@ -3522,6 +3522,38 @@ async def test_acompletion_mid_stream_fallback_walks_every_entry_of_the_configur assert attempted_model_groups == ["primary", "fb1", "fb2"] +def test_refusal_on_the_last_fallback_hop_is_returned_instead_of_raised(): + """LIT-7400 follow-up: a refusal on the final hop of an exhausted list passes through.""" + from litellm.router_utils.fallback_event_handlers import AttemptedFallbackTargets + + router = litellm.Router( + model_list=[ + {"model_name": "primary", "litellm_params": {"model": "openai/primary-model", "api_key": "fake-key"}}, + {"model_name": "fb1", "litellm_params": {"model": "openai/fb1-model", "api_key": "fake-key"}}, + {"model_name": "fb2", "litellm_params": {"model": "openai/fb2-model", "api_key": "fake-key"}}, + ], + fallbacks=[{"primary": ["fb1", "fb2"]}], + num_retries=0, + ) + + attempted: Final = AttemptedFallbackTargets() + attempted.record("primary") + attempted.record("fb1") + attempted.record("fb2") + kwargs: Final = { + "attempted_targets": attempted, + "metadata": {"model_group": "fb2", "original_model_group": "primary"}, + } + + assert router._refusal_fallback_available("fb2", kwargs) is False + assert ( + router._refusal_fallback_available( + "fb1", {"metadata": {"model_group": "fb1", "original_model_group": "primary"}} + ) + is True + ) + + def test_completion_streaming_iterator_adopts_fallback_response_headers(): """LIT-6767, sync counterpart of the fallback-adoption test.""" from unittest.mock import MagicMock, patch From 45d22dc5e133dbfb6ee76edc79572f0c82f3ec40 Mon Sep 17 00:00:00 2001 From: Yuneng Jiang Date: Mon, 21 Sep 2026 12:50:58 -0700 Subject: [PATCH 20/44] test(migrations): cover the release-to-release upgrade path The migration e2e harness only ever used one image: it seeded the database with the candidate build and then applied synthetic migrations on top. That proves the migration machinery (locking, crash recovery, legacy baselining, pooling) but never executes the real schema of release N against the real migrations of release N+1, which is the path operators actually run. Adds a baseline image alongside the candidate, so a test can seed with a published release and upgrade with the build under test. Suites: - test_upgrade.py: the candidate applies the pending release migrations, keys minted by the baseline release survive, and concurrent replicas upgrade a baseline database exactly once. - test_rolling_upgrade.py: a baseline replica keeps serving virtual-key auth while the candidate migrates underneath it, and both releases serve and resolve each other's keys during the overlap. This is the reported failure: a new column on LiteLLM_VerificationToken invalidates prepared plans on pods still running the old release, which the proxy reads whole-row, and auth starts failing until those pods leave service. - test_shaped_database.py: the upgrade completes and preserves rows on a populated spend log, rather than on the empty database every other migration test starts from. Every upgrade assertion is gated on the candidate having actually applied migrations the baseline had not, so a stale pin fails loudly instead of passing on an empty delta. CI adds two jobs to the migration_startup workflow. The baseline defaults to a committed release pin and is overridable per pipeline, matching how migration_candidate_image already works; only the upgrade jobs pull it. Verified against a real v1.101.0 -> v1.102.0 upgrade: 6 passed, with the baseline seeding 165 migrations and the candidate applying the 6 that landed between the two releases. --- .circleci/config.yml | 29 ++++- .circleci/scripts/run_migration_tests.py | 3 + tests/e2e/migrations/conftest.py | 29 +++++ tests/e2e/migrations/containers.py | 5 +- tests/e2e/migrations/test_rolling_upgrade.py | 55 +++++++++ tests/e2e/migrations/test_shaped_database.py | 43 +++++++ tests/e2e/migrations/test_upgrade.py | 47 ++++++++ tests/e2e/migrations/upgrade.py | 117 +++++++++++++++++++ 8 files changed, 326 insertions(+), 2 deletions(-) create mode 100644 tests/e2e/migrations/test_rolling_upgrade.py create mode 100644 tests/e2e/migrations/test_shaped_database.py create mode 100644 tests/e2e/migrations/test_upgrade.py create mode 100644 tests/e2e/migrations/upgrade.py diff --git a/.circleci/config.yml b/.circleci/config.yml index cc9aa7fe1c4..83acf0ac1c8 100644 --- a/.circleci/config.yml +++ b/.circleci/config.yml @@ -6,6 +6,9 @@ parameters: migration_candidate_image: type: string default: "" + migration_baseline_image: + type: string + default: "ghcr.io/berriai/litellm-database:v1.102.0" migration_source_sha: type: string default: "" @@ -2946,7 +2949,10 @@ jobs: parameters: suite: type: enum - enum: [startup, recovery, legacy] + enum: [startup, recovery, legacy, upgrade, shaped] + baseline: + type: boolean + default: false machine: image: ubuntu-2204:2024.04.1 resource_class: large @@ -2954,6 +2960,7 @@ jobs: environment: LITELLM_MIGRATION_TESTS: "1" LITELLM_MIGRATION_TEST_IMAGE: litellm-docker-database:ci + LITELLM_MIGRATION_BASELINE_IMAGE: << pipeline.parameters.migration_baseline_image >> MIGRATION_TEST_ADMIN_URL: postgresql://postgres:postgres@127.0.0.1:5432/postgres MIGRATION_TEST_CONTAINER_ADMIN_URL: postgresql://postgres:postgres@host.docker.internal:5432/postgres MIGRATION_TEST_OUTPUT: /tmp/migration-results @@ -2981,6 +2988,16 @@ jobs: - wait_for_service: url: tcp://localhost:5432 timeout: "60" + - when: + condition: << parameters.baseline >> + steps: + - run: + name: Pull the baseline release the upgrade starts from + environment: + BASELINE_IMAGE: << pipeline.parameters.migration_baseline_image >> + command: | + [[ "$BASELINE_IMAGE" =~ ^ghcr.io/berriai/[a-z0-9._/-]+(@sha256:[0-9a-f]{64}|:v[0-9][0-9a-z.-]*)$ ]] || exit 1 + docker pull "$BASELINE_IMAGE" - run: name: Run migration startup regressions environment: @@ -3188,6 +3205,16 @@ workflows: name: migration-legacy-and-pooling suite: legacy requires: [build_docker_database_image] + - migration_startup_tests: + name: migration-upgrade + suite: upgrade + baseline: true + requires: [build_docker_database_image] + - migration_startup_tests: + name: migration-upgrade-shaped + suite: shaped + baseline: true + requires: [build_docker_database_image] migration_startup_scheduled: triggers: - schedule: diff --git a/.circleci/scripts/run_migration_tests.py b/.circleci/scripts/run_migration_tests.py index 56029c406fb..5a73c54e3f5 100644 --- a/.circleci/scripts/run_migration_tests.py +++ b/.circleci/scripts/run_migration_tests.py @@ -13,6 +13,8 @@ SUITES: Final = { "startup": (("test_startup.py",), 12), "recovery": (("test_recovery.py",), 15), "legacy": (("test_legacy.py", "test_pooling.py"), 11), + "upgrade": (("test_upgrade.py", "test_rolling_upgrade.py"), 5), + "shaped": (("test_shaped_database.py",), 1), } @@ -93,6 +95,7 @@ def main() -> int: { **metadata, "suite": suite, + "baseline_image": os.environ.get("LITELLM_MIGRATION_BASELINE_IMAGE", ""), "expected_cases": expected, "passed": passed, "pytest_exit_code": result.returncode, diff --git a/tests/e2e/migrations/conftest.py b/tests/e2e/migrations/conftest.py index 735adeedbdb..b7604a4fdda 100644 --- a/tests/e2e/migrations/conftest.py +++ b/tests/e2e/migrations/conftest.py @@ -60,3 +60,32 @@ def containers(migration_image: str, tmp_path: Path, request: SubRequest) -> Con output: Final = Path(configured) / request.node.name if configured else tmp_path output.mkdir(parents=True, exist_ok=True) return Containers(migration_image, output) + + +@pytest.fixture(scope="session") +def baseline_image(tmp_path_factory: pytest.TempPathFactory) -> str: + configured: Final = os.environ.get("LITELLM_MIGRATION_BASELINE_IMAGE") + assert configured, "LITELLM_MIGRATION_BASELINE_IMAGE must name the released image the upgrade starts from" + image: Final = docker("image", "inspect", configured, "--format", "{{.Id}}") + assert image.startswith("sha256:"), "Unable to identify the baseline image" + output: Final = Path(os.environ.get("MIGRATION_TEST_OUTPUT", str(tmp_path_factory.getbasetemp()))) + output.mkdir(parents=True, exist_ok=True) + (output / "baseline-image.json").write_text(json.dumps({"requested": configured, "image_id": image})) + return image + + +@pytest.fixture(scope="session") +def baseline_template( + databases: Databases, baseline_image: str, tmp_path_factory: pytest.TempPathFactory +) -> Iterator[Database]: + output: Final = Path(os.environ.get("MIGRATION_TEST_OUTPUT", str(tmp_path_factory.getbasetemp()))) / "baseline-seed" + with databases.create() as database: + with Containers(baseline_image, output).start(database) as replica: + ready((replica,), database) + yield database + + +@pytest.fixture +def baseline_database(databases: Databases, baseline_template: Database) -> Iterator[Database]: + with databases.create(baseline_template) as database: + yield database diff --git a/tests/e2e/migrations/containers.py b/tests/e2e/migrations/containers.py index 0f5793b81dd..dd126b994d3 100644 --- a/tests/e2e/migrations/containers.py +++ b/tests/e2e/migrations/containers.py @@ -5,7 +5,7 @@ import subprocess import time from collections.abc import Callable, Generator, Mapping from contextlib import contextmanager -from dataclasses import dataclass +from dataclasses import dataclass, replace from pathlib import Path from typing import Final from uuid import uuid4 @@ -123,6 +123,9 @@ class Containers: image: str output: Path + def using(self, image: str) -> "Containers": + return replace(self, image=image) + @contextmanager def start( self, diff --git a/tests/e2e/migrations/test_rolling_upgrade.py b/tests/e2e/migrations/test_rolling_upgrade.py new file mode 100644 index 00000000000..4d60d20df4a --- /dev/null +++ b/tests/e2e/migrations/test_rolling_upgrade.py @@ -0,0 +1,55 @@ +from typing import Final + +import pytest + +from .containers import Containers, ready +from .database import Database +from .upgrade import ( + CACHED_PLAN, + assert_history_clean, + assert_upgraded, + auth_traffic, + confirm, + keep_serving, + migration_names, + provision, +) + +pytestmark: Final = [pytest.mark.e2e, pytest.mark.migration_startup] + + +class TestRollingUpgrade: + def test_baseline_replica_keeps_serving_while_the_candidate_migrates( + self, containers: Containers, baseline_image: str, baseline_database: Database + ) -> None: + with containers.using(baseline_image).start(baseline_database) as old: + ready((old,), baseline_database) + key, _ = provision(old) + before: Final = migration_names(baseline_database) + with auth_traffic(old, key) as traffic: + keep_serving(traffic, "the baseline replica authenticating before the upgrade") + with containers.start(baseline_database) as new: + ready((new,), baseline_database) + assert_upgraded(before, migration_names(baseline_database)) + keep_serving(traffic, "the baseline replica authenticating after the schema moved") + assert_history_clean(baseline_database) + assert CACHED_PLAN not in old.logs(), "The baseline replica hit a stale prepared statement" + assert old.state().Running, "The baseline replica died during the upgrade" + + def test_both_releases_serve_and_share_keys_during_the_overlap( + self, containers: Containers, baseline_image: str, baseline_database: Database + ) -> None: + with containers.using(baseline_image).start(baseline_database) as old: + ready((old,), baseline_database) + old_key, old_alias = provision(old) + before: Final = migration_names(baseline_database) + with containers.start(baseline_database) as new: + ready((new,), baseline_database) + assert_upgraded(before, migration_names(baseline_database)) + new_key, new_alias = provision(new) + with auth_traffic(old, old_key) as old_traffic, auth_traffic(new, new_key) as new_traffic: + keep_serving(old_traffic, "the baseline replica serving through the overlap") + keep_serving(new_traffic, "the candidate replica serving through the overlap") + confirm(old, new_key, new_alias) + confirm(new, old_key, old_alias) + assert CACHED_PLAN not in old.logs(), "The baseline replica hit a stale prepared statement" diff --git a/tests/e2e/migrations/test_shaped_database.py b/tests/e2e/migrations/test_shaped_database.py new file mode 100644 index 00000000000..20c4368ae33 --- /dev/null +++ b/tests/e2e/migrations/test_shaped_database.py @@ -0,0 +1,43 @@ +from typing import Final + +import pytest + +from .containers import Containers, ready +from .database import Database +from .upgrade import assert_history_clean, assert_upgraded, confirm, migration_names, provision + +SPEND_ROWS: Final = 20_000 + +pytestmark: Final = [pytest.mark.e2e, pytest.mark.migration_startup] + + +def seed_spend_logs(database: Database, rows: int) -> None: + database.execute( + 'INSERT INTO "LiteLLM_SpendLogs" (request_id, call_type, "startTime", "endTime") ' + "SELECT 'upgrade-shape-' || g, 'acompletion', now() - (g || ' seconds')::interval, " + "now() - (g || ' seconds')::interval FROM generate_series(1, %s) AS g", + (rows,), + ) + assert database.query('SELECT count(*) FROM "LiteLLM_SpendLogs"') == ((rows,),) + + +class TestPopulatedDatabaseUpgrade: + def test_upgrade_completes_and_preserves_a_populated_spend_log( + self, containers: Containers, baseline_image: str, baseline_database: Database + ) -> None: + with containers.using(baseline_image).start(baseline_database) as old: + ready((old,), baseline_database) + key, alias = provision(old) + seed_spend_logs(baseline_database, SPEND_ROWS) + before: Final = migration_names(baseline_database) + with containers.start(baseline_database) as new: + ready((new,), baseline_database) + assert_upgraded(before, migration_names(baseline_database)) + confirm(new, key, alias) + assert_history_clean(baseline_database) + assert baseline_database.query('SELECT count(*) FROM "LiteLLM_SpendLogs"') == ((SPEND_ROWS,),), ( + "The upgrade lost spend rows" + ) + assert baseline_database.query( + 'SELECT count(*) FROM "LiteLLM_SpendLogs" WHERE "startTime" IS NULL OR "endTime" IS NULL' + ) == ((0,),), "The upgrade nulled timestamps on existing spend rows" diff --git a/tests/e2e/migrations/test_upgrade.py b/tests/e2e/migrations/test_upgrade.py new file mode 100644 index 00000000000..23f0bbe9124 --- /dev/null +++ b/tests/e2e/migrations/test_upgrade.py @@ -0,0 +1,47 @@ +from contextlib import ExitStack +from typing import Final + +import pytest + +from .checks import start_replicas +from .containers import Containers, ready +from .database import Database +from .upgrade import assert_history_clean, assert_upgraded, confirm, migration_names, provision + +pytestmark: Final = [pytest.mark.e2e, pytest.mark.migration_startup] + + +class TestReleaseUpgrade: + def test_candidate_applies_the_pending_release_migrations( + self, containers: Containers, baseline_database: Database + ) -> None: + before: Final = migration_names(baseline_database) + with containers.start(baseline_database) as replica: + ready((replica,), baseline_database) + assert_upgraded(before, migration_names(baseline_database)) + assert_history_clean(baseline_database) + + def test_upgrade_preserves_keys_minted_by_the_baseline_release( + self, containers: Containers, baseline_image: str, baseline_database: Database + ) -> None: + with containers.using(baseline_image).start(baseline_database) as old: + ready((old,), baseline_database) + key, alias = provision(old) + confirm(old, key, alias) + before: Final = migration_names(baseline_database) + with containers.start(baseline_database) as new: + ready((new,), baseline_database) + assert_upgraded(before, migration_names(baseline_database)) + confirm(new, key, alias) + + def test_concurrent_replicas_upgrade_a_baseline_database_once( + self, containers: Containers, baseline_database: Database + ) -> None: + before: Final = migration_names(baseline_database) + with ExitStack() as stack: + ready(start_replicas(stack, containers, baseline_database), baseline_database) + assert_upgraded(before, migration_names(baseline_database)) + assert_history_clean(baseline_database) + assert baseline_database.query("SELECT count(*) FROM _prisma_migrations WHERE applied_steps_count > 1") == ( + (0,), + ), "A migration was executed more than once across the upgrading replicas" diff --git a/tests/e2e/migrations/upgrade.py b/tests/e2e/migrations/upgrade.py new file mode 100644 index 00000000000..2758d0fb974 --- /dev/null +++ b/tests/e2e/migrations/upgrade.py @@ -0,0 +1,117 @@ +from __future__ import annotations + +import threading +from collections.abc import Generator +from contextlib import contextmanager +from dataclasses import dataclass, field +from typing import Final +from uuid import uuid4 + +from e2e_http import Result, Success, unwrap +from models import ( + KeyGenerateBody, + KeyGenerateResponse, + KeyInfoParams, + KeyInfoResponse, + ModelsListParams, + ModelsListResponse, +) +from pydantic import BaseModel + +from .containers import Replica, until +from .database import Database + +CACHED_PLAN: Final = "cached plan must not change result type" + + +def provision(replica: Replica) -> tuple[str, str]: + alias: Final = f"upgrade-{uuid4().hex}" + key: Final = unwrap( + replica.transport.post( + "/key/generate", + headers=replica.transport.master, + json=KeyGenerateBody(key_alias=alias), + response_type=KeyGenerateResponse, + ) + ).key + return key, alias + + +def confirm(replica: Replica, key: str, alias: str) -> None: + info: Final = unwrap( + replica.transport.get( + "/key/info", + headers=replica.transport.master, + params=KeyInfoParams(key=key), + response_type=KeyInfoResponse, + ) + ) + assert info.info.key_alias == alias, "Key minted on one release did not resolve on the other" + + +@dataclass(slots=True) +class Outcomes: + served: int = 0 + failures: list[str] = field(default_factory=list) + + def record(self, result: Result[BaseModel]) -> None: + match result: + case Success(): + self.served += 1 + case _: + self.failures.append(result.model_dump_json()) + + +@contextmanager +def auth_traffic(replica: Replica, key: str, interval: float = 0.05) -> Generator[Outcomes]: + outcomes: Final = Outcomes() + stop: Final = threading.Event() + + def drive() -> None: + while not stop.is_set(): + outcomes.record( + replica.transport.get( + "/v1/models", + headers=replica.transport.bearer(key), + params=ModelsListParams(), + response_type=ModelsListResponse, + timeout=10, + ) + ) + stop.wait(interval) + + thread: Final = threading.Thread(target=drive, name="upgrade-auth-traffic", daemon=True) + thread.start() + try: + yield outcomes + finally: + stop.set() + thread.join(30) + assert not thread.is_alive(), "Auth traffic thread did not stop" + + +def keep_serving(outcomes: Outcomes, description: str, calls: int = 20) -> int: + target: Final = outcomes.served + calls + until(description, lambda: outcomes.served >= target or bool(outcomes.failures)) + assert not outcomes.failures, f"Virtual-key auth failed during {description}: {outcomes.failures[:5]}" + return outcomes.served + + +def migration_names(database: Database) -> frozenset[str]: + return frozenset(str(row[0]) for row in database.query("SELECT migration_name FROM _prisma_migrations")) + + +def assert_history_clean(database: Database) -> None: + assert database.query( + "SELECT count(*) FROM _prisma_migrations WHERE finished_at IS NULL OR rolled_back_at IS NOT NULL" + ) == ((0,),), "The upgrade left an unfinished or rolled-back migration behind" + + +def assert_upgraded(before: frozenset[str], after: frozenset[str]) -> frozenset[str]: + applied: Final = after - before + assert applied, ( + "The candidate applied no migrations the baseline release had not: the pinned " + "LITELLM_MIGRATION_BASELINE_IMAGE is at or ahead of the candidate, so this suite proves nothing" + ) + assert not before - after, "The upgrade removed migration history the baseline release had already applied" + return applied From 9002749e29bcff3de2f8991ee988f923ffc5a3c5 Mon Sep 17 00:00:00 2001 From: Joshua Valluru <326636767+joshua-berri@users.noreply.github.com> Date: Mon, 21 Sep 2026 12:58:06 -0700 Subject: [PATCH 21/44] fix(mcp): preserve Python 3.10 imports and integration test seams --- .../_experimental/mcp_server/operations.py | 4 ++-- tests/mcp_tests/test_mcp_logging.py | 6 ++--- tests/mcp_tests/test_mcp_server.py | 24 ++++++++++--------- 3 files changed, 18 insertions(+), 16 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/operations.py b/litellm/proxy/_experimental/mcp_server/operations.py index 3ee4f37add7..a579c33e702 100644 --- a/litellm/proxy/_experimental/mcp_server/operations.py +++ b/litellm/proxy/_experimental/mcp_server/operations.py @@ -6,7 +6,7 @@ import types import uuid from collections.abc import Mapping, Sequence from datetime import datetime -from typing import Any, Final, NoReturn, TypeAlias, assert_never, overload +from typing import Any, Final, NoReturn, TypeAlias, overload from fastapi import HTTPException from mcp import ReadResourceResult, Resource @@ -34,7 +34,7 @@ from mcp.types import ( ) from mcp.types import Tool as MCPTool from pydantic import AnyUrl, ConfigDict, Field, TypeAdapter -from typing_extensions import ReadOnly, TypedDict +from typing_extensions import ReadOnly, TypedDict, assert_never from litellm._logging import verbose_logger from litellm.constants import ( diff --git a/tests/mcp_tests/test_mcp_logging.py b/tests/mcp_tests/test_mcp_logging.py index ed8829945e5..41d0e2cb59b 100644 --- a/tests/mcp_tests/test_mcp_logging.py +++ b/tests/mcp_tests/test_mcp_logging.py @@ -142,7 +142,7 @@ async def test_mcp_cost_tracking(): local_mcp_server_manager, ), patch( - "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager", + "litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager", local_mcp_server_manager, ), ): @@ -293,7 +293,7 @@ async def test_mcp_cost_tracking_per_tool(): local_mcp_server_manager, ), patch( - "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager", + "litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager", local_mcp_server_manager, ), ): @@ -451,7 +451,7 @@ async def test_mcp_tool_call_hook(): local_mcp_server_manager, ), patch( - "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager", + "litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager", local_mcp_server_manager, ), ): diff --git a/tests/mcp_tests/test_mcp_server.py b/tests/mcp_tests/test_mcp_server.py index 94cf35b675d..2b92367f186 100644 --- a/tests/mcp_tests/test_mcp_server.py +++ b/tests/mcp_tests/test_mcp_server.py @@ -922,7 +922,7 @@ async def test_get_tools_from_mcp_servers(): ) with patch( - "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager", + "litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager", mock_manager, ): # Test with specific servers @@ -950,6 +950,7 @@ async def test_get_tools_from_mcp_servers(): extra_headers=None, add_prefix=False, raw_headers=None, + client_ip=None, user_api_key_auth=None, oauth2_headers=None, ): @@ -966,7 +967,7 @@ async def test_get_tools_from_mcp_servers(): ) with patch( - "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager", + "litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager", mock_manager_2, ): result = await _get_tools_from_mcp_servers( @@ -998,7 +999,7 @@ async def test_get_tools_from_mcp_servers(): ) with patch( - "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager", + "litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager", mock_manager, ): with patch( @@ -1981,6 +1982,7 @@ async def test_get_tools_for_single_server(): extra_headers=None, add_prefix=False, raw_headers=None, + client_ip=None, user_api_key_auth=None, ) @@ -2076,7 +2078,7 @@ async def test_rest_listing_hides_key_grants_dispatch_would_refuse(): with patch( "litellm.proxy._experimental.mcp_server.rest_endpoints.global_mcp_server_manager" ) as mock_manager, patch( - "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager" + "litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager" ) as mock_server_manager, patch.object( MCPRequestHandler, "get_allowed_tools_for_server", @@ -2473,7 +2475,7 @@ async def test_filter_tools_by_allowed_tools_integration(): # Mock the global MCP server manager with patch( - "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager" + "litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager" ) as mock_manager: # Mock manager methods mock_manager.get_allowed_mcp_servers = AsyncMock( @@ -2588,7 +2590,7 @@ async def test_filter_tools_by_disallowed_tools_integration(): # Mock the global MCP server manager with patch( - "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager" + "litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager" ) as mock_manager: # Mock manager methods mock_manager.get_allowed_mcp_servers = AsyncMock( @@ -2689,7 +2691,7 @@ async def test_filter_tools_no_restrictions_integration(): # Mock the global MCP server manager with patch( - "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager" + "litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager" ) as mock_manager: # Mock manager methods mock_manager.get_allowed_mcp_servers = AsyncMock( @@ -2970,10 +2972,10 @@ async def test_call_mcp_tool_uses_manager_permission_lookup(): return_value=mock_server, ) as mock_get_server, patch( - "litellm.proxy._experimental.mcp_server.server.global_mcp_tool_registry" + "litellm.proxy._experimental.mcp_server.operations.global_mcp_tool_registry" ) as mock_tool_registry, patch( - "litellm.proxy._experimental.mcp_server.server._handle_managed_mcp_tool", + "litellm.proxy._experimental.mcp_server.operations._handle_managed_mcp_tool", new_callable=AsyncMock, ) as mock_handle_managed, patch( @@ -3046,10 +3048,10 @@ async def test_call_mcp_tool_resolves_unprefixed_tool_name_and_checks_permission return_value=mock_server, ) as mock_get_server, patch( - "litellm.proxy._experimental.mcp_server.server.global_mcp_tool_registry" + "litellm.proxy._experimental.mcp_server.operations.global_mcp_tool_registry" ) as mock_tool_registry, patch( - "litellm.proxy._experimental.mcp_server.server._handle_managed_mcp_tool", + "litellm.proxy._experimental.mcp_server.operations._handle_managed_mcp_tool", new_callable=AsyncMock, ) as mock_handle_managed, patch( From db7d52eedaf31d9a6d9b964048134d3e9d9bd0e2 Mon Sep 17 00:00:00 2001 From: shivam Date: Mon, 21 Sep 2026 20:01:50 +0000 Subject: [PATCH 22/44] test(router): track attempted fallback groups via the mock call log Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- tests/test_litellm/test_router.py | 11 ++++++----- 1 file changed, 6 insertions(+), 5 deletions(-) diff --git a/tests/test_litellm/test_router.py b/tests/test_litellm/test_router.py index e31316009fa..e642520bdbd 100644 --- a/tests/test_litellm/test_router.py +++ b/tests/test_litellm/test_router.py @@ -3460,8 +3460,6 @@ async def test_acompletion_mid_stream_fallback_walks_every_entry_of_the_configur from litellm.exceptions import MidStreamFallbackError from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper - attempted_model_groups: list[str] = [] - class FailingStream(CustomStreamWrapper): def __init__(self, model: str): super().__init__( @@ -3497,7 +3495,6 @@ async def test_acompletion_mid_stream_fallback_walks_every_entry_of_the_configur raise StopAsyncIteration from None async def fake_acompletion(**kwargs): - attempted_model_groups.append(kwargs["metadata"]["model_group"]) if "fb2" in kwargs["model"]: return OkStream(kwargs["model"]) return FailingStream(kwargs["model"]) @@ -3512,14 +3509,18 @@ async def test_acompletion_mid_stream_fallback_walks_every_entry_of_the_configur num_retries=0, ) - with patch("litellm.acompletion", side_effect=fake_acompletion): + with patch("litellm.acompletion", side_effect=fake_acompletion) as mock_acompletion: response = await router.acompletion(model="primary", messages=[{"role": "user", "content": "hi"}], stream=True) content: Final = "".join( [chunk.choices[0].delta.content or "" async for chunk in response if chunk is not None] ) assert content == "ok-from-openai/fb2-model" - assert attempted_model_groups == ["primary", "fb1", "fb2"] + assert [c.kwargs["metadata"]["model_group"] for c in mock_acompletion.call_args_list] == [ + "primary", + "fb1", + "fb2", + ] def test_refusal_on_the_last_fallback_hop_is_returned_instead_of_raised(): From ef67412e50c9f3250eb473a0a14fbf00617e36ea Mon Sep 17 00:00:00 2001 From: Joshua Valluru <326636767+joshua-berri@users.noreply.github.com> Date: Mon, 21 Sep 2026 13:13:20 -0700 Subject: [PATCH 23/44] fix(mcp): keep OAuth prefetch failure logs free of caller data --- .../_experimental/mcp_server/operations.py | 6 ++--- .../proxy/_experimental/mcp_server/server.py | 3 +++ .../mcp_server/test_operations.py | 22 +++++++++++++++++++ 3 files changed, 28 insertions(+), 3 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/operations.py b/litellm/proxy/_experimental/mcp_server/operations.py index a579c33e702..fcee3483e15 100644 --- a/litellm/proxy/_experimental/mcp_server/operations.py +++ b/litellm/proxy/_experimental/mcp_server/operations.py @@ -763,8 +763,8 @@ async def _prefetch_oauth_creds_for_user( ) creds: Final = await list_user_oauth_credentials(prisma_client, user_id) return {c["server_id"]: c for c in creds if "server_id" in c} - except Exception as e: - verbose_logger.warning("_prefetch_oauth_creds_for_user: failed to prefetch for user=%s: %s", user_id, e) + except Exception: + verbose_logger.warning("_prefetch_oauth_creds_for_user: failed to prefetch OAuth credentials") return {} @@ -3099,4 +3099,4 @@ class GatewayOperations: case ReadResourceRequest(params=params): return await _execute_read_resource(context, params, self._host_progress_callback) case _: - assert_never(operation) + return assert_never(operation) diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 33978bb9182..a57eed04419 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -416,7 +416,10 @@ def _proxy_exception_to_http_exception(exc: ProxyException) -> HTTPException: if MCP_AVAILABLE: __all__ = ( "_MCP_CREDENTIAL_REQUEST_FIELDS", + "BlobResourceContents", "ListMCPToolsRestAPIResponseObject", + "ResourceTemplate", + "TextResourceContents", "_McpDeniedDetail", "_aggregate_server_key", "_build_virtual_call_logging_obj", diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_operations.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_operations.py index 870f2cd5f8f..abb925ddc77 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_operations.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_operations.py @@ -8,6 +8,28 @@ from litellm.proxy._experimental.mcp_server.operations import GatewayOperations, from litellm.proxy._types import UserAPIKeyAuth +@pytest.mark.asyncio +async def test_oauth_prefetch_failure_does_not_log_caller_or_exception_text(caplog): + from litellm.proxy._experimental.mcp_server.operations import _prefetch_oauth_creds_for_user + + user_id = "caller\nFORGED-USER-LINE" + fetch = AsyncMock(side_effect=RuntimeError("database\nFORGED-ERROR-LINE")) + database = object() + with ( + patch("litellm.proxy.utils.get_prisma_client_or_throw", return_value=database), + patch("litellm.proxy._experimental.mcp_server.db.list_user_oauth_credentials", fetch), + caplog.at_level("WARNING", logger="LiteLLM"), + ): + result = await _prefetch_oauth_creds_for_user(UserAPIKeyAuth(user_id=user_id)) + assert result == {} + fetch.assert_awaited_once_with(database, user_id) + warnings = [record.getMessage() for record in caplog.records if "prefetch" in record.getMessage()] + assert len(warnings) == 1 + assert "failed" in warnings[0] + assert "\n" not in warnings[0] + assert "FORGED" not in warnings[0] + + @pytest.mark.asyncio async def test_dispatch_uses_explicit_context_when_ambient_caller_differs(): from mcp.server.auth.middleware.auth_context import auth_context_var From dc85812971ef1ba5a143f9ee87a8ef5f395852dd Mon Sep 17 00:00:00 2001 From: Yuneng Jiang Date: Mon, 21 Sep 2026 13:15:46 -0700 Subject: [PATCH 24/44] test(migrations): close the gaps the upgrade assertions left open Three holes in the new suite, all of which let a test pass without proving what its name claims: - A migration recorded twice, once per replica, each with applied_steps_count = 1, slipped past both the step-count check and migration_names(), which collapses the history into a set. Reject duplicate migration_name rows outright. - auth_traffic only asserted the failures it had seen by the time keep_serving hit its target. A request failing after that, or on the other replica while the test waited on one stream, was recorded and never read. Assert the recorded failures once the thread has joined. - The rolling test warmed the baseline replica's virtual-key cache before the upgrade, and that cache holds for 60 seconds by default (UserAPIKeyCacheTTLEnum.in_memory_cache_ttl). The candidate migrates well inside that window, so the post-upgrade requests could be served from cache without ever repeating the whole-row token lookup that the stale prepared statement breaks. Drive the baseline replica with a key minted after the schema moved, which it has never seen and must resolve from the database. Re-ran against v1.101.0 -> v1.102.0: 6 passed. --- tests/e2e/migrations/test_rolling_upgrade.py | 2 ++ tests/e2e/migrations/upgrade.py | 7 +++++++ 2 files changed, 9 insertions(+) diff --git a/tests/e2e/migrations/test_rolling_upgrade.py b/tests/e2e/migrations/test_rolling_upgrade.py index 4d60d20df4a..5ad74e0ba8c 100644 --- a/tests/e2e/migrations/test_rolling_upgrade.py +++ b/tests/e2e/migrations/test_rolling_upgrade.py @@ -32,6 +32,8 @@ class TestRollingUpgrade: ready((new,), baseline_database) assert_upgraded(before, migration_names(baseline_database)) keep_serving(traffic, "the baseline replica authenticating after the schema moved") + with auth_traffic(old, provision(new)[0]) as uncached: + keep_serving(uncached, "the baseline replica resolving a key minted after the schema moved") assert_history_clean(baseline_database) assert CACHED_PLAN not in old.logs(), "The baseline replica hit a stale prepared statement" assert old.state().Running, "The baseline replica died during the upgrade" diff --git a/tests/e2e/migrations/upgrade.py b/tests/e2e/migrations/upgrade.py index 2758d0fb974..2123f86450a 100644 --- a/tests/e2e/migrations/upgrade.py +++ b/tests/e2e/migrations/upgrade.py @@ -88,6 +88,9 @@ def auth_traffic(replica: Replica, key: str, interval: float = 0.05) -> Generato stop.set() thread.join(30) assert not thread.is_alive(), "Auth traffic thread did not stop" + assert not outcomes.failures, ( + f"Virtual-key auth failed on {replica.name} after the traffic window closed: {outcomes.failures[:5]}" + ) def keep_serving(outcomes: Outcomes, description: str, calls: int = 20) -> int: @@ -105,6 +108,10 @@ def assert_history_clean(database: Database) -> None: assert database.query( "SELECT count(*) FROM _prisma_migrations WHERE finished_at IS NULL OR rolled_back_at IS NOT NULL" ) == ((0,),), "The upgrade left an unfinished or rolled-back migration behind" + assert database.query( + "SELECT count(*) FROM (SELECT migration_name FROM _prisma_migrations GROUP BY migration_name " + "HAVING count(*) > 1) duplicated" + ) == ((0,),), "A migration was recorded more than once, so it ran on more than one replica" def assert_upgraded(before: frozenset[str], after: frozenset[str]) -> frozenset[str]: From 56d1f042ef3de13f27d98a9acecf9bc7b5f5c3b3 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Mon, 21 Sep 2026 13:16:58 -0700 Subject: [PATCH 25/44] fix(types): read upstream headers through a typed helper --- litellm/llms/custom_httpx/http_handler.py | 5 +++++ .../_experimental/mcp_server/openapi_to_mcp_generator.py | 3 ++- litellm/rag/ingestion/gemini_ingestion.py | 3 ++- 3 files changed, 9 insertions(+), 2 deletions(-) diff --git a/litellm/llms/custom_httpx/http_handler.py b/litellm/llms/custom_httpx/http_handler.py index 6b90394043f..b7b2477e85c 100644 --- a/litellm/llms/custom_httpx/http_handler.py +++ b/litellm/llms/custom_httpx/http_handler.py @@ -479,6 +479,11 @@ def _safe_get_response_text(response: httpx.Response) -> str: return "" +def header_value(headers: Mapping[str, str], name: str) -> str | None: + """Read one header as ``str | None``; ``httpx.Headers.get`` itself is typed ``Any``.""" + return headers.get(name) + + async def _safe_aread_response(response: httpx.Response, timeout: float | None = None) -> bytes: """Safely read async response body, falling back to empty bytes on errors.""" try: diff --git a/litellm/proxy/_experimental/mcp_server/openapi_to_mcp_generator.py b/litellm/proxy/_experimental/mcp_server/openapi_to_mcp_generator.py index 3c26937bfac..1247ff1ac28 100644 --- a/litellm/proxy/_experimental/mcp_server/openapi_to_mcp_generator.py +++ b/litellm/proxy/_experimental/mcp_server/openapi_to_mcp_generator.py @@ -49,6 +49,7 @@ from litellm.litellm_core_utils.url_utils import async_safe_get from litellm.llms.custom_httpx.http_handler import ( AsyncHTTPHandler, get_async_httpx_client, + header_value, httpxSpecialProvider, ) from litellm.proxy._experimental.mcp_server.tool_registry import ( @@ -457,7 +458,7 @@ def _raise_for_upstream_failure( if response.status_code == 401 and relays_upstream_auth: raise MCPUpstreamAuthError( status_code=response.status_code, - www_authenticate=dict(response.headers).get("www-authenticate"), + www_authenticate=header_value(response.headers, "www-authenticate"), server_name=upstream, ) raise MCPOpenApiUpstreamError(response.status_code, upstream) diff --git a/litellm/rag/ingestion/gemini_ingestion.py b/litellm/rag/ingestion/gemini_ingestion.py index b81c2cc0ebe..1cf5db549e4 100644 --- a/litellm/rag/ingestion/gemini_ingestion.py +++ b/litellm/rag/ingestion/gemini_ingestion.py @@ -12,6 +12,7 @@ from typing import TYPE_CHECKING, Final, cast from litellm._logging import verbose_logger from litellm.llms.custom_httpx.http_handler import ( get_async_httpx_client, + header_value, httpxSpecialProvider, ) from litellm.llms.gemini.common_utils import GeminiModelInfo @@ -277,7 +278,7 @@ class GeminiRAGIngestion(BaseRAGIngestion): raise Exception(error_msg) verbose_logger.debug("Initiate resumable upload response: %s", response.headers) # Extract upload URL from response headers - upload_url: Final = dict(response.headers).get("x-goog-upload-url") + upload_url: Final = header_value(response.headers, "x-goog-upload-url") if not upload_url: raise Exception("No upload URL returned in response headers") From ae69a8c79a82d54d50436c1eb06ab9cca3625574 Mon Sep 17 00:00:00 2001 From: Yujong Lee Date: Mon, 21 Sep 2026 20:30:48 +0000 Subject: [PATCH 26/44] feat(rust): add Azure Key Vault secret manager backend Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .github/workflows/test-rust.yml | 2 +- litellm-rust/Cargo.lock | 21 ++ litellm-rust/Cargo.toml | 1 + litellm-rust/crates/secrets-azure/Cargo.toml | 24 ++ .../crates/secrets-azure/src/error.rs | 25 ++ .../crates/secrets-azure/src/key_vault.rs | 122 ++++++++++ litellm-rust/crates/secrets-azure/src/lib.rs | 7 + .../tests/fixtures/key_vault_parity.json | 8 + .../crates/secrets-azure/tests/key_vault.rs | 220 ++++++++++++++++++ .../crates/secrets-azure/tests/live.rs | 30 +++ litellm-rust/crates/secrets/Cargo.toml | 2 + litellm-rust/crates/secrets/src/error.rs | 3 + litellm-rust/crates/secrets/src/handler.rs | 9 + litellm-rust/crates/secrets/src/lib.rs | 2 + litellm-rust/crates/secrets/tests/handler.rs | 61 +++++ .../test_secret_manager_handler.py | 105 +++++++++ 16 files changed, 641 insertions(+), 1 deletion(-) create mode 100644 litellm-rust/crates/secrets-azure/Cargo.toml create mode 100644 litellm-rust/crates/secrets-azure/src/error.rs create mode 100644 litellm-rust/crates/secrets-azure/src/key_vault.rs create mode 100644 litellm-rust/crates/secrets-azure/src/lib.rs create mode 100644 litellm-rust/crates/secrets-azure/tests/fixtures/key_vault_parity.json create mode 100644 litellm-rust/crates/secrets-azure/tests/key_vault.rs create mode 100644 litellm-rust/crates/secrets-azure/tests/live.rs create mode 100644 tests/test_litellm/secret_managers/test_secret_manager_handler.py diff --git a/.github/workflows/test-rust.yml b/.github/workflows/test-rust.yml index 278fa7c425f..e7cb9984676 100644 --- a/.github/workflows/test-rust.yml +++ b/.github/workflows/test-rust.yml @@ -130,7 +130,7 @@ jobs: - name: Test secret manager feature combinations run: | cargo test -p litellm-auth-gcp --locked --no-default-features - for features in '' aws google aws,google; do + for features in '' aws google azure aws,google aws,azure google,azure aws,google,azure; do cargo test -p litellm-secrets --locked --no-default-features --features "$features" done diff --git a/litellm-rust/Cargo.lock b/litellm-rust/Cargo.lock index ed4ae4e3353..f7b1e371c0b 100644 --- a/litellm-rust/Cargo.lock +++ b/litellm-rust/Cargo.lock @@ -2697,6 +2697,7 @@ dependencies = [ "jsonwebtoken", "litellm-core-utils", "litellm-secrets-aws", + "litellm-secrets-azure", "litellm-secrets-google", "litellm-secrets-types", "moka", @@ -2731,6 +2732,26 @@ dependencies = [ "wiremock", ] +[[package]] +name = "litellm-secrets-azure" +version = "0.1.0" +dependencies = [ + "litellm-auth-azure", + "litellm-auth-types", + "litellm-core-utils", + "litellm-secrets-types", + "percent-encoding", + "reqwest 0.12.28", + "rstest", + "serde", + "serde_json", + "sha2 0.10.9", + "thiserror 2.0.19", + "tokio", + "veil", + "wiremock", +] + [[package]] name = "litellm-secrets-google" version = "0.1.0" diff --git a/litellm-rust/Cargo.toml b/litellm-rust/Cargo.toml index 570d0dd3568..802a29898d3 100644 --- a/litellm-rust/Cargo.toml +++ b/litellm-rust/Cargo.toml @@ -22,6 +22,7 @@ litellm-secrets = { path = "crates/secrets" } litellm-secrets-types = { path = "crates/secrets-types" } litellm-secrets-aws = { path = "crates/secrets-aws" } litellm-secrets-google = { path = "crates/secrets-google" } +litellm-secrets-azure = { path = "crates/secrets-azure" } litellm-http = { path = "crates/http" } litellm-llms = { path = "crates/llms" } litellm-types = { path = "crates/types" } diff --git a/litellm-rust/crates/secrets-azure/Cargo.toml b/litellm-rust/crates/secrets-azure/Cargo.toml new file mode 100644 index 00000000000..d559592cb34 --- /dev/null +++ b/litellm-rust/crates/secrets-azure/Cargo.toml @@ -0,0 +1,24 @@ +[package] +name = "litellm-secrets-azure" +version = "0.1.0" +edition.workspace = true +license.workspace = true +repository.workspace = true + +[dependencies] +litellm-auth-azure.workspace = true +litellm-auth-types.workspace = true +litellm-secrets-types.workspace = true +litellm-core-utils.workspace = true +reqwest.workspace = true +serde.workspace = true +serde_json.workspace = true +thiserror.workspace = true +veil.workspace = true +percent-encoding = "2.3" + +[dev-dependencies] +tokio.workspace = true +wiremock = "0.6.5" +rstest.workspace = true +sha2.workspace = true diff --git a/litellm-rust/crates/secrets-azure/src/error.rs b/litellm-rust/crates/secrets-azure/src/error.rs new file mode 100644 index 00000000000..9b20efe4f7c --- /dev/null +++ b/litellm-rust/crates/secrets-azure/src/error.rs @@ -0,0 +1,25 @@ +#[derive(thiserror::Error, veil::Redact)] +pub enum Error { + #[error("{0} environment variable is missing")] + MissingEnvironment(&'static str), + #[error("AZURE_KEY_VAULT_URI is not a valid https vault URL")] + VaultUri, + #[error("Azure Key Vault credentials are not configured")] + MissingCredentials, + #[error(transparent)] + Auth( + #[from] + #[redact] + litellm_auth_types::Error, + ), + #[error("Azure Key Vault request failed")] + Http( + #[source] + #[redact] + reqwest::Error, + ), + #[error("Azure Key Vault returned HTTP {0}")] + Status(u16), + #[error("Azure Key Vault response is missing the secret value")] + MissingValue, +} diff --git a/litellm-rust/crates/secrets-azure/src/key_vault.rs b/litellm-rust/crates/secrets-azure/src/key_vault.rs new file mode 100644 index 00000000000..262b19386df --- /dev/null +++ b/litellm-rust/crates/secrets-azure/src/key_vault.rs @@ -0,0 +1,122 @@ +use std::sync::Arc; + +use litellm_auth_azure::{AzureAuthInputs, AzureAuthService}; +use litellm_core_utils::settings::Lookup; +use litellm_secrets_types::{Secret, SecretValue}; +use serde::Deserialize; + +use crate::Error; + +const AZURE_KEY_VAULT_URI: &str = "AZURE_KEY_VAULT_URI"; +const API_VERSION: &str = "7.4"; + +#[derive(Clone)] +pub struct AzureKeyVault { + client: reqwest::Client, + vault: reqwest::Url, + auth: Arc, + inputs: AzureAuthInputs, + environment: Arc, +} + +#[derive(Deserialize)] +struct SecretResponse { + value: Option, +} + +impl AzureKeyVault { + pub fn with_client( + client: reqwest::Client, + vault: reqwest::Url, + environment: Arc, + ) -> Result { + if vault.host_str().is_none() { + return Err(Error::VaultUri); + } + let scope = scope_for(&vault); + let inputs = AzureAuthInputs::from_sourced_optional_params( + serde_json::json!({ + "azure_scope": scope, + "enable_azure_ad_token_refresh": true, + }) + .as_object() + .expect("static Azure auth inputs object"), + &std::collections::BTreeMap::new(), + )?; + Ok(Self { + client, + vault, + auth: Arc::new(AzureAuthService::default()), + inputs, + environment, + }) + } + + pub fn new(environment: Arc) -> Result { + let value = environment + .get(AZURE_KEY_VAULT_URI) + .ok_or(Error::MissingEnvironment(AZURE_KEY_VAULT_URI))?; + let vault = reqwest::Url::parse(&value).map_err(|_| Error::VaultUri)?; + if vault.scheme() != "https" || vault.host_str().is_none() { + return Err(Error::VaultUri); + } + Self::with_client(reqwest::Client::new(), vault, environment) + } + + pub fn scope(&self) -> &str { + self.inputs + .azure_scope + .as_value() + .map(|value| value.value().as_str()) + .unwrap_or_default() + } + + pub async fn get_secret_from_azure_key_vault( + &self, + name: &str, + ) -> Result, Error> { + let token = self + .auth + .get_azure_ad_token(&self.inputs, &|key| self.environment.get(key)) + .await? + .ok_or(Error::MissingCredentials)?; + let encoded_name = encode_name(name); + let url = self + .vault + .join(&format!("secrets/{encoded_name}?api-version={API_VERSION}")) + .map_err(|_| Error::VaultUri)?; + let response = self + .client + .get(url) + .bearer_auth(token.value().secret().expose()) + .send() + .await + .map_err(Error::Http)?; + if response.status() == reqwest::StatusCode::NOT_FOUND { + return Ok(None); + } + if response.status() != reqwest::StatusCode::OK { + return Err(Error::Status(response.status().as_u16())); + } + let payload: SecretResponse = response.json().await.map_err(Error::Http)?; + let value = payload.value.ok_or(Error::MissingValue)?; + Ok(Some(Secret::String(SecretValue::new(value)))) + } +} + +fn scope_for(vault: &reqwest::Url) -> String { + let host = vault.host_str().unwrap_or_default(); + let resource = host + .split_once('.') + .map_or(host, |(_, remainder)| remainder); + format!("https://{resource}/.default") +} + +fn encode_name(name: &str) -> String { + percent_encoding::utf8_percent_encode(name, percent_encoding::NON_ALPHANUMERIC) + .to_string() + .replace("%2D", "-") + .replace("%2E", ".") + .replace("%5F", "_") + .replace("%7E", "~") +} diff --git a/litellm-rust/crates/secrets-azure/src/lib.rs b/litellm-rust/crates/secrets-azure/src/lib.rs new file mode 100644 index 00000000000..c0094fc033b --- /dev/null +++ b/litellm-rust/crates/secrets-azure/src/lib.rs @@ -0,0 +1,7 @@ +#![forbid(unsafe_code)] + +mod error; +mod key_vault; + +pub use error::Error; +pub use key_vault::AzureKeyVault; diff --git a/litellm-rust/crates/secrets-azure/tests/fixtures/key_vault_parity.json b/litellm-rust/crates/secrets-azure/tests/fixtures/key_vault_parity.json new file mode 100644 index 00000000000..c4a83cd150a --- /dev/null +++ b/litellm-rust/crates/secrets-azure/tests/fixtures/key_vault_parity.json @@ -0,0 +1,8 @@ +{ + "cases": [ + {"name": "plain_value", "secret_name": "OPENAI-API-KEY", "response": {"status": 200, "body": {"value": "sk-parity-1", "id": "https://example.vault.azure.net/secrets/OPENAI-API-KEY/abc"}}, "expected": {"value": "sk-parity-1"}}, + {"name": "json_value_is_kept_as_string", "secret_name": "JSON-SECRET", "response": {"status": 200, "body": {"value": "{\"api_key\": \"nested\"}", "id": "https://example.vault.azure.net/secrets/JSON-SECRET/abc"}}, "expected": {"value": "{\"api_key\": \"nested\"}"}}, + {"name": "missing_secret", "secret_name": "MISSING", "response": {"status": 404, "body": {"error": {"code": "SecretNotFound", "message": "not found"}}}, "expected": {"missing": true}}, + {"name": "forbidden", "secret_name": "FORBIDDEN", "response": {"status": 403, "body": {"error": {"code": "Forbidden", "message": "denied"}}}, "expected": {"error": true}} + ] +} diff --git a/litellm-rust/crates/secrets-azure/tests/key_vault.rs b/litellm-rust/crates/secrets-azure/tests/key_vault.rs new file mode 100644 index 00000000000..aae09369243 --- /dev/null +++ b/litellm-rust/crates/secrets-azure/tests/key_vault.rs @@ -0,0 +1,220 @@ +use std::sync::Arc; + +use litellm_secrets_azure::{AzureKeyVault, Error}; +use litellm_secrets_types::{Secret, SecretValue}; +use serde::Deserialize; +use wiremock::{ + Mock, MockServer, ResponseTemplate, + matchers::{header, path, query_param}, +}; + +fn manager(server: &MockServer) -> AzureKeyVault { + AzureKeyVault::with_client( + reqwest::Client::new(), + server.uri().parse().unwrap(), + Arc::new(|name: &str| (name == "AZURE_AD_TOKEN").then(|| "fake".to_owned())), + ) + .unwrap() +} + +#[tokio::test] +async fn reads_secret_with_bearer_token_and_api_version() { + let server = MockServer::start().await; + Mock::given(path("/secrets/OPENAI-API-KEY")) + .and(query_param("api-version", "7.4")) + .and(header("authorization", "Bearer fake")) + .respond_with( + ResponseTemplate::new(200) + .set_body_json(serde_json::json!({"value": "s3cret", "id": "secret-id"})), + ) + .expect(1) + .mount(&server) + .await; + + let secret = manager(&server) + .get_secret_from_azure_key_vault("OPENAI-API-KEY") + .await + .unwrap() + .unwrap(); + + assert_eq!(secret, Secret::String(SecretValue::new("s3cret"))); +} + +#[tokio::test] +async fn percent_encodes_secret_name_path_segment() { + let server = MockServer::start().await; + Mock::given(path("/secrets/name%2Fwith%20spaces")) + .and(query_param("api-version", "7.4")) + .respond_with( + ResponseTemplate::new(200).set_body_json(serde_json::json!({"value": "value"})), + ) + .expect(1) + .mount(&server) + .await; + + let secret = manager(&server) + .get_secret_from_azure_key_vault("name/with spaces") + .await + .unwrap() + .unwrap(); + + assert_eq!(secret.as_str(), Some("value")); +} + +#[rstest::rstest] +#[case::not_found(404, None)] +#[case::forbidden(403, Some(403))] +#[tokio::test] +async fn handles_statuses(#[case] status: u16, #[case] expected_status: Option) { + let server = MockServer::start().await; + Mock::given(path("/secrets/NAME")) + .respond_with(ResponseTemplate::new(status)) + .expect(1) + .mount(&server) + .await; + + let result = manager(&server) + .get_secret_from_azure_key_vault("NAME") + .await; + + match expected_status { + None => assert_eq!(result.unwrap(), None), + Some(status) => assert!(matches!(result, Err(Error::Status(actual)) if actual == status)), + } +} + +#[tokio::test] +async fn missing_value_is_an_error() { + let server = MockServer::start().await; + Mock::given(path("/secrets/NAME")) + .respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({}))) + .expect(1) + .mount(&server) + .await; + + assert!(matches!( + manager(&server) + .get_secret_from_azure_key_vault("NAME") + .await, + Err(Error::MissingValue) + )); +} + +#[test] +fn new_validates_vault_environment() { + assert!(matches!( + AzureKeyVault::new(Arc::new(|_: &str| None)), + Err(Error::MissingEnvironment("AZURE_KEY_VAULT_URI")) + )); + assert!(matches!( + AzureKeyVault::new(Arc::new(|name: &str| { + (name == "AZURE_KEY_VAULT_URI").then(|| "http://vault.example".to_owned()) + })), + Err(Error::VaultUri) + )); + assert!(matches!( + AzureKeyVault::new(Arc::new(|name: &str| { + (name == "AZURE_KEY_VAULT_URI").then(|| "vault.example".to_owned()) + })), + Err(Error::VaultUri) + )); +} + +#[rstest::rstest] +#[case("https://myvault.vault.azure.net", "https://vault.azure.net/.default")] +#[case( + "https://v.vault.usgovcloudapi.net/", + "https://vault.usgovcloudapi.net/.default" +)] +#[case("http://localhost:8080", "https://localhost/.default")] +#[test] +fn derives_scope_from_vault_host(#[case] uri: &str, #[case] expected: &str) { + let manager = AzureKeyVault::with_client( + reqwest::Client::new(), + uri.parse().unwrap(), + Arc::new(|_: &str| None), + ) + .unwrap(); + + assert_eq!(manager.scope(), expected); +} + +#[tokio::test] +async fn missing_credentials_do_not_request_vault() { + let server = MockServer::start().await; + Mock::given(path("/secrets/NAME")) + .respond_with(ResponseTemplate::new(200)) + .expect(0) + .mount(&server) + .await; + + assert!( + manager_without_credentials(&server) + .get_secret_from_azure_key_vault("NAME") + .await + .is_err() + ); +} + +fn manager_without_credentials(server: &MockServer) -> AzureKeyVault { + AzureKeyVault::with_client( + reqwest::Client::new(), + server.uri().parse().unwrap(), + Arc::new(|_: &str| None), + ) + .unwrap() +} + +#[derive(Deserialize)] +struct Fixture { + cases: Vec, +} + +#[derive(Deserialize)] +struct FixtureCase { + secret_name: String, + response: FixtureResponse, + expected: FixtureExpected, +} + +#[derive(Deserialize)] +struct FixtureResponse { + status: u16, + body: serde_json::Value, +} + +#[derive(Deserialize)] +struct FixtureExpected { + value: Option, + missing: Option, + error: Option, +} + +#[tokio::test] +async fn parity_fixture_matches_python_backend_contract() { + let fixture: Fixture = + serde_json::from_str(include_str!("fixtures/key_vault_parity.json")).unwrap(); + for case in fixture.cases { + let server = MockServer::start().await; + Mock::given(path(format!("/secrets/{}", case.secret_name))) + .respond_with( + ResponseTemplate::new(case.response.status).set_body_json(case.response.body), + ) + .expect(1) + .mount(&server) + .await; + let result = manager(&server) + .get_secret_from_azure_key_vault(&case.secret_name) + .await; + if case.expected.missing == Some(true) { + assert_eq!(result.unwrap(), None); + } else if case.expected.error == Some(true) { + assert!(result.is_err()); + } else { + assert_eq!( + result.unwrap().unwrap().as_str(), + case.expected.value.as_deref() + ); + } + } +} diff --git a/litellm-rust/crates/secrets-azure/tests/live.rs b/litellm-rust/crates/secrets-azure/tests/live.rs new file mode 100644 index 00000000000..a062ba95070 --- /dev/null +++ b/litellm-rust/crates/secrets-azure/tests/live.rs @@ -0,0 +1,30 @@ +use std::sync::Arc; + +use litellm_core_utils::settings::ProcessEnvironment; +use litellm_secrets_azure::AzureKeyVault; +use litellm_secrets_types::Secret; + +#[tokio::test] +#[ignore] +async fn reads_a_real_secret() { + let environment = Arc::new(ProcessEnvironment); + let manager = AzureKeyVault::new(environment).unwrap(); + let name = std::env::var("AZURE_KEY_VAULT_LIVE_SECRET_NAME").unwrap(); + let secret = manager + .get_secret_from_azure_key_vault(&name) + .await + .unwrap() + .unwrap(); + assert!(matches!(&secret, Secret::String(_))); + let host = std::env::var("AZURE_KEY_VAULT_URI") + .unwrap() + .parse::() + .unwrap() + .host_str() + .unwrap() + .to_owned(); + let value_len = secret.as_str().unwrap().len(); + println!( + "native provider=litellm-secrets-azure vault_host={host} secret={name} value_len={value_len}" + ); +} diff --git a/litellm-rust/crates/secrets/Cargo.toml b/litellm-rust/crates/secrets/Cargo.toml index a7e7ec80636..f56ad100337 100644 --- a/litellm-rust/crates/secrets/Cargo.toml +++ b/litellm-rust/crates/secrets/Cargo.toml @@ -9,11 +9,13 @@ repository.workspace = true default = [] aws = ["dep:litellm-secrets-aws"] google = ["dep:litellm-secrets-google"] +azure = ["dep:litellm-secrets-azure"] [dependencies] litellm-secrets-types.workspace = true litellm-secrets-aws = { workspace = true, optional = true } litellm-secrets-google = { workspace = true, optional = true } +litellm-secrets-azure = { workspace = true, optional = true } litellm-core-utils.workspace = true base64.workspace = true serde.workspace = true diff --git a/litellm-rust/crates/secrets/src/error.rs b/litellm-rust/crates/secrets/src/error.rs index 0c6e681b8aa..3729b5c2f2c 100644 --- a/litellm-rust/crates/secrets/src/error.rs +++ b/litellm-rust/crates/secrets/src/error.rs @@ -30,4 +30,7 @@ pub enum Error { #[cfg(feature = "google")] #[error(transparent)] Google(#[from] litellm_secrets_google::Error), + #[cfg(feature = "azure")] + #[error(transparent)] + Azure(#[from] litellm_secrets_azure::Error), } diff --git a/litellm-rust/crates/secrets/src/handler.rs b/litellm-rust/crates/secrets/src/handler.rs index 943ffdf6158..84360b39f26 100644 --- a/litellm-rust/crates/secrets/src/handler.rs +++ b/litellm-rust/crates/secrets/src/handler.rs @@ -13,6 +13,8 @@ pub enum SecretManager { GoogleKms(crate::google::GoogleKms), #[cfg(feature = "google")] GoogleSecretManager(crate::google::GoogleSecretManager), + #[cfg(feature = "azure")] + AzureKeyVault(crate::azure::AzureKeyVault), } impl SecretManager { @@ -27,6 +29,8 @@ impl SecretManager { Self::GoogleKms(_) => KeyManagementSystem::GoogleKms, #[cfg(feature = "google")] Self::GoogleSecretManager(_) => KeyManagementSystem::GoogleSecretManager, + #[cfg(feature = "azure")] + Self::AzureKeyVault(_) => KeyManagementSystem::AzureKeyVault, } } } @@ -78,6 +82,11 @@ pub async fn get_secret_from_manager( .get_secret_from_google_secret_manager(secret_name) .await .map_err(Error::from), + #[cfg(feature = "azure")] + SecretManager::AzureKeyVault(client) => client + .get_secret_from_azure_key_vault(secret_name) + .await + .map_err(Error::from), } } diff --git a/litellm-rust/crates/secrets/src/lib.rs b/litellm-rust/crates/secrets/src/lib.rs index ff2e95f7b2f..52251bb593a 100644 --- a/litellm-rust/crates/secrets/src/lib.rs +++ b/litellm-rust/crates/secrets/src/lib.rs @@ -17,5 +17,7 @@ pub use state::{SecretManagerState, secret_manager_would_be_consulted}; #[cfg(feature = "aws")] pub use litellm_secrets_aws as aws; +#[cfg(feature = "azure")] +pub use litellm_secrets_azure as azure; #[cfg(feature = "google")] pub use litellm_secrets_google as google; diff --git a/litellm-rust/crates/secrets/tests/handler.rs b/litellm-rust/crates/secrets/tests/handler.rs index a2cbbd843e1..b5c7b9e0cfb 100644 --- a/litellm-rust/crates/secrets/tests/handler.rs +++ b/litellm-rust/crates/secrets/tests/handler.rs @@ -105,3 +105,64 @@ async fn google_handler_requires_canonical_base64_and_preserves_plaintext_whites Err(Error::MissingCiphertext) )); } + +#[cfg(feature = "azure")] +#[tokio::test] +async fn azure_handler_reads_missing_and_failed_secrets() { + use litellm_secrets::{ + Error, KeyManagementSettings, KeyManagementSystem, SecretManager, azure::AzureKeyVault, + get_secret_from_manager, + }; + use wiremock::{ + Mock, MockServer, ResponseTemplate, + matchers::{path, query_param}, + }; + + let server = MockServer::start().await; + Mock::given(path("/secrets/KEY")) + .and(query_param("api-version", "7.4")) + .respond_with( + ResponseTemplate::new(200).set_body_json(serde_json::json!({"value": "value"})), + ) + .expect(1) + .mount(&server) + .await; + let manager = SecretManager::AzureKeyVault( + AzureKeyVault::with_client( + reqwest::Client::new(), + server.uri().parse().unwrap(), + std::sync::Arc::new(|name: &str| (name == "AZURE_AD_TOKEN").then(|| "fake".to_owned())), + ) + .unwrap(), + ); + assert_eq!(manager.system(), KeyManagementSystem::AzureKeyVault); + let settings = KeyManagementSettings::default(); + let value = get_secret_from_manager(&manager, "KEY", &settings, &|_: &str| None) + .await + .unwrap() + .unwrap(); + assert_eq!(value.as_str(), Some("value")); + + let not_found = Mock::given(path("/secrets/MISSING")) + .respond_with(ResponseTemplate::new(404)) + .expect(1) + .mount_as_scoped(&server) + .await; + assert_eq!( + get_secret_from_manager(&manager, "MISSING", &settings, &|_: &str| None) + .await + .unwrap(), + None + ); + drop(not_found); + + Mock::given(path("/secrets/FAILED")) + .respond_with(ResponseTemplate::new(500)) + .expect(1) + .mount(&server) + .await; + assert!(matches!( + get_secret_from_manager(&manager, "FAILED", &settings, &|_: &str| None).await, + Err(Error::Azure(_)) + )); +} diff --git a/tests/test_litellm/secret_managers/test_secret_manager_handler.py b/tests/test_litellm/secret_managers/test_secret_manager_handler.py new file mode 100644 index 00000000000..0925198992b --- /dev/null +++ b/tests/test_litellm/secret_managers/test_secret_manager_handler.py @@ -0,0 +1,105 @@ +from dataclasses import dataclass +from pathlib import Path +from typing import Final + +import pytest +from pydantic import BaseModel, ConfigDict + +from litellm.secret_managers.secret_manager_handler import get_secret_from_manager +from litellm.types.secret_managers.main import KeyManagementSystem + + +def _azure_exception_types() -> tuple[type[Exception], type[Exception]]: + try: + from azure.core.exceptions import ( + HttpResponseError, + ResourceNotFoundError, + ) + except ImportError: + return Exception, Exception + return HttpResponseError, ResourceNotFoundError + + +_AZURE_EXCEPTION_TYPES: Final[tuple[type[Exception], type[Exception]]] = _azure_exception_types() +AzureHttpResponseError: Final[type[Exception]] = _AZURE_EXCEPTION_TYPES[0] +AzureResourceNotFoundError: Final[type[Exception]] = _AZURE_EXCEPTION_TYPES[1] + + +class FixtureResponse(BaseModel): + model_config = ConfigDict(frozen=True) + + status: int + body: dict[str, object] + + +class FixtureExpected(BaseModel): + model_config = ConfigDict(frozen=True) + + value: str | None = None + missing: bool = False + error: bool = False + + +class FixtureCase(BaseModel): + model_config = ConfigDict(frozen=True) + + name: str + secret_name: str + response: FixtureResponse + expected: FixtureExpected + + +class Fixture(BaseModel): + model_config = ConfigDict(frozen=True) + + cases: tuple[FixtureCase, ...] + + +@dataclass(frozen=True, slots=True) +class FakeSecret: + value: str | None + + +@dataclass(frozen=True, slots=True) +class FakeAzureKeyVaultClient: + status: int + value: str | None + + def get_secret(self, name: str) -> FakeSecret: + if self.status == 404: + raise AzureResourceNotFoundError() + if self.status != 200: + raise AzureHttpResponseError() + return FakeSecret(value=self.value) + + +FIXTURE_PATH: Path = ( + Path(__file__).parents[3] + / "litellm-rust/crates/secrets-azure/tests/fixtures/key_vault_parity.json" +) + + +def test_azure_key_vault_matches_rust_parity_fixture() -> None: + fixture: Fixture = Fixture.model_validate_json(FIXTURE_PATH.read_text()) + for case in fixture.cases: + value: object = case.response.body.get("value") + secret: str | None = value if isinstance(value, str) else None + client: FakeAzureKeyVaultClient = FakeAzureKeyVaultClient( + status=case.response.status, + value=secret, + ) + if case.expected.missing or case.expected.error: + with pytest.raises(Exception): + get_secret_from_manager( + secret_name=case.secret_name, + key_manager=KeyManagementSystem.AZURE_KEY_VAULT.value, + client=client, + ) + continue + + result: str | None = get_secret_from_manager( + secret_name=case.secret_name, + key_manager=KeyManagementSystem.AZURE_KEY_VAULT.value, + client=client, + ) + assert result == case.expected.value From bfb4a8a2b33ed60fd083d59806f33aea4e5b5f59 Mon Sep 17 00:00:00 2001 From: Yujong Lee Date: Mon, 21 Sep 2026 20:36:10 +0000 Subject: [PATCH 27/44] refactor(rust): build Azure Key Vault auth inputs directly and keep credential tests offline Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm-rust/crates/auth-azure/src/lib.rs | 2 +- litellm-rust/crates/secrets-azure/Cargo.toml | 2 +- .../crates/secrets-azure/src/key_vault.rs | 42 +++++++++---------- .../crates/secrets-azure/tests/key_vault.rs | 4 +- 4 files changed, 24 insertions(+), 26 deletions(-) diff --git a/litellm-rust/crates/auth-azure/src/lib.rs b/litellm-rust/crates/auth-azure/src/lib.rs index e76227d6aa2..5c7c654b69d 100644 --- a/litellm-rust/crates/auth-azure/src/lib.rs +++ b/litellm-rust/crates/auth-azure/src/lib.rs @@ -4,4 +4,4 @@ mod resolve; mod types; pub use resolve::AzureAuthService; -pub use types::AzureAuthInputs; +pub use types::{AzureAuthInputs, ConfigValue}; diff --git a/litellm-rust/crates/secrets-azure/Cargo.toml b/litellm-rust/crates/secrets-azure/Cargo.toml index d559592cb34..96db7f235ef 100644 --- a/litellm-rust/crates/secrets-azure/Cargo.toml +++ b/litellm-rust/crates/secrets-azure/Cargo.toml @@ -12,7 +12,6 @@ litellm-secrets-types.workspace = true litellm-core-utils.workspace = true reqwest.workspace = true serde.workspace = true -serde_json.workspace = true thiserror.workspace = true veil.workspace = true percent-encoding = "2.3" @@ -21,4 +20,5 @@ percent-encoding = "2.3" tokio.workspace = true wiremock = "0.6.5" rstest.workspace = true +serde_json.workspace = true sha2.workspace = true diff --git a/litellm-rust/crates/secrets-azure/src/key_vault.rs b/litellm-rust/crates/secrets-azure/src/key_vault.rs index 262b19386df..e12289b83f5 100644 --- a/litellm-rust/crates/secrets-azure/src/key_vault.rs +++ b/litellm-rust/crates/secrets-azure/src/key_vault.rs @@ -1,21 +1,28 @@ use std::sync::Arc; -use litellm_auth_azure::{AzureAuthInputs, AzureAuthService}; +use litellm_auth_azure::{AzureAuthInputs, AzureAuthService, ConfigValue}; +use litellm_auth_types::{InputSource, Sourced}; use litellm_core_utils::settings::Lookup; use litellm_secrets_types::{Secret, SecretValue}; +use percent_encoding::{AsciiSet, NON_ALPHANUMERIC}; use serde::Deserialize; use crate::Error; const AZURE_KEY_VAULT_URI: &str = "AZURE_KEY_VAULT_URI"; const API_VERSION: &str = "7.4"; +const PATH_SEGMENT: &AsciiSet = &NON_ALPHANUMERIC + .remove(b'-') + .remove(b'.') + .remove(b'_') + .remove(b'~'); #[derive(Clone)] pub struct AzureKeyVault { client: reqwest::Client, vault: reqwest::Url, auth: Arc, - inputs: AzureAuthInputs, + inputs: Arc, environment: Arc, } @@ -33,21 +40,19 @@ impl AzureKeyVault { if vault.host_str().is_none() { return Err(Error::VaultUri); } - let scope = scope_for(&vault); - let inputs = AzureAuthInputs::from_sourced_optional_params( - serde_json::json!({ - "azure_scope": scope, - "enable_azure_ad_token_refresh": true, - }) - .as_object() - .expect("static Azure auth inputs object"), - &std::collections::BTreeMap::new(), - )?; + let inputs = AzureAuthInputs { + azure_scope: ConfigValue::Value(Sourced::new( + scope_for(&vault), + InputSource::Deployment, + )), + enable_azure_ad_token_refresh: Sourced::new(true, InputSource::Deployment), + ..AzureAuthInputs::default() + }; Ok(Self { client, vault, auth: Arc::new(AzureAuthService::default()), - inputs, + inputs: Arc::new(inputs), environment, }) } @@ -80,7 +85,7 @@ impl AzureKeyVault { .get_azure_ad_token(&self.inputs, &|key| self.environment.get(key)) .await? .ok_or(Error::MissingCredentials)?; - let encoded_name = encode_name(name); + let encoded_name = percent_encoding::utf8_percent_encode(name, PATH_SEGMENT); let url = self .vault .join(&format!("secrets/{encoded_name}?api-version={API_VERSION}")) @@ -111,12 +116,3 @@ fn scope_for(vault: &reqwest::Url) -> String { .map_or(host, |(_, remainder)| remainder); format!("https://{resource}/.default") } - -fn encode_name(name: &str) -> String { - percent_encoding::utf8_percent_encode(name, percent_encoding::NON_ALPHANUMERIC) - .to_string() - .replace("%2D", "-") - .replace("%2E", ".") - .replace("%5F", "_") - .replace("%7E", "~") -} diff --git a/litellm-rust/crates/secrets-azure/tests/key_vault.rs b/litellm-rust/crates/secrets-azure/tests/key_vault.rs index aae09369243..cf9102d0b45 100644 --- a/litellm-rust/crates/secrets-azure/tests/key_vault.rs +++ b/litellm-rust/crates/secrets-azure/tests/key_vault.rs @@ -160,7 +160,9 @@ fn manager_without_credentials(server: &MockServer) -> AzureKeyVault { AzureKeyVault::with_client( reqwest::Client::new(), server.uri().parse().unwrap(), - Arc::new(|_: &str| None), + Arc::new(|name: &str| { + (name == "AZURE_CREDENTIAL").then(|| "ClientSecretCredential".to_owned()) + }), ) .unwrap() } From df8966591939e6aa4f747429b762ec33a6721543 Mon Sep 17 00:00:00 2001 From: Yujong Lee Date: Mon, 21 Sep 2026 20:37:57 +0000 Subject: [PATCH 28/44] test: narrow Azure parity exception assertion Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../secret_managers/test_secret_manager_handler.py | 4 +++- 1 file changed, 3 insertions(+), 1 deletion(-) diff --git a/tests/test_litellm/secret_managers/test_secret_manager_handler.py b/tests/test_litellm/secret_managers/test_secret_manager_handler.py index 0925198992b..b4838912d27 100644 --- a/tests/test_litellm/secret_managers/test_secret_manager_handler.py +++ b/tests/test_litellm/secret_managers/test_secret_manager_handler.py @@ -89,7 +89,9 @@ def test_azure_key_vault_matches_rust_parity_fixture() -> None: value=secret, ) if case.expected.missing or case.expected.error: - with pytest.raises(Exception): + with pytest.raises( + AzureResourceNotFoundError if case.expected.missing else AzureHttpResponseError + ): get_secret_from_manager( secret_name=case.secret_name, key_manager=KeyManagementSystem.AZURE_KEY_VAULT.value, From b8c793b9625467d5647e9a49e251fe5480953350 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Mon, 21 Sep 2026 13:44:58 -0700 Subject: [PATCH 29/44] fix(types): drop restating docstrings, import search tool types at runtime, close the spend-table match --- .../proxy/common_utils/check_batch_cost.py | 6 ------ .../proxy/common_utils/check_responses_cost.py | 3 --- litellm/proxy/db/db_spend_update_writer.py | 5 +++-- litellm/router_utils/search_api_router.py | 12 ++++-------- 4 files changed, 7 insertions(+), 19 deletions(-) diff --git a/enterprise/litellm_enterprise/proxy/common_utils/check_batch_cost.py b/enterprise/litellm_enterprise/proxy/common_utils/check_batch_cost.py index 3ea9b7d9bfd..41974c26158 100644 --- a/enterprise/litellm_enterprise/proxy/common_utils/check_batch_cost.py +++ b/enterprise/litellm_enterprise/proxy/common_utils/check_batch_cost.py @@ -43,8 +43,6 @@ TERMINAL_MANAGED_OBJECT_STATUSES: Final[Tuple[str, ...]] = ( class _ManagedObjectRow(Protocol): - """The managed-object row fields this poller reads off whatever the DB hands back.""" - @property def id(self) -> str: ... @@ -59,19 +57,16 @@ class _ManagedObjectRow(Protocol): def _managed_object_table(prisma_client: "PrismaClient") -> "TableActions[_ManagedObjectRow]": - """The managed-object table's prisma actions, typed to the row fields this module reads.""" table: Final[TableActions[_ManagedObjectRow]] = prisma_client.db.litellm_managedobjecttable return table def _user_table(prisma_client: "PrismaClient") -> "TableActions[prisma_models.LiteLLM_UserTable]": - """The user table's prisma actions.""" table: Final[TableActions[prisma_models.LiteLLM_UserTable]] = prisma_client.db.litellm_usertable return table def _token_table(prisma_client: "PrismaClient") -> "TableActions[prisma_models.LiteLLM_VerificationToken]": - """The virtual-key table's prisma actions.""" table: Final[TableActions[prisma_models.LiteLLM_VerificationToken]] = ( prisma_client.db.litellm_verificationtoken ) @@ -79,7 +74,6 @@ def _token_table(prisma_client: "PrismaClient") -> "TableActions[prisma_models.L def _team_table(prisma_client: "PrismaClient") -> "TableActions[prisma_models.LiteLLM_TeamTable]": - """The team table's prisma actions.""" table: Final[TableActions[prisma_models.LiteLLM_TeamTable]] = prisma_client.db.litellm_teamtable return table diff --git a/enterprise/litellm_enterprise/proxy/common_utils/check_responses_cost.py b/enterprise/litellm_enterprise/proxy/common_utils/check_responses_cost.py index 1bc41f2aa5b..cdeea0d3d4b 100644 --- a/enterprise/litellm_enterprise/proxy/common_utils/check_responses_cost.py +++ b/enterprise/litellm_enterprise/proxy/common_utils/check_responses_cost.py @@ -29,8 +29,6 @@ TERMINAL_RESPONSE_STATUSES = frozenset({"completed", "failed", "cancelled", "inc class _ManagedObjectRow(Protocol): - """The managed-object row fields this poller reads off whatever the DB hands back.""" - @property def id(self) -> str: ... @@ -45,7 +43,6 @@ class _ManagedObjectRow(Protocol): def _managed_object_table(prisma_client: "PrismaClient") -> "TableActions[_ManagedObjectRow]": - """The managed-object table's prisma actions, typed to the row fields this poller reads.""" table: Final[TableActions[_ManagedObjectRow]] = prisma_client.db.litellm_managedobjecttable return table diff --git a/litellm/proxy/db/db_spend_update_writer.py b/litellm/proxy/db/db_spend_update_writer.py index 0516f5460a7..865a8c39f03 100644 --- a/litellm/proxy/db/db_spend_update_writer.py +++ b/litellm/proxy/db/db_spend_update_writer.py @@ -19,7 +19,7 @@ from typing import TYPE_CHECKING, Any, Final, Literal, Protocol, TypeAlias, Type from urllib.parse import quote, unquote from pydantic import TypeAdapter -from typing_extensions import LiteralString, ReadOnly, TypedDict +from typing_extensions import LiteralString, ReadOnly, TypedDict, assert_never import litellm from litellm._logging import verbose_proxy_logger @@ -142,7 +142,6 @@ _EntitySpendTable: TypeAlias = Literal[ def _entity_spend_table(batcher: _SpendBatch, table_accessor: _EntitySpendTable) -> BatchTable: - """The batch table an entity type's spend increments are written to.""" match table_accessor: case "litellm_tagtable": return batcher.litellm_tagtable @@ -152,6 +151,8 @@ def _entity_spend_table(batcher: _SpendBatch, table_accessor: _EntitySpendTable) return batcher.litellm_modelaccessgroupbudgettable case "litellm_projecttable": return batcher.litellm_projecttable + case _ as unreachable: + assert_never(unreachable) class _SpendBatchManager(Protocol): diff --git a/litellm/router_utils/search_api_router.py b/litellm/router_utils/search_api_router.py index 76e833563ba..1cfb311d796 100644 --- a/litellm/router_utils/search_api_router.py +++ b/litellm/router_utils/search_api_router.py @@ -10,18 +10,16 @@ import traceback from collections.abc import Callable from functools import partial from types import MappingProxyType -from typing import TYPE_CHECKING, Any, Final, Protocol +from typing import Any, Final, Protocol from litellm._logging import verbose_router_logger - -if TYPE_CHECKING: - from litellm.types.router import SearchToolLiteLLMParams, SearchToolTypedDict +from litellm.types.router import SearchToolLiteLLMParams, SearchToolTypedDict class _SearchToolsRouter(Protocol): """The one router attribute the search-tool helpers read and replace.""" - search_tools: "list[SearchToolTypedDict]" + search_tools: list[SearchToolTypedDict] class SearchAPIRouter: @@ -34,7 +32,7 @@ class SearchAPIRouter: @staticmethod def _resolve_search_provider_credentials( *, - tool_litellm_params: "SearchToolLiteLLMParams", + tool_litellm_params: SearchToolLiteLLMParams, ) -> tuple[str | None, str | None]: """ Resolve search provider credentials from tool configuration ONLY. @@ -65,8 +63,6 @@ class SearchAPIRouter: search_tools: List of search tool configurations from the database """ try: - from litellm.types.router import SearchToolTypedDict - verbose_router_logger.debug("Adding %s search tools to router", len(search_tools)) # Convert search tools to the format expected by the router From 5f84e8331689c95c333d0309e89ad566cd613122 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Mon, 21 Sep 2026 13:48:11 -0700 Subject: [PATCH 30/44] chore(responses): drop narrative comments from the MCP streaming iterator --- litellm/responses/mcp/mcp_streaming_iterator.py | 15 --------------- 1 file changed, 15 deletions(-) diff --git a/litellm/responses/mcp/mcp_streaming_iterator.py b/litellm/responses/mcp/mcp_streaming_iterator.py index 1c892eec803..fb05e7bff99 100644 --- a/litellm/responses/mcp/mcp_streaming_iterator.py +++ b/litellm/responses/mcp/mcp_streaming_iterator.py @@ -341,11 +341,6 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator): self._error_event_emitted = False self._last_sequence_number = 0 - # Every auto-execute round is a distinct upstream response, but the - # client is reading one stream. Fold the rounds into one public - # lifecycle: one response.created, one response.completed whose - # output holds every round's items, and output indexes that are - # never reused for a different item. self._round_index = 0 self._output_index_offset = 0 self._round_max_output_index = -1 @@ -438,8 +433,6 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator): chunk: Final = await self._anext_impl() sequence_number: Final = getattr(chunk, "sequence_number", None) if isinstance(sequence_number, int): - # Follow-up rounds and gateway events restart their numbering. - # Keep the public stream strictly increasing. if sequence_number <= self._last_sequence_number and self._last_sequence_number > 0: self._last_sequence_number += 1 _set_event_field(chunk, "sequence_number", self._last_sequence_number) @@ -560,8 +553,6 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator): self.phase = "mcp_discovery" return await self._compose_round_chunk(chunk) - # None means the chunk was folded into the single public - # lifecycle; fall through so phase 4 runs the follow-up. return await self._compose_round_chunk(chunk) except StopAsyncIteration: if self.should_auto_execute and self.collected_response: @@ -683,7 +674,6 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator): composed: Final = await self._compose_round_chunk(chunk) if composed is None: - # The chunk stays internal; hand the next public event back instead. return await self._anext_impl() return composed @@ -749,10 +739,6 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator): return self.tool_call_round += 1 - # Each executed tool is one mcp_call output item of the single - # public response. Announce it at an output_index past the items - # this round already streamed, and keep that item id for the - # completion events below. from litellm.types.llms.openai import OutputItemAddedEvent next_output_index = self._output_index_offset + self._round_output_width( # rebind-ok: advances per item @@ -870,7 +856,6 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator): item=mcp_call_item, ) self.tool_execution_events.append(output_item_done_event) - # The response model accepts output items as dicts, not as the generic event object. self._pending_mcp_call_items.append(mcp_call_item.model_dump()) # Store tool results for follow-up call From 239bff3315ee9836559b7947157abe71c3cdcd0c Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Mon, 21 Sep 2026 13:57:50 -0700 Subject: [PATCH 31/44] test(responses): type the MCP lifecycle test helpers --- .../responses/mcp/test_mcp_streaming_iterator.py | 14 +++++++++----- 1 file changed, 9 insertions(+), 5 deletions(-) diff --git a/tests/test_litellm/responses/mcp/test_mcp_streaming_iterator.py b/tests/test_litellm/responses/mcp/test_mcp_streaming_iterator.py index 57ccbbfa493..1587557f3d2 100644 --- a/tests/test_litellm/responses/mcp/test_mcp_streaming_iterator.py +++ b/tests/test_litellm/responses/mcp/test_mcp_streaming_iterator.py @@ -11,7 +11,11 @@ from litellm.responses.mcp.mcp_streaming_iterator import ( MAX_MCP_TOOL_CALL_ROUNDS, MCPEnhancedStreamingIterator, ) -from litellm.types.llms.openai import ResponsesAPIResponse, ResponsesAPIStreamEvents +from litellm.types.llms.openai import ( + BaseLiteLLMOpenAIResponseObject, + ResponsesAPIResponse, + ResponsesAPIStreamEvents, +) # `litellm.__init__` re-exports a function named `responses`, which shadows the # `litellm.responses` subpackage as an attribute — `import litellm.responses.main` @@ -57,8 +61,8 @@ def _text_message(text: str): return {"type": "message", "role": "assistant", "content": [{"type": "output_text", "text": text}]} -def _item_type(item) -> str: - return item["type"] if isinstance(item, dict) else item.type +def _item_type(item: dict[str, object] | BaseLiteLLMOpenAIResponseObject) -> str: + return str(item["type"]) if isinstance(item, dict) else str(item.type) def _tool_call_stream(call_id: str, tool_name: str, response_id: str = "resp-1") -> _FakeAsyncStream: @@ -354,11 +358,11 @@ async def test_streaming_follow_up_keeps_previous_response_id_when_stored(monkey assert not [item for item in follow_up_kwargs["input"] if item.get("type") == "reasoning"] -def _event(event_type, **fields): +def _event(event_type: ResponsesAPIStreamEvents, **fields: object) -> SimpleNamespace: return SimpleNamespace(type=event_type, **fields) -def _lifecycle_round(response_id: str, item: dict, sequence_start: int = 0): +def _lifecycle_round(response_id: str, item: dict[str, object], sequence_start: int = 0) -> list[SimpleNamespace]: """One upstream Responses round as a provider streams it: its own id, indexes from 0, numbering from 0.""" return [ _event( From 5acfcfa4febefce2f9df7d6876915278bf91d229 Mon Sep 17 00:00:00 2001 From: Yujong Lee Date: Mon, 21 Sep 2026 21:00:48 +0000 Subject: [PATCH 32/44] feat(rust): native Azure Blob response cache backend Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm-rust/Cargo.lock | 82 +++ litellm-rust/Cargo.toml | 1 + litellm-rust/crates/auth-azure/src/types.rs | 15 + .../crates/cache-azure-blob/Cargo.toml | 22 + .../crates/cache-azure-blob/src/cache.rs | 246 +++++++ .../cache-azure-blob/src/cache/tests.rs | 692 ++++++++++++++++++ .../crates/cache-azure-blob/src/credential.rs | 84 +++ .../crates/cache-azure-blob/src/lib.rs | 5 + .../crates/cache-azure-blob/src/tests.rs | 0 litellm-rust/crates/python-bridge/Cargo.toml | 1 + .../crates/python-bridge/src/cache/config.rs | 49 +- .../crates/python-bridge/src/cache/facade.rs | 88 ++- .../crates/python-bridge/src/cache/handle.rs | 17 +- .../crates/python-bridge/src/cache/native.rs | 44 +- tests/test_litellm_rust/test_cache.py | 117 +++ 15 files changed, 1439 insertions(+), 24 deletions(-) create mode 100644 litellm-rust/crates/cache-azure-blob/Cargo.toml create mode 100644 litellm-rust/crates/cache-azure-blob/src/cache.rs create mode 100644 litellm-rust/crates/cache-azure-blob/src/cache/tests.rs create mode 100644 litellm-rust/crates/cache-azure-blob/src/credential.rs create mode 100644 litellm-rust/crates/cache-azure-blob/src/lib.rs create mode 100644 litellm-rust/crates/cache-azure-blob/src/tests.rs diff --git a/litellm-rust/Cargo.lock b/litellm-rust/Cargo.lock index ed4ae4e3353..5b112519add 100644 --- a/litellm-rust/Cargo.lock +++ b/litellm-rust/Cargo.lock @@ -115,6 +115,28 @@ dependencies = [ "pin-project-lite", ] +[[package]] +name = "async-stream" +version = "0.3.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0b5a71a6f37880a80d1d7f19efd781e4b5de42c88f0722cc13bcb6cc2cfe8476" +dependencies = [ + "async-stream-impl", + "futures-core", + "pin-project-lite", +] + +[[package]] +name = "async-stream-impl" +version = "0.3.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c7c24de15d275a1ecfd47a380fb4d5ec9bfe0933f309ed5e705b775596a3574d" +dependencies = [ + "proc-macro2", + "quote", + "syn 2.0.119", +] + [[package]] name = "async-trait" version = "0.1.91" @@ -599,6 +621,37 @@ dependencies = [ "url", ] +[[package]] +name = "azure_storage_blob" +version = "1.1.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "17b10207ecf7d666df6940b50051f433b3cd5d2b9b1dd190613208d7a84e7eed" +dependencies = [ + "async-stream", + "async-trait", + "azure_core", + "azure_storage_common", + "bytes", + "futures", + "percent-encoding", + "pin-project", + "serde", + "serde_json", + "time", + "tokio", +] + +[[package]] +name = "azure_storage_common" +version = "1.0.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0af2e6aeb8d76b17fc998f453c320913f73787b944e3cc29509d19411fa0321d" +dependencies = [ + "azure_core", + "serde", + "time", +] + [[package]] name = "base64" version = "0.13.1" @@ -2464,6 +2517,23 @@ dependencies = [ "tokio", ] +[[package]] +name = "litellm-cache-azure-blob" +version = "0.1.0" +dependencies = [ + "async-trait", + "azure_core", + "azure_storage_blob", + "futures-util", + "litellm-auth-azure", + "litellm-auth-types", + "litellm-cache", + "litellm-cache-response", + "serde_json", + "tokio", + "url", +] + [[package]] name = "litellm-cache-memory" version = "0.1.0" @@ -2666,6 +2736,7 @@ dependencies = [ "litellm-auth", "litellm-auth-gcp", "litellm-cache", + "litellm-cache-azure-blob", "litellm-cache-memory", "litellm-cache-redis", "litellm-cache-response", @@ -3487,6 +3558,16 @@ version = "1.2.3" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "a1d01941d82fa2ab50be1e79e6714289dd7cde78eba4c074bc5a4374f650dfe0" +[[package]] +name = "quick-xml" +version = "0.41.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e660451e55124f798a69a5af3f49ccfbefbd41910eefd25caf2393e1f3473ec1" +dependencies = [ + "memchr", + "serde", +] + [[package]] name = "quinn" version = "0.11.11" @@ -5030,6 +5111,7 @@ dependencies = [ "base64 0.22.1", "bytes", "futures", + "quick-xml", "serde", "serde_json", "url", diff --git a/litellm-rust/Cargo.toml b/litellm-rust/Cargo.toml index 570d0dd3568..58fc1ba1d05 100644 --- a/litellm-rust/Cargo.toml +++ b/litellm-rust/Cargo.toml @@ -27,6 +27,7 @@ litellm-llms = { path = "crates/llms" } litellm-types = { path = "crates/types" } litellm-core-utils = { path = "crates/core-utils" } litellm-cache = { path = "crates/cache" } +litellm-cache-azure-blob = { path = "crates/cache-azure-blob" } litellm-cache-memory = { path = "crates/cache-memory" } litellm-cache-redis = { path = "crates/cache-redis" } litellm-cache-response = { path = "crates/cache-response" } diff --git a/litellm-rust/crates/auth-azure/src/types.rs b/litellm-rust/crates/auth-azure/src/types.rs index d5a00f09751..a3a898f000f 100644 --- a/litellm-rust/crates/auth-azure/src/types.rs +++ b/litellm-rust/crates/auth-azure/src/types.rs @@ -51,6 +51,21 @@ pub struct AzureAuthInputs { } impl AzureAuthInputs { + pub fn default_credential_for_scope(scope: &str) -> Self { + Self { + azure_scope: ConfigValue::Value(Sourced::new( + scope.to_string(), + InputSource::Deployment, + )), + azure_credential: ConfigValue::Value(Sourced::new( + "DefaultAzureCredential".to_string(), + InputSource::Deployment, + )), + enable_azure_ad_token_refresh: Sourced::new(true, InputSource::Deployment), + ..Self::default() + } + } + pub fn or_configured_token_refresh(self, enabled: bool) -> Self { if *self.enable_azure_ad_token_refresh.value() || !enabled { return self; diff --git a/litellm-rust/crates/cache-azure-blob/Cargo.toml b/litellm-rust/crates/cache-azure-blob/Cargo.toml new file mode 100644 index 00000000000..55abaff1975 --- /dev/null +++ b/litellm-rust/crates/cache-azure-blob/Cargo.toml @@ -0,0 +1,22 @@ +[package] +name = "litellm-cache-azure-blob" +version = "0.1.0" +edition.workspace = true +license.workspace = true +repository.workspace = true + +[dependencies] +litellm-auth-azure.workspace = true +litellm-auth-types.workspace = true +litellm-cache.workspace = true + +async-trait = "0.1" +azure_core = "1.1.0" +azure_storage_blob = "1.1.0" +futures-util.workspace = true +tokio.workspace = true +url.workspace = true + +[dev-dependencies] +litellm-cache-response.workspace = true +serde_json.workspace = true diff --git a/litellm-rust/crates/cache-azure-blob/src/cache.rs b/litellm-rust/crates/cache-azure-blob/src/cache.rs new file mode 100644 index 00000000000..f28b8c9a641 --- /dev/null +++ b/litellm-rust/crates/cache-azure-blob/src/cache.rs @@ -0,0 +1,246 @@ +use std::{sync::Arc, time::Duration}; + +use azure_core::{ + credentials::TokenCredential, + error::ErrorKind, + http::{ClientOptions, RequestContent}, +}; +use azure_storage_blob::{ + BlobContainerClient, BlobContainerClientOptions, + models::{BlobClientUploadOptions, StorageErrorCode}, +}; +use futures_util::{TryStreamExt, future::try_join_all}; +use litellm_cache::{ + BaseCache, BatchCache, CacheCodec, CacheConnectionResult, CacheConnectionStatus, Error, + ExactCacheContext, FlushCache, +}; +use tokio::runtime::Handle; +use url::Url; + +use crate::credential::AzureBlobCredential; + +/// Synchronous methods block on `runtime` and therefore must run outside of it +pub struct AzureBlobCache { + container: BlobContainerClient, + codec: C, + runtime: Handle, + account_url: String, + container_name: String, +} + +impl AzureBlobCache { + pub async fn connect( + account_url: &str, + container: &str, + codec: C, + runtime: Handle, + ) -> Result { + Self::connect_with_options( + account_url, + container, + Some(Arc::new(AzureBlobCredential::default())), + ClientOptions::default(), + codec, + runtime, + ) + .await + } + + pub async fn connect_with_options( + account_url: &str, + container: &str, + credential: Option>, + client_options: ClientOptions, + codec: C, + runtime: Handle, + ) -> Result { + let mut url = Url::parse(account_url).map_err(|_| Error::Unavailable)?; + let account_url = url.as_str().trim_end_matches('/').to_string(); + url.path_segments_mut() + .map_err(|()| Error::Unavailable)? + .pop_if_empty() + .push(container); + let client = BlobContainerClient::new( + url, + credential, + Some(BlobContainerClientOptions { + client_options, + ..BlobContainerClientOptions::default() + }), + ) + .map_err(|_| Error::Unavailable)?; + let cache = Self { + container: client, + codec, + runtime, + account_url, + container_name: container.to_string(), + }; + cache.create_container().await?; + Ok(cache) + } + + pub fn account_url(&self) -> &str { + &self.account_url + } + + pub fn container_name(&self) -> &str { + &self.container_name + } + + async fn create_container(&self) -> Result<(), Error> { + match self.container.create(None).await { + Ok(_) => Ok(()), + Err(error) if is_storage_error(&error, StorageErrorCode::ContainerAlreadyExists) => { + Ok(()) + } + Err(_) => Err(Error::Unavailable), + } + } + + async fn upload(&self, key: &str, value: &C::Value, overwrite: bool) -> Result<(), Error> { + let payload = self.codec.encode(value)?; + let options = (!overwrite).then(|| BlobClientUploadOptions::default().if_not_exists()); + match self + .container + .blob_client(key) + .upload(RequestContent::from(payload), options) + .await + { + Ok(_) => Ok(()), + Err(error) if is_storage_error(&error, StorageErrorCode::BlobAlreadyExists) => Ok(()), + Err(_) => Err(Error::Unavailable), + } + } + + async fn download(&self, key: &str) -> Result, Error> { + let response = match self.container.blob_client(key).download(None).await { + Ok(response) => response, + Err(error) if is_storage_error(&error, StorageErrorCode::BlobNotFound) => { + return Ok(None); + } + Err(_) => return Err(Error::Unavailable), + }; + let bytes = response + .body + .collect() + .await + .map_err(|_| Error::Unavailable)?; + self.codec.decode(&bytes).map(Some) + } + + async fn delete_all_blobs(&self) -> Result<(), Error> { + let mut pages = self + .container + .list_blobs(None) + .map_err(|_| Error::Unavailable)? + .into_pages(); + while let Some(page) = pages.try_next().await.map_err(|_| Error::Unavailable)? { + let page = page.into_model().map_err(|_| Error::Unavailable)?; + for name in page.blob_items.into_iter().filter_map(|item| item.name) { + self.container + .blob_client(&name) + .delete(None) + .await + .map_err(|_| Error::Unavailable)?; + } + } + Ok(()) + } + + fn block_on(&self, future: impl Future) -> T { + self.runtime.block_on(future) + } +} + +fn is_storage_error(error: &azure_core::Error, code: StorageErrorCode) -> bool { + matches!( + error.kind(), + ErrorKind::HttpResponse { + error_code: Some(error_code), + .. + } if error_code == code.as_ref() + ) +} + +impl BaseCache for AzureBlobCache { + type Value = C::Value; + type Context = ExactCacheContext; + + fn get_ttl(&self, _: &ExactCacheContext) -> Option { + None + } + + fn set_cache(&self, key: &str, value: C::Value, _: &ExactCacheContext) -> Result<(), Error> { + self.block_on(self.upload(key, &value, false)) + } + + fn get_cache(&self, key: &str, _: &ExactCacheContext) -> Result, Error> { + self.block_on(self.download(key)) + } + + async fn async_set_cache( + &self, + key: &str, + value: C::Value, + _: ExactCacheContext, + ) -> Result<(), Error> { + self.upload(key, &value, true).await + } + + async fn async_get_cache( + &self, + key: &str, + _: &ExactCacheContext, + ) -> Result, Error> { + self.download(key).await + } + + async fn async_set_cache_pipeline( + &self, + entries: Vec<(String, C::Value)>, + _: ExactCacheContext, + ) -> Result<(), Error> { + try_join_all( + entries + .iter() + .map(|(key, value)| self.upload(key, value, true)), + ) + .await + .map(drop) + } + + async fn disconnect(&self) -> Result<(), Error> { + Ok(()) + } + + async fn test_connection(&self) -> Result { + Ok(match self.container.get_properties(None).await { + Ok(_) => CacheConnectionResult { + status: CacheConnectionStatus::Success, + message: "Azure Blob cache connection test successful".into(), + error: None, + }, + Err(error) => CacheConnectionResult { + status: CacheConnectionStatus::Failed, + message: format!("Azure Blob connection failed: {error}"), + error: Some(error.to_string()), + }, + }) + } +} + +impl BatchCache for AzureBlobCache {} + +impl FlushCache for AzureBlobCache { + fn flush_cache(&self) -> Result<(), Error> { + self.block_on(self.delete_all_blobs()) + } + + async fn async_flush_cache(&self) -> Result<(), Error> { + self.delete_all_blobs().await + } +} + +#[cfg(test)] +mod tests; diff --git a/litellm-rust/crates/cache-azure-blob/src/cache/tests.rs b/litellm-rust/crates/cache-azure-blob/src/cache/tests.rs new file mode 100644 index 00000000000..fd116a0e28f --- /dev/null +++ b/litellm-rust/crates/cache-azure-blob/src/cache/tests.rs @@ -0,0 +1,692 @@ +use std::{ + collections::BTreeMap, + sync::{Arc, Mutex}, + time::Duration, +}; + +use azure_core::http::{ + AsyncRawResponse, Body, ClientOptions, HttpClient, Method, Request, StatusCode, Transport, + headers::{HeaderName, Headers}, +}; +use litellm_cache::{ + BaseCache, BatchCache, BatchEntry, CacheConnectionStatus, Error, ExactCacheContext, FlushCache, +}; +use litellm_cache_response::{ + CacheEntry, CacheKeyField, CacheKeyInput, ResponseCache, ResponseCacheCodec, + ResponseCacheRequest, cache_key, +}; +use serde_json::json; +use tokio::runtime::Runtime; + +use super::AzureBlobCache; + +const ACCOUNT_URL: &str = "https://example.blob.core.windows.net"; +const CONTAINER: &str = "litellm-cache"; +const IF_NONE_MATCH: HeaderName = HeaderName::from_static("if-none-match"); +const ERROR_CODE: HeaderName = HeaderName::from_static("x-ms-error-code"); + +#[derive(Clone, Debug, PartialEq, Eq)] +struct RecordedRequest { + method: Method, + path: String, + query: String, + if_none_match: Option, +} + +#[derive(Default)] +struct FakeState { + container_exists: bool, + blobs: BTreeMap>, + requests: Vec, + failing: bool, +} + +#[derive(Clone, Default)] +struct FakeBlobService { + state: Arc>, +} + +impl std::fmt::Debug for FakeBlobService { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.write_str("FakeBlobService") + } +} + +impl FakeBlobService { + fn with_existing_container() -> Self { + let service = Self::default(); + service.state.lock().unwrap().container_exists = true; + service + } + + fn blob(&self, name: &str) -> Option> { + self.state.lock().unwrap().blobs.get(name).cloned() + } + + fn blob_names(&self) -> Vec { + self.state.lock().unwrap().blobs.keys().cloned().collect() + } + + fn seed_blob(&self, name: &str, bytes: &[u8]) { + self.state + .lock() + .unwrap() + .blobs + .insert(name.to_string(), bytes.to_vec()); + } + + fn set_failing(&self, failing: bool) { + self.state.lock().unwrap().failing = failing; + } + + fn requests(&self) -> Vec { + self.state.lock().unwrap().requests.clone() + } + + fn container_exists(&self) -> bool { + self.state.lock().unwrap().container_exists + } + + fn respond(status: StatusCode, error_code: Option<&str>, body: Vec) -> AsyncRawResponse { + let mut headers = Headers::new(); + if let Some(code) = error_code { + headers.insert(ERROR_CODE, code.to_string()); + } + AsyncRawResponse::from_bytes(status, headers, body) + } + + fn list_body(state: &FakeState) -> Vec { + let mut xml = String::from( + r#""#, + ); + for name in state.blobs.keys() { + xml.push_str(&format!( + "{name}BlockBlob" + )); + } + xml.push_str(""); + xml.into_bytes() + } +} + +#[async_trait::async_trait] +impl HttpClient for FakeBlobService { + async fn execute_request(&self, request: &Request) -> azure_core::Result { + let mut state = self.state.lock().unwrap(); + let path = request.url().path().to_string(); + let query = request.url().query().unwrap_or_default().to_string(); + let if_none_match = request + .headers() + .get_optional_str(&IF_NONE_MATCH) + .map(str::to_owned); + state.requests.push(RecordedRequest { + method: request.method(), + path: path.clone(), + query: query.clone(), + if_none_match: if_none_match.clone(), + }); + if state.failing { + return Ok(Self::respond( + StatusCode::Forbidden, + Some("AuthorizationFailure"), + Vec::new(), + )); + } + let container_path = format!("/{CONTAINER}"); + let blob_name = path + .strip_prefix(&format!("{container_path}/")) + .map(str::to_owned); + let is_container = path == container_path && query.contains("restype=container"); + let response = match (request.method(), is_container, blob_name) { + (Method::Put, true, None) if state.container_exists => Self::respond( + StatusCode::Conflict, + Some("ContainerAlreadyExists"), + Vec::new(), + ), + (Method::Put, true, None) => { + state.container_exists = true; + Self::respond(StatusCode::Created, None, Vec::new()) + } + (Method::Get, true, None) if query.contains("comp=list") => { + Self::respond(StatusCode::Ok, None, Self::list_body(&state)) + } + (Method::Get, true, None) if state.container_exists => { + Self::respond(StatusCode::Ok, None, Vec::new()) + } + (Method::Get, true, None) => { + Self::respond(StatusCode::NotFound, Some("ContainerNotFound"), Vec::new()) + } + (Method::Put, false, Some(name)) => { + if if_none_match.as_deref() == Some("*") && state.blobs.contains_key(&name) { + Self::respond(StatusCode::Conflict, Some("BlobAlreadyExists"), Vec::new()) + } else { + let bytes = match request.body() { + Body::Bytes(bytes) => bytes.to_vec(), + Body::SeekableStream(_) => panic!("unexpected streaming upload"), + }; + state.blobs.insert(name, bytes); + Self::respond(StatusCode::Created, None, Vec::new()) + } + } + (Method::Get, false, Some(name)) => match state.blobs.get(&name) { + Some(bytes) => Self::respond(StatusCode::Ok, None, bytes.clone()), + None => Self::respond(StatusCode::NotFound, Some("BlobNotFound"), Vec::new()), + }, + (Method::Delete, false, Some(name)) => match state.blobs.remove(&name) { + Some(_) => Self::respond(StatusCode::Accepted, None, Vec::new()), + None => Self::respond(StatusCode::NotFound, Some("BlobNotFound"), Vec::new()), + }, + (method, _, _) => panic!("unexpected request {method:?} {path}?{query}"), + }; + Ok(response) + } +} + +struct Fixture { + runtime: Runtime, + service: FakeBlobService, + cache: Arc>, +} + +impl Fixture { + fn new(service: FakeBlobService) -> Self { + let runtime = Runtime::new().unwrap(); + let cache = runtime + .block_on(Self::connect(&service, runtime.handle().clone())) + .unwrap(); + Self { + runtime, + service, + cache: Arc::new(cache), + } + } + + async fn connect( + service: &FakeBlobService, + handle: tokio::runtime::Handle, + ) -> Result, Error> { + AzureBlobCache::connect_with_options( + ACCOUNT_URL, + CONTAINER, + None, + ClientOptions { + transport: Some(Transport::new(Arc::new(service.clone()))), + ..ClientOptions::default() + }, + ResponseCacheCodec, + handle, + ) + .await + } + + fn response_cache(&self) -> ResponseCache> { + ResponseCache::new(self.cache.clone()) + } + + fn stored_json(&self, key: &str) -> serde_json::Value { + serde_json::from_slice(&self.service.blob(key).expect("blob should exist")).unwrap() + } +} + +fn request(model: &str) -> ResponseCacheRequest { + ResponseCacheRequest::new(CacheKeyInput { + fields: vec![CacheKeyField { + name: "model".into(), + value: Some(model.into()), + api_parameter: true, + internal_parameter: false, + }], + preset: None, + namespace: None, + include_provider_parameters: false, + }) +} + +fn now() -> Duration { + Duration::from_secs(1_700_000_000) +} + +fn entry(value: serde_json::Value) -> CacheEntry { + CacheEntry { + timestamp: Some(1_700_000_000.5), + response: value, + } +} + +fn no_ttl() -> ExactCacheContext { + ExactCacheContext::default() +} + +fn with_ttl(seconds: u64) -> ExactCacheContext { + ExactCacheContext { + ttl: Some(Duration::from_secs(seconds)), + } +} + +#[test] +fn connect_creates_the_container_once() { + let fixture = Fixture::new(FakeBlobService::default()); + assert!(fixture.service.container_exists()); + assert_eq!( + fixture.service.requests(), + vec![RecordedRequest { + method: Method::Put, + path: format!("/{CONTAINER}"), + query: "restype=container".into(), + if_none_match: None, + }] + ); + assert_eq!(fixture.cache.account_url(), ACCOUNT_URL); + assert_eq!(fixture.cache.container_name(), CONTAINER); +} + +#[test] +fn connect_accepts_an_existing_container() { + let fixture = Fixture::new(FakeBlobService::with_existing_container()); + assert!(fixture.service.container_exists()); + assert_eq!(fixture.service.requests().len(), 1); +} + +#[test] +fn connect_accepts_account_urls_with_trailing_slash() { + let runtime = Runtime::new().unwrap(); + let service = FakeBlobService::default(); + let cache = runtime + .block_on(AzureBlobCache::connect_with_options( + "https://example.blob.core.windows.net/", + CONTAINER, + None, + ClientOptions { + transport: Some(Transport::new(Arc::new(service.clone()))), + ..ClientOptions::default() + }, + ResponseCacheCodec, + runtime.handle().clone(), + )) + .unwrap(); + assert_eq!(service.requests()[0].path, format!("/{CONTAINER}")); + assert_eq!(cache.account_url(), "https://example.blob.core.windows.net"); +} + +#[test] +fn connect_surfaces_service_failures() { + let runtime = Runtime::new().unwrap(); + let service = FakeBlobService::default(); + service.set_failing(true); + let result = runtime.block_on(Fixture::connect(&service, runtime.handle().clone())); + assert!(matches!(result, Err(Error::Unavailable))); +} + +#[test] +fn sync_set_and_get_round_trip_python_json_shape() { + let fixture = Fixture::new(FakeBlobService::default()); + let value = entry(json!({"choices": [{"message": {"content": "héllo 🌍"}}]})); + fixture + .cache + .set_cache("key-1", value.clone(), &no_ttl()) + .unwrap(); + + assert_eq!( + fixture.stored_json("key-1"), + json!({ + "timestamp": 1_700_000_000.5, + "response": {"choices": [{"message": {"content": "héllo 🌍"}}]} + }) + ); + assert_eq!( + fixture.cache.get_cache("key-1", &no_ttl()).unwrap(), + Some(value) + ); +} + +#[test] +fn sync_set_does_not_overwrite_an_existing_blob() { + let fixture = Fixture::new(FakeBlobService::default()); + fixture + .cache + .set_cache("key", entry(json!({"v": "first"})), &no_ttl()) + .unwrap(); + fixture + .cache + .set_cache("key", entry(json!({"v": "second"})), &no_ttl()) + .unwrap(); + + assert_eq!( + fixture.stored_json("key")["response"], + json!({"v": "first"}) + ); + let uploads: Vec<_> = fixture + .service + .requests() + .into_iter() + .filter(|request| request.method == Method::Put && request.path.ends_with("/key")) + .collect(); + assert_eq!(uploads.len(), 2); + assert!( + uploads + .iter() + .all(|request| request.if_none_match.as_deref() == Some("*")) + ); +} + +#[test] +fn async_set_overwrites_an_existing_blob() { + let fixture = Fixture::new(FakeBlobService::default()); + fixture.runtime.block_on(async { + fixture + .cache + .async_set_cache("key", entry(json!({"v": "first"})), no_ttl()) + .await + .unwrap(); + fixture + .cache + .async_set_cache("key", entry(json!({"v": "second"})), no_ttl()) + .await + .unwrap(); + assert_eq!( + fixture + .cache + .async_get_cache("key", &no_ttl()) + .await + .unwrap(), + Some(entry(json!({"v": "second"}))) + ); + }); + assert_eq!( + fixture.stored_json("key")["response"], + json!({"v": "second"}) + ); + assert!( + fixture + .service + .requests() + .iter() + .filter(|request| request.method == Method::Put && request.path.ends_with("/key")) + .all(|request| request.if_none_match.is_none()) + ); +} + +#[test] +fn missing_blobs_are_misses() { + let fixture = Fixture::new(FakeBlobService::default()); + assert_eq!(fixture.cache.get_cache("absent", &no_ttl()).unwrap(), None); + assert_eq!( + fixture + .runtime + .block_on(fixture.cache.async_get_cache("absent", &no_ttl())) + .unwrap(), + None + ); +} + +#[test] +fn ttl_is_ignored_and_entries_never_expire() { + let fixture = Fixture::new(FakeBlobService::default()); + assert_eq!(fixture.cache.get_ttl(&with_ttl(1)), None); + assert_eq!(fixture.cache.get_ttl(&no_ttl()), None); + + fixture + .cache + .set_cache("key", entry(json!("value")), &with_ttl(1)) + .unwrap(); + std::thread::sleep(Duration::from_millis(1100)); + assert_eq!( + fixture.cache.get_cache("key", &with_ttl(1)).unwrap(), + Some(entry(json!("value"))) + ); + assert!( + fixture + .service + .requests() + .iter() + .all(|request| !request.query.contains("expiry")) + ); +} + +#[test] +fn malformed_blobs_are_invalid_entries_and_response_cache_misses() { + let fixture = Fixture::new(FakeBlobService::default()); + fixture.service.seed_blob("broken-json", b"{not json"); + fixture + .service + .seed_blob("broken-utf8", &[0xff, 0xfe, 0x22]); + fixture + .service + .seed_blob("wrong-shape", br#"{"timestamp": "yesterday"}"#); + + for key in ["broken-json", "broken-utf8", "wrong-shape"] { + assert!(matches!( + fixture.cache.get_cache(key, &no_ttl()), + Err(Error::InvalidEntry) + )); + } + + let response_cache = fixture.response_cache(); + let broken = request("broken"); + fixture + .service + .seed_blob(&cache_key(&broken.key), b"{not json"); + assert_eq!(response_cache.lookup(&broken, now()).unwrap(), None); + assert_eq!( + fixture + .runtime + .block_on(response_cache.async_lookup(&broken, now())) + .unwrap(), + None + ); +} + +#[test] +fn batch_get_preserves_order_and_marks_misses_and_invalid_entries() { + let fixture = Fixture::new(FakeBlobService::default()); + fixture + .cache + .set_cache("a", entry(json!("A")), &no_ttl()) + .unwrap(); + fixture + .cache + .set_cache("c", entry(json!("C")), &no_ttl()) + .unwrap(); + fixture.service.seed_blob("bad", b"nope"); + let keys = ["c", "missing", "a", "bad"].map(String::from); + + let sync = fixture.cache.batch_get_cache(&keys, &no_ttl()).unwrap(); + assert_eq!( + sync, + vec![ + BatchEntry::Hit(entry(json!("C"))), + BatchEntry::Miss, + BatchEntry::Hit(entry(json!("A"))), + BatchEntry::Invalid, + ] + ); + + let asynchronous = fixture + .runtime + .block_on(fixture.cache.async_batch_get_cache(keys.to_vec(), no_ttl())) + .unwrap(); + assert_eq!(asynchronous, sync); + + let response_cache = fixture.response_cache(); + let requests = [request("hit"), request("missing"), request("bad")]; + response_cache + .store(&requests[0], json!("HIT"), now()) + .unwrap(); + fixture + .service + .seed_blob(&cache_key(&requests[2].key), b"nope"); + let hits = response_cache.lookup_batch(&requests, now()).unwrap(); + assert_eq!(hits.values, vec![Some(json!("HIT")), None, None]); + assert_eq!(hits.missing_indices, vec![1, 2]); + let async_hits = fixture + .runtime + .block_on(response_cache.async_lookup_batch(&requests, now())) + .unwrap(); + assert_eq!(async_hits.values, hits.values); +} + +#[test] +fn async_pipeline_writes_every_entry_with_overwrite() { + let fixture = Fixture::new(FakeBlobService::default()); + fixture.service.seed_blob("k2", b"stale"); + fixture + .runtime + .block_on(fixture.cache.async_set_cache_pipeline( + vec![ + ("k1".into(), entry(json!({"n": 1}))), + ("k2".into(), entry(json!({"n": 2}))), + ("k3".into(), entry(json!({"n": 3}))), + ], + with_ttl(30), + )) + .unwrap(); + assert_eq!(fixture.service.blob_names(), ["k1", "k2", "k3"]); + assert_eq!(fixture.stored_json("k2")["response"], json!({"n": 2})); +} + +#[test] +fn flush_deletes_every_blob_in_the_container() { + let fixture = Fixture::new(FakeBlobService::default()); + for key in ["x", "y", "z"] { + fixture + .cache + .set_cache(key, entry(json!(key)), &no_ttl()) + .unwrap(); + } + fixture.cache.flush_cache().unwrap(); + assert!(fixture.service.blob_names().is_empty()); + assert!(fixture.service.container_exists()); + + fixture + .cache + .set_cache("again", entry(json!(1)), &no_ttl()) + .unwrap(); + fixture + .runtime + .block_on(fixture.cache.async_flush_cache()) + .unwrap(); + assert!(fixture.service.blob_names().is_empty()); +} + +#[test] +fn service_failures_map_to_unavailable() { + let fixture = Fixture::new(FakeBlobService::default()); + fixture.service.set_failing(true); + assert!(matches!( + fixture.cache.get_cache("key", &no_ttl()), + Err(Error::Unavailable) + )); + assert!(matches!( + fixture.cache.set_cache("key", entry(json!(1)), &no_ttl()), + Err(Error::Unavailable) + )); + assert!(matches!( + fixture.cache.flush_cache(), + Err(Error::Unavailable) + )); + assert!(matches!( + fixture.runtime.block_on( + fixture + .cache + .async_set_cache_pipeline(vec![("k".into(), entry(json!(1)))], no_ttl()) + ), + Err(Error::Unavailable) + )); +} + +#[test] +fn test_connection_reports_container_reachability() { + let fixture = Fixture::new(FakeBlobService::default()); + let ok = fixture + .runtime + .block_on(fixture.cache.test_connection()) + .unwrap(); + assert_eq!(ok.status, CacheConnectionStatus::Success); + assert!(ok.error.is_none()); + + fixture.service.set_failing(true); + let failed = fixture + .runtime + .block_on(fixture.cache.test_connection()) + .unwrap(); + assert_eq!(failed.status, CacheConnectionStatus::Failed); + assert!(failed.error.is_some()); +} + +#[test] +fn disconnect_is_idempotent_and_keeps_data() { + let fixture = Fixture::new(FakeBlobService::default()); + fixture + .cache + .set_cache("key", entry(json!(1)), &no_ttl()) + .unwrap(); + fixture.runtime.block_on(async { + fixture.cache.disconnect().await.unwrap(); + fixture.cache.disconnect().await.unwrap(); + }); + assert_eq!( + fixture.cache.get_cache("key", &no_ttl()).unwrap(), + Some(entry(json!(1))) + ); +} + +#[test] +fn response_cache_stores_and_reads_through_the_backend() { + let fixture = Fixture::new(FakeBlobService::default()); + let response_cache = fixture.response_cache(); + let mut request = request("gpt"); + request.context = with_ttl(60); + let response = json!({"id": "chatcmpl-1"}); + response_cache + .store(&request, response.clone(), now()) + .unwrap(); + assert_eq!( + fixture.stored_json(&cache_key(&request.key)), + json!({"timestamp": 1_700_000_000.0, "response": {"id": "chatcmpl-1"}}) + ); + assert_eq!( + response_cache + .lookup(&request, now() + Duration::from_secs(3600)) + .unwrap(), + Some(response.clone()) + ); + assert_eq!( + fixture + .runtime + .block_on(response_cache.async_lookup(&request, now() + Duration::from_secs(3600))) + .unwrap(), + Some(response.clone()) + ); + fixture.runtime.block_on(async { + response_cache + .async_store(&request, json!("replaced"), now()) + .await + .unwrap(); + assert_eq!( + response_cache.async_lookup(&request, now()).await.unwrap(), + Some(json!("replaced")) + ); + response_cache.async_flush().await.unwrap(); + assert_eq!( + response_cache.async_lookup(&request, now()).await.unwrap(), + None + ); + }); +} + +#[test] +fn non_object_responses_are_written_serialized_like_python() { + let fixture = Fixture::new(FakeBlobService::default()); + fixture + .cache + .set_cache("s", entry(json!("plain")), &no_ttl()) + .unwrap(); + assert_eq!( + fixture.stored_json("s"), + json!({"timestamp": 1_700_000_000.5, "response": "\"plain\""}) + ); + assert_eq!( + fixture.cache.get_cache("s", &no_ttl()).unwrap(), + Some(entry(json!("plain"))) + ); +} diff --git a/litellm-rust/crates/cache-azure-blob/src/credential.rs b/litellm-rust/crates/cache-azure-blob/src/credential.rs new file mode 100644 index 00000000000..d1a3d0e44ec --- /dev/null +++ b/litellm-rust/crates/cache-azure-blob/src/credential.rs @@ -0,0 +1,84 @@ +use std::{ + fmt, + sync::Arc, + time::{Duration, SystemTime}, +}; + +use azure_core::{ + credentials::{AccessToken, TokenCredential, TokenRequestOptions}, + error::ErrorKind, + time::OffsetDateTime, +}; +use litellm_auth_azure::{AzureAuthInputs, AzureAuthService}; +use litellm_auth_types::ResolvedCredential; + +const STATIC_TOKEN_LIFETIME: Duration = Duration::from_secs(300); +const LLM_TOKEN_ENV: &str = "AZURE_AD_TOKEN"; + +type EnvLookup = Arc Option + Send + Sync>; + +pub struct AzureBlobCredential { + service: AzureAuthService, + env_lookup: EnvLookup, +} + +impl fmt::Debug for AzureBlobCredential { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.write_str("AzureBlobCredential") + } +} + +impl Default for AzureBlobCredential { + fn default() -> Self { + Self::new( + AzureAuthService::default(), + Arc::new(|name| std::env::var(name).ok()), + ) + } +} + +impl AzureBlobCredential { + pub fn new(service: AzureAuthService, env_lookup: EnvLookup) -> Self { + Self { + service, + env_lookup, + } + } +} + +#[async_trait::async_trait] +impl TokenCredential for AzureBlobCredential { + async fn get_token( + &self, + scopes: &[&str], + _options: Option>, + ) -> azure_core::Result { + let env_lookup = &self.env_lookup; + let lookup = move |name: &str| (name != LLM_TOKEN_ENV).then(|| env_lookup(name)).flatten(); + let credential = self + .service + .get_azure_ad_token( + &AzureAuthInputs::default_credential_for_scope(&scopes.join(" ")), + &lookup, + ) + .await + .map_err(|error| { + azure_core::Error::with_message(ErrorKind::Credential, error.to_string()) + })? + .ok_or_else(|| { + azure_core::Error::with_message( + ErrorKind::Credential, + "no Azure credential is available for blob storage", + ) + })?; + let (token, expires_on) = match credential.into_value() { + ResolvedCredential::AccessToken { token, expires_on } => (token, expires_on), + ResolvedCredential::Static(token) => (token, None), + }; + let expires_on = expires_on.unwrap_or_else(|| SystemTime::now() + STATIC_TOKEN_LIFETIME); + Ok(AccessToken::new( + token.expose().to_string(), + OffsetDateTime::from(expires_on), + )) + } +} diff --git a/litellm-rust/crates/cache-azure-blob/src/lib.rs b/litellm-rust/crates/cache-azure-blob/src/lib.rs new file mode 100644 index 00000000000..5ae752c111d --- /dev/null +++ b/litellm-rust/crates/cache-azure-blob/src/lib.rs @@ -0,0 +1,5 @@ +mod cache; +mod credential; + +pub use cache::AzureBlobCache; +pub use credential::AzureBlobCredential; diff --git a/litellm-rust/crates/cache-azure-blob/src/tests.rs b/litellm-rust/crates/cache-azure-blob/src/tests.rs new file mode 100644 index 00000000000..e69de29bb2d diff --git a/litellm-rust/crates/python-bridge/Cargo.toml b/litellm-rust/crates/python-bridge/Cargo.toml index 1eb2ec28036..7e48874a30b 100644 --- a/litellm-rust/crates/python-bridge/Cargo.toml +++ b/litellm-rust/crates/python-bridge/Cargo.toml @@ -21,6 +21,7 @@ tiktoken = ["litellm-token-counter/tiktoken"] [dependencies] bytes.workspace = true litellm-cache.workspace = true +litellm-cache-azure-blob.workspace = true litellm-cache-memory.workspace = true litellm-cache-redis.workspace = true litellm-cache-response.workspace = true diff --git a/litellm-rust/crates/python-bridge/src/cache/config.rs b/litellm-rust/crates/python-bridge/src/cache/config.rs index 0e7d6aee11d..91f56157bdc 100644 --- a/litellm-rust/crates/python-bridge/src/cache/config.rs +++ b/litellm-rust/crates/python-bridge/src/cache/config.rs @@ -73,9 +73,15 @@ pub(super) struct RedisCacheConfig { pub(super) connection: RedisConnectionConfig, } +pub(super) struct AzureBlobCacheConfig { + pub(super) account_url: String, + pub(super) container: String, +} + pub(super) enum CacheBackendConfig { Memory(MemoryCacheConfig), Redis(Box), + AzureBlob(AzureBlobCacheConfig), } #[allow(dead_code, reason = "consumed by the cache activation follow-up")] @@ -142,13 +148,18 @@ impl NativeCacheConfig { }))), Err(reason) => Ok(CacheConfigProjection::Unsupported(reason)), }, + Some(CacheType::AzureBlob) => project_azure_blob(&backend).map(|backend| { + CacheConfigProjection::Native(Box::new(Self { + policy, + backend: CacheBackendConfig::AzureBlob(backend), + })) + }), Some( CacheType::RedisSemantic | CacheType::ValkeySemantic | CacheType::S3 | CacheType::Disk | CacheType::QdrantSemantic - | CacheType::AzureBlob | CacheType::Gcs, ) | None => Ok(CacheConfigProjection::Unsupported( @@ -158,12 +169,12 @@ impl NativeCacheConfig { } pub(super) fn service_mismatch(&self, service: &NativeResponseCache) -> Option<&'static str> { - if service.default_ttl() - != Some(match &self.backend { - CacheBackendConfig::Memory(config) => config.default_ttl, - CacheBackendConfig::Redis(config) => config.default_ttl, - }) - { + let default_ttl = match &self.backend { + CacheBackendConfig::Memory(config) => Some(config.default_ttl), + CacheBackendConfig::Redis(config) => Some(config.default_ttl), + CacheBackendConfig::AzureBlob(_) => None, + }; + if service.default_ttl() != default_ttl { return Some("facade and native backend default TTLs must match"); } match &self.backend { @@ -185,10 +196,34 @@ impl NativeCacheConfig { CacheBackendConfig::Redis(config) => (service.namespace() != config.namespace.as_deref()) .then_some("facade and native backend namespaces must match"), + CacheBackendConfig::AzureBlob(config) => match service.azure_blob_identity() { + None => Some("facade and native backend types must match"), + Some((account_url, container)) + if account_url != config.account_url || container != config.container => + { + Some("facade and native backend containers must match") + } + Some(_) => None, + }, } } } +#[inline(never)] +fn project_azure_blob(backend: &Bound<'_, PyAny>) -> PyResult { + let client = backend.getattr("container_client")?; + let container = client.getattr("container_name")?.extract::()?; + let url = client.getattr("url")?.extract::()?; + let account_url = url + .strip_suffix(container.as_str()) + .and_then(|url| url.strip_suffix('/')) + .ok_or_else(|| PyValueError::new_err("Azure Blob container URL is malformed"))?; + Ok(AzureBlobCacheConfig { + account_url: account_url.to_string(), + container, + }) +} + #[inline(never)] fn project_memory(backend: &Bound<'_, PyAny>) -> PyResult { let max_size_kib = backend.getattr("max_size_per_item")?.extract::()?; diff --git a/litellm-rust/crates/python-bridge/src/cache/facade.rs b/litellm-rust/crates/python-bridge/src/cache/facade.rs index f2f86c14b37..1df6312503f 100644 --- a/litellm-rust/crates/python-bridge/src/cache/facade.rs +++ b/litellm-rust/crates/python-bridge/src/cache/facade.rs @@ -32,10 +32,23 @@ struct RedisPoolGuard { max_connections: usize, } +struct AzureBlobClientGuard { + sync_client: Py, + async_client: Py, + url: String, + container_name: String, +} + +enum ConnectionGuard { + None, + RedisPool(RedisPoolGuard), + AzureBlob(AzureBlobClientGuard), +} + pub(super) struct FacadeGuard { outer: ObjectGuard, backend: ObjectGuard, - redis_pool: Option, + connection: ConnectionGuard, } impl ObjectGuard { @@ -176,6 +189,60 @@ impl RedisPoolGuard { } } +impl AzureBlobClientGuard { + fn capture(backend: &Bound<'_, PyAny>) -> PyResult { + let sync_client = backend.getattr("container_client")?; + Ok(Self { + url: sync_client.getattr("url")?.extract::()?, + container_name: sync_client.getattr("container_name")?.extract::()?, + sync_client: sync_client.unbind(), + async_client: backend.getattr("async_container_client")?.unbind(), + }) + } + + fn matches(&self, py: Python<'_>, backend: &Bound<'_, PyAny>) -> PyResult { + let sync_client = backend.getattr("container_client")?; + Ok(self.sync_client.bind(py).is(&sync_client) + && self + .async_client + .bind(py) + .is(&backend.getattr("async_container_client")?) + && self.url == sync_client.getattr("url")?.extract::()? + && self.container_name == sync_client.getattr("container_name")?.extract::()?) + } + + fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> { + visit.call(&self.sync_client)?; + visit.call(&self.async_client) + } +} + +impl ConnectionGuard { + fn capture(kind: &str, backend: &Bound<'_, PyAny>) -> PyResult { + Ok(match kind { + "redis" => Self::RedisPool(RedisPoolGuard::capture(backend)?), + "azure-blob" => Self::AzureBlob(AzureBlobClientGuard::capture(backend)?), + _ => Self::None, + }) + } + + fn matches(&self, py: Python<'_>, backend: &Bound<'_, PyAny>) -> PyResult { + match self { + Self::None => Ok(true), + Self::RedisPool(guard) => guard.matches(py, backend), + Self::AzureBlob(guard) => guard.matches(py, backend), + } + } + + fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> { + match self { + Self::None => Ok(()), + Self::RedisPool(guard) => guard.traverse(visit), + Self::AzureBlob(guard) => guard.traverse(visit), + } + } +} + impl FacadeGuard { pub(super) fn capture( py: Python<'_>, @@ -192,6 +259,11 @@ impl FacadeGuard { let (module, name, cache_kind) = match kind { "memory" => ("litellm.caching.in_memory_cache", "InMemoryCache", "local"), "redis" => ("litellm.caching.redis_cache", "RedisCache", "redis"), + "azure-blob" => ( + "litellm.caching.azure_blob_cache", + "AzureBlobCache", + "azure-blob", + ), _ => unreachable!(), }; let backend = facade.getattr("cache")?; @@ -237,9 +309,7 @@ impl FacadeGuard { "redis_flush_size", ], )?, - redis_pool: (kind == "redis") - .then(|| RedisPoolGuard::capture(&backend)) - .transpose()?, + connection: ConnectionGuard::capture(kind, &backend)?, }) } @@ -251,19 +321,13 @@ impl FacadeGuard { if !self.backend.matches(py, &backend)? { return Ok(false); } - match &self.redis_pool { - Some(guard) => guard.matches(py, &backend), - None => Ok(true), - } + self.connection.matches(py, &backend) } pub(super) fn traverse(&self, visit: PyVisit<'_>) -> Result<(), PyTraverseError> { self.outer.traverse(&visit)?; self.backend.traverse(&visit)?; - if let Some(guard) = &self.redis_pool { - guard.traverse(&visit)?; - } - Ok(()) + self.connection.traverse(&visit) } } diff --git a/litellm-rust/crates/python-bridge/src/cache/handle.rs b/litellm-rust/crates/python-bridge/src/cache/handle.rs index 8251b3df06c..69988980d64 100644 --- a/litellm-rust/crates/python-bridge/src/cache/handle.rs +++ b/litellm-rust/crates/python-bridge/src/cache/handle.rs @@ -1,4 +1,4 @@ -use litellm_host_python::release_gil; +use litellm_host_python::{release_gil, run_sync_value}; use pyo3::{PyTraverseError, PyVisit, exceptions::PyRuntimeError, prelude::*}; use super::{cache_error, facade::FacadeGuard, native::NativeResponseCache, request::duration}; @@ -51,6 +51,21 @@ impl CacheTestHandle { }) } + #[staticmethod] + #[pyo3(signature = (account_url, container))] + fn azure_blob(py: Python<'_>, account_url: String, container: String) -> PyResult { + let service = run_sync_value(py, async move { + NativeResponseCache::azure_blob(&account_url, &container) + .await + .map_err(cache_error) + })?; + Ok(Self { + service, + guard: None, + pid: std::process::id(), + }) + } + #[getter] fn backend(&self) -> &'static str { self.service.kind() diff --git a/litellm-rust/crates/python-bridge/src/cache/native.rs b/litellm-rust/crates/python-bridge/src/cache/native.rs index a9475429e45..08c65905e84 100644 --- a/litellm-rust/crates/python-bridge/src/cache/native.rs +++ b/litellm-rust/crates/python-bridge/src/cache/native.rs @@ -1,6 +1,7 @@ use std::{sync::Arc, time::Duration}; use litellm_cache::{CacheCodec, CacheConnectionResult, Error}; +use litellm_cache_azure_blob::AzureBlobCache; use litellm_cache_memory::InMemoryCache; use litellm_cache_redis::RedisCache; use litellm_cache_response::{ @@ -15,6 +16,7 @@ pub(super) enum NativeResponseCache { cache: Arc>>, buffer: Option>, }, + AzureBlob(Arc>>), } impl NativeResponseCache { @@ -43,6 +45,29 @@ impl NativeResponseCache { buffer: None, }) } + + pub async fn azure_blob(account_url: &str, container: &str) -> Result { + let backend = AzureBlobCache::connect( + account_url, + container, + ResponseCacheCodec, + tokio::runtime::Handle::current(), + ) + .await?; + Ok(Self::AzureBlob(Arc::new(ResponseCache::new(Arc::new( + backend, + ))))) + } + + pub fn azure_blob_identity(&self) -> Option<(&str, &str)> { + match self { + Self::AzureBlob(cache) => Some(( + cache.backend().account_url(), + cache.backend().container_name(), + )), + Self::Memory(_) | Self::Redis { .. } => None, + } + } } impl NativeResponseCache { @@ -50,6 +75,7 @@ impl NativeResponseCache { match self { Self::Memory(_) => "memory", Self::Redis { .. } => "redis", + Self::AzureBlob(_) => "azure-blob", } } @@ -57,12 +83,13 @@ impl NativeResponseCache { match self { Self::Memory(cache) => cache.default_ttl(), Self::Redis { cache, .. } => cache.default_ttl(), + Self::AzureBlob(cache) => cache.default_ttl(), } } pub fn namespace(&self) -> Option<&str> { match self { - Self::Memory(_) => None, + Self::Memory(_) | Self::AzureBlob(_) => None, Self::Redis { cache, .. } => cache.backend().namespace(), } } @@ -70,14 +97,14 @@ impl NativeResponseCache { pub fn capacity(&self) -> Option { match self { Self::Memory(cache) => Some(cache.backend().max_size_in_memory()), - Self::Redis { .. } => None, + Self::Redis { .. } | Self::AzureBlob(_) => None, } } pub fn max_entry_bytes(&self) -> Option { match self { Self::Memory(cache) => cache.backend().max_entry_bytes(), - Self::Redis { .. } => None, + Self::Redis { .. } | Self::AzureBlob(_) => None, } } @@ -87,7 +114,7 @@ impl NativeResponseCache { cache, buffer: flush_size.map(|flush_size| Arc::new(WriteBuffer::new(flush_size))), }, - memory => memory, + other => other, } } @@ -99,6 +126,7 @@ impl NativeResponseCache { match self { Self::Memory(cache) => cache.lookup(request, now), Self::Redis { cache, .. } => cache.lookup(request, now), + Self::AzureBlob(cache) => cache.lookup(request, now), } } @@ -111,6 +139,7 @@ impl NativeResponseCache { match self { Self::Memory(cache) => cache.store(request, response, now), Self::Redis { cache, .. } => cache.store(request, response, now), + Self::AzureBlob(cache) => cache.store(request, response, now), } } @@ -122,6 +151,7 @@ impl NativeResponseCache { match self { Self::Memory(cache) => cache.lookup_batch(requests, now), Self::Redis { cache, .. } => cache.lookup_batch(requests, now), + Self::AzureBlob(cache) => cache.lookup_batch(requests, now), } } @@ -133,6 +163,7 @@ impl NativeResponseCache { match self { Self::Memory(cache) => cache.async_lookup(request, now).await, Self::Redis { cache, .. } => cache.async_lookup(request, now).await, + Self::AzureBlob(cache) => cache.async_lookup(request, now).await, } } @@ -152,6 +183,7 @@ impl NativeResponseCache { cache, buffer: Some(buffer), } => buffer.async_store(cache, request, response, now).await, + Self::AzureBlob(cache) => cache.async_store(request, response, now).await, } } @@ -163,6 +195,7 @@ impl NativeResponseCache { match self { Self::Memory(cache) => cache.async_lookup_batch(requests, now).await, Self::Redis { cache, .. } => cache.async_lookup_batch(requests, now).await, + Self::AzureBlob(cache) => cache.async_lookup_batch(requests, now).await, } } @@ -174,6 +207,7 @@ impl NativeResponseCache { match self { Self::Memory(cache) => cache.async_store_batch(entries, now).await, Self::Redis { cache, .. } => cache.async_store_batch(entries, now).await, + Self::AzureBlob(cache) => cache.async_store_batch(entries, now).await, } } @@ -186,6 +220,7 @@ impl NativeResponseCache { } cache.async_flush().await } + Self::AzureBlob(cache) => cache.async_flush().await, } } @@ -193,6 +228,7 @@ impl NativeResponseCache { match self { Self::Memory(cache) => cache.test_connection().await, Self::Redis { cache, .. } => cache.test_connection().await, + Self::AzureBlob(cache) => cache.test_connection().await, } } } diff --git a/tests/test_litellm_rust/test_cache.py b/tests/test_litellm_rust/test_cache.py index c35cb1a20fb..99e371d8c2e 100644 --- a/tests/test_litellm_rust/test_cache.py +++ b/tests/test_litellm_rust/test_cache.py @@ -2,8 +2,10 @@ import asyncio import contextvars import gc import json +import os import threading import time +import uuid import weakref from collections.abc import Generator from types import SimpleNamespace @@ -13,8 +15,10 @@ from urllib.parse import urlparse import fakeredis import pytest import redis +from azure.storage.blob import ContainerClient import litellm +from litellm.caching.azure_blob_cache import AzureBlobCache from litellm.caching.caching import Cache, disable_cache, enable_cache, update_cache from litellm.caching.in_memory_cache import InMemoryCache from litellm.rust_bridge import _native @@ -45,6 +49,36 @@ def redis_url() -> Generator[str]: worker.join(timeout=5) +@pytest.fixture +def azure_blob_facade() -> Generator[Cache]: + account_url: Final = os.environ.get("AZURE_BLOB_CACHE_ACCOUNT_URL") + if account_url is None: + pytest.skip( + "live Azure Blob parity needs AZURE_BLOB_CACHE_ACCOUNT_URL plus DefaultAzureCredential inputs in the environment" + ) + facade: Final = Cache( + type=LiteLLMCacheType.AZURE_BLOB, + azure_account_url=account_url, + azure_blob_container=f"litellm-parity-{uuid.uuid4().hex[:12]}", + ) + backend: Final = facade.cache + assert isinstance(backend, AzureBlobCache) + try: + yield facade + finally: + backend.container_client.delete_container() + asyncio.run(backend.disconnect()) + + +def azure_blob_handle(facade: Cache) -> _native._CacheTestHandle: + backend: Final = facade.cache + assert isinstance(backend, AzureBlobCache) + return _native._CacheTestHandle.azure_blob( + backend.container_client.url.removesuffix(f"/{backend.container_client.container_name}"), + backend.container_client.container_name, + ) + + def test_existing_constructor_and_global_are_unchanged() -> None: facade: Final = Cache(type=LiteLLMCacheType.LOCAL) assert type(facade.cache) is InMemoryCache @@ -361,6 +395,89 @@ def test_facade_registration_rejects_mismatched_capacity() -> None: _native._CacheTestHandle.memory(capacity=7)._bind_facade(facade) +def test_azure_blob_facade_serves_natively_and_python_reads_the_same_blobs(azure_blob_facade: Cache) -> None: + backend: Final = azure_blob_facade.cache + assert isinstance(backend, AzureBlobCache) + handle: Final = azure_blob_handle(azure_blob_facade) + assert handle.backend == "azure-blob" + account_url: Final = backend.container_client.url.removesuffix(f"/{backend.container_client.container_name}") + with pytest.raises(TypeError, match="containers must match"): + _native._CacheTestHandle.azure_blob(account_url, f"{backend.container_client.container_name}-other")._bind_facade( + azure_blob_facade + ) + handle._bind_facade(azure_blob_facade) + resolver: Final = _native._CacheTestResolver(SimpleNamespace(cache=azure_blob_facade)) + native: Final = resolver.resolve() + assert native.kind == "native" + + response: Final = {"choices": [{"text": "caf\u00e9 \u2603"}], "usage": {"total_tokens": 3}, "flag": True, "empty": None} + native.store({**request("sync"), "ttl_seconds": 0.001}, response) + native.store(request("sync"), {"choices": [{"text": "second"}]}) + time.sleep(0.01) + stored: Final = json.loads(backend.container_client.download_blob("sync").readall()) + assert stored["response"] == response + assert isinstance(stored["timestamp"], float) + assert native.lookup(request("sync")) == response + assert cast(CacheLookup, azure_blob_facade).get_cache(cache_key="sync") == response + + backend.set_cache("python", {"timestamp": time.time(), "response": response}) + backend.set_cache("legacy", "bare legacy value") + backend.container_client.upload_blob("invalid", b"{not json", overwrite=True) + assert native.lookup(request("python")) == response + assert native.lookup(request("legacy")) == cast(CacheLookup, azure_blob_facade).get_cache(cache_key="legacy") + assert native.lookup_batch([request("python"), request("missing"), request("invalid"), request("sync")]) == { + "values": [response, None, None, response], + "missing_indices": [1, 2], + } + + with rebound(azure_blob_facade, "ttl", 12): + assert resolver.resolve().kind == "python_callback" + with rebound(backend, "container_client", ContainerClient.from_container_url(backend.container_client.url)): + assert resolver.resolve().kind == "python_callback" + + def custom_get(*_args: object, **_kwargs: object) -> None: + return None + + with rebound(backend, "get_cache", custom_get): + assert resolver.resolve().kind == "python_callback" + assert resolver.resolve().kind == "python_callback" + assert cast(CacheLookup, azure_blob_facade).get_cache(cache_key="sync") == response + + class CustomBlobCache(AzureBlobCache): + pass + + with rebound(azure_blob_facade, "cache", CustomBlobCache(account_url, backend.container_client.container_name)): + assert resolver.resolve().kind == "python_callback" + with pytest.raises(TypeError): + azure_blob_handle(azure_blob_facade)._bind_facade(azure_blob_facade) + + +async def test_azure_blob_native_async_writes_overwrite_batch_and_flush_like_python(azure_blob_facade: Cache) -> None: + backend: Final = azure_blob_facade.cache + assert isinstance(backend, AzureBlobCache) + azure_blob_handle(azure_blob_facade)._bind_facade(azure_blob_facade) + binding: Final = _native._CacheTestResolver(SimpleNamespace(cache=azure_blob_facade)).resolve() + assert binding.kind == "native" + ping: Final = cast(dict[str, object], await binding.ping()) + assert ping["status"] == "success", ping + + await binding.async_store(request("async"), {"value": 1}) + await binding.async_store({**request("async"), "ttl_seconds": 0.001}, {"value": 2}) + time.sleep(0.01) + assert await binding.async_lookup(request("async")) == {"value": 2} + assert await backend.async_get_cache("async") == json.loads(backend.container_client.download_blob("async").readall()) + assert cast(CacheLookup, azure_blob_facade).get_cache(cache_key="async") == {"value": 2} + + await binding.async_store_batch([request("first"), request("second")], [{"value": 3}, {"value": 4}]) + assert await binding.async_lookup_batch([request("second"), request("missing"), request("first")]) == { + "values": [{"value": 4}, None, {"value": 3}], + "missing_indices": [1], + } + await binding.async_flush() + assert [blob.name for blob in backend.container_client.list_blobs()] == [] + assert await binding.async_lookup(request("async")) is None + + async def test_redis_facade_buffers_native_async_writes(redis_url: str) -> None: parsed: Final = urlparse(redis_url) with rebound(litellm, "default_redis_ttl", 60): From a64febb3e7d32aef73dc7b90665c33d0b65b4805 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Mon, 21 Sep 2026 14:01:08 -0700 Subject: [PATCH 33/44] fix(streaming): keep an explicit provider prompt_tokens=0 or completion_tokens=0 in streamed usage The stream chunk builder started its per-chunk accumulators at 0 and adopted only nonzero counts, then fell back to litellm's tokenizer whenever the accumulated value was falsy, so a provider that reported an explicit 0 for prompt or completion tokens was billed the estimate instead. The accumulators now start at None, a usage chunk that reports a count marks it reported (a later chunk's 0 never replaces a reported nonzero), and the estimate only runs when no chunk reported the count. The Anthropic message_start cursor reset now yields None so the estimate still covers a cancelled stream, and Ollama chat streaming only attaches usage on the done chunk when both counts are present instead of inventing 0/0 on every chunk --- .../streaming_chunk_builder_utils.py | 46 ++++++----- litellm/llms/ollama/chat/transformation.py | 22 +++-- .../streaming_chunk_builder_utils.py | 4 +- .../test_streaming_chunk_builder_cursor.py | 26 +++--- .../test_streaming_chunk_builder_utils.py | 80 +++++++++++++++++++ .../ollama/test_ollama_chat_transformation.py | 47 ++++++++++- 6 files changed, 183 insertions(+), 42 deletions(-) diff --git a/litellm/litellm_core_utils/streaming_chunk_builder_utils.py b/litellm/litellm_core_utils/streaming_chunk_builder_utils.py index fcd55c844c6..025db65a7ce 100644 --- a/litellm/litellm_core_utils/streaming_chunk_builder_utils.py +++ b/litellm/litellm_core_utils/streaming_chunk_builder_utils.py @@ -839,15 +839,17 @@ class ChunkProcessor: UsagePerChunk, ) - # # Update usage information if needed - prompt_tokens = 0 - completion_tokens = 0 + # None means no usage chunk reported the count, which is the only case + # calculate_usage() estimates with the tokenizer. An explicit provider 0 + # is a reported count and stays 0; a reported count is never replaced by + # a later chunk's 0 (Ollama sends 0/0 on every chunk before the done one). + prompt_tokens: int | None = None + completion_tokens: int | None = None # Anthropic's `message_start` SSE event carries usage.output_tokens=1 as a # cursor/placeholder; the real value only arrives in `message_delta`. - # If a stream is cancelled before `message_delta` lands, the last-wins - # accumulator below leaves completion_tokens stuck at 1 — which then - # bypasses the `completion_tokens or token_counter(...)` fallback in - # calculate_usage() because 1 is truthy. Count the completion-bearing + # If a stream is cancelled before `message_delta` lands, the accumulator + # below leaves completion_tokens stuck at 1, a reported count that + # calculate_usage() would keep. Count the completion-bearing # usage events so `_reset_anthropic_cursor_completion_tokens` can tell a # legitimate single-token reply (Anthropic emits 1 in BOTH message_start # AND message_delta, so >=2 events is positive evidence message_delta @@ -875,10 +877,15 @@ class ChunkProcessor: if usage_chunk is not None: usage_chunk_dict = self._usage_chunk_calculation_helper(usage_chunk) - if usage_chunk_dict["prompt_tokens"] is not None and usage_chunk_dict["prompt_tokens"] > 0: + if usage_chunk_dict["prompt_tokens"] is not None and ( + usage_chunk_dict["prompt_tokens"] > 0 or prompt_tokens is None + ): prompt_tokens = usage_chunk_dict["prompt_tokens"] - if usage_chunk_dict["completion_tokens"] is not None and usage_chunk_dict["completion_tokens"] > 0: + if usage_chunk_dict["completion_tokens"] is not None and ( + usage_chunk_dict["completion_tokens"] > 0 or completion_tokens is None + ): completion_tokens = usage_chunk_dict["completion_tokens"] + if usage_chunk_dict["completion_tokens"] is not None and usage_chunk_dict["completion_tokens"] > 0: completion_usage_updates += 1 if usage_chunk_dict["cache_creation_input_tokens"] is not None and ( usage_chunk_dict["cache_creation_input_tokens"] > 0 or cache_creation_input_tokens is None @@ -995,10 +1002,10 @@ class ChunkProcessor: @staticmethod def _reset_anthropic_cursor_completion_tokens( chunks: Sequence["_UsageBearingChunk | ModelResponse"], - completion_tokens: int, + completion_tokens: int | None, completion_usage_updates: int, - ) -> int: - """Reset a stale Anthropic ``message_start`` cursor placeholder to 0. + ) -> int | None: + """Reset a stale Anthropic ``message_start`` cursor placeholder to unreported. See the ``completion_usage_updates`` comment in ``_calculate_usage_per_chunk``. The accumulated value is NOT a stale @@ -1006,8 +1013,8 @@ class ChunkProcessor: carried a ``finish_reason`` (positive evidence ``message_delta`` arrived). Otherwise the only completion update we ever saw was the Anthropic ``message_start`` cursor, a small placeholder whose magnitude - varies per request (1 and 8 both observed live), so reset to 0 and let - ``calculate_usage()``'s ``or token_counter(...)`` fallback estimate from + varies per request (1 and 8 both observed live), so reset to None and let + ``calculate_usage()``'s ``token_counter(...)`` fallback estimate from the actually-received text and reasoning instead. Gated on ``custom_llm_provider == "anthropic"`` so the heuristic (which encodes Anthropic's specific message_start SSE shape) does not silently affect @@ -1028,7 +1035,7 @@ class ChunkProcessor: custom_llm_provider = hp.get("custom_llm_provider") if custom_llm_provider == "anthropic": - return 0 + return None return completion_tokens def calculate_usage( @@ -1063,15 +1070,18 @@ class ChunkProcessor: cost: Final[float | None] = calculated_usage_per_chunk["cost"] try: - returned_usage.prompt_tokens = prompt_tokens or ( - count_prompt_tokens() if count_prompt_tokens else token_counter(model=model, messages=messages) + returned_usage.prompt_tokens = ( + prompt_tokens + if prompt_tokens is not None + else (count_prompt_tokens() if count_prompt_tokens else token_counter(model=model, messages=messages)) ) except Exception: # don't allow this failing to block a complete streaming response from being returned print_verbose("token_counter failed, assuming prompt tokens is 0") returned_usage.prompt_tokens = 0 returned_usage.completion_tokens = ( completion_tokens - or ( + if completion_tokens is not None + else ( token_counter( model=model, text=completion_output, diff --git a/litellm/llms/ollama/chat/transformation.py b/litellm/llms/ollama/chat/transformation.py index 181894646e3..d39b4049903 100644 --- a/litellm/llms/ollama/chat/transformation.py +++ b/litellm/llms/ollama/chat/transformation.py @@ -1,6 +1,6 @@ import json import time -from collections.abc import AsyncIterator, Iterator +from collections.abc import AsyncIterator, Iterator, Mapping from typing import TYPE_CHECKING, Any, Final, cast from httpx._models import Headers, Response @@ -420,6 +420,18 @@ class OllamaChatConfig(BaseConfig): ) +def _done_chunk_usage(chunk: Mapping[str, object]) -> ChatCompletionUsageBlock | None: + prompt_eval_count: Final = chunk.get("prompt_eval_count") + eval_count: Final = chunk.get("eval_count") + if chunk.get("done") is not True or not isinstance(prompt_eval_count, int) or not isinstance(eval_count, int): + return None + return ChatCompletionUsageBlock( + prompt_tokens=prompt_eval_count, + completion_tokens=eval_count, + total_tokens=prompt_eval_count + eval_count, + ) + + class OllamaChatCompletionResponseIterator(BaseModelResponseIterator): started_reasoning_content: bool = False finished_reasoning_content: bool = False @@ -528,17 +540,11 @@ class OllamaChatCompletionResponseIterator(BaseModelResponseIterator): ) ] - usage: Final = ChatCompletionUsageBlock( - prompt_tokens=chunk.get("prompt_eval_count", 0), - completion_tokens=chunk.get("eval_count", 0), - total_tokens=chunk.get("prompt_eval_count", 0) + chunk.get("eval_count", 0), - ) - return ModelResponseStream( id=str(uuid.uuid4()), object="chat.completion.chunk", created=int(time.time()), # ollama created_at is in UTC - usage=usage, + usage=_done_chunk_usage(chunk), model=chunk["model"], choices=choices, ) diff --git a/litellm/types/litellm_core_utils/streaming_chunk_builder_utils.py b/litellm/types/litellm_core_utils/streaming_chunk_builder_utils.py index f981089d370..3deb307881f 100644 --- a/litellm/types/litellm_core_utils/streaming_chunk_builder_utils.py +++ b/litellm/types/litellm_core_utils/streaming_chunk_builder_utils.py @@ -4,8 +4,8 @@ from ..utils import CompletionTokensDetails, PromptTokensDetailsWrapper, ServerT class UsagePerChunk(TypedDict): - prompt_tokens: int - completion_tokens: int + prompt_tokens: ReadOnly[int | None] + completion_tokens: ReadOnly[int | None] cache_creation_input_tokens: int | None cache_read_input_tokens: int | None server_tool_use: ServerToolUse | None diff --git a/tests/test_litellm/litellm_core_utils/test_streaming_chunk_builder_cursor.py b/tests/test_litellm/litellm_core_utils/test_streaming_chunk_builder_cursor.py index 8617c5b81e8..b4a1f7733e7 100644 --- a/tests/test_litellm/litellm_core_utils/test_streaming_chunk_builder_cursor.py +++ b/tests/test_litellm/litellm_core_utils/test_streaming_chunk_builder_cursor.py @@ -15,7 +15,7 @@ calculate_usage() never fires, and the request is billed for 1 output token even when several thousand tokens of text were actually streamed. These tests pin the post-fix behavior: completion_tokens should reset -to 0 when the only update we saw was the cursor, allowing the +to None when the only update we saw was the cursor, allowing the text-based fallback to estimate from the real completion text. """ @@ -63,10 +63,10 @@ def _make_chunk( class TestAnthropicCursorBug: """The core regression: completion_tokens=1 cursor must not leak through.""" - def test_only_message_start_cursor_resets_completion_to_zero(self): + def test_only_message_start_cursor_resets_completion_to_unreported(self): """ Stream cancelled before message_delta — only the message_start cursor - (output_tokens=1) was seen. Per-chunk accumulator must reset to 0 so + (output_tokens=1) was seen. Per-chunk accumulator must reset to None so token_counter fallback can estimate from completion text. """ # Anthropic message_start: input_tokens accurate, output_tokens=1 cursor @@ -83,11 +83,11 @@ class TestAnthropicCursorBug: result = processor._calculate_usage_per_chunk(chunks=chunks) assert result["prompt_tokens"] == 1024 - # The cursor value of 1 must NOT leak through — should be reset to 0 + # The cursor value of 1 must NOT leak through — should be reset to None # so the text-based fallback estimates the real completion length. - assert result["completion_tokens"] == 0, ( + assert result["completion_tokens"] is None, ( "completion_tokens=1 from message_start cursor leaked through. " - "Should reset to 0 when only cursor was seen, so token_counter " + "Should reset to None when only cursor was seen, so token_counter " "fallback in calculate_usage() can estimate from completion text." ) @@ -233,10 +233,10 @@ class TestAnthropicCursorBug: result = processor._calculate_usage_per_chunk(chunks=chunks) assert result["cache_read_input_tokens"] == 4096 - assert result["completion_tokens"] == 0, ( + assert result["completion_tokens"] is None, ( "cache chunks alone don't count as completion progress — only " "completion_tokens > 0 in a usage event proves real output happened. " - "Reset to 0 forces token_counter fallback." + "Reset to None forces token_counter fallback." ) @pytest.mark.parametrize("placeholder", [1, 3, 8]) @@ -326,7 +326,7 @@ class TestAnthropicCursorBug: ] processor = ChunkProcessor(chunks=chunks, messages=[]) result = processor._calculate_usage_per_chunk(chunks=chunks) - assert result["completion_tokens"] == 0 + assert result["completion_tokens"] is None assert result["completion_tokens_details"] is None def test_estimated_reasoning_is_capped_to_trusted_completion_total(self): @@ -403,11 +403,11 @@ class TestNonAnthropicStreamingIntact: result = processor._calculate_usage_per_chunk(chunks=chunks) assert result["completion_tokens"] == 5 - def test_no_usage_chunks_leaves_zero(self): - """Stream with zero usage info → completion_tokens stays 0 + def test_no_usage_chunks_leaves_unreported(self): + """Stream with zero usage info → both counts stay None (token_counter fallback will handle it).""" chunks = [_make_chunk(content="hi"), _make_chunk(content=" there")] processor = ChunkProcessor(chunks=chunks, messages=[]) result = processor._calculate_usage_per_chunk(chunks=chunks) - assert result["prompt_tokens"] == 0 - assert result["completion_tokens"] == 0 + assert result["prompt_tokens"] is None + assert result["completion_tokens"] is None diff --git a/tests/test_litellm/litellm_core_utils/test_streaming_chunk_builder_utils.py b/tests/test_litellm/litellm_core_utils/test_streaming_chunk_builder_utils.py index 8e7ed52fade..5d75c6699cf 100644 --- a/tests/test_litellm/litellm_core_utils/test_streaming_chunk_builder_utils.py +++ b/tests/test_litellm/litellm_core_utils/test_streaming_chunk_builder_utils.py @@ -1648,3 +1648,83 @@ def test_calculate_usage_falls_back_to_prompt_counter_when_mock_stream_has_no_ad ) assert usage.prompt_tokens == 77 + + +_ZERO_USAGE_TEXT_CHUNKS: Final = ( + _openai_chunk(choices=[{"index": 0, "delta": {"role": "assistant", "content": "Hi"}, "finish_reason": None}]), + _openai_chunk(choices=[{"index": 0, "delta": {"content": " there"}, "finish_reason": None}]), + _openai_chunk(choices=[{"index": 0, "delta": {}, "finish_reason": "stop"}]), +) + + +@pytest.mark.parametrize( + "reported", + [ + pytest.param({"prompt_tokens": 0, "completion_tokens": 17, "total_tokens": 17}, id="zero_prompt"), + pytest.param({"prompt_tokens": 5, "completion_tokens": 0, "total_tokens": 5}, id="zero_completion"), + pytest.param({"prompt_tokens": 0, "completion_tokens": 0, "total_tokens": 0}, id="all_zero"), + ], +) +def test_calculate_usage_keeps_an_explicit_provider_zero(reported: Mapping[str, int]) -> None: + chunks: Final = [*_ZERO_USAGE_TEXT_CHUNKS, _openai_chunk(choices=[], usage=reported)] + + usage: Final = ChunkProcessor(chunks=chunks).calculate_usage( + chunks=chunks, + model="gpt-5.4-mini", + completion_output="Hi there", + messages=[{"role": "user", "content": "hi"}], + count_prompt_tokens=lambda: 999, + ) + + assert (usage.prompt_tokens, usage.completion_tokens, usage.total_tokens) == ( + reported["prompt_tokens"], + reported["completion_tokens"], + reported["prompt_tokens"] + reported["completion_tokens"], + ) + + +def test_stream_chunk_builder_keeps_an_explicit_zero_prompt_count_end_to_end() -> None: + reported: Final = {"prompt_tokens": 0, "completion_tokens": 17, "total_tokens": 17} + chunks: Final = [*_ZERO_USAGE_TEXT_CHUNKS, _openai_chunk(choices=[], usage=reported)] + + response: Final = stream_chunk_builder(chunks=chunks, messages=[{"role": "user", "content": "hi"}]) + + assert response is not None + assert (response.usage.prompt_tokens, response.usage.completion_tokens, response.usage.total_tokens) == (0, 17, 17) + + +def test_calculate_usage_estimates_only_when_no_chunk_reported_usage() -> None: + chunks: Final = list(_ZERO_USAGE_TEXT_CHUNKS) + + usage: Final = ChunkProcessor(chunks=chunks).calculate_usage( + chunks=chunks, + model="gpt-5.4-mini", + completion_output="Hi there", + count_prompt_tokens=lambda: 77, + ) + + assert usage.prompt_tokens == 77 + assert usage.completion_tokens > 0 + assert usage.total_tokens == 77 + usage.completion_tokens + + +def test_calculate_usage_keeps_a_reported_count_over_a_later_chunks_zero() -> None: + chunks: Final = [ + _openai_chunk( + choices=[{"index": 0, "delta": {"role": "assistant", "content": "Hi"}, "finish_reason": None}], + usage={"prompt_tokens": 5, "completion_tokens": 0, "total_tokens": 5}, + ), + _openai_chunk( + choices=[{"index": 0, "delta": {"content": " there"}, "finish_reason": "stop"}], + usage={"prompt_tokens": 0, "completion_tokens": 17, "total_tokens": 17}, + ), + ] + + usage: Final = ChunkProcessor(chunks=chunks).calculate_usage( + chunks=chunks, + model="gpt-5.4-mini", + completion_output="Hi there", + count_prompt_tokens=lambda: 999, + ) + + assert (usage.prompt_tokens, usage.completion_tokens, usage.total_tokens) == (5, 17, 22) diff --git a/tests/test_litellm/llms/ollama/test_ollama_chat_transformation.py b/tests/test_litellm/llms/ollama/test_ollama_chat_transformation.py index 25f9645faa0..c7fba21d222 100644 --- a/tests/test_litellm/llms/ollama/test_ollama_chat_transformation.py +++ b/tests/test_litellm/llms/ollama/test_ollama_chat_transformation.py @@ -22,7 +22,7 @@ import json from unittest.mock import MagicMock import litellm -from litellm.types.utils import Choices, Message, ModelResponse +from litellm.types.utils import Choices, Message, ModelResponse, ModelResponseStream class TestEvent(BaseModel): @@ -944,3 +944,48 @@ class TestOllamaToolCallTransformation: assert tool_msg["content"] == "Sunny, 72°F" assert "tool_call_id" in tool_msg, "tool_call_id must be forwarded to Ollama" assert tool_msg["tool_call_id"] == "call_abc123" + + +class TestOllamaStreamingUsage: + @staticmethod + def _parse(chunk: dict) -> ModelResponseStream: + iterator = OllamaChatCompletionResponseIterator(streaming_response=iter([]), sync_stream=True) + return iterator.chunk_parser(chunk) + + def test_done_chunk_reports_the_counts_ollama_sent(self): + result = self._parse( + { + "model": "qwen3:0.6b", + "message": {"role": "assistant", "content": ""}, + "done": True, + "done_reason": "stop", + "prompt_eval_count": 100, + "eval_count": 50, + } + ) + + assert result.usage is not None + assert (result.usage.prompt_tokens, result.usage.completion_tokens, result.usage.total_tokens) == (100, 50, 150) + + def test_done_chunk_without_counts_reports_no_usage_instead_of_zeros(self): + result = self._parse( + { + "model": "qwen3:0.6b", + "message": {"role": "assistant", "content": ""}, + "done": True, + "done_reason": "stop", + } + ) + + assert result.usage is None + + def test_chunk_before_done_reports_no_usage(self): + result = self._parse( + { + "model": "qwen3:0.6b", + "message": {"role": "assistant", "content": "Hi"}, + "done": False, + } + ) + + assert result.usage is None From 1b568319d070875c5137cc7f942cba85c80e6133 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Mon, 21 Sep 2026 14:10:09 -0700 Subject: [PATCH 34/44] fix(types): blank non-string datadog tool text fields, look the spend table up by name --- .../integrations/datadog/datadog_llm_obs.py | 24 +++++++++++-------- litellm/proxy/db/db_spend_update_writer.py | 22 ++++++++--------- .../datadog/test_datadog_llm_obs.py | 17 +++++++++++++ 3 files changed, 41 insertions(+), 22 deletions(-) diff --git a/litellm/integrations/datadog/datadog_llm_obs.py b/litellm/integrations/datadog/datadog_llm_obs.py index 3afe38ea075..98aac7336bf 100644 --- a/litellm/integrations/datadog/datadog_llm_obs.py +++ b/litellm/integrations/datadog/datadog_llm_obs.py @@ -254,13 +254,17 @@ def _reasoning_output_tokens(usage_object: Mapping[str, object] | None) -> float ) -def _mapping_field(source: Mapping[str, object], key: str) -> Mapping[str, Any]: +def _mapping_field(source: Mapping[str, object], key: str) -> Mapping[str, object]: """The value at `key` when it is a mapping, else an empty one.""" value: Final = source.get(key) return value if isinstance(value, dict) else _EMPTY_MAPPING -def _content_blocks(message: Mapping[str, object]) -> tuple[Mapping[str, Any], ...]: +def _text_field(source: Mapping[str, object], key: str, default: str = "") -> str: + return _safe_identifier(source.get(key, default)) + + +def _content_blocks(message: Mapping[str, object]) -> tuple[Mapping[str, object], ...]: content: Final = message.get("content") if not isinstance(content, list): return () @@ -293,10 +297,10 @@ def _to_dd_tool_calls(message: Mapping[str, object]) -> tuple[ToolCall, ...]: raw_tool_calls: Final = message.get("tool_calls") openai_calls: Final = tuple( ToolCall( - name=function.get("name", ""), + name=_text_field(function, "name"), arguments=_to_dd_arguments(function.get("arguments", "")), - tool_id=tool_call.get("id", ""), - type=tool_call.get("type", "function"), + tool_id=_text_field(tool_call, "id"), + type=_text_field(tool_call, "type", "function"), ) for tool_call in (raw_tool_calls if isinstance(raw_tool_calls, list) else ()) if isinstance(tool_call, dict) @@ -304,9 +308,9 @@ def _to_dd_tool_calls(message: Mapping[str, object]) -> tuple[ToolCall, ...]: ) anthropic_calls: Final = tuple( ToolCall( - name=block.get("name", ""), + name=_text_field(block, "name"), arguments=_to_dd_arguments(block.get("input") or {}), - tool_id=block.get("id", ""), + tool_id=_text_field(block, "id"), type="tool_use", ) for block in _content_blocks(message) @@ -402,12 +406,12 @@ def _to_dd_messages(messages: object) -> tuple[Message, ...]: def _to_dd_tool_definition(entry: Mapping[str, object]) -> ToolDefinition | None: function: Final = entry.get("function") - declared: Final[Mapping[str, Any]] = function if isinstance(function, dict) else entry - name: Final = declared.get("name") + declared: Final[Mapping[str, object]] = function if isinstance(function, dict) else entry + name: Final = _text_field(declared, "name") if not name: return None schema: Final = declared.get("parameters") or declared.get("input_schema") - description: Final = declared.get("description", "") + description: Final = _text_field(declared, "description") if not isinstance(schema, dict): return ToolDefinition(name=name, description=description) return ToolDefinition(name=name, description=description, schema=schema) diff --git a/litellm/proxy/db/db_spend_update_writer.py b/litellm/proxy/db/db_spend_update_writer.py index 865a8c39f03..5381071172f 100644 --- a/litellm/proxy/db/db_spend_update_writer.py +++ b/litellm/proxy/db/db_spend_update_writer.py @@ -19,7 +19,7 @@ from typing import TYPE_CHECKING, Any, Final, Literal, Protocol, TypeAlias, Type from urllib.parse import quote, unquote from pydantic import TypeAdapter -from typing_extensions import LiteralString, ReadOnly, TypedDict, assert_never +from typing_extensions import LiteralString, ReadOnly, TypedDict import litellm from litellm._logging import verbose_proxy_logger @@ -142,17 +142,15 @@ _EntitySpendTable: TypeAlias = Literal[ def _entity_spend_table(batcher: _SpendBatch, table_accessor: _EntitySpendTable) -> BatchTable: - match table_accessor: - case "litellm_tagtable": - return batcher.litellm_tagtable - case "litellm_agentstable": - return batcher.litellm_agentstable - case "litellm_modelaccessgroupbudgettable": - return batcher.litellm_modelaccessgroupbudgettable - case "litellm_projecttable": - return batcher.litellm_projecttable - case _ as unreachable: - assert_never(unreachable) + tables: Final[Mapping[_EntitySpendTable, BatchTable]] = MappingProxyType( + { + "litellm_tagtable": batcher.litellm_tagtable, + "litellm_agentstable": batcher.litellm_agentstable, + "litellm_modelaccessgroupbudgettable": batcher.litellm_modelaccessgroupbudgettable, + "litellm_projecttable": batcher.litellm_projecttable, + } + ) + return tables[table_accessor] class _SpendBatchManager(Protocol): diff --git a/tests/test_litellm/integrations/datadog/test_datadog_llm_obs.py b/tests/test_litellm/integrations/datadog/test_datadog_llm_obs.py index a2f81091893..77d2518696c 100644 --- a/tests/test_litellm/integrations/datadog/test_datadog_llm_obs.py +++ b/tests/test_litellm/integrations/datadog/test_datadog_llm_obs.py @@ -110,6 +110,23 @@ def test_output_tool_calls_use_the_datadog_tool_call_schema(logger: DataDogLLMOb assert "function" not in message["tool_calls"][0] +def test_tool_call_identifiers_that_are_not_strings_are_blanked_not_stringified(logger: DataDogLLMObsLogger) -> None: + payload = build( + logger, + response_message={ + "role": "assistant", + "content": None, + "tool_calls": [ + {"id": {"nested": "call_1"}, "type": 7, "function": {"name": ["get_weather"], "arguments": "{}"}} + ], + }, + ) + + assert payload["meta"]["output"]["messages"][0]["tool_calls"] == [ + {"name": "", "arguments": {}, "tool_id": "", "type": ""} + ] + + def test_tool_calls_are_not_duplicated_into_metadata(logger: DataDogLLMObsLogger) -> None: """The flat `output_tool_calls.*` keys were a second copy of a fact that now has its own field.""" payload = build( From b5548082c621f536b6a9ab243508f40d9da34fad Mon Sep 17 00:00:00 2001 From: Yujong Lee Date: Mon, 21 Sep 2026 21:12:07 +0000 Subject: [PATCH 35/44] build(rust): raise native extension size limit to 30 MB for the Azure Blob SDK Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .github/scripts/verify_linux_native_wheel.py | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/.github/scripts/verify_linux_native_wheel.py b/.github/scripts/verify_linux_native_wheel.py index 4fb8f068eb0..0adbc015ad0 100644 --- a/.github/scripts/verify_linux_native_wheel.py +++ b/.github/scripts/verify_linux_native_wheel.py @@ -205,7 +205,7 @@ def main( native_module: Final = load_native_module(native_path) native_module_loads: Final = native_module is not None panic_test_hook_absent: Final = native_module is not None and not hasattr(native_module, "_panic_for_test") - native_size_limit: Final = 25_000_000 + native_size_limit: Final = 30_000_000 native_size_within_limit: Final = native_member.file_size <= native_size_limit validations: Final = ( (f"Python tag is {EXPECTED_PYTHON_TAG}", python_tag == EXPECTED_PYTHON_TAG), @@ -222,7 +222,7 @@ def main( ("Python extension entry point is present", extension_entry_point_present), ("Native module loads", native_module_loads), ("Production module omits the panic test hook", panic_test_hook_absent), - ("Native extension does not exceed 25 MB", native_size_within_limit), + ("Native extension does not exceed 30 MB", native_size_within_limit), ("Wheel contents are valid", not unexpected_members), ) @@ -267,7 +267,7 @@ def main( ), ( not native_size_within_limit, - f"native extension exceeds 20 MB: {native_member.file_size / 1_000_000:.2f} MB", + f"native extension exceeds 30 MB: {native_member.file_size / 1_000_000:.2f} MB", ), (bool(unexpected_members), f"wheel contains unexpected build artifacts: {', '.join(unexpected_members)}"), ) From 0e281d458fa71ae5b9c12006e5986115776f5f50 Mon Sep 17 00:00:00 2001 From: Yujong Lee Date: Mon, 21 Sep 2026 21:16:19 +0000 Subject: [PATCH 36/44] fix(rust): treat ConditionNotMet as an existing blob on sync Azure Blob writes Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../crates/cache-azure-blob/src/cache.rs | 23 ++++++++----- .../cache-azure-blob/src/cache/tests.rs | 34 ++++++++++++++++++- 2 files changed, 47 insertions(+), 10 deletions(-) diff --git a/litellm-rust/crates/cache-azure-blob/src/cache.rs b/litellm-rust/crates/cache-azure-blob/src/cache.rs index f28b8c9a641..efad0bc7f82 100644 --- a/litellm-rust/crates/cache-azure-blob/src/cache.rs +++ b/litellm-rust/crates/cache-azure-blob/src/cache.rs @@ -19,7 +19,6 @@ use url::Url; use crate::credential::AzureBlobCredential; -/// Synchronous methods block on `runtime` and therefore must run outside of it pub struct AzureBlobCache { container: BlobContainerClient, codec: C, @@ -54,14 +53,15 @@ impl AzureBlobCache { codec: C, runtime: Handle, ) -> Result { - let mut url = Url::parse(account_url).map_err(|_| Error::Unavailable)?; - let account_url = url.as_str().trim_end_matches('/').to_string(); - url.path_segments_mut() - .map_err(|()| Error::Unavailable)? - .pop_if_empty() - .push(container); + let account_url = Url::parse(account_url) + .map_err(|_| Error::Unavailable)? + .as_str() + .trim_end_matches('/') + .to_string(); + let container_url = + Url::parse(&format!("{account_url}/{container}")).map_err(|_| Error::Unavailable)?; let client = BlobContainerClient::new( - url, + container_url, credential, Some(BlobContainerClientOptions { client_options, @@ -108,7 +108,7 @@ impl AzureBlobCache { .await { Ok(_) => Ok(()), - Err(error) if is_storage_error(&error, StorageErrorCode::BlobAlreadyExists) => Ok(()), + Err(error) if !overwrite && is_already_present(&error) => Ok(()), Err(_) => Err(Error::Unavailable), } } @@ -153,6 +153,11 @@ impl AzureBlobCache { } } +fn is_already_present(error: &azure_core::Error) -> bool { + is_storage_error(error, StorageErrorCode::BlobAlreadyExists) + || is_storage_error(error, StorageErrorCode::ConditionNotMet) +} + fn is_storage_error(error: &azure_core::Error, code: StorageErrorCode) -> bool { matches!( error.kind(), diff --git a/litellm-rust/crates/cache-azure-blob/src/cache/tests.rs b/litellm-rust/crates/cache-azure-blob/src/cache/tests.rs index fd116a0e28f..580674450c3 100644 --- a/litellm-rust/crates/cache-azure-blob/src/cache/tests.rs +++ b/litellm-rust/crates/cache-azure-blob/src/cache/tests.rs @@ -39,6 +39,7 @@ struct FakeState { blobs: BTreeMap>, requests: Vec, failing: bool, + precondition_conflicts: bool, } #[derive(Clone, Default)] @@ -79,6 +80,10 @@ impl FakeBlobService { self.state.lock().unwrap().failing = failing; } + fn set_precondition_conflicts(&self, enabled: bool) { + self.state.lock().unwrap().precondition_conflicts = enabled; + } + fn requests(&self) -> Vec { self.state.lock().unwrap().requests.clone() } @@ -158,7 +163,15 @@ impl HttpClient for FakeBlobService { } (Method::Put, false, Some(name)) => { if if_none_match.as_deref() == Some("*") && state.blobs.contains_key(&name) { - Self::respond(StatusCode::Conflict, Some("BlobAlreadyExists"), Vec::new()) + if state.precondition_conflicts { + Self::respond( + StatusCode::PreconditionFailed, + Some("ConditionNotMet"), + Vec::new(), + ) + } else { + Self::respond(StatusCode::Conflict, Some("BlobAlreadyExists"), Vec::new()) + } } else { let bytes = match request.body() { Body::Bytes(bytes) => bytes.to_vec(), @@ -369,6 +382,25 @@ fn sync_set_does_not_overwrite_an_existing_blob() { ); } +#[test] +fn sync_set_treats_a_precondition_conflict_as_an_existing_blob() { + let fixture = Fixture::new(FakeBlobService::default()); + fixture.service.set_precondition_conflicts(true); + fixture + .cache + .set_cache("key", entry(json!({"v": "first"})), &no_ttl()) + .unwrap(); + fixture + .cache + .set_cache("key", entry(json!({"v": "second"})), &no_ttl()) + .unwrap(); + + assert_eq!( + fixture.stored_json("key")["response"], + json!({"v": "first"}) + ); +} + #[test] fn async_set_overwrites_an_existing_blob() { let fixture = Fixture::new(FakeBlobService::default()); From a2ae80ec9bb61057ae186a4930152f5ba7f41d61 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Mon, 21 Sep 2026 14:19:52 -0700 Subject: [PATCH 37/44] fix(router): wrap every Responses and Messages fallback hop for mid-stream failover The /v1/responses and /v1/messages streaming wrappers only ever wrapped the primary's stream, so a hop reached through the regular fallback chain had no mid-stream handler: its failure re-raised, or the outer wrapper retried the same entry with a fresh attempted set and never reached the rest of the list. Every attempt of the chain now runs through a per-endpoint attempt function that wraps its own stream, mirroring chat completions, and the per-request fallback and retry overrides ride a frozen carrier so each hop's re-entry still sees them after the retry layer pops them. --- litellm/router.py | 170 ++++++++---------- .../router_utils/fallback_event_handlers.py | 74 +++++++- ...st_router_aresponses_streaming_fallback.py | 133 +++++++++++++- tests/test_litellm/test_router.py | 55 +++++- 4 files changed, 328 insertions(+), 104 deletions(-) diff --git a/litellm/router.py b/litellm/router.py index d9dd48a2c6a..9a5c770e78a 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -184,14 +184,17 @@ from litellm.router_utils.cooldown_handlers import ( is_caller_timeout_408, ) from litellm.router_utils.fallback_event_handlers import ( + MID_STREAM_FALLBACK_CONTROLS_KEY, AttemptedFallbackTargets, _check_non_standard_fallback_format, + carry_over_pre_routing_selection, clear_pre_routing_selection, fallback_lookup_groups, fallbacks_disabled_for_request, get_fallback_model_group_for_lookup_groups, - get_pre_routing_selection, has_unattempted_fallback_target, + mid_stream_fallback_hop_kwargs, + per_request_fallback_controls, record_disable_fallbacks, record_pre_routing_selection, run_async_fallback, @@ -3305,12 +3308,7 @@ class Router: content_policy_fallbacks: Final[list | None] = initial_kwargs.get( "content_policy_fallbacks", self.content_policy_fallbacks ) - # Re-enter via the per-attempt helper so the fallback chain - # picks deployments through - # _ageneric_api_call_with_fallbacks_helper. - # original_generic_function is preserved by the caller so - # the helper knows what underlying API to invoke per attempt. - initial_kwargs["original_function"] = self._ageneric_api_call_with_fallbacks_helper + initial_kwargs["original_function"] = self._ageneric_api_call_with_fallbacks_responses_attempt if e.is_pre_first_chunk or not e.generated_content: # No content generated before the error — retry with the # original input. Adding a continuation prompt would @@ -5141,22 +5139,28 @@ class Router: request_kwargs=None, ) - async def _ageneric_api_call_with_fallbacks(self, model: str, original_function: Callable, **kwargs): + async def _ageneric_api_call_with_fallbacks( + self, model: str, original_function: Callable, attempt_function: Callable | None = None, **kwargs + ): """ Helper function to make a generic LLM API call through the router, this allows you to use retries/fallbacks with litellm router + + attempt_function runs every attempt of the chain instead of the plain helper, so a streaming + endpoint can wrap each attempt's stream with its own mid-stream fallback handling. """ try: kwargs["model"] = model kwargs["original_generic_function"] = original_function - kwargs["original_function"] = self._ageneric_api_call_with_fallbacks_helper + kwargs["original_function"] = attempt_function or self._ageneric_api_call_with_fallbacks_helper + if attempt_function is not None: + controls: Final = per_request_fallback_controls(kwargs) + kwargs[MID_STREAM_FALLBACK_CONTROLS_KEY] = controls # rebind-ok: forwarded to every hop self._update_kwargs_before_fallbacks(model=model, kwargs=kwargs, metadata_variable_name="litellm_metadata") verbose_router_logger.debug( "Inside ageneric_api_call_with_fallbacks() - model: %s; kwargs: %s", model, kwargs ) response: Final = await self.async_function_with_fallbacks(**kwargs) return response - - return response except Exception as e: asyncio.create_task( send_llm_exception_alert( @@ -5277,61 +5281,42 @@ class Router: self, original_function: Callable, **kwargs: Any ) -> Union["ResponsesAPIResponse", "BaseResponsesAPIStreamingIterator"]: """ - _ageneric_api_call_with_fallbacks for the Responses API, with the - addition of mid-stream fallback handling. - - When stream=True and the underlying call returns a - BaseResponsesAPIStreamingIterator, wrap it with - _aresponses_streaming_iterator so MidStreamFallbackError raised - during iteration triggers the Router's cross-provider fallback chain. + _ageneric_api_call_with_fallbacks for the Responses API, with every attempt's stream + carrying its own mid-stream fallback handling + (see _ageneric_api_call_with_fallbacks_responses_attempt). + """ + return await self._ageneric_api_call_with_fallbacks( + original_function=original_function, + attempt_function=self._ageneric_api_call_with_fallbacks_responses_attempt, + **kwargs, + ) + + async def _ageneric_api_call_with_fallbacks_responses_attempt( + self, + model: str, + original_generic_function: Callable, + **kwargs: object, # kwargs-ok: forwarded verbatim to the per-attempt helper, shape varies per call site + ) -> Union["ResponsesAPIResponse", "BaseResponsesAPIStreamingIterator"]: + """ + One attempt of the Responses API fallback chain. A streaming result is wrapped with + _aresponses_streaming_iterator over this attempt's own kwargs, so a fallback hop that + fails mid-stream resumes the original group's chain instead of re-raising; the name keeps + _get_router_metadata_variable_name resolving to litellm_metadata for every hop. """ - from litellm.litellm_core_utils.core_helpers import safe_deep_copy from litellm.responses.streaming_iterator import ( BaseResponsesAPIStreamingIterator, ) - # Snapshot the request kwargs before _ageneric_api_call_with_fallbacks - # mutates them. A shallow copy alone is not enough: the primary - # attempt mutates nested dicts in place — notably `litellm_metadata`, - # which `_update_kwargs_with_deployment` populates with - # deployment-specific fields (`deployment`, `model_info`, `api_base`, - # tags, etc.). Without an explicit copy of that dict, the shallow - # copy would still share its reference, leaking primary-deployment - # metadata into the mid-stream fallback request. - # - # We avoid deep-copying the full kwargs because it can contain - # non-deepcopyable objects (logging handles, async clients, etc.); - # `safe_deep_copy` deep-copies the metadata dicts key-by-key with a - # fallback to the original reference for any non-picklable value. - # The original_generic_function is preserved so the per-attempt - # helper knows which underlying API to call on fallback. - # The pre-routing hook stamps its tier selection into this bucket during the primary - # attempt; seeding it before the snapshot gives both the live kwargs and the copy a - # bucket, so the post-call carry-over below always has somewhere to read and write. - kwargs.setdefault("litellm_metadata", {}) # mutable-ok: shared bucket # rebind-ok: stamp must be readable here - - fallback_kwargs: Final[dict[str, object]] = kwargs.copy() - if isinstance(fallback_kwargs.get("litellm_metadata"), dict): - fallback_kwargs["litellm_metadata"] = safe_deep_copy(fallback_kwargs["litellm_metadata"]) - if isinstance(fallback_kwargs.get("metadata"), dict): - fallback_kwargs["metadata"] = safe_deep_copy(fallback_kwargs["metadata"]) - fallback_kwargs["original_generic_function"] = original_function - - response: Final = await self._ageneric_api_call_with_fallbacks(original_function=original_function, **kwargs) - - # The snapshot predates the pre-routing hook, so the tier it stamped into the live kwargs - # is carried over write-or-clear: a stale or caller-supplied selection left in the copy - # would key the mid-stream fallback lookup off a tier this attempt never routed to. - clear_pre_routing_selection(fallback_kwargs) - live_pre_routing_selection: Final = get_pre_routing_selection(kwargs) - if live_pre_routing_selection is not None: - record_pre_routing_selection(fallback_kwargs, live_pre_routing_selection) - + controls: Final = kwargs.pop(MID_STREAM_FALLBACK_CONTROLS_KEY, None) + hop_kwargs: Final = mid_stream_fallback_hop_kwargs( + model=model, original_generic_function=original_generic_function, controls=controls, kwargs=kwargs + ) + response: Final = await self._ageneric_api_call_with_fallbacks_helper( + model=model, original_generic_function=original_generic_function, **kwargs + ) + carry_over_pre_routing_selection(live_kwargs=kwargs, snapshot=hop_kwargs) if kwargs.get("stream") and isinstance(response, BaseResponsesAPIStreamingIterator): - return await self._aresponses_streaming_iterator( - response=response, - initial_kwargs=fallback_kwargs, - ) + return await self._aresponses_streaming_iterator(response=response, initial_kwargs=hop_kwargs) return response async def _aanthropic_messages_streaming_iterator( @@ -5560,7 +5545,7 @@ class Router: content_policy_fallbacks: Final[list | None] = initial_kwargs.get( # mutable-ok: matches the param below "content_policy_fallbacks", self.content_policy_fallbacks ) - initial_kwargs["original_function"] = self._ageneric_api_call_with_fallbacks_helper + initial_kwargs["original_function"] = self._ageneric_api_call_with_fallbacks_anthropic_messages_attempt self._update_kwargs_before_fallbacks( model=model_group, kwargs=initial_kwargs, @@ -5614,46 +5599,41 @@ class Router: **kwargs: object, # kwargs-ok: forwarded verbatim to original_function, shape varies per call site ) -> Union["AnthropicMessagesResponse", AsyncIterator[bytes]]: """ - _ageneric_api_call_with_fallbacks for anthropic_messages, with the - addition of mid-stream fallback handling (see - _aanthropic_messages_streaming_iterator). Parity with + _ageneric_api_call_with_fallbacks for anthropic_messages, with every attempt's stream + carrying its own mid-stream fallback handling + (see _ageneric_api_call_with_fallbacks_anthropic_messages_attempt). Parity with _aresponses_with_streaming_fallbacks for the Responses API. """ - from litellm.litellm_core_utils.core_helpers import safe_deep_copy - - # Snapshot the request kwargs before the primary attempt mutates them - # in place: _update_kwargs_with_deployment writes deployment-specific - # fields (deployment, model_info, api_base, tags, ...) into the - # SAME litellm_metadata/metadata dicts a shallow .copy() would still - # share, leaking primary-deployment metadata into the mid-stream - # fallback request. safe_deep_copy avoids deep-copying the full - # kwargs (which can hold non-deepcopyable logging handles/clients). - # The pre-routing hook stamps its tier selection into this bucket during the primary - # attempt; seeding it before the snapshot gives both the live kwargs and the copy a - # bucket, so the post-call carry-over below always has somewhere to read and write. - kwargs.setdefault("litellm_metadata", {}) # mutable-ok: shared bucket # rebind-ok: stamp must be readable here - - fallback_kwargs: Final[dict[str, object]] = kwargs.copy() # mutable-ok: mutated below before re-entry - if isinstance(fallback_kwargs.get("litellm_metadata"), dict): - fallback_kwargs["litellm_metadata"] = safe_deep_copy(fallback_kwargs["litellm_metadata"]) - if isinstance(fallback_kwargs.get("metadata"), dict): - fallback_kwargs["metadata"] = safe_deep_copy(fallback_kwargs["metadata"]) - fallback_kwargs["original_generic_function"] = original_function - - response: Final = await self._ageneric_api_call_with_fallbacks(original_function=original_function, **kwargs) - - # The snapshot predates the pre-routing hook, so the tier it stamped into the live kwargs - # is carried over write-or-clear: a stale or caller-supplied selection left in the copy - # would key the mid-stream fallback lookup off a tier this attempt never routed to. - clear_pre_routing_selection(fallback_kwargs) - live_pre_routing_selection: Final = get_pre_routing_selection(kwargs) - if live_pre_routing_selection is not None: - record_pre_routing_selection(fallback_kwargs, live_pre_routing_selection) + return await self._ageneric_api_call_with_fallbacks( + original_function=original_function, + attempt_function=self._ageneric_api_call_with_fallbacks_anthropic_messages_attempt, + **kwargs, + ) + async def _ageneric_api_call_with_fallbacks_anthropic_messages_attempt( + self, + model: str, + original_generic_function: Callable, + **kwargs: object, # kwargs-ok: forwarded verbatim to the per-attempt helper, shape varies per call site + ) -> Union["AnthropicMessagesResponse", AsyncIterator[bytes]]: + """ + One attempt of the anthropic_messages fallback chain. A streaming result is wrapped with + _aanthropic_messages_streaming_iterator over this attempt's own kwargs, so a fallback hop + that fails mid-stream resumes the original group's chain instead of re-raising; the name + keeps _get_router_metadata_variable_name resolving to litellm_metadata for every hop. + """ + controls: Final = kwargs.pop(MID_STREAM_FALLBACK_CONTROLS_KEY, None) + hop_kwargs: Final = mid_stream_fallback_hop_kwargs( + model=model, original_generic_function=original_generic_function, controls=controls, kwargs=kwargs + ) + response: Final = await self._ageneric_api_call_with_fallbacks_helper( + model=model, original_generic_function=original_generic_function, **kwargs + ) + carry_over_pre_routing_selection(live_kwargs=kwargs, snapshot=hop_kwargs) if kwargs.get("stream") and hasattr(response, "__aiter__"): return await self._aanthropic_messages_streaming_iterator( response=cast("AsyncIterator[bytes]", response), # cast-ok: stream=True always returns a byte iterator - initial_kwargs=fallback_kwargs, + initial_kwargs=hop_kwargs, ) return response diff --git a/litellm/router_utils/fallback_event_handlers.py b/litellm/router_utils/fallback_event_handlers.py index 08b9246e562..f3ab2c5e493 100644 --- a/litellm/router_utils/fallback_event_handlers.py +++ b/litellm/router_utils/fallback_event_handlers.py @@ -1,6 +1,6 @@ import hashlib import json -from collections.abc import Mapping, Sequence +from collections.abc import Callable, Mapping, Sequence from dataclasses import dataclass from datetime import datetime from enum import Enum @@ -10,7 +10,7 @@ from typing import TYPE_CHECKING, Any, Final import litellm from litellm._logging import verbose_router_logger from litellm.integrations.custom_logger import CustomLogger -from litellm.litellm_core_utils.core_helpers import get_metadata_variable_name_from_kwargs +from litellm.litellm_core_utils.core_helpers import get_metadata_variable_name_from_kwargs, safe_deep_copy from litellm.litellm_core_utils.sensitive_data_masker import mask_sensitive_structure from litellm.router_utils.add_retry_fallback_headers import ( add_fallback_headers_to_response, @@ -284,6 +284,76 @@ def get_pre_routing_selection(kwargs: Mapping[str, object]) -> str | None: return next((selected for selected in selections if isinstance(selected, str) and selected), None) +def carry_over_pre_routing_selection(live_kwargs: Mapping[str, object], snapshot: Mapping[str, object]) -> None: + """ + Replace whatever selection the snapshot carries with the one the pre-routing hook stamped + into the live kwargs while routing this attempt, so a mid-stream fallback keys its lookup + off the tier this attempt actually routed to. + """ + clear_pre_routing_selection(snapshot) + live_selection: Final = get_pre_routing_selection(live_kwargs) + if live_selection is not None: + record_pre_routing_selection(snapshot, live_selection) + + +MID_STREAM_FALLBACK_CONTROLS_KEY: Final = "_mid_stream_fallback_controls" +_PER_REQUEST_FALLBACK_CONTROL_KEYS: Final = ( + "fallbacks", + "context_window_fallbacks", + "content_policy_fallbacks", + "num_retries", + "model_group_retry_policy", +) + + +@dataclass(frozen=True, slots=True) +class MidStreamFallbackControls: + """ + The per-request fallback and retry overrides every streaming attempt must see again. + + async_function_with_retries pops them before the attempt function runs, so without this + carrier a fallback hop's own mid-stream re-entry would fall back to the router-level settings. + """ + + overrides: Mapping[str, object] + + +_NO_FALLBACK_CONTROLS: Final = MidStreamFallbackControls(MappingProxyType({})) + + +def per_request_fallback_controls(kwargs: Mapping[str, object]) -> MidStreamFallbackControls: + return MidStreamFallbackControls( + MappingProxyType({key: kwargs[key] for key in _PER_REQUEST_FALLBACK_CONTROL_KEYS if key in kwargs}) + ) + + +def mid_stream_fallback_hop_kwargs( + model: str, + original_generic_function: Callable[..., object], + controls: object, + kwargs: Mapping[str, object], +) -> dict[str, object]: # mutable-ok: the streaming iterators rewrite it in place when they re-enter the chain + """ + The kwargs one streaming attempt re-enters the fallback chain with if its stream fails. + + A shallow copy keeps ``attempted_targets`` shared with the outer chain, so entries this + request already tried are never retried; the metadata buckets are copied key by key because + the attempt writes deployment-specific fields into them in place. + """ + hop_controls: Final = controls if isinstance(controls, MidStreamFallbackControls) else _NO_FALLBACK_CONTROLS + copied_buckets: Final = MappingProxyType( + {name: safe_deep_copy(kwargs[name]) for name in _ROUTER_METADATA_BUCKETS if isinstance(kwargs.get(name), dict)} + ) + return { # mutable-ok: handed to the streaming iterator as its initial_kwargs, which it rewrites on re-entry + **kwargs, + **copied_buckets, + **hop_controls.overrides, + MID_STREAM_FALLBACK_CONTROLS_KEY: hop_controls, + "model": model, + "original_generic_function": original_generic_function, + } + + DISABLE_FALLBACKS_METADATA_KEY: Final = "_disable_fallbacks" diff --git a/tests/router_unit_tests/test_router_aresponses_streaming_fallback.py b/tests/router_unit_tests/test_router_aresponses_streaming_fallback.py index 0d33435cf7a..18b10e9c8e5 100644 --- a/tests/router_unit_tests/test_router_aresponses_streaming_fallback.py +++ b/tests/router_unit_tests/test_router_aresponses_streaming_fallback.py @@ -251,7 +251,7 @@ async def test_aresponses_with_streaming_fallbacks_non_streaming_passthrough(): with patch.object( router, - "_ageneric_api_call_with_fallbacks", + "_ageneric_api_call_with_fallbacks_helper", new=AsyncMock(return_value=plain_response), ): out = await router._aresponses_with_streaming_fallbacks( @@ -278,7 +278,7 @@ async def test_aresponses_with_streaming_fallbacks_wraps_streaming_iterator(): with patch.object( router, - "_ageneric_api_call_with_fallbacks", + "_ageneric_api_call_with_fallbacks_helper", new=AsyncMock(return_value=streaming_iter), ), patch.object( router, @@ -294,6 +294,135 @@ async def test_aresponses_with_streaming_fallbacks_wraps_streaming_iterator(): mock_wrap.assert_awaited_once() +# -------- every fallback entry stays reachable across hops -------- + + +def _make_three_tier_router(**router_kwargs) -> Router: + return Router( + model_list=[ + {"model_name": "primary", "litellm_params": {"model": "openai/primary-model", "api_key": "sk-test"}}, + {"model_name": "fb1", "litellm_params": {"model": "openai/fb1-model", "api_key": "sk-test"}}, + {"model_name": "fb2", "litellm_params": {"model": "openai/fb2-model", "api_key": "sk-test"}}, + ], + num_retries=0, + **router_kwargs, + ) + + +def _mid_stream_failure(model: str): + import litellm + from litellm.exceptions import MidStreamFallbackError + + return MidStreamFallbackError( + message="stream dropped", + model=model, + llm_provider="openai", + original_exception=litellm.InternalServerError(message="stream dropped", llm_provider="openai", model=model), + is_pre_first_chunk=True, + ) + + +def _scripted_responses_stream(events: list, error: Exception | None = None): + from litellm.responses.streaming_iterator import BaseResponsesAPIStreamingIterator + + class _ScriptedStream(BaseResponsesAPIStreamingIterator): + def __init__(self) -> None: + self._events = list(events) + self._hidden_params: dict = {} + self.completed_response = None + + def __aiter__(self): + return self + + async def __anext__(self): + if self._events: + return self._events.pop(0) + if error is not None: + raise error + raise StopAsyncIteration + + async def aclose(self) -> None: + return None + + return _ScriptedStream() + + +def _three_tier_original(calls: list, primary_fails_pre_stream: bool): + import litellm + + completed_event = _make_completed_event(1, 1, 2) + + async def fake_original(**kwargs): + model = kwargs["model"] + calls.append(model) + if model == "openai/primary-model": + if primary_fails_pre_stream: + raise litellm.InternalServerError(message="primary down", llm_provider="openai", model=model) + return _scripted_responses_stream([], _mid_stream_failure(model)) + if model == "openai/fb1-model": + return _scripted_responses_stream([], _mid_stream_failure(model)) + return _scripted_responses_stream([completed_event]) + + return fake_original, completed_event + + +@pytest.mark.asyncio +async def test_aresponses_pre_stream_primary_failure_then_hop_stream_failure_reaches_second_entry(): + """Regression: fallbacks=[{"primary": ["fb1", "fb2"]}]. The primary fails before streaming, + fb1 is reached through the regular fallback chain and then fails mid-stream. Only the + primary's stream used to be wrapped, so fb1's mid-stream failure either re-raised or + re-tried fb1 itself; fb2 was unreachable.""" + router = _make_three_tier_router(fallbacks=[{"primary": ["fb1", "fb2"]}]) + calls: list = [] + fake_original, completed_event = _three_tier_original(calls, primary_fails_pre_stream=True) + + stream = await router._aresponses_with_streaming_fallbacks( + original_function=fake_original, model="primary", stream=True, input="hi" + ) + collected = [event async for event in stream] + + assert calls == ["openai/primary-model", "openai/fb1-model", "openai/fb2-model"] + assert collected == [completed_event] + + +@pytest.mark.asyncio +async def test_aresponses_two_consecutive_mid_stream_failures_reach_second_entry(): + """Regression: the primary and fb1 both fail mid-stream; fb2 must still be tried.""" + router = _make_three_tier_router(fallbacks=[{"primary": ["fb1", "fb2"]}]) + calls: list = [] + fake_original, completed_event = _three_tier_original(calls, primary_fails_pre_stream=False) + + stream = await router._aresponses_with_streaming_fallbacks( + original_function=fake_original, model="primary", stream=True, input="hi" + ) + collected = [event async for event in stream] + + assert calls == ["openai/primary-model", "openai/fb1-model", "openai/fb2-model"] + assert collected == [completed_event] + + +@pytest.mark.asyncio +async def test_aresponses_per_request_fallbacks_survive_into_hop_streams(): + """Regression: a request-level fallbacks list (key or team router_settings) is popped + before each attempt runs, so a hop's mid-stream re-entry used to see only the router's + own (empty) list and gave up after fb1.""" + router = _make_three_tier_router() + calls: list = [] + fake_original, completed_event = _three_tier_original(calls, primary_fails_pre_stream=False) + + stream = await router._aresponses_with_streaming_fallbacks( + original_function=fake_original, + model="primary", + stream=True, + input="hi", + fallbacks=[{"primary": ["fb1", "fb2"]}], + ) + collected = [event async for event in stream] + + assert calls == ["openai/primary-model", "openai/fb1-model", "openai/fb2-model"] + assert collected == [completed_event] + + @pytest.mark.asyncio async def test_aresponses_fallback_on_in_stream_error_event(): """A retriable in-stream error event (429) must trigger the router's mid-stream diff --git a/tests/test_litellm/test_router.py b/tests/test_litellm/test_router.py index e642520bdbd..1ff04a1722f 100644 --- a/tests/test_litellm/test_router.py +++ b/tests/test_litellm/test_router.py @@ -4234,7 +4234,7 @@ async def test_aresponses_streaming_iterator_fallback(): call_kwargs = mock_fallback_utils.call_args.kwargs fbk = call_kwargs["kwargs"] # Bound methods compare equal when they share the same instance + __func__. - assert fbk["original_function"] == router._ageneric_api_call_with_fallbacks_helper + assert fbk["original_function"] == router._ageneric_api_call_with_fallbacks_responses_attempt assert fbk["original_generic_function"] is litellm.aresponses assert call_kwargs["model_group"] == "anthropic/claude-sonnet-4-6" assert call_kwargs["disable_fallbacks"] is False @@ -13819,7 +13819,7 @@ async def test_aanthropic_messages_with_streaming_fallbacks_non_streaming_passth with patch.object( router, - "_ageneric_api_call_with_fallbacks", + "_ageneric_api_call_with_fallbacks_helper", new=AsyncMock(return_value=plain_response), ): out = await router._aanthropic_messages_with_streaming_fallbacks( @@ -13843,7 +13843,7 @@ async def test_aanthropic_messages_with_streaming_fallbacks_wraps_streaming_iter with ( patch.object( router, - "_ageneric_api_call_with_fallbacks", + "_ageneric_api_call_with_fallbacks_helper", new=AsyncMock(return_value=streaming_iter), ), patch.object( @@ -14128,7 +14128,7 @@ async def test_aanthropic_messages_with_streaming_fallbacks_deep_copies_nested_m ): with patch.object( router, - "_ageneric_api_call_with_fallbacks", + "_ageneric_api_call_with_fallbacks_helper", new=AsyncMock(side_effect=fake_original), ): await router._aanthropic_messages_with_streaming_fallbacks( @@ -14162,7 +14162,7 @@ async def test_aanthropic_messages_with_streaming_fallbacks_deep_copies_metadata ): with patch.object( router, - "_ageneric_api_call_with_fallbacks", + "_ageneric_api_call_with_fallbacks_helper", new=AsyncMock(side_effect=fake_original), ): await router._aanthropic_messages_with_streaming_fallbacks( @@ -14177,6 +14177,51 @@ async def test_aanthropic_messages_with_streaming_fallbacks_deep_copies_metadata assert "deployment" not in fallback_kwargs["metadata"] +@pytest.mark.asyncio +async def test_anthropic_messages_hop_stream_failure_reaches_second_fallback_entry(): + """Regression: fallbacks=[{"primary": ["fb1", "fb2"]}]. The primary fails before + streaming, fb1 is reached through the regular fallback chain and then sends an + error frame mid-stream. Only the primary's stream used to be wrapped, so the outer + wrapper re-tried fb1 with a fresh attempted set and forwarded fb1's error frame to + the client on an HTTP 200; fb2 was unreachable.""" + router = Router( + model_list=[ + {"model_name": "primary", "litellm_params": {"model": "anthropic/primary-model", "api_key": "sk-test"}}, + {"model_name": "fb1", "litellm_params": {"model": "anthropic/fb1-model", "api_key": "sk-test"}}, + {"model_name": "fb2", "litellm_params": {"model": "anthropic/fb2-model", "api_key": "sk-test"}}, + ], + num_retries=0, + fallbacks=[{"primary": ["fb1", "fb2"]}], + ) + calls: list = [] + + async def fake_original(**kwargs): + model = kwargs["model"] + calls.append(model) + if model == "anthropic/primary-model": + raise litellm.InternalServerError(message="primary down", llm_provider="anthropic", model=model) + if model == "anthropic/fb1-model": + return _AnthropicMessagesFakeByteStream( + [_anthropic_messages_message_start_chunk(), _anthropic_messages_overloaded_error_chunk()] + ) + return _AnthropicMessagesFakeByteStream( + [_anthropic_messages_message_start_chunk(), _anthropic_messages_content_chunk("from fb2")] + ) + + stream = await router._aanthropic_messages_with_streaming_fallbacks( + original_function=fake_original, + model="primary", + stream=True, + messages=[{"role": "user", "content": "hi"}], + max_tokens=10, + ) + body = b"".join([chunk async for chunk in stream]) + + assert calls == ["anthropic/primary-model", "anthropic/fb1-model", "anthropic/fb2-model"] + assert b"from fb2" in body + assert b"overloaded_error" not in body + + @pytest.mark.asyncio async def test_anthropic_messages_fallback_triggers_after_lifecycle_only_frame(): """Regression: Anthropic routinely sends a message_start lifecycle frame From 406365816802151bb9b9b7d99d6bd28a7ff79d22 Mon Sep 17 00:00:00 2001 From: Yujong Lee Date: Mon, 21 Sep 2026 21:22:02 +0000 Subject: [PATCH 38/44] build(rust): drop native extension size limit bump, deferred to #42300 Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .github/scripts/verify_linux_native_wheel.py | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/.github/scripts/verify_linux_native_wheel.py b/.github/scripts/verify_linux_native_wheel.py index 0adbc015ad0..4fb8f068eb0 100644 --- a/.github/scripts/verify_linux_native_wheel.py +++ b/.github/scripts/verify_linux_native_wheel.py @@ -205,7 +205,7 @@ def main( native_module: Final = load_native_module(native_path) native_module_loads: Final = native_module is not None panic_test_hook_absent: Final = native_module is not None and not hasattr(native_module, "_panic_for_test") - native_size_limit: Final = 30_000_000 + native_size_limit: Final = 25_000_000 native_size_within_limit: Final = native_member.file_size <= native_size_limit validations: Final = ( (f"Python tag is {EXPECTED_PYTHON_TAG}", python_tag == EXPECTED_PYTHON_TAG), @@ -222,7 +222,7 @@ def main( ("Python extension entry point is present", extension_entry_point_present), ("Native module loads", native_module_loads), ("Production module omits the panic test hook", panic_test_hook_absent), - ("Native extension does not exceed 30 MB", native_size_within_limit), + ("Native extension does not exceed 25 MB", native_size_within_limit), ("Wheel contents are valid", not unexpected_members), ) @@ -267,7 +267,7 @@ def main( ), ( not native_size_within_limit, - f"native extension exceeds 30 MB: {native_member.file_size / 1_000_000:.2f} MB", + f"native extension exceeds 20 MB: {native_member.file_size / 1_000_000:.2f} MB", ), (bool(unexpected_members), f"wheel contains unexpected build artifacts: {', '.join(unexpected_members)}"), ) From 51eaad657baf790379ed7d047d45c08364f2de40 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Mon, 21 Sep 2026 14:42:23 -0700 Subject: [PATCH 39/44] fix(proxy): look the entity spend table up lazily so only the selected batcher table is touched --- litellm/proxy/db/db_spend_update_writer.py | 20 +++++++++++--------- 1 file changed, 11 insertions(+), 9 deletions(-) diff --git a/litellm/proxy/db/db_spend_update_writer.py b/litellm/proxy/db/db_spend_update_writer.py index 5381071172f..d9c8b271646 100644 --- a/litellm/proxy/db/db_spend_update_writer.py +++ b/litellm/proxy/db/db_spend_update_writer.py @@ -141,16 +141,18 @@ _EntitySpendTable: TypeAlias = Literal[ ] +_ENTITY_SPEND_TABLES: Final[Mapping[_EntitySpendTable, Callable[[_SpendBatch], BatchTable]]] = MappingProxyType( + { + "litellm_tagtable": lambda batcher: batcher.litellm_tagtable, + "litellm_agentstable": lambda batcher: batcher.litellm_agentstable, + "litellm_modelaccessgroupbudgettable": lambda batcher: batcher.litellm_modelaccessgroupbudgettable, + "litellm_projecttable": lambda batcher: batcher.litellm_projecttable, + } +) + + def _entity_spend_table(batcher: _SpendBatch, table_accessor: _EntitySpendTable) -> BatchTable: - tables: Final[Mapping[_EntitySpendTable, BatchTable]] = MappingProxyType( - { - "litellm_tagtable": batcher.litellm_tagtable, - "litellm_agentstable": batcher.litellm_agentstable, - "litellm_modelaccessgroupbudgettable": batcher.litellm_modelaccessgroupbudgettable, - "litellm_projecttable": batcher.litellm_projecttable, - } - ) - return tables[table_accessor] + return _ENTITY_SPEND_TABLES[table_accessor](batcher) class _SpendBatchManager(Protocol): From b0305d0a31a4eb4dd1ff6204ef2e1e5f84db07c4 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Mon, 21 Sep 2026 14:48:42 -0700 Subject: [PATCH 40/44] test(router): call both mid-stream fallback attempt functions directly The router coverage gate wants every router.py function reached by name from a router test. The two per-endpoint attempt functions were only reached through their callers, so each now has a direct test proving the per-request controls carrier never reaches the provider call and every hop's stream comes back wrapped. --- ...st_router_aresponses_streaming_fallback.py | 38 ++++++++++++++++ tests/test_litellm/test_router.py | 44 +++++++++++++++++++ 2 files changed, 82 insertions(+) diff --git a/tests/router_unit_tests/test_router_aresponses_streaming_fallback.py b/tests/router_unit_tests/test_router_aresponses_streaming_fallback.py index 18b10e9c8e5..5370089eef5 100644 --- a/tests/router_unit_tests/test_router_aresponses_streaming_fallback.py +++ b/tests/router_unit_tests/test_router_aresponses_streaming_fallback.py @@ -423,6 +423,44 @@ async def test_aresponses_per_request_fallbacks_survive_into_hop_streams(): assert collected == [completed_event] +@pytest.mark.asyncio +async def test_aresponses_attempt_strips_the_controls_carrier_and_wraps_every_hop_stream(): + """Each attempt of the chain, not only the primary's, comes back wrapped for mid-stream + failover, and the per-request controls carrier rides into the wrapper's re-entry kwargs + without ever reaching the provider call.""" + from types import MappingProxyType + + from litellm.router_utils.fallback_event_handlers import ( + MID_STREAM_FALLBACK_CONTROLS_KEY, + MidStreamFallbackControls, + ) + + router = _make_three_tier_router() + completed_event = _make_completed_event(1, 1, 2) + hop_stream = _scripted_responses_stream([completed_event]) + seen: dict = {} + + async def fake_original(**kwargs): + seen.update(kwargs) + return hop_stream + + controls = MidStreamFallbackControls(MappingProxyType({"fallbacks": [{"primary": ["fb1", "fb2"]}]})) + stream = await router._ageneric_api_call_with_fallbacks_responses_attempt( + model="fb1", + original_generic_function=fake_original, + stream=True, + input="hi", + **{MID_STREAM_FALLBACK_CONTROLS_KEY: controls}, + ) + collected = [event async for event in stream] + + assert seen["model"] == "openai/fb1-model" + assert MID_STREAM_FALLBACK_CONTROLS_KEY not in seen + assert "fallbacks" not in seen + assert stream is not hop_stream + assert collected == [completed_event] + + @pytest.mark.asyncio async def test_aresponses_fallback_on_in_stream_error_event(): """A retriable in-stream error event (429) must trigger the router's mid-stream diff --git a/tests/test_litellm/test_router.py b/tests/test_litellm/test_router.py index 1ff04a1722f..9d2e37b4fb5 100644 --- a/tests/test_litellm/test_router.py +++ b/tests/test_litellm/test_router.py @@ -14222,6 +14222,50 @@ async def test_anthropic_messages_hop_stream_failure_reaches_second_fallback_ent assert b"overloaded_error" not in body +@pytest.mark.asyncio +async def test_anthropic_messages_attempt_strips_the_controls_carrier_and_wraps_every_hop_stream(): + """Each attempt of the chain, not only the primary's, comes back wrapped for mid-stream + failover, and the per-request controls carrier never reaches the provider call.""" + from types import MappingProxyType + + from litellm.router_utils.fallback_event_handlers import ( + MID_STREAM_FALLBACK_CONTROLS_KEY, + MidStreamFallbackControls, + ) + + router = Router( + model_list=[ + {"model_name": "fb1", "litellm_params": {"model": "anthropic/fb1-model", "api_key": "sk-test"}}, + ], + num_retries=0, + ) + hop_stream = _AnthropicMessagesFakeByteStream( + [_anthropic_messages_message_start_chunk(), _anthropic_messages_content_chunk("from fb1")] + ) + seen: dict = {} + + async def fake_original(**kwargs): + seen.update(kwargs) + return hop_stream + + controls = MidStreamFallbackControls(MappingProxyType({"fallbacks": [{"primary": ["fb1", "fb2"]}]})) + stream = await router._ageneric_api_call_with_fallbacks_anthropic_messages_attempt( + model="fb1", + original_generic_function=fake_original, + stream=True, + messages=[{"role": "user", "content": "hi"}], + max_tokens=10, + **{MID_STREAM_FALLBACK_CONTROLS_KEY: controls}, + ) + body = b"".join([chunk async for chunk in stream]) + + assert seen["model"] == "anthropic/fb1-model" + assert MID_STREAM_FALLBACK_CONTROLS_KEY not in seen + assert "fallbacks" not in seen + assert stream is not hop_stream + assert b"from fb1" in body + + @pytest.mark.asyncio async def test_anthropic_messages_fallback_triggers_after_lifecycle_only_frame(): """Regression: Anthropic routinely sends a message_start lifecycle frame From e7557ada57ad649fc0b0063aa2d89e1f7186f979 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Mon, 21 Sep 2026 14:56:02 -0700 Subject: [PATCH 41/44] fix(responses): list executed MCP calls as completed mcp_call items in the final output --- .../responses/mcp/mcp_streaming_iterator.py | 23 ++++++++- .../mcp/test_mcp_streaming_iterator.py | 48 +++++++++++++++---- 2 files changed, 60 insertions(+), 11 deletions(-) diff --git a/litellm/responses/mcp/mcp_streaming_iterator.py b/litellm/responses/mcp/mcp_streaming_iterator.py index fb05e7bff99..101000b14a9 100644 --- a/litellm/responses/mcp/mcp_streaming_iterator.py +++ b/litellm/responses/mcp/mcp_streaming_iterator.py @@ -39,6 +39,19 @@ def _output_items(response: ResponsesAPIResponse) -> Sequence[object]: return tuple(cast("Sequence[object]", response.output)) # cast-ok: items are only carried, never inspected +def _function_call_id(item: object) -> str | None: + """The call id of a function_call item, None for every other item kind.""" + item_type: Final[object] = item.get("type") if isinstance(item, dict) else getattr(item, "type", None) + if item_type != "function_call": + return None + call_id: Final[object] = ( + item.get("call_id") or item.get("id") + if isinstance(item, dict) + else getattr(item, "call_id", None) or getattr(item, "id", None) + ) + return call_id if isinstance(call_id, str) else None + + def _set_event_field(event: ResponsesAPIStreamingResponse, name: str, value: object) -> None: """Events are pydantic models with extra fields allowed, so any event type can carry the field.""" setattr(event, name, value) @@ -588,9 +601,14 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator): return max(len(_output_items(response)), self._round_max_output_index + 1) def _absorb_round(self, response: ResponsesAPIResponse) -> None: - """Bank a finished round's items so the final response.completed can list them.""" + """Bank a finished round's items, each function_call the gateway answered replaced by its mcp_call.""" width: Final = self._round_output_width(response) - self._composed_output.extend(_output_items(response)) + answered_call_ids: Final = frozenset( + call_id for result in self.tool_results if (call_id := result.get("tool_call_id")) is not None + ) + self._composed_output.extend( + item for item in _output_items(response) if _function_call_id(item) not in answered_call_ids + ) self._composed_output.extend(self._pending_mcp_call_items) self._output_index_offset += width + len(self._pending_mcp_call_items) self._pending_mcp_call_items.clear() @@ -842,6 +860,7 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator): **{ # mutable-ok: consumed once by the model constructor "id": item_id, "type": "mcp_call", + "status": "completed", "approval_request_id": f"mcpr_{uuid.uuid4().hex[:8]}", "arguments": tool_arguments, "error": None, diff --git a/tests/test_litellm/responses/mcp/test_mcp_streaming_iterator.py b/tests/test_litellm/responses/mcp/test_mcp_streaming_iterator.py index 1587557f3d2..1e869745dc3 100644 --- a/tests/test_litellm/responses/mcp/test_mcp_streaming_iterator.py +++ b/tests/test_litellm/responses/mcp/test_mcp_streaming_iterator.py @@ -143,17 +143,12 @@ async def test_second_round_tool_call_is_executed_and_reaches_final_text(monkeyp # The stream reached round 3 and produced the final text response instead # of stopping after round 1 or round 2. The client sees one lifecycle whose - # final output lists every round's items in order. + # final output lists every round's items in order, each executed call as + # the gateway's mcp_call rather than the function_call the model emitted. completed_chunks = [c for c in chunks if getattr(c, "type", None) == ResponsesAPIStreamEvents.RESPONSE_COMPLETED] assert len(completed_chunks) == 1 final_output = completed_chunks[-1].response.output - assert [_item_type(item) for item in final_output] == [ - "function_call", - "mcp_call", - "function_call", - "mcp_call", - "message", - ] + assert [_item_type(item) for item in final_output] == ["mcp_call", "mcp_call", "message"] assert final_output[-1]["content"][0]["text"] == "Here's what I found after retrying." @@ -422,7 +417,7 @@ async def test_auto_execute_rounds_share_one_public_lifecycle(monkeypatch): # The single completed event lists every round's items and keeps the final round's id for continuation. completed = chunks[-1] assert completed.response.id == "resp-final" - assert [_item_type(item) for item in completed.response.output] == ["function_call", "mcp_call", "message"] + assert [_item_type(item) for item in completed.response.output] == ["mcp_call", "message"] assert completed.response.output[-1]["content"][0]["text"] == "Alpha." # The proxy serializes every chunk; the merged output must still be a valid response. assert '"type":"mcp_call"' in completed.response.model_dump_json(exclude_none=True, exclude_unset=True) @@ -433,6 +428,41 @@ async def test_auto_execute_rounds_share_one_public_lifecycle(monkeypatch): assert len(set(sequence_numbers)) == len(sequence_numbers) +@pytest.mark.asyncio +async def test_final_output_lists_executed_call_as_completed_mcp_call(monkeypatch): + """ + A function_call the gateway executed must not reach the final output: an + agent framework reading it (the OpenAI Agents SDK) tries to run a tool the + caller never declared and aborts the run. The final output lists the + gateway's completed mcp_call in its place, next to the round's other items. + """ + _mock_mcp_environment(monkeypatch) + + follow_up = _FakeAsyncStream(_lifecycle_round("resp-final", _text_message("Alpha."))) + monkeypatch.setattr(responses_main_module, "aresponses", AsyncMock(side_effect=[follow_up])) + + reasoning = {"type": "reasoning", "id": "rs_1", "summary": []} + iterator = _make_iterator( + [ + _created_chunk("resp-interim"), + _completed_chunk([reasoning, _function_call("call_1", "read_wiki_contents")], response_id="resp-interim"), + ] + ) + chunks = [chunk async for chunk in iterator] + + final_output = chunks[-1].response.output + assert [_item_type(item) for item in final_output] == ["reasoning", "mcp_call", "message"] + executed_call = final_output[1] + assert executed_call["status"] == "completed" + assert executed_call["name"] == "read_wiki_contents" + assert executed_call["arguments"] == "{}" + + done_mcp_items = [ + c.item for c in chunks if c.type == ResponsesAPIStreamEvents.OUTPUT_ITEM_DONE and _item_type(c.item) == "mcp_call" + ] + assert [item.status for item in done_mcp_items] == ["completed"] + + @pytest.mark.asyncio async def test_stream_without_auto_execute_is_forwarded_unchanged(monkeypatch): """With approval required there is one round, and it passes through untouched.""" From 5c6854510400ca53a8718401adf0ca911f65e62a Mon Sep 17 00:00:00 2001 From: Yuneng Jiang Date: Mon, 21 Sep 2026 15:00:54 -0700 Subject: [PATCH 42/44] test: fix stale budget-status and bad-database-url assertions MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Two CI checks were asserting behaviour the proxy no longer has. Neither was catching anything; both now fail for the right reason. budget_exceeded (tests/otel_tests/test_e2e_budgeting.py) bf804f5188 made 422 the default for budget refusals and added budget_exceeded_status_code to restore 429 for callers that need it. The e2e budget tests still asserted 429, so all six have been failing on a status change that was deliberate. Assert 422, the documented default, rather than reading litellm.budget_exceeded_status_code back — a test that asks the code what it does would have passed straight through this change and through the next one. The helpers also caught bare Exception, so a connection error reached `e.body` and surfaced as an AttributeError instead of a failed assertion. Narrow both to openai.APIStatusError, which is what a refusal actually raises (UnprocessableEntityError for 422, RateLimitError for 429), and let anything else propagate as itself. test_bad_database_url (.circleci/config.yml) The check required "Database setup failed after multiple retries", which only the v1 resolver emits, OR uvicorn's "Application startup failed. Exiting.". With v2 the default, the first branch is dead and the whole assertion rests on an incidental uvicorn line that HEAD's run did not emit at all. Assert the behaviour instead of the wording: the container exits non-zero, the log names the unreachable server (P1001), it never reaches "Application startup complete", and it is not left running. The exit code was previously discarded by `|| true`, so the one thing the job most needed to check was never checked. Verified by running the bad-DATABASE_URL container: exit 3, P1001 present, no startup-complete line, container stopped — the new check passes and the old one passed only by the uvicorn line's accident. --- .circleci/config.yml | 23 ++++++++++++----------- tests/otel_tests/test_e2e_budgeting.py | 13 +++++++------ 2 files changed, 19 insertions(+), 17 deletions(-) diff --git a/.circleci/config.yml b/.circleci/config.yml index cc9aa7fe1c4..c69ed6a3510 100644 --- a/.circleci/config.yml +++ b/.circleci/config.yml @@ -3033,28 +3033,29 @@ jobs: - run: name: Run Docker container with bad DATABASE_URL command: | + set +e docker run --name my-app \ -p 4000:4000 \ -e LITELLM_DANGEROUSLY_PERMIT_WEAK_OR_UNSET_MASTER_KEY=true \ -e DEFAULT_NUM_WORKERS_LITELLM_PROXY=1 \ -e DATABASE_URL="postgresql://wrong:wrong@wrong:5432/wrong" \ myapp:latest \ - --port 4000 > docker_output.log 2>&1 || true + --port 4000 > docker_output.log 2>&1 + echo "$?" > docker_exit_code + set -e - run: name: Display Docker logs command: cat docker_output.log - run: - name: Check for expected error + name: Proxy must refuse to serve on an unreachable database command: | - if grep -q "Error: P1001: Can't reach database server at" docker_output.log && \ - (grep -q "Database setup failed after multiple retries" docker_output.log || \ - grep -q "ERROR: Application startup failed. Exiting." docker_output.log); then - echo "Expected error found. Test passed." - else - echo "Expected error not found. Test failed." - cat docker_output.log - exit 1 - fi + fail() { echo "FAILED: $1"; cat docker_output.log; exit 1; } + exit_code="$(cat docker_exit_code)" + [ "$exit_code" -ne 0 ] || fail "proxy exited 0 with an unreachable database" + grep -q "P1001" docker_output.log || fail "log does not name the unreachable database server" + ! grep -q "Application startup complete" docker_output.log || fail "proxy reached serving state" + ! docker exec my-app true 2>/dev/null || fail "container is still running" + echo "Proxy refused to serve (exit $exit_code) and never reached startup. Test passed." provider_replay_harness: docker: diff --git a/tests/otel_tests/test_e2e_budgeting.py b/tests/otel_tests/test_e2e_budgeting.py index ca5058818e4..ae8f0ddc3ec 100644 --- a/tests/otel_tests/test_e2e_budgeting.py +++ b/tests/otel_tests/test_e2e_budgeting.py @@ -5,6 +5,7 @@ import uuid from typing import Any, Optional import aiohttp +import openai import pytest from httpx import AsyncClient @@ -23,7 +24,7 @@ async def make_calls_until_budget_exceeded(session, key: str, call_function, **k call_count += 1 await asyncio.sleep(0.1) # allow spend tracking to catch up pytest.fail(f"Budget was not exceeded after {MAX_CALLS} calls") - except Exception as e: + except openai.APIStatusError as e: print("vars: ", vars(e)) print("e.body: ", e.body) @@ -32,8 +33,8 @@ async def make_calls_until_budget_exceeded(session, key: str, call_function, **k # Check error structure and values that should be consistent assert ( - error_dict["code"] == "429" - ), f"Expected error code 429, got: {error_dict['code']}" + error_dict["code"] == "422" + ), f"Expected error code 422, got: {error_dict['code']}" assert ( error_dict["type"] == "budget_exceeded" ), f"Expected error type budget_exceeded, got: {error_dict['type']}" @@ -506,9 +507,9 @@ async def make_calls_until_team_budget_exceeded_cli_sso( call_count += 1 await asyncio.sleep(0.1) pytest.fail(f"Budget was not exceeded after {MAX_CALLS} calls") - except Exception as e: + except openai.APIStatusError as e: error_dict = e.body - assert error_dict["code"] == "429" + assert error_dict["code"] == "422" assert error_dict["type"] == "budget_exceeded" message = error_dict["message"] assert "Budget has been exceeded!" in message @@ -556,7 +557,7 @@ async def test_team_budget_enforcement_cli_sso_token(): 1. Create team with a tiny max_budget and a user on that team 2. Obtain a CLI SSO JWT (HTTP poll flow when Redis is shared, else mint) 3. Make chat completion calls until the team budget is exceeded - 4. Verify HTTP 429 budget_exceeded names the team + 4. Verify HTTP 422 budget_exceeded names the team """ user_id = f"cli-budget-user-{uuid.uuid4().hex[:8]}" user_email = f"{user_id}@example.com" From 0eea01427d99f3d17844e1688be3b4c7be4891b4 Mon Sep 17 00:00:00 2001 From: "berriai-litellm-provider-info-sync[bot]" <328147090+berriai-litellm-provider-info-sync[bot]@users.noreply.github.com> Date: Mon, 21 Sep 2026 22:01:13 +0000 Subject: [PATCH 43/44] chore(prices): sync OpenRouter prices: 3 models openrouter/~deepseek/deepseek-v4-flash-latest: output_cost_per_token openrouter/deepseek/deepseek-v4-flash-0731: output_cost_per_token openrouter/deepseek/deepseek-v4-pro: input_cost_per_token, output_cost_per_token, cache_read_input_token_cost --- litellm/model_prices_and_context_window_backup.json | 10 +++++----- model_prices_and_context_window.json | 10 +++++----- 2 files changed, 10 insertions(+), 10 deletions(-) diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index 4288a87491c..5aea9cffcec 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -43011,21 +43011,21 @@ "supports_web_search": false }, "openrouter/deepseek/deepseek-v4-pro": { - "input_cost_per_token": 9.00798e-07, + "input_cost_per_token": 8.95578e-07, "input_cost_per_token_cache_hit": 4.4e-08, "litellm_provider": "openrouter", "max_input_tokens": 1048576, "max_output_tokens": 384000, "max_tokens": 384000, "mode": "chat", - "output_cost_per_token": 1.801596e-06, + "output_cost_per_token": 1.791156e-06, "source": "https://openrouter.ai/api/v1/models", "supports_function_calling": true, "supports_prompt_caching": true, "supports_reasoning": true, "supports_response_schema": true, "supports_tool_choice": true, - "cache_read_input_token_cost": 7.50665e-08, + "cache_read_input_token_cost": 7.46315e-08, "supports_audio_input": false, "supports_pdf_input": false, "supports_vision": false, @@ -68346,7 +68346,7 @@ }, "openrouter/deepseek/deepseek-v4-flash-0731": { "input_cost_per_token": 4e-08, - "output_cost_per_token": 3.2e-07, + "output_cost_per_token": 6.4e-07, "cache_read_input_token_cost": 1.6e-08, "litellm_provider": "openrouter", "max_input_tokens": 1310720, @@ -73001,7 +73001,7 @@ "max_output_tokens": 943718, "max_tokens": 943718, "mode": "chat", - "output_cost_per_token": 3.2e-07, + "output_cost_per_token": 6.4e-07, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index 4288a87491c..5aea9cffcec 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -43011,21 +43011,21 @@ "supports_web_search": false }, "openrouter/deepseek/deepseek-v4-pro": { - "input_cost_per_token": 9.00798e-07, + "input_cost_per_token": 8.95578e-07, "input_cost_per_token_cache_hit": 4.4e-08, "litellm_provider": "openrouter", "max_input_tokens": 1048576, "max_output_tokens": 384000, "max_tokens": 384000, "mode": "chat", - "output_cost_per_token": 1.801596e-06, + "output_cost_per_token": 1.791156e-06, "source": "https://openrouter.ai/api/v1/models", "supports_function_calling": true, "supports_prompt_caching": true, "supports_reasoning": true, "supports_response_schema": true, "supports_tool_choice": true, - "cache_read_input_token_cost": 7.50665e-08, + "cache_read_input_token_cost": 7.46315e-08, "supports_audio_input": false, "supports_pdf_input": false, "supports_vision": false, @@ -68346,7 +68346,7 @@ }, "openrouter/deepseek/deepseek-v4-flash-0731": { "input_cost_per_token": 4e-08, - "output_cost_per_token": 3.2e-07, + "output_cost_per_token": 6.4e-07, "cache_read_input_token_cost": 1.6e-08, "litellm_provider": "openrouter", "max_input_tokens": 1310720, @@ -73001,7 +73001,7 @@ "max_output_tokens": 943718, "max_tokens": 943718, "mode": "chat", - "output_cost_per_token": 3.2e-07, + "output_cost_per_token": 6.4e-07, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, From b305928422e511dd11e288c4b4f19ae743dd6f35 Mon Sep 17 00:00:00 2001 From: "berriai-litellm-provider-info-sync[bot]" <328147090+berriai-litellm-provider-info-sync[bot]@users.noreply.github.com> Date: Mon, 21 Sep 2026 22:01:14 +0000 Subject: [PATCH 44/44] chore(prices): sync AWS Bedrock prices: 6 models [enrichment failed: AWS Bedrock, 4 held] minimax.minimax-m2: max_tokens, supports_vision, max_input_tokens, max_output_tokens, supports_audio_input, supports_response_schema, supports_function_calling minimax.minimax-m2.1: max_tokens, supports_vision, max_output_tokens, supports_audio_input, supports_response_schema minimax.minimax-m2.5: max_tokens, supports_vision, max_input_tokens, max_output_tokens, supports_audio_input, supports_response_schema moonshot.kimi-k2-thinking: max_tokens, supports_vision, max_input_tokens, max_output_tokens, supports_audio_input, supports_response_schema, supports_function_calling openai.gpt-oss-safeguard-120b: max_tokens, supports_vision, max_output_tokens, supports_audio_input, supports_response_schema, supports_function_calling openai.gpt-oss-safeguard-20b: max_tokens, supports_vision, max_output_tokens, supports_audio_input, supports_response_schema, supports_function_calling --- ...odel_prices_and_context_window_backup.json | 68 +++++++++++++------ model_prices_and_context_window.json | 68 +++++++++++++------ 2 files changed, 94 insertions(+), 42 deletions(-) diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index 4288a87491c..5526867ce74 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -38214,39 +38214,50 @@ "minimax.minimax-m2": { "input_cost_per_token": 3e-07, "litellm_provider": "bedrock_converse", - "max_input_tokens": 128000, - "max_output_tokens": 8192, - "max_tokens": 8192, + "max_input_tokens": 1000000, + "max_output_tokens": 8000, + "max_tokens": 8000, "mode": "chat", "output_cost_per_token": 1.2e-06, + "source": "https://aws.amazon.com/bedrock/pricing/", + "supports_audio_input": false, + "supports_function_calling": true, "supports_system_messages": true, - "supports_native_structured_output": true + "supports_native_structured_output": true, + "supports_response_schema": true, + "supports_vision": false }, "minimax.minimax-m2.1": { "input_cost_per_token": 3e-07, "litellm_provider": "bedrock_converse", "max_input_tokens": 196000, - "max_output_tokens": 8192, - "max_tokens": 8192, + "max_output_tokens": 8000, + "max_tokens": 8000, "mode": "chat", "output_cost_per_token": 1.2e-06, "supports_function_calling": true, "supports_system_messages": true, "supports_tool_choice": true, - "source": "https://aws.amazon.com/bedrock/pricing/" + "source": "https://aws.amazon.com/bedrock/pricing/", + "supports_audio_input": false, + "supports_response_schema": true, + "supports_vision": false }, "minimax.minimax-m2.5": { "input_cost_per_token": 3e-07, "output_cost_per_token": 1.2e-06, "litellm_provider": "bedrock_converse", - "max_input_tokens": 1000000, - "max_output_tokens": 8192, - "max_tokens": 8192, + "max_input_tokens": 196000, + "max_output_tokens": 8000, + "max_tokens": 8000, "mode": "chat", "supports_function_calling": true, "supports_system_messages": true, "supports_tool_choice": true, - "source": "https://aws.amazon.com/bedrock/pricing/" + "source": "https://aws.amazon.com/bedrock/pricing/", + "supports_audio_input": false, + "supports_response_schema": true, + "supports_vision": false }, "minimax/speech-02-hd": { "input_cost_per_character": 0.0001, @@ -39705,14 +39716,19 @@ "moonshot.kimi-k2-thinking": { "input_cost_per_token": 6e-07, "litellm_provider": "bedrock_converse", - "max_input_tokens": 128000, - "max_output_tokens": 8192, - "max_tokens": 8192, + "max_input_tokens": 256000, + "max_output_tokens": 16000, + "max_tokens": 16000, "mode": "chat", "output_cost_per_token": 2.5e-06, + "source": "https://aws.amazon.com/bedrock/pricing/", + "supports_audio_input": false, + "supports_function_calling": true, "supports_reasoning": true, "supports_system_messages": true, - "supports_native_structured_output": true + "supports_native_structured_output": true, + "supports_response_schema": true, + "supports_vision": false }, "moonshotai.kimi-k2.5": { "input_cost_per_token": 6e-07, @@ -42492,21 +42508,31 @@ "input_cost_per_token": 1.5e-07, "litellm_provider": "bedrock_converse", "max_input_tokens": 128000, - "max_output_tokens": 8192, - "max_tokens": 8192, + "max_output_tokens": 16000, + "max_tokens": 16000, "mode": "chat", "output_cost_per_token": 6e-07, - "supports_system_messages": true + "source": "https://aws.amazon.com/bedrock/pricing/", + "supports_audio_input": false, + "supports_function_calling": true, + "supports_response_schema": true, + "supports_system_messages": true, + "supports_vision": false }, "openai.gpt-oss-safeguard-20b": { "input_cost_per_token": 7e-08, "litellm_provider": "bedrock_converse", "max_input_tokens": 128000, - "max_output_tokens": 8192, - "max_tokens": 8192, + "max_output_tokens": 16000, + "max_tokens": 16000, "mode": "chat", "output_cost_per_token": 2e-07, - "supports_system_messages": true + "source": "https://aws.amazon.com/bedrock/pricing/", + "supports_audio_input": false, + "supports_function_calling": true, + "supports_response_schema": true, + "supports_system_messages": true, + "supports_vision": false }, "openrouter/anthropic/claude-3-haiku": { "cache_creation_input_token_cost": 3e-07, diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index 4288a87491c..5526867ce74 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -38214,39 +38214,50 @@ "minimax.minimax-m2": { "input_cost_per_token": 3e-07, "litellm_provider": "bedrock_converse", - "max_input_tokens": 128000, - "max_output_tokens": 8192, - "max_tokens": 8192, + "max_input_tokens": 1000000, + "max_output_tokens": 8000, + "max_tokens": 8000, "mode": "chat", "output_cost_per_token": 1.2e-06, + "source": "https://aws.amazon.com/bedrock/pricing/", + "supports_audio_input": false, + "supports_function_calling": true, "supports_system_messages": true, - "supports_native_structured_output": true + "supports_native_structured_output": true, + "supports_response_schema": true, + "supports_vision": false }, "minimax.minimax-m2.1": { "input_cost_per_token": 3e-07, "litellm_provider": "bedrock_converse", "max_input_tokens": 196000, - "max_output_tokens": 8192, - "max_tokens": 8192, + "max_output_tokens": 8000, + "max_tokens": 8000, "mode": "chat", "output_cost_per_token": 1.2e-06, "supports_function_calling": true, "supports_system_messages": true, "supports_tool_choice": true, - "source": "https://aws.amazon.com/bedrock/pricing/" + "source": "https://aws.amazon.com/bedrock/pricing/", + "supports_audio_input": false, + "supports_response_schema": true, + "supports_vision": false }, "minimax.minimax-m2.5": { "input_cost_per_token": 3e-07, "output_cost_per_token": 1.2e-06, "litellm_provider": "bedrock_converse", - "max_input_tokens": 1000000, - "max_output_tokens": 8192, - "max_tokens": 8192, + "max_input_tokens": 196000, + "max_output_tokens": 8000, + "max_tokens": 8000, "mode": "chat", "supports_function_calling": true, "supports_system_messages": true, "supports_tool_choice": true, - "source": "https://aws.amazon.com/bedrock/pricing/" + "source": "https://aws.amazon.com/bedrock/pricing/", + "supports_audio_input": false, + "supports_response_schema": true, + "supports_vision": false }, "minimax/speech-02-hd": { "input_cost_per_character": 0.0001, @@ -39705,14 +39716,19 @@ "moonshot.kimi-k2-thinking": { "input_cost_per_token": 6e-07, "litellm_provider": "bedrock_converse", - "max_input_tokens": 128000, - "max_output_tokens": 8192, - "max_tokens": 8192, + "max_input_tokens": 256000, + "max_output_tokens": 16000, + "max_tokens": 16000, "mode": "chat", "output_cost_per_token": 2.5e-06, + "source": "https://aws.amazon.com/bedrock/pricing/", + "supports_audio_input": false, + "supports_function_calling": true, "supports_reasoning": true, "supports_system_messages": true, - "supports_native_structured_output": true + "supports_native_structured_output": true, + "supports_response_schema": true, + "supports_vision": false }, "moonshotai.kimi-k2.5": { "input_cost_per_token": 6e-07, @@ -42492,21 +42508,31 @@ "input_cost_per_token": 1.5e-07, "litellm_provider": "bedrock_converse", "max_input_tokens": 128000, - "max_output_tokens": 8192, - "max_tokens": 8192, + "max_output_tokens": 16000, + "max_tokens": 16000, "mode": "chat", "output_cost_per_token": 6e-07, - "supports_system_messages": true + "source": "https://aws.amazon.com/bedrock/pricing/", + "supports_audio_input": false, + "supports_function_calling": true, + "supports_response_schema": true, + "supports_system_messages": true, + "supports_vision": false }, "openai.gpt-oss-safeguard-20b": { "input_cost_per_token": 7e-08, "litellm_provider": "bedrock_converse", "max_input_tokens": 128000, - "max_output_tokens": 8192, - "max_tokens": 8192, + "max_output_tokens": 16000, + "max_tokens": 16000, "mode": "chat", "output_cost_per_token": 2e-07, - "supports_system_messages": true + "source": "https://aws.amazon.com/bedrock/pricing/", + "supports_audio_input": false, + "supports_function_calling": true, + "supports_response_schema": true, + "supports_system_messages": true, + "supports_vision": false }, "openrouter/anthropic/claude-3-haiku": { "cache_creation_input_token_cost": 3e-07,