From 40f02b2eb520f2b8178cd1a3c56ca4c3549639d7 Mon Sep 17 00:00:00 2001
From: Tin Chi Lo
Date: Tue, 14 Jul 2026 00:21:22 -0700
Subject: [PATCH 001/296] refactor(mcp): consolidate exception-tree walkers
into one shared faults traversal
---
.../mcp_server/faults/__init__.py | 2 +
.../mcp_server/faults/traversal.py | 35 ++++++++++++++
.../mcp_server/mcp_server_manager.py | 42 ++++------------
.../mcp_server/semantic_tool_filter.py | 22 ++++-----
ruff-strict-budget.json | 4 +-
.../mcp_server/faults/test_traversal.py | 48 +++++++++++++++++++
.../test_mcp_oauth_passthrough_tools.py | 30 ++++++++++++
.../mcp_server/test_semantic_tool_filter.py | 28 +++++++++++
type-discipline-budget.json | 2 +-
9 files changed, 164 insertions(+), 49 deletions(-)
create mode 100644 litellm/proxy/_experimental/mcp_server/faults/traversal.py
create mode 100644 tests/test_litellm/proxy/_experimental/mcp_server/faults/test_traversal.py
diff --git a/litellm/proxy/_experimental/mcp_server/faults/__init__.py b/litellm/proxy/_experimental/mcp_server/faults/__init__.py
index da078f0e242..1b9ee77d795 100644
--- a/litellm/proxy/_experimental/mcp_server/faults/__init__.py
+++ b/litellm/proxy/_experimental/mcp_server/faults/__init__.py
@@ -15,6 +15,7 @@ from litellm.proxy._experimental.mcp_server.faults.render_oauth import (
dcr_fault_detail,
render_token_fault,
)
+from litellm.proxy._experimental.mcp_server.faults.traversal import iter_exception_tree
from litellm.proxy._experimental.mcp_server.faults.types import (
CallerRejected,
CredentialSource,
@@ -34,5 +35,6 @@ __all__ = [
"classify_upstream_dcr_rejection",
"classify_upstream_token_rejection",
"dcr_fault_detail",
+ "iter_exception_tree",
"render_token_fault",
]
diff --git a/litellm/proxy/_experimental/mcp_server/faults/traversal.py b/litellm/proxy/_experimental/mcp_server/faults/traversal.py
new file mode 100644
index 00000000000..78e94e22e70
--- /dev/null
+++ b/litellm/proxy/_experimental/mcp_server/faults/traversal.py
@@ -0,0 +1,35 @@
+"""Shared exception-tree traversal for fault classification.
+
+Failures cross the MCP SDK's anyio task groups wrapped in ``ExceptionGroup``s and chained through
+``raise ... from`` causes, so every classifier that needs an exception buried in the tree (an
+upstream ``httpx.Response``, a context-window overflow) has to walk the same shapes. One traversal
+with one deliberate order keeps blame assignment consistent across classifiers: explicit links are
+searched before incidental ones, so an exception raised while handling the real failure can never
+shadow the failure itself.
+"""
+
+from __future__ import annotations
+
+from collections.abc import Iterator
+
+
+def iter_exception_tree(exc: BaseException) -> Iterator[BaseException]:
+ """Yield ``exc`` and every exception reachable from it, explicit links first: each node's
+ ``raise ... from`` cause subtree, then ``ExceptionGroup`` members in raise order, then the
+ incidental ``__context__`` chain last. Cycle-safe via identity tracking, and iterative so a
+ deep chain cannot overflow the interpreter stack."""
+ seen: set[int] = set()
+ stack = [exc]
+ while stack:
+ current = stack.pop()
+ if id(current) in seen:
+ continue
+ seen.add(id(current))
+ yield current
+ if current.__context__ is not None:
+ stack.append(current.__context__)
+ exceptions = getattr(current, "exceptions", None)
+ if isinstance(exceptions, tuple):
+ stack.extend(reversed(exceptions))
+ if current.__cause__ is not None:
+ stack.append(current.__cause__)
diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py
index 1d681b43b9e..f2d3f568635 100644
--- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py
+++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py
@@ -51,6 +51,7 @@ from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import (
MCPRequestHandler,
)
from litellm.proxy._experimental.mcp_server.exceptions import MCPUpstreamAuthError
+from litellm.proxy._experimental.mcp_server.faults import iter_exception_tree
from litellm.proxy._experimental.mcp_server.elicitation_handler import (
MCP_ELICITATION_AVAILABLE,
)
@@ -440,44 +441,17 @@ def _extract_upstream_auth_failure(
upstream MCP server.
The MCP SDK wraps transport errors in anyio ``ExceptionGroup`` objects and
- may chain through ``__cause__`` / ``__context__``. We inspect all of those
- layers for an ``httpx.Response``-bearing exception (typically
- ``httpx.HTTPStatusError``) and extract the status code and any upstream
- ``WWW-Authenticate`` header.
+ may chain through ``__cause__`` / ``__context__``; ``iter_exception_tree``
+ visits all of those layers, explicit links first. The first exception
+ bearing a real ``httpx.Response`` with a 401/403 wins, and its status code
+ and upstream ``WWW-Authenticate`` header are extracted.
Returns ``(status_code, www_authenticate)`` on match, else ``None``.
"""
- seen: set[int] = set()
- stack: list[BaseException] = [exc]
- while stack:
- current = stack.pop()
- if id(current) in seen:
- continue
- seen.add(id(current))
-
+ for current in iter_exception_tree(exc):
response = getattr(current, "response", None)
- if response is not None:
- status_code = getattr(response, "status_code", None)
- if isinstance(status_code, int) and status_code in (401, 403):
- www_authenticate: Optional[str] = None
- headers = getattr(response, "headers", None)
- if headers is not None:
- try:
- www_authenticate = headers.get("www-authenticate")
- except Exception:
- www_authenticate = None
- return status_code, www_authenticate
-
- # anyio / PEP 654 ExceptionGroup
- sub_exceptions = getattr(current, "exceptions", None)
- if sub_exceptions:
- stack.extend(sub_exceptions)
-
- if current.__cause__ is not None:
- stack.append(current.__cause__)
- if current.__context__ is not None and current.__context__ is not current.__cause__:
- stack.append(current.__context__)
-
+ if isinstance(response, httpx.Response) and response.status_code in (401, 403):
+ return response.status_code, response.headers.get("www-authenticate")
return None
diff --git a/litellm/proxy/_experimental/mcp_server/semantic_tool_filter.py b/litellm/proxy/_experimental/mcp_server/semantic_tool_filter.py
index e12c6cdbd56..b22dd64e7fc 100644
--- a/litellm/proxy/_experimental/mcp_server/semantic_tool_filter.py
+++ b/litellm/proxy/_experimental/mcp_server/semantic_tool_filter.py
@@ -9,6 +9,7 @@ from typing import TYPE_CHECKING, Any, Dict, List, Optional
from litellm._logging import verbose_logger
from litellm.exceptions import ContextWindowExceededError
from litellm.litellm_core_utils.exception_mapping_utils import ExceptionCheckers
+from litellm.proxy._experimental.mcp_server.faults import iter_exception_tree
from litellm.proxy._experimental.mcp_server.utils import MCP_TOOL_PREFIX_SEPARATOR
if TYPE_CHECKING:
@@ -33,18 +34,15 @@ class SemanticToolFilterContextWindowError(Exception):
)
-def _is_context_window_error(error: Optional[BaseException], max_depth: int = 5) -> bool:
- """Detect a context-window overflow anywhere in an exception's cause chain."""
- current = error
- for _ in range(max_depth):
- if current is None:
- return False
- if isinstance(current, ContextWindowExceededError):
- return True
- if ExceptionCheckers.is_error_str_context_window_exceeded(str(current)):
- return True
- current = current.__cause__ or current.__context__
- return False
+def _is_context_window_error(error: Optional[BaseException]) -> bool:
+ """Detect a context-window overflow anywhere in an exception's tree."""
+ if error is None:
+ return False
+ return any(
+ isinstance(current, ContextWindowExceededError)
+ or ExceptionCheckers.is_error_str_context_window_exceeded(str(current))
+ for current in iter_exception_tree(error)
+ )
class SemanticMCPToolFilter:
diff --git a/ruff-strict-budget.json b/ruff-strict-budget.json
index dcde6fd1641..448a0079674 100644
--- a/ruff-strict-budget.json
+++ b/ruff-strict-budget.json
@@ -60,7 +60,7 @@
"limit": 4
},
"BLE001": {
- "limit": 2903
+ "limit": 2902
},
"C401": {
"limit": 11
@@ -363,6 +363,6 @@
"limit": 105
},
"UP045": {
- "limit": 18462
+ "limit": 18461
}
}
diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/faults/test_traversal.py b/tests/test_litellm/proxy/_experimental/mcp_server/faults/test_traversal.py
new file mode 100644
index 00000000000..a12c02339e6
--- /dev/null
+++ b/tests/test_litellm/proxy/_experimental/mcp_server/faults/test_traversal.py
@@ -0,0 +1,48 @@
+"""Traversal contract for the shared exception-tree walk: the root is yielded first, explicit
+links win (the ``raise ... from`` cause subtree, then ExceptionGroup members in raise order,
+then the incidental ``__context__`` chain last), and adversarial shapes terminate."""
+
+from litellm.proxy._experimental.mcp_server.faults import iter_exception_tree
+
+
+def test_yields_the_root_itself_first():
+ exc = ValueError("root")
+ assert list(iter_exception_tree(exc)) == [exc]
+
+
+def test_cause_subtree_is_exhausted_before_context():
+ deep = KeyError("deep")
+ cause = RuntimeError("cause")
+ cause.__cause__ = deep
+ context = OSError("context")
+ root = ValueError("root")
+ root.__cause__ = cause
+ root.__context__ = context
+ assert list(iter_exception_tree(root)) == [root, cause, deep, context]
+
+
+def test_group_members_yield_in_raise_order_between_cause_and_context():
+ first = KeyError("first")
+ second = IndexError("second")
+ group = BaseExceptionGroup("group", [first, second])
+ cause = RuntimeError("cause")
+ context = OSError("context")
+ group.__cause__ = cause
+ group.__context__ = context
+ assert list(iter_exception_tree(group)) == [group, cause, first, second, context]
+
+
+def test_terminates_on_a_cause_cycle():
+ a = ValueError("a")
+ b = RuntimeError("b")
+ a.__cause__ = b
+ b.__cause__ = a
+ assert list(iter_exception_tree(a)) == [a, b]
+
+
+def test_node_reachable_as_both_cause_and_context_yields_once():
+ inner = KeyError("inner")
+ root = ValueError("root")
+ root.__cause__ = inner
+ root.__context__ = inner
+ assert list(iter_exception_tree(root)) == [root, inner]
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 16c36af5156..617b5aa17fc 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
@@ -50,6 +50,36 @@ def test_extract_upstream_auth_failure_returns_none_for_non_auth():
assert _extract_upstream_auth_failure(RuntimeError("boom")) is None
+def _auth_status_error(status_code: int, www_authenticate: str) -> httpx.HTTPStatusError:
+ response = httpx.Response(
+ status_code=status_code,
+ headers={"www-authenticate": www_authenticate},
+ request=httpx.Request("GET", "https://upstream/mcp"),
+ )
+ return httpx.HTTPStatusError(str(status_code), request=response.request, response=response)
+
+
+def test_extract_upstream_auth_failure_finds_401_behind_cause_chain():
+ wrapper = RuntimeError("wrapped")
+ wrapper.__cause__ = _auth_status_error(401, "Bearer")
+ assert _extract_upstream_auth_failure(wrapper) == (401, "Bearer")
+
+
+def test_extract_upstream_auth_failure_finds_401_behind_context_chain():
+ wrapper = RuntimeError("wrapped")
+ wrapper.__context__ = _auth_status_error(401, "Bearer")
+ assert _extract_upstream_auth_failure(wrapper) == (401, "Bearer")
+
+
+def test_extract_upstream_auth_failure_prefers_causal_chain_over_context():
+ """A 403 raised incidentally while handling the real 401 (surviving only as ``__context__``)
+ must not shadow the 401 on the explicit ``raise ... from`` chain."""
+ wrapper = RuntimeError("wrapped")
+ wrapper.__cause__ = _auth_status_error(401, "Bearer realm=real")
+ wrapper.__context__ = _auth_status_error(403, "Bearer realm=incidental")
+ assert _extract_upstream_auth_failure(wrapper) == (401, "Bearer realm=real")
+
+
@pytest.mark.asyncio
async def test_fetch_tools_from_passthrough_raises_on_upstream_401():
manager = MCPServerManager()
diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_semantic_tool_filter.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_semantic_tool_filter.py
index 9bc0a525326..28ee7435702 100644
--- a/tests/test_litellm/proxy/_experimental/mcp_server/test_semantic_tool_filter.py
+++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_semantic_tool_filter.py
@@ -1664,3 +1664,31 @@ def test_is_context_window_error_detection_variants():
assert _is_context_window_error(ValueError("Invalid 'input[0]': maximum input length is 8192 tokens."))
assert not _is_context_window_error(ValueError("A generic API error occurred."))
assert not _is_context_window_error(None)
+
+
+def test_is_context_window_error_sees_through_trees_the_chain_walk_missed():
+ """Overflow shapes the old single-path depth-5 chain walk could not reach: hidden in
+ ``__context__`` behind a non-matching ``__cause__``, buried inside an anyio-style
+ ``ExceptionGroup``, and chained deeper than five links."""
+ import litellm
+ from litellm.proxy._experimental.mcp_server.semantic_tool_filter import (
+ _is_context_window_error,
+ )
+
+ def _cwe() -> litellm.ContextWindowExceededError:
+ return litellm.ContextWindowExceededError(message="overflow", model="m", llm_provider="openai")
+
+ shadowed = ValueError("wrapper")
+ shadowed.__cause__ = TypeError("unrelated failure")
+ shadowed.__context__ = _cwe()
+ assert _is_context_window_error(shadowed)
+
+ grouped = BaseExceptionGroup("task group", [RuntimeError("sibling"), _cwe()])
+ assert _is_context_window_error(grouped)
+
+ deep: BaseException = _cwe()
+ for depth in range(6):
+ wrapper = ValueError(f"layer {depth}")
+ wrapper.__cause__ = deep
+ deep = wrapper
+ assert _is_context_window_error(deep)
diff --git a/type-discipline-budget.json b/type-discipline-budget.json
index 87b4c96e323..83291fac388 100644
--- a/type-discipline-budget.json
+++ b/type-discipline-budget.json
@@ -1,6 +1,6 @@
{
"LIT001": {
- "limit": 23409
+ "limit": 23408
},
"LIT002": {
"limit": 27511
From ebb0f7e4cf1bfde5f720d89b76c4311176851878 Mon Sep 17 00:00:00 2001
From: Krrish Dholakia
Date: Wed, 15 Jul 2026 18:40:32 +0000
Subject: [PATCH 002/296] fix(responses): preserve reasoning through prompt
hooks
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
---
litellm/responses/main.py | 18 +++-
litellm/responses/utils.py | 50 ++++++++++
.../test_responses_prompt_management.py | 94 ++++++++++++++++---
3 files changed, 148 insertions(+), 14 deletions(-)
diff --git a/litellm/responses/main.py b/litellm/responses/main.py
index 12f9be970c7..453676937fc 100644
--- a/litellm/responses/main.py
+++ b/litellm/responses/main.py
@@ -494,7 +494,14 @@ async def aresponses(
prompt_label=kwargs.get("prompt_label", None),
prompt_version=kwargs.get("prompt_version", None),
)
- input = cast(Union[str, ResponseInputParam], merged_input)
+ input = cast(
+ Union[str, ResponseInputParam],
+ ResponsesAPIRequestUtils.merge_prompt_management_input(
+ original_input=input,
+ client_input=client_input,
+ merged_input=merged_input,
+ ),
+ )
if model != original_model:
_, custom_llm_provider, _, _ = litellm.get_llm_provider(model=model)
kwargs.pop("prompt_id", None)
@@ -609,7 +616,14 @@ def _apply_prompt_management_to_responses_call(
prompt_label=kwargs.get("prompt_label", None),
prompt_version=kwargs.get("prompt_version", None),
)
- input = cast(Union[str, ResponseInputParam], merged_input)
+ input = cast(
+ Union[str, ResponseInputParam],
+ ResponsesAPIRequestUtils.merge_prompt_management_input(
+ original_input=input,
+ client_input=client_input,
+ merged_input=merged_input,
+ ),
+ )
local_vars["input"] = input
local_vars["model"] = model
if model != original_model:
diff --git a/litellm/responses/utils.py b/litellm/responses/utils.py
index 234eb777aca..6d42e33a268 100644
--- a/litellm/responses/utils.py
+++ b/litellm/responses/utils.py
@@ -19,7 +19,9 @@ import litellm
from litellm._logging import verbose_logger
from litellm.llms.base_llm.responses.transformation import BaseResponsesAPIConfig
from litellm.types.llms.openai import (
+ AllMessageValues,
ResponseAPIUsage,
+ ResponseInputParam,
ResponsesAPIOptionalRequestParams,
ResponsesAPIResponse,
ResponseText,
@@ -36,6 +38,54 @@ from litellm.types.utils import (
class ResponsesAPIRequestUtils:
"""Helper utils for constructing ResponseAPI requests"""
+ @staticmethod
+ def merge_prompt_management_input(
+ original_input: str | ResponseInputParam,
+ client_input: list[AllMessageValues],
+ merged_input: list[AllMessageValues],
+ ) -> list[object]:
+ if isinstance(original_input, str):
+ return [*merged_input]
+
+ original_items = tuple(original_input)
+ client_item_ids = frozenset(id(item) for item in client_input)
+ message_positions = tuple(index for index, item in enumerate(original_items) if id(item) in client_item_ids)
+
+ if len(message_positions) == len(original_items):
+ return [*merged_input]
+ if not message_positions:
+ return [*merged_input, *original_items]
+
+ corresponding_messages = len(client_input) == len(merged_input) and all(
+ original.get("role") == merged.get("role")
+ and (not isinstance(original.get("id"), str) or original.get("id") == merged.get("id"))
+ for original, merged in zip(client_input, merged_input)
+ )
+ if corresponding_messages:
+ merged_by_position = dict(zip(message_positions, merged_input))
+ return [
+ merged_by_position[index] if index in merged_by_position else item
+ for index, item in enumerate(original_items)
+ ]
+
+ all_messages_preserved = all(any(original is merged for merged in merged_input) for original in client_input)
+ if all_messages_preserved:
+ prefixes = {
+ id(original_items[position]): original_items[
+ message_positions[index - 1] + 1 if index else 0 : position
+ ]
+ for index, position in enumerate(message_positions)
+ }
+ trailing_items = original_items[message_positions[-1] + 1 :]
+ return [item for merged in merged_input for item in (*prefixes.get(id(merged), ()), merged)] + list(
+ trailing_items
+ )
+
+ verbose_logger.warning(
+ "Prompt management hook replaced Responses API messages; non-message input items were dropped"
+ )
+ return [*merged_input]
+
@staticmethod
def _check_valid_arg(
supported_params: Optional[List[str]],
diff --git a/tests/test_litellm/responses/test_responses_prompt_management.py b/tests/test_litellm/responses/test_responses_prompt_management.py
index 84e98390268..e4207b292da 100644
--- a/tests/test_litellm/responses/test_responses_prompt_management.py
+++ b/tests/test_litellm/responses/test_responses_prompt_management.py
@@ -14,13 +14,19 @@ Covers:
"""
import asyncio
-from typing import List
+from typing import List, cast
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
+from litellm.integrations.anthropic_cache_control_hook import (
+ AnthropicCacheControlHook,
+)
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
-from litellm.types.llms.openai import AllMessageValues
+from litellm.types.llms.openai import (
+ AllMessageValues,
+ ResponseInputParam,
+)
# ---------------------------------------------------------------------------
# Helpers
@@ -54,18 +60,15 @@ def _patch_responses_dispatch():
return_value=("gpt-4o", "openai", None, None),
),
patch(
- "litellm.responses.mcp.litellm_proxy_mcp_handler."
- "LiteLLM_Proxy_MCP_Handler._should_use_litellm_mcp_gateway",
+ "litellm.responses.mcp.litellm_proxy_mcp_handler.LiteLLM_Proxy_MCP_Handler._should_use_litellm_mcp_gateway",
return_value=False,
),
patch(
- "litellm.responses.main.ProviderConfigManager"
- ".get_provider_responses_api_config",
+ "litellm.responses.main.ProviderConfigManager.get_provider_responses_api_config",
return_value=None,
),
patch(
- "litellm.responses.main.litellm_completion_transformation_handler"
- ".response_api_handler",
+ "litellm.responses.main.litellm_completion_transformation_handler.response_api_handler",
return_value=MagicMock(),
),
]
@@ -77,7 +80,6 @@ def _patch_responses_dispatch():
class TestResponsesAPIPromptManagement:
-
def test_str_input_coerced_and_merged(self):
"""[A] str input is wrapped into a message list before being passed to the hook."""
template_messages: List[AllMessageValues] = [
@@ -108,9 +110,7 @@ class TestResponsesAPIPromptManagement:
logging_obj.get_chat_completion_prompt.assert_called_once()
call_kwargs = logging_obj.get_chat_completion_prompt.call_args.kwargs
# str was coerced to a single user message before being passed to the hook
- assert call_kwargs["messages"] == [
- {"role": "user", "content": "Tell me about AI."}
- ]
+ assert call_kwargs["messages"] == [{"role": "user", "content": "Tell me about AI."}]
assert call_kwargs["prompt_id"] == "summariser-prompt"
def test_list_input_merged_with_template(self):
@@ -256,6 +256,76 @@ class TestResponsesAPIPromptManagement:
assert all(isinstance(m, dict) and "role" in m for m in passed_messages)
assert len(passed_messages) == 1
+ def test_cache_control_hook_preserves_reasoning_items(self):
+ system_message = cast(
+ AllMessageValues,
+ {"role": "system", "content": "Analyze the request"},
+ )
+ assistant_message = cast(
+ AllMessageValues,
+ {
+ "type": "message",
+ "id": "msg_1",
+ "role": "assistant",
+ "status": "completed",
+ "content": [
+ {
+ "type": "output_text",
+ "text": "The code has a bug",
+ "annotations": [],
+ }
+ ],
+ },
+ )
+ user_message = cast(
+ AllMessageValues,
+ {"role": "user", "content": "Check for security issues"},
+ )
+ reasoning_item = {
+ "type": "reasoning",
+ "id": "rs_1",
+ "summary": [],
+ "encrypted_content": "encrypted",
+ }
+ original_input = cast(
+ ResponseInputParam,
+ [system_message, reasoning_item, assistant_message, user_message],
+ )
+ _, merged_messages, _ = AnthropicCacheControlHook().get_chat_completion_prompt(
+ model="azure/gpt-5-codex",
+ messages=[system_message, assistant_message, user_message],
+ non_default_params={"cache_control_injection_points": [{"location": "message", "role": "system"}]},
+ prompt_id=None,
+ prompt_variables=None,
+ dynamic_callback_params={},
+ )
+ logging_obj = _make_logging_obj(
+ merged_model="azure/gpt-5-codex",
+ merged_messages=merged_messages,
+ )
+
+ patches = _patch_responses_dispatch()
+ with patches[0], patches[1], patches[2], patches[3] as mock_handler:
+ import litellm
+
+ litellm.responses(
+ input=original_input,
+ model="azure/gpt-5-codex",
+ litellm_logging_obj=logging_obj,
+ cache_control_injection_points=[{"location": "message", "role": "system"}],
+ )
+
+ sent_input = mock_handler.call_args.kwargs["input"]
+ assert [item.get("type") for item in sent_input] == [
+ None,
+ "reasoning",
+ "message",
+ None,
+ ]
+ assert sent_input[0]["cache_control"] == {"type": "ephemeral"}
+ assert sent_input[1] == reasoning_item
+ assert sent_input[2]["id"] == "msg_1"
+
def test_model_override_re_resolves_provider(self):
"""[G] When the prompt template overrides the model to a different provider,
custom_llm_provider is re-resolved so downstream routing uses the correct provider.
From 4baee71bdd1e82db5225122ad5d9a1a7ae925af2 Mon Sep 17 00:00:00 2001
From: Krrish Dholakia
Date: Wed, 15 Jul 2026 18:41:02 +0000
Subject: [PATCH 003/296] chore(responses): minimize regression test diff
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
---
.../responses/test_responses_prompt_management.py | 14 ++++++++++----
1 file changed, 10 insertions(+), 4 deletions(-)
diff --git a/tests/test_litellm/responses/test_responses_prompt_management.py b/tests/test_litellm/responses/test_responses_prompt_management.py
index e4207b292da..b3ba81ee2e8 100644
--- a/tests/test_litellm/responses/test_responses_prompt_management.py
+++ b/tests/test_litellm/responses/test_responses_prompt_management.py
@@ -60,15 +60,18 @@ def _patch_responses_dispatch():
return_value=("gpt-4o", "openai", None, None),
),
patch(
- "litellm.responses.mcp.litellm_proxy_mcp_handler.LiteLLM_Proxy_MCP_Handler._should_use_litellm_mcp_gateway",
+ "litellm.responses.mcp.litellm_proxy_mcp_handler."
+ "LiteLLM_Proxy_MCP_Handler._should_use_litellm_mcp_gateway",
return_value=False,
),
patch(
- "litellm.responses.main.ProviderConfigManager.get_provider_responses_api_config",
+ "litellm.responses.main.ProviderConfigManager"
+ ".get_provider_responses_api_config",
return_value=None,
),
patch(
- "litellm.responses.main.litellm_completion_transformation_handler.response_api_handler",
+ "litellm.responses.main.litellm_completion_transformation_handler"
+ ".response_api_handler",
return_value=MagicMock(),
),
]
@@ -80,6 +83,7 @@ def _patch_responses_dispatch():
class TestResponsesAPIPromptManagement:
+
def test_str_input_coerced_and_merged(self):
"""[A] str input is wrapped into a message list before being passed to the hook."""
template_messages: List[AllMessageValues] = [
@@ -110,7 +114,9 @@ class TestResponsesAPIPromptManagement:
logging_obj.get_chat_completion_prompt.assert_called_once()
call_kwargs = logging_obj.get_chat_completion_prompt.call_args.kwargs
# str was coerced to a single user message before being passed to the hook
- assert call_kwargs["messages"] == [{"role": "user", "content": "Tell me about AI."}]
+ assert call_kwargs["messages"] == [
+ {"role": "user", "content": "Tell me about AI."}
+ ]
assert call_kwargs["prompt_id"] == "summariser-prompt"
def test_list_input_merged_with_template(self):
From 4db0bdf465c9f07e8561136cf055ca9e6854f086 Mon Sep 17 00:00:00 2001
From: Krrish Dholakia
Date: Wed, 15 Jul 2026 18:51:07 +0000
Subject: [PATCH 004/296] fix(responses): handle non-message-only prompt input
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
---
litellm/responses/utils.py | 5 +-
.../test_responses_prompt_management.py | 154 +++++++++++++-----
2 files changed, 116 insertions(+), 43 deletions(-)
diff --git a/litellm/responses/utils.py b/litellm/responses/utils.py
index 6d42e33a268..a5203c4ee6a 100644
--- a/litellm/responses/utils.py
+++ b/litellm/responses/utils.py
@@ -54,7 +54,10 @@ class ResponsesAPIRequestUtils:
if len(message_positions) == len(original_items):
return [*merged_input]
if not message_positions:
- return [*merged_input, *original_items]
+ verbose_logger.warning(
+ "Prompt management hook returned messages without Responses API input messages; merged messages were ignored"
+ )
+ return [*original_items]
corresponding_messages = len(client_input) == len(merged_input) and all(
original.get("role") == merged.get("role")
diff --git a/tests/test_litellm/responses/test_responses_prompt_management.py b/tests/test_litellm/responses/test_responses_prompt_management.py
index b3ba81ee2e8..7044d8384f8 100644
--- a/tests/test_litellm/responses/test_responses_prompt_management.py
+++ b/tests/test_litellm/responses/test_responses_prompt_management.py
@@ -77,6 +77,56 @@ def _patch_responses_dispatch():
]
+def _make_cache_control_case() -> tuple[
+ ResponseInputParam,
+ list[AllMessageValues],
+ dict[str, object],
+]:
+ system_message = cast(
+ AllMessageValues,
+ {"role": "system", "content": "Analyze the request"},
+ )
+ assistant_message = cast(
+ AllMessageValues,
+ {
+ "type": "message",
+ "id": "msg_1",
+ "role": "assistant",
+ "status": "completed",
+ "content": [
+ {
+ "type": "output_text",
+ "text": "The code has a bug",
+ "annotations": [],
+ }
+ ],
+ },
+ )
+ user_message = cast(
+ AllMessageValues,
+ {"role": "user", "content": "Check for security issues"},
+ )
+ reasoning_item = {
+ "type": "reasoning",
+ "id": "rs_1",
+ "summary": [],
+ "encrypted_content": "encrypted",
+ }
+ original_input = cast(
+ ResponseInputParam,
+ [system_message, reasoning_item, assistant_message, user_message],
+ )
+ _, merged_messages, _ = AnthropicCacheControlHook().get_chat_completion_prompt(
+ model="azure/gpt-5-codex",
+ messages=[system_message, assistant_message, user_message],
+ non_default_params={"cache_control_injection_points": [{"location": "message", "role": "system"}]},
+ prompt_id=None,
+ prompt_variables=None,
+ dynamic_callback_params={},
+ )
+ return original_input, merged_messages, reasoning_item
+
+
# ---------------------------------------------------------------------------
# Tests
# ---------------------------------------------------------------------------
@@ -263,48 +313,7 @@ class TestResponsesAPIPromptManagement:
assert len(passed_messages) == 1
def test_cache_control_hook_preserves_reasoning_items(self):
- system_message = cast(
- AllMessageValues,
- {"role": "system", "content": "Analyze the request"},
- )
- assistant_message = cast(
- AllMessageValues,
- {
- "type": "message",
- "id": "msg_1",
- "role": "assistant",
- "status": "completed",
- "content": [
- {
- "type": "output_text",
- "text": "The code has a bug",
- "annotations": [],
- }
- ],
- },
- )
- user_message = cast(
- AllMessageValues,
- {"role": "user", "content": "Check for security issues"},
- )
- reasoning_item = {
- "type": "reasoning",
- "id": "rs_1",
- "summary": [],
- "encrypted_content": "encrypted",
- }
- original_input = cast(
- ResponseInputParam,
- [system_message, reasoning_item, assistant_message, user_message],
- )
- _, merged_messages, _ = AnthropicCacheControlHook().get_chat_completion_prompt(
- model="azure/gpt-5-codex",
- messages=[system_message, assistant_message, user_message],
- non_default_params={"cache_control_injection_points": [{"location": "message", "role": "system"}]},
- prompt_id=None,
- prompt_variables=None,
- dynamic_callback_params={},
- )
+ original_input, merged_messages, reasoning_item = _make_cache_control_case()
logging_obj = _make_logging_obj(
merged_model="azure/gpt-5-codex",
merged_messages=merged_messages,
@@ -332,6 +341,37 @@ class TestResponsesAPIPromptManagement:
assert sent_input[1] == reasoning_item
assert sent_input[2]["id"] == "msg_1"
+ def test_all_non_message_input_items_remain_unchanged(self):
+ reasoning_item = {
+ "type": "reasoning",
+ "id": "rs_1",
+ "summary": [],
+ "encrypted_content": "encrypted",
+ }
+ original_input = cast(ResponseInputParam, [reasoning_item])
+ logging_obj = _make_logging_obj(
+ merged_model="openai/gpt-4o",
+ merged_messages=[
+ cast(
+ AllMessageValues,
+ {"role": "system", "content": "Analyze the request"},
+ )
+ ],
+ )
+
+ patches = _patch_responses_dispatch()
+ with patches[0], patches[1], patches[2], patches[3] as mock_handler:
+ import litellm
+
+ litellm.responses(
+ input=original_input,
+ model="gpt-4o",
+ prompt_id="all-non-message",
+ litellm_logging_obj=logging_obj,
+ )
+
+ assert mock_handler.call_args.kwargs["input"] == original_input
+
def test_model_override_re_resolves_provider(self):
"""[G] When the prompt template overrides the model to a different provider,
custom_llm_provider is re-resolved so downstream routing uses the correct provider.
@@ -469,3 +509,33 @@ class TestAsyncResponsesAPIPromptManagement:
passed_messages = call_kwargs["messages"]
assert all(isinstance(m, dict) and "role" in m for m in passed_messages)
assert len(passed_messages) == 1
+
+ @pytest.mark.asyncio
+ async def test_async_cache_control_hook_preserves_reasoning_items(self):
+ original_input, merged_messages, reasoning_item = _make_cache_control_case()
+ logging_obj = _make_logging_obj(
+ merged_model="azure/gpt-5-codex",
+ merged_messages=merged_messages,
+ )
+
+ patches = _patch_responses_dispatch()
+ with patches[0], patches[1], patches[2], patches[3] as mock_handler:
+ import litellm
+
+ await litellm.aresponses(
+ input=original_input,
+ model="azure/gpt-5-codex",
+ litellm_logging_obj=logging_obj,
+ cache_control_injection_points=[{"location": "message", "role": "system"}],
+ )
+
+ sent_input = mock_handler.call_args.kwargs["input"]
+ assert [item.get("type") for item in sent_input] == [
+ None,
+ "reasoning",
+ "message",
+ None,
+ ]
+ assert sent_input[0]["cache_control"] == {"type": "ephemeral"}
+ assert sent_input[1] == reasoning_item
+ assert sent_input[2]["id"] == "msg_1"
From ae952ce971ff52c91a17abb5e89bd1062383820a Mon Sep 17 00:00:00 2001
From: Tin Chi Lo
Date: Thu, 16 Jul 2026 18:35:58 -0700
Subject: [PATCH 005/296] feat(mcp): support MCP servers on the Anthropic
/v1/messages API
MCP tool calling worked on /v1/chat/completions and /v1/responses but not on
/v1/messages. Those are the only two surfaces with an MCP gateway entry point,
so a litellm_proxy MCP reference reached Anthropic verbatim inside tools and the
API rejected the request with "Input tag 'mcp' found using 'type' does not match
any of the expected tags". The playground never surfaced this because it dropped
the reference before sending, and disabled the MCP selector for the endpoint.
Add the third entry point in anthropic_messages_handler, ahead of the provider
branch so it covers the native path and both bridges from one place. The gateway
expands the reference against the caller's own credentials and access control,
which is the whole point of routing it through litellm rather than handing the
url to the provider.
/v1/messages needs Anthropic's own tool shape, so transform_mcp_tool_to_anthropic_tool
joins the OpenAI chat and Responses transforms alongside it. The tool loop speaks
tool_use and tool_result rather than OpenAI tool_calls, and reuses the existing
FakeAnthropicMessagesStreamIterator to re-stream the result, the same pattern the
websearch interception already uses on this route. Argument extraction moves into
the shared extractor: an Anthropic tool_use block carries its arguments under
`input`, and reading only `arguments` failed silently, executing the tool with
every argument dropped.
On the frontend the request builder declared selectedMCPTools and never read it,
so no tools key was ever sent. Wire it through a shared block builder and add the
endpoint to MCP_SUPPORTED_ENDPOINTS, which is what greys the selector out.
Resolves LIT-4517
Resolves LIT-4518
---
litellm/experimental_mcp_client/tools.py | 13 ++
.../messages/handler.py | 35 ++++
.../messages/mcp_handler.py | 174 ++++++++++++++++++
.../mcp/litellm_proxy_mcp_handler.py | 5 +
.../experimental_mcp_client/test_tools.py | 49 +++++
.../messages/test_mcp_handler.py | 122 ++++++++++++
.../mcp/test_litellm_proxy_mcp_handler.py | 42 +++++
.../playground/components/chat_ui/ChatUI.tsx | 12 +-
.../llm_calls/anthropic_messages.tsx | 14 +-
.../components/llm_calls/mcp_tool_blocks.ts | 79 ++++++++
10 files changed, 542 insertions(+), 3 deletions(-)
create mode 100644 litellm/llms/anthropic/experimental_pass_through/messages/mcp_handler.py
create mode 100644 tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_mcp_handler.py
create mode 100644 ui/litellm-dashboard/src/components/llm_calls/mcp_tool_blocks.ts
diff --git a/litellm/experimental_mcp_client/tools.py b/litellm/experimental_mcp_client/tools.py
index c65b266bd02..1bd65847616 100644
--- a/litellm/experimental_mcp_client/tools.py
+++ b/litellm/experimental_mcp_client/tools.py
@@ -9,6 +9,7 @@ from openai.types.chat import ChatCompletionToolParam
from openai.types.responses.function_tool_param import FunctionToolParam
from openai.types.shared_params.function_definition import FunctionDefinition
+from litellm.types.llms.anthropic import AnthropicInputSchema, AnthropicMessagesTool
from litellm.types.utils import ChatCompletionMessageToolCall
@@ -75,6 +76,18 @@ def transform_mcp_tool_to_openai_responses_api_tool(
)
+def transform_mcp_tool_to_anthropic_tool(mcp_tool: MCPTool) -> AnthropicMessagesTool:
+ """Convert an MCP tool to an Anthropic Messages API tool."""
+ normalized_parameters = _normalize_mcp_input_schema(mcp_tool.inputSchema)
+
+ return AnthropicMessagesTool(
+ name=mcp_tool.name,
+ description=mcp_tool.description or "",
+ input_schema=AnthropicInputSchema(**normalized_parameters),
+ type="custom",
+ )
+
+
async def load_mcp_tools(
session: ClientSession, format: Literal["mcp", "openai"] = "mcp"
) -> Union[List[MCPTool], List[ChatCompletionToolParam]]:
diff --git a/litellm/llms/anthropic/experimental_pass_through/messages/handler.py b/litellm/llms/anthropic/experimental_pass_through/messages/handler.py
index dd983f0c344..499f5bc486c 100644
--- a/litellm/llms/anthropic/experimental_pass_through/messages/handler.py
+++ b/litellm/llms/anthropic/experimental_pass_through/messages/handler.py
@@ -477,6 +477,41 @@ def anthropic_messages_handler(
mock_response=litellm_params.mock_response,
)
+ # Expand litellm_proxy MCP references through the MCP gateway before dispatch, so every
+ # downstream path (native passthrough and both bridges) gets real tools rather than a
+ # reference the provider cannot resolve. Popped from kwargs so it never reaches the provider.
+ skip_mcp_handler = kwargs.pop("_skip_mcp_handler", False)
+ if not skip_mcp_handler and tools:
+ from litellm.llms.anthropic.experimental_pass_through.messages.mcp_handler import (
+ anthropic_messages_with_mcp,
+ )
+ from litellm.responses.mcp.litellm_proxy_mcp_handler import (
+ LiteLLM_Proxy_MCP_Handler,
+ )
+
+ if LiteLLM_Proxy_MCP_Handler._should_use_litellm_mcp_gateway(tools=tools):
+ return anthropic_messages_with_mcp(
+ max_tokens=max_tokens,
+ messages=messages,
+ model=model,
+ metadata=metadata,
+ stop_sequences=stop_sequences,
+ stream=stream,
+ system=system,
+ temperature=temperature,
+ thinking=thinking,
+ tool_choice=tool_choice,
+ tools=tools,
+ top_k=top_k,
+ top_p=top_p,
+ container=container,
+ api_key=api_key,
+ api_base=api_base,
+ client=client,
+ custom_llm_provider=custom_llm_provider,
+ **kwargs,
+ )
+
anthropic_messages_provider_config: Optional[BaseAnthropicMessagesConfig] = None
if custom_llm_provider is not None and custom_llm_provider in [provider.value for provider in LlmProviders]:
diff --git a/litellm/llms/anthropic/experimental_pass_through/messages/mcp_handler.py b/litellm/llms/anthropic/experimental_pass_through/messages/mcp_handler.py
new file mode 100644
index 00000000000..392b9e2e02d
--- /dev/null
+++ b/litellm/llms/anthropic/experimental_pass_through/messages/mcp_handler.py
@@ -0,0 +1,174 @@
+"""
+MCP gateway support for the Anthropic `/v1/messages` API.
+
+Mirrors ``litellm.responses.mcp.chat_completions_handler`` but speaks the
+Anthropic Messages shapes: tools carry an ``input_schema``, the model asks for a
+tool through a ``tool_use`` content block, and results are fed back as
+``tool_result`` blocks in a user message.
+"""
+
+from typing import Any, AsyncIterator, Mapping, Sequence, Union
+
+from litellm._logging import verbose_logger
+from litellm.types.llms.anthropic import (
+ AnthropicMessagesTool,
+ AnthropicMessagesToolResultParam,
+ AnthropicMessagesUserMessageParam,
+)
+from litellm.types.llms.anthropic_messages.anthropic_response import (
+ AnthropicMessagesResponse,
+)
+
+MAX_MCP_TOOL_USE_ITERATIONS = 10
+
+
+def _get_response_content(response: AnthropicMessagesResponse) -> Sequence[Mapping[str, Any]]:
+ content = response.get("content")
+ if not isinstance(content, list):
+ return ()
+ return tuple(block for block in content if isinstance(block, dict))
+
+
+def _extract_tool_use_blocks(response: AnthropicMessagesResponse) -> Sequence[Mapping[str, Any]]:
+ """Return the ``tool_use`` content blocks the model emitted."""
+ return tuple(block for block in _get_response_content(response) if block.get("type") == "tool_use")
+
+
+def _get_stop_reason(response: AnthropicMessagesResponse) -> Union[str, None]:
+ stop_reason = response.get("stop_reason")
+ return stop_reason if isinstance(stop_reason, str) else None
+
+
+def _build_tool_result_message(tool_results: Sequence[Mapping[str, Any]]) -> AnthropicMessagesUserMessageParam:
+ """Turn executed tool results into the user message Anthropic expects."""
+ return AnthropicMessagesUserMessageParam(
+ role="user",
+ content=tuple(
+ AnthropicMessagesToolResultParam(
+ type="tool_result",
+ tool_use_id=str(result.get("tool_call_id") or ""),
+ content=str(result.get("result") or ""),
+ )
+ for result in tool_results
+ ),
+ )
+
+
+def _resolve_user_api_key_auth(
+ kwargs: Mapping[str, Any],
+) -> Any: # any-ok: UserAPIKeyAuth is proxy-only, importing it here would create a cycle
+ """`/v1/messages` is a LITELLM_METADATA_ROUTE, so the auth object rides in litellm_metadata."""
+ litellm_metadata = kwargs.get("litellm_metadata") or {}
+ metadata = kwargs.get("metadata") or {}
+ return (
+ kwargs.get("user_api_key_auth")
+ or litellm_metadata.get("user_api_key_auth")
+ or metadata.get("user_api_key_auth")
+ )
+
+
+async def anthropic_messages_with_mcp(
+ max_tokens: int,
+ messages: Sequence[Mapping[str, Any]],
+ model: str,
+ tools: Union[Sequence[Mapping[str, Any]], None] = None,
+ **kwargs: Any, # kwargs-ok: forwarded verbatim to litellm.anthropic_messages, which owns the param contract
+) -> Union[AnthropicMessagesResponse, AsyncIterator[Any]]:
+ """
+ Expand litellm_proxy MCP references for `/v1/messages` and run the tool loop.
+
+ The MCP gateway owns the expansion so the reference resolves against the
+ caller's own credentials and access control, rather than being handed to the
+ upstream provider as a url it cannot reach.
+ """
+ import litellm
+ from litellm.experimental_mcp_client.tools import (
+ transform_mcp_tool_to_anthropic_tool,
+ )
+ from litellm.responses.mcp.litellm_proxy_mcp_handler import (
+ LiteLLM_Proxy_MCP_Handler,
+ )
+
+ mcp_references, other_tools = LiteLLM_Proxy_MCP_Handler._parse_mcp_tools(tools)
+
+ if not mcp_references:
+ return await litellm.anthropic_messages(
+ max_tokens=max_tokens,
+ messages=list(messages),
+ model=model,
+ tools=list(tools) if tools else None,
+ _skip_mcp_handler=True,
+ **kwargs,
+ )
+
+ user_api_key_auth = _resolve_user_api_key_auth(kwargs)
+
+ (
+ deduplicated_mcp_tools,
+ tool_server_map,
+ ) = await LiteLLM_Proxy_MCP_Handler._process_mcp_tools_without_openai_transform(
+ user_api_key_auth,
+ mcp_references,
+ litellm_trace_id=kwargs.get("litellm_trace_id"),
+ )
+
+ anthropic_tools: Sequence[AnthropicMessagesTool] = tuple(
+ transform_mcp_tool_to_anthropic_tool(mcp_tool) for mcp_tool in deduplicated_mcp_tools
+ )
+ all_tools = [*anthropic_tools, *(other_tools or ())]
+
+ should_auto_execute = LiteLLM_Proxy_MCP_Handler._should_auto_execute_tools(
+ mcp_tools_with_litellm_proxy=mcp_references
+ )
+ stream = bool(kwargs.pop("stream", False))
+
+ base_call_args: Mapping[str, Any] = {
+ "max_tokens": max_tokens,
+ "model": model,
+ "tools": all_tools or None,
+ "_skip_mcp_handler": True,
+ **kwargs,
+ }
+
+ if not should_auto_execute:
+ return await litellm.anthropic_messages(messages=list(messages), stream=stream, **base_call_args)
+
+ working_messages: Sequence[Mapping[str, Any]] = tuple(messages)
+ response: AnthropicMessagesResponse = await litellm.anthropic_messages(
+ messages=list(working_messages), stream=False, **base_call_args
+ )
+
+ for _ in range(MAX_MCP_TOOL_USE_ITERATIONS):
+ if _get_stop_reason(response) != "tool_use":
+ break
+
+ tool_use_blocks = _extract_tool_use_blocks(response)
+ if not tool_use_blocks:
+ break
+
+ tool_results = await LiteLLM_Proxy_MCP_Handler._execute_tool_calls(
+ tool_server_map=tool_server_map,
+ tool_calls=list(tool_use_blocks),
+ user_api_key_auth=user_api_key_auth,
+ litellm_trace_id=kwargs.get("litellm_trace_id"),
+ )
+
+ working_messages = (
+ *working_messages,
+ {"role": "assistant", "content": list(_get_response_content(response))},
+ _build_tool_result_message(tool_results),
+ )
+ response = await litellm.anthropic_messages(messages=list(working_messages), stream=False, **base_call_args)
+ else:
+ verbose_logger.warning(
+ f"MCP tool loop hit its {MAX_MCP_TOOL_USE_ITERATIONS} iteration cap for model {model}; "
+ "returning the last response"
+ )
+
+ if stream:
+ from litellm.llms.anthropic.experimental_pass_through.messages.fake_stream_iterator import (
+ FakeAnthropicMessagesStreamIterator,
+ )
+
+ return FakeAnthropicMessagesStreamIterator(response)
+ return response
diff --git a/litellm/responses/mcp/litellm_proxy_mcp_handler.py b/litellm/responses/mcp/litellm_proxy_mcp_handler.py
index e03f0296109..d2c9f220690 100644
--- a/litellm/responses/mcp/litellm_proxy_mcp_handler.py
+++ b/litellm/responses/mcp/litellm_proxy_mcp_handler.py
@@ -541,7 +541,10 @@ class LiteLLM_Proxy_MCP_Handler:
tool_arguments = function_block.get("arguments")
else:
tool_name = tool_call.get("name")
+ # Anthropic tool_use blocks carry the arguments under `input`
tool_arguments = tool_call.get("arguments")
+ if tool_arguments is None:
+ tool_arguments = tool_call.get("input")
else:
tool_call_id = getattr(tool_call, "call_id", None) or getattr(tool_call, "id", None)
@@ -552,6 +555,8 @@ class LiteLLM_Proxy_MCP_Handler:
else:
tool_name = getattr(tool_call, "name", None)
tool_arguments = getattr(tool_call, "arguments", None)
+ if tool_arguments is None:
+ tool_arguments = getattr(tool_call, "input", None)
return tool_name, tool_arguments, tool_call_id
diff --git a/tests/test_litellm/experimental_mcp_client/test_tools.py b/tests/test_litellm/experimental_mcp_client/test_tools.py
index 786bbf7dcc9..625ab56951f 100644
--- a/tests/test_litellm/experimental_mcp_client/test_tools.py
+++ b/tests/test_litellm/experimental_mcp_client/test_tools.py
@@ -18,6 +18,7 @@ from mcp.types import (
from mcp.types import Tool as MCPTool
from litellm.experimental_mcp_client.tools import (
+ transform_mcp_tool_to_anthropic_tool,
_get_function_arguments,
_normalize_mcp_input_schema,
call_mcp_tool,
@@ -250,3 +251,51 @@ def test_transform_mcp_tool_to_openai_responses_api_tool():
assert "query" in openai_tool["parameters"]["properties"]
assert openai_tool["parameters"]["required"] == ["query"]
assert openai_tool["parameters"]["additionalProperties"] == False
+
+
+def test_transform_mcp_tool_to_anthropic_tool():
+ """
+ Regression test (LIT-4517): MCP tools must reach /v1/messages in Anthropic's
+ own tool shape.
+
+ Given: An MCP tool
+ When: It is transformed for the Anthropic Messages API
+ Then: It carries name/description/input_schema, the shape that endpoint
+ accepts, rather than an OpenAI function block
+
+ /v1/messages rejects an OpenAI-shaped tool outright ("Input tag 'function'
+ does not match any of the expected tags"), so reusing either OpenAI
+ transform here loses every MCP tool.
+ """
+ tool = MCPTool(
+ name="read_wiki_structure",
+ description="Get a list of documentation topics",
+ inputSchema={
+ "type": "object",
+ "properties": {"repoName": {"type": "string"}},
+ "required": ["repoName"],
+ },
+ )
+
+ anthropic_tool = transform_mcp_tool_to_anthropic_tool(tool)
+
+ assert anthropic_tool["name"] == "read_wiki_structure"
+ assert anthropic_tool["description"] == "Get a list of documentation topics"
+ assert anthropic_tool["type"] == "custom"
+ assert anthropic_tool["input_schema"]["type"] == "object"
+ assert "repoName" in anthropic_tool["input_schema"]["properties"]
+ assert anthropic_tool["input_schema"]["required"] == ["repoName"]
+ assert "function" not in anthropic_tool, "Anthropic tools must not carry an OpenAI function block"
+ assert "parameters" not in anthropic_tool, "Anthropic names the schema input_schema, not parameters"
+
+
+def test_transform_mcp_tool_to_anthropic_tool_normalizes_empty_schema():
+ """A tool with no declared arguments must still present a valid object schema."""
+ anthropic_tool = transform_mcp_tool_to_anthropic_tool(
+ MCPTool(name="noargs", description=None, inputSchema={})
+ )
+
+ assert anthropic_tool["name"] == "noargs"
+ assert anthropic_tool["description"] == ""
+ assert anthropic_tool["input_schema"]["type"] == "object"
+ assert anthropic_tool["input_schema"]["properties"] == {}
diff --git a/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_mcp_handler.py b/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_mcp_handler.py
new file mode 100644
index 00000000000..3faa6b1e4e2
--- /dev/null
+++ b/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_mcp_handler.py
@@ -0,0 +1,122 @@
+import os
+import sys
+from unittest.mock import AsyncMock, patch
+
+import pytest
+
+sys.path.insert(0, os.path.abspath("../../../../../.."))
+
+from litellm.llms.anthropic.experimental_pass_through.messages.handler import (
+ anthropic_messages_handler,
+)
+from litellm.llms.anthropic.experimental_pass_through.messages.mcp_handler import (
+ _build_tool_result_message,
+ _extract_tool_use_blocks,
+)
+
+MCP_REFERENCE = {
+ "type": "mcp",
+ "server_label": "litellm",
+ "server_url": "litellm_proxy/mcp/deepwiki",
+ "require_approval": "never",
+}
+
+
+def test_anthropic_messages_handler_routes_litellm_proxy_mcp_to_the_gateway():
+ """
+ Regression test (LIT-4517): /v1/messages must expand a litellm_proxy MCP
+ reference through the MCP gateway.
+
+ Given: A /v1/messages request whose tools carry a litellm_proxy MCP reference
+ When: The handler dispatches
+ Then: It hands off to the MCP gateway instead of the provider
+
+ Without this hook the reference is forwarded to Anthropic verbatim and the API
+ rejects the request ("Input tag 'mcp' found using 'type' does not match any of
+ the expected tags"), because only /v1/chat/completions and /v1/responses ever
+ had a gateway entry point. This pins the wiring, not the helper: deleting the
+ dispatch makes the whole feature unreachable while every unit test still passes.
+ """
+ with patch(
+ "litellm.llms.anthropic.experimental_pass_through.messages.mcp_handler.anthropic_messages_with_mcp",
+ new=AsyncMock(return_value={"routed": True}),
+ ) as routed:
+ result = anthropic_messages_handler(
+ max_tokens=100,
+ messages=[{"role": "user", "content": "hi"}],
+ model="claude-sonnet-4-5",
+ tools=[MCP_REFERENCE],
+ custom_llm_provider="anthropic",
+ )
+
+ assert routed.called, "A litellm_proxy MCP reference must be dispatched to the MCP gateway"
+ assert routed.call_args.kwargs["tools"] == [MCP_REFERENCE]
+ assert routed.call_args.kwargs["model"] == "claude-sonnet-4-5"
+ assert result is not None
+
+
+def test_anthropic_messages_handler_skips_the_gateway_on_recursion():
+ """The gateway's own follow-up call must not re-enter the gateway."""
+ with patch(
+ "litellm.llms.anthropic.experimental_pass_through.messages.mcp_handler.anthropic_messages_with_mcp",
+ new=AsyncMock(return_value={"routed": True}),
+ ) as routed:
+ with pytest.raises(Exception):
+ anthropic_messages_handler(
+ max_tokens=100,
+ messages=[{"role": "user", "content": "hi"}],
+ model="claude-sonnet-4-5",
+ tools=[MCP_REFERENCE],
+ custom_llm_provider="anthropic",
+ _skip_mcp_handler=True,
+ )
+
+ assert not routed.called, "_skip_mcp_handler must stop the gateway from recursing"
+
+
+def test_anthropic_messages_handler_leaves_native_tools_alone():
+ """A plain Anthropic tool is not an MCP reference and must not reach the gateway."""
+ with patch(
+ "litellm.llms.anthropic.experimental_pass_through.messages.mcp_handler.anthropic_messages_with_mcp",
+ new=AsyncMock(return_value={"routed": True}),
+ ) as routed:
+ with pytest.raises(Exception):
+ anthropic_messages_handler(
+ max_tokens=100,
+ messages=[{"role": "user", "content": "hi"}],
+ model="claude-sonnet-4-5",
+ tools=[{"name": "get_weather", "input_schema": {"type": "object"}}],
+ custom_llm_provider="anthropic",
+ )
+
+ assert not routed.called, "Only litellm_proxy MCP references belong to the gateway"
+
+
+def test_extract_tool_use_blocks_ignores_text_blocks():
+ """Only tool_use blocks drive the loop; text blocks are the model's prose."""
+ response = {
+ "content": [
+ {"type": "text", "text": "let me look that up"},
+ {"type": "tool_use", "id": "toolu_1", "name": "read_wiki_structure", "input": {"repoName": "a/b"}},
+ ]
+ }
+
+ blocks = _extract_tool_use_blocks(response)
+
+ assert len(blocks) == 1
+ assert blocks[0]["name"] == "read_wiki_structure"
+
+
+def test_build_tool_result_message_uses_anthropic_tool_result_blocks():
+ """
+ Results must go back as tool_result blocks in a user message.
+
+ Anthropic pairs each result to its request by tool_use_id; the OpenAI shape
+ (a role="tool" message keyed by tool_call_id) is rejected here.
+ """
+ message = _build_tool_result_message([{"tool_call_id": "toolu_1", "result": "9 sections", "name": "read_wiki"}])
+
+ assert message["role"] == "user"
+ assert list(message["content"]) == [
+ {"type": "tool_result", "tool_use_id": "toolu_1", "content": "9 sections"}
+ ]
diff --git a/tests/test_litellm/responses/mcp/test_litellm_proxy_mcp_handler.py b/tests/test_litellm/responses/mcp/test_litellm_proxy_mcp_handler.py
index 6fdbb0741aa..1c23b1a8b98 100644
--- a/tests/test_litellm/responses/mcp/test_litellm_proxy_mcp_handler.py
+++ b/tests/test_litellm/responses/mcp/test_litellm_proxy_mcp_handler.py
@@ -605,3 +605,45 @@ def test_completion_with_function_tools_works_without_fastapi_installed():
timeout=120,
)
assert result.returncode == 0, result.stderr
+
+
+def test_extract_tool_call_details_reads_anthropic_tool_use_input():
+ """
+ Regression test (LIT-4517): an Anthropic tool_use block carries its arguments
+ under `input`, not `arguments`.
+
+ Given: A tool_use content block as /v1/messages returns it
+ When: The shared extractor reads it
+ Then: The arguments come back, so the MCP tool is called with them
+
+ Reading only `arguments` fails silently rather than loudly: _parse_tool_arguments
+ turns the resulting None into {}, so the tool still executes, just with every
+ argument dropped.
+ """
+ tool_use_block = {
+ "type": "tool_use",
+ "id": "toolu_01ABC",
+ "name": "read_wiki_structure",
+ "input": {"repoName": "BerriAI/litellm"},
+ }
+
+ name, arguments, call_id = LiteLLM_Proxy_MCP_Handler._extract_tool_call_details(tool_use_block)
+
+ assert name == "read_wiki_structure"
+ assert call_id == "toolu_01ABC"
+ assert arguments == {"repoName": "BerriAI/litellm"}
+ assert LiteLLM_Proxy_MCP_Handler._parse_tool_arguments(arguments) == {"repoName": "BerriAI/litellm"}
+
+
+def test_extract_tool_call_details_still_prefers_openai_arguments():
+ """The OpenAI chat shape must keep winning; `input` is only the fallback."""
+ openai_tool_call = {
+ "id": "call_123",
+ "function": {"name": "get_weather", "arguments": '{"city": "Paris"}'},
+ }
+
+ name, arguments, call_id = LiteLLM_Proxy_MCP_Handler._extract_tool_call_details(openai_tool_call)
+
+ assert name == "get_weather"
+ assert call_id == "call_123"
+ assert arguments == '{"city": "Paris"}'
diff --git a/ui/litellm-dashboard/src/app/(dashboard)/playground/components/chat_ui/ChatUI.tsx b/ui/litellm-dashboard/src/app/(dashboard)/playground/components/chat_ui/ChatUI.tsx
index d8f927b9a63..d2cf27e0c8b 100644
--- a/ui/litellm-dashboard/src/app/(dashboard)/playground/components/chat_ui/ChatUI.tsx
+++ b/ui/litellm-dashboard/src/app/(dashboard)/playground/components/chat_ui/ChatUI.tsx
@@ -98,7 +98,12 @@ interface ChatUIProps {
fixedModel?: string;
}
-const MCP_SUPPORTED_ENDPOINTS = new Set([EndpointType.CHAT, EndpointType.RESPONSES, EndpointType.MCP]);
+const MCP_SUPPORTED_ENDPOINTS = new Set([
+ EndpointType.CHAT,
+ EndpointType.RESPONSES,
+ EndpointType.MCP,
+ EndpointType.ANTHROPIC_MESSAGES,
+]);
const CUSTOM_MODEL_DEBOUNCE_WAIT_MS = 500;
@@ -870,8 +875,11 @@ const ChatUI: React.FC = ({
selectedVectorStores.length > 0 ? selectedVectorStores : undefined,
selectedGuardrails.length > 0 ? selectedGuardrails : undefined,
selectedPolicies.length > 0 ? selectedPolicies : undefined,
- selectedMCPServers, // Pass the selected tools array
+ selectedMCPServers,
customProxyBaseUrl || undefined,
+ mcpServers,
+ mcpServerToolRestrictions,
+ mcpToolsets,
);
} else if (endpointType === EndpointType.EMBEDDINGS) {
await makeOpenAIEmbeddingsRequest(
diff --git a/ui/litellm-dashboard/src/app/(dashboard)/playground/llm_calls/anthropic_messages.tsx b/ui/litellm-dashboard/src/app/(dashboard)/playground/llm_calls/anthropic_messages.tsx
index ed2b4280b79..4319315396a 100644
--- a/ui/litellm-dashboard/src/app/(dashboard)/playground/llm_calls/anthropic_messages.tsx
+++ b/ui/litellm-dashboard/src/app/(dashboard)/playground/llm_calls/anthropic_messages.tsx
@@ -1,6 +1,8 @@
import Anthropic from "@anthropic-ai/sdk";
import { MessageType } from "@/components/chat_ui/types";
import { TokenUsage } from "@/components/chat_ui/ResponseMetrics";
+import { buildMcpToolBlocks } from "@/components/llm_calls/mcp_tool_blocks";
+import { MCPServer, MCPToolset } from "@/components/mcp_tools/types";
import { getProxyBaseUrl } from "@/components/networking";
import NotificationManager from "@/components/molecules/notifications_manager";
@@ -18,8 +20,11 @@ export async function makeAnthropicMessagesRequest(
vector_store_ids?: string[],
guardrails?: string[],
policies?: string[],
- selectedMCPTools?: string[],
+ selectedMCPServers?: string[],
customBaseUrl?: string,
+ mcpServers?: MCPServer[],
+ mcpServerToolRestrictions?: Record,
+ mcpToolsets?: MCPToolset[],
) {
if (!accessToken) {
throw new Error("Virtual Key is required");
@@ -58,6 +63,13 @@ export async function makeAnthropicMessagesRequest(
litellm_trace_id: traceId,
};
+ const tools = buildMcpToolBlocks({
+ selectedMCPServers,
+ mcpServers,
+ mcpToolsets,
+ mcpServerToolRestrictions,
+ });
+ if (tools.length > 0) requestBody.tools = tools;
if (vector_store_ids) requestBody.vector_store_ids = vector_store_ids;
if (guardrails) requestBody.guardrails = guardrails;
if (policies) requestBody.policies = policies;
diff --git a/ui/litellm-dashboard/src/components/llm_calls/mcp_tool_blocks.ts b/ui/litellm-dashboard/src/components/llm_calls/mcp_tool_blocks.ts
new file mode 100644
index 00000000000..42fa94d8208
--- /dev/null
+++ b/ui/litellm-dashboard/src/components/llm_calls/mcp_tool_blocks.ts
@@ -0,0 +1,79 @@
+import { MCPServer, MCPToolset } from "@/components/mcp_tools/types";
+
+export const ALL_MCP_SERVERS_SENTINEL = "__all__";
+const TOOLSET_PREFIX = "toolset:";
+
+export interface McpToolBlock {
+ type: "mcp";
+ server_label: string;
+ server_url: string;
+ require_approval: "never";
+ allowed_tools?: string[];
+}
+
+export interface BuildMcpToolBlocksArgs {
+ selectedMCPServers?: string[];
+ mcpServers?: MCPServer[];
+ mcpToolsets?: MCPToolset[];
+ mcpServerToolRestrictions?: Record;
+}
+
+/**
+ * Build the litellm_proxy MCP reference blocks for a playground request.
+ *
+ * Every endpoint that supports MCP sends the same reference shape; the gateway
+ * expands it server side and each endpoint's own transformation decides the
+ * final tool shape. Keeping one builder here stops the endpoints from drifting
+ * apart on routing name, label uniqueness, or escaping.
+ *
+ * server_name is used for both routing and labelling because it is the unique
+ * registered identifier; aliases can collide across servers, and a duplicated
+ * server_label causes silent tool-routing failures.
+ */
+export function buildMcpToolBlocks({
+ selectedMCPServers,
+ mcpServers,
+ mcpToolsets,
+ mcpServerToolRestrictions,
+}: BuildMcpToolBlocksArgs): McpToolBlock[] {
+ if (!selectedMCPServers || selectedMCPServers.length === 0) {
+ return [];
+ }
+
+ if (selectedMCPServers.includes(ALL_MCP_SERVERS_SENTINEL)) {
+ return [
+ {
+ type: "mcp",
+ server_label: "litellm",
+ server_url: "litellm_proxy/mcp",
+ require_approval: "never",
+ },
+ ];
+ }
+
+ return selectedMCPServers.map((serverId) => {
+ if (serverId.startsWith(TOOLSET_PREFIX)) {
+ const toolsetId = serverId.slice(TOOLSET_PREFIX.length);
+ const toolset = mcpToolsets?.find((t) => t.toolset_id === toolsetId);
+ const toolsetName = toolset?.toolset_name || toolsetId;
+ return {
+ type: "mcp",
+ server_label: toolsetName,
+ server_url: `litellm_proxy/mcp/${encodeURIComponent(toolsetName)}`,
+ require_approval: "never",
+ };
+ }
+
+ const server = mcpServers?.find((s) => s.server_id === serverId);
+ const routeName = server?.server_name || serverId;
+ const allowedTools = mcpServerToolRestrictions?.[serverId] || [];
+
+ return {
+ type: "mcp",
+ server_label: routeName,
+ server_url: `litellm_proxy/mcp/${encodeURIComponent(routeName)}`,
+ require_approval: "never",
+ ...(allowedTools.length > 0 ? { allowed_tools: allowedTools } : {}),
+ };
+ });
+}
From cd3ac05a1fb6eedbbc078b38024f66693d3ef779 Mon Sep 17 00:00:00 2001
From: Tin Chi Lo
Date: Thu, 16 Jul 2026 19:00:41 -0700
Subject: [PATCH 006/296] fix(mcp): forward the caller's MCP credentials from
every gateway surface
The /v1/messages handler resolved only the auth object and the trace id, so tool
listing and tool execution ran without the caller's MCP auth headers. That fails
quietly rather than loudly: the tool still executes, just with no credentials, so
every server behind interactive OAuth, a bearer token or per-user env vars returns
nothing while the model reports it has no access. Only a no-auth server looks
healthy, which is exactly what the first proof used.
Threading the missing arguments would have left the real problem in place. Each
gateway surface rebuilds the same context by hand (responses/main.py twice,
chat_completions_handler, mcp_streaming_iterator), which is why a new surface
drops fields; this adds a fifth that dropped six of eight. Resolve it once into a
frozen MCPRequestContext and have the handlers take that, so a field cannot be
forgotten at a call site. chat_completions_handler now uses it too, and the
resolver reads user_api_key_auth from both metadata keys because
LITELLM_METADATA_ROUTES carry it in litellm_metadata while chat uses metadata.
Also stop the loop when every tool call was skipped. tool_results is empty then,
and the tool_result message built from it has empty content, which Anthropic
rejects; the caller saw a 400 from mid-loop instead of the model's own answer.
Tests pin both: dropping the headers from either listing or execution fails, and
so does removing the empty-results guard.
---
.../messages/mcp_handler.py | 38 +++---
.../responses/mcp/chat_completions_handler.py | 23 ++--
litellm/responses/mcp/request_context.py | 73 +++++++++++
.../messages/test_mcp_handler.py | 121 ++++++++++++++++++
4 files changed, 222 insertions(+), 33 deletions(-)
create mode 100644 litellm/responses/mcp/request_context.py
diff --git a/litellm/llms/anthropic/experimental_pass_through/messages/mcp_handler.py b/litellm/llms/anthropic/experimental_pass_through/messages/mcp_handler.py
index 392b9e2e02d..813d4a62089 100644
--- a/litellm/llms/anthropic/experimental_pass_through/messages/mcp_handler.py
+++ b/litellm/llms/anthropic/experimental_pass_through/messages/mcp_handler.py
@@ -10,6 +10,7 @@ tool through a ``tool_use`` content block, and results are fed back as
from typing import Any, AsyncIterator, Mapping, Sequence, Union
from litellm._logging import verbose_logger
+from litellm.responses.mcp.request_context import MCPRequestContext
from litellm.types.llms.anthropic import (
AnthropicMessagesTool,
AnthropicMessagesToolResultParam,
@@ -54,19 +55,6 @@ def _build_tool_result_message(tool_results: Sequence[Mapping[str, Any]]) -> Ant
)
-def _resolve_user_api_key_auth(
- kwargs: Mapping[str, Any],
-) -> Any: # any-ok: UserAPIKeyAuth is proxy-only, importing it here would create a cycle
- """`/v1/messages` is a LITELLM_METADATA_ROUTE, so the auth object rides in litellm_metadata."""
- litellm_metadata = kwargs.get("litellm_metadata") or {}
- metadata = kwargs.get("metadata") or {}
- return (
- kwargs.get("user_api_key_auth")
- or litellm_metadata.get("user_api_key_auth")
- or metadata.get("user_api_key_auth")
- )
-
-
async def anthropic_messages_with_mcp(
max_tokens: int,
messages: Sequence[Mapping[str, Any]],
@@ -101,15 +89,18 @@ async def anthropic_messages_with_mcp(
**kwargs,
)
- user_api_key_auth = _resolve_user_api_key_auth(kwargs)
+ context = MCPRequestContext.resolve(kwargs=dict(kwargs), tools=tools)
(
deduplicated_mcp_tools,
tool_server_map,
) = await LiteLLM_Proxy_MCP_Handler._process_mcp_tools_without_openai_transform(
- user_api_key_auth,
+ context.user_api_key_auth,
mcp_references,
- litellm_trace_id=kwargs.get("litellm_trace_id"),
+ litellm_trace_id=context.litellm_trace_id,
+ mcp_auth_header=context.mcp_auth_header,
+ mcp_server_auth_headers=context.mcp_server_auth_headers,
+ request_tags=list(context.request_tags) if context.request_tags else None,
)
anthropic_tools: Sequence[AnthropicMessagesTool] = tuple(
@@ -149,10 +140,21 @@ async def anthropic_messages_with_mcp(
tool_results = await LiteLLM_Proxy_MCP_Handler._execute_tool_calls(
tool_server_map=tool_server_map,
tool_calls=list(tool_use_blocks),
- user_api_key_auth=user_api_key_auth,
- litellm_trace_id=kwargs.get("litellm_trace_id"),
+ user_api_key_auth=context.user_api_key_auth,
+ mcp_auth_header=context.mcp_auth_header,
+ mcp_server_auth_headers=context.mcp_server_auth_headers,
+ oauth2_headers=context.oauth2_headers,
+ raw_headers=context.raw_headers,
+ litellm_call_id=context.litellm_call_id,
+ litellm_trace_id=context.litellm_trace_id,
+ request_tags=list(context.request_tags) if context.request_tags else None,
)
+ # Every tool call was skipped, so there is nothing to feed back; a
+ # tool_result message with empty content is rejected by Anthropic.
+ if not tool_results:
+ break
+
working_messages = (
*working_messages,
{"role": "assistant", "content": list(_get_response_content(response))},
diff --git a/litellm/responses/mcp/chat_completions_handler.py b/litellm/responses/mcp/chat_completions_handler.py
index f2ccfd430ae..5c3e0cf0902 100644
--- a/litellm/responses/mcp/chat_completions_handler.py
+++ b/litellm/responses/mcp/chat_completions_handler.py
@@ -12,7 +12,7 @@ from typing import (
from litellm.responses.mcp.litellm_proxy_mcp_handler import (
LiteLLM_Proxy_MCP_Handler,
)
-from litellm.responses.utils import ResponsesAPIRequestUtils
+from litellm.responses.mcp.request_context import MCPRequestContext
from litellm.types.utils import ModelResponse
from litellm.utils import CustomStreamWrapper
@@ -114,20 +114,13 @@ async def acompletion_with_mcp(
**kwargs,
)
- # Extract user_api_key_auth from metadata or kwargs
- user_api_key_auth = kwargs.get("user_api_key_auth") or ((kwargs.get("metadata", {}) or {}).get("user_api_key_auth"))
- request_tags = LiteLLM_Proxy_MCP_Handler._get_parent_request_tags(kwargs)
-
- # Extract MCP auth headers before fetching tools (needed for dynamic auth)
- (
- mcp_auth_header,
- mcp_server_auth_headers,
- oauth2_headers,
- raw_headers,
- ) = ResponsesAPIRequestUtils.extract_mcp_headers_from_request(
- secret_fields=kwargs.get("secret_fields"),
- tools=tools,
- )
+ context = MCPRequestContext.resolve(kwargs=kwargs, tools=tools)
+ user_api_key_auth = context.user_api_key_auth
+ request_tags = list(context.request_tags) if context.request_tags else None
+ mcp_auth_header = context.mcp_auth_header
+ mcp_server_auth_headers = context.mcp_server_auth_headers
+ oauth2_headers = context.oauth2_headers
+ raw_headers = context.raw_headers
# Process MCP tools (pass auth headers for dynamic auth)
(
diff --git a/litellm/responses/mcp/request_context.py b/litellm/responses/mcp/request_context.py
new file mode 100644
index 00000000000..fa03e677b39
--- /dev/null
+++ b/litellm/responses/mcp/request_context.py
@@ -0,0 +1,73 @@
+"""
+The per-request context an MCP gateway handler needs.
+
+Listing and executing MCP tools both need the caller's identity, their MCP auth
+headers, and the request's trace/tag identifiers. Every gateway surface resolves
+the same set from its own kwargs, so resolving it in one place keeps a new
+surface from silently dropping a field: omitting the auth headers, for instance,
+still executes the tool, just with no credentials.
+"""
+
+from dataclasses import dataclass
+from typing import Any, Iterable, Mapping, Sequence, Union
+
+
+@dataclass(frozen=True, slots=True)
+class MCPRequestContext:
+ """Everything a gateway handler must forward to MCP tool listing and execution."""
+
+ user_api_key_auth: Any # any-ok: UserAPIKeyAuth is proxy-only; importing it here would create a cycle
+ mcp_auth_header: Union[str, None] = None
+ mcp_server_auth_headers: Union[Mapping[str, Mapping[str, str]], None] = None
+ oauth2_headers: Union[Mapping[str, str], None] = None
+ raw_headers: Union[Mapping[str, str], None] = None
+ request_tags: Union[Sequence[str], None] = None
+ litellm_trace_id: Union[str, None] = None
+ litellm_call_id: Union[str, None] = None
+
+ @classmethod
+ def resolve(
+ cls,
+ kwargs: Mapping[str, Any],
+ tools: Union[Iterable[Any], None],
+ ) -> "MCPRequestContext":
+ """
+ Build the context from a gateway handler's kwargs.
+
+ ``user_api_key_auth`` is read from both metadata keys because routes differ:
+ LITELLM_METADATA_ROUTES (``/v1/messages``, ``/responses``) carry it in
+ ``litellm_metadata`` while ``/chat/completions`` uses ``metadata``.
+ """
+ from litellm.responses.mcp.litellm_proxy_mcp_handler import (
+ LiteLLM_Proxy_MCP_Handler,
+ )
+ from litellm.responses.utils import ResponsesAPIRequestUtils
+
+ litellm_metadata = kwargs.get("litellm_metadata") or {}
+ metadata = kwargs.get("metadata") or {}
+ user_api_key_auth = (
+ kwargs.get("user_api_key_auth")
+ or litellm_metadata.get("user_api_key_auth")
+ or metadata.get("user_api_key_auth")
+ )
+
+ (
+ mcp_auth_header,
+ mcp_server_auth_headers,
+ oauth2_headers,
+ raw_headers,
+ ) = ResponsesAPIRequestUtils.extract_mcp_headers_from_request(
+ secret_fields=kwargs.get("secret_fields"),
+ tools=tools,
+ )
+
+ return cls(
+ 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,
+ request_tags=LiteLLM_Proxy_MCP_Handler._get_parent_request_tags(dict(kwargs)),
+ litellm_trace_id=kwargs.get("litellm_trace_id"),
+ litellm_call_id=kwargs.get("litellm_call_id"),
+ )
diff --git a/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_mcp_handler.py b/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_mcp_handler.py
index 3faa6b1e4e2..060c3e459d0 100644
--- a/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_mcp_handler.py
+++ b/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_mcp_handler.py
@@ -120,3 +120,124 @@ def test_build_tool_result_message_uses_anthropic_tool_result_blocks():
assert list(message["content"]) == [
{"type": "tool_result", "tool_use_id": "toolu_1", "content": "9 sections"}
]
+
+
+@pytest.mark.asyncio
+async def test_anthropic_messages_with_mcp_forwards_the_callers_mcp_credentials():
+ """
+ Regression test (LIT-4517): the caller's MCP auth must reach both tool listing
+ and tool execution on /v1/messages.
+
+ Given: A request carrying MCP auth headers and request tags
+ When: The gateway lists and then executes an MCP tool
+ Then: Both calls receive the caller's credentials, tags and trace ids
+
+ Dropping them does not fail loudly; the tool still executes, just with no
+ credentials, so every auth-requiring MCP server (interactive OAuth, bearer
+ token, per-user env) silently returns nothing while the model claims it has
+ no access. Only a no-auth server would look healthy.
+ """
+ from litellm.llms.anthropic.experimental_pass_through.messages import mcp_handler
+ from litellm.responses.mcp.request_context import MCPRequestContext
+
+ context = MCPRequestContext(
+ user_api_key_auth="auth-object",
+ mcp_auth_header="legacy-header",
+ mcp_server_auth_headers={"deepwiki": {"authorization": "Bearer per-server"}},
+ oauth2_headers={"authorization": "Bearer oauth"},
+ raw_headers={"x-trace": "abc"},
+ request_tags=["team-a"],
+ litellm_trace_id="trace-123",
+ litellm_call_id="call-456",
+ )
+
+ process = AsyncMock(return_value=([], {}))
+ execute = AsyncMock(return_value=[{"tool_call_id": "toolu_1", "result": "ok", "name": "t"}])
+ responses = [
+ {"stop_reason": "tool_use", "content": [{"type": "tool_use", "id": "toolu_1", "name": "t", "input": {}}]},
+ {"stop_reason": "end_turn", "content": [{"type": "text", "text": "done"}]},
+ ]
+
+ with patch.object(MCPRequestContext, "resolve", return_value=context), patch.object(
+ mcp_handler.LiteLLM_Proxy_MCP_Handler
+ if hasattr(mcp_handler, "LiteLLM_Proxy_MCP_Handler")
+ else __import__(
+ "litellm.responses.mcp.litellm_proxy_mcp_handler", fromlist=["LiteLLM_Proxy_MCP_Handler"]
+ ).LiteLLM_Proxy_MCP_Handler,
+ "_process_mcp_tools_without_openai_transform",
+ new=process,
+ ), patch(
+ "litellm.responses.mcp.litellm_proxy_mcp_handler.LiteLLM_Proxy_MCP_Handler._execute_tool_calls",
+ new=execute,
+ ), patch(
+ "litellm.anthropic_messages", new=AsyncMock(side_effect=responses)
+ ):
+ await mcp_handler.anthropic_messages_with_mcp(
+ max_tokens=100,
+ messages=[{"role": "user", "content": "hi"}],
+ model="claude-sonnet-4-5",
+ tools=[MCP_REFERENCE],
+ )
+
+ listing = process.call_args.kwargs
+ assert listing["mcp_auth_header"] == "legacy-header", "tool listing must use the caller's MCP auth"
+ assert listing["mcp_server_auth_headers"] == {"deepwiki": {"authorization": "Bearer per-server"}}
+ assert listing["request_tags"] == ["team-a"]
+ assert listing["litellm_trace_id"] == "trace-123"
+
+ execution = execute.call_args.kwargs
+ assert execution["user_api_key_auth"] == "auth-object"
+ assert execution["mcp_auth_header"] == "legacy-header", "tool execution must use the caller's MCP auth"
+ assert execution["mcp_server_auth_headers"] == {"deepwiki": {"authorization": "Bearer per-server"}}
+ assert execution["oauth2_headers"] == {"authorization": "Bearer oauth"}
+ assert execution["raw_headers"] == {"x-trace": "abc"}
+ assert execution["litellm_call_id"] == "call-456"
+ assert execution["litellm_trace_id"] == "trace-123"
+ assert execution["request_tags"] == ["team-a"]
+
+
+@pytest.mark.asyncio
+async def test_anthropic_messages_with_mcp_stops_when_every_tool_call_is_skipped():
+ """
+ Regression test (LIT-4517): a tool_use turn whose calls all get skipped must
+ end the loop, not send an empty tool_result message.
+
+ Given: The model asks for a tool but the executor skips it (unresolvable name)
+ When: The gateway loop handles the empty result set
+ Then: It returns the last response instead of calling the model again
+
+ _build_tool_result_message([]) produces a user message with empty content, and
+ Anthropic rejects that, so the caller would get an unhandled 400 from the middle
+ of the loop rather than the model's own answer.
+ """
+ from litellm.llms.anthropic.experimental_pass_through.messages import mcp_handler
+ from litellm.responses.mcp.request_context import MCPRequestContext
+
+ tool_use_response = {
+ "stop_reason": "tool_use",
+ "content": [{"type": "tool_use", "id": "toolu_1", "name": "gone", "input": {}}],
+ }
+ anthropic_messages_mock = AsyncMock(return_value=tool_use_response)
+
+ with patch.object(
+ MCPRequestContext, "resolve", return_value=MCPRequestContext(user_api_key_auth="auth")
+ ), patch(
+ "litellm.responses.mcp.litellm_proxy_mcp_handler.LiteLLM_Proxy_MCP_Handler._process_mcp_tools_without_openai_transform",
+ new=AsyncMock(return_value=([], {})),
+ ), patch(
+ "litellm.responses.mcp.litellm_proxy_mcp_handler.LiteLLM_Proxy_MCP_Handler._execute_tool_calls",
+ new=AsyncMock(return_value=[]),
+ ), patch(
+ "litellm.anthropic_messages", new=anthropic_messages_mock
+ ):
+ result = await mcp_handler.anthropic_messages_with_mcp(
+ max_tokens=100,
+ messages=[{"role": "user", "content": "hi"}],
+ model="claude-sonnet-4-5",
+ tools=[MCP_REFERENCE],
+ )
+
+ assert anthropic_messages_mock.await_count == 1, (
+ "With no tool results there is nothing to send back, so the loop must not call the model again"
+ )
+ assert result == tool_use_response
From bea11ddedd414db960cbc57670e4370c08ef624b Mon Sep 17 00:00:00 2001
From: Shivam Rawat
Date: Thu, 16 Jul 2026 21:14:45 -0700
Subject: [PATCH 007/296] fix(proxy): resolve team wildcard credentials for
vector store files
Team-scoped wildcard deployments like openai/* are indexed separately from
global router models, so vector store file requests failed with api_key=None
when a team also had other yaml/db models. Pass team_id into credential
lookup and consult team model indexes and pattern routers.
Co-authored-by: Cursor
---
.../vector_store_files_endpoints/endpoints.py | 8 ++++++--
litellm/router.py | 17 ++++++++++++++++-
2 files changed, 22 insertions(+), 3 deletions(-)
diff --git a/litellm/proxy/vector_store_files_endpoints/endpoints.py b/litellm/proxy/vector_store_files_endpoints/endpoints.py
index 890db2f73a4..44935fc57c9 100644
--- a/litellm/proxy/vector_store_files_endpoints/endpoints.py
+++ b/litellm/proxy/vector_store_files_endpoints/endpoints.py
@@ -227,6 +227,8 @@ async def _update_request_data_with_model_routing_hint(
model_hint = data.get("model") or user_controlled_model_hint
should_authorize_model_hint = isinstance(model_hint, str) and model_hint == user_controlled_model_hint
+ caller_team_id = getattr(user_api_key_dict, "team_id", None) if user_api_key_dict else None
+
should_route = False
credentials = None
if isinstance(model_hint, str) and "*" in model_hint:
@@ -237,7 +239,9 @@ async def _update_request_data_with_model_routing_hint(
llm_router=llm_router,
user_api_key_dict=user_api_key_dict,
)
- credentials = llm_router.get_deployment_credentials_with_provider(model_id=model_hint)
+ credentials = llm_router.get_deployment_credentials_with_provider(
+ model_id=model_hint, team_id=caller_team_id
+ )
should_route = credentials is not None
else:
if isinstance(model_hint, str) and should_authorize_model_hint:
@@ -285,7 +289,7 @@ async def _update_request_data_with_model_routing_hint(
openai_credentials = None
for model_name in model_names_to_check:
- credentials = llm_router.get_deployment_credentials_with_provider(model_id=model_name)
+ credentials = llm_router.get_deployment_credentials_with_provider(model_id=model_name, team_id=caller_team_id)
if credentials is None:
continue
diff --git a/litellm/router.py b/litellm/router.py
index 78e156801f8..dbc6da106e7 100644
--- a/litellm/router.py
+++ b/litellm/router.py
@@ -8459,7 +8459,9 @@ class Router:
raise Exception("Model Name invalid - {}".format(type(model)))
return None
- def get_deployment_credentials_with_provider(self, model_id: str) -> Optional[Dict[str, Any]]:
+ def get_deployment_credentials_with_provider(
+ self, model_id: str, team_id: Optional[str] = None
+ ) -> Optional[Dict[str, Any]]:
"""
Get API credentials and provider info from a model name in model_list.
Useful for passthrough endpoints (files, batches, etc.) that need credentials.
@@ -8469,6 +8471,9 @@ class Router:
Args:
model_id: Model ID or model name from model_list (e.g., "gpt-4o-litellm")
+ team_id: Optional team id of the caller. When set, team-scoped
+ deployments (indexed by team public model name, including team
+ wildcard models like "openai/*") are also considered.
Returns:
Dictionary containing api_key, api_base, custom_llm_provider, etc.
@@ -8487,9 +8492,19 @@ class Router:
if deployment is None:
deployment = self.get_deployment_by_model_group_name(model_group_name=model_id)
+ # If not found, check team-scoped deployments (team public model names,
+ # e.g. team wildcard models like "openai/*", live in a separate index).
+ if deployment is None and team_id is not None:
+ team_indices = self.team_model_to_deployment_indices.get((team_id, model_id), [])
+ if team_indices:
+ team_model = self.model_list[team_indices[0]]
+ deployment = Deployment(**team_model) if isinstance(team_model, dict) else team_model
+
# If still not found, check for wildcard pattern matches
if deployment is None:
potential_wildcard_models = self.pattern_router.route(model_id) or []
+ if not potential_wildcard_models and team_id is not None and team_id in self.team_pattern_routers:
+ potential_wildcard_models = self.team_pattern_routers[team_id].route(model_id) or []
if potential_wildcard_models:
# Use the first matching wildcard deployment
deployment_dict = potential_wildcard_models[0]
From 56cda9f674d815a5f2686e29df9fb0b105a836f3 Mon Sep 17 00:00:00 2001
From: Tin Chi Lo
Date: Fri, 17 Jul 2026 10:33:28 -0700
Subject: [PATCH 008/296] fix(mcp): sanitize Anthropic tool schemas and stop
encoding gateway names
Two review findings, both a chat-vs-messages divergence.
transform_mcp_tool_to_anthropic_tool sent the MCP inputSchema to Anthropic almost
as-is, while the chat path (_map_tool_helper) coerces the type to object, inlines
legacy definitions with unpack_legacy_defs, and allow-lists keys to
AnthropicInputSchema. So a tool whose schema carried $schema, legacy definitions
or oneOf worked on /chat/completions and 400d on /v1/messages; a clean-schema
server hid it. Both paths now run the same sanitize_input_schema_for_anthropic,
extracted next to unpack_legacy_defs so they cannot drift again, and the chat
path is refactored onto it rather than keeping its own copy.
buildMcpToolBlocks percent-encoded the server and toolset names inside
litellm_proxy/mcp/... urls, but the gateway resolves the name with a raw
server_url.split("/")[-1] and never url-decodes, so a name with a space failed
lookup. The already-working chat path does not encode; the shared builder now
matches it.
Tests pin both: reverting the transform to the unfiltered schema fails, and
re-adding encodeURIComponent fails the builder test.
---
litellm/experimental_mcp_client/tools.py | 8 ++-
.../prompt_templates/common_utils.py | 26 ++++++++
litellm/llms/anthropic/chat/transformation.py | 29 ++-------
.../experimental_mcp_client/test_tools.py | 42 +++++++++++++
.../llm_calls/mcp_tool_blocks.test.ts | 63 +++++++++++++++++++
.../components/llm_calls/mcp_tool_blocks.ts | 8 ++-
6 files changed, 147 insertions(+), 29 deletions(-)
create mode 100644 ui/litellm-dashboard/src/components/llm_calls/mcp_tool_blocks.test.ts
diff --git a/litellm/experimental_mcp_client/tools.py b/litellm/experimental_mcp_client/tools.py
index 1bd65847616..500d226752b 100644
--- a/litellm/experimental_mcp_client/tools.py
+++ b/litellm/experimental_mcp_client/tools.py
@@ -9,7 +9,7 @@ from openai.types.chat import ChatCompletionToolParam
from openai.types.responses.function_tool_param import FunctionToolParam
from openai.types.shared_params.function_definition import FunctionDefinition
-from litellm.types.llms.anthropic import AnthropicInputSchema, AnthropicMessagesTool
+from litellm.types.llms.anthropic import AnthropicMessagesTool
from litellm.types.utils import ChatCompletionMessageToolCall
@@ -78,12 +78,14 @@ def transform_mcp_tool_to_openai_responses_api_tool(
def transform_mcp_tool_to_anthropic_tool(mcp_tool: MCPTool) -> AnthropicMessagesTool:
"""Convert an MCP tool to an Anthropic Messages API tool."""
- normalized_parameters = _normalize_mcp_input_schema(mcp_tool.inputSchema)
+ from litellm.litellm_core_utils.prompt_templates.common_utils import (
+ sanitize_input_schema_for_anthropic,
+ )
return AnthropicMessagesTool(
name=mcp_tool.name,
description=mcp_tool.description or "",
- input_schema=AnthropicInputSchema(**normalized_parameters),
+ input_schema=sanitize_input_schema_for_anthropic(mcp_tool.inputSchema),
type="custom",
)
diff --git a/litellm/litellm_core_utils/prompt_templates/common_utils.py b/litellm/litellm_core_utils/prompt_templates/common_utils.py
index 538d5f650ef..c43089950ee 100644
--- a/litellm/litellm_core_utils/prompt_templates/common_utils.py
+++ b/litellm/litellm_core_utils/prompt_templates/common_utils.py
@@ -42,6 +42,7 @@ from litellm.types.utils import (
)
if TYPE_CHECKING: # newer pattern to avoid importing pydantic objects on __init__.py
+ from litellm.types.llms.anthropic import AnthropicInputSchema
from litellm.types.llms.openai import ChatCompletionImageObject
DEFAULT_USER_CONTINUE_MESSAGE = ChatCompletionUserMessage(content="Please continue.", role="user")
@@ -1046,6 +1047,31 @@ def unpack_legacy_defs(
return schema
+def sanitize_input_schema_for_anthropic(input_schema: dict) -> "AnthropicInputSchema":
+ """Coerce an arbitrary tool input_schema into the shape Anthropic accepts.
+
+ Anthropic requires ``type == "object"``, only recognises ``$defs`` (legacy
+ ``definitions`` / OpenAPI ``components.schemas`` refs must be inlined first),
+ and rejects keys outside ``AnthropicInputSchema``. Both the chat
+ (``AnthropicConfig._map_tool_helper``) and Anthropic Messages MCP paths run
+ a schema through here so an external MCP schema cannot succeed on one route
+ and 400 on the other.
+ """
+ from litellm.types.llms.anthropic import AnthropicInputSchema
+
+ normalized = dict(input_schema) if input_schema else {}
+ if normalized.get("type") != "object":
+ normalized["type"] = "object"
+ if "properties" not in normalized:
+ normalized["properties"] = {}
+
+ normalized = unpack_legacy_defs(normalized, copy=True)
+
+ allowed_keys = set(AnthropicInputSchema.__annotations__.keys())
+ filtered = {key: value for key, value in normalized.items() if key in allowed_keys}
+ return AnthropicInputSchema(**filtered)
+
+
def _get_image_mime_type_from_url(url: str) -> Optional[str]:
"""
Get mime type for common image URLs
diff --git a/litellm/llms/anthropic/chat/transformation.py b/litellm/llms/anthropic/chat/transformation.py
index 0ec1f3eae13..5a0f274e3ca 100644
--- a/litellm/llms/anthropic/chat/transformation.py
+++ b/litellm/llms/anthropic/chat/transformation.py
@@ -29,7 +29,9 @@ from litellm.constants import (
RESPONSE_FORMAT_TOOL_NAME,
)
from litellm.litellm_core_utils.core_helpers import map_finish_reason
-from litellm.litellm_core_utils.prompt_templates.common_utils import unpack_legacy_defs
+from litellm.litellm_core_utils.prompt_templates.common_utils import (
+ sanitize_input_schema_for_anthropic,
+)
from litellm.llms.base_llm.base_utils import type_to_response_format_param
from litellm.llms.base_llm.chat.transformation import BaseConfig, BaseLLMException
from litellm.types.llms.anthropic import (
@@ -634,7 +636,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
mcp_server: Optional[AnthropicMcpServerTool] = None
if tool["type"] == "function" or tool["type"] == "custom":
- _input_schema: dict = tool["function"].get(
+ _input_schema = tool["function"].get(
"parameters",
{
"type": "object",
@@ -642,28 +644,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
},
)
- # Anthropic requires input_schema.type to be "object". Normalize
- # schemas from external sources (MCP servers, OpenAI callers) that
- # may omit the type field or use a non-object type.
- if _input_schema.get("type") != "object":
- litellm.verbose_logger.debug(
- "_map_tool_helper: coercing input_schema type from %r to "
- "'object' for Anthropic compatibility (tool: %s)",
- _input_schema.get("type"),
- tool["function"].get("name"),
- )
- _input_schema = dict(_input_schema) # avoid mutating caller's dict
- _input_schema["type"] = "object"
- if "properties" not in _input_schema:
- _input_schema["properties"] = {}
-
- # Inline legacy / OpenAPI $refs before the allow-list filter strips
- # their backing def blocks (https://github.com/BerriAI/litellm/issues/26692).
- _input_schema = unpack_legacy_defs(_input_schema, copy=True)
-
- _allowed_properties = set(AnthropicInputSchema.__annotations__.keys())
- input_schema_filtered = {k: v for k, v in _input_schema.items() if k in _allowed_properties}
- input_anthropic_schema: AnthropicInputSchema = AnthropicInputSchema(**input_schema_filtered)
+ input_anthropic_schema = sanitize_input_schema_for_anthropic(_input_schema)
_tool = AnthropicMessagesTool(
name=tool["function"]["name"],
diff --git a/tests/test_litellm/experimental_mcp_client/test_tools.py b/tests/test_litellm/experimental_mcp_client/test_tools.py
index 625ab56951f..804e99b6f4e 100644
--- a/tests/test_litellm/experimental_mcp_client/test_tools.py
+++ b/tests/test_litellm/experimental_mcp_client/test_tools.py
@@ -299,3 +299,45 @@ def test_transform_mcp_tool_to_anthropic_tool_normalizes_empty_schema():
assert anthropic_tool["description"] == ""
assert anthropic_tool["input_schema"]["type"] == "object"
assert anthropic_tool["input_schema"]["properties"] == {}
+
+
+def test_transform_mcp_tool_to_anthropic_tool_strips_keys_anthropic_rejects():
+ """
+ Regression test (LIT-4517): an MCP schema with keys Anthropic does not accept
+ must be sanitized, so the same tool cannot succeed on /chat/completions and 400
+ on /v1/messages.
+
+ Given: An MCP tool whose inputSchema carries $schema, legacy definitions and oneOf
+ When: It is transformed for the Anthropic Messages API
+ Then: Only keys in AnthropicInputSchema survive, matching the chat path
+
+ The chat path runs the schema through the same sanitizer, so before this the two
+ routes diverged: a clean-schema server (deepwiki) worked on both, but a server
+ with a richer schema would be rejected only on messages.
+ """
+ from litellm.types.llms.anthropic import AnthropicInputSchema
+
+ tool = MCPTool(
+ name="rich",
+ description="tool with a dirty schema",
+ inputSchema={
+ "type": "object",
+ "properties": {"q": {"type": "string"}},
+ "required": ["q"],
+ "$schema": "http://json-schema.org/draft-07/schema#",
+ "definitions": {"D": {"type": "string"}},
+ "oneOf": [{"required": ["q"]}],
+ },
+ )
+
+ anthropic_tool = transform_mcp_tool_to_anthropic_tool(tool)
+ schema_keys = set(anthropic_tool["input_schema"].keys())
+
+ assert schema_keys <= set(AnthropicInputSchema.__annotations__.keys()), (
+ f"schema must only carry keys Anthropic accepts, got {schema_keys}"
+ )
+ assert "$schema" not in schema_keys
+ assert "definitions" not in schema_keys
+ assert "oneOf" not in schema_keys
+ assert anthropic_tool["input_schema"]["properties"] == {"q": {"type": "string"}}
+ assert anthropic_tool["input_schema"]["required"] == ["q"]
diff --git a/ui/litellm-dashboard/src/components/llm_calls/mcp_tool_blocks.test.ts b/ui/litellm-dashboard/src/components/llm_calls/mcp_tool_blocks.test.ts
new file mode 100644
index 00000000000..62dd4d2631c
--- /dev/null
+++ b/ui/litellm-dashboard/src/components/llm_calls/mcp_tool_blocks.test.ts
@@ -0,0 +1,63 @@
+import { describe, it, expect } from "vitest";
+import { buildMcpToolBlocks } from "./mcp_tool_blocks";
+import { MCPServer, MCPToolset } from "@/components/mcp_tools/types";
+
+const server = (over: Partial): MCPServer =>
+ ({
+ server_id: "id-1",
+ server_name: "deepwiki",
+ alias: "wiki",
+ url: "",
+ transport: "http",
+ auth_type: "none",
+ ...over,
+ }) as any;
+
+describe("buildMcpToolBlocks", () => {
+ it("returns no blocks when nothing is selected", () => {
+ expect(buildMcpToolBlocks({ selectedMCPServers: [] })).toEqual([]);
+ expect(buildMcpToolBlocks({ selectedMCPServers: undefined })).toEqual([]);
+ });
+
+ it("routes by server_name, not alias, so colliding aliases cannot cross-route", () => {
+ const [block] = buildMcpToolBlocks({
+ selectedMCPServers: ["id-1"],
+ mcpServers: [server({})],
+ });
+ expect(block.server_url).toBe("litellm_proxy/mcp/deepwiki");
+ expect(block.server_label).toBe("deepwiki");
+ });
+
+ it("does not percent-encode the name; the gateway splits the raw path and never decodes", () => {
+ const [block] = buildMcpToolBlocks({
+ selectedMCPServers: ["id-1"],
+ mcpServers: [server({ server_name: "my server" }) as any],
+ });
+ expect(block.server_url).toBe("litellm_proxy/mcp/my server");
+ expect(block.server_url).not.toContain("%20");
+ });
+
+ it("passes per-server tool restrictions through as allowed_tools", () => {
+ const [block] = buildMcpToolBlocks({
+ selectedMCPServers: ["id-1"],
+ mcpServers: [server({})],
+ mcpServerToolRestrictions: { "id-1": ["read_wiki_structure"] },
+ });
+ expect(block.allowed_tools).toEqual(["read_wiki_structure"]);
+ });
+
+ it("collapses the all-servers sentinel to a single proxy-wide block", () => {
+ expect(buildMcpToolBlocks({ selectedMCPServers: ["__all__", "id-1"] })).toEqual([
+ { type: "mcp", server_label: "litellm", server_url: "litellm_proxy/mcp", require_approval: "never" },
+ ]);
+ });
+
+ it("routes a toolset by its name", () => {
+ const toolset = { toolset_id: "ts-1", toolset_name: "docs" } as MCPToolset;
+ const [block] = buildMcpToolBlocks({
+ selectedMCPServers: ["toolset:ts-1"],
+ mcpToolsets: [toolset],
+ });
+ expect(block.server_url).toBe("litellm_proxy/mcp/docs");
+ });
+});
diff --git a/ui/litellm-dashboard/src/components/llm_calls/mcp_tool_blocks.ts b/ui/litellm-dashboard/src/components/llm_calls/mcp_tool_blocks.ts
index 42fa94d8208..401d9fd9c84 100644
--- a/ui/litellm-dashboard/src/components/llm_calls/mcp_tool_blocks.ts
+++ b/ui/litellm-dashboard/src/components/llm_calls/mcp_tool_blocks.ts
@@ -29,6 +29,10 @@ export interface BuildMcpToolBlocksArgs {
* server_name is used for both routing and labelling because it is the unique
* registered identifier; aliases can collide across servers, and a duplicated
* server_label causes silent tool-routing failures.
+ *
+ * The name is not percent-encoded: the gateway resolves it with a raw
+ * `server_url.split("/")[-1]` and never url-decodes, so an encoded name would
+ * fail server lookup rather than round-trip.
*/
export function buildMcpToolBlocks({
selectedMCPServers,
@@ -59,7 +63,7 @@ export function buildMcpToolBlocks({
return {
type: "mcp",
server_label: toolsetName,
- server_url: `litellm_proxy/mcp/${encodeURIComponent(toolsetName)}`,
+ server_url: `litellm_proxy/mcp/${toolsetName}`,
require_approval: "never",
};
}
@@ -71,7 +75,7 @@ export function buildMcpToolBlocks({
return {
type: "mcp",
server_label: routeName,
- server_url: `litellm_proxy/mcp/${encodeURIComponent(routeName)}`,
+ server_url: `litellm_proxy/mcp/${routeName}`,
require_approval: "never",
...(allowedTools.length > 0 ? { allowed_tools: allowedTools } : {}),
};
From 31f293a9fc60a5ef7ff8c40ebfe4eab3fc5d11f2 Mon Sep 17 00:00:00 2001
From: mateo-berri <277851410+mateo-berri@users.noreply.github.com>
Date: Fri, 17 Jul 2026 13:49:02 -0400
Subject: [PATCH 009/296] feat(bedrock): forward bedrock_tags to
CreateModelInvocationJob for batch jobs
---
.../llms/bedrock/batches/transformation.py | 18 ++++
litellm/types/llms/bedrock.py | 7 +-
.../bedrock/batches/test_transformation.py | 87 +++++++++++++++++++
3 files changed, 111 insertions(+), 1 deletion(-)
diff --git a/litellm/llms/bedrock/batches/transformation.py b/litellm/llms/bedrock/batches/transformation.py
index 4fcf7cf91cb..8648d6586e8 100644
--- a/litellm/llms/bedrock/batches/transformation.py
+++ b/litellm/llms/bedrock/batches/transformation.py
@@ -4,6 +4,7 @@ import time
from typing import Any, Dict, List, Literal, Optional, Union, cast
from httpx import Headers, Response
+from pydantic import TypeAdapter, ValidationError
from litellm.litellm_core_utils.cloud_storage_security import (
BEDROCK_MANAGED_S3_BATCH_PREFIX,
@@ -19,6 +20,7 @@ from litellm.types.llms.bedrock import (
BedrockOutputDataConfig,
BedrockS3InputDataConfig,
BedrockS3OutputDataConfig,
+ BedrockTag,
)
from litellm.types.llms.openai import (
AllMessageValues,
@@ -38,6 +40,18 @@ _S3_BATCH_FILE_UUID_SUFFIX_PATTERN = re.compile(
r"-[0-9a-fA-F]{8}-[0-9a-fA-F]{4}-[0-9a-fA-F]{4}-[0-9a-fA-F]{4}-[0-9a-fA-F]{12}\.jsonl$"
)
+_BEDROCK_TAGS_ADAPTER: TypeAdapter[list[BedrockTag]] = TypeAdapter(list[BedrockTag])
+
+
+def _validate_bedrock_tags(raw_tags: object) -> list[BedrockTag]:
+ try:
+ return _BEDROCK_TAGS_ADAPTER.validate_python(raw_tags, strict=True)
+ except ValidationError as e:
+ raise ValueError(
+ "Invalid 'bedrock_tags' value. Expected a list of {'key': , 'value': } dicts, "
+ f"e.g. [{{'key': 'team', 'value': 'genai'}}]. Got: {raw_tags!r}"
+ ) from e
+
class BedrockBatchesConfig(BaseAWSLLM, BaseBatchesConfig):
"""
@@ -201,6 +215,10 @@ class BedrockBatchesConfig(BaseAWSLLM, BaseBatchesConfig):
"roleArn": role_arn,
}
+ bedrock_tags = litellm_params.get("bedrock_tags") or optional_params.get("bedrock_tags")
+ if bedrock_tags is not None:
+ bedrock_request["tags"] = _validate_bedrock_tags(bedrock_tags)
+
# Add optional parameters if provided
completion_window = create_batch_data.get("completion_window")
if completion_window:
diff --git a/litellm/types/llms/bedrock.py b/litellm/types/llms/bedrock.py
index bdf6b8fefed..d9f8229dbed 100644
--- a/litellm/types/llms/bedrock.py
+++ b/litellm/types/llms/bedrock.py
@@ -985,6 +985,11 @@ class BedrockOutputDataConfig(TypedDict):
s3OutputDataConfig: BedrockS3OutputDataConfig
+class BedrockTag(TypedDict):
+ key: str
+ value: str
+
+
class BedrockCreateBatchRequest(TypedDict, total=False):
"""
Request structure for creating a Bedrock batch inference job.
@@ -999,7 +1004,7 @@ class BedrockCreateBatchRequest(TypedDict, total=False):
outputDataConfig: BedrockOutputDataConfig
timeoutDurationInHours: Optional[int]
clientRequestToken: Optional[str]
- tags: Optional[List[dict]]
+ tags: Optional[List[BedrockTag]]
BedrockBatchJobStatus = Literal["Submitted", "InProgress", "Completed", "Failed", "Stopping", "Stopped"]
diff --git a/tests/test_litellm/llms/bedrock/batches/test_transformation.py b/tests/test_litellm/llms/bedrock/batches/test_transformation.py
index d1ad5943ae6..b38d271e210 100644
--- a/tests/test_litellm/llms/bedrock/batches/test_transformation.py
+++ b/tests/test_litellm/llms/bedrock/batches/test_transformation.py
@@ -258,6 +258,93 @@ def test_create_request_no_timeout_for_non_24h_window(config):
assert "timeoutDurationInHours" not in mock_sign.call_args.kwargs["data"]
+def test_create_request_forwards_bedrock_tags_from_litellm_params(config):
+ tags = [
+ {"key": "application", "value": "genai-proxy"},
+ {"key": "team", "value": "ml-platform"},
+ ]
+ with patch.object(
+ config.common_utils,
+ "generate_unique_job_name",
+ return_value="litellm-batch-1",
+ ), patch.object(config.common_utils, "sign_aws_request") as mock_sign:
+ mock_sign.return_value = ({}, b"{}")
+ config.transform_create_batch_request(
+ model="m",
+ create_batch_data={"input_file_id": "s3://b/in.jsonl"},
+ optional_params={},
+ litellm_params={
+ "aws_batch_role_arn": "arn:aws:iam::1:role/r",
+ "bedrock_tags": tags,
+ },
+ )
+ assert mock_sign.call_args.kwargs["data"]["tags"] == tags
+
+
+def test_create_request_forwards_bedrock_tags_from_optional_params(config):
+ tags = [{"key": "env", "value": "prod"}]
+ with patch.object(
+ config.common_utils,
+ "generate_unique_job_name",
+ return_value="litellm-batch-1",
+ ), patch.object(config.common_utils, "sign_aws_request") as mock_sign:
+ mock_sign.return_value = ({}, b"{}")
+ config.transform_create_batch_request(
+ model="m",
+ create_batch_data={"input_file_id": "s3://b/in.jsonl"},
+ optional_params={"bedrock_tags": tags},
+ litellm_params={"aws_batch_role_arn": "arn:aws:iam::1:role/r"},
+ )
+ assert mock_sign.call_args.kwargs["data"]["tags"] == tags
+
+
+def test_create_request_omits_tags_when_bedrock_tags_absent(config):
+ with patch.object(
+ config.common_utils,
+ "generate_unique_job_name",
+ return_value="litellm-batch-1",
+ ), patch.object(config.common_utils, "sign_aws_request") as mock_sign:
+ mock_sign.return_value = ({}, b"{}")
+ config.transform_create_batch_request(
+ model="m",
+ create_batch_data={"input_file_id": "s3://b/in.jsonl"},
+ optional_params={},
+ litellm_params={"aws_batch_role_arn": "arn:aws:iam::1:role/r"},
+ )
+ assert "tags" not in mock_sign.call_args.kwargs["data"]
+
+
+@pytest.mark.parametrize(
+ "bad_tags",
+ [
+ ["application=genai-proxy"],
+ [{"key": "application"}],
+ [{"value": "genai-proxy"}],
+ [{"key": "application", "value": 42}],
+ {"key": "application", "value": "genai-proxy"},
+ "application=genai-proxy",
+ ],
+)
+def test_create_request_rejects_malformed_bedrock_tags(config, bad_tags):
+ with patch.object(
+ config.common_utils,
+ "generate_unique_job_name",
+ return_value="litellm-batch-1",
+ ), patch.object(config.common_utils, "sign_aws_request") as mock_sign:
+ mock_sign.return_value = ({}, b"{}")
+ with pytest.raises(ValueError, match="Invalid 'bedrock_tags' value"):
+ config.transform_create_batch_request(
+ model="m",
+ create_batch_data={"input_file_id": "s3://b/in.jsonl"},
+ optional_params={},
+ litellm_params={
+ "aws_batch_role_arn": "arn:aws:iam::1:role/r",
+ "bedrock_tags": bad_tags,
+ },
+ )
+ mock_sign.assert_not_called()
+
+
# --------------------------------------------------------------------------- #
# transform_create_batch_response - status mapping + LiteLLMBatch shape
# --------------------------------------------------------------------------- #
From 12919628501340c8b7b596d33492bd9f5ef6eff0 Mon Sep 17 00:00:00 2001
From: Tin Chi Lo
Date: Thu, 16 Jul 2026 13:23:41 -0700
Subject: [PATCH 010/296] feat(ui): configure Anthropic automatic prompt
caching from the Admin UI
Register enable_anthropic_prompt_caching and anthropic_prompt_caching_ttl on the
General Settings table so caching can be turned on without hand-writing config.
The registry could not express either field: validation was hardcoded to a float in
(0, 1], reset set every field to None (not a bool for a boolean flag), and the listing
reported any non-None value as 'In Config', which a False default would always trip.
Validation now dispatches on the declared type and reset restores each field's own
default. ConfigList carries field_options so the table can render a Select for enums
instead of no editor at all.
---
litellm/proxy/_types.py | 3 +-
litellm/proxy/proxy_server.py | 96 +++++++--
tests/test_litellm/proxy/test_proxy_server.py | 204 ++++++++++++++++++
.../_components/general_settings.tsx | 15 +-
ui/litellm-dashboard/src/lib/http/schema.d.ts | 2 +
5 files changed, 300 insertions(+), 20 deletions(-)
diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py
index d102c1d1e37..e07c7b9ae78 100644
--- a/litellm/proxy/_types.py
+++ b/litellm/proxy/_types.py
@@ -1011,10 +1011,10 @@ class LiteLLM_ObjectPermissionBase(LiteLLMPydanticObjectBase):
mcp_tool_search_enabled: Optional[bool] = None
+from litellm.models.team import BudgetLimitEntry as BudgetLimitEntry # noqa: E402
from litellm.types.object_permission import ( # noqa: E402
ObjectPermissionDict as ObjectPermissionDict,
)
-from litellm.models.team import BudgetLimitEntry as BudgetLimitEntry # noqa: E402
class GenerateRequestBase(LiteLLMPydanticObjectBase):
@@ -2122,6 +2122,7 @@ class ConfigList(LiteLLMPydanticObjectBase):
field_default_value: Any
premium_field: bool = False
nested_fields: Optional[List[FieldDetail]] = None # For nested dictionary or Pydantic fields
+ field_options: Optional[list[str]] = None # Allowed values, for field_type == "Select"
class UserHeaderMapping(LiteLLMPydanticObjectBase):
diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py
index dbdfdd5fdd3..ae91eb70427 100644
--- a/litellm/proxy/proxy_server.py
+++ b/litellm/proxy/proxy_server.py
@@ -28,6 +28,7 @@ from typing import (
Optional,
Set,
Tuple,
+ TypedDict,
Union,
cast,
get_args,
@@ -39,6 +40,7 @@ import anyio
import websockets
import websockets.exceptions
from pydantic import BaseModel, Json, JsonValue
+from typing_extensions import NotRequired, assert_never
from litellm._uuid import uuid
from litellm.constants import (
@@ -363,15 +365,15 @@ from litellm.proxy.management_endpoints.cache_settings_endpoints import (
from litellm.proxy.management_endpoints.callback_management_endpoints import (
router as callback_management_endpoints_router,
)
-from litellm.proxy.management_endpoints.coordination_redis_endpoints import (
- get_persisted_coordination_redis_settings,
- router as coordination_redis_settings_router,
-)
from litellm.proxy.management_endpoints.common_utils import (
_user_has_admin_privileges,
_user_has_admin_view,
admin_can_invite_user,
)
+from litellm.proxy.management_endpoints.coordination_redis_endpoints import (
+ get_persisted_coordination_redis_settings,
+ router as coordination_redis_settings_router,
+)
from litellm.proxy.management_endpoints.cost_tracking_settings import (
router as cost_tracking_settings_router,
)
@@ -14800,7 +14802,16 @@ async def get_config_general_settings(
)
-_GENERAL_SETTINGS_UI_LITELLM_FIELDS: dict[str, dict[str, str]] = {
+GeneralSettingsUILiteLLMValue = Union[float, bool, str, None]
+
+
+class GeneralSettingsUILiteLLMFieldSpec(TypedDict):
+ type: Literal["Float", "Boolean", "Select"]
+ description: str
+ options: NotRequired[tuple[str, ...]]
+
+
+_GENERAL_SETTINGS_UI_LITELLM_FIELDS: dict[str, GeneralSettingsUILiteLLMFieldSpec] = {
"budget_exceeded_throttle_percentage": {
"type": "Float",
"description": (
@@ -14809,18 +14820,64 @@ _GENERAL_SETTINGS_UI_LITELLM_FIELDS: dict[str, dict[str, str]] = {
"over-budget keys."
),
},
+ "enable_anthropic_prompt_caching": {
+ "type": "Boolean",
+ "description": (
+ "Automatically add Anthropic cache_control breakpoints to the system prompt and the "
+ "trailing turn, for Anthropic and Bedrock Claude models that support prompt caching. "
+ "Lets clients that never set cache_control themselves still get cached prompts. "
+ "Requests that already carry their own cache_control are left untouched."
+ ),
+ },
+ "anthropic_prompt_caching_ttl": {
+ "type": "Select",
+ "options": ("5m", "1h"),
+ "description": (
+ "Cache lifetime for the breakpoints added by 'enable_anthropic_prompt_caching'. "
+ "Leave empty for Anthropic's 5m default. 1h suits long agentic sessions but doubles "
+ "the cache write premium."
+ ),
+ },
}
-def _validate_general_settings_ui_litellm_value(field_name: str, value: Any) -> Optional[float]:
+def _general_settings_ui_litellm_default(
+ field_type: Literal["Float", "Boolean", "Select"],
+) -> GeneralSettingsUILiteLLMValue:
+ """The value a field falls back to when it is cleared or reset."""
+ return False if field_type == "Boolean" else None
+
+
+def _validate_general_settings_ui_litellm_value(field_name: str, value: Any) -> GeneralSettingsUILiteLLMValue:
+ spec = _GENERAL_SETTINGS_UI_LITELLM_FIELDS[field_name]
+ field_type = spec["type"]
if value is None or value == "":
- return None
- if isinstance(value, bool) or not isinstance(value, (int, float)) or not (0 < float(value) <= 1):
- raise HTTPException(
- status_code=400,
- detail={"error": f"{field_name} must be a number in (0, 1] or empty"},
- )
- return float(value)
+ return _general_settings_ui_litellm_default(field_type)
+ match field_type:
+ case "Boolean":
+ if not isinstance(value, bool):
+ raise HTTPException(
+ status_code=400,
+ detail={"error": f"{field_name} must be true or false"},
+ )
+ return value
+ case "Select":
+ options = spec.get("options", ())
+ if value not in options:
+ raise HTTPException(
+ status_code=400,
+ detail={"error": f"{field_name} must be one of: {', '.join(options)}, or empty"},
+ )
+ return cast(str, value) # cast-ok: membership in options proves it is one of the option strings
+ case "Float":
+ if isinstance(value, bool) or not isinstance(value, (int, float)) or not (0 < float(value) <= 1):
+ raise HTTPException(
+ status_code=400,
+ detail={"error": f"{field_name} must be a number in (0, 1] or empty"},
+ )
+ return float(value)
+ case _:
+ assert_never(field_type)
async def _persist_general_settings_ui_litellm_field(
@@ -14841,11 +14898,12 @@ async def _persist_general_settings_ui_litellm_field(
async def _reset_general_settings_ui_litellm_field(field_name: str, user_api_key_dict: UserAPIKeyAuth) -> dict:
config = await proxy_config.get_config()
before_value = config.get("litellm_settings", {}).get(field_name)
- setattr(litellm, field_name, None)
+ default_value = _general_settings_ui_litellm_default(_GENERAL_SETTINGS_UI_LITELLM_FIELDS[field_name]["type"])
+ setattr(litellm, field_name, default_value)
if "litellm_settings" in config:
config["litellm_settings"].pop(field_name, None)
await proxy_config.save_config(new_config=config)
- asyncio.create_task(create_config_audit_log(field_name, "deleted", before_value, None, user_api_key_dict))
+ asyncio.create_task(create_config_audit_log(field_name, "deleted", before_value, default_value, user_api_key_dict))
return {"message": f"Field {field_name} reset", "status": "success"}
@@ -15013,11 +15071,12 @@ async def get_config_list(
else {}
)
for litellm_field_name, spec in _GENERAL_SETTINGS_UI_LITELLM_FIELDS.items():
- current_value: Optional[float] = getattr(litellm, litellm_field_name, None)
+ current_value: GeneralSettingsUILiteLLMValue = getattr(litellm, litellm_field_name, None)
+ default_value = _general_settings_ui_litellm_default(spec["type"])
stored_in_db_litellm: Optional[bool]
if litellm_field_name in db_litellm_settings:
stored_in_db_litellm = True
- elif current_value is not None:
+ elif current_value != default_value:
stored_in_db_litellm = False
else:
stored_in_db_litellm = None
@@ -15028,7 +15087,8 @@ async def get_config_list(
field_description=spec["description"],
field_value=current_value,
stored_in_db=stored_in_db_litellm,
- field_default_value=None,
+ field_default_value=default_value,
+ field_options=list(spec.get("options", ())) or None,
nested_fields=None,
)
)
diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py
index 54db0c0fd4f..86e807447c1 100644
--- a/tests/test_litellm/proxy/test_proxy_server.py
+++ b/tests/test_litellm/proxy/test_proxy_server.py
@@ -8984,6 +8984,210 @@ async def test_update_config_field_throttle_persists_to_litellm_settings(monkeyp
assert saved["litellm_settings"]["budget_exceeded_throttle_percentage"] == 0.1
+def test_get_config_list_includes_anthropic_prompt_caching_fields(monkeypatch):
+ """The auto prompt caching flag and its ttl are litellm_settings globals surfaced on the
+ General Settings table, so an admin can turn caching on without hand-writing config. The
+ ttl is a Select and must ship its allowed values, or the table renders no editor for it."""
+ import types
+ from unittest.mock import AsyncMock, MagicMock
+
+ from fastapi.testclient import TestClient
+
+ import litellm.proxy.proxy_server as ps
+ from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
+ from litellm.proxy.proxy_server import app
+
+ mock_prisma = MagicMock()
+ mock_config_table = MagicMock()
+ mock_config_table.find_first = AsyncMock(return_value=None)
+ mock_prisma.db = types.SimpleNamespace(litellm_config=mock_config_table)
+ monkeypatch.setattr(ps, "prisma_client", mock_prisma)
+ monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", True)
+ monkeypatch.setattr(litellm, "anthropic_prompt_caching_ttl", "1h")
+ app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth(
+ user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN
+ )
+ try:
+ client = TestClient(app)
+ resp = client.get("/config/list", params={"config_type": "general_settings"})
+ assert resp.status_code == 200, resp.text
+ fields = {item["field_name"]: item for item in resp.json()}
+
+ assert fields["enable_anthropic_prompt_caching"]["field_type"] == "Boolean"
+ assert fields["enable_anthropic_prompt_caching"]["field_value"] is True
+
+ assert fields["anthropic_prompt_caching_ttl"]["field_type"] == "Select"
+ assert fields["anthropic_prompt_caching_ttl"]["field_value"] == "1h"
+ assert fields["anthropic_prompt_caching_ttl"]["field_options"] == ["5m", "1h"]
+ finally:
+ app.dependency_overrides.clear()
+
+
+def test_get_config_list_marks_untouched_prompt_caching_flag_as_not_set(monkeypatch):
+ """The flag defaults to False rather than None, so a plain 'is not None' check would
+ report the default as 'In Config' and imply an admin had set it."""
+ import types
+ from unittest.mock import AsyncMock, MagicMock
+
+ from fastapi.testclient import TestClient
+
+ import litellm.proxy.proxy_server as ps
+ from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
+ from litellm.proxy.proxy_server import app
+
+ mock_prisma = MagicMock()
+ mock_config_table = MagicMock()
+ mock_config_table.find_first = AsyncMock(return_value=None)
+ mock_prisma.db = types.SimpleNamespace(litellm_config=mock_config_table)
+ monkeypatch.setattr(ps, "prisma_client", mock_prisma)
+ monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", False)
+ app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth(
+ user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN
+ )
+ try:
+ client = TestClient(app)
+ resp = client.get("/config/list", params={"config_type": "general_settings"})
+ fields = {item["field_name"]: item for item in resp.json()}
+ assert fields["enable_anthropic_prompt_caching"]["stored_in_db"] is None
+ finally:
+ app.dependency_overrides.clear()
+
+
+@pytest.mark.parametrize(
+ "field_name, field_value",
+ [
+ ("enable_anthropic_prompt_caching", True),
+ ("enable_anthropic_prompt_caching", False),
+ ("anthropic_prompt_caching_ttl", "5m"),
+ ("anthropic_prompt_caching_ttl", "1h"),
+ ],
+)
+@pytest.mark.asyncio
+async def test_update_config_field_prompt_caching_persists_to_litellm_settings(monkeypatch, field_name, field_value):
+ """Toggling either row must set litellm. live and persist under litellm_settings,
+ so the running proxy caches immediately and still does after a restart."""
+ from unittest.mock import MagicMock
+
+ import litellm.proxy.proxy_server as ps
+ from litellm.proxy._types import (
+ ConfigFieldUpdate,
+ LitellmUserRoles,
+ UserAPIKeyAuth,
+ )
+ from litellm.proxy.proxy_server import update_config_general_settings
+
+ saved: dict = {}
+
+ async def fake_get_config():
+ return {"litellm_settings": {}}
+
+ async def fake_save_config(new_config=None):
+ saved.update(new_config or {})
+
+ monkeypatch.setattr(ps.proxy_config, "get_config", fake_get_config)
+ monkeypatch.setattr(ps.proxy_config, "save_config", fake_save_config)
+ monkeypatch.setattr(ps, "prisma_client", MagicMock())
+ monkeypatch.setattr(litellm, "store_audit_logs", False)
+ monkeypatch.setattr(litellm, field_name, None)
+
+ admin = UserAPIKeyAuth(api_key="k", user_id="a", user_role=LitellmUserRoles.PROXY_ADMIN)
+ await update_config_general_settings(
+ data=ConfigFieldUpdate(field_name=field_name, field_value=field_value, config_type="general_settings"),
+ user_api_key_dict=admin,
+ )
+
+ assert getattr(litellm, field_name) == field_value
+ assert saved["litellm_settings"][field_name] == field_value
+
+
+@pytest.mark.parametrize(
+ "field_name, bad_value",
+ [
+ ("enable_anthropic_prompt_caching", "yes"),
+ ("enable_anthropic_prompt_caching", 1),
+ ("anthropic_prompt_caching_ttl", "10m"),
+ ("anthropic_prompt_caching_ttl", "1H"),
+ ("anthropic_prompt_caching_ttl", 3600),
+ ],
+)
+@pytest.mark.asyncio
+async def test_update_config_field_prompt_caching_rejects_invalid(monkeypatch, field_name, bad_value):
+ """An unsupported ttl must be refused here rather than reaching Anthropic verbatim."""
+ from unittest.mock import MagicMock
+
+ from fastapi import HTTPException
+
+ import litellm.proxy.proxy_server as ps
+ from litellm.proxy._types import (
+ ConfigFieldUpdate,
+ LitellmUserRoles,
+ UserAPIKeyAuth,
+ )
+ from litellm.proxy.proxy_server import update_config_general_settings
+
+ async def fake_get_config():
+ return {"litellm_settings": {}}
+
+ monkeypatch.setattr(ps.proxy_config, "get_config", fake_get_config)
+ monkeypatch.setattr(ps, "prisma_client", MagicMock())
+ monkeypatch.setattr(litellm, field_name, None)
+
+ admin = UserAPIKeyAuth(api_key="k", user_id="a", user_role=LitellmUserRoles.PROXY_ADMIN)
+ with pytest.raises(HTTPException) as exc:
+ await update_config_general_settings(
+ data=ConfigFieldUpdate(field_name=field_name, field_value=bad_value, config_type="general_settings"),
+ user_api_key_dict=admin,
+ )
+ assert exc.value.status_code == 400
+ assert getattr(litellm, field_name) is None
+
+
+@pytest.mark.parametrize(
+ "field_name, expected_default",
+ [
+ ("enable_anthropic_prompt_caching", False),
+ ("anthropic_prompt_caching_ttl", None),
+ ("budget_exceeded_throttle_percentage", None),
+ ],
+)
+@pytest.mark.asyncio
+async def test_reset_config_field_restores_type_default(monkeypatch, field_name, expected_default):
+ """Reset must restore each field's own default. Blanket None would leave the boolean flag
+ set to None, which is not a bool and would read as neither on nor off."""
+ from unittest.mock import MagicMock
+
+ import litellm.proxy.proxy_server as ps
+ from litellm.proxy._types import (
+ ConfigFieldDelete,
+ LitellmUserRoles,
+ UserAPIKeyAuth,
+ )
+ from litellm.proxy.proxy_server import delete_config_general_settings
+
+ saved: dict = {}
+
+ async def fake_get_config():
+ return {"litellm_settings": {field_name: "stale"}}
+
+ async def fake_save_config(new_config=None):
+ saved.update(new_config or {})
+
+ monkeypatch.setattr(ps.proxy_config, "get_config", fake_get_config)
+ monkeypatch.setattr(ps.proxy_config, "save_config", fake_save_config)
+ monkeypatch.setattr(ps, "prisma_client", MagicMock())
+ monkeypatch.setattr(litellm, "store_audit_logs", False)
+ monkeypatch.setattr(litellm, field_name, "stale")
+
+ admin = UserAPIKeyAuth(api_key="k", user_id="a", user_role=LitellmUserRoles.PROXY_ADMIN)
+ await delete_config_general_settings(
+ data=ConfigFieldDelete(field_name=field_name, config_type="general_settings"),
+ user_api_key_dict=admin,
+ )
+
+ assert getattr(litellm, field_name) is expected_default
+ assert field_name not in saved["litellm_settings"]
+
+
@pytest.mark.parametrize("bad_value", [0, -0.1, 1.5, True])
@pytest.mark.asyncio
async def test_update_config_field_throttle_rejects_invalid(monkeypatch, bad_value):
diff --git a/ui/litellm-dashboard/src/app/(dashboard)/router-settings/_components/general_settings.tsx b/ui/litellm-dashboard/src/app/(dashboard)/router-settings/_components/general_settings.tsx
index 3955e80f5e9..a1ef8252afa 100644
--- a/ui/litellm-dashboard/src/app/(dashboard)/router-settings/_components/general_settings.tsx
+++ b/ui/litellm-dashboard/src/app/(dashboard)/router-settings/_components/general_settings.tsx
@@ -14,7 +14,7 @@ import {
} from "@tremor/react";
import { TabPanel, TabPanels, TabGroup, TabList, Tab } from "@tremor/react";
import { getGeneralSettingsCall, updateConfigFieldSetting, deleteConfigFieldSetting } from "@/components/networking";
-import { InputNumber } from "antd";
+import { InputNumber, Select as AntdSelect } from "antd";
import { TrashIcon } from "@heroicons/react/outline";
import { StatusBadge } from "@/components/shared/table_cells";
@@ -33,6 +33,7 @@ interface generalSettingsItem {
field_value: any;
field_description: string;
stored_in_db: boolean | null;
+ field_options?: string[] | null;
}
const GeneralSettings: React.FC = ({ accessToken, userRole, userID }) => {
@@ -169,6 +170,18 @@ const GeneralSettings: React.FC = ({ accessToken, user
value={value.field_value}
onChange={(newValue) => handleInputChange(value.field_name, newValue)}
/>
+ ) : value.field_type == "Select" ? (
+ ({
+ label: option,
+ value: option,
+ }))}
+ onChange={(newValue) => handleInputChange(value.field_name, newValue ?? "")}
+ />
) : null}
diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts
index 0d8f55164f9..45b06c9e44f 100644
--- a/ui/litellm-dashboard/src/lib/http/schema.d.ts
+++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts
@@ -22711,6 +22711,8 @@ export interface components {
field_description: string;
/** Field Name */
field_name: string;
+ /** Field Options */
+ field_options?: string[] | null;
/** Field Type */
field_type: string;
/** Field Value */
From 9f7f53a82a938b0471e41c89be967d49b3275434 Mon Sep 17 00:00:00 2001
From: Tin Chi Lo
Date: Thu, 16 Jul 2026 13:28:34 -0700
Subject: [PATCH 011/296] refactor(ui): extract the General Settings value
editor into a component
The value cell was a ternary chain over field_type; adding Select made it a fourth
level and tripped no-nested-ternary. Early returns read better than a deeper chain
and let the suppression baseline ratchet down.
---
ui/litellm-dashboard/eslint-suppressions.json | 2 +-
.../_components/general_settings.tsx | 80 +++++++++++--------
2 files changed, 49 insertions(+), 33 deletions(-)
diff --git a/ui/litellm-dashboard/eslint-suppressions.json b/ui/litellm-dashboard/eslint-suppressions.json
index 32e9a03da95..90b0c84244e 100644
--- a/ui/litellm-dashboard/eslint-suppressions.json
+++ b/ui/litellm-dashboard/eslint-suppressions.json
@@ -1079,7 +1079,7 @@
},
"src/app/(dashboard)/router-settings/_components/general_settings.tsx": {
"no-nested-ternary": {
- "count": 3
+ "count": 1
},
"no-restricted-imports": {
"count": 2
diff --git a/ui/litellm-dashboard/src/app/(dashboard)/router-settings/_components/general_settings.tsx b/ui/litellm-dashboard/src/app/(dashboard)/router-settings/_components/general_settings.tsx
index a1ef8252afa..af6bbdde8b9 100644
--- a/ui/litellm-dashboard/src/app/(dashboard)/router-settings/_components/general_settings.tsx
+++ b/ui/litellm-dashboard/src/app/(dashboard)/router-settings/_components/general_settings.tsx
@@ -36,6 +36,53 @@ interface generalSettingsItem {
field_options?: string[] | null;
}
+const SettingValueEditor: React.FC<{
+ setting: generalSettingsItem;
+ onChange: (fieldName: string, newValue: any) => void;
+}> = ({ setting, onChange }) => {
+ if (setting.field_type === "Integer") {
+ return (
+ onChange(setting.field_name, newValue)}
+ />
+ );
+ }
+ if (setting.field_type === "Boolean") {
+ return (
+ onChange(setting.field_name, checked)}
+ />
+ );
+ }
+ if (setting.field_type === "Float") {
+ return (
+ onChange(setting.field_name, newValue)}
+ />
+ );
+ }
+ if (setting.field_type === "Select") {
+ return (
+ ({ label: option, value: option }))}
+ onChange={(newValue) => onChange(setting.field_name, newValue ?? "")}
+ />
+ );
+ }
+ return null;
+};
+
const GeneralSettings: React.FC = ({ accessToken, userRole, userID }) => {
const [generalSettings, setGeneralSettings] = useState([]);
@@ -151,38 +198,7 @@ const GeneralSettings: React.FC = ({ accessToken, user
- {value.field_type == "Integer" ? (
- handleInputChange(value.field_name, newValue)}
- />
- ) : value.field_type == "Boolean" ? (
- handleInputChange(value.field_name, checked)}
- />
- ) : value.field_type == "Float" ? (
- handleInputChange(value.field_name, newValue)}
- />
- ) : value.field_type == "Select" ? (
- ({
- label: option,
- value: option,
- }))}
- onChange={(newValue) => handleInputChange(value.field_name, newValue ?? "")}
- />
- ) : null}
+
{value.stored_in_db == true ? (
From 16e39542a095c0f4aaf1a7835d8598b664dd6716 Mon Sep 17 00:00:00 2001
From: Tin Chi Lo
Date: Thu, 16 Jul 2026 15:54:53 -0700
Subject: [PATCH 012/296] docs(ui): state that Anthropic prompt caches are
shared per upstream credential
The provider caches a prefix against the credentials that sent it, not per end user, so
turning the flag on makes every caller's prompts cacheable on that shared account. Surface
that where the toggle is, since it is the operator's call to make.
---
litellm/proxy/proxy_server.py | 6 +++++-
1 file changed, 5 insertions(+), 1 deletion(-)
diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py
index ae91eb70427..ad1617017f9 100644
--- a/litellm/proxy/proxy_server.py
+++ b/litellm/proxy/proxy_server.py
@@ -14826,7 +14826,11 @@ _GENERAL_SETTINGS_UI_LITELLM_FIELDS: dict[str, GeneralSettingsUILiteLLMFieldSpec
"Automatically add Anthropic cache_control breakpoints to the system prompt and the "
"trailing turn, for Anthropic and Bedrock Claude models that support prompt caching. "
"Lets clients that never set cache_control themselves still get cached prompts. "
- "Requests that already carry their own cache_control are left untouched."
+ "Requests that already carry their own cache_control are left untouched. "
+ "The provider caches a prefix against the upstream credentials that sent it, not per "
+ "end user, so this makes every caller's prompts cacheable on that shared account. "
+ "Leave this off if callers sharing a set of credentials must not learn whether "
+ "another caller recently sent a given prompt."
),
},
"anthropic_prompt_caching_ttl": {
From cf23df94313ae308484a41772e1e8aab23ded4d6 Mon Sep 17 00:00:00 2001
From: Tin Chi Lo
Date: Fri, 17 Jul 2026 11:34:08 -0700
Subject: [PATCH 013/296] fix(mcp): require every reference to opt in before
auto-executing tools
_should_auto_execute_tools returned True as soon as any MCP reference set
require_approval="never", so a request that mixed a "never" reference with an
"always" or "manual" one auto-executed every tool call the model produced,
including the approval-gated ones. A prompt could name the approval-required
tool and have it run with no approval.
Make the gate fail closed: auto-execute only when every reference opts in with
"never". A single approval-required reference (including the object form or an
unset value) returns the model's tool calls to the caller instead of running
them, so an approval-gated tool can never be auto-invoked. This is the shared
decision behind /chat/completions, /responses, the streaming iterator and the
new /v1/messages path, so all four fail closed from one change. The common case,
every reference "never", is unchanged.
The alternative, executing the "never" calls and returning only the
approval-required ones, needs partial execution that the Anthropic tool loop
cannot express without fabricating tool_result blocks for the calls it withheld,
so the whole-request fail-closed gate is the safe minimum. A future change can
add per-call partial execution if a caller needs it.
Test covers the mixed and manual cases; reverting to "any never" fails it.
---
.../mcp/litellm_proxy_mcp_handler.py | 28 ++++++++++++-------
.../mcp_tests/test_aresponses_api_with_mcp.py | 9 ++++++
2 files changed, 27 insertions(+), 10 deletions(-)
diff --git a/litellm/responses/mcp/litellm_proxy_mcp_handler.py b/litellm/responses/mcp/litellm_proxy_mcp_handler.py
index d2c9f220690..a94cd2413d8 100644
--- a/litellm/responses/mcp/litellm_proxy_mcp_handler.py
+++ b/litellm/responses/mcp/litellm_proxy_mcp_handler.py
@@ -478,17 +478,25 @@ class LiteLLM_Proxy_MCP_Handler:
) -> bool:
"""Check if we should auto-execute tool calls.
- Only auto-execute tools if user passed a MCP tool with require_approval set to "never".
-
-
+ Auto-execution requires EVERY MCP reference to opt in with
+ ``require_approval="never"``. A single reference that requires approval
+ ("always", "manual", the object form, or an unset value) disables
+ auto-execution for the whole request. This fails closed: when an
+ approval-required reference shares a request with a "never" one, the
+ model's tool calls are returned to the caller instead of being run, so
+ an approval-gated tool can never be invoked without approval. Returns
+ False for an empty list.
"""
- for tool in mcp_tools_with_litellm_proxy:
- if isinstance(tool, dict):
- if tool.get("require_approval") == "never":
- return True
- elif getattr(tool, "require_approval", None) == "never":
- return True
- return False
+ references = list(mcp_tools_with_litellm_proxy or [])
+ if not references:
+ return False
+ for tool in references:
+ approval = (
+ tool.get("require_approval") if isinstance(tool, dict) else getattr(tool, "require_approval", None)
+ )
+ if approval != "never":
+ return False
+ return True
@staticmethod
def _extract_tool_calls_from_response(response: ResponsesAPIResponse) -> List[Any]:
diff --git a/tests/mcp_tests/test_aresponses_api_with_mcp.py b/tests/mcp_tests/test_aresponses_api_with_mcp.py
index 9cd45f3d6fc..32295310005 100644
--- a/tests/mcp_tests/test_aresponses_api_with_mcp.py
+++ b/tests/mcp_tests/test_aresponses_api_with_mcp.py
@@ -86,6 +86,15 @@ async def test_mcp_helper_methods():
LiteLLM_Proxy_MCP_Handler._should_auto_execute_tools(mcp_tools_always) == False
)
+ # A single approval-required reference must disable auto-execution for the
+ # whole request; otherwise a "never" reference alongside an "always" one
+ # would let the approval-gated tool run without approval.
+ mcp_tools_mixed = [{"require_approval": "never"}, {"require_approval": "always"}]
+ assert LiteLLM_Proxy_MCP_Handler._should_auto_execute_tools(mcp_tools_mixed) == False
+ mcp_tools_manual = [{"require_approval": "never"}, {"require_approval": "manual"}]
+ assert LiteLLM_Proxy_MCP_Handler._should_auto_execute_tools(mcp_tools_manual) == False
+ assert LiteLLM_Proxy_MCP_Handler._should_auto_execute_tools([]) == False
+
print("✓ MCP helper methods test passed!")
From 73cbbdd51defe21a6db5beadf2b3ee73454677be Mon Sep 17 00:00:00 2001
From: Tin Chi Lo
Date: Fri, 17 Jul 2026 11:38:24 -0700
Subject: [PATCH 014/296] feat(ui): move Anthropic prompt caching to its own
Router Settings tab
Rather than mixing the flag and its ttl into the generic General settings table
(which also surfaced the confusing Not Set / In Config / In DB provenance badges),
give prompt caching a dedicated tab with a purpose-built toggle and ttl dropdown.
Each registry field gains an optional tab, surfaced as ConfigList.field_tab, so
the General tab renders the ungrouped fields and the caching fields render on
their own tab. The update, persist and reset endpoints are unchanged.
---
litellm/proxy/_types.py | 1 +
litellm/proxy/proxy_server.py | 4 +
tests/test_litellm/proxy/test_proxy_server.py | 6 ++
.../_components/general_settings.tsx | 78 ++++++++++++++++++-
ui/litellm-dashboard/src/lib/http/schema.d.ts | 2 +
5 files changed, 90 insertions(+), 1 deletion(-)
diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py
index e07c7b9ae78..b47b43411c5 100644
--- a/litellm/proxy/_types.py
+++ b/litellm/proxy/_types.py
@@ -2123,6 +2123,7 @@ class ConfigList(LiteLLMPydanticObjectBase):
premium_field: bool = False
nested_fields: Optional[List[FieldDetail]] = None # For nested dictionary or Pydantic fields
field_options: Optional[list[str]] = None # Allowed values, for field_type == "Select"
+ field_tab: Optional[str] = None # Admin UI sub-tab this field renders under; None groups it with the rest
class UserHeaderMapping(LiteLLMPydanticObjectBase):
diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py
index ad1617017f9..6725ecdb584 100644
--- a/litellm/proxy/proxy_server.py
+++ b/litellm/proxy/proxy_server.py
@@ -14809,6 +14809,7 @@ class GeneralSettingsUILiteLLMFieldSpec(TypedDict):
type: Literal["Float", "Boolean", "Select"]
description: str
options: NotRequired[tuple[str, ...]]
+ tab: NotRequired[str] # Admin UI sub-tab this field renders under; None groups it with the rest
_GENERAL_SETTINGS_UI_LITELLM_FIELDS: dict[str, GeneralSettingsUILiteLLMFieldSpec] = {
@@ -14822,6 +14823,7 @@ _GENERAL_SETTINGS_UI_LITELLM_FIELDS: dict[str, GeneralSettingsUILiteLLMFieldSpec
},
"enable_anthropic_prompt_caching": {
"type": "Boolean",
+ "tab": "prompt_caching",
"description": (
"Automatically add Anthropic cache_control breakpoints to the system prompt and the "
"trailing turn, for Anthropic and Bedrock Claude models that support prompt caching. "
@@ -14836,6 +14838,7 @@ _GENERAL_SETTINGS_UI_LITELLM_FIELDS: dict[str, GeneralSettingsUILiteLLMFieldSpec
"anthropic_prompt_caching_ttl": {
"type": "Select",
"options": ("5m", "1h"),
+ "tab": "prompt_caching",
"description": (
"Cache lifetime for the breakpoints added by 'enable_anthropic_prompt_caching'. "
"Leave empty for Anthropic's 5m default. 1h suits long agentic sessions but doubles "
@@ -15093,6 +15096,7 @@ async def get_config_list(
stored_in_db=stored_in_db_litellm,
field_default_value=default_value,
field_options=list(spec.get("options", ())) or None,
+ field_tab=spec.get("tab"),
nested_fields=None,
)
)
diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py
index 86e807447c1..56cf213f103 100644
--- a/tests/test_litellm/proxy/test_proxy_server.py
+++ b/tests/test_litellm/proxy/test_proxy_server.py
@@ -9019,6 +9019,12 @@ def test_get_config_list_includes_anthropic_prompt_caching_fields(monkeypatch):
assert fields["anthropic_prompt_caching_ttl"]["field_type"] == "Select"
assert fields["anthropic_prompt_caching_ttl"]["field_value"] == "1h"
assert fields["anthropic_prompt_caching_ttl"]["field_options"] == ["5m", "1h"]
+
+ # Both caching fields carry their sub-tab so the Admin UI can render them on a
+ # dedicated Prompt Caching tab, while ungrouped fields stay on General.
+ assert fields["enable_anthropic_prompt_caching"]["field_tab"] == "prompt_caching"
+ assert fields["anthropic_prompt_caching_ttl"]["field_tab"] == "prompt_caching"
+ assert fields["budget_exceeded_throttle_percentage"]["field_tab"] is None
finally:
app.dependency_overrides.clear()
diff --git a/ui/litellm-dashboard/src/app/(dashboard)/router-settings/_components/general_settings.tsx b/ui/litellm-dashboard/src/app/(dashboard)/router-settings/_components/general_settings.tsx
index af6bbdde8b9..8cea529d25d 100644
--- a/ui/litellm-dashboard/src/app/(dashboard)/router-settings/_components/general_settings.tsx
+++ b/ui/litellm-dashboard/src/app/(dashboard)/router-settings/_components/general_settings.tsx
@@ -7,6 +7,7 @@ import {
TableHeaderCell,
TableCell,
TableBody,
+ Title,
Text,
Button,
Icon,
@@ -21,6 +22,11 @@ import { StatusBadge } from "@/components/shared/table_cells";
import RouterSettings from "@/components/router_settings";
import Fallbacks from "@/components/Settings/RouterSettings/Fallbacks/Fallbacks";
import RoutingGroups from "@/components/routing_groups";
+
+const PROMPT_CACHING_TAB = "prompt_caching";
+const ENABLE_ANTHROPIC_PROMPT_CACHING = "enable_anthropic_prompt_caching";
+const ANTHROPIC_PROMPT_CACHING_TTL = "anthropic_prompt_caching_ttl";
+
interface GeneralSettingsPageProps {
accessToken: string | null;
userRole: string | null;
@@ -34,6 +40,7 @@ interface generalSettingsItem {
field_description: string;
stored_in_db: boolean | null;
field_options?: string[] | null;
+ field_tab?: string | null;
}
const SettingValueEditor: React.FC<{
@@ -83,6 +90,71 @@ const SettingValueEditor: React.FC<{
return null;
};
+const PromptCachingPanel: React.FC<{
+ accessToken: string;
+ settings: generalSettingsItem[];
+ onChange: (fieldName: string, newValue: any) => void;
+}> = ({ accessToken, settings, onChange }) => {
+ const enableSetting = settings.find((s) => s.field_name === ENABLE_ANTHROPIC_PROMPT_CACHING);
+ const ttlSetting = settings.find((s) => s.field_name === ANTHROPIC_PROMPT_CACHING_TTL);
+
+ // The two rows come from the same registry the General tab reads; if they
+ // are not loaded yet there is nothing to render.
+ if (!enableSetting) {
+ return null;
+ }
+
+ const enabled = enableSetting.field_value === true || enableSetting.field_value === "true";
+
+ // Apply immediately: a toggle and a dropdown are direct controls, so there is
+ // no separate Update button. Clearing the ttl resets it to the provider default.
+ const persist = (fieldName: string, value: any) => {
+ onChange(fieldName, value);
+ if (value === "" || value === null || value === undefined) {
+ deleteConfigFieldSetting(accessToken, fieldName);
+ } else {
+ updateConfigFieldSetting(accessToken, fieldName, value);
+ }
+ };
+
+ return (
+
+ Prompt Caching
+
+ Automatically inject Anthropic prompt caching for every Anthropic and Bedrock Claude model, so clients that
+ never set cache_control themselves still get cached prompts. This is a single
+ gateway-wide switch; there is no per-model setup.
+
+
+
@@ -181,7 +257,7 @@ const GeneralSettings: React.FC = ({ accessToken, user
{generalSettings
- .filter((value) => value.field_type !== "TypedDictionary")
+ .filter((value) => value.field_type !== "TypedDictionary" && value.field_tab !== PROMPT_CACHING_TAB)
.map((value, index) => (
diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts
index 45b06c9e44f..6dc63762e5c 100644
--- a/ui/litellm-dashboard/src/lib/http/schema.d.ts
+++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts
@@ -22713,6 +22713,8 @@ export interface components {
field_name: string;
/** Field Options */
field_options?: string[] | null;
+ /** Field Tab */
+ field_tab?: string | null;
/** Field Type */
field_type: string;
/** Field Value */
From 4e5f4884523ea124c6a252563624d104b4dc394c Mon Sep 17 00:00:00 2001
From: Tin Chi Lo
Date: Fri, 17 Jul 2026 12:16:09 -0700
Subject: [PATCH 015/296] feat(ui): tighten the Prompt Caching descriptions
The toggle and ttl descriptions were a wall of text, with a panel intro that
mostly repeated the toggle description. Drop the intro and cut both descriptions
to one or two lines, keeping a one-clause note that the cache is shared across
callers on the same upstream credentials.
---
litellm/proxy/proxy_server.py | 16 +++-------------
.../_components/general_settings.tsx | 5 -----
2 files changed, 3 insertions(+), 18 deletions(-)
diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py
index 6725ecdb584..7d87207f7a9 100644
--- a/litellm/proxy/proxy_server.py
+++ b/litellm/proxy/proxy_server.py
@@ -14825,25 +14825,15 @@ _GENERAL_SETTINGS_UI_LITELLM_FIELDS: dict[str, GeneralSettingsUILiteLLMFieldSpec
"type": "Boolean",
"tab": "prompt_caching",
"description": (
- "Automatically add Anthropic cache_control breakpoints to the system prompt and the "
- "trailing turn, for Anthropic and Bedrock Claude models that support prompt caching. "
- "Lets clients that never set cache_control themselves still get cached prompts. "
- "Requests that already carry their own cache_control are left untouched. "
- "The provider caches a prefix against the upstream credentials that sent it, not per "
- "end user, so this makes every caller's prompts cacheable on that shared account. "
- "Leave this off if callers sharing a set of credentials must not learn whether "
- "another caller recently sent a given prompt."
+ "Auto-adds cache_control to the system prompt and trailing turn for supported Anthropic "
+ "and Bedrock Claude models. The cache is shared across callers on the same upstream credentials."
),
},
"anthropic_prompt_caching_ttl": {
"type": "Select",
"options": ("5m", "1h"),
"tab": "prompt_caching",
- "description": (
- "Cache lifetime for the breakpoints added by 'enable_anthropic_prompt_caching'. "
- "Leave empty for Anthropic's 5m default. 1h suits long agentic sessions but doubles "
- "the cache write premium."
- ),
+ "description": "Empty uses Anthropic's 5m default. 1h suits long sessions but doubles the cache write cost.",
},
}
diff --git a/ui/litellm-dashboard/src/app/(dashboard)/router-settings/_components/general_settings.tsx b/ui/litellm-dashboard/src/app/(dashboard)/router-settings/_components/general_settings.tsx
index 8cea529d25d..1e8658d5104 100644
--- a/ui/litellm-dashboard/src/app/(dashboard)/router-settings/_components/general_settings.tsx
+++ b/ui/litellm-dashboard/src/app/(dashboard)/router-settings/_components/general_settings.tsx
@@ -120,11 +120,6 @@ const PromptCachingPanel: React.FC<{
return (
Prompt Caching
-
- Automatically inject Anthropic prompt caching for every Anthropic and Bedrock Claude model, so clients that
- never set cache_control themselves still get cached prompts. This is a single
- gateway-wide switch; there is no per-model setup.
-
From d9661222492a098555f40cb8b50014054bea5ab8 Mon Sep 17 00:00:00 2001
From: Tin Chi Lo
Date: Fri, 17 Jul 2026 17:19:17 -0700
Subject: [PATCH 016/296] fix(fireworks_ai): correct glm-5p2 prompt-cache read
price to $0.14/1M
glm-5p2 (and its fireworks_ai/glm-5p2 alias) carried cache_read_input_token_cost
of 2.6e-07, the GLM 5.1 rate; the entry was seeded from the wrong row. Fireworks'
standard serverless rate for GLM 5.2 is $0.14/1M = 1.4e-07, so every prompt-cache
hit was billed at nearly double the real rate.
Corrects the value in both the canonical map and the bundled backup. The existing
fireworks cost-calculator test now reads the cached rate from the map instead of
hardcoding it, so it tracks the shipped value.
---
litellm/model_prices_and_context_window_backup.json | 4 ++--
model_prices_and_context_window.json | 4 ++--
.../llms/fireworks_ai/test_fireworks_ai_cost_calculator.py | 5 ++++-
tests/test_litellm/test_utils.py | 2 +-
4 files changed, 9 insertions(+), 6 deletions(-)
diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json
index ee996198b28..e5afc81b641 100644
--- a/litellm/model_prices_and_context_window_backup.json
+++ b/litellm/model_prices_and_context_window_backup.json
@@ -16272,7 +16272,7 @@
"supports_vision": false
},
"fireworks_ai/accounts/fireworks/models/glm-5p2": {
- "cache_read_input_token_cost": 2.6e-07,
+ "cache_read_input_token_cost": 1.4e-07,
"input_cost_per_token": 1.4e-06,
"litellm_provider": "fireworks_ai",
"max_input_tokens": 1048576,
@@ -16686,7 +16686,7 @@
"supports_vision": false
},
"fireworks_ai/glm-5p2": {
- "cache_read_input_token_cost": 2.6e-07,
+ "cache_read_input_token_cost": 1.4e-07,
"input_cost_per_token": 1.4e-06,
"litellm_provider": "fireworks_ai",
"max_input_tokens": 1048576,
diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json
index b1a87c444c8..5cf99ba8bac 100644
--- a/model_prices_and_context_window.json
+++ b/model_prices_and_context_window.json
@@ -16272,7 +16272,7 @@
"supports_vision": false
},
"fireworks_ai/accounts/fireworks/models/glm-5p2": {
- "cache_read_input_token_cost": 2.6e-07,
+ "cache_read_input_token_cost": 1.4e-07,
"input_cost_per_token": 1.4e-06,
"litellm_provider": "fireworks_ai",
"max_input_tokens": 1048576,
@@ -16686,7 +16686,7 @@
"supports_vision": false
},
"fireworks_ai/glm-5p2": {
- "cache_read_input_token_cost": 2.6e-07,
+ "cache_read_input_token_cost": 1.4e-07,
"input_cost_per_token": 1.4e-06,
"litellm_provider": "fireworks_ai",
"max_input_tokens": 1048576,
diff --git a/tests/test_litellm/llms/fireworks_ai/test_fireworks_ai_cost_calculator.py b/tests/test_litellm/llms/fireworks_ai/test_fireworks_ai_cost_calculator.py
index 99dcaa36c75..3297750fa6e 100644
--- a/tests/test_litellm/llms/fireworks_ai/test_fireworks_ai_cost_calculator.py
+++ b/tests/test_litellm/llms/fireworks_ai/test_fireworks_ai_cost_calculator.py
@@ -5,12 +5,15 @@ import pytest
sys.path.insert(0, os.path.abspath("../../../../.."))
+import litellm
from litellm.llms.fireworks_ai.cost_calculator import cost_per_token
from litellm.types.utils import PromptTokensDetailsWrapper, Usage
MODEL = "accounts/fireworks/models/glm-5p2"
INPUT_COST = 1.4e-06
-CACHE_READ_COST = 2.6e-07
+# Read the cached rate from the price map so this test tracks the shipped value
+# (glm-5p2 is $0.14/1M) instead of hardcoding a number that breaks when it changes.
+CACHE_READ_COST = litellm.get_model_info(model=MODEL, custom_llm_provider="fireworks_ai")["cache_read_input_token_cost"]
OUTPUT_COST = 4.4e-06
diff --git a/tests/test_litellm/test_utils.py b/tests/test_litellm/test_utils.py
index edd93cbebe0..a1a9448cc58 100644
--- a/tests/test_litellm/test_utils.py
+++ b/tests/test_litellm/test_utils.py
@@ -4317,7 +4317,7 @@ _FIREWORKS_MODELS = [
"accounts/fireworks/models/glm-5p2",
1.4e-06,
4.4e-06,
- 2.6e-07,
+ 1.4e-07,
1048576,
131072,
False,
From 836bf0807b62fe346697e3a1b987cc5b05afbbf9 Mon Sep 17 00:00:00 2001
From: Shivam Rawat
Date: Fri, 17 Jul 2026 18:19:02 -0700
Subject: [PATCH 017/296] fix(router): keep team wildcard routers fresh and
prioritize them over global patterns
team_pattern_routers retained deleted/replaced deployments, so team users could
keep resolving stale credentials; now set_model_list resets the registry and
deployment removal prunes it. Also consult the team wildcard router before the
global pattern_router in get_deployment_credentials_with_provider so a global
pattern like "openai/*" no longer shadows the team's own entry
Co-authored-by: Cursor
---
litellm/router.py | 26 ++++--
.../router_utils/pattern_match_deployments.py | 11 +++
tests/test_litellm/test_router.py | 92 +++++++++++++++++++
3 files changed, 121 insertions(+), 8 deletions(-)
diff --git a/litellm/router.py b/litellm/router.py
index dbc6da106e7..6a055f54b0e 100644
--- a/litellm/router.py
+++ b/litellm/router.py
@@ -7876,6 +7876,7 @@ class Router:
self.model_id_to_deployment_index_map = {} # Reset the index
self.model_name_to_deployment_indices = {} # Reset the model_name index
self.team_model_to_deployment_indices = {} # Reset the team_model index
+ self.team_pattern_routers = {}
self.team_public_model_names = frozenset()
# Reset per-strategy router registries so hot-reload doesn't leave
# stale routers pointing at the old model_list.
@@ -8232,6 +8233,12 @@ class Router:
public_model_name for _, public_model_name in self.team_model_to_deployment_indices
)
+ for team_id in list(self.team_pattern_routers.keys()):
+ team_pattern_router = self.team_pattern_routers[team_id]
+ team_pattern_router.remove_deployment(model_id)
+ if not team_pattern_router.patterns:
+ del self.team_pattern_routers[team_id]
+
def _update_team_model_index(self, model: dict, idx: int) -> None:
"""
Helper to update team_model_to_deployment_indices for a single deployment.
@@ -8460,8 +8467,8 @@ class Router:
return None
def get_deployment_credentials_with_provider(
- self, model_id: str, team_id: Optional[str] = None
- ) -> Optional[Dict[str, Any]]:
+ self, model_id: str, team_id: str | None = None
+ ) -> dict[str, Any] | None:
"""
Get API credentials and provider info from a model name in model_list.
Useful for passthrough endpoints (files, batches, etc.) that need credentials.
@@ -8492,19 +8499,22 @@ class Router:
if deployment is None:
deployment = self.get_deployment_by_model_group_name(model_group_name=model_id)
- # If not found, check team-scoped deployments (team public model names,
- # e.g. team wildcard models like "openai/*", live in a separate index).
+ # If not found, check team-scoped deployments whose team public model
+ # name exactly matches model_id (wildcard team names are matched via
+ # team_pattern_routers below).
if deployment is None and team_id is not None:
team_indices = self.team_model_to_deployment_indices.get((team_id, model_id), [])
if team_indices:
team_model = self.model_list[team_indices[0]]
deployment = Deployment(**team_model) if isinstance(team_model, dict) else team_model
- # If still not found, check for wildcard pattern matches
+ # If still not found, check for wildcard pattern matches. Team wildcard
+ # matches take priority so a global pattern (e.g. "openai/*") doesn't
+ # shadow the team's own entry.
if deployment is None:
- potential_wildcard_models = self.pattern_router.route(model_id) or []
- if not potential_wildcard_models and team_id is not None and team_id in self.team_pattern_routers:
- potential_wildcard_models = self.team_pattern_routers[team_id].route(model_id) or []
+ team_pattern_router = self.team_pattern_routers.get(team_id) if team_id is not None else None
+ team_wildcard_models = (team_pattern_router.route(model_id) or []) if team_pattern_router else []
+ potential_wildcard_models = team_wildcard_models or self.pattern_router.route(model_id) or []
if potential_wildcard_models:
# Use the first matching wildcard deployment
deployment_dict = potential_wildcard_models[0]
diff --git a/litellm/router_utils/pattern_match_deployments.py b/litellm/router_utils/pattern_match_deployments.py
index c08f8e95cf4..7e1ed739ef8 100644
--- a/litellm/router_utils/pattern_match_deployments.py
+++ b/litellm/router_utils/pattern_match_deployments.py
@@ -73,6 +73,17 @@ class PatternMatchRouter:
self.patterns[regex] = []
self.patterns[regex].append(llm_deployment)
+ def remove_deployment(self, model_id: str) -> None:
+ """
+ Remove every deployment with the given model id from the pattern registry,
+ dropping any pattern whose deployment list becomes empty.
+ """
+ self.patterns = {
+ regex: remaining
+ for regex, deployments in self.patterns.items()
+ if (remaining := [d for d in deployments if (d.get("model_info") or {}).get("id") != model_id])
+ }
+
def _pattern_to_regex(self, pattern: str) -> str:
"""
Convert a wildcard pattern to a regex pattern
diff --git a/tests/test_litellm/test_router.py b/tests/test_litellm/test_router.py
index c2c98c8869c..0fe855151b0 100644
--- a/tests/test_litellm/test_router.py
+++ b/tests/test_litellm/test_router.py
@@ -3535,6 +3535,98 @@ def test_get_deployment_credentials_with_provider_resolves_credential_name():
litellm.credential_list = []
+def _team_wildcard_model(api_key: str, model_id: str = "team-wildcard-id") -> dict:
+ return {
+ "model_name": f"model_name_team-1_{model_id}",
+ "litellm_params": {"model": "openai/*", "api_key": api_key},
+ "model_info": {
+ "id": model_id,
+ "team_id": "team-1",
+ "team_public_model_name": "openai/*",
+ },
+ }
+
+
+def test_get_deployment_credentials_with_provider_team_wildcard_priority():
+ """
+ Regression: a global wildcard pattern (e.g. "openai/*") must not shadow a
+ team's own wildcard entry. When team_id is provided, the team wildcard
+ deployment's credentials win; without team_id the global one is used.
+ """
+ router = litellm.Router(
+ model_list=[
+ {
+ "model_name": "openai/*",
+ "litellm_params": {"model": "openai/*", "api_key": "global-key"},
+ },
+ _team_wildcard_model(api_key="team-key"),
+ ],
+ )
+
+ team_credentials = router.get_deployment_credentials_with_provider(
+ model_id="openai/gpt-5.2", team_id="team-1"
+ )
+ assert team_credentials is not None
+ assert team_credentials["api_key"] == "team-key"
+
+ global_credentials = router.get_deployment_credentials_with_provider(
+ model_id="openai/gpt-5.2"
+ )
+ assert global_credentials is not None
+ assert global_credentials["api_key"] == "global-key"
+
+
+def test_team_wildcard_credentials_not_usable_after_delete_deployment():
+ """
+ Regression: team_pattern_routers retained deleted deployments, so a team
+ user could keep resolving credentials of a deleted wildcard deployment.
+ """
+ router = litellm.Router(model_list=[_team_wildcard_model(api_key="old-key")])
+
+ assert (
+ router.get_deployment_credentials_with_provider(
+ model_id="openai/gpt-5.2", team_id="team-1"
+ )
+ is not None
+ )
+
+ router.delete_deployment(id="team-wildcard-id")
+
+ assert (
+ router.get_deployment_credentials_with_provider(
+ model_id="openai/gpt-5.2", team_id="team-1"
+ )
+ is None
+ )
+
+
+def test_team_wildcard_credentials_refreshed_on_upsert_and_set_model_list():
+ """
+ Regression: replacing a team wildcard deployment (upsert or model list
+ reload) must serve the new credentials, not the stale cached ones.
+ """
+ from litellm.types.router import Deployment
+
+ router = litellm.Router(model_list=[_team_wildcard_model(api_key="old-key")])
+
+ router.upsert_deployment(
+ deployment=Deployment(**_team_wildcard_model(api_key="new-key"))
+ )
+ credentials = router.get_deployment_credentials_with_provider(
+ model_id="openai/gpt-5.2", team_id="team-1"
+ )
+ assert credentials is not None
+ assert credentials["api_key"] == "new-key"
+
+ router.set_model_list(model_list=[])
+ assert (
+ router.get_deployment_credentials_with_provider(
+ model_id="openai/gpt-5.2", team_id="team-1"
+ )
+ is None
+ )
+
+
def test_get_available_guardrail_single_deployment():
"""
Test get_available_guardrail returns the single guardrail when only one exists.
From b792fd7c5fb1e448fba5ae910d4c6ba230fd04a9 Mon Sep 17 00:00:00 2001
From: Shivam Rawat
Date: Fri, 17 Jul 2026 18:24:58 -0700
Subject: [PATCH 018/296] test(router): cover
PatternMatchRouter.remove_deployment for router code coverage gate
Co-authored-by: Cursor
---
tests/test_litellm/test_router.py | 27 +++++++++++++++++++++++++++
1 file changed, 27 insertions(+)
diff --git a/tests/test_litellm/test_router.py b/tests/test_litellm/test_router.py
index 0fe855151b0..f9360abea51 100644
--- a/tests/test_litellm/test_router.py
+++ b/tests/test_litellm/test_router.py
@@ -3600,6 +3600,33 @@ def test_team_wildcard_credentials_not_usable_after_delete_deployment():
)
+def test_pattern_match_router_remove_deployment():
+ """
+ remove_deployment must drop only the deployment with the given model id and
+ delete patterns whose deployment list becomes empty.
+ """
+ from litellm.router_utils.pattern_match_deployments import PatternMatchRouter
+
+ pattern_router = PatternMatchRouter()
+ pattern_router.add_pattern(
+ "openai/*",
+ {"litellm_params": {"model": "openai/*", "api_key": "key-a"}, "model_info": {"id": "dep-a"}},
+ )
+ pattern_router.add_pattern(
+ "openai/*",
+ {"litellm_params": {"model": "openai/*", "api_key": "key-b"}, "model_info": {"id": "dep-b"}},
+ )
+
+ pattern_router.remove_deployment(model_id="dep-a")
+ matches = pattern_router.route("openai/gpt-5.2")
+ assert matches is not None
+ assert [m["model_info"]["id"] for m in matches] == ["dep-b"]
+
+ pattern_router.remove_deployment(model_id="dep-b")
+ assert pattern_router.patterns == {}
+ assert pattern_router.route("openai/gpt-5.2") is None
+
+
def test_team_wildcard_credentials_refreshed_on_upsert_and_set_model_list():
"""
Regression: replacing a team wildcard deployment (upsert or model list
From 47ba9e76121dd7dbf572e112d7df5ebad5414bac Mon Sep 17 00:00:00 2001
From: Tin Chi Lo
Date: Fri, 17 Jul 2026 19:38:42 -0700
Subject: [PATCH 019/296] fix(proxy): propagate the caching flag across workers
via the safe-override allowlist
enable_anthropic_prompt_caching and anthropic_prompt_caching_ttl are set as live
litellm attributes on the worker that handles the UI save, exactly like
budget_exceeded_throttle_percentage, but they were missing from
LITELLM_SETTINGS_SAFE_DB_OVERRIDES, so a peer worker's config reload merged the DB
value without applying it to the live attribute and stayed stale.
Add both to the allowlist so they behave like the sibling field, and add
test_general_settings_ui_fields_are_db_overridable so the UI registry and the
override allowlist cannot drift again (the exact omission that caused this), plus
a regression test that the flag flips on a simulated peer-worker reload.
---
litellm/constants.py | 6 +++
tests/test_litellm/proxy/test_proxy_server.py | 45 +++++++++++++++++++
2 files changed, 51 insertions(+)
diff --git a/litellm/constants.py b/litellm/constants.py
index e104c937a9b..6432e2176c7 100644
--- a/litellm/constants.py
+++ b/litellm/constants.py
@@ -1517,6 +1517,12 @@ LITELLM_SETTINGS_SAFE_DB_OVERRIDES = [
"cost_discount_config",
"cost_margin_config",
"budget_exceeded_throttle_percentage",
+ # Every field editable from the Admin UI (proxy_server._GENERAL_SETTINGS_UI_LITELLM_FIELDS)
+ # must be listed here so a DB write from one worker overrides the live litellm attribute on
+ # the others when config reloads; otherwise peer workers stay on their startup value.
+ # test_general_settings_ui_fields_are_db_overridable enforces that pairing.
+ "enable_anthropic_prompt_caching",
+ "anthropic_prompt_caching_ttl",
]
SPECIAL_LITELLM_AUTH_TOKEN = ["ui-token"]
DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL = int(os.getenv("DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL", 60))
diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py
index 56cf213f103..a100e7837f4 100644
--- a/tests/test_litellm/proxy/test_proxy_server.py
+++ b/tests/test_litellm/proxy/test_proxy_server.py
@@ -9029,6 +9029,51 @@ def test_get_config_list_includes_anthropic_prompt_caching_fields(monkeypatch):
app.dependency_overrides.clear()
+def test_general_settings_ui_fields_are_db_overridable():
+ """Every field the Admin UI can edit is a `litellm.` set via setattr on the handling
+ worker (`_persist_general_settings_ui_litellm_field`). Unless it is also in
+ LITELLM_SETTINGS_SAFE_DB_OVERRIDES, a config reload on a peer worker merges the DB value but
+ never applies it to the live attribute, so peer workers stay on their startup value.
+
+ This invariant is the guard against the two registries drifting: adding a UI-editable field
+ without enrolling it in the DB-override allowlist silently breaks cross-worker propagation.
+ """
+ from litellm.constants import LITELLM_SETTINGS_SAFE_DB_OVERRIDES
+ from litellm.proxy.proxy_server import _GENERAL_SETTINGS_UI_LITELLM_FIELDS
+
+ missing = set(_GENERAL_SETTINGS_UI_LITELLM_FIELDS) - set(LITELLM_SETTINGS_SAFE_DB_OVERRIDES)
+ assert not missing, (
+ f"UI-editable litellm_settings fields missing from LITELLM_SETTINGS_SAFE_DB_OVERRIDES: {sorted(missing)}. "
+ "Add them, or they will not propagate to other workers when changed from the UI."
+ )
+
+
+@pytest.mark.parametrize(
+ "field_name, db_value",
+ [
+ ("enable_anthropic_prompt_caching", True),
+ ("anthropic_prompt_caching_ttl", "1h"),
+ ],
+)
+def test_prompt_caching_settings_propagate_on_config_reload(monkeypatch, field_name, db_value):
+ """A UI toggle on one worker persists to the DB; a peer worker picks it up only when the
+ config reload applies the safe-override allowlist. Regression for the fields being absent
+ from that allowlist, which left peer workers stale."""
+ import litellm.proxy.proxy_server as ps
+
+ # peer worker booted with the opposite/absent value
+ monkeypatch.setattr(litellm, field_name, False if isinstance(db_value, bool) else None)
+
+ pc = ps.ProxyConfig()
+ pc._update_config_fields(
+ current_config={"litellm_settings": {}},
+ param_name="litellm_settings",
+ db_param_value={field_name: db_value},
+ )
+
+ assert getattr(litellm, field_name) == db_value
+
+
def test_get_config_list_marks_untouched_prompt_caching_flag_as_not_set(monkeypatch):
"""The flag defaults to False rather than None, so a plain 'is not None' check would
report the default as 'In Config' and imply an admin had set it."""
From 99b85a3f2cac8fff501a07e6301274cc387ef245 Mon Sep 17 00:00:00 2001
From: Tin Chi Lo
Date: Fri, 17 Jul 2026 12:10:12 -0700
Subject: [PATCH 020/296] fix(mcp): persist config.yaml DCR clients in a
server-scoped store
Config.yaml-declared OAuth2 MCP servers using Dynamic Client Registration have no LiteLLM_MCPServerTable row, so the DCR persist path called update_mcp_server, which returns None for a missing row, then update_server(None), which dereferenced .approval_status and raised AttributeError. The exception was swallowed to a warning while /register still returned 200, so the minted client was never stored and every access-token expiry forced a full re-authorization
Persist the acquired DCR client (client_id, client_secret, token_endpoint_auth_method, redirect_uris, encrypted at rest) in a dedicated LiteLLM_MCPServerOAuthClient store keyed by server_id when the server has no row, overlay it onto the in-memory config server so the refresh_token grant can authenticate within the process, and rehydrate it when the registry syncs from the database (which runs after the DB connects, unlike config load) so restarts and other pods pick it up. The store is encrypted at rest and is re-encrypted by the master-key rotation path alongside the server rows, through a shared helper so the two sites cannot diverge. The DB-backed server path is unchanged, and guarding the None return removes the swallowed-crash footgun
Resolves the config.yaml DCR persistence regression introduced in v1.92.0 by #31912
---
.../migration.sql | 9 +
.../litellm_proxy_extras/schema.prisma | 7 +
litellm/proxy/_experimental/mcp_server/db.py | 82 +++-
.../mcp_server/discoverable_endpoints.py | 144 +++++--
.../mcp_server/mcp_server_manager.py | 34 ++
litellm/proxy/schema.prisma | 7 +
litellm/repositories/table_repositories.py | 4 +
schema.prisma | 7 +
.../mcp_server/test_db_credentials.py | 64 +++
.../mcp_server/test_discoverable_endpoints.py | 370 ++++++++++++++++++
.../mcp_server/test_mcp_sigv4_auth.py | 2 +
11 files changed, 679 insertions(+), 51 deletions(-)
create mode 100644 litellm-proxy-extras/litellm_proxy_extras/migrations/20260717000000_add_mcp_server_oauth_client_table/migration.sql
diff --git a/litellm-proxy-extras/litellm_proxy_extras/migrations/20260717000000_add_mcp_server_oauth_client_table/migration.sql b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260717000000_add_mcp_server_oauth_client_table/migration.sql
new file mode 100644
index 00000000000..7aa6cdb1e33
--- /dev/null
+++ b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260717000000_add_mcp_server_oauth_client_table/migration.sql
@@ -0,0 +1,9 @@
+-- CreateTable
+CREATE TABLE IF NOT EXISTS "LiteLLM_MCPServerOAuthClient" (
+ "server_id" TEXT NOT NULL,
+ "credentials" JSONB,
+ "created_at" TIMESTAMP(3) NOT NULL DEFAULT CURRENT_TIMESTAMP,
+ "updated_at" TIMESTAMP(3) NOT NULL DEFAULT CURRENT_TIMESTAMP,
+
+ CONSTRAINT "LiteLLM_MCPServerOAuthClient_pkey" PRIMARY KEY ("server_id")
+);
diff --git a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma
index f842bf13da9..a99cec49417 100644
--- a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma
+++ b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma
@@ -396,6 +396,13 @@ model LiteLLM_MCPUserEnvVars {
@@index([server_id])
}
+model LiteLLM_MCPServerOAuthClient {
+ server_id String @id
+ credentials Json?
+ created_at DateTime @default(now()) @map("created_at")
+ updated_at DateTime @default(now()) @updatedAt @map("updated_at")
+}
+
// Generate Tokens for Proxy
model LiteLLM_VerificationToken {
token String @id
diff --git a/litellm/proxy/_experimental/mcp_server/db.py b/litellm/proxy/_experimental/mcp_server/db.py
index d55eb3ac014..7129582ff2a 100644
--- a/litellm/proxy/_experimental/mcp_server/db.py
+++ b/litellm/proxy/_experimental/mcp_server/db.py
@@ -33,6 +33,7 @@ from litellm.proxy.common_utils.encrypt_decrypt_utils import (
from litellm.proxy.utils import PrismaClient
from litellm.repositories.object_permission_repository import ObjectPermissionRepository
from litellm.repositories.table_repositories import (
+ MCPServerOAuthClientRepository,
MCPServerRepository,
MCPUserCredentialsRepository,
)
@@ -639,6 +640,7 @@ async def delete_mcp_server(
for model, label in (
(prisma_client.db.litellm_mcpusercredentials, "credential"),
(prisma_client.db.litellm_mcpuserenvvars, "env var"),
+ (prisma_client.db.litellm_mcpserveroauthclient, "OAuth client"),
):
try:
await model.delete_many(where={"server_id": server_id})
@@ -823,26 +825,66 @@ async def update_mcp_server(
return updated_mcp_server
-async def rotate_mcp_server_credentials_master_key(prisma_client: PrismaClient, touched_by: str, new_master_key: str):
+async def get_mcp_server_oauth_client_credentials(prisma_client: PrismaClient, server_id: str) -> object | None:
+ """Read the persisted (encrypted) DCR OAuth client blob for a server from the
+ server-scoped store, or None. Config.yaml-declared servers have no
+ LiteLLM_MCPServerTable row, so their dynamically registered client lives here keyed
+ by server_id. The returned value is the raw credentials blob for
+ ``_get_persisted_dcr_credentials`` to parse."""
+ row = await MCPServerOAuthClientRepository(prisma_client).table.find_unique(where={"server_id": server_id})
+ if row is None:
+ return None
+ return row.credentials
+
+
+async def upsert_mcp_server_oauth_client_credentials(
+ prisma_client: PrismaClient, server_id: str, credentials: MCPCredentials
+) -> None:
+ """Persist a server's dynamically registered OAuth client (RFC 7591 DCR) in the
+ server-scoped store keyed by server_id, independent of any LiteLLM_MCPServerTable row.
+ client_id/client_secret are encrypted at rest with the same salt key used for the
+ server row's credentials blob, so ``_apply_persisted_dcr_credentials`` decrypts them the
+ same way regardless of which store a server's client came from."""
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
+ encrypted = encrypt_credentials(credentials=dict(credentials), encryption_key=_get_salt_key())
+ blob = safe_dumps(encrypted)
+ await MCPServerOAuthClientRepository(prisma_client).table.upsert(
+ where={"server_id": server_id},
+ data={
+ "create": {"server_id": server_id, "credentials": blob},
+ "update": {"credentials": blob},
+ },
+ )
+
+
+def _reencrypt_mcp_credentials_blob(credentials: object, new_master_key: str) -> str | None:
+ """Decrypt an at-rest MCP credentials blob with the current key and re-encrypt it under
+ new_master_key, returning the serialized blob or None when there is nothing to rotate. Shared by
+ every table that stores an encrypted MCP credentials blob so a master-key rotation covers them
+ uniformly and cannot silently skip one."""
+ if not credentials:
+ return None
+ from litellm.litellm_core_utils.safe_json_dumps import safe_dumps # noqa: PLC0415 # avoids circular import
+
+ creds_dict = json.loads(credentials) if isinstance(credentials, str) else dict(credentials)
+ decrypted = decrypt_credentials(credentials=cast(MCPCredentials, creds_dict))
+ encrypted = encrypt_credentials(credentials=decrypted, encryption_key=new_master_key)
+ return safe_dumps(encrypted)
+
+
+async def rotate_mcp_server_credentials_master_key(prisma_client: PrismaClient, touched_by: str, new_master_key: str):
+ from litellm.litellm_core_utils.safe_json_dumps import safe_dumps # noqa: PLC0415 # avoids circular import
+
mcp_servers = await MCPServerRepository(prisma_client).table.find_many()
updated = 0
for mcp_server in mcp_servers:
update_data: Dict[str, Any] = {}
- credentials = mcp_server.credentials
- if credentials:
- # Decrypt with current key first, then re-encrypt with new key
- decrypted_credentials = decrypt_credentials(
- credentials=cast(MCPCredentials, dict(credentials)),
- )
- encrypted_credentials = encrypt_credentials(
- credentials=decrypted_credentials,
- encryption_key=new_master_key,
- )
- update_data["credentials"] = safe_dumps(encrypted_credentials)
+ rotated_credentials = _reencrypt_mcp_credentials_blob(mcp_server.credentials, new_master_key)
+ if rotated_credentials is not None:
+ update_data["credentials"] = rotated_credentials
rotated_env_vars = _reencrypt_global_env_var_values(mcp_server.env_vars, new_master_key)
if rotated_env_vars is not None:
@@ -857,9 +899,23 @@ async def rotate_mcp_server_credentials_master_key(prisma_client: PrismaClient,
data=update_data,
)
updated += 1
+
+ oauth_clients = await MCPServerOAuthClientRepository(prisma_client).table.find_many()
+ oauth_updated = 0
+ for oauth_client in oauth_clients:
+ rotated_credentials = _reencrypt_mcp_credentials_blob(oauth_client.credentials, new_master_key)
+ if rotated_credentials is None:
+ continue
+ await MCPServerOAuthClientRepository(prisma_client).table.update(
+ where={"server_id": oauth_client.server_id},
+ data={"credentials": rotated_credentials},
+ )
+ oauth_updated += 1
+
verbose_proxy_logger.info(
- "rotate_mcp_server_credentials_master_key: rotated %d MCP server row(s)",
+ "rotate_mcp_server_credentials_master_key: rotated %d MCP server row(s) and %d OAuth-client row(s)",
updated,
+ oauth_updated,
)
diff --git a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py
index 54aff86aab2..1af64749304 100644
--- a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py
+++ b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py
@@ -971,43 +971,93 @@ def _apply_persisted_dcr_credentials(mcp_server: MCPServer, credentials: _Persis
return True
-async def _get_persisted_mcp_server_with_dcr_client_id(
- mcp_server: MCPServer,
-) -> Optional[tuple["LiteLLM_MCPServerTable", _PersistedDcrCredentials]]:
- from litellm.proxy._experimental.mcp_server.db import get_mcp_server # noqa: PLC0415
- from litellm.proxy.utils import get_prisma_client_or_throw # noqa: PLC0415
+async def _load_store_dcr_credentials(mcp_server: MCPServer) -> _PersistedDcrCredentials | None:
+ """DCR client persisted in the server-scoped OAuth-client store for a config-declared server
+ (which has no LiteLLM_MCPServerTable row). Returns None when the store has no usable client_id
+ or the DB is unreachable."""
+ from litellm.proxy._experimental.mcp_server.db import ( # noqa: PLC0415 # avoids circular import
+ get_mcp_server_oauth_client_credentials,
+ )
+ from litellm.proxy.utils import get_prisma_client_or_throw # noqa: PLC0415 # avoids circular import
try:
prisma_client = get_prisma_client_or_throw("Database not connected. Cannot read MCP OAuth client registration.")
- persisted_mcp_server = await get_mcp_server(
- prisma_client=prisma_client,
- server_id=mcp_server.server_id,
+ blob = await get_mcp_server_oauth_client_credentials(
+ prisma_client=prisma_client, server_id=mcp_server.server_id
)
- except Exception as exc: # noqa: BLE001
+ except Exception as exc: # noqa: BLE001 # best-effort read; DB may be unreachable
verbose_logger.debug(
- "register_client_with_server: failed to read persisted DCR client registration for server_id=%s: %s",
+ "register_client_with_server: failed to read stored DCR client for server_id=%s: %s",
mcp_server.server_id,
exc,
)
return None
- if persisted_mcp_server is None:
- return None
-
- credentials = _get_persisted_dcr_credentials(persisted_mcp_server.credentials)
+ credentials = _get_persisted_dcr_credentials(blob)
if credentials is None or not credentials.client_id:
return None
+ return credentials
- return persisted_mcp_server, credentials
+
+async def hydrate_config_server_dcr_client(mcp_server: MCPServer) -> bool:
+ """Overlay a config-declared server's persisted DCR client onto its in-memory object so token
+ refresh can authenticate. Config.yaml servers have no LiteLLM_MCPServerTable row, so their
+ minted client lives in the server-scoped store; without this overlay the in-memory server
+ carries no client_id after a restart. An explicit client_id set in config.yaml wins and is never
+ overwritten by a persisted store client."""
+ if mcp_server.client_id:
+ return False
+ credentials = await _load_store_dcr_credentials(mcp_server)
+ if credentials is None:
+ return False
+ return _apply_persisted_dcr_credentials(mcp_server, credentials)
+
+
+async def _resolve_persisted_dcr_client(
+ mcp_server: MCPServer,
+) -> tuple[Optional["LiteLLM_MCPServerTable"], _PersistedDcrCredentials | None]:
+ """Resolve a server's persisted DCR client using the same two-level rule the write path uses, so
+ read and write always agree. First, whether the server HAS a LiteLLM_MCPServerTable row: a row is
+ always resolved to that row and the store is never consulted for a server that has a row, so a
+ caller-chosen server_id colliding with a config-declared server cannot inherit that config
+ server's client, and a row that exists but carries no usable client_id yields (row, None) rather
+ than a store fallback. Second, among rowless servers: a config-declared server keeps its client in
+ the server-scoped store, while a rowless non-config server is a throwaway temp/session server with
+ no persisted client. Returns (row_or_None, credentials_or_None); the row is only needed by the
+ reuse path to refresh the registry for a DB-declared server."""
+ from litellm.proxy._experimental.mcp_server.db import get_mcp_server # noqa: PLC0415 # avoids circular import
+ from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( # noqa: PLC0415 # avoids circular import
+ global_mcp_server_manager,
+ )
+ from litellm.proxy.utils import get_prisma_client_or_throw # noqa: PLC0415 # avoids circular import
+
+ try:
+ prisma_client = get_prisma_client_or_throw("Database not connected. Cannot read MCP OAuth client registration.")
+ row = await get_mcp_server(prisma_client=prisma_client, server_id=mcp_server.server_id)
+ except Exception as exc: # noqa: BLE001 # best-effort read; DB may be unreachable
+ verbose_logger.debug(
+ "register_client_with_server: failed to read persisted DCR client for server_id=%s: %s",
+ mcp_server.server_id,
+ exc,
+ )
+ return None, None
+
+ if row is not None:
+ credentials = _get_persisted_dcr_credentials(row.credentials)
+ if credentials is not None and credentials.client_id:
+ return row, credentials
+ return row, None
+ if global_mcp_server_manager.is_config_declared_server(mcp_server.server_id):
+ return None, await _load_store_dcr_credentials(mcp_server)
+ return None, None
async def _reuse_persisted_dcr_client_if_available(
mcp_server: MCPServer, current_redirect_uri: Optional[str] = None
) -> bool:
- persisted = await _get_persisted_mcp_server_with_dcr_client_id(mcp_server)
- if persisted is None:
+ persisted_mcp_server, credentials = await _resolve_persisted_dcr_client(mcp_server)
+ if credentials is None:
return False
- persisted_mcp_server, credentials = persisted
if current_redirect_uri is not None and _redirect_uri_not_registered(credentials, current_redirect_uri):
verbose_logger.debug(
"register_client_with_server: not reusing persisted DCR client for server_id=%s; its registered "
@@ -1021,18 +1071,19 @@ async def _reuse_persisted_dcr_client_if_available(
if not _apply_persisted_dcr_credentials(mcp_server, credentials):
return False
- from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( # noqa: PLC0415
- global_mcp_server_manager,
- )
-
- try:
- await global_mcp_server_manager.update_server(persisted_mcp_server)
- except Exception as exc: # noqa: BLE001
- verbose_logger.warning(
- "register_client_with_server: failed to refresh persisted DCR client registration for server_id=%s: %s",
- mcp_server.server_id,
- exc,
+ if persisted_mcp_server is not None:
+ from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( # noqa: PLC0415 # avoids circular import
+ global_mcp_server_manager,
)
+
+ try:
+ await global_mcp_server_manager.update_server(persisted_mcp_server)
+ except Exception as exc: # noqa: BLE001 # best-effort registry refresh
+ verbose_logger.warning(
+ "register_client_with_server: failed to refresh persisted DCR client registration for server_id=%s: %s",
+ mcp_server.server_id,
+ exc,
+ )
return bool(mcp_server.client_id)
@@ -1044,10 +1095,9 @@ async def _persisted_dcr_redirect_uri_is_stale(mcp_server: MCPServer, current_re
otherwise short-circuits registration before any redirect check can run. Servers
without a persisted DCR recording (admin-configured client_id, or registered before
redirect_uris were recorded) are never reported stale."""
- persisted = await _get_persisted_mcp_server_with_dcr_client_id(mcp_server)
- if persisted is None:
+ _, credentials = await _resolve_persisted_dcr_client(mcp_server)
+ if credentials is None:
return False
- _, credentials = persisted
if not _redirect_uri_not_registered(credentials, current_redirect_uri):
return False
verbose_logger.warning(
@@ -1067,7 +1117,10 @@ DcrRegistrationPersistenceResult = Literal["persisted", "reused", "skipped", "fa
async def _persist_dcr_client_registration(
mcp_server: MCPServer, registration_response: object, current_redirect_uri: str
) -> DcrRegistrationPersistenceResult:
- """Persist the dynamically registered OAuth client (RFC 7591) onto the MCP server row.
+ """Persist the dynamically registered OAuth client (RFC 7591) to its single home: the server's
+ ``LiteLLM_MCPServerTable`` row when it has one, otherwise the server-scoped store when the server
+ is config-declared. A rowless server that is not config-declared is a throwaway temp/session
+ server, so its client is overlaid in memory only and not persisted.
The interactive authorization_code flow mints a ``client_id`` via Dynamic Client
Registration that discovery cannot re-derive; without persisting it the autonomous
@@ -1106,16 +1159,20 @@ async def _persist_dcr_client_registration(
if await _reuse_persisted_dcr_client_if_available(mcp_server, current_redirect_uri=current_redirect_uri):
return "reused"
+ token_endpoint_auth_method = (
+ "client_secret_basic" if registration.token_endpoint_auth_method == "client_secret_basic" else None
+ )
credentials: MCPCredentials = {
"client_id": registration.client_id,
"client_secret": registration.client_secret,
- "token_endpoint_auth_method": (
- "client_secret_basic" if registration.token_endpoint_auth_method == "client_secret_basic" else None
- ),
+ "token_endpoint_auth_method": token_endpoint_auth_method,
"redirect_uris": [current_redirect_uri],
}
- from litellm.proxy._experimental.mcp_server.db import update_mcp_server # noqa: PLC0415
+ from litellm.proxy._experimental.mcp_server.db import ( # noqa: PLC0415 # avoids circular import
+ update_mcp_server,
+ upsert_mcp_server_oauth_client_credentials,
+ )
from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( # noqa: PLC0415
global_mcp_server_manager,
)
@@ -1136,7 +1193,18 @@ async def _persist_dcr_client_registration(
),
touched_by="mcp_oauth_dcr",
)
- await global_mcp_server_manager.update_server(updated_row)
+ if updated_row is not None:
+ await global_mcp_server_manager.update_server(updated_row)
+ return "persisted"
+ if global_mcp_server_manager.is_config_declared_server(mcp_server.server_id):
+ await upsert_mcp_server_oauth_client_credentials(
+ prisma_client=prisma_client,
+ server_id=mcp_server.server_id,
+ credentials=credentials,
+ )
+ mcp_server.client_id = registration.client_id
+ mcp_server.client_secret = registration.client_secret
+ mcp_server.token_endpoint_auth_method = token_endpoint_auth_method
return "persisted"
except Exception as exc: # noqa: BLE001
verbose_logger.warning(
diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py
index 115ff2e492c..941bd45f98c 100644
--- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py
+++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py
@@ -1127,6 +1127,14 @@ class MCPServerManager:
"""
return self.config_mcp_servers | self.registry
+ def is_config_declared_server(self, server_id: str) -> bool:
+ """True when server_id was declared in config.yaml (present in the in-memory config map).
+ Config servers are rowless and persistent, so their DCR client belongs in the server-scoped
+ store; a rowless server that is NOT config-declared is a throwaway temp/session server whose
+ client must not be persisted. This never overrides the row-existence check: a server that has
+ a LiteLLM_MCPServerTable row is always resolved to that row first."""
+ return server_id in self.config_mcp_servers
+
async def load_servers_from_config(
self,
mcp_servers_config: dict[str, Any],
@@ -1367,8 +1375,32 @@ class MCPServerManager:
verbose_logger.debug(f"Loaded MCP Servers: {json.dumps(self.config_mcp_servers, indent=4, default=str)}")
+ await self._hydrate_config_servers_dcr_clients()
+
self.initialize_tool_name_to_mcp_server_name_mapping()
+ async def _hydrate_config_servers_dcr_clients(self) -> None:
+ """Overlay each config-declared server's persisted DCR client (from the server-scoped
+ store) onto its in-memory object so token refresh authenticates after a restart. A
+ best-effort no-op when the DB is unreachable at config-load time."""
+ from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( # noqa: PLC0415 # circular import
+ hydrate_config_server_dcr_client,
+ )
+
+ for server in self.config_mcp_servers.values():
+ try:
+ if await hydrate_config_server_dcr_client(server):
+ verbose_logger.debug(
+ "hydrated persisted DCR client onto config MCP server server_id=%s",
+ server.server_id,
+ )
+ except Exception as exc: # noqa: BLE001 # best-effort hydration; never fail config load
+ verbose_logger.debug(
+ "load_servers_from_config: failed to hydrate DCR client for server_id=%s: %s",
+ server.server_id,
+ exc,
+ )
+
async def _register_openapi_tools(self, spec_path: str, server: MCPServer, base_url: str):
"""
Register tools from an OpenAPI specification for a given server.
@@ -4968,6 +5000,8 @@ class MCPServerManager:
verbose_logger.debug("MCP registry refreshed (%s servers in registry)", len(registered_registry))
+ await self._hydrate_config_servers_dcr_clients()
+
def get_mcp_servers_from_ids(self, server_ids: list[str]) -> list[MCPServer]:
servers = []
registry = self.get_registry()
diff --git a/litellm/proxy/schema.prisma b/litellm/proxy/schema.prisma
index f842bf13da9..a99cec49417 100644
--- a/litellm/proxy/schema.prisma
+++ b/litellm/proxy/schema.prisma
@@ -396,6 +396,13 @@ model LiteLLM_MCPUserEnvVars {
@@index([server_id])
}
+model LiteLLM_MCPServerOAuthClient {
+ server_id String @id
+ credentials Json?
+ created_at DateTime @default(now()) @map("created_at")
+ updated_at DateTime @default(now()) @updatedAt @map("updated_at")
+}
+
// Generate Tokens for Proxy
model LiteLLM_VerificationToken {
token String @id
diff --git a/litellm/repositories/table_repositories.py b/litellm/repositories/table_repositories.py
index 7ce4607e1ca..dc2a7d25259 100644
--- a/litellm/repositories/table_repositories.py
+++ b/litellm/repositories/table_repositories.py
@@ -77,6 +77,10 @@ class MCPUserCredentialsRepository(PrismaTableRepository):
table_name = "litellm_mcpusercredentials"
+class MCPServerOAuthClientRepository(PrismaTableRepository):
+ table_name = "litellm_mcpserveroauthclient"
+
+
class PromptRepository(PrismaTableRepository):
table_name = "litellm_prompttable"
diff --git a/schema.prisma b/schema.prisma
index f842bf13da9..a99cec49417 100644
--- a/schema.prisma
+++ b/schema.prisma
@@ -396,6 +396,13 @@ model LiteLLM_MCPUserEnvVars {
@@index([server_id])
}
+model LiteLLM_MCPServerOAuthClient {
+ server_id String @id
+ credentials Json?
+ created_at DateTime @default(now()) @map("created_at")
+ updated_at DateTime @default(now()) @updatedAt @map("updated_at")
+}
+
// Generate Tokens for Proxy
model LiteLLM_VerificationToken {
token String @id
diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_db_credentials.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_db_credentials.py
index 7269774442b..a245200c4d1 100644
--- a/tests/test_litellm/proxy/_experimental/mcp_server/test_db_credentials.py
+++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_db_credentials.py
@@ -978,3 +978,67 @@ def test_prepare_mcp_server_data_update_carries_token_exchange_columns():
assert data["audience"] == "https://upstream.example.com"
assert data["subject_token_type"] == "urn:ietf:params:oauth:token-type:jwt"
assert data["token_exchange_profile"] == "entra_obo"
+
+
+@pytest.mark.asyncio
+async def test_master_key_rotation_reencrypts_oauth_client_store(monkeypatch):
+ """The server-scoped DCR client store (LiteLLM_MCPServerOAuthClient) is encrypted at rest, so a
+ master-key rotation must re-encrypt it alongside the server rows. Skipping it leaves
+ config-declared DCR clients under the retired key, where they decrypt back to ciphertext and
+ force a full re-authorization."""
+ import litellm.proxy.common_utils.encrypt_decrypt_utils as enc
+ from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
+ from litellm.proxy._experimental.mcp_server.db import (
+ decrypt_credentials,
+ encrypt_credentials,
+ rotate_mcp_server_credentials_master_key,
+ )
+
+ key_old, key_new = "salt-old-key", "salt-new-key"
+
+ blob_old = safe_dumps(
+ encrypt_credentials(
+ credentials={"client_id": "cid-123", "client_secret": "sec-456"},
+ encryption_key=key_old,
+ )
+ )
+
+ monkeypatch.setattr(enc, "_get_salt_key", lambda: key_old)
+
+ prisma = MagicMock()
+ prisma.db.litellm_mcpservertable.find_many = AsyncMock(return_value=[])
+ prisma.db.litellm_mcpserveroauthclient.find_many = AsyncMock(
+ return_value=[SimpleNamespace(server_id="config_faros", credentials=blob_old)]
+ )
+ store_update = AsyncMock()
+ prisma.db.litellm_mcpserveroauthclient.update = store_update
+
+ await rotate_mcp_server_credentials_master_key(prisma, touched_by="test", new_master_key=key_new)
+
+ store_update.assert_awaited_once()
+ assert store_update.await_args.kwargs["where"] == {"server_id": "config_faros"}
+ rotated_blob = store_update.await_args.kwargs["data"]["credentials"]
+
+ monkeypatch.setattr(enc, "_get_salt_key", lambda: key_new)
+ recovered = decrypt_credentials(credentials=json.loads(rotated_blob))
+ assert recovered["client_id"] == "cid-123"
+ assert recovered["client_secret"] == "sec-456"
+
+
+@pytest.mark.asyncio
+async def test_delete_mcp_server_cleans_oauth_client_store():
+ """Deleting a server must remove its server-scoped DCR client store entry alongside the per-user
+ credential and env-var rows, or a re-created server reusing the same server_id would inherit the
+ deleted server's OAuth client."""
+ from litellm.proxy._experimental.mcp_server.db import delete_mcp_server
+
+ prisma = MagicMock()
+ prisma.db.litellm_mcpservertable.delete = AsyncMock(return_value=SimpleNamespace(server_id="s1"))
+ prisma.db.litellm_mcpusercredentials.find_many = AsyncMock(return_value=[])
+ prisma.db.litellm_mcpusercredentials.delete_many = AsyncMock()
+ prisma.db.litellm_mcpuserenvvars.delete_many = AsyncMock()
+ prisma.db.litellm_mcpserveroauthclient.delete_many = AsyncMock()
+
+ await delete_mcp_server(prisma, "s1", invalidate_token_cache=AsyncMock())
+
+ prisma.db.litellm_mcpserveroauthclient.delete_many.assert_awaited_once_with(where={"server_id": "s1"})
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 f5ac229d119..6f2f24df8fa 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
@@ -7130,3 +7130,373 @@ async def test_token_exchange_unreadable_body_still_renders_oauth_fault():
assert response.status_code == 502
body = json.loads(response.body)
assert body == {"error": "server_error", "error_description": "upstream token endpoint returned HTTP 400"}
+
+
+@pytest.mark.asyncio
+async def test_persist_dcr_client_for_config_server_uses_side_store():
+ """A config.yaml-declared OAuth2 DCR server has no LiteLLM_MCPServerTable row, so
+ update_mcp_server returns None. The minted client must then persist to the server-scoped
+ OAuth-client store keyed by server_id (never a shadow server row), overlay onto the in-memory
+ server so refresh can authenticate this process, and never call update_server(None) (which
+ previously raised AttributeError on .approval_status, was swallowed, and reported a 200 that
+ persisted nothing)."""
+ from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
+ _persist_dcr_client_registration,
+ )
+ from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
+ global_mcp_server_manager,
+ )
+ from litellm.proxy._types import MCPTransport
+ from litellm.types.mcp_server.mcp_server_manager import MCPServer
+
+ config_server = MCPServer(
+ server_id="config_faros",
+ name="config_faros",
+ server_name="config_faros",
+ transport=MCPTransport.http,
+ auth_type=MCPAuth.oauth2,
+ client_id=None,
+ client_secret=None,
+ authorization_url="https://provider.example/oauth/authorize",
+ token_url="https://provider.example/oauth/token",
+ registration_url="https://provider.example/oauth/register",
+ )
+
+ mock_upsert = AsyncMock()
+ mock_update_server = AsyncMock()
+
+ with (
+ patch.object(global_mcp_server_manager, "is_config_declared_server", return_value=True),
+ patch("litellm.proxy.utils.get_prisma_client_or_throw", return_value=MagicMock()),
+ patch(
+ "litellm.proxy._experimental.mcp_server.db.update_mcp_server",
+ new=AsyncMock(return_value=None),
+ ),
+ patch(
+ "litellm.proxy._experimental.mcp_server.db.get_mcp_server",
+ new=AsyncMock(return_value=None),
+ ),
+ patch(
+ "litellm.proxy._experimental.mcp_server.db.get_mcp_server_oauth_client_credentials",
+ new=AsyncMock(return_value=None),
+ ),
+ patch(
+ "litellm.proxy._experimental.mcp_server.db.upsert_mcp_server_oauth_client_credentials",
+ new=mock_upsert,
+ ),
+ patch.object(global_mcp_server_manager, "update_server", new=mock_update_server),
+ ):
+ result = await _persist_dcr_client_registration(
+ mcp_server=config_server,
+ registration_response={
+ "client_id": "minted-client",
+ "client_secret": "minted-secret",
+ "token_endpoint_auth_method": "client_secret_basic",
+ },
+ current_redirect_uri="https://proxy.litellm.example/callback",
+ )
+
+ assert result == "persisted"
+
+ mock_upsert.assert_called_once()
+ assert mock_upsert.call_args.kwargs["server_id"] == "config_faros"
+ stored = mock_upsert.call_args.kwargs["credentials"]
+ assert stored["client_id"] == "minted-client"
+ assert stored["client_secret"] == "minted-secret"
+ assert stored["token_endpoint_auth_method"] == "client_secret_basic"
+ assert stored["redirect_uris"] == ["https://proxy.litellm.example/callback"]
+
+ assert config_server.client_id == "minted-client"
+ assert config_server.client_secret == "minted-secret"
+ assert config_server.token_endpoint_auth_method == "client_secret_basic"
+
+ mock_update_server.assert_not_called()
+
+
+@pytest.mark.asyncio
+async def test_hydrate_config_server_applies_stored_dcr_client(monkeypatch):
+ """On restart a config server's in-memory object has no client_id; hydration overlays the
+ persisted DCR client from the server-scoped store, decrypting the encrypted-at-rest blob, so the
+ refresh_token grant can authenticate as the registered client instead of re-authenticating."""
+ import litellm.proxy.common_utils.encrypt_decrypt_utils as enc
+ from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
+ from litellm.proxy._experimental.mcp_server.db import encrypt_credentials
+ from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
+ hydrate_config_server_dcr_client,
+ )
+ from litellm.proxy._types import MCPTransport
+ from litellm.types.mcp_server.mcp_server_manager import MCPServer
+
+ server = MCPServer(
+ server_id="config_faros",
+ name="config_faros",
+ server_name="config_faros",
+ transport=MCPTransport.http,
+ auth_type=MCPAuth.oauth2,
+ client_id=None,
+ )
+
+ monkeypatch.setattr(enc, "_get_salt_key", lambda: "salt-hydrate-key")
+ stored_blob = safe_dumps(
+ encrypt_credentials(
+ credentials={
+ "client_id": "stored-client",
+ "client_secret": "stored-secret",
+ "token_endpoint_auth_method": "client_secret_basic",
+ "redirect_uris": ["https://proxy.litellm.example/callback"],
+ },
+ encryption_key="salt-hydrate-key",
+ )
+ )
+ assert "stored-client" not in stored_blob and "stored-secret" not in stored_blob
+
+ with (
+ patch("litellm.proxy.utils.get_prisma_client_or_throw", return_value=MagicMock()),
+ patch(
+ "litellm.proxy._experimental.mcp_server.db.get_mcp_server_oauth_client_credentials",
+ new=AsyncMock(return_value=stored_blob),
+ ),
+ ):
+ applied = await hydrate_config_server_dcr_client(server)
+
+ assert applied is True
+ assert server.client_id == "stored-client"
+ assert server.client_secret == "stored-secret"
+ assert server.token_endpoint_auth_method == "client_secret_basic"
+
+
+@pytest.mark.asyncio
+async def test_reuse_config_server_reads_store_with_real_crypto(monkeypatch):
+ """A config-declared server (rowless) keeps its DCR client in the store, so the reuse read
+ resolves it from the store and decrypts the encrypted-at-rest client, mirroring the write path so
+ a re-authorize reuses the client instead of re-minting one."""
+ import litellm.proxy.common_utils.encrypt_decrypt_utils as enc
+ from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
+ from litellm.proxy._experimental.mcp_server.db import encrypt_credentials
+ from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
+ _reuse_persisted_dcr_client_if_available,
+ )
+ from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
+ global_mcp_server_manager,
+ )
+ from litellm.proxy._types import MCPTransport
+ from litellm.types.mcp_server.mcp_server_manager import MCPServer
+
+ server = MCPServer(
+ server_id="config_faros",
+ name="config_faros",
+ server_name="config_faros",
+ transport=MCPTransport.http,
+ auth_type=MCPAuth.oauth2,
+ client_id=None,
+ )
+
+ monkeypatch.setattr(enc, "_get_salt_key", lambda: "salt-reuse-key")
+ blob = safe_dumps(
+ encrypt_credentials(
+ credentials={"client_id": "stored-client", "client_secret": "sec", "redirect_uris": ["https://x/callback"]},
+ encryption_key="salt-reuse-key",
+ )
+ )
+ assert "stored-client" not in blob
+ store_lookup = AsyncMock(return_value=blob)
+ with (
+ patch.object(global_mcp_server_manager, "is_config_declared_server", return_value=True),
+ patch("litellm.proxy.utils.get_prisma_client_or_throw", return_value=MagicMock()),
+ patch("litellm.proxy._experimental.mcp_server.db.get_mcp_server", new=AsyncMock(return_value=None)),
+ patch(
+ "litellm.proxy._experimental.mcp_server.db.get_mcp_server_oauth_client_credentials",
+ new=store_lookup,
+ ),
+ ):
+ result = await _reuse_persisted_dcr_client_if_available(server, current_redirect_uri="https://x/callback")
+
+ assert result is True
+ assert server.client_id == "stored-client"
+ store_lookup.assert_awaited_once()
+
+
+@pytest.mark.asyncio
+async def test_temp_server_is_not_persisted_to_store():
+ """A rowless server that is NOT config-declared (a throwaway /server/oauth/session server) must
+ not leave a permanent store row on persist, and the read must never consult the store for it. Its
+ minted client is overlaid in memory for the session only."""
+ from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
+ _persist_dcr_client_registration,
+ _reuse_persisted_dcr_client_if_available,
+ )
+ from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
+ global_mcp_server_manager,
+ )
+ from litellm.proxy._types import MCPTransport
+ from litellm.types.mcp_server.mcp_server_manager import MCPServer
+
+ temp = MCPServer(
+ server_id="temp-uuid",
+ name="temp",
+ server_name="temp",
+ transport=MCPTransport.http,
+ auth_type=MCPAuth.oauth2,
+ client_id=None,
+ authorization_url="https://p.example/authorize",
+ token_url="https://p.example/token",
+ registration_url="https://p.example/register",
+ )
+
+ upsert = AsyncMock()
+ store_read = AsyncMock(return_value=None)
+ with (
+ patch.object(global_mcp_server_manager, "is_config_declared_server", return_value=False),
+ patch("litellm.proxy.utils.get_prisma_client_or_throw", return_value=MagicMock()),
+ patch("litellm.proxy._experimental.mcp_server.db.update_mcp_server", new=AsyncMock(return_value=None)),
+ patch("litellm.proxy._experimental.mcp_server.db.get_mcp_server", new=AsyncMock(return_value=None)),
+ patch("litellm.proxy._experimental.mcp_server.db.upsert_mcp_server_oauth_client_credentials", new=upsert),
+ patch(
+ "litellm.proxy._experimental.mcp_server.db.get_mcp_server_oauth_client_credentials",
+ new=store_read,
+ ),
+ patch.object(global_mcp_server_manager, "update_server", new=AsyncMock()),
+ ):
+ result = await _persist_dcr_client_registration(
+ temp, {"client_id": "temp-client", "client_secret": "s"}, "https://x/callback"
+ )
+ reused = await _reuse_persisted_dcr_client_if_available(
+ MCPServer(
+ server_id="temp-uuid",
+ name="temp",
+ server_name="temp",
+ transport=MCPTransport.http,
+ auth_type=MCPAuth.oauth2,
+ client_id=None,
+ ),
+ current_redirect_uri="https://x/callback",
+ )
+
+ assert result == "persisted"
+ assert temp.client_id == "temp-client"
+ upsert.assert_not_called()
+ store_read.assert_not_called()
+ assert reused is False
+
+
+@pytest.mark.asyncio
+async def test_hydrate_does_not_overwrite_explicit_config_client_id():
+ """An explicit client_id set in config.yaml wins: hydration must not overwrite it with a stale
+ persisted store client, and must not even read the store when config already supplied a client."""
+ from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
+ hydrate_config_server_dcr_client,
+ )
+ from litellm.proxy._types import MCPTransport
+ from litellm.types.mcp_server.mcp_server_manager import MCPServer
+
+ server = MCPServer(
+ server_id="config_static",
+ name="config_static",
+ server_name="config_static",
+ transport=MCPTransport.http,
+ auth_type=MCPAuth.oauth2,
+ client_id="explicit-from-config",
+ )
+ store_read = AsyncMock(
+ return_value={"client_id": "stale-store-client", "client_secret": "x", "redirect_uris": []}
+ )
+ with (
+ patch("litellm.proxy.utils.get_prisma_client_or_throw", return_value=MagicMock()),
+ patch(
+ "litellm.proxy._experimental.mcp_server.db.get_mcp_server_oauth_client_credentials",
+ new=store_read,
+ ),
+ ):
+ applied = await hydrate_config_server_dcr_client(server)
+
+ assert applied is False
+ assert server.client_id == "explicit-from-config"
+ store_read.assert_not_called()
+
+
+@pytest.mark.asyncio
+async def test_reuse_does_not_inherit_store_client_when_a_row_exists():
+ """Security: a server that HAS a LiteLLM_MCPServerTable row reads its DCR client only from that
+ row, never from the server-scoped store. server_id is caller-settable on create, so a submitted
+ server whose id collides with a config-declared server must not be able to load that config
+ server's client from the store and send it to its own token endpoint. A row that exists but has
+ no client_id yields no reusable client and must not fall back to the store."""
+ from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
+ _reuse_persisted_dcr_client_if_available,
+ )
+ from litellm.proxy._types import MCPTransport
+ from litellm.types.mcp_server.mcp_server_manager import MCPServer
+
+ submitted = MCPServer(
+ server_id="collides_with_config",
+ name="submitted",
+ server_name="submitted",
+ transport=MCPTransport.http,
+ auth_type=MCPAuth.oauth2,
+ client_id=None,
+ )
+
+ row_without_client = MagicMock()
+ row_without_client.credentials = None
+ row_without_client.server_id = "collides_with_config"
+ store_lookup = AsyncMock(
+ return_value={"client_id": "config-secret-client", "client_secret": "leak", "redirect_uris": []}
+ )
+ with (
+ patch("litellm.proxy.utils.get_prisma_client_or_throw", return_value=MagicMock()),
+ patch(
+ "litellm.proxy._experimental.mcp_server.db.get_mcp_server",
+ new=AsyncMock(return_value=row_without_client),
+ ),
+ patch(
+ "litellm.proxy._experimental.mcp_server.db.get_mcp_server_oauth_client_credentials",
+ new=store_lookup,
+ ),
+ ):
+ result = await _reuse_persisted_dcr_client_if_available(submitted, current_redirect_uri="https://x/callback")
+
+ assert result is False
+ assert submitted.client_id is None
+ store_lookup.assert_not_called()
+
+
+@pytest.mark.asyncio
+async def test_load_servers_from_config_hydrates_dcr_clients():
+ """load_servers_from_config must invoke DCR-client hydration so config servers pick up their
+ persisted client on startup; deleting the call site leaves a restarted server with no client_id
+ and forces re-authentication on every token expiry."""
+ from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
+ global_mcp_server_manager,
+ )
+
+ hydrate_spy = AsyncMock()
+ with patch.object(global_mcp_server_manager, "_hydrate_config_servers_dcr_clients", new=hydrate_spy):
+ await global_mcp_server_manager.load_servers_from_config({})
+
+ hydrate_spy.assert_awaited_once()
+
+
+@pytest.mark.asyncio
+async def test_reload_servers_from_database_hydrates_dcr_clients():
+ """load_servers_from_config runs before the DB connects at startup, so its hydration no-ops;
+ reload_servers_from_database runs after the DB connects and must hydrate config servers' persisted
+ DCR clients too, or a fresh pod has no client_id for a config server and forces re-authentication
+ on the first token refresh."""
+ from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
+ global_mcp_server_manager,
+ )
+
+ prisma = MagicMock()
+ prisma.db.litellm_mcpservertable.find_many = AsyncMock(return_value=[])
+
+ hydrate_spy = AsyncMock()
+ with (
+ patch(
+ "litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw",
+ return_value=prisma,
+ ),
+ patch.object(global_mcp_server_manager, "_hydrate_config_servers_dcr_clients", new=hydrate_spy),
+ ):
+ await global_mcp_server_manager.reload_servers_from_database()
+
+ hydrate_spy.assert_awaited_once()
diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_sigv4_auth.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_sigv4_auth.py
index 5992fd1814f..f6b61c1d9f7 100644
--- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_sigv4_auth.py
+++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_sigv4_auth.py
@@ -989,6 +989,7 @@ class TestRotateCredentials:
mock_prisma = MagicMock()
mock_prisma.db.litellm_mcpservertable.find_many = AsyncMock(return_value=[server])
mock_prisma.db.litellm_mcpservertable.update = AsyncMock()
+ mock_prisma.db.litellm_mcpserveroauthclient.find_many = AsyncMock(return_value=[])
with (
patch(
@@ -1036,6 +1037,7 @@ class TestRotateCredentials:
mock_prisma = MagicMock()
mock_prisma.db.litellm_mcpservertable.find_many = AsyncMock(return_value=[server])
mock_prisma.db.litellm_mcpservertable.update = AsyncMock()
+ mock_prisma.db.litellm_mcpserveroauthclient.find_many = AsyncMock(return_value=[])
with (
patch(
From dbb5b813c1e70329ea977ceec59831cba4ef4522 Mon Sep 17 00:00:00 2001
From: ryan-crabbe-berri
Date: Fri, 17 Jul 2026 20:02:22 -0700
Subject: [PATCH 021/296] test(e2e): budget reset diagonal for team, org, user,
and #32005 team-member keys (#33771)
* test(e2e): budget reset diagonal for team, org, user, and #32005 team-member keys
Adds E2E-7/8/10/11 from the budget-level x key-kind coverage matrix: each budget
level serves traffic again after its budget_duration window elapses, walking the
same ladder as the enforcement diagonal. New registry rows and tests cover the
team, organization, and internal-user reset rungs, plus the #32005 interplay
where a team-member key frozen by its owner's user budget comes back when the
user's window renews; the bare-key and per-team-member rungs already had coverage
Each case isolates the cap to one entity, drives spend to a budget_exceeded
block, then polls past the window until a call succeeds, holding every refusal
as a budget block so a reset that no-ops (stays blocked forever) or crashes
(leaks a 5xx) fails the test. budget_duration becomes an optional param on the
budget_client create_team / create_user / create_org helpers
* test(e2e): fold the reset diagonal into test_budget_reset_e2e.py and address greptile nits
Move the team / org / user / #32005 reset cases out of the standalone
test_budget_reset_diagonal_e2e.py and into test_budget_reset_e2e.py, absorbing
the pre-existing bare-key reset into the same TestBudgetResetDiagonal spec class
so the whole reset ladder reads as one file (mirroring how the enforcement
diagonal lives in test_budget_enforcement_e2e.py) and the drive/poll helpers are
defined once instead of duplicated across reset files.
Greptile nits: bound the drive phase to under one window (12 attempts x 2s < 30s)
so a block is observed before the reset job can fire, and replace the bare assert
in the poll loop with a pytest.fail that prints the HTTP status, so a provider 429
or a crashed reset path is distinguishable from a budget block at a glance.
* test(e2e): trim reset diagonal docstrings back to the file's original style
* test(e2e): inline single-use drive-loop bounds
* test(e2e): cut the reset module docstring to one line
* test(e2e): make the org reset test wait for a scheduled window (bugbot)
/organization/new stores budget_duration without scheduling budget_reset_at, so
the reset job's NULL catch-up branch zeroes org spend on its first 5-10s tick;
the org reset test could pass off that catch-up instead of a real window roll
(tracked as LIT-4570). The test now reads the org's budget_id and polls
/budget/info until budget_reset_at is scheduled before driving spend, so the
recovery it observes can only come from a genuine window expiry. Verified live:
the org case now runs ~33s (a full window) instead of beating the rescheduler
---
.../coverage_registry/quota_management.yaml | 3 +
.../quota_management/budgets/budget_client.py | 41 +++++-
.../budgets/test_budget_reset_e2e.py | 127 +++++++++++++-----
3 files changed, 132 insertions(+), 39 deletions(-)
diff --git a/tests/e2e/coverage_registry/quota_management.yaml b/tests/e2e/coverage_registry/quota_management.yaml
index 0d61f48703d..633351ca97c 100644
--- a/tests/e2e/coverage_registry/quota_management.yaml
+++ b/tests/e2e/coverage_registry/quota_management.yaml
@@ -18,6 +18,9 @@
- {id: quota_management.budget.model_max.isolates_per_model, module: quota_management, tier: P1, behavior: budget, variant: model_max, assertions: [isolates_per_model], exercised_on: [chat_completions], source: "proxy/hooks/model_max_budget_limiter.py", rationale: "model_max_budget caps one model without touching a sibling's budget"}
- {id: quota_management.budget.soft.alerts_without_blocking, module: quota_management, tier: P1, behavior: budget, variant: soft, assertions: [alerts_without_blocking], exercised_on: [chat_completions], source: "proxy/auth/auth_checks.py", rationale: "soft_budget alerts but never blocks traffic"}
- {id: quota_management.budget.key.resets_after_window, module: quota_management, tier: P1, behavior: budget, variant: key, assertions: [resets_after_window], exercised_on: [chat_completions], source: "proxy/common_utils/reset_budget_job.py", rationale: "budget_duration zeroes key spend after the window; a blocked key serves again"}
+- {id: quota_management.budget.team.resets_after_window, module: quota_management, tier: P1, behavior: budget, variant: team, assertions: [resets_after_window], exercised_on: [chat_completions], source: "proxy/common_utils/reset_budget_job.py", rationale: "budget_duration zeroes a team's spend after the window; every key on the team serves again"}
+- {id: quota_management.budget.organization.resets_after_window, module: quota_management, tier: P1, behavior: budget, variant: organization, assertions: [resets_after_window], exercised_on: [chat_completions], source: "proxy/common_utils/reset_budget_job.py", rationale: "An org budget resets after its window; keys under the org serve again"}
+- {id: quota_management.budget.internal_user.resets_after_window, module: quota_management, tier: P1, behavior: budget, variant: internal_user, assertions: [resets_after_window], exercised_on: [chat_completions], source: "proxy/common_utils/reset_budget_job.py", rationale: "An internal user's budget resets after its window; their personal and team-member keys serve again"}
- {id: quota_management.budget.team_member.resets_after_window, module: quota_management, tier: P1, behavior: budget, variant: team_member, assertions: [resets_after_window], exercised_on: [chat_completions], source: "proxy/common_utils/reset_budget_job.py", rationale: "Member per-team budget reset keeps advancing window after window"}
- {id: quota_management.budget.key_multi_window.blocks_then_resets, module: quota_management, tier: P1, behavior: budget, variant: key_multi_window, assertions: [blocks_then_resets], exercised_on: [chat_completions], source: "proxy/common_utils/reset_budget_job.py", rationale: "budget_limits enforce within a short window and serve again in the next"}
- {id: quota_management.budget.key_multi_window.resets_windows_independently, module: quota_management, tier: P2, behavior: budget, variant: key_multi_window, assertions: [resets_windows_independently], exercised_on: [chat_completions], source: "proxy/common_utils/reset_budget_job.py", rationale: "Each window of a multi-window budget resets on its own schedule"}
diff --git a/tests/e2e/quota_management/budgets/budget_client.py b/tests/e2e/quota_management/budgets/budget_client.py
index 01e4d63c1c3..c2f5dcdd57c 100644
--- a/tests/e2e/quota_management/budgets/budget_client.py
+++ b/tests/e2e/quota_management/budgets/budget_client.py
@@ -33,6 +33,7 @@ _TEAM_READY_SLEEP_SECONDS = 0.4
class UserNewBody(BaseModel):
max_budget: float
+ budget_duration: str | None = None
class UserNewResponse(BaseModel):
@@ -64,6 +65,7 @@ class CustomerNewBody(BaseModel):
class OrgNewBody(BaseModel):
organization_alias: str
max_budget: float
+ budget_duration: str | None = None
class OrgNewResponse(BaseModel):
@@ -74,6 +76,14 @@ class OrgDeleteBody(BaseModel):
organization_ids: list[str]
+class OrgInfoParams(BaseModel):
+ organization_id: str
+
+
+class OrgInfoResponse(BaseModel):
+ budget_id: str | None = None
+
+
class TeamMember(BaseModel):
role: str
user_id: str
@@ -82,6 +92,7 @@ class TeamMember(BaseModel):
class TeamNewBody(BaseModel):
team_alias: str
max_budget: float | None = None
+ budget_duration: str | None = None
organization_id: str | None = None
budget_limits: list[BudgetWindow] | None = None
@@ -257,12 +268,12 @@ class BudgetClient:
# ---- internal user --------------------------------------------------
- def create_user(self, *, max_budget: float) -> str:
+ def create_user(self, *, max_budget: float, budget_duration: str | None = None) -> str:
return unwrap(
self.gateway.transport.post(
"/user/new",
headers=self.gateway.transport.master,
- json=UserNewBody(max_budget=max_budget),
+ json=UserNewBody(max_budget=max_budget, budget_duration=budget_duration),
response_type=UserNewResponse,
)
).user_id
@@ -301,16 +312,36 @@ class BudgetClient:
# ---- organization ---------------------------------------------------
- def create_org(self, *, max_budget: float, alias: str) -> str:
+ def create_org(self, *, max_budget: float, alias: str, budget_duration: str | None = None) -> str:
return unwrap(
self.gateway.transport.post(
"/organization/new",
headers=self.gateway.transport.master,
- json=OrgNewBody(organization_alias=alias, max_budget=max_budget),
+ json=OrgNewBody(
+ organization_alias=alias,
+ max_budget=max_budget,
+ budget_duration=budget_duration,
+ ),
response_type=OrgNewResponse,
)
).organization_id
+ def org_budget_id(self, org_id: str) -> str | None:
+ """The id of the budget row backing an org; its budget_reset_at is read via
+ budget_info (LIT-4570: /organization/new stores budget_duration without
+ scheduling budget_reset_at, so the reset job's first tick schedules it)."""
+ result = self.gateway.transport.get(
+ "/organization/info",
+ headers=self.gateway.transport.master,
+ params=OrgInfoParams(organization_id=org_id),
+ response_type=OrgInfoResponse,
+ )
+ match result:
+ case Success(data=data):
+ return data.budget_id
+ case _:
+ return None
+
def delete_org(self, org_id: str) -> None:
_ = self.gateway.transport.delete(
"/organization/delete",
@@ -326,6 +357,7 @@ class BudgetClient:
*,
alias: str,
max_budget: float | None = None,
+ budget_duration: str | None = None,
organization_id: str | None = None,
budget_limits: list[BudgetWindow] | None = None,
) -> str:
@@ -336,6 +368,7 @@ class BudgetClient:
json=TeamNewBody(
team_alias=alias,
max_budget=max_budget,
+ budget_duration=budget_duration,
organization_id=organization_id,
budget_limits=budget_limits,
),
diff --git a/tests/e2e/quota_management/budgets/test_budget_reset_e2e.py b/tests/e2e/quota_management/budgets/test_budget_reset_e2e.py
index bdbee027f28..b57ad23ebf5 100644
--- a/tests/e2e/quota_management/budgets/test_budget_reset_e2e.py
+++ b/tests/e2e/quota_management/budgets/test_budget_reset_e2e.py
@@ -1,12 +1,4 @@
-"""Live e2e: a key budget resets (zeroes spend) after its budget_duration.
-
-Short budget_duration (30s) + the fast-rescheduled reset job: a key blocked for
-exceeding its max_budget starts succeeding again once the duration elapses and the
-reset job zeroes key.spend. Closes the reset-zeroing gap in
-BUDGET_TEST_COVERAGE_MATRIX.md (reset_budget_for_litellm_keys), which the unit
-suite covers but no live test did - distinct from the per-window reset in
-test_multi_window_budget_e2e.py.
-"""
+"""Live e2e: an entity blocked over its max_budget serves again after its budget_duration window."""
import time
@@ -19,42 +11,107 @@ from lifecycle import ResourceManager
pytestmark = pytest.mark.e2e
+TINY_CAP = 3e-6
+WINDOW = "30s"
+RESET_DEADLINE_SECONDS = 150
+
def _call(client: BudgetClient, key: str):
- return client.chat(
- key, "claude-haiku-4-5", f"reset {unique_marker()}", max_tokens=16
- )
+ return client.chat(key, "claude-haiku-4-5", f"reset {unique_marker()}", max_tokens=16)
-@pytest.mark.covers("quota_management.budget.key.resets_after_window")
-def test_key_budget_resets_after_duration(
- client: BudgetClient, resources: ResourceManager
-) -> None:
- key = client.generate_key(max_budget=3e-6, budget_duration="30s")
- resources.defer(lambda: client.delete_key(key))
-
- # 1. exceed the budget -> litellm returns budget_exceeded
- blocked = False
- for _ in range(20):
+def _drive_to_block(client: BudgetClient, key: str) -> None:
+ """Spend until the cap blocks a call, staying under one window so the block
+ is observed before the reset job can fire; fail hard if enforcement never trips."""
+ for _ in range(12):
result = _call(client, key)
if is_budget_block(result):
- blocked = True
- break
+ return
require_successful_call(result)
time.sleep(2)
- assert blocked, "key budget never enforced"
+ pytest.fail("budget never enforced before the window could reset")
- # 2. once the 30s duration elapses + the reset job runs, key.spend zeroes and
- # calls flow again. The window is wall-clock-aligned, so the reset lands up to
- # a window later, then the rescheduler (~15-20s) zeroes the spend; allow
- # generous headroom over that. A stuck rescheduler is caught by the wait-loop
- # timeout, not this elapsed bound.
- start = time.monotonic()
- while time.monotonic() < start + 150:
+
+def _poll_until_serves_again(client: BudgetClient, key: str) -> None:
+ """Poll past the window until the blocked key serves again; every refusal must
+ stay a budget block, so a crashed reset path or provider error fails loudly."""
+ deadline = time.monotonic() + RESET_DEADLINE_SECONDS
+ while time.monotonic() < deadline:
time.sleep(5)
result = _call(client, key)
if result.ok:
- assert time.monotonic() - start < 120, "reset too slow for a 30s budget"
return
- assert is_budget_block(result), f"non-budget error: {result.body[:200]}"
- pytest.fail("key budget never reset within 150s")
+ if not is_budget_block(result):
+ pytest.fail(f"non-budget error during reset wait: HTTP {result.status_code}: {result.body[:200]}")
+ pytest.fail(f"budget never reset within {RESET_DEADLINE_SECONDS}s")
+
+
+class TestBudgetResetDiagonal:
+ @pytest.mark.covers("quota_management.budget.key.resets_after_window")
+ def test_bare_key_budget_resets_after_window(self, client: BudgetClient, resources: ResourceManager) -> None:
+ key = client.generate_key(max_budget=TINY_CAP, budget_duration=WINDOW)
+ resources.defer(lambda: client.delete_key(key))
+
+ _drive_to_block(client, key)
+ _poll_until_serves_again(client, key)
+
+ @pytest.mark.covers("quota_management.budget.team.resets_after_window")
+ def test_team_budget_resets_after_window(self, client: BudgetClient, resources: ResourceManager) -> None:
+ team_id = client.create_team(
+ alias=f"e2e-team-reset-{unique_marker()}", max_budget=TINY_CAP, budget_duration=WINDOW
+ )
+ resources.defer(lambda: client.delete_team(team_id))
+ key = client.generate_key(team_id=team_id)
+ resources.defer(lambda: client.delete_key(key))
+
+ _drive_to_block(client, key)
+ _poll_until_serves_again(client, key)
+
+ @pytest.mark.covers("quota_management.budget.organization.resets_after_window")
+ def test_org_budget_resets_after_window(self, client: BudgetClient, resources: ResourceManager) -> None:
+ org_id = client.create_org(
+ max_budget=TINY_CAP, alias=f"e2e-org-reset-{unique_marker()}", budget_duration=WINDOW
+ )
+ resources.defer(lambda: client.delete_org(org_id))
+ team_id = client.create_team(alias=f"e2e-org-team-{unique_marker()}", organization_id=org_id)
+ resources.defer(lambda: client.delete_team(team_id))
+ key = client.generate_key(team_id=team_id)
+ resources.defer(lambda: client.delete_key(key))
+
+ budget_id = client.org_budget_id(org_id)
+ assert budget_id, "org created without a budget row"
+ deadline = time.monotonic() + 30
+ while not any(row.budget_reset_at for row in client.budget_info(budget_id)):
+ if time.monotonic() > deadline:
+ pytest.fail("org budget window never scheduled by the reset job")
+ time.sleep(2)
+
+ _drive_to_block(client, key)
+ _poll_until_serves_again(client, key)
+
+ @pytest.mark.covers("quota_management.budget.internal_user.resets_after_window")
+ def test_personal_key_user_budget_resets_after_window(
+ self, client: BudgetClient, resources: ResourceManager
+ ) -> None:
+ user_id = client.create_user(max_budget=TINY_CAP, budget_duration=WINDOW)
+ resources.defer(lambda: client.delete_user(user_id))
+ key = client.generate_key(user_id=user_id)
+ resources.defer(lambda: client.delete_key(key))
+
+ _drive_to_block(client, key)
+ _poll_until_serves_again(client, key)
+
+ @pytest.mark.covers("quota_management.budget.internal_user.resets_after_window")
+ def test_team_member_key_user_budget_resets_after_window(
+ self, client: BudgetClient, resources: ResourceManager
+ ) -> None:
+ user_id = client.create_user(max_budget=TINY_CAP, budget_duration=WINDOW)
+ resources.defer(lambda: client.delete_user(user_id))
+ team_id = client.create_team(alias=f"e2e-user-team-reset-{unique_marker()}")
+ resources.defer(lambda: client.delete_team(team_id))
+ client.add_team_member(team_id, user_id, max_budget_in_team=100.0)
+ key = client.generate_key(team_id=team_id, user_id=user_id)
+ resources.defer(lambda: client.delete_key(key))
+
+ _drive_to_block(client, key)
+ _poll_until_serves_again(client, key)
From 8536e3b80eabf0e0bed9f0f4eda18936137ebad2 Mon Sep 17 00:00:00 2001
From: "devin-ai-integration[bot]"
<158243242+devin-ai-integration[bot]@users.noreply.github.com>
Date: Fri, 17 Jul 2026 20:04:18 -0700
Subject: [PATCH 022/296] fix(proxy): source /v1/models token limits from the
cost map instead of Router.get_model_group_info (#33721)
* fix(proxy): source /v1/models token limits from cost map instead of Router.get_model_group_info
Resolves the per-model get_model_group_info fan-out on GET /v1/models
(and /models) that pegged the event loop on wildcard listings (#33636).
create_model_info_response now reads max_input_tokens/max_output_tokens
from litellm.get_model_info (the static cost map) rather than the router,
which aggregated and deepcopied every deployment in a group per listed
model.
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
* test(proxy): inject model-info lookup into create_model_info_response for deterministic coverage
Inject the cost-map lookup (defaulting to litellm.get_model_info) so the
except and max_output_tokens branches are exercised deterministically and
the token-limit tests no longer hardcode mutable cost-map values.
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
* feat(proxy): surface custom deployment token limits on /v1/models via cheap index lookup
Add Router.get_configured_token_limits, an O(1) model-name index lookup that
reads a concrete deployment's configured max_input_tokens/max_output_tokens
without triggering pattern matching or deep copies. create_model_info_response
layers this over the cost map so custom deployments absent from the cost map
still surface their limits, and admin-configured limits override cost-map
defaults, while wildcard-expanded names stay on the fast path.
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
---------
Co-authored-by: ryan
Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
---
litellm/proxy/utils.py | 52 +++--
litellm/router.py | 21 ++
tests/test_litellm/proxy/test_proxy_utils.py | 186 ++++++++++--------
.../proxy/utils/helpers/test_model_access.py | 12 +-
tests/test_litellm/test_router.py | 47 +++++
5 files changed, 209 insertions(+), 109 deletions(-)
diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py
index 48164ce913a..7a52fdfdb87 100644
--- a/litellm/proxy/utils.py
+++ b/litellm/proxy/utils.py
@@ -19,6 +19,7 @@ from typing import (
Any,
AsyncGenerator,
Awaitable,
+ Callable,
ClassVar,
Dict,
List,
@@ -49,7 +50,7 @@ from litellm.proxy._types import (
from litellm.proxy.spend_tracking.spend_log_error_logger import spend_log_error
from litellm.types.guardrails import GuardrailEventHooks
from litellm.types.proxy.model_listing import ModelInfoResponse
-from litellm.types.utils import CallTypes, CallTypesLiteral
+from litellm.types.utils import CallTypes, CallTypesLiteral, ModelInfo
try:
from litellm_enterprise.enterprise_callbacks.send_emails.base_email import (
@@ -6096,6 +6097,7 @@ def create_model_info_response(
include_metadata: bool = False,
fallback_type: Optional[str] = None,
llm_router: Optional["Router"] = None,
+ get_model_info: Callable[[str], ModelInfo] = litellm.get_model_info,
) -> ModelInfoResponse:
"""
Create a standardized OpenAI-compatible model object.
@@ -6113,25 +6115,37 @@ def create_model_info_response(
"owned_by": provider,
}
- # Surface context-window limits for OpenAI-compatible discovery clients.
- # Only emitted when known, so wildcard routes and limitless backends stay clean.
- # Limits are best-effort enrichment, so a single malformed deployment degrades
- # to the base response rather than 500-ing the whole listing.
+ try:
+ model_cost_info: ModelInfo | None = get_model_info(model_id)
+ except Exception as e:
+ verbose_proxy_logger.debug(
+ "create_model_info_response: cost map lookup failed for %s: %s",
+ model_id,
+ e,
+ )
+ model_cost_info = None
+
+ max_input_tokens: int | None = None
+ max_output_tokens: int | None = None
+ if model_cost_info is not None:
+ cost_map_input = model_cost_info.get("max_input_tokens")
+ if cost_map_input is not None:
+ max_input_tokens = int(cost_map_input)
+ cost_map_output = model_cost_info.get("max_output_tokens")
+ if cost_map_output is not None:
+ max_output_tokens = int(cost_map_output)
+
if llm_router is not None:
- try:
- model_group_info = llm_router.get_model_group_info(model_id)
- except Exception as e:
- verbose_proxy_logger.debug(
- "create_model_info_response: get_model_group_info failed for %s: %s",
- model_id,
- e,
- )
- model_group_info = None
- if model_group_info is not None:
- if model_group_info.max_input_tokens is not None:
- base["max_input_tokens"] = int(model_group_info.max_input_tokens)
- if model_group_info.max_output_tokens is not None:
- base["max_output_tokens"] = int(model_group_info.max_output_tokens)
+ configured_input, configured_output = llm_router.get_configured_token_limits(model_id)
+ if configured_input is not None:
+ max_input_tokens = configured_input
+ if configured_output is not None:
+ max_output_tokens = configured_output
+
+ if max_input_tokens is not None:
+ base["max_input_tokens"] = max_input_tokens
+ if max_output_tokens is not None:
+ base["max_output_tokens"] = max_output_tokens
if not include_metadata:
return base
diff --git a/litellm/router.py b/litellm/router.py
index 186f382654f..b1a5405ebf1 100644
--- a/litellm/router.py
+++ b/litellm/router.py
@@ -8521,6 +8521,27 @@ class Router:
raise Exception("Model Name invalid - {}".format(type(model)))
return None
+ def get_configured_token_limits(self, model_name: str) -> "tuple[int | None, int | None]":
+ """
+ Return (max_input_tokens, max_output_tokens) explicitly configured in a concrete
+ deployment's model_info for model_name, via O(1) index lookup.
+
+ Returns (None, None) for wildcard-expanded or unknown names. Unlike
+ get_model_group_info, this never triggers pattern matching or deep copies, so it
+ is safe to call per listed model on the /v1/models hot path.
+ """
+ deployment = self.get_deployment_by_model_group_name(model_group_name=model_name)
+ if deployment is None:
+ return (None, None)
+
+ model_info = deployment.model_info
+ max_input = model_info.get("max_input_tokens")
+ max_output = model_info.get("max_output_tokens")
+ return (
+ int(max_input) if max_input is not None else None,
+ int(max_output) if max_output is not None else None,
+ )
+
def get_deployment_credentials_with_provider(self, model_id: str) -> Optional[Dict[str, Any]]:
"""
Get API credentials and provider info from a model name in model_list.
diff --git a/tests/test_litellm/proxy/test_proxy_utils.py b/tests/test_litellm/proxy/test_proxy_utils.py
index a909c510581..d2bdb1764a4 100644
--- a/tests/test_litellm/proxy/test_proxy_utils.py
+++ b/tests/test_litellm/proxy/test_proxy_utils.py
@@ -476,101 +476,118 @@ class TestPostCallFailureHookLiftsRecoveredPartialSpend:
assert "response_cost" not in request_data
+from typing import cast
+
from litellm.proxy.utils import create_model_info_response
-from litellm.types.router import ModelGroupInfo
+from litellm.types.utils import ModelInfo
-def _router_returning(model_group_info):
- router = MagicMock()
- router.get_model_group_info = MagicMock(return_value=model_group_info)
- return router
+def _fake_model_info(**fields: int) -> ModelInfo:
+ return cast(ModelInfo, dict(fields))
-def test_create_model_info_response_includes_max_tokens_when_available():
- router = _router_returning(
- ModelGroupInfo(
- model_group="qwen-vllm",
- providers=["hosted_vllm"],
- max_input_tokens=32768,
- max_output_tokens=8192,
- )
+def _raise_unmapped(model_id: str) -> ModelInfo:
+ raise ValueError(f"This model isn't mapped yet: {model_id}")
+
+
+def test_create_model_info_response_includes_max_tokens_from_lookup():
+ response = create_model_info_response(
+ model_id="some-model",
+ provider="openai",
+ llm_router=None,
+ get_model_info=lambda _model: _fake_model_info(
+ max_input_tokens=128000, max_output_tokens=16384
+ ),
)
+ assert response["id"] == "some-model"
+ assert response["object"] == "model"
+ assert response["max_input_tokens"] == 128000
+ assert response["max_output_tokens"] == 16384
+
+
+def test_create_model_info_response_does_not_call_router_group_info():
+ router = MagicMock()
+ router.get_configured_token_limits.return_value = (None, None)
+
response = create_model_info_response(
- model_id="qwen-vllm", provider="openai", llm_router=router
+ model_id="some-model",
+ provider="openai",
+ llm_router=router,
+ get_model_info=lambda _model: _fake_model_info(
+ max_input_tokens=128000, max_output_tokens=16384
+ ),
)
- router.get_model_group_info.assert_called_once_with("qwen-vllm")
- assert response["id"] == "qwen-vllm"
- assert response["object"] == "model"
- assert response["max_input_tokens"] == 32768
- assert response["max_output_tokens"] == 8192
+ router.get_model_group_info.assert_not_called()
+ assert response["max_input_tokens"] == 128000
+
+
+def test_create_model_info_response_uses_deployment_limits_when_not_in_cost_map():
+ router = MagicMock()
+ router.get_configured_token_limits.return_value = (32000, 8000)
+
+ response = create_model_info_response(
+ model_id="my-custom-deployment",
+ provider="openai",
+ llm_router=router,
+ get_model_info=_raise_unmapped,
+ )
+
+ router.get_model_group_info.assert_not_called()
+ assert response["max_input_tokens"] == 32000
+ assert response["max_output_tokens"] == 8000
+
+
+def test_create_model_info_response_deployment_limits_override_cost_map():
+ router = MagicMock()
+ router.get_configured_token_limits.return_value = (200000, None)
+
+ response = create_model_info_response(
+ model_id="gpt-4o",
+ provider="openai",
+ llm_router=router,
+ get_model_info=lambda _model: _fake_model_info(
+ max_input_tokens=128000, max_output_tokens=16384
+ ),
+ )
+
+ assert response["max_input_tokens"] == 200000
+ assert response["max_output_tokens"] == 16384
def test_create_model_info_response_emits_integer_token_counts():
- # ModelGroupInfo types the limits as float; OpenAI-compatible clients expect
- # plain integers, so the response must not leak 128000.0.
- router = _router_returning(
- ModelGroupInfo(
- model_group="gpt-4o",
- providers=["openai"],
- max_input_tokens=128000.0,
- max_output_tokens=16384.0,
- )
- )
-
response = create_model_info_response(
- model_id="gpt-4o", provider="openai", llm_router=router
+ model_id="some-model",
+ provider="openai",
+ llm_router=None,
+ get_model_info=lambda _model: _fake_model_info(
+ max_input_tokens=128000, max_output_tokens=16384
+ ),
)
- assert response["max_input_tokens"] == 128000
assert isinstance(response["max_input_tokens"], int)
- assert response["max_output_tokens"] == 16384
assert isinstance(response["max_output_tokens"], int)
def test_create_model_info_response_omits_unknown_individual_limit():
- router = _router_returning(
- ModelGroupInfo(
- model_group="partial",
- providers=["openai"],
- max_input_tokens=4096,
- max_output_tokens=None,
- )
- )
-
response = create_model_info_response(
- model_id="partial", provider="openai", llm_router=router
+ model_id="some-embedding",
+ provider="openai",
+ llm_router=None,
+ get_model_info=lambda _model: _fake_model_info(max_input_tokens=8191),
)
- assert response["max_input_tokens"] == 4096
+ assert response["max_input_tokens"] == 8191
assert "max_output_tokens" not in response
-def test_create_model_info_response_omits_limits_when_both_none():
- router = _router_returning(
- ModelGroupInfo(
- model_group="no-limits",
- providers=["openai"],
- max_input_tokens=None,
- max_output_tokens=None,
- )
- )
-
+def test_create_model_info_response_omits_limits_when_lookup_raises():
response = create_model_info_response(
- model_id="no-limits", provider="openai", llm_router=router
- )
-
- assert "max_input_tokens" not in response
- assert "max_output_tokens" not in response
-
-
-def test_create_model_info_response_omits_limits_when_group_unknown():
- # Wildcard routes / access groups have no ModelGroupInfo.
- router = _router_returning(None)
-
- response = create_model_info_response(
- model_id="openai/*", provider="openai", llm_router=router
+ model_id="openai/*",
+ provider="openai",
+ llm_router=None,
+ get_model_info=_raise_unmapped,
)
assert response["id"] == "openai/*"
@@ -578,32 +595,33 @@ def test_create_model_info_response_omits_limits_when_group_unknown():
assert "max_output_tokens" not in response
-def test_create_model_info_response_degrades_when_group_info_raises():
- # A malformed deployment must not turn the listing into a 500; the entry
- # falls back to the base fields without limits.
- router = MagicMock()
- router.get_model_group_info = MagicMock(side_effect=ValueError("bad deployment"))
-
- response = create_model_info_response(
- model_id="broken", provider="openai", llm_router=router
- )
-
- assert response["id"] == "broken"
- assert "max_input_tokens" not in response
- assert "max_output_tokens" not in response
-
-
def test_create_model_info_response_no_router_keeps_base_fields():
response = create_model_info_response(
- model_id="some-model", provider="openai", llm_router=None
+ model_id="totally-unknown-model-xyz",
+ provider="openai",
+ llm_router=None,
+ get_model_info=_raise_unmapped,
)
assert response == {
- "id": "some-model",
+ "id": "totally-unknown-model-xyz",
"object": "model",
"created": response["created"],
"owned_by": "openai",
}
+
+
+def test_create_model_info_response_reads_real_cost_map():
+ response = create_model_info_response(
+ model_id="gpt-4o", provider="openai", llm_router=None
+ )
+
+ assert isinstance(response["max_input_tokens"], int)
+ assert response["max_input_tokens"] > 0
+ assert isinstance(response["max_output_tokens"], int)
+ assert response["max_output_tokens"] > 0
+
+
class TestPostCallFailureHookLLMExceptionAlerting:
"""The llm_exceptions alert is for infra / LLM-API failures, not user
errors (https://github.com/BerriAI/litellm/issues/3395). Already-normalized
diff --git a/tests/test_litellm/proxy/utils/helpers/test_model_access.py b/tests/test_litellm/proxy/utils/helpers/test_model_access.py
index b8e4013c960..59268e1427b 100644
--- a/tests/test_litellm/proxy/utils/helpers/test_model_access.py
+++ b/tests/test_litellm/proxy/utils/helpers/test_model_access.py
@@ -103,18 +103,16 @@ def test_is_known_vector_store_index_error_path_no_registry(monkeypatch):
def test_create_model_info_response_happy_path_no_metadata():
result = create_model_info_response(model_id="gpt-4o", provider="openai")
- assert result == {
- "id": "gpt-4o",
- "object": "model",
- "created": result["created"],
- "owned_by": "openai",
- }
snapshot = {
"id": result["id"],
"object": result["object"],
"owned_by": result["owned_by"],
"created_is_int": isinstance(result["created"], int),
"metadata_absent": "metadata" not in result,
+ "max_input_tokens_positive_int": isinstance(result["max_input_tokens"], int)
+ and result["max_input_tokens"] > 0,
+ "max_output_tokens_positive_int": isinstance(result["max_output_tokens"], int)
+ and result["max_output_tokens"] > 0,
}
assert snapshot == {
"id": "gpt-4o",
@@ -122,6 +120,8 @@ def test_create_model_info_response_happy_path_no_metadata():
"owned_by": "openai",
"created_is_int": True,
"metadata_absent": True,
+ "max_input_tokens_positive_int": True,
+ "max_output_tokens_positive_int": True,
}
diff --git a/tests/test_litellm/test_router.py b/tests/test_litellm/test_router.py
index c2c98c8869c..55c09e6cac4 100644
--- a/tests/test_litellm/test_router.py
+++ b/tests/test_litellm/test_router.py
@@ -5430,3 +5430,50 @@ class TestRouterRequestTimeoutPropagation:
)
== 60
)
+
+
+def test_get_configured_token_limits_reads_deployment_model_info():
+ router = litellm.Router(
+ model_list=[
+ {
+ "model_name": "my-custom-model",
+ "litellm_params": {"model": "openai/some-unmapped-model"},
+ "model_info": {"max_input_tokens": 32000, "max_output_tokens": 8000},
+ }
+ ]
+ )
+
+ assert router.get_configured_token_limits("my-custom-model") == (32000, 8000)
+
+
+def test_get_configured_token_limits_returns_none_for_unset_or_unknown():
+ router = litellm.Router(
+ model_list=[
+ {
+ "model_name": "no-limits-model",
+ "litellm_params": {"model": "openai/some-unmapped-model"},
+ }
+ ]
+ )
+
+ assert router.get_configured_token_limits("no-limits-model") == (None, None)
+ assert router.get_configured_token_limits("not-a-real-model") == (None, None)
+
+
+def test_get_configured_token_limits_skips_wildcard_pattern_matching():
+ router = litellm.Router(
+ model_list=[
+ {
+ "model_name": "bedrock/*",
+ "litellm_params": {"model": "bedrock/*"},
+ "model_info": {"max_input_tokens": 12345},
+ }
+ ]
+ )
+
+ with patch.object(
+ router.pattern_router, "route", side_effect=AssertionError("pattern route called")
+ ):
+ assert router.get_configured_token_limits(
+ "bedrock/anthropic.claude-3-5-sonnet-20240620-v1:0"
+ ) == (None, None)
From 93afde8605829159dc8ff0117a5d6f66cba2ff67 Mon Sep 17 00:00:00 2001
From: "devin-ai-integration[bot]"
<158243242+devin-ai-integration[bot]@users.noreply.github.com>
Date: Fri, 17 Jul 2026 20:29:42 -0700
Subject: [PATCH 023/296] feat(proxy): add x-litellm-model-name response header
with deployment model string (#33698)
The proxy already returns x-litellm-model-id (the deployment id) and x-litellm-model-group (the requested model-group alias), but never surfaces the concrete underlying model that served the request; the router rewrites the response model field to the group alias, so callers had no way to read the actual deployment model like anthropic/claude-haiku-4-5. Expose it as x-litellm-model-name, sourced from the deployment recorded in litellm_params metadata.
Co-authored-by: Krrish Dholakia
Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
---
litellm/proxy/common_request_processing.py | 24 +++++++++
.../proxy/test_model_id_header_propagation.py | 54 +++++++++++++++++++
2 files changed, 78 insertions(+)
diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py
index c7c9397d850..1dc0ee3f947 100644
--- a/litellm/proxy/common_request_processing.py
+++ b/litellm/proxy/common_request_processing.py
@@ -925,9 +925,12 @@ class ProxyBaseLLMRequestProcessing:
# If conversion fails, use original spend
pass
+ model_name = ProxyBaseLLMRequestProcessing._get_deployment_model_name(litellm_logging_obj)
+
headers = {
"x-litellm-call-id": call_id,
"x-litellm-model-id": model_id,
+ "x-litellm-model-name": model_name,
"x-litellm-cache-key": cache_key,
"x-litellm-model-api-base": (
api_base.split("?")[0] if api_base else None
@@ -1396,6 +1399,27 @@ class ProxyBaseLLMRequestProcessing:
model_id = model_info.get("id", "") or ""
return model_id
+ @staticmethod
+ def _get_deployment_model_name(
+ litellm_logging_obj: LiteLLMLoggingObj | None,
+ ) -> str | None:
+ """Extract the underlying deployment model string (e.g. ``azure/gpt-4o``).
+
+ The router rewrites the response ``model`` field to the model-group alias
+ the client requested, so neither the response body nor the existing
+ headers expose the concrete deployment model. The router records it under
+ ``litellm_params`` metadata as ``deployment``, so read it back from there.
+ """
+ litellm_params = getattr(litellm_logging_obj, "litellm_params", None)
+ if not isinstance(litellm_params, dict):
+ return None
+ for key in ("litellm_metadata", "metadata"):
+ metadata = litellm_params.get(key, {}) or {}
+ deployment = metadata.get("deployment")
+ if deployment:
+ return deployment
+ return None
+
@staticmethod
def _response_cost_from_logging_obj(
*,
diff --git a/tests/test_litellm/proxy/test_model_id_header_propagation.py b/tests/test_litellm/proxy/test_model_id_header_propagation.py
index e48168f89b3..f7f3eabae6d 100644
--- a/tests/test_litellm/proxy/test_model_id_header_propagation.py
+++ b/tests/test_litellm/proxy/test_model_id_header_propagation.py
@@ -200,6 +200,60 @@ def test_get_custom_headers_without_model_id():
assert headers["x-litellm-model-id"] in [None, ""]
+class _FakeLoggingObj:
+ def __init__(self, litellm_params):
+ self.litellm_params = litellm_params
+ self.litellm_call_id = "test-call-id"
+
+
+@pytest.mark.parametrize("metadata_key", ["metadata", "litellm_metadata"])
+def test_get_custom_headers_includes_deployment_model_name(metadata_key):
+ """
+ x-litellm-model-name should expose the underlying deployment model string,
+ which the router records under litellm_params[metadata]["deployment"].
+ """
+ mock_user_api_key_dict = MagicMock()
+ mock_user_api_key_dict.tpm_limit = 1000
+ mock_user_api_key_dict.rpm_limit = 100
+
+ logging_obj = _FakeLoggingObj(
+ litellm_params={metadata_key: {"deployment": "azure/gpt-4o-2024-08-06"}}
+ )
+
+ headers = ProxyBaseLLMRequestProcessing.get_custom_headers(
+ user_api_key_dict=mock_user_api_key_dict,
+ model_id="deployment-uuid",
+ request_data={},
+ hidden_params={},
+ litellm_logging_obj=logging_obj,
+ )
+
+ assert headers["x-litellm-model-name"] == "azure/gpt-4o-2024-08-06"
+ assert headers["x-litellm-model-id"] == "deployment-uuid"
+
+
+def test_get_custom_headers_omits_model_name_when_deployment_missing():
+ """
+ Without a deployment model string, x-litellm-model-name must not be emitted
+ (rather than leaking an empty/None value).
+ """
+ mock_user_api_key_dict = MagicMock()
+ mock_user_api_key_dict.tpm_limit = 1000
+ mock_user_api_key_dict.rpm_limit = 100
+
+ logging_obj = _FakeLoggingObj(litellm_params={"metadata": {}})
+
+ headers = ProxyBaseLLMRequestProcessing.get_custom_headers(
+ user_api_key_dict=mock_user_api_key_dict,
+ model_id="deployment-uuid",
+ request_data={},
+ hidden_params={},
+ litellm_logging_obj=logging_obj,
+ )
+
+ assert "x-litellm-model-name" not in headers
+
+
def test_get_custom_headers_with_empty_string_model_id():
"""
Test that get_custom_headers handles empty string model_id correctly.
From f759c75466f0475362a5016ad4cd03c1ecbbd515 Mon Sep 17 00:00:00 2001
From: yucheng-berri
Date: Fri, 17 Jul 2026 20:31:29 -0700
Subject: [PATCH 024/296] feat: add Straiker guardrail integration (#33781)
* feat: add Straiker guardrail integration
Implements LLM security guardrails via Straiker with prompt and response inspection, multi-mode execution (pre_call, post_call), and configurable blocking or redaction of flagged content across providers, streaming, images, and tool calls.
* fix(guardrails): harden straiker source attribution and error-path consistency
Use the operator-configured source for Straiker application attribution instead of a caller-supplied agent_id metadata value, so a caller cannot spoof which application a detection is attributed to. Make _fail reuse _block so a post_call error raises ModifyResponseException like a deliberate post_call block rather than GuardrailRaisedException, and type the blocking helper as NoReturn so the type checker enforces that execution never falls through the BLOCKED branch. Serialize the webhook payload once and send it as raw content to avoid re-serializing on the size check and on every retry.
* fix(guardrails): read straiker config and metadata from all supported shapes
Handle a dict optional_params in _get_config_value so nested guardrail
settings loaded from YAML or the DB (timeout, unreachable_fallback, and
the rest) are applied instead of silently falling back to defaults;
previously only attribute-style access was supported. Build the webhook
metadata bag from the merged metadata so client tags stored under
litellm_metadata on routes like /v1/messages reach Straiker the same way
identity and application fields already do, and widen the internal-key
skip prefix to user_api so proxy-injected budget values are not
forwarded.
* fix(guardrails): fail safe on straiker interventions without redactions
Block instead of passing content through when Straiker returns
GUARDRAIL_INTERVENED without replacement texts, so a positive
intervention verdict can never silently forward the original flagged
content. Fix the streamed-request detection to read the request body
from proxy_server_request.body, where the proxy stores it, instead of a
top-level body key that is never populated; the previous fallback was
dead, so a streamed response whose stream flag was not lifted to the top
level would have been redacted rather than blocked while buffering
replayed the original chunks.
* revert(guardrails): restore straiker caller agent_id application attribution
Restore the original behavior where a request-scoped agent_id in metadata
sets the Straiker application source, falling back to the configured
source. This is the integration's intended per-application attribution;
litellm already resolves a key-owned agent_id ahead of any caller-supplied
value, so a configured key cannot be spoofed.
* revert(guardrails): restore straiker webhook metadata scoping
Restore the original behavior where the Straiker webhook metadata bag is
built from request-scoped metadata only. Forwarding litellm_metadata was
a scope change to what the integration sends to Straiker; keep the
author's intended scoping.
* fix(guardrails): keep proxy key material out of straiker webhook metadata
Widen the internal-key skip prefix from user_api_key_ to user_api so the
proxy-injected user_api_key hash and user_api_end_user_max_budget are not
copied into the Straiker webhook metadata bag. The narrower prefix missed
the bare user_api_key name, leaking the hashed key to the vendor. Keeps
the request-scoped metadata source unchanged.
---------
Co-authored-by: cs-mehta
---
.../guardrail_hooks/straiker/__init__.py | 71 ++
.../guardrail_hooks/straiker/straiker.py | 541 +++++++++++++
litellm/types/guardrails.py | 1 +
.../guardrails/guardrail_hooks/straiker.py | 169 ++++
.../guardrail_hooks/test_straiker.py | 733 ++++++++++++++++++
.../public/assets/logos/straiker.svg | 9 +
.../_components/guardrail_garden_configs.ts | 6 +
.../_components/guardrail_garden_data.ts | 10 +
.../_components/guardrail_info_helpers.tsx | 1 +
9 files changed, 1541 insertions(+)
create mode 100644 litellm/proxy/guardrails/guardrail_hooks/straiker/__init__.py
create mode 100644 litellm/proxy/guardrails/guardrail_hooks/straiker/straiker.py
create mode 100644 litellm/types/proxy/guardrails/guardrail_hooks/straiker.py
create mode 100644 tests/test_litellm/proxy/guardrails/guardrail_hooks/test_straiker.py
create mode 100644 ui/litellm-dashboard/public/assets/logos/straiker.svg
diff --git a/litellm/proxy/guardrails/guardrail_hooks/straiker/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/straiker/__init__.py
new file mode 100644
index 00000000000..ba4c712764e
--- /dev/null
+++ b/litellm/proxy/guardrails/guardrail_hooks/straiker/__init__.py
@@ -0,0 +1,71 @@
+from typing import TYPE_CHECKING
+
+import litellm
+from litellm.types.guardrails import SupportedGuardrailIntegrations
+
+from .straiker import StraikerGuardrail
+
+if TYPE_CHECKING:
+ from litellm.types.guardrails import Guardrail, LitellmParams
+
+_OPTIONAL_INIT_FIELDS = (
+ "timeout",
+ "max_retries",
+ "initial_backoff",
+ "max_backoff",
+ "unreachable_fallback",
+ "fail_on_error",
+ "max_payload_bytes",
+ "custom_headers",
+ "metadata",
+ "verbose",
+)
+
+
+def _get_config_value(litellm_params: "LitellmParams", optional_params: object, attribute_name: str) -> object:
+ if optional_params is not None:
+ if isinstance(optional_params, dict):
+ value = optional_params.get(attribute_name)
+ else:
+ value = getattr(optional_params, attribute_name, None)
+ if value is not None:
+ return value
+ return getattr(litellm_params, attribute_name, None)
+
+
+def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"):
+ optional_params = getattr(litellm_params, "optional_params", None)
+ api_key = litellm_params.api_key
+ if not api_key:
+ raise ValueError("api_key is required for straiker")
+
+ api_base = litellm_params.api_base or "https://api.prod.straiker.ai"
+ default_app = getattr(litellm_params, "default_app", None) or getattr(litellm_params, "source", None)
+ source = default_app if isinstance(default_app, str) and default_app else "LiteLLM Gateway"
+ kwargs: dict[str, object] = {
+ field: value
+ for field in _OPTIONAL_INIT_FIELDS
+ for value in [_get_config_value(litellm_params, optional_params, field)]
+ if value is not None
+ }
+ _callback = StraikerGuardrail(
+ api_key=api_key,
+ api_base=api_base if isinstance(api_base, str) else "https://api.prod.straiker.ai",
+ source=source,
+ guardrail_name=guardrail.get("guardrail_name", "straiker"),
+ event_hook=litellm_params.mode,
+ default_on=litellm_params.default_on,
+ **kwargs,
+ )
+
+ litellm.logging_callback_manager.add_litellm_callback(_callback)
+ return _callback
+
+
+guardrail_initializer_registry = {
+ SupportedGuardrailIntegrations.STRAIKER.value: initialize_guardrail,
+}
+
+guardrail_class_registry = {
+ SupportedGuardrailIntegrations.STRAIKER.value: StraikerGuardrail,
+}
diff --git a/litellm/proxy/guardrails/guardrail_hooks/straiker/straiker.py b/litellm/proxy/guardrails/guardrail_hooks/straiker/straiker.py
new file mode 100644
index 00000000000..5c9f93fc2cd
--- /dev/null
+++ b/litellm/proxy/guardrails/guardrail_hooks/straiker/straiker.py
@@ -0,0 +1,541 @@
+from __future__ import annotations
+
+import asyncio
+import json
+import random
+from dataclasses import dataclass
+from typing import TYPE_CHECKING, Any, Literal, NoReturn
+from urllib.parse import urlsplit
+
+import httpx
+from pydantic import ValidationError
+
+from litellm._logging import verbose_proxy_logger
+from litellm._version import version as litellm_version
+from litellm.exceptions import (
+ BadRequestError,
+ GuardrailRaisedException,
+ ModifyResponseException,
+ Timeout,
+)
+from litellm.integrations.custom_guardrail import (
+ CustomGuardrail,
+ get_session_id_from_request_data,
+ log_guardrail_information,
+)
+from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider
+from litellm.llms.custom_httpx.http_handler import (
+ get_async_httpx_client,
+ httpxSpecialProvider,
+)
+from litellm.types.guardrails import GuardrailEventHooks
+from litellm.types.proxy.guardrails.guardrail_hooks.straiker import (
+ STRAIKER_WEBHOOK_SCHEMA_VERSION,
+ StraikerGuardrailConfigModel,
+ StraikerWebhookApplication,
+ StraikerWebhookContent,
+ StraikerWebhookContext,
+ StraikerWebhookEvent,
+ StraikerWebhookIdentity,
+ StraikerWebhookRequest,
+ StraikerWebhookResponse,
+ StraikerWebhookStream,
+ StraikerWebhookUsage,
+)
+from litellm.types.utils import GenericGuardrailAPIInputs, Usage
+
+if TYPE_CHECKING:
+ from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
+ from litellm.types.proxy.guardrails.guardrail_hooks.base import GuardrailConfigModel
+
+GUARDRAIL_NAME = "straiker"
+DEFAULT_BLOCK_MESSAGE = "Content violates policy"
+DEFAULT_API_BASE = "https://api.prod.straiker.ai"
+DEFAULT_MAX_PAYLOAD_BYTES = 524288
+WEBHOOK_PATH = "/api/v1/detect/webhook"
+RETRY_STATUS = frozenset({408, 429, 500, 502, 503, 504})
+UNREACHABLE_STATUS = frozenset({502, 503, 504})
+_APPLICATION_METADATA_KEYS = frozenset({"agent_id", "app_name"})
+_OPAQUE_METADATA_SCALAR_TYPES = (str, int, float, bool)
+
+
+@dataclass(frozen=True, slots=True)
+class _WebhookFailure:
+ message: str
+ is_unreachable: bool
+
+
+def _as_dict(value: object) -> dict:
+ return value if isinstance(value, dict) else {}
+
+
+def _merged_metadata(request_data: dict) -> dict:
+ return {
+ **_as_dict(request_data.get("metadata")),
+ **_as_dict(request_data.get("litellm_metadata")),
+ }
+
+
+def _as_optional_str(value: object) -> str | None:
+ return value if isinstance(value, str) and value else None
+
+
+def _build_webhook_metadata(request_data: dict, default_metadata: dict[str, str]) -> dict[str, object] | None:
+ out: dict[str, object] = {}
+ for key, value in _as_dict(request_data.get("metadata")).items():
+ if key in _APPLICATION_METADATA_KEYS or key.startswith("user_api"):
+ continue
+ if key == "session_id":
+ continue
+ if isinstance(value, _OPAQUE_METADATA_SCALAR_TYPES):
+ out[key] = value
+ out.update(default_metadata)
+ return out or None
+
+
+def _extract_identity(request_data: dict) -> StraikerWebhookIdentity:
+ meta = _merged_metadata(request_data)
+ return StraikerWebhookIdentity(
+ litellm_key=_as_optional_str(meta.get("user_api_key_alias"))
+ or _as_optional_str(meta.get("user_api_key_hash"))
+ or _as_optional_str(meta.get("user_api_key_token")),
+ litellm_team=_as_optional_str(meta.get("user_api_key_team_alias"))
+ or _as_optional_str(meta.get("user_api_key_team_id")),
+ litellm_user_id=_as_optional_str(meta.get("user_api_key_user_id")),
+ litellm_user_email=_as_optional_str(meta.get("user_api_key_user_email")),
+ litellm_org_id=_as_optional_str(meta.get("user_api_key_org_id")),
+ end_user_id=_as_optional_str(meta.get("user_api_key_end_user_id")),
+ )
+
+
+def _resolve_provider(request_data: dict, model: str | None) -> str | None:
+ litellm_params = _as_dict(request_data.get("litellm_params"))
+ custom_llm_provider = request_data.get("custom_llm_provider") or litellm_params.get("custom_llm_provider")
+ if custom_llm_provider:
+ return custom_llm_provider
+ if not model:
+ return None
+ try:
+ _, provider, _, _ = get_llm_provider(
+ model=model,
+ api_base=request_data.get("api_base") or litellm_params.get("api_base"),
+ api_key=request_data.get("api_key") or litellm_params.get("api_key"),
+ )
+ except BadRequestError:
+ return None
+ return provider or None
+
+
+def _resolve_destination(request_data: dict) -> str | None:
+ litellm_params = _as_dict(request_data.get("litellm_params"))
+ api_base = request_data.get("api_base") or litellm_params.get("api_base")
+ if not isinstance(api_base, str):
+ return None
+ try:
+ return urlsplit(api_base).hostname
+ except ValueError:
+ return None
+
+
+def _resolve_call_surface(logging_obj: LiteLLMLoggingObj | None, request_data: dict) -> str:
+ call_type = (
+ (getattr(logging_obj, "call_type", None) if logging_obj is not None else None)
+ or request_data.get("call_type")
+ or request_data.get("litellm_call_type")
+ )
+ return call_type if isinstance(call_type, str) and call_type else "unknown"
+
+
+def _response_finish_reason(response: Any) -> str | None:
+ choices = getattr(response, "choices", None)
+ if not isinstance(choices, list):
+ return None
+ for choice in choices:
+ reason = getattr(choice, "finish_reason", None)
+ if isinstance(reason, str) and reason:
+ return reason
+ return None
+
+
+def _build_usage(response: object) -> StraikerWebhookUsage | None:
+ usage = getattr(response, "usage", None)
+ if not isinstance(usage, Usage):
+ return None
+ input_tokens = usage.prompt_tokens
+ output_tokens = usage.completion_tokens
+ if input_tokens is None and output_tokens is None:
+ return None
+ return StraikerWebhookUsage(input_tokens=input_tokens, output_tokens=output_tokens)
+
+
+def _is_streamed_request(request_data: dict) -> bool:
+ if request_data.get("stream") is True:
+ return True
+ body = _as_dict(_as_dict(request_data.get("proxy_server_request")).get("body"))
+ return body.get("stream") is True
+
+
+class StraikerGuardrail(CustomGuardrail):
+ @staticmethod
+ def get_config_model() -> type[GuardrailConfigModel]:
+ return StraikerGuardrailConfigModel
+
+ @classmethod
+ def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]:
+ return [
+ GuardrailEventHooks.pre_call,
+ GuardrailEventHooks.post_call,
+ ]
+
+ def __init__(
+ self,
+ api_key: str,
+ api_base: str = DEFAULT_API_BASE,
+ source: str = "LiteLLM Gateway",
+ timeout: float = 5.0,
+ max_retries: int = 2,
+ initial_backoff: float = 0.1,
+ max_backoff: float = 2.0,
+ unreachable_fallback: Literal["fail_open", "fail_closed"] = "fail_closed",
+ fail_on_error: bool = True,
+ max_payload_bytes: int = DEFAULT_MAX_PAYLOAD_BYTES,
+ custom_headers: dict[str, str] | None = None,
+ metadata: dict[str, str] | None = None,
+ verbose: bool = False,
+ async_handler: httpx.AsyncClient | None = None,
+ **kwargs: object,
+ ) -> None:
+ if not api_key:
+ raise ValueError("api_key must be non-empty")
+ if unreachable_fallback not in ("fail_open", "fail_closed"):
+ raise ValueError(f"unreachable_fallback must be 'fail_open' or 'fail_closed'; got {unreachable_fallback!r}")
+
+ self.api_key = api_key
+ self.api_base = api_base.rstrip("/")
+ self.source = source
+ self.timeout = float(timeout)
+ self.max_retries = max(0, int(max_retries))
+ self.initial_backoff = max(0.0, float(initial_backoff))
+ self.max_backoff = max(self.initial_backoff, float(max_backoff))
+ self.unreachable_fallback = unreachable_fallback
+ self.fail_on_error = fail_on_error
+ self.max_payload_bytes = int(max_payload_bytes)
+ self.custom_headers = dict(custom_headers) if custom_headers else {}
+ self.default_metadata = dict(metadata) if metadata else {}
+ self.verbose = bool(verbose)
+
+ self.streaming_end_of_stream_only = True
+ self.streaming_buffer_until_moderated = True
+
+ self.async_handler = async_handler or get_async_httpx_client(
+ llm_provider=httpxSpecialProvider.GuardrailCallback,
+ )
+
+ kwargs.setdefault("supported_event_hooks", list(self.get_supported_event_hooks()))
+ super().__init__(**kwargs)
+
+ def _webhook_url(self) -> str:
+ return f"{self.api_base}{WEBHOOK_PATH}"
+
+ def _headers(self) -> dict[str, str]:
+ reserved = {"authorization", "content-type", "x-straiker-webhook-format"}
+ extra = {k: v for k, v in self.custom_headers.items() if k.lower() not in reserved}
+ return {
+ "Authorization": f"Bearer {self.api_key}",
+ "Content-Type": "application/json",
+ "X-Straiker-Webhook-Format": "litellm",
+ **extra,
+ }
+
+ def _build_application(self, request_data: dict) -> StraikerWebhookApplication:
+ meta = _merged_metadata(request_data)
+ agent_id = _as_optional_str(meta.get("agent_id"))
+ return StraikerWebhookApplication(
+ source=agent_id or self.source,
+ name=_as_optional_str(meta.get("app_name")),
+ )
+
+ def _build_context(
+ self,
+ request_data: dict,
+ model: str | None,
+ logging_obj: LiteLLMLoggingObj | None,
+ ) -> StraikerWebhookContext:
+ return StraikerWebhookContext(
+ call_surface=_resolve_call_surface(logging_obj, request_data),
+ model=model,
+ model_provider=_resolve_provider(request_data, model),
+ destination=_resolve_destination(request_data),
+ session_id=get_session_id_from_request_data(request_data),
+ litellm_call_id=getattr(logging_obj, "litellm_call_id", None) if logging_obj else None,
+ litellm_trace_id=getattr(logging_obj, "litellm_trace_id", None) if logging_obj else None,
+ litellm_version=litellm_version,
+ )
+
+ def _build_envelope(
+ self,
+ *,
+ inputs: GenericGuardrailAPIInputs,
+ request_data: dict,
+ input_type: Literal["request", "response"],
+ logging_obj: LiteLLMLoggingObj | None,
+ ) -> StraikerWebhookRequest:
+ model = inputs.get("model") or request_data.get("model")
+ call_id = getattr(logging_obj, "litellm_call_id", None) if logging_obj else None
+ event_id = f"{call_id or 'litellm'}:{input_type}"
+
+ content = StraikerWebhookContent(
+ texts=list(inputs.get("texts") or []),
+ images=list(inputs.get("images") or []),
+ structured_messages=inputs.get("structured_messages"),
+ tools=inputs.get("tools"),
+ tool_calls=inputs.get("tool_calls"),
+ )
+
+ if input_type == "request":
+ event = StraikerWebhookEvent(type="pre_call", id=event_id)
+ return StraikerWebhookRequest(
+ event=event,
+ request=content,
+ context=self._build_context(request_data, model, logging_obj),
+ identity=_extract_identity(request_data),
+ application=self._build_application(request_data),
+ metadata=_build_webhook_metadata(request_data, self.default_metadata),
+ )
+
+ response_obj = request_data.get("response")
+ content.finish_reason = _response_finish_reason(response_obj)
+ original_messages = request_data.get("messages")
+ request_content = StraikerWebhookContent(
+ structured_messages=original_messages if isinstance(original_messages, list) else None,
+ )
+ phase: Literal["none", "assembled"] = "assembled" if _is_streamed_request(request_data) else "none"
+ event = StraikerWebhookEvent(type="post_call", id=event_id, stream=StraikerWebhookStream(phase=phase))
+ return StraikerWebhookRequest(
+ event=event,
+ request=request_content,
+ response=content,
+ context=self._build_context(request_data, model, logging_obj),
+ identity=_extract_identity(request_data),
+ application=self._build_application(request_data),
+ usage=_build_usage(response_obj),
+ metadata=_build_webhook_metadata(request_data, self.default_metadata),
+ )
+
+ async def _post_webhook(self, payload: dict) -> tuple[StraikerWebhookResponse | None, _WebhookFailure | None]:
+ try:
+ body = json.dumps(payload).encode("utf-8")
+ except (TypeError, ValueError, OverflowError) as error:
+ return None, _WebhookFailure(f"request serialization failed: {error}", is_unreachable=False)
+ body_bytes = len(body)
+ if body_bytes > self.max_payload_bytes:
+ return None, _WebhookFailure(
+ f"payload {body_bytes}B exceeds max_payload_bytes {self.max_payload_bytes}",
+ is_unreachable=False,
+ )
+
+ url = self._webhook_url()
+ headers = self._headers()
+ attempts = self.max_retries + 1
+ last_failure: _WebhookFailure | None = None
+
+ if self.verbose:
+ verbose_proxy_logger.info(
+ json.dumps(
+ {
+ "event": "straiker.webhook_request",
+ "url": url,
+ "bytes": body_bytes,
+ "payload": payload,
+ },
+ default=str,
+ )
+ )
+
+ for attempt in range(attempts):
+ try:
+ resp = await self.async_handler.post(url, content=body, headers=headers, timeout=self.timeout)
+ if resp.status_code == 200:
+ try:
+ body = resp.json()
+ parsed = StraikerWebhookResponse.model_validate(body)
+ except (ValidationError, json.JSONDecodeError) as ve:
+ return None, _WebhookFailure(f"invalid response schema: {ve}", is_unreachable=False)
+ if self.verbose:
+ verbose_proxy_logger.info(
+ json.dumps(
+ {
+ "event": "straiker.webhook_response",
+ "status_code": resp.status_code,
+ "body": body,
+ },
+ default=str,
+ )
+ )
+ return parsed, None
+ last_failure = _WebhookFailure(
+ f"HTTP {resp.status_code}: {resp.text[:200]}",
+ is_unreachable=resp.status_code in UNREACHABLE_STATUS,
+ )
+ if resp.status_code not in RETRY_STATUS:
+ return None, last_failure
+ except (httpx.RequestError, asyncio.TimeoutError, Timeout) as e:
+ last_failure = _WebhookFailure(f"{type(e).__name__}: {e}", is_unreachable=True)
+ except (json.JSONDecodeError, TypeError, ValueError) as e:
+ return None, _WebhookFailure(f"{type(e).__name__}: {e}", is_unreachable=False)
+
+ if attempt < attempts - 1:
+ backoff = min(self.initial_backoff * (2**attempt), self.max_backoff)
+ await asyncio.sleep(random.uniform(0, backoff))
+
+ return None, last_failure or _WebhookFailure("unknown error", is_unreachable=True)
+
+ def _record(
+ self,
+ *,
+ request_data: dict,
+ logging_obj: LiteLLMLoggingObj | None,
+ parsed: StraikerWebhookResponse,
+ ) -> None:
+ if not self.verbose:
+ return
+ response_obj = request_data.get("response")
+ hidden = getattr(response_obj, "_hidden_params", None)
+ if isinstance(hidden, dict):
+ straiker_hidden = hidden.setdefault("straiker", {})
+ if isinstance(straiker_hidden, dict):
+ straiker_hidden.update({"action": parsed.action, "turn_id": parsed.turn_id})
+
+ def _fail(
+ self,
+ *,
+ inputs: GenericGuardrailAPIInputs,
+ request_data: dict,
+ input_type: Literal["request", "response"],
+ error: str,
+ is_unreachable: bool,
+ ) -> GenericGuardrailAPIInputs:
+ fail_open = (is_unreachable and self.unreachable_fallback == "fail_open") or not self.fail_on_error
+ verbose_proxy_logger.error(
+ json.dumps(
+ {
+ "event": "straiker.error",
+ "input_type": input_type,
+ "error": error,
+ "fail_open": fail_open,
+ },
+ default=str,
+ )
+ )
+ if fail_open:
+ return inputs
+ self._block(
+ request_data=request_data,
+ input_type=input_type,
+ message=f"Straiker detection unavailable: {error}",
+ )
+
+ def _block(
+ self,
+ *,
+ request_data: dict,
+ input_type: Literal["request", "response"],
+ message: str,
+ ) -> NoReturn:
+ if input_type == "request":
+ raise GuardrailRaisedException(
+ guardrail_name=self.guardrail_name or GUARDRAIL_NAME,
+ message=message,
+ should_wrap_with_default_message=False,
+ )
+ raise ModifyResponseException(
+ message=message,
+ model=request_data.get("model", "unknown") or "unknown",
+ request_data=request_data,
+ guardrail_name=self.guardrail_name or GUARDRAIL_NAME,
+ original_response=request_data.get("response"),
+ )
+
+ @staticmethod
+ def _intervened_inputs(
+ inputs: GenericGuardrailAPIInputs,
+ parsed: StraikerWebhookResponse,
+ ) -> GenericGuardrailAPIInputs:
+ return_inputs: GenericGuardrailAPIInputs = {}
+ return_inputs.update(inputs)
+ if parsed.texts is not None:
+ return_inputs["texts"] = parsed.texts
+ return return_inputs
+
+ @log_guardrail_information
+ async def apply_guardrail(
+ self,
+ inputs: GenericGuardrailAPIInputs,
+ request_data: dict,
+ input_type: Literal["request", "response"],
+ logging_obj: LiteLLMLoggingObj | None = None,
+ ) -> GenericGuardrailAPIInputs:
+ try:
+ envelope = self._build_envelope(
+ inputs=inputs,
+ request_data=request_data,
+ input_type=input_type,
+ logging_obj=logging_obj,
+ )
+ payload = envelope.model_dump(mode="json", exclude_none=True)
+ except (ValidationError, TypeError, ValueError) as error:
+ return self._fail(
+ inputs=inputs,
+ request_data=request_data,
+ input_type=input_type,
+ error=str(error),
+ is_unreachable=False,
+ )
+
+ parsed, failure = await self._post_webhook(payload)
+ if failure is not None:
+ return self._fail(
+ inputs=inputs,
+ request_data=request_data,
+ input_type=input_type,
+ error=failure.message,
+ is_unreachable=failure.is_unreachable,
+ )
+
+ if parsed is None:
+ return self._fail(
+ inputs=inputs,
+ request_data=request_data,
+ input_type=input_type,
+ error="empty response from Straiker",
+ is_unreachable=False,
+ )
+ self._record(request_data=request_data, logging_obj=logging_obj, parsed=parsed)
+
+ if parsed.schema_version is not None and parsed.schema_version != STRAIKER_WEBHOOK_SCHEMA_VERSION:
+ verbose_proxy_logger.warning(
+ json.dumps(
+ {
+ "event": "straiker.schema_drift",
+ "expected": STRAIKER_WEBHOOK_SCHEMA_VERSION,
+ "received": parsed.schema_version,
+ }
+ )
+ )
+
+ if parsed.action == "BLOCKED":
+ self._block(
+ request_data=request_data,
+ input_type=input_type,
+ message=parsed.blocked_reason or DEFAULT_BLOCK_MESSAGE,
+ )
+ if parsed.action == "GUARDRAIL_INTERVENED":
+ is_streamed_response = input_type == "response" and _is_streamed_request(request_data)
+ if parsed.texts is None or is_streamed_response:
+ self._block(
+ request_data=request_data,
+ input_type=input_type,
+ message=parsed.blocked_reason or DEFAULT_BLOCK_MESSAGE,
+ )
+ return self._intervened_inputs(inputs, parsed)
+ return inputs
diff --git a/litellm/types/guardrails.py b/litellm/types/guardrails.py
index 86e69467dbf..5b611971154 100644
--- a/litellm/types/guardrails.py
+++ b/litellm/types/guardrails.py
@@ -131,6 +131,7 @@ class SupportedGuardrailIntegrations(Enum):
SINGULR = "singulr"
HEADROOM = "headroom"
COMPRESR = "compresr"
+ STRAIKER = "straiker"
class Role(Enum):
diff --git a/litellm/types/proxy/guardrails/guardrail_hooks/straiker.py b/litellm/types/proxy/guardrails/guardrail_hooks/straiker.py
new file mode 100644
index 00000000000..b4375237917
--- /dev/null
+++ b/litellm/types/proxy/guardrails/guardrail_hooks/straiker.py
@@ -0,0 +1,169 @@
+from __future__ import annotations
+
+from typing import Literal
+
+from pydantic import BaseModel, ConfigDict, Field
+
+from litellm.types.llms.openai import AllMessageValues, ChatCompletionToolCallChunk
+from litellm.types.utils import ChatCompletionMessageToolCall
+
+from .base import GuardrailConfigModel
+
+StraikerWebhookEventType = Literal["pre_call", "post_call"]
+StraikerWebhookStreamPhase = Literal["none", "assembled"]
+StraikerWebhookAction = Literal["NONE", "BLOCKED", "GUARDRAIL_INTERVENED"]
+
+STRAIKER_WEBHOOK_SCHEMA_VERSION = "1"
+
+
+class StraikerWebhookStream(BaseModel):
+ phase: StraikerWebhookStreamPhase = "none"
+ index: int | None = None
+
+
+class StraikerWebhookEvent(BaseModel):
+ type: StraikerWebhookEventType
+ id: str
+ stream: StraikerWebhookStream = Field(default_factory=StraikerWebhookStream)
+
+
+class StraikerWebhookContent(BaseModel):
+ model_config = ConfigDict(extra="ignore")
+
+ texts: list[str] = Field(default_factory=list)
+ images: list[str] = Field(default_factory=list)
+ structured_messages: list[AllMessageValues] | None = None
+ tools: list[dict[str, object]] | None = None
+ tool_calls: list[ChatCompletionToolCallChunk] | list[ChatCompletionMessageToolCall] | None = None
+ finish_reason: str | None = None
+
+
+class StraikerWebhookUsage(BaseModel):
+ input_tokens: int | None = None
+ output_tokens: int | None = None
+
+
+class StraikerWebhookContext(BaseModel):
+ call_surface: str
+ model: str | None = None
+ model_provider: str | None = None
+ destination: str | None = None
+ session_id: str | None = None
+ litellm_call_id: str | None = None
+ litellm_trace_id: str | None = None
+ litellm_version: str | None = None
+
+
+class StraikerWebhookIdentity(BaseModel):
+ litellm_key: str | None = None
+ litellm_team: str | None = None
+ litellm_user_id: str | None = None
+ litellm_user_email: str | None = None
+ litellm_org_id: str | None = None
+ end_user_id: str | None = None
+
+
+class StraikerWebhookApplication(BaseModel):
+ source: str
+ name: str | None = None
+
+
+class StraikerWebhookRequest(BaseModel):
+ schema_version: str = STRAIKER_WEBHOOK_SCHEMA_VERSION
+ event: StraikerWebhookEvent
+ request: StraikerWebhookContent
+ response: StraikerWebhookContent | None = None
+ context: StraikerWebhookContext
+ identity: StraikerWebhookIdentity
+ application: StraikerWebhookApplication
+ usage: StraikerWebhookUsage | None = None
+ metadata: dict[str, object] | None = None
+
+
+class StraikerWebhookResponse(BaseModel):
+ model_config = ConfigDict(extra="allow")
+
+ action: StraikerWebhookAction = "NONE"
+ blocked_reason: str | None = None
+ texts: list[str] | None = None
+ schema_version: str | None = None
+ turn_id: str | None = Field(default=None, alias="turnId")
+
+
+class StraikerGuardrailConfigModelOptionalParams(BaseModel):
+ timeout: float | None = Field(
+ default=5.0,
+ gt=0.0,
+ description="Per-attempt HTTP timeout in seconds.",
+ )
+ max_retries: int | None = Field(
+ default=2,
+ ge=0,
+ description="Retries on transient HTTP (408/429/5xx) and network errors.",
+ )
+ initial_backoff: float | None = Field(
+ default=0.1,
+ ge=0.0,
+ description="Initial retry backoff in seconds.",
+ )
+ max_backoff: float | None = Field(
+ default=2.0,
+ ge=0.0,
+ description="Maximum retry backoff in seconds.",
+ )
+ unreachable_fallback: Literal["fail_open", "fail_closed"] | None = Field(
+ default="fail_closed",
+ description="Behavior when Straiker is unreachable after retries.",
+ )
+ fail_on_error: bool | None = Field(
+ default=True,
+ description=(
+ "Behavior on any guardrail error, not just unreachability. True (default) blocks "
+ "the request on error; False logs and allows the request to proceed."
+ ),
+ )
+ max_payload_bytes: int | None = Field(
+ default=524288,
+ gt=0,
+ description="Maximum serialized webhook payload size sent to Straiker.",
+ )
+ custom_headers: dict[str, str] | None = Field(
+ default=None,
+ description="Additional HTTP headers sent to Straiker, excluding Authorization and the webhook-format header.",
+ )
+ metadata: dict[str, str] | None = Field(
+ default=None,
+ description=(
+ "Default metadata key/values added to the webhook metadata bag on every request. "
+ "On key conflict with request-derived metadata, these configured values win."
+ ),
+ )
+ verbose: bool | None = Field(
+ default=False,
+ description="Log webhook request/response payloads and record action/turn_id in response hidden params.",
+ )
+
+
+class StraikerGuardrailConfigModel(GuardrailConfigModel[StraikerGuardrailConfigModelOptionalParams]):
+ api_key: str = Field(
+ min_length=1,
+ description="Straiker DefendAI environment API key (Bearer token). Env: STRAIKER_API_KEY.",
+ json_schema_extra={"secret": True},
+ )
+
+ api_base: str | None = Field(
+ default="https://api.prod.straiker.ai",
+ description="Straiker API base URL. Use the regional variant for non-US tenants.",
+ )
+
+ default_app: str | None = Field(
+ default="LiteLLM Gateway",
+ description=(
+ "Default application registered in the Straiker Defend Console. "
+ "Overridden per-request by metadata.agent_id when present."
+ ),
+ )
+
+ @staticmethod
+ def ui_friendly_name() -> str:
+ return "Straiker"
diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_straiker.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_straiker.py
new file mode 100644
index 00000000000..ca57118ee9d
--- /dev/null
+++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_straiker.py
@@ -0,0 +1,733 @@
+import json
+from unittest.mock import AsyncMock, MagicMock
+
+import httpx
+import pytest
+
+from litellm.exceptions import GuardrailRaisedException, ModifyResponseException
+from litellm.proxy.guardrails.guardrail_hooks.straiker import initialize_guardrail
+from litellm.proxy.guardrails.guardrail_hooks.straiker.straiker import (
+ StraikerGuardrail,
+)
+from litellm.proxy.guardrails.guardrail_registry import (
+ guardrail_class_registry,
+ guardrail_initializer_registry,
+)
+from litellm.types.proxy.guardrails.guardrail_hooks.straiker import (
+ StraikerGuardrailConfigModel,
+ StraikerGuardrailConfigModelOptionalParams,
+)
+from litellm.types.utils import Choices, Message, ModelResponse, Usage
+
+
+def _mock_response(action: str, turn_id: str = "turn-1", schema_version: str = "1", **extra) -> MagicMock:
+ resp = MagicMock(spec=httpx.Response)
+ resp.status_code = 200
+ resp.json.return_value = {
+ "schema_version": schema_version,
+ "action": action,
+ "turn_id": turn_id,
+ **extra,
+ }
+ resp.text = ""
+ return resp
+
+
+def _make_guardrail(**overrides) -> StraikerGuardrail:
+ defaults = dict(
+ api_key="test-key",
+ api_base="https://test.straiker.ai",
+ max_retries=0,
+ guardrail_name="straiker",
+ event_hook="pre_call",
+ async_handler=MagicMock(spec=httpx.AsyncClient),
+ )
+ defaults.update(overrides)
+ g = StraikerGuardrail(**defaults)
+ g.async_handler.post = AsyncMock()
+ return g
+
+
+def _logging_obj() -> MagicMock:
+ obj = MagicMock()
+ obj.litellm_call_id = "call-123"
+ obj.litellm_trace_id = "trace-456"
+ obj.call_type = "acompletion"
+ return obj
+
+
+def _posted_payload(g: StraikerGuardrail) -> dict:
+ return json.loads(g.async_handler.post.call_args.kwargs["content"])
+
+
+def test_registry_membership():
+ assert "straiker" in guardrail_initializer_registry
+ assert guardrail_class_registry["straiker"] is StraikerGuardrail
+
+
+def test_config_model_wiring():
+ assert StraikerGuardrailConfigModel.ui_friendly_name() == "Straiker"
+ assert StraikerGuardrail.get_config_model() is StraikerGuardrailConfigModel
+ fields = StraikerGuardrailConfigModel.model_fields
+ assert "api_key" in fields
+ assert "api_base" in fields
+ assert "default_app" in fields
+ assert "source" not in fields
+ assert "optional_params" in fields
+ assert "timeout" not in fields
+ assert "verbose" not in fields
+
+
+def test_init_rejects_empty_api_key():
+ with pytest.raises(ValueError):
+ StraikerGuardrail(api_key="")
+
+
+def test_init_rejects_invalid_fallback():
+ with pytest.raises(ValueError):
+ StraikerGuardrail(api_key="k", unreachable_fallback="nope")
+
+
+def test_supported_hooks_limited_to_pre_and_post():
+ from litellm.types.guardrails import GuardrailEventHooks
+
+ assert StraikerGuardrail.get_supported_event_hooks() == [
+ GuardrailEventHooks.pre_call,
+ GuardrailEventHooks.post_call,
+ ]
+
+
+def test_during_call_mode_rejected_at_init():
+ with pytest.raises(ValueError):
+ StraikerGuardrail(api_key="k", event_hook="during_call")
+
+
+def test_streaming_attrs_hardcoded_to_buffered():
+ g = _make_guardrail()
+ assert g.streaming_buffer_until_moderated is True
+ assert g.streaming_end_of_stream_only is True
+
+
+def test_streaming_flags_not_configurable():
+ fields = StraikerGuardrailConfigModelOptionalParams.model_fields
+ assert "streaming_buffer_until_moderated" not in fields
+ assert "streaming_end_of_stream_only" not in fields
+ assert "streaming_sampling_rate" not in fields
+
+
+def test_initializer_builds_working_callback():
+ from litellm.types.guardrails import LitellmParams
+
+ params = LitellmParams(guardrail="straiker", mode="pre_call", api_key="abc", api_base="https://x.straiker.ai")
+ callback = initialize_guardrail(params, {"guardrail_name": "straiker"})
+ assert isinstance(callback, StraikerGuardrail)
+ assert callback.api_base == "https://x.straiker.ai"
+
+
+def test_initializer_maps_default_app_to_source():
+ from litellm.types.guardrails import LitellmParams
+
+ params = LitellmParams(
+ guardrail="straiker",
+ mode="pre_call",
+ api_key="abc",
+ default_app="My App",
+ )
+ callback = initialize_guardrail(params, {"guardrail_name": "straiker"})
+ assert callback.source == "My App"
+
+
+def test_initializer_reads_optional_params_flattened_like_ui():
+ from litellm.types.guardrails import LitellmParams
+
+ params = LitellmParams(
+ guardrail="straiker",
+ mode="pre_call",
+ api_key="abc",
+ api_base="https://x.straiker.ai",
+ timeout=9.5,
+ verbose=True,
+ unreachable_fallback="fail_open",
+ )
+ callback = initialize_guardrail(params, {"guardrail_name": "straiker"})
+ assert isinstance(callback, StraikerGuardrail)
+ assert callback.timeout == 9.5
+ assert callback.verbose is True
+ assert callback.unreachable_fallback == "fail_open"
+ assert callback.api_base == "https://x.straiker.ai"
+
+
+def test_initializer_reads_nested_optional_params():
+ from types import SimpleNamespace
+
+ from litellm.types.guardrails import LitellmParams
+
+ params = LitellmParams.model_construct(
+ guardrail="straiker",
+ mode="pre_call",
+ api_key="abc",
+ api_base="https://x.straiker.ai",
+ optional_params=SimpleNamespace(
+ timeout=7.25,
+ verbose=True,
+ unreachable_fallback="fail_open",
+ ),
+ )
+ callback = initialize_guardrail(params, {"guardrail_name": "straiker"})
+ assert isinstance(callback, StraikerGuardrail)
+ assert callback.timeout == 7.25
+ assert callback.verbose is True
+ assert callback.unreachable_fallback == "fail_open"
+
+
+def test_initializer_reads_dict_optional_params():
+ from litellm.types.guardrails import LitellmParams
+
+ params = LitellmParams.model_construct(
+ guardrail="straiker",
+ mode="pre_call",
+ api_key="abc",
+ api_base="https://x.straiker.ai",
+ optional_params={"timeout": 7.25, "verbose": True, "unreachable_fallback": "fail_open"},
+ )
+ callback = initialize_guardrail(params, {"guardrail_name": "straiker"})
+ assert isinstance(callback, StraikerGuardrail)
+ assert callback.timeout == 7.25
+ assert callback.verbose is True
+ assert callback.unreachable_fallback == "fail_open"
+
+
+@pytest.mark.asyncio
+async def test_request_envelope_transport_and_shape():
+ g = _make_guardrail()
+ g.async_handler.post.return_value = _mock_response("NONE")
+ inputs = {"texts": ["hello world"], "model": "gpt-4o-mini"}
+ request_data = {
+ "model": "gpt-4o-mini",
+ "messages": [{"role": "user", "content": "hello world"}],
+ "metadata": {"user_api_key_alias": "team-key", "agent_id": "chatbot-app", "app_name": "Chatbot"},
+ }
+
+ out = await g.apply_guardrail(inputs=inputs, request_data=request_data, input_type="request", logging_obj=_logging_obj())
+
+ assert out is inputs
+ url = g.async_handler.post.call_args.args[0]
+ assert url == "https://test.straiker.ai/api/v1/detect/webhook"
+ headers = g.async_handler.post.call_args.kwargs["headers"]
+ assert headers["X-Straiker-Webhook-Format"] == "litellm"
+ assert headers["Authorization"] == "Bearer test-key"
+
+ payload = _posted_payload(g)
+ assert payload["schema_version"] == "1"
+ assert payload["event"]["type"] == "pre_call"
+ assert payload["event"]["id"] == "call-123:request"
+ assert payload["request"]["texts"] == ["hello world"]
+ assert payload["context"]["litellm_call_id"] == "call-123"
+ assert payload["identity"]["litellm_key"] == "team-key"
+ assert payload["application"] == {"source": "chatbot-app", "name": "Chatbot"}
+ assert "session_id" not in payload["application"]
+ assert "user_name" not in payload["application"]
+ assert "user_role" not in payload["application"]
+ assert "response" not in payload
+ assert "metadata" not in payload
+
+
+@pytest.mark.asyncio
+async def test_webhook_metadata_session_id_and_opaque_passthrough():
+ g = _make_guardrail()
+ g.async_handler.post.return_value = _mock_response("NONE")
+ await g.apply_guardrail(
+ inputs={"texts": ["x"]},
+ request_data={
+ "model": "m",
+ "litellm_session_id": "sess-from-litellm",
+ "metadata": {
+ "agent_id": "chatbot-app",
+ "app_name": "Chatbot",
+ "user_api_key_alias": "team-key",
+ "custom_tag": "experiment-7",
+ "client_ip": "10.0.0.1",
+ },
+ },
+ input_type="request",
+ logging_obj=_logging_obj(),
+ )
+ payload = _posted_payload(g)
+ assert payload["application"] == {"source": "chatbot-app", "name": "Chatbot"}
+ assert payload["identity"]["litellm_key"] == "team-key"
+ assert payload["context"]["session_id"] == "sess-from-litellm"
+ assert "session_id" not in payload["metadata"]
+ assert payload["metadata"] == {
+ "custom_tag": "experiment-7",
+ "client_ip": "10.0.0.1",
+ }
+
+
+@pytest.mark.asyncio
+async def test_webhook_metadata_never_forwards_proxy_internal_keys():
+ g = _make_guardrail()
+ g.async_handler.post.return_value = _mock_response("NONE")
+ await g.apply_guardrail(
+ inputs={"texts": ["x"]},
+ request_data={
+ "model": "m",
+ "metadata": {
+ "custom_tag": "experiment-7",
+ "user_api_key": "sk-hashed-secret",
+ "user_api_end_user_max_budget": 12.5,
+ },
+ },
+ input_type="request",
+ logging_obj=_logging_obj(),
+ )
+ assert _posted_payload(g)["metadata"] == {"custom_tag": "experiment-7"}
+
+
+@pytest.mark.asyncio
+async def test_default_metadata_injected_and_config_wins_on_clash():
+ g = _make_guardrail(metadata={"tenant": "acme", "custom_tag": "config-value"})
+ g.async_handler.post.return_value = _mock_response("NONE")
+ await g.apply_guardrail(
+ inputs={"texts": ["x"]},
+ request_data={
+ "model": "m",
+ "metadata": {"custom_tag": "request-value", "client_ip": "10.0.0.1"},
+ },
+ input_type="request",
+ logging_obj=_logging_obj(),
+ )
+ assert _posted_payload(g)["metadata"] == {
+ "client_ip": "10.0.0.1",
+ "custom_tag": "config-value",
+ "tenant": "acme",
+ }
+
+
+@pytest.mark.asyncio
+async def test_default_metadata_present_without_request_metadata():
+ g = _make_guardrail(metadata={"tenant": "acme"})
+ g.async_handler.post.return_value = _mock_response("NONE")
+ await g.apply_guardrail(
+ inputs={"texts": ["x"]},
+ request_data={"model": "m"},
+ input_type="request",
+ logging_obj=_logging_obj(),
+ )
+ assert _posted_payload(g)["metadata"] == {"tenant": "acme"}
+
+
+@pytest.mark.asyncio
+async def test_context_session_id_from_request_metadata():
+ g = _make_guardrail()
+ g.async_handler.post.return_value = _mock_response("NONE")
+ await g.apply_guardrail(
+ inputs={"texts": ["x"]},
+ request_data={"model": "m", "metadata": {"session_id": "sess-meta"}},
+ input_type="request",
+ logging_obj=_logging_obj(),
+ )
+ payload = _posted_payload(g)
+ assert payload["context"]["session_id"] == "sess-meta"
+ assert "metadata" not in payload
+
+
+@pytest.mark.asyncio
+async def test_identity_key_and_team_coalesce_alias_over_id():
+ g = _make_guardrail()
+ g.async_handler.post.return_value = _mock_response("NONE")
+ await g.apply_guardrail(
+ inputs={"texts": ["x"]},
+ request_data={
+ "model": "m",
+ "metadata": {
+ "user_api_key_alias": "prod-key",
+ "user_api_key_hash": "hash-abc",
+ "user_api_key_team_alias": "growth",
+ "user_api_key_team_id": "team-9",
+ },
+ },
+ input_type="request",
+ logging_obj=_logging_obj(),
+ )
+ identity = _posted_payload(g)["identity"]
+ assert identity["litellm_key"] == "prod-key"
+ assert identity["litellm_team"] == "growth"
+ assert "key" not in identity
+ assert "team" not in identity
+
+
+@pytest.mark.asyncio
+async def test_identity_key_and_team_fall_back_to_hash_and_id():
+ g = _make_guardrail()
+ g.async_handler.post.return_value = _mock_response("NONE")
+ await g.apply_guardrail(
+ inputs={"texts": ["x"]},
+ request_data={
+ "model": "m",
+ "metadata": {
+ "user_api_key_hash": "hash-abc",
+ "user_api_key_team_id": "team-9",
+ },
+ },
+ input_type="request",
+ logging_obj=_logging_obj(),
+ )
+ identity = _posted_payload(g)["identity"]
+ assert identity["litellm_key"] == "hash-abc"
+ assert identity["litellm_team"] == "team-9"
+
+
+@pytest.mark.asyncio
+async def test_identity_end_user_from_resolved_metadata():
+ g = _make_guardrail()
+ g.async_handler.post.return_value = _mock_response("NONE")
+
+ await g.apply_guardrail(
+ inputs={"texts": ["x"]},
+ request_data={
+ "model": "m",
+ "metadata": {
+ "user_api_key_end_user_id": "eu-meta",
+ "user_api_key_user_id": "default_user_id",
+ },
+ "user": "eu-body",
+ },
+ input_type="request",
+ logging_obj=_logging_obj(),
+ )
+ identity = _posted_payload(g)["identity"]
+ assert identity["end_user_id"] == "eu-meta"
+ assert identity["litellm_user_id"] == "default_user_id"
+ assert _posted_payload(g)["application"] == {"source": g.source}
+
+
+@pytest.mark.asyncio
+async def test_identity_end_user_absent_without_resolved_metadata():
+ g = _make_guardrail()
+ g.async_handler.post.return_value = _mock_response("NONE")
+ await g.apply_guardrail(
+ inputs={"texts": ["x"]},
+ request_data={"model": "m", "user": "eu-body", "metadata": {"user_api_key_user_id": "default_user_id"}},
+ input_type="request",
+ logging_obj=_logging_obj(),
+ )
+ assert "end_user_id" not in _posted_payload(g)["identity"]
+
+
+@pytest.mark.asyncio
+async def test_application_source_from_agent_id():
+ g = _make_guardrail(source="litellm")
+ g.async_handler.post.return_value = _mock_response("NONE")
+ await g.apply_guardrail(
+ inputs={"texts": ["x"]},
+ request_data={"model": "m", "metadata": {"agent_id": "analytics-app", "app_name": "Analytics"}},
+ input_type="request",
+ logging_obj=_logging_obj(),
+ )
+ assert _posted_payload(g)["application"] == {"source": "analytics-app", "name": "Analytics"}
+
+@pytest.mark.asyncio
+async def test_request_block_raises_guardrail_exception_with_reason():
+ g = _make_guardrail()
+ g.async_handler.post.return_value = _mock_response("BLOCKED", blocked_reason="prompt injection")
+ with pytest.raises(GuardrailRaisedException) as exc:
+ await g.apply_guardrail(
+ inputs={"texts": ["attack"]}, request_data={"model": "m"}, input_type="request", logging_obj=_logging_obj()
+ )
+ assert "prompt injection" in str(exc.value)
+
+
+@pytest.mark.asyncio
+async def test_guardrail_intervened_writes_back_modified_text_only():
+ g = _make_guardrail()
+ g.async_handler.post.return_value = _mock_response("GUARDRAIL_INTERVENED", texts=["[redacted]"])
+ inputs = {"texts": ["my ssn is 123"], "images": ["img-a"]}
+ out = await g.apply_guardrail(
+ inputs=inputs, request_data={"model": "m"}, input_type="request", logging_obj=_logging_obj()
+ )
+ assert out["texts"] == ["[redacted]"]
+ assert out["images"] == ["img-a"]
+
+
+@pytest.mark.asyncio
+async def test_streamed_response_intervention_converts_to_block():
+ g = _make_guardrail()
+ g.async_handler.post.return_value = _mock_response("GUARDRAIL_INTERVENED", texts=["[redacted]"])
+ response = ModelResponse(
+ choices=[Choices(finish_reason="stop", index=0, message=Message(content="secret", role="assistant"))],
+ model="gpt-4o-mini",
+ )
+ request_data = {
+ "model": "gpt-4o-mini",
+ "messages": [{"role": "user", "content": "p"}],
+ "stream": True,
+ "response": response,
+ }
+ with pytest.raises(ModifyResponseException):
+ await g.apply_guardrail(
+ inputs={"texts": ["secret"], "model": "gpt-4o-mini"},
+ request_data=request_data,
+ input_type="response",
+ logging_obj=_logging_obj(),
+ )
+
+
+@pytest.mark.asyncio
+async def test_non_streamed_response_intervention_redacts():
+ g = _make_guardrail()
+ g.async_handler.post.return_value = _mock_response("GUARDRAIL_INTERVENED", texts=["[redacted]"])
+ response = ModelResponse(
+ choices=[Choices(finish_reason="stop", index=0, message=Message(content="secret", role="assistant"))],
+ model="gpt-4o-mini",
+ )
+ request_data = {
+ "model": "gpt-4o-mini",
+ "messages": [{"role": "user", "content": "p"}],
+ "response": response,
+ }
+ out = await g.apply_guardrail(
+ inputs={"texts": ["secret"], "model": "gpt-4o-mini"},
+ request_data=request_data,
+ input_type="response",
+ logging_obj=_logging_obj(),
+ )
+ assert out["texts"] == ["[redacted]"]
+
+
+@pytest.mark.asyncio
+async def test_guardrail_intervened_without_texts_blocks():
+ g = _make_guardrail()
+ g.async_handler.post.return_value = _mock_response("GUARDRAIL_INTERVENED")
+ with pytest.raises(GuardrailRaisedException):
+ await g.apply_guardrail(
+ inputs={"texts": ["my ssn is 123"]},
+ request_data={"model": "m"},
+ input_type="request",
+ logging_obj=_logging_obj(),
+ )
+
+
+@pytest.mark.asyncio
+async def test_streamed_via_proxy_server_request_body_converts_to_block():
+ g = _make_guardrail()
+ g.async_handler.post.return_value = _mock_response("GUARDRAIL_INTERVENED", texts=["[redacted]"])
+ response = ModelResponse(
+ choices=[Choices(finish_reason="stop", index=0, message=Message(content="secret", role="assistant"))],
+ model="gpt-4o-mini",
+ )
+ request_data = {
+ "model": "gpt-4o-mini",
+ "proxy_server_request": {"body": {"stream": True}},
+ "response": response,
+ }
+ with pytest.raises(ModifyResponseException):
+ await g.apply_guardrail(
+ inputs={"texts": ["secret"], "model": "gpt-4o-mini"},
+ request_data=request_data,
+ input_type="response",
+ logging_obj=_logging_obj(),
+ )
+
+
+@pytest.mark.asyncio
+async def test_response_envelope_and_block_replaces_response():
+ g = _make_guardrail(verbose=True)
+ g.async_handler.post.return_value = _mock_response("BLOCKED")
+ response = ModelResponse(
+ choices=[Choices(finish_reason="stop", index=0, message=Message(content="secret", role="assistant"))],
+ model="gpt-4o-mini",
+ )
+ request_data = {
+ "model": "gpt-4o-mini",
+ "messages": [{"role": "user", "content": "original prompt"}],
+ "stream": True,
+ "response": response,
+ }
+ with pytest.raises(ModifyResponseException) as exc:
+ await g.apply_guardrail(
+ inputs={"texts": ["secret"], "model": "gpt-4o-mini"},
+ request_data=request_data,
+ input_type="response",
+ logging_obj=_logging_obj(),
+ )
+
+ assert exc.value.original_response is response
+ payload = _posted_payload(g)
+ assert payload["event"]["type"] == "post_call"
+ assert payload["event"]["stream"]["phase"] == "assembled"
+ assert payload["response"]["texts"] == ["secret"]
+ assert payload["response"]["finish_reason"] == "stop"
+ assert payload["request"]["structured_messages"] == [{"role": "user", "content": "original prompt"}]
+
+
+@pytest.mark.asyncio
+async def test_post_call_fail_closed_raises_modify_response_exception():
+ g = _make_guardrail(unreachable_fallback="fail_closed")
+ g.async_handler.post.side_effect = httpx.ConnectError("boom")
+ response = ModelResponse(
+ choices=[Choices(finish_reason="stop", index=0, message=Message(content="secret", role="assistant"))],
+ model="gpt-4o-mini",
+ )
+ request_data = {"model": "gpt-4o-mini", "response": response}
+ with pytest.raises(ModifyResponseException) as exc:
+ await g.apply_guardrail(
+ inputs={"texts": ["secret"], "model": "gpt-4o-mini"},
+ request_data=request_data,
+ input_type="response",
+ logging_obj=_logging_obj(),
+ )
+ assert exc.value.original_response is response
+
+
+@pytest.mark.asyncio
+async def test_usage_tokens_on_post_call():
+ g = _make_guardrail()
+ g.async_handler.post.return_value = _mock_response("NONE")
+ response = ModelResponse(
+ choices=[Choices(finish_reason="stop", index=0, message=Message(content="hi", role="assistant"))],
+ model="gpt-4o-mini",
+ usage=Usage(prompt_tokens=11, completion_tokens=7, total_tokens=18),
+ )
+ await g.apply_guardrail(
+ inputs={"texts": ["hi"], "model": "gpt-4o-mini"},
+ request_data={"model": "gpt-4o-mini", "messages": [{"role": "user", "content": "hey"}], "response": response},
+ input_type="response",
+ logging_obj=_logging_obj(),
+ )
+ usage = _posted_payload(g)["usage"]
+ assert usage == {"input_tokens": 11, "output_tokens": 7}
+
+
+@pytest.mark.asyncio
+async def test_usage_absent_on_pre_call():
+ g = _make_guardrail()
+ g.async_handler.post.return_value = _mock_response("NONE")
+ await g.apply_guardrail(
+ inputs={"texts": ["hi"]},
+ request_data={"model": "gpt-4o-mini"},
+ input_type="request",
+ logging_obj=_logging_obj(),
+ )
+ assert "usage" not in _posted_payload(g)
+
+
+@pytest.mark.asyncio
+async def test_allow_returns_inputs_unchanged():
+ g = _make_guardrail()
+ g.async_handler.post.return_value = _mock_response("NONE")
+ inputs = {"texts": ["fine"]}
+ out = await g.apply_guardrail(
+ inputs=inputs, request_data={"model": "m"}, input_type="request", logging_obj=_logging_obj()
+ )
+ assert out is inputs
+
+
+@pytest.mark.asyncio
+async def test_unreachable_fail_closed_blocks():
+ g = _make_guardrail(unreachable_fallback="fail_closed")
+ g.async_handler.post.side_effect = httpx.ConnectError("boom")
+ with pytest.raises(GuardrailRaisedException):
+ await g.apply_guardrail(
+ inputs={"texts": ["x"]}, request_data={"model": "m"}, input_type="request", logging_obj=_logging_obj()
+ )
+
+
+@pytest.mark.asyncio
+async def test_unreachable_fail_open_passes_through():
+ g = _make_guardrail(unreachable_fallback="fail_open")
+ g.async_handler.post.side_effect = httpx.ConnectError("boom")
+ inputs = {"texts": ["x"]}
+ out = await g.apply_guardrail(
+ inputs=inputs, request_data={"model": "m"}, input_type="request", logging_obj=_logging_obj()
+ )
+ assert out is inputs
+
+
+@pytest.mark.asyncio
+async def test_fail_on_error_false_allows_on_bad_status():
+ g = _make_guardrail(unreachable_fallback="fail_closed", fail_on_error=False)
+ bad = MagicMock(spec=httpx.Response)
+ bad.status_code = 400
+ bad.text = "bad request"
+ g.async_handler.post.return_value = bad
+ inputs = {"texts": ["x"]}
+ out = await g.apply_guardrail(
+ inputs=inputs, request_data={"model": "m"}, input_type="request", logging_obj=_logging_obj()
+ )
+ assert out is inputs
+
+
+@pytest.mark.asyncio
+async def test_non_retryable_status_fail_closed_blocks():
+ g = _make_guardrail(unreachable_fallback="fail_closed", fail_on_error=True)
+ bad = MagicMock(spec=httpx.Response)
+ bad.status_code = 401
+ bad.text = "unauthorized"
+ g.async_handler.post.return_value = bad
+ with pytest.raises(GuardrailRaisedException):
+ await g.apply_guardrail(
+ inputs={"texts": ["x"]}, request_data={"model": "m"}, input_type="request", logging_obj=_logging_obj()
+ )
+
+
+@pytest.mark.asyncio
+async def test_payload_size_guard_fails_closed():
+ g = _make_guardrail(max_payload_bytes=10)
+ inputs = {"texts": ["x" * 5000]}
+ with pytest.raises(GuardrailRaisedException):
+ await g.apply_guardrail(
+ inputs=inputs, request_data={"model": "m"}, input_type="request", logging_obj=_logging_obj()
+ )
+ g.async_handler.post.assert_not_called()
+
+
+@pytest.mark.asyncio
+async def test_payload_size_guard_blocks_even_with_fail_open():
+ g = _make_guardrail(max_payload_bytes=10, unreachable_fallback="fail_open")
+ inputs = {"texts": ["x" * 5000]}
+ with pytest.raises(GuardrailRaisedException):
+ await g.apply_guardrail(
+ inputs=inputs, request_data={"model": "m"}, input_type="request", logging_obj=_logging_obj()
+ )
+ g.async_handler.post.assert_not_called()
+
+
+@pytest.mark.asyncio
+async def test_invalid_response_schema_blocks_even_with_fail_open():
+ g = _make_guardrail(unreachable_fallback="fail_open")
+ bad = MagicMock(spec=httpx.Response)
+ bad.status_code = 200
+ bad.json.return_value = {"action": "NOT_A_VALID_ACTION"}
+ bad.text = ""
+ g.async_handler.post.return_value = bad
+ with pytest.raises(GuardrailRaisedException):
+ await g.apply_guardrail(
+ inputs={"texts": ["x"]}, request_data={"model": "m"}, input_type="request", logging_obj=_logging_obj()
+ )
+
+
+@pytest.mark.asyncio
+async def test_unreachable_http_status_fail_open_passes():
+ g = _make_guardrail(unreachable_fallback="fail_open")
+ resp = MagicMock(spec=httpx.Response)
+ resp.status_code = 503
+ resp.text = "service unavailable"
+ g.async_handler.post.return_value = resp
+ inputs = {"texts": ["x"]}
+ out = await g.apply_guardrail(
+ inputs=inputs, request_data={"model": "m"}, input_type="request", logging_obj=_logging_obj()
+ )
+ assert out is inputs
+
+
+@pytest.mark.asyncio
+async def test_unreachable_http_status_fail_closed_blocks():
+ g = _make_guardrail(unreachable_fallback="fail_closed")
+ resp = MagicMock(spec=httpx.Response)
+ resp.status_code = 503
+ resp.text = "service unavailable"
+ g.async_handler.post.return_value = resp
+ with pytest.raises(GuardrailRaisedException):
+ await g.apply_guardrail(
+ inputs={"texts": ["x"]}, request_data={"model": "m"}, input_type="request", logging_obj=_logging_obj()
+ )
diff --git a/ui/litellm-dashboard/public/assets/logos/straiker.svg b/ui/litellm-dashboard/public/assets/logos/straiker.svg
new file mode 100644
index 00000000000..bdfe0405736
--- /dev/null
+++ b/ui/litellm-dashboard/public/assets/logos/straiker.svg
@@ -0,0 +1,9 @@
+
diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_configs.ts b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_configs.ts
index 2ad5819b5f0..a40587cb3ae 100644
--- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_configs.ts
+++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_configs.ts
@@ -300,4 +300,10 @@ export const GUARDRAIL_PRESETS: Record = {
mode: "pre_call",
defaultOn: false,
},
+ straiker: {
+ provider: "Straiker",
+ guardrailNameSuggestion: "Straiker Guardrail",
+ mode: "pre_call",
+ defaultOn: false,
+ },
};
diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_data.ts b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_data.ts
index f81277f13c3..ba11d3d400d 100644
--- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_data.ts
+++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_garden_data.ts
@@ -442,6 +442,16 @@ export const PARTNER_GUARDRAIL_CARDS: GuardrailCardInfo[] = [
tags: ["Security", "Policy", "Prompt Injection"],
providerKey: "Repelloai",
},
+ {
+ id: "straiker",
+ name: "Straiker",
+ description:
+ "Defend AI Agentic Guardrails: Indirect/Direct Prompt Injection, Tool Misuse, Malicious MCP and Skills",
+ category: "partner",
+ logo: `${ASSET_PREFIX}straiker.svg`,
+ tags: ["Agentic", "Prompt Injection", "Tool Misuse", "MCP", "Skills"],
+ providerKey: "Straiker",
+ },
];
export const ALL_CARDS = [...LITELLM_CONTENT_FILTER_CARDS, ...PARTNER_GUARDRAIL_CARDS];
diff --git a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info_helpers.tsx b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info_helpers.tsx
index e8c4810ca69..a2873797096 100644
--- a/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info_helpers.tsx
+++ b/ui/litellm-dashboard/src/app/(dashboard)/guardrails/_components/guardrail_info_helpers.tsx
@@ -166,6 +166,7 @@ export const guardrailLogoMap: Record = {
Akto: `${asset_logos_folder}akto.svg`,
"Qostodian Nexus": `${asset_logos_folder}qohash.jpg`,
"RepelloAI Argus": `${asset_logos_folder}repelloai.png`,
+ Straiker: `${asset_logos_folder}straiker.svg`,
};
export const getGuardrailLogoAndName = (guardrailValue: string): { logo: string; displayName: string } => {
From 07e07e6e2b0dd27f9bd50180ed8d916cc32068f0 Mon Sep 17 00:00:00 2001
From: "devin-ai-integration[bot]"
<158243242+devin-ai-integration[bot]@users.noreply.github.com>
Date: Fri, 17 Jul 2026 21:17:49 -0700
Subject: [PATCH 025/296] fix(vertex_ai): exclude Gemini Google Search
grounding tokens from input token billing (#33742)
* fix(vertex_ai): exclude Google Search grounding tokens from Gemini input token billing
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
* test(proxy): stub get_configured_token_limits on mocked routers
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
---------
Co-authored-by: Krrish Dholakia
Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
---
.../vertex_and_google_ai_studio_gemini.py | 32 +++++-
...test_vertex_and_google_ai_studio_gemini.py | 97 +++++++++++++++++++
.../test_model_management_endpoints.py | 2 +
.../test_team_model_name_translation.py | 6 ++
4 files changed, 136 insertions(+), 1 deletion(-)
diff --git a/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py b/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py
index 3193b72a7d9..624190a0b61 100644
--- a/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py
+++ b/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py
@@ -1744,6 +1744,30 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
)
return non_thinking_tokens == usage_metadata.get("totalTokenCount", 0)
+ @staticmethod
+ def _response_has_search_grounding(
+ completion_response: Union[GenerateContentResponseBody, BidiGenerateContentServerMessage],
+ ) -> bool:
+ """
+ Whether the response used Grounding with Google Search, detected via
+ groundingMetadata.webSearchQueries (an actual web search was performed).
+
+ Google bills grounding-with-Google-Search retrieved tokens separately (a per-request /
+ per-query search fee) and excludes them from input token billing, unlike URL context /
+ File Search / code execution whose tool-use tokens are charged at the input token rate.
+ URL context also emits groundingMetadata (with groundingChunks but no webSearchQueries),
+ so presence of groundingMetadata alone is not a sufficient signal.
+ See https://ai.google.dev/gemini-api/docs/pricing and
+ https://github.com/BerriAI/litellm/discussions/33198
+ """
+ if "candidates" not in completion_response:
+ return False
+ for candidate in completion_response["candidates"] or []:
+ grounding_metadata, _, _, _ = VertexGeminiConfig._extract_candidate_metadata(candidate)
+ if VertexGeminiConfig._calculate_web_search_requests(grounding_metadata):
+ return True
+ return False
+
@staticmethod
def _calculate_usage(
completion_response: Union[GenerateContentResponseBody, BidiGenerateContentServerMessage],
@@ -1899,12 +1923,18 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
tool_use_tokens=tool_use_prompt_tokens,
)
+ billable_tool_use_prompt_tokens = (
+ 0
+ if VertexGeminiConfig._response_has_search_grounding(completion_response)
+ else (tool_use_prompt_tokens or 0)
+ )
+
completion_tokens = response_tokens or completion_response["usageMetadata"].get("candidatesTokenCount", 0)
if not VertexGeminiConfig.is_candidate_token_count_inclusive(usage_metadata) and reasoning_tokens:
completion_tokens = reasoning_tokens + completion_tokens
## GET USAGE ##
usage = Usage(
- prompt_tokens=usage_metadata.get("promptTokenCount", 0) + (tool_use_prompt_tokens or 0),
+ prompt_tokens=usage_metadata.get("promptTokenCount", 0) + billable_tool_use_prompt_tokens,
completion_tokens=completion_tokens,
total_tokens=usage_metadata.get("totalTokenCount", 0),
prompt_tokens_details=prompt_tokens_details,
diff --git a/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_and_google_ai_studio_gemini.py b/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_and_google_ai_studio_gemini.py
index 5adc5b76990..95e8e6561f1 100644
--- a/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_and_google_ai_studio_gemini.py
+++ b/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_and_google_ai_studio_gemini.py
@@ -547,6 +547,103 @@ def test_vertex_ai_non_grounded_usage_omits_tool_use_tokens():
assert not hasattr(usage.prompt_tokens_details, "tool_use_tokens")
+def test_response_has_search_grounding_detection():
+ """
+ Only groundingMetadata.webSearchQueries signals an actual Google Search. URL context also
+ emits groundingMetadata (groundingChunks but no webSearchQueries) and must not be treated
+ as search grounding.
+ """
+ assert (
+ VertexGeminiConfig._response_has_search_grounding(
+ {"candidates": [{"groundingMetadata": {"webSearchQueries": ["latest nobel physics"]}}]}
+ )
+ is True
+ )
+ assert (
+ VertexGeminiConfig._response_has_search_grounding(
+ {
+ "candidates": [
+ {
+ "urlContextMetadata": {"urlMetadata": []},
+ "groundingMetadata": {
+ "groundingChunks": [{"web": {"uri": "https://example.com", "title": "Example"}}]
+ },
+ }
+ ]
+ }
+ )
+ is False
+ )
+ assert (
+ VertexGeminiConfig._response_has_search_grounding({"candidates": [{"groundingMetadata": {"webSearchQueries": []}}]})
+ is False
+ )
+ assert VertexGeminiConfig._response_has_search_grounding({"candidates": []}) is False
+ assert VertexGeminiConfig._response_has_search_grounding({}) is False
+
+
+def test_vertex_ai_search_grounding_tool_use_tokens_excluded_from_prompt_tokens():
+ """
+ Grounding with Google Search retrieved tokens are not billed at the input token rate
+ (Google charges a separate per-request / per-query search fee), so toolUsePromptTokenCount
+ must be surfaced on prompt_tokens_details.tool_use_tokens but excluded from prompt_tokens.
+ See https://ai.google.dev/gemini-api/docs/pricing and
+ https://github.com/BerriAI/litellm/discussions/33198
+ """
+ v = VertexGeminiConfig()
+ completion_response = {
+ "candidates": [{"groundingMetadata": {"webSearchQueries": ["latest nobel physics"]}}],
+ "usageMetadata": UsageMetadata(
+ promptTokenCount=19,
+ candidatesTokenCount=304,
+ thoughtsTokenCount=122,
+ toolUsePromptTokenCount=142,
+ totalTokenCount=587,
+ ),
+ }
+
+ usage = v._calculate_usage(completion_response=completion_response)
+
+ assert usage.prompt_tokens == 19
+ assert usage.completion_tokens == 304 + 122
+ assert usage.total_tokens == 587
+ assert usage.prompt_tokens_details.tool_use_tokens == 142
+ assert usage.total_tokens - usage.prompt_tokens - usage.completion_tokens == 142
+
+
+def test_vertex_ai_url_context_tool_use_tokens_billed_as_input_tokens():
+ """
+ URL context / File Search / code execution tool-use tokens are billed as input tokens, so
+ toolUsePromptTokenCount is folded into prompt_tokens when the response is not search grounded.
+ """
+ v = VertexGeminiConfig()
+ completion_response = {
+ "candidates": [
+ {
+ "urlContextMetadata": {"urlMetadata": []},
+ "groundingMetadata": {
+ "groundingChunks": [{"web": {"uri": "https://example.com", "title": "Example"}}]
+ },
+ }
+ ],
+ "usageMetadata": UsageMetadata(
+ promptTokenCount=19,
+ candidatesTokenCount=304,
+ thoughtsTokenCount=122,
+ toolUsePromptTokenCount=142,
+ totalTokenCount=587,
+ ),
+ }
+
+ usage = v._calculate_usage(completion_response=completion_response)
+
+ assert usage.prompt_tokens == 19 + 142
+ assert usage.completion_tokens == 304 + 122
+ assert usage.total_tokens == 587
+ assert usage.prompt_tokens_details.tool_use_tokens == 142
+ assert usage.total_tokens - usage.prompt_tokens - usage.completion_tokens == 0
+
+
def test_streaming_chunk_includes_reasoning_tokens():
from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import (
ModelResponseIterator,
diff --git a/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py
index 8c6bdefedae..79c5f3ea549 100644
--- a/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py
+++ b/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py
@@ -1727,6 +1727,7 @@ class TestModelInfoEndpoint:
"gpt-3.5-turbo",
]
mock_router.get_model_access_groups.return_value = {}
+ mock_router.get_configured_token_limits.return_value = (None, None)
mock_get_key_models.return_value = ["gpt-4", "claude-3"]
mock_get_team_models.return_value = ["gpt-3.5-turbo"]
mock_get_complete_models.return_value = [
@@ -1812,6 +1813,7 @@ class TestModelInfoEndpoint:
# Setup mocks
mock_router.get_model_names.return_value = ["team-model-1"]
mock_router.get_model_access_groups.return_value = {}
+ mock_router.get_configured_token_limits.return_value = (None, None)
mock_get_key_models.return_value = []
mock_get_team_models.return_value = ["team-model-1"]
mock_get_complete_models.return_value = ["team-model-1"]
diff --git a/tests/test_litellm/proxy/proxy_server/test_team_model_name_translation.py b/tests/test_litellm/proxy/proxy_server/test_team_model_name_translation.py
index 25e84fb59a7..577af3dcffc 100644
--- a/tests/test_litellm/proxy/proxy_server/test_team_model_name_translation.py
+++ b/tests/test_litellm/proxy/proxy_server/test_team_model_name_translation.py
@@ -725,6 +725,7 @@ async def test_v1_models_translates_team_model_for_access_group_key(monkeypatch)
router.get_model_names.return_value = ["model_name_teamX_uuid9"]
router.get_model_access_groups.return_value = {"grp-a": ["model_name_teamX_uuid9"]}
router.get_fully_blocked_model_names.return_value = set()
+ router.get_configured_token_limits.return_value = (None, None)
router.model_list = [team_dep]
router.get_model_list.return_value = [team_dep]
@@ -766,6 +767,7 @@ async def test_v1_models_keeps_internal_names_when_public_name_flag_disabled(
router.get_model_names.return_value = ["model_name_teamX_uuid9"]
router.get_model_access_groups.return_value = {"grp-a": ["model_name_teamX_uuid9"]}
router.get_fully_blocked_model_names.return_value = set()
+ router.get_configured_token_limits.return_value = (None, None)
router.model_list = [team_dep]
router.get_model_list.return_value = [team_dep]
@@ -800,6 +802,7 @@ async def test_v1_models_translates_team_model_with_metadata(monkeypatch):
router.get_model_names.return_value = ["model_name_teamX_uuid9"]
router.get_model_access_groups.return_value = {"grp-a": ["model_name_teamX_uuid9"]}
router.get_fully_blocked_model_names.return_value = set()
+ router.get_configured_token_limits.return_value = (None, None)
router.model_list = [team_dep]
router.get_model_list.return_value = [team_dep]
router.get_model_group_info.return_value = None
@@ -845,6 +848,7 @@ async def test_v1_models_metadata_fallbacks_use_internal_routing_key(monkeypatch
router.get_model_names.return_value = ["model_name_teamX_uuid9"]
router.get_model_access_groups.return_value = {"grp-a": ["model_name_teamX_uuid9"]}
router.get_fully_blocked_model_names.return_value = set()
+ router.get_configured_token_limits.return_value = (None, None)
router.model_list = [team_dep]
router.get_model_list.return_value = [team_dep]
# Fallbacks are keyed on the internal routing name, as the router stores them.
@@ -901,6 +905,7 @@ async def test_v1_models_metadata_does_not_leak_other_team_fallbacks(monkeypatch
router.get_model_names.return_value = ["model_name_teamX_uuid9"]
router.get_model_access_groups.return_value = {"grp-a": ["model_name_teamX_uuid9"]}
router.get_fully_blocked_model_names.return_value = set()
+ router.get_configured_token_limits.return_value = (None, None)
router.model_list = [team_x, team_y]
router.get_model_list.return_value = [team_x, team_y]
router.fallbacks = [
@@ -1155,6 +1160,7 @@ def test_translate_team_model_names_for_listing_respects_legacy_flag():
def _public_named_router(*team_rows: dict) -> MagicMock:
router = MagicMock()
router.get_model_list.return_value = list(team_rows)
+ router.get_configured_token_limits.return_value = (None, None)
return router
From b3d05bd10b9a044ea08a1f1ce0e165ee5ba1ef35 Mon Sep 17 00:00:00 2001
From: "devin-ai-integration[bot]"
<158243242+devin-ai-integration[bot]@users.noreply.github.com>
Date: Fri, 17 Jul 2026 21:33:34 -0700
Subject: [PATCH 026/296] feat(fireworks_ai): map litellm session id to
x-session-affinity header for prompt caching (#33717)
* feat(fireworks_ai): map litellm session id to x-session-affinity header for prompt caching
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
* fix(proxy): normalize cached usage in spend logs
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
* fix(fireworks_ai): initialize chat config base class
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
* fix(fireworks_ai): normalize cached usage for spend logs
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
* fix(fireworks_ai): cover cached usage normalization
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
* fix(proxy): normalize cached usage in spend logs
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
* test(fireworks_ai): cover session id precedence
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
---------
Co-authored-by: yuneng-jiang
Co-authored-by: Mateo Wang <277851410+mateo-berri@users.noreply.github.com>
Co-authored-by: Krrish Dholakia
Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
---
.../llms/fireworks_ai/chat/transformation.py | 14 ++-
litellm/llms/fireworks_ai/common_utils.py | 24 ++++-
.../spend_tracking/spend_tracking_utils.py | 6 ++
.../test_fireworks_ai_chat_transformation.py | 100 ++++++++++++++++++
.../test_spend_tracking_utils.py | 63 +++++++++++
5 files changed, 204 insertions(+), 3 deletions(-)
diff --git a/litellm/llms/fireworks_ai/chat/transformation.py b/litellm/llms/fireworks_ai/chat/transformation.py
index d4258557fe7..319f03fea89 100644
--- a/litellm/llms/fireworks_ai/chat/transformation.py
+++ b/litellm/llms/fireworks_ai/chat/transformation.py
@@ -48,7 +48,7 @@ from ...openai.chat.gpt_transformation import (
OpenAIChatCompletionStreamingHandler,
OpenAIGPTConfig,
)
-from ..common_utils import FireworksAIException
+from ..common_utils import FireworksAIMixin, FireworksAIException
def _extract_fireworks_hidden_params(payload: dict) -> dict:
@@ -70,7 +70,7 @@ def _extract_fireworks_hidden_params(payload: dict) -> dict:
return {**top_level, **per_choice}
-class FireworksAIConfig(OpenAIGPTConfig):
+class FireworksAIConfig(FireworksAIMixin, OpenAIGPTConfig):
"""
Reference: https://docs.fireworks.ai/api-reference/post-chatcompletions
@@ -114,6 +114,16 @@ class FireworksAIConfig(OpenAIGPTConfig):
prompt_truncate_len: Optional[int] = None,
context_length_exceeded_behavior: Optional[Literal["error", "truncate"]] = None,
) -> None:
+ OpenAIGPTConfig.__init__(
+ self,
+ frequency_penalty=frequency_penalty,
+ max_tokens=max_tokens,
+ n=n,
+ stop=stop,
+ temperature=temperature,
+ top_p=top_p,
+ response_format=response_format,
+ )
locals_ = locals().copy()
for key, value in locals_.items():
if key != "self" and value is not None:
diff --git a/litellm/llms/fireworks_ai/common_utils.py b/litellm/llms/fireworks_ai/common_utils.py
index a1b6309d1e0..4e22445bcc0 100644
--- a/litellm/llms/fireworks_ai/common_utils.py
+++ b/litellm/llms/fireworks_ai/common_utils.py
@@ -12,6 +12,23 @@ class FireworksAIException(BaseLLMException):
pass
+def get_fireworks_session_id(litellm_params: dict) -> str | None:
+ params = litellm_params
+ for key in ("litellm_session_id", "session_id"):
+ value = params.get(key)
+ if value:
+ return str(value)
+ metadata = params.get("metadata")
+ if isinstance(metadata, dict):
+ value = metadata.get("session_id")
+ if value:
+ return str(value)
+ value = params.get("litellm_trace_id")
+ if value:
+ return str(value)
+ return None
+
+
class FireworksAIMixin:
"""
Common Base Config functions across Fireworks AI Endpoints
@@ -47,4 +64,9 @@ class FireworksAIMixin:
if api_key is None:
raise ValueError("FIREWORKS_API_KEY is not set")
- return {"Authorization": "Bearer {}".format(api_key), **headers}
+ validated_headers = {"Authorization": "Bearer {}".format(api_key), **headers}
+ if not any(key.lower() == "x-session-affinity" for key in validated_headers):
+ session_id = get_fireworks_session_id(litellm_params)
+ if session_id:
+ validated_headers["x-session-affinity"] = session_id
+ return validated_headers
diff --git a/litellm/proxy/spend_tracking/spend_tracking_utils.py b/litellm/proxy/spend_tracking/spend_tracking_utils.py
index b38d5e39800..23e7711b223 100644
--- a/litellm/proxy/spend_tracking/spend_tracking_utils.py
+++ b/litellm/proxy/spend_tracking/spend_tracking_utils.py
@@ -373,6 +373,12 @@ def get_logging_payload(kwargs, response_obj, start_time, end_time) -> SpendLogs
if isinstance(v, BaseModel):
v = v.model_dump()
additional_usage_values.update({k: v})
+ if "cache_read_input_tokens" not in additional_usage_values:
+ prompt_tokens_details = additional_usage_values.get("prompt_tokens_details")
+ if isinstance(prompt_tokens_details, dict):
+ cached_tokens = prompt_tokens_details.get("cached_tokens")
+ if isinstance(cached_tokens, int) and cached_tokens > 0:
+ additional_usage_values["cache_read_input_tokens"] = cached_tokens
clean_metadata["additional_usage_values"] = additional_usage_values
if litellm.cache is not None:
diff --git a/tests/test_litellm/llms/fireworks_ai/chat/test_fireworks_ai_chat_transformation.py b/tests/test_litellm/llms/fireworks_ai/chat/test_fireworks_ai_chat_transformation.py
index 03e763a4161..6809799d34f 100644
--- a/tests/test_litellm/llms/fireworks_ai/chat/test_fireworks_ai_chat_transformation.py
+++ b/tests/test_litellm/llms/fireworks_ai/chat/test_fireworks_ai_chat_transformation.py
@@ -13,6 +13,7 @@ sys.path.insert(
from litellm import get_model_info, supports_reasoning, supports_vision
from litellm.llms.fireworks_ai.chat.transformation import FireworksAIConfig
+from litellm.llms.fireworks_ai.common_utils import get_fireworks_session_id
from litellm.types.utils import (
ChatCompletionMessageToolCall,
Function,
@@ -32,6 +33,105 @@ def force_local_model_cost(monkeypatch):
litellm.model_cost = get_model_cost_map(url=litellm.model_cost_map_url)
+def test_validate_environment_sets_session_affinity_from_litellm_session_id():
+ config = FireworksAIConfig()
+
+ headers = config.validate_environment(
+ headers={},
+ model="accounts/fireworks/models/test-model",
+ messages=[],
+ optional_params={},
+ litellm_params={"litellm_session_id": "session-123"},
+ api_key="test-key",
+ )
+
+ assert headers["x-session-affinity"] == "session-123"
+
+
+def test_validate_environment_sets_session_affinity_from_metadata_session_id():
+ config = FireworksAIConfig()
+
+ headers = config.validate_environment(
+ headers={},
+ model="accounts/fireworks/models/test-model",
+ messages=[],
+ optional_params={},
+ litellm_params={"metadata": {"session_id": "metadata-session-123"}},
+ api_key="test-key",
+ )
+
+ assert headers["x-session-affinity"] == "metadata-session-123"
+
+
+def test_validate_environment_sets_session_affinity_from_session_id():
+ config = FireworksAIConfig()
+
+ headers = config.validate_environment(
+ headers={},
+ model="accounts/fireworks/models/test-model",
+ messages=[],
+ optional_params={},
+ litellm_params={"session_id": "session-id-123"},
+ api_key="test-key",
+ )
+
+ assert headers["x-session-affinity"] == "session-id-123"
+
+
+def test_validate_environment_sets_session_affinity_from_trace_id():
+ config = FireworksAIConfig()
+
+ headers = config.validate_environment(
+ headers={},
+ model="accounts/fireworks/models/test-model",
+ messages=[],
+ optional_params={},
+ litellm_params={"litellm_trace_id": "trace-id-123"},
+ api_key="test-key",
+ )
+
+ assert headers["x-session-affinity"] == "trace-id-123"
+
+
+def test_validate_environment_does_not_set_session_affinity_without_session_id():
+ config = FireworksAIConfig()
+
+ headers = config.validate_environment(
+ headers={},
+ model="accounts/fireworks/models/test-model",
+ messages=[],
+ optional_params={},
+ litellm_params={},
+ api_key="test-key",
+ )
+
+ assert "x-session-affinity" not in headers
+
+
+def test_validate_environment_preserves_explicit_session_affinity_header():
+ config = FireworksAIConfig()
+
+ headers = config.validate_environment(
+ headers={"x-session-affinity": "explicit-session"},
+ model="accounts/fireworks/models/test-model",
+ messages=[],
+ optional_params={},
+ litellm_params={"litellm_session_id": "session-123"},
+ api_key="test-key",
+ )
+
+ assert headers["x-session-affinity"] == "explicit-session"
+
+
+def test_get_fireworks_session_id_prefers_litellm_session_id_over_trace_id():
+ assert (
+ get_fireworks_session_id(
+ {"litellm_session_id": "session-123", "litellm_trace_id": "trace-123"}
+ )
+ == "session-123"
+ )
+
+
def test_handle_message_content_with_tool_calls():
config = FireworksAIConfig()
message = Message(
diff --git a/tests/test_litellm/proxy/spend_tracking/test_spend_tracking_utils.py b/tests/test_litellm/proxy/spend_tracking/test_spend_tracking_utils.py
index 9a8f8146d6f..51d72aa2ab2 100644
--- a/tests/test_litellm/proxy/spend_tracking/test_spend_tracking_utils.py
+++ b/tests/test_litellm/proxy/spend_tracking/test_spend_tracking_utils.py
@@ -46,6 +46,69 @@ from litellm.types.utils import (
)
+def _get_additional_usage_values_for_usage(usage: litellm.Usage) -> dict:
+ payload = get_logging_payload(
+ kwargs={
+ "model": "gpt-4o-mini",
+ "litellm_params": {"metadata": {"user_api_key": "test-key"}},
+ },
+ response_obj=litellm.ModelResponse(
+ id="chatcmpl-test",
+ choices=[],
+ usage=usage,
+ ),
+ start_time=datetime.datetime.now(timezone.utc),
+ end_time=datetime.datetime.now(timezone.utc),
+ )
+ metadata = json.loads(payload["metadata"])
+ return metadata["additional_usage_values"]
+
+
+def test_get_logging_payload_maps_openai_cached_tokens_to_cache_read_input_tokens():
+ additional_usage_values = _get_additional_usage_values_for_usage(
+ litellm.Usage(
+ prompt_tokens=10,
+ completion_tokens=2,
+ total_tokens=12,
+ prompt_tokens_details={"cached_tokens": 123},
+ )
+ )
+
+ assert additional_usage_values["cache_read_input_tokens"] == 123
+ assert additional_usage_values["prompt_tokens_details"]["cached_tokens"] == 123
+
+
+def test_get_logging_payload_preserves_anthropic_cache_read_input_tokens():
+ additional_usage_values = _get_additional_usage_values_for_usage(
+ litellm.Usage(
+ prompt_tokens=10,
+ completion_tokens=2,
+ total_tokens=12,
+ prompt_tokens_details={"cached_tokens": 123},
+ cache_read_input_tokens=456,
+ )
+ )
+
+ assert additional_usage_values["cache_read_input_tokens"] == 456
+
+
+@pytest.mark.parametrize(
+ "prompt_tokens_details",
+ [None, {"cached_tokens": 0}],
+)
+def test_get_logging_payload_does_not_map_missing_or_zero_cached_tokens(prompt_tokens_details):
+ additional_usage_values = _get_additional_usage_values_for_usage(
+ litellm.Usage(
+ prompt_tokens=10,
+ completion_tokens=2,
+ total_tokens=12,
+ prompt_tokens_details=prompt_tokens_details,
+ )
+ )
+
+ assert "cache_read_input_tokens" not in additional_usage_values
+
+
def test_sanitize_request_body_for_spend_logs_payload_basic():
request_body = {
"messages": [{"role": "user", "content": "Hello, how are you?"}],
From 010b20072d20f043650ab654e2c0190b1c9da1fb Mon Sep 17 00:00:00 2001
From: "devin-ai-integration[bot]"
<158243242+devin-ai-integration[bot]@users.noreply.github.com>
Date: Sat, 18 Jul 2026 10:26:48 -0700
Subject: [PATCH 027/296] fix(router): enforce context-window pre-call checks
for Responses API input (#33706)
* fix(router): enforce context-window pre-call checks for Responses API input
* test(router): cover _count_pre_call_check_tokens across API surfaces
* fix(router): count Responses instructions and skip pre-call token count when no input
* fix(router): forward Responses input into deployment selection for context-window checks
---------
Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
---
litellm/router.py | 56 +++++++++-
tests/test_litellm/test_router.py | 179 ++++++++++++++++++++++++++++++
2 files changed, 229 insertions(+), 6 deletions(-)
diff --git a/litellm/router.py b/litellm/router.py
index b1a5405ebf1..0b1471dc527 100644
--- a/litellm/router.py
+++ b/litellm/router.py
@@ -4461,6 +4461,7 @@ class Router:
model=model,
request_kwargs=kwargs,
messages=kwargs.get("messages", None),
+ input=kwargs.get("input", None),
specific_deployment=kwargs.pop("specific_deployment", None),
)
except Exception as e:
@@ -4608,6 +4609,7 @@ class Router:
deployment = self.get_available_deployment(
model=model,
messages=kwargs.get("messages", None),
+ input=kwargs.get("input", None),
specific_deployment=kwargs.pop("specific_deployment", None),
request_kwargs=kwargs,
)
@@ -10002,11 +10004,44 @@ class Router:
client = self.cache.get_cache(key=cache_key, parent_otel_span=parent_otel_span)
return client
+ def _count_pre_call_check_tokens(
+ self,
+ messages: list[dict[str, str]] | None,
+ input: str | list | None,
+ instructions: str | None = None,
+ ) -> int:
+ """
+ Count input tokens for context-window pre-call checks.
+
+ Chat Completions send `messages`; the Responses API sends `input` (a string or
+ a list of Responses input items) plus an optional `instructions` system prompt.
+ The Responses payload is normalized to chat messages via the shared
+ LiteLLMCompletionResponsesConfig transform so the same token_counter path covers
+ both API surfaces and `instructions` tokens are included in the count.
+ """
+ if messages is not None:
+ return litellm.token_counter(messages=messages)
+ if input is not None:
+ from openai.types.responses.response_create_params import ResponseInputParam
+
+ from litellm.responses.litellm_completion_transformation.transformation import (
+ LiteLLMCompletionResponsesConfig,
+ )
+
+ typed_input = cast(str | ResponseInputParam, input) # cast-ok: str | list matches transform input
+ input_messages = LiteLLMCompletionResponsesConfig.transform_responses_api_input_to_messages(
+ input=typed_input,
+ responses_api_request={"instructions": instructions} if instructions is not None else {},
+ )
+ return litellm.token_counter(messages=cast(list, input_messages)) # cast-ok: transformed chat messages
+ raise ValueError("Either messages or input must be provided to count tokens")
+
def _pre_call_checks(
self,
model: str,
healthy_deployments: List,
- messages: List[Dict[str, str]],
+ messages: list[dict[str, str]] | None = None,
+ input: str | list | None = None,
request_kwargs: Optional[dict] = None,
):
"""
@@ -10036,6 +10071,10 @@ class Router:
_rate_limit_error = False
parent_otel_span = _get_parent_otel_span_from_kwargs(request_kwargs)
+ raw_instructions = request_kwargs.get("instructions") if request_kwargs else None
+ instructions = raw_instructions if isinstance(raw_instructions, str) else None
+ has_countable_input = messages is not None or input is not None
+
## get model group RPM ##
dt = get_utc_datetime()
current_minute = dt.strftime("%H-%M")
@@ -10058,10 +10097,12 @@ class Router:
_deployment_model = base_model or _litellm_params.get("model", None)
max_input_tokens = model_info.get("max_input_tokens") if isinstance(model_info, dict) else None
- if isinstance(max_input_tokens, int):
+ if isinstance(max_input_tokens, int) and has_countable_input:
if input_tokens is None:
try:
- input_tokens = litellm.token_counter(messages=messages)
+ input_tokens = self._count_pre_call_check_tokens(
+ messages=messages, input=input, instructions=instructions
+ )
except Exception as e:
verbose_router_logger.error(
"litellm.router.py::_pre_call_checks: failed to count tokens. Returning initial list of deployments. Got - {}".format(
@@ -10526,11 +10567,12 @@ class Router:
parent_otel_span=parent_otel_span,
)
- if self.enable_pre_call_checks and messages is not None:
+ if self.enable_pre_call_checks and (messages is not None or input is not None):
healthy_deployments = self._pre_call_checks(
model=model,
healthy_deployments=cast(List[Dict], healthy_deployments),
messages=messages,
+ input=input,
request_kwargs=request_kwargs,
)
# check if user wants to do tag based routing
@@ -11041,11 +11083,12 @@ class Router:
healthy_deployments = self._filter_blocked_deployments(healthy_deployments)
# filter pre-call checks
- if self.enable_pre_call_checks and messages is not None:
+ if self.enable_pre_call_checks and (messages is not None or input is not None):
healthy_deployments = self._pre_call_checks(
model=model,
healthy_deployments=healthy_deployments,
messages=messages,
+ input=input,
request_kwargs=request_kwargs,
)
@@ -11195,11 +11238,12 @@ class Router:
pass_through_deployments = self._filter_blocked_deployments(pass_through_deployments)
# 5. Apply pre-call checks (if enabled)
- if self.enable_pre_call_checks and messages is not None:
+ if self.enable_pre_call_checks and (messages is not None or input is not None):
pass_through_deployments = self._pre_call_checks(
model=model,
healthy_deployments=pass_through_deployments,
messages=messages,
+ input=input,
request_kwargs=request_kwargs,
)
diff --git a/tests/test_litellm/test_router.py b/tests/test_litellm/test_router.py
index 55c09e6cac4..76a6e3c1bbe 100644
--- a/tests/test_litellm/test_router.py
+++ b/tests/test_litellm/test_router.py
@@ -2864,6 +2864,185 @@ def test_pre_call_checks_counts_once_and_filters_on_max_input_tokens(monkeypatch
assert calls == [1]
+def test_pre_call_checks_counts_tokens_from_responses_input_string(monkeypatch):
+ """
+ Responses API calls pass `input` (str) instead of `messages`. Context-window
+ checks must count tokens from `input` and filter deployments over the limit. Uses
+ the real token_counter so the transform + counting path is a true regression guard.
+ """
+ router = litellm.Router(
+ model_list=[
+ {"model_name": "m", "litellm_params": {"model": "gpt-3.5-turbo"}},
+ ],
+ enable_pre_call_checks=True,
+ )
+ monkeypatch.setattr(
+ router, "get_router_model_info", lambda **kwargs: {"max_input_tokens": 1}
+ )
+
+ deployments = [
+ {"litellm_params": {"model": "gpt-3.5-turbo"}, "model_info": {"id": "d1"}},
+ ]
+ with pytest.raises(litellm.ContextWindowExceededError):
+ router._pre_call_checks(
+ model="m",
+ healthy_deployments=deployments,
+ input="a very long prompt that exceeds the tiny context window",
+ )
+
+
+def test_pre_call_checks_counts_tokens_from_responses_input_list(monkeypatch):
+ """
+ Responses API `input` can be a list of input items. It must be normalized to
+ chat messages and counted so oversized requests are filtered out. Uses the real
+ token_counter (no mock) so the transform + counting path is a true regression guard.
+ """
+ router = litellm.Router(
+ model_list=[
+ {"model_name": "m", "litellm_params": {"model": "gpt-3.5-turbo"}},
+ ],
+ enable_pre_call_checks=True,
+ )
+ monkeypatch.setattr(
+ router, "get_router_model_info", lambda **kwargs: {"max_input_tokens": 1}
+ )
+
+ deployments = [
+ {"litellm_params": {"model": "gpt-3.5-turbo"}, "model_info": {"id": "d1"}},
+ ]
+ with pytest.raises(litellm.ContextWindowExceededError):
+ router._pre_call_checks(
+ model="m",
+ healthy_deployments=deployments,
+ input=[
+ {"role": "user", "content": "count these tokens against the one token limit please"},
+ ],
+ )
+
+
+def test_pre_call_checks_counts_responses_instructions_tokens(monkeypatch):
+ """
+ Responses API `instructions` become a system message the model receives, so their
+ tokens must be counted too. A request whose `input` alone fits under the limit but
+ whose `input` + `instructions` exceeds it must be filtered (regression for the
+ context-window check under-filtering when instructions were ignored).
+ """
+ router = litellm.Router(
+ model_list=[
+ {"model_name": "m", "litellm_params": {"model": "gpt-3.5-turbo"}},
+ ],
+ enable_pre_call_checks=True,
+ )
+
+ deployments = [
+ {"litellm_params": {"model": "gpt-3.5-turbo"}, "model_info": {"id": "d1"}},
+ ]
+
+ short_input = "hi"
+ long_instructions = "you are a helpful assistant. " * 20
+
+ input_only_tokens = router._count_pre_call_check_tokens(messages=None, input=short_input)
+ with_instructions_tokens = router._count_pre_call_check_tokens(
+ messages=None, input=short_input, instructions=long_instructions
+ )
+ assert with_instructions_tokens > input_only_tokens
+
+ monkeypatch.setattr(
+ router, "get_router_model_info", lambda **kwargs: {"max_input_tokens": input_only_tokens}
+ )
+ with pytest.raises(litellm.ContextWindowExceededError):
+ router._pre_call_checks(
+ model="m",
+ healthy_deployments=deployments,
+ input=short_input,
+ request_kwargs={"instructions": long_instructions},
+ )
+
+
+def test_count_pre_call_check_tokens_across_api_surfaces():
+ """
+ _count_pre_call_check_tokens must count tokens from chat `messages`, a Responses
+ API string `input`, and a Responses API list `input`, and raise when given neither.
+ """
+ router = litellm.Router(
+ model_list=[
+ {"model_name": "m", "litellm_params": {"model": "gpt-3.5-turbo"}},
+ ],
+ )
+
+ messages_tokens = router._count_pre_call_check_tokens(
+ messages=[{"role": "user", "content": "hello world"}], input=None
+ )
+ string_input_tokens = router._count_pre_call_check_tokens(messages=None, input="hello world")
+ list_input_tokens = router._count_pre_call_check_tokens(
+ messages=None, input=[{"role": "user", "content": "hello world"}]
+ )
+
+ assert messages_tokens > 0
+ assert string_input_tokens > 0
+ assert list_input_tokens > 0
+
+ with pytest.raises(ValueError):
+ router._count_pre_call_check_tokens(messages=None, input=None)
+
+
+def test_pre_call_checks_no_messages_or_input_does_not_crash(monkeypatch):
+ """
+ When neither messages nor input is provided (e.g. endpoints without prompt text),
+ token counting is skipped gracefully and all deployments are returned.
+ """
+ router = litellm.Router(
+ model_list=[
+ {"model_name": "m", "litellm_params": {"model": "gpt-3.5-turbo"}},
+ ],
+ enable_pre_call_checks=True,
+ )
+ monkeypatch.setattr(
+ router, "get_router_model_info", lambda **kwargs: {"max_input_tokens": 5}
+ )
+
+ counted: list[dict] = []
+ original = router._count_pre_call_check_tokens
+ monkeypatch.setattr(
+ router,
+ "_count_pre_call_check_tokens",
+ lambda **kwargs: counted.append(kwargs) or original(**kwargs),
+ )
+
+ deployments = [
+ {"litellm_params": {"model": "gpt-3.5-turbo"}, "model_info": {"id": "d1"}},
+ ]
+ result = router._pre_call_checks(model="m", healthy_deployments=deployments)
+ assert len(result) == 1
+ assert counted == [] # token counting skipped entirely, so no misleading error is logged
+
+
+@pytest.mark.asyncio
+async def test_aresponses_enforces_context_window_pre_call_check():
+ """
+ End-to-end router regression: a Responses API call whose `input` exceeds the
+ deployment's max_input_tokens must be filtered by the pre-call check, raising
+ ContextWindowExceededError instead of being silently routed. This guards the
+ wiring that forwards `input` from the generic-call path into deployment selection
+ (the deployment uses mock_response, so the check must trip before any real call).
+ """
+ router = litellm.Router(
+ model_list=[
+ {
+ "model_name": "small-ctx",
+ "litellm_params": {"model": "gpt-3.5-turbo", "mock_response": "hi"},
+ "model_info": {"max_input_tokens": 5},
+ }
+ ],
+ enable_pre_call_checks=True,
+ )
+ with pytest.raises(litellm.ContextWindowExceededError):
+ await router.aresponses(
+ model="small-ctx",
+ input="this responses input is definitely much longer than five tokens for sure",
+ )
+
+
def test_get_deployment_model_info_base_model_flow():
"""Test that get_deployment_model_info correctly handles the base model flow"""
from unittest.mock import patch
From 4a297dd6114cdaa1ba6795c33b45bc9b8f5fdd8b Mon Sep 17 00:00:00 2001
From: "devin-ai-integration[bot]"
<158243242+devin-ai-integration[bot]@users.noreply.github.com>
Date: Sat, 18 Jul 2026 10:52:27 -0700
Subject: [PATCH 028/296] fix(otel): restore proxy-level error.* attributes on
v2 failure spans (LIT-4179) (#33664)
* fix(otel): restore proxy-level error.* attributes on v2 failure spans (LIT-4179)
* refactor(otel): narrow v2 failure hook return type to drop fastapi import (LIT-4179)
---------
Co-authored-by: yucheng-berri
---
litellm/integrations/otel/emitter.py | 55 ++++++---
litellm/integrations/otel/logger.py | 79 +++++++++++-
litellm/proxy/proxy_server.py | 21 ++--
.../integrations/otel/test_otel_v2_emitter.py | 42 +++++++
.../integrations/otel/test_otel_v2_logger.py | 114 ++++++++++++++++++
.../proxy_server/test_exception_handlers.py | 41 +++++++
6 files changed, 328 insertions(+), 24 deletions(-)
diff --git a/litellm/integrations/otel/emitter.py b/litellm/integrations/otel/emitter.py
index f97f8b8394c..8651cf586cd 100644
--- a/litellm/integrations/otel/emitter.py
+++ b/litellm/integrations/otel/emitter.py
@@ -72,6 +72,42 @@ def _stamp_litellm_error_attributes(span: Span, error: SpanError) -> None:
span.set_attribute(LiteLLMError.LLM_PROVIDER, error.llm_provider)
+def stamp_error(
+ span: Span,
+ error: SpanError,
+ *,
+ record_event: bool = True,
+ set_status: bool = True,
+) -> tuple[str, str] | None:
+ """Stamp the full v2 error attribute set on ``span`` and return the resolved
+ ``(error_type, message)`` pair, or ``None`` when the error carries neither a
+ type nor a message.
+
+ Shared by the LLM-call span (``finish_span``) and the proxy-level failure
+ spans (the FastAPI SERVER span and the ``auth`` phase span) so every v2 error
+ span carries identical keys. The semconv ``exception`` event rides alongside
+ the attributes so backends that map unknown string attrs to a truncated
+ ``keyword`` (e.g. Elasticsearch's 1024-char ``ignore_above``) still see the
+ full untruncated message on the recognized event field. ``record_event`` and
+ ``set_status`` are opt-outs for callers whose span lifecycle (``use_span``) or
+ owner (the FastAPI instrumentor) already records the event or the status.
+ """
+ if not (error.error_type or error.message):
+ return None
+ error_type = error.error_type or "error"
+ message = error.message or error.error_type or "error"
+ _stamp_otel_error_attributes(span, error_type, message)
+ _stamp_litellm_error_attributes(span, error)
+ if set_status:
+ span.set_status(Status(StatusCode.ERROR, message))
+ if record_event:
+ span.add_event(
+ ExceptionEvent.NAME,
+ {ExceptionEvent.TYPE: error_type, ExceptionEvent.MESSAGE: message},
+ )
+ return error_type, message
+
+
class SpanEmitter:
def __init__(
self,
@@ -212,21 +248,10 @@ class SpanEmitter:
)
else None
)
- if error and (error.error_type or error.message):
- error_type = error.error_type or "error"
- message = error.message or error.error_type or "error"
- _stamp_otel_error_attributes(span, error_type, message)
- _stamp_litellm_error_attributes(span, error)
- span.set_status(Status(StatusCode.ERROR, message))
- # Also emit the semconv ``exception`` event so backends that
- # dynamic-map unknown string span attrs to ``keyword`` (e.g.
- # Elasticsearch with a 1024-char ``ignore_above``) still see the
- # full untruncated message on the recognized event field.
- span.add_event(
- ExceptionEvent.NAME,
- {ExceptionEvent.TYPE: error_type, ExceptionEvent.MESSAGE: message},
- )
- if self._event_recorder is not None and role is SpanRole.LLM_CALL:
+ if error:
+ stamped = stamp_error(span, error)
+ if stamped is not None and self._event_recorder is not None and role is SpanRole.LLM_CALL:
+ error_type, message = stamped
self._event_recorder.record_operation_exception(
span_context=span.get_span_context(),
error_type=error_type,
diff --git a/litellm/integrations/otel/logger.py b/litellm/integrations/otel/logger.py
index be72fabd387..778f5342e90 100644
--- a/litellm/integrations/otel/logger.py
+++ b/litellm/integrations/otel/logger.py
@@ -24,7 +24,7 @@ from litellm.integrations.otel.plumbing.context import (
set_request_baggage,
set_request_root_span,
)
-from litellm.integrations.otel.emitter import SpanEmitter
+from litellm.integrations.otel.emitter import SpanEmitter, stamp_error
from litellm.integrations.otel.mappers import resolve_mappers
from litellm.integrations.otel.model.metadata import (
LLMCallEvent,
@@ -59,6 +59,7 @@ from litellm.integrations.otel.model.spans import SpanRole, span_role_for_servic
from litellm.integrations.otel.model.utils import to_ns
if TYPE_CHECKING:
+ from litellm.proxy._types import UserAPIKeyAuth
from litellm.types.utils import (
StandardLoggingGuardrailInformation,
StandardLoggingPayload,
@@ -66,6 +67,33 @@ if TYPE_CHECKING:
LITELLM_TRACER_NAME = "litellm"
+
+def _span_error_from_exception(
+ exception: "Exception | None",
+ *,
+ status_code: int | None = None,
+ traceback_str: str | None = None,
+) -> SpanError:
+ """A ``SpanError`` for a proxy-level failure that never produced a
+ ``StandardLoggingPayload`` (auth / validation / malformed-body rejections),
+ mirroring ``_parse_error``'s field mapping so it stamps the same v2 keys a
+ failed LLM call does. ``status_code`` pins ``error.code`` to the real response
+ status, matching v1's SERVER-span behavior."""
+ from litellm.litellm_core_utils.litellm_logging import StandardLoggingPayloadSetup
+
+ info = StandardLoggingPayloadSetup.get_error_information(
+ original_exception=exception,
+ traceback_str=traceback_str,
+ )
+ return SpanError(
+ error_type=info.get("error_class") or info.get("error_code") or None,
+ message=info.get("error_message") or None,
+ code=str(status_code) if status_code is not None else (info.get("error_code") or None),
+ stack_trace=info.get("traceback") or None,
+ llm_provider=info.get("llm_provider") or None,
+ )
+
+
# Any callback whose class belongs to one of these modules is "the OTel
# callback" for proxy-global-registration purposes.
_OTEL_MODULES = (
@@ -558,7 +586,12 @@ class OpenTelemetryV2(CustomLogger):
def start_phase_span(self, name: str) -> "Iterator[Span]":
span = self._emitter.start_span(SpanRole.SERVICE, name)
with use_span(span, end_on_exit=True):
- yield span
+ try:
+ yield span
+ except Exception as exc:
+ if is_recordable_span(span):
+ stamp_error(span, _span_error_from_exception(exc), record_event=False, set_status=False)
+ raise
async def async_pre_call_hook(
self,
@@ -573,6 +606,48 @@ class OpenTelemetryV2(CustomLogger):
)
return data
+ def record_error_attributes_on_span(
+ self,
+ span: "Span | None",
+ exception: "Exception | None",
+ status_code: int,
+ ) -> None:
+ """Stamp the v2 error.* attributes on the FastAPI-owned SERVER span for a
+ failure that dies before any LLM-call span exists (malformed body, auth /
+ validation rejection). Called from the proxy's global exception handler via
+ ``_close_dangling_otel_server_span``. The instrumentor still owns the span's
+ status and lifecycle, so this only decorates it — never sets status, never
+ ends it — and emits no exception event, matching v1's SERVER-span behavior
+ and avoiding a duplicate of the event ``async_post_call_failure_hook`` or
+ the ``auth`` phase span already records."""
+ if span is None or not is_recordable_span(span):
+ return
+ stamp_error(
+ span,
+ _span_error_from_exception(exception, status_code=status_code),
+ record_event=False,
+ set_status=False,
+ )
+
+ async def async_post_call_failure_hook(
+ self,
+ request_data: dict,
+ original_exception: Exception,
+ user_api_key_dict: "UserAPIKeyAuth",
+ traceback_str: "str | None" = None,
+ ) -> None:
+ """Stamp error.* on the request's root SERVER span for a proxy-level
+ failure that never reached an LLM call (empty body rejected in the
+ endpoint, auth failure), so the failed request carries the same error keys
+ a failed LLM call does. v1's ``OpenTelemetry`` implemented this same hook;
+ v2 lost it when it stopped subclassing ``OpenTelemetry``, which is the
+ LIT-4179 regression for pre-call failures."""
+ span = request_root_span() or user_api_key_dict.parent_otel_span
+ if span is None or not is_recordable_span(span):
+ return None
+ stamp_error(span, _span_error_from_exception(original_exception, traceback_str=traceback_str))
+ return None
+
def emit_guardrail_span(self, entry: "StandardLoggingGuardrailInformation") -> None:
# Emitted by the guardrail-recording code the moment a guardrail finishes,
# not from a post-call hook — that hook does not fire on every path (a
diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py
index 3b640ee54fd..aed345c5db4 100644
--- a/litellm/proxy/proxy_server.py
+++ b/litellm/proxy/proxy_server.py
@@ -1395,19 +1395,25 @@ def _close_dangling_otel_server_span(request: Request, status_code: int, exc: Op
if open_telemetry_logger is None:
return
# Under OTel V2 the FastAPI instrumentor owns the server span (parent_otel_span
- # is that same span), and it records the error + ends it itself. Ending it here
- # would end it early — losing the http.* attributes the instrumentor stamps on
- # completion — and double-end it. Leave it to the instrumentor.
+ # is that same span) and ends it itself with the http.* attributes stamped on
+ # completion. The instrumentor only records an error when the exception reaches
+ # it uncaught, but these handlers swallow it into a JSONResponse, so it never
+ # does; stamp the error.* attributes here (without ending or re-statusing the
+ # span, which the instrumentor still owns) so pre-call failures carry the error
+ # like v1 did. Otherwise close and annotate the dangling span ourselves.
try:
from litellm.integrations.otel.model.config import is_otel_v2_enabled
- if is_otel_v2_enabled():
- return
+ v2_enabled = is_otel_v2_enabled()
except Exception:
- pass
+ v2_enabled = False
try:
from opentelemetry.trace import Status, StatusCode
+ if v2_enabled:
+ if status_code >= 400:
+ open_telemetry_logger.record_error_attributes_on_span(parent_otel_span, exc, status_code)
+ return
open_telemetry_logger.set_response_status_code_attribute(parent_otel_span, status_code)
if status_code >= 400:
open_telemetry_logger.record_error_attributes_on_span(parent_otel_span, exc, status_code)
@@ -1416,7 +1422,8 @@ def _close_dangling_otel_server_span(request: Request, status_code: int, exc: Op
except Exception as e:
verbose_proxy_logger.debug("Error closing dangling OTEL SERVER span: %s", str(e))
finally:
- request.state.parent_otel_span = None
+ if not v2_enabled:
+ request.state.parent_otel_span = None
@app.exception_handler(RequestValidationError)
diff --git a/tests/test_litellm/integrations/otel/test_otel_v2_emitter.py b/tests/test_litellm/integrations/otel/test_otel_v2_emitter.py
index 48190a798da..6b1da4c2952 100644
--- a/tests/test_litellm/integrations/otel/test_otel_v2_emitter.py
+++ b/tests/test_litellm/integrations/otel/test_otel_v2_emitter.py
@@ -16,10 +16,12 @@ from litellm.integrations.otel import ( # noqa: E402
from litellm.integrations.otel.plumbing import context as ctx_mod # noqa: E402
from litellm.integrations.otel.plumbing import providers # noqa: E402
from litellm.integrations.otel.emitter import SpanEmitter # noqa: E402
+from litellm.integrations.otel.emitter import stamp_error # noqa: E402
from litellm.integrations.otel.model.payloads import ( # noqa: E402
GuardrailSpanData,
LLMCallSpanData,
ServiceSpanData,
+ SpanError,
)
from litellm.integrations.otel.model.spans import SPAN_REGISTRY, SpanRole # noqa: E402
@@ -155,6 +157,46 @@ def test_error_span_sets_status_and_error_type():
assert span.attributes["error.type"] == "RateLimitError"
+def test_stamp_error_writes_full_attribute_set_and_event():
+ engine, exporter = _engine()
+ span = engine.start_span(SpanRole.PROXY_REQUEST, "POST /chat/completions")
+ result = stamp_error(
+ span, SpanError("ProxyException", "boom", code="401", stack_trace="tb", llm_provider="anthropic")
+ )
+ span.end()
+ (s,) = exporter.get_finished_spans()
+ assert result == ("ProxyException", "boom")
+ assert s.attributes["error.type"] == "ProxyException"
+ assert s.attributes["error.message"] == "boom"
+ assert s.attributes["litellm.provider.error.code"] == "401"
+ assert s.attributes["litellm.provider.error.stack_trace"] == "tb"
+ assert s.attributes["litellm.provider.error.llm_provider"] == "anthropic"
+ assert s.status.status_code is StatusCode.ERROR
+ assert [e.name for e in s.events] == ["exception"]
+
+
+def test_stamp_error_opt_outs_skip_status_and_event():
+ engine, exporter = _engine()
+ span = engine.start_span(SpanRole.PROXY_REQUEST, "POST /chat/completions")
+ stamp_error(span, SpanError("ProxyException", "boom", code="401"), record_event=False, set_status=False)
+ span.end()
+ (s,) = exporter.get_finished_spans()
+ assert s.attributes["error.type"] == "ProxyException"
+ assert s.attributes["litellm.provider.error.code"] == "401"
+ assert s.status.status_code is StatusCode.UNSET
+ assert s.events == ()
+
+
+def test_stamp_error_without_type_or_message_is_noop():
+ engine, exporter = _engine()
+ span = engine.start_span(SpanRole.PROXY_REQUEST, "POST /chat/completions")
+ assert stamp_error(span, SpanError()) is None
+ span.end()
+ (s,) = exporter.get_finished_spans()
+ assert "error.type" not in s.attributes
+ assert s.status.status_code is StatusCode.UNSET
+
+
def test_hierarchy_and_kinds_match_registry():
engine, exporter = _engine()
data = LLMCallSpanData.from_standard_logging_payload(_payload())
diff --git a/tests/test_litellm/integrations/otel/test_otel_v2_logger.py b/tests/test_litellm/integrations/otel/test_otel_v2_logger.py
index b5e077e3561..5f6002f4cdf 100644
--- a/tests/test_litellm/integrations/otel/test_otel_v2_logger.py
+++ b/tests/test_litellm/integrations/otel/test_otel_v2_logger.py
@@ -860,6 +860,120 @@ def test_guardrail_span_anchors_to_root_inside_active_phase_span():
assert guard.parent.span_id != auth_span.get_span_context().span_id
+# --------------------------------------------------------------------------- #
+# LIT-4179 — proxy-level failures that never reach an LLM call must still stamp
+# the structured error.* attributes onto the request's spans, restoring the v1
+# behavior v2 dropped when it stopped subclassing ``OpenTelemetry``.
+# --------------------------------------------------------------------------- #
+
+
+def _proxy_exc(message, code):
+ from litellm.proxy._types import ProxyException
+
+ return ProxyException(message=message, type="bad_request_error", param=None, code=code)
+
+
+def test_async_post_call_failure_hook_stamps_error_on_root_span():
+ """PATH B: an endpoint-level failure (empty body rejected before dispatch)
+ reaches ``async_post_call_failure_hook``; it must stamp error.* + an exception
+ event on the anchored request root span."""
+ from litellm.proxy._types import UserAPIKeyAuth
+
+ logger, exporter = _logger()
+ server = logger._emitter.start_span(SpanRole.PROXY_REQUEST, LITELLM_PROXY_REQUEST_SPAN_NAME)
+ set_request_root_span(server)
+ exc = _proxy_exc("litellm.BadRequestError: messages is required", 400)
+ result = asyncio.run(
+ logger.async_post_call_failure_hook(
+ request_data={}, original_exception=exc, user_api_key_dict=UserAPIKeyAuth()
+ )
+ )
+ server.end()
+ assert result is None
+ (span,) = exporter.get_finished_spans()
+ assert span.attributes["error.type"] == "ProxyException"
+ assert "messages is required" in span.attributes["error.message"]
+ assert span.attributes["litellm.provider.error.code"] == "400"
+ assert span.status.status_code is StatusCode.ERROR
+ assert any(e.name == "exception" for e in span.events)
+
+
+def test_async_post_call_failure_hook_falls_back_to_user_api_key_parent_span():
+ """With no anchor set (a path that never captured the root), the hook must fall
+ back to ``user_api_key_dict.parent_otel_span`` rather than dropping the error."""
+ from litellm.proxy._types import UserAPIKeyAuth
+
+ logger, exporter = _logger()
+ server = logger._emitter.start_span(SpanRole.PROXY_REQUEST, LITELLM_PROXY_REQUEST_SPAN_NAME)
+ asyncio.run(
+ logger.async_post_call_failure_hook(
+ request_data={},
+ original_exception=_proxy_exc("boom", 401),
+ user_api_key_dict=UserAPIKeyAuth(parent_otel_span=server),
+ )
+ )
+ server.end()
+ (span,) = exporter.get_finished_spans()
+ assert span.attributes["error.type"] == "ProxyException"
+ assert span.attributes["litellm.provider.error.code"] == "401"
+
+
+def test_record_error_attributes_on_span_decorates_without_ending():
+ """PATH A: a failure that dies before any LLM-call span (malformed body,
+ validation) is stamped onto the instrumentor-owned SERVER span. The method must
+ not end the span or emit a duplicate exception event, and must pin error.code
+ to the real response status (not the exception's own code)."""
+ logger, exporter = _logger()
+ server = logger._emitter.start_span(SpanRole.PROXY_REQUEST, LITELLM_PROXY_REQUEST_SPAN_NAME)
+ logger.record_error_attributes_on_span(server, _proxy_exc("Invalid JSON body", 400), 422)
+ assert server.is_recording()
+ server.end()
+ (span,) = exporter.get_finished_spans()
+ assert span.attributes["error.type"] == "ProxyException"
+ assert span.attributes["error.message"] == "Invalid JSON body"
+ assert span.attributes["litellm.provider.error.code"] == "422"
+ assert all(e.name != "exception" for e in span.events)
+
+
+def test_record_error_attributes_on_span_ignores_below_400_and_missing_span():
+ logger, _ = _logger()
+ server = logger._emitter.start_span(SpanRole.PROXY_REQUEST, LITELLM_PROXY_REQUEST_SPAN_NAME)
+ logger.record_error_attributes_on_span(None, _proxy_exc("boom", 400), 400) # no span → no-op
+ logger.record_error_attributes_on_span(server, None, 400) # no exception → no-op
+ server.end()
+ assert "error.type" not in (server.attributes or {})
+
+
+def test_start_phase_span_stamps_error_attributes_on_failure():
+ """An ``auth`` phase span that dies (expired key) must carry the structured
+ error.* attributes, not only the exception event ``use_span`` records."""
+ logger, exporter = _logger()
+ server = logger._emitter.start_span(SpanRole.PROXY_REQUEST, LITELLM_PROXY_REQUEST_SPAN_NAME)
+ set_request_root_span(server)
+ exc = _proxy_exc("Authentication Error, ExpiredToken", 401)
+ with trace.use_span(server, end_on_exit=False):
+ with contextlib.suppress(Exception):
+ with logger.start_phase_span("auth /chat/completions"):
+ raise exc
+ server.end()
+ by_name = {s.name: s for s in exporter.get_finished_spans()}
+ auth = by_name["auth /chat/completions"]
+ assert auth.attributes["error.type"] == "ProxyException"
+ assert "ExpiredToken" in auth.attributes["error.message"]
+ assert auth.attributes["litellm.provider.error.code"] == "401"
+ assert auth.status.status_code is StatusCode.ERROR
+ assert any(e.name == "exception" for e in auth.events)
+
+
+def test_start_phase_span_success_carries_no_error():
+ logger, exporter = _logger()
+ with logger.start_phase_span("auth /chat/completions"):
+ pass
+ (span,) = exporter.get_finished_spans()
+ assert "error.type" not in span.attributes
+ assert span.status.status_code is not StatusCode.ERROR
+
+
def test_real_logging_pre_call_opens_span_end_to_end():
"""Regression guard: a real ``LiteLLMLoggingObj.pre_call`` must fire
``log_pre_api_call`` on the V2 logger (via ``litellm.input_callback``), so the
diff --git a/tests/test_litellm/proxy/proxy_server/test_exception_handlers.py b/tests/test_litellm/proxy/proxy_server/test_exception_handlers.py
index cf92f9cd12b..e4bf06991b4 100644
--- a/tests/test_litellm/proxy/proxy_server/test_exception_handlers.py
+++ b/tests/test_litellm/proxy/proxy_server/test_exception_handlers.py
@@ -124,6 +124,47 @@ def test_close_dangling_otel_server_span_records_status_and_ends(monkeypatch):
}
+def test_close_dangling_otel_server_span_v2_stamps_error_without_ending(monkeypatch):
+ """LIT-4179: under OTel v2 the FastAPI instrumentor owns the SERVER span, so
+ the handler must only stamp error.* on it (via record_error_attributes_on_span)
+ and must NOT set status, end the span, or clear request state — otherwise the
+ instrumentor's http.* attributes and span close are lost."""
+ import litellm.integrations.otel.model.config as otel_config
+ import litellm.proxy.proxy_server as ps
+
+ span = MagicMock()
+ fake_logger = MagicMock()
+ monkeypatch.setattr(ps, "open_telemetry_logger", fake_logger, raising=False)
+ monkeypatch.setattr(otel_config, "is_otel_v2_enabled", lambda: True)
+ request = _make_request(parent_otel_span=span)
+ exc = ProxyException(message="bad", type="bad_request_error", param=None, code=400)
+
+ _close_dangling_otel_server_span(request=request, status_code=422, exc=exc)
+
+ fake_logger.record_error_attributes_on_span.assert_called_once_with(span, exc, 422)
+ assert not span.end.called
+ assert not span.set_status.called
+ assert not fake_logger.set_response_status_code_attribute.called
+ assert request.state.parent_otel_span is span
+
+
+def test_close_dangling_otel_server_span_v2_success_does_not_stamp(monkeypatch):
+ """Under v2 a sub-400 status must not stamp an error onto the SERVER span."""
+ import litellm.integrations.otel.model.config as otel_config
+ import litellm.proxy.proxy_server as ps
+
+ span = MagicMock()
+ fake_logger = MagicMock()
+ monkeypatch.setattr(ps, "open_telemetry_logger", fake_logger, raising=False)
+ monkeypatch.setattr(otel_config, "is_otel_v2_enabled", lambda: True)
+ request = _make_request(parent_otel_span=span)
+
+ _close_dangling_otel_server_span(request=request, status_code=200)
+
+ assert not fake_logger.record_error_attributes_on_span.called
+ assert not span.end.called
+
+
def test_close_dangling_otel_server_span_missing_span_is_noop_error():
"""When parent_otel_span is missing the call short-circuits — no error."""
request = _make_request(parent_otel_span=None)
From 6f4f4f69df2e2369e95235eaa8a6c0e1aea5a6aa Mon Sep 17 00:00:00 2001
From: ryan-crabbe-berri
Date: Sat, 18 Jul 2026 11:24:03 -0700
Subject: [PATCH 029/296] refactor(ui): consolidate Add/Edit credential modals
into one CredentialModal (#32572)
* refactor(ui): consolidate Add/Edit credential modals into one CredentialModal
AddCredentialModal and EditCredentialModal were ~90% identical: the same
provider select, ProviderSpecificFields, and submit/filter logic, differing
only in title, button text, edit-mode prefill, and the disabled credential
name. Replace both with a single CredentialModal driven by a mode: 'add' |
'edit' prop, and point the two call sites in credentials.tsx at it.
Removes ~120 lines of duplication and drops the no-explicit-any and
no-restricted-imports baselines. The two per-file tests merge into one
CredentialModal.test.tsx covering both modes (add: editable empty name;
edit: prefilled, disabled name; provider fields render).
* refactor(ui): derive credential name disabled state from mode, not data
The disabled flag on the credential name field was tied to whether
existingCredential?.credential_name is truthy, an artifact of the old
EditCredentialModal. Drive it from the isEdit flag like the rest of the
component so mode='add' with a stray existingCredential can't disable the
field and mode='edit' with an empty name can't leave it editable. Behavior
is unchanged for real call sites; adds a regression test for the edit-with-
empty-name case.
* refactor(ui): prefill credential form declaratively instead of via useEffect
The edit-mode form was seeded with an imperative form.setFieldsValue inside
a useEffect that also set React state (setSelectedProvider), an antd anti-
pattern carried over from the old EditCredentialModal. Both call sites mount
the modal fresh with existingCredential already present (conditional && plus
destroyOnHidden), so there is no 'prop arrives after mount' case to handle.
Replace it with antd's declarative initialValues on the Form and a lazy
useState initializer for the provider. Removes the effect, its
react-hooks/set-state-in-effect suppression and exhaustive-deps warning, and
one any cast; behavior is unchanged (edit now shows the real provider on
first paint instead of flashing the default). Existing tests cover prefill
and the disabled name field.
---
ui/litellm-dashboard/eslint-suppressions.json | 10 +-
.../model_add/AddCredentialModal.test.tsx | 108 -------------
.../model_add/CredentialModal.test.tsx | 140 ++++++++++++++++
...redentialModal.tsx => CredentialModal.tsx} | 70 ++++----
.../model_add/EditCredentialModal.test.tsx | 123 --------------
.../model_add/EditCredentialModal.tsx | 150 ------------------
.../src/components/model_add/credentials.tsx | 13 +-
7 files changed, 191 insertions(+), 423 deletions(-)
delete mode 100644 ui/litellm-dashboard/src/components/model_add/AddCredentialModal.test.tsx
create mode 100644 ui/litellm-dashboard/src/components/model_add/CredentialModal.test.tsx
rename ui/litellm-dashboard/src/components/model_add/{AddCredentialModal.tsx => CredentialModal.tsx} (71%)
delete mode 100644 ui/litellm-dashboard/src/components/model_add/EditCredentialModal.test.tsx
delete mode 100644 ui/litellm-dashboard/src/components/model_add/EditCredentialModal.tsx
diff --git a/ui/litellm-dashboard/eslint-suppressions.json b/ui/litellm-dashboard/eslint-suppressions.json
index 90b0c84244e..dcf482450e9 100644
--- a/ui/litellm-dashboard/eslint-suppressions.json
+++ b/ui/litellm-dashboard/eslint-suppressions.json
@@ -1885,19 +1885,11 @@
"count": 1
}
},
- "src/components/model_add/AddCredentialModal.tsx": {
+ "src/components/model_add/CredentialModal.tsx": {
"no-restricted-imports": {
"count": 1
}
},
- "src/components/model_add/EditCredentialModal.tsx": {
- "no-restricted-imports": {
- "count": 1
- },
- "react-hooks/set-state-in-effect": {
- "count": 1
- }
- },
"src/components/model_add/credentials.tsx": {
"no-restricted-imports": {
"count": 1
diff --git a/ui/litellm-dashboard/src/components/model_add/AddCredentialModal.test.tsx b/ui/litellm-dashboard/src/components/model_add/AddCredentialModal.test.tsx
deleted file mode 100644
index aee7a0cdd1d..00000000000
--- a/ui/litellm-dashboard/src/components/model_add/AddCredentialModal.test.tsx
+++ /dev/null
@@ -1,108 +0,0 @@
-import { QueryClient, QueryClientProvider } from "@tanstack/react-query";
-import { render, screen, waitFor } from "@testing-library/react";
-import { describe, expect, it, vi } from "vitest";
-import { Providers } from "../provider_info_helpers";
-import AddCredentialModal from "./AddCredentialModal";
-
-vi.mock("../networking", async () => {
- const actual = await vi.importActual("../networking");
- return {
- ...actual,
- getProviderCreateMetadata: vi.fn().mockResolvedValue([
- {
- provider: "OpenAI",
- provider_display_name: Providers.OpenAI,
- litellm_provider: "openai",
- default_model_placeholder: "gpt-3.5-turbo",
- credential_fields: [
- {
- key: "api_key",
- label: "OpenAI API Key",
- field_type: "password",
- required: true,
- },
- {
- key: "api_base",
- label: "API Base",
- field_type: "text",
- placeholder: "https://api.openai.com/v1",
- },
- ],
- },
- {
- provider: "Anthropic",
- provider_display_name: Providers.Anthropic,
- litellm_provider: "anthropic",
- default_model_placeholder: "claude-3-opus-20240229",
- credential_fields: [
- {
- key: "api_key",
- label: "Anthropic API Key",
- field_type: "password",
- required: true,
- },
- ],
- },
- ]),
- };
-});
-
-const createQueryClient = () =>
- new QueryClient({
- defaultOptions: {
- queries: {
- retry: false,
- gcTime: 0,
- },
- },
- });
-
-const mockUploadProps = {
- beforeUpload: vi.fn(),
- onChange: vi.fn(),
-};
-
-describe("AddCredentialModal", () => {
- it("should render", () => {
- const queryClient = createQueryClient();
- const onCancel = vi.fn();
- const onAddCredential = vi.fn();
-
- render(
-
-
- ,
- );
-
- expect(screen.getByText("Add New Credential")).toBeInTheDocument();
- expect(screen.getByLabelText("Credential Name:")).toBeInTheDocument();
- expect(screen.getByLabelText("Provider:")).toBeInTheDocument();
- });
-
- it("should show the correct provider fields", async () => {
- const queryClient = createQueryClient();
- const onCancel = vi.fn();
- const onAddCredential = vi.fn();
-
- render(
-
-
- ,
- );
-
- await waitFor(() => {
- expect(screen.getByLabelText("OpenAI API Key")).toBeInTheDocument();
- expect(screen.getByPlaceholderText("https://api.openai.com/v1")).toBeInTheDocument();
- });
- });
-});
diff --git a/ui/litellm-dashboard/src/components/model_add/CredentialModal.test.tsx b/ui/litellm-dashboard/src/components/model_add/CredentialModal.test.tsx
new file mode 100644
index 00000000000..6804d0cba92
--- /dev/null
+++ b/ui/litellm-dashboard/src/components/model_add/CredentialModal.test.tsx
@@ -0,0 +1,140 @@
+import { QueryClient, QueryClientProvider } from "@tanstack/react-query";
+import { render, screen, waitFor } from "@testing-library/react";
+import { describe, expect, it, vi } from "vitest";
+import { Providers } from "../provider_info_helpers";
+import { CredentialItem } from "../networking";
+import CredentialModal from "./CredentialModal";
+
+vi.mock("../networking", async () => {
+ const actual = await vi.importActual("../networking");
+ return {
+ ...actual,
+ getProviderCreateMetadata: vi.fn().mockResolvedValue([
+ {
+ provider: "OpenAI",
+ provider_display_name: Providers.OpenAI,
+ litellm_provider: "openai",
+ default_model_placeholder: "gpt-3.5-turbo",
+ credential_fields: [
+ {
+ key: "api_key",
+ label: "OpenAI API Key",
+ field_type: "password",
+ required: true,
+ },
+ {
+ key: "api_base",
+ label: "API Base",
+ field_type: "text",
+ placeholder: "https://api.openai.com/v1",
+ },
+ ],
+ },
+ {
+ provider: "Anthropic",
+ provider_display_name: Providers.Anthropic,
+ litellm_provider: "anthropic",
+ default_model_placeholder: "claude-3-opus-20240229",
+ credential_fields: [
+ {
+ key: "api_key",
+ label: "Anthropic API Key",
+ field_type: "password",
+ required: true,
+ },
+ ],
+ },
+ ]),
+ };
+});
+
+const createQueryClient = () =>
+ new QueryClient({
+ defaultOptions: {
+ queries: {
+ retry: false,
+ gcTime: 0,
+ },
+ },
+ });
+
+const mockUploadProps = {
+ beforeUpload: vi.fn(),
+ onChange: vi.fn(),
+};
+
+const mockCredential: CredentialItem = {
+ credential_name: "test-credential",
+ credential_values: {
+ api_key: "test-api-key",
+ api_base: "https://api.test.com",
+ },
+ credential_info: {
+ custom_llm_provider: Providers.OpenAI,
+ },
+};
+
+const renderModal = (props: Partial> = {}) =>
+ render(
+
+
+ ,
+ );
+
+describe("CredentialModal", () => {
+ describe("add mode", () => {
+ it("renders the add title and an editable credential name", () => {
+ renderModal({ mode: "add" });
+
+ expect(screen.getByText("Add New Credential")).toBeInTheDocument();
+ expect(screen.getByText("Add Credential")).toBeInTheDocument();
+ const nameInput = screen.getByLabelText("Credential Name:") as HTMLInputElement;
+ expect(nameInput.value).toBe("");
+ expect(nameInput.disabled).toBe(false);
+ });
+
+ it("shows provider-specific fields for the selected provider", async () => {
+ renderModal({ mode: "add" });
+
+ await waitFor(() => {
+ expect(screen.getByLabelText("OpenAI API Key")).toBeInTheDocument();
+ expect(screen.getByPlaceholderText("https://api.openai.com/v1")).toBeInTheDocument();
+ });
+ });
+ });
+
+ describe("edit mode", () => {
+ it("renders the edit title and update button", () => {
+ renderModal({ mode: "edit", existingCredential: mockCredential });
+
+ expect(screen.getByText("Edit Credential")).toBeInTheDocument();
+ expect(screen.getByText("Update Credential")).toBeInTheDocument();
+ });
+
+ it("prefills the credential name and disables it", async () => {
+ renderModal({ mode: "edit", existingCredential: mockCredential });
+
+ await waitFor(() => {
+ const nameInput = screen.getByLabelText("Credential Name:") as HTMLInputElement;
+ expect(nameInput.value).toBe("test-credential");
+ expect(nameInput.disabled).toBe(true);
+ });
+ });
+
+ it("disables the name from the mode, not the credential's name value", () => {
+ renderModal({
+ mode: "edit",
+ existingCredential: { ...mockCredential, credential_name: "" },
+ });
+
+ expect((screen.getByLabelText("Credential Name:") as HTMLInputElement).disabled).toBe(true);
+ });
+ });
+});
diff --git a/ui/litellm-dashboard/src/components/model_add/AddCredentialModal.tsx b/ui/litellm-dashboard/src/components/model_add/CredentialModal.tsx
similarity index 71%
rename from ui/litellm-dashboard/src/components/model_add/AddCredentialModal.tsx
rename to ui/litellm-dashboard/src/components/model_add/CredentialModal.tsx
index b86a379d3d1..c92a4a90578 100644
--- a/ui/litellm-dashboard/src/components/model_add/AddCredentialModal.tsx
+++ b/ui/litellm-dashboard/src/components/model_add/CredentialModal.tsx
@@ -1,23 +1,47 @@
import { TextInput } from "@tremor/react";
import { Select as AntdSelect, Button, Form, Modal, Tooltip, Typography } from "antd";
import type { UploadProps } from "antd/es/upload";
-import React, { useState } from "react";
+import { useState } from "react";
import ProviderSpecificFields from "../add_model/provider_specific_fields";
+import { CredentialItem } from "../networking";
import { Providers, providerLogoMap } from "../provider_info_helpers";
import { resolveLogoSrc } from "@/lib/assetPaths";
import { resetCredentialFormOnProviderChange } from "./credential_form_helpers";
+
const { Link } = Typography;
-interface AddCredentialsModalProps {
+interface CredentialModalProps {
open: boolean;
onCancel: () => void;
- onAddCredential: (values: any) => void;
+ onSubmit: (values: any) => void;
uploadProps: UploadProps;
+ mode: "add" | "edit";
+ existingCredential?: CredentialItem | null;
}
-const AddCredentialsModal: React.FC = ({ open, onCancel, onAddCredential, uploadProps }) => {
+export default function CredentialModal({
+ open,
+ onCancel,
+ onSubmit,
+ uploadProps,
+ mode,
+ existingCredential = null,
+}: CredentialModalProps) {
+ const isEdit = mode === "edit";
const [form] = Form.useForm();
- const [selectedProvider, setSelectedProvider] = useState(Providers.OpenAI);
+ const [selectedProvider, setSelectedProvider] = useState(
+ (existingCredential?.credential_info.custom_llm_provider as Providers) ?? Providers.OpenAI,
+ );
+
+ const initialValues = existingCredential
+ ? {
+ credential_name: existingCredential.credential_name,
+ custom_llm_provider: existingCredential.credential_info.custom_llm_provider,
+ ...Object.fromEntries(
+ Object.entries(existingCredential.credential_values || {}).map(([key, value]) => [key, value ?? null]),
+ ),
+ }
+ : undefined;
const handleSubmit = (values: any) => {
const filteredValues = Object.entries(values).reduce((acc, [key, value]) => {
@@ -26,32 +50,33 @@ const AddCredentialsModal: React.FC = ({ open, onCance
}
return acc;
}, {} as any);
- onAddCredential(filteredValues);
+ onSubmit(filteredValues);
+ form.resetFields();
+ };
+
+ const closeAndReset = () => {
+ onCancel();
form.resetFields();
};
return (
{
- onCancel();
- form.resetFields();
- }}
+ onCancel={closeAndReset}
footer={null}
width={600}
+ destroyOnHidden={isEdit}
>
-
-
+
- {/* Provider Selection */}
= ({ open, onCance
- {/* Modal Footer */}