Merge pull request #32091 from BerriAI/litellm_backport_191_rc1

chore(release): backport #31912/#31920/#31921 (+#31923/#31929 parity, #31635 prereq) onto patch-1.91.0rc1
This commit is contained in:
yuneng-jiang 2026-07-03 18:03:52 -07:00 committed by GitHub
commit 465a7b22d5
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
41 changed files with 2768 additions and 249 deletions

View file

@ -5010,15 +5010,18 @@ def _bedrock_tools_pt(tools: List, model: Optional[str] = None) -> List[BedrockT
]
"""
from litellm.llms.bedrock.common_utils import (
get_bedrock_base_model,
bedrock_converse_supports_strict_tools,
normalize_json_schema_custom_types_to_object,
)
from litellm.litellm_core_utils.prompt_templates.common_utils import unpack_defs
_valid_json_schema_root_types = frozenset(("array", "boolean", "integer", "null", "number", "object", "string"))
# Only Claude on Bedrock honours strict tool schemas; other families
# (Nova, Llama, GPT-OSS) reject the strict field outright.
supports_strict_tools = bool(model and get_bedrock_base_model(model).startswith("anthropic"))
# (Nova, Llama, GPT-OSS) reject the strict field outright. Opus 4.7/4.8
# also reject `strict` on Bedrock Converse (see #31582) — their validator
# maps toolSpec to the native Anthropic tool shape, which has no strict
# field, even though Anthropic's native API accepts it as a top-level key.
supports_strict_tools = bool(model and bedrock_converse_supports_strict_tools(model))
tool_block_list: List[BedrockToolBlock] = []
for tool_idx, tool in enumerate(tools):
# Check if tool is already a BedrockToolBlock (e.g., systemTool for Nova grounding)

View file

@ -76,6 +76,7 @@ from litellm.utils import (
from ..common_utils import (
BedrockError,
BedrockModelInfo,
bedrock_converse_supports_parallel_tool_use_config,
get_anthropic_beta_from_headers,
get_bedrock_tool_name,
is_claude_4_5_on_bedrock,
@ -1106,18 +1107,28 @@ class AmazonConverseConfig(BaseConfig):
if cache_control is None:
return None
cache_point = CachePointBlock(type="default")
if isinstance(cache_control, dict) and "ttl" in cache_control:
ttl = cache_control["ttl"]
if ttl in ["5m", "1h"] and model is not None:
if is_claude_4_5_on_bedrock(model):
cache_point["ttl"] = ttl
cache_point = self._build_cache_point_block(cache_control, model)
if block_type == "system":
return SystemContentBlock(cachePoint=cache_point)
else:
return ContentBlock(cachePoint=cache_point)
@staticmethod
def _build_cache_point_block(control: Optional[dict], model: Optional[str] = None) -> CachePointBlock:
"""Build a Bedrock ``cachePoint`` block from an OpenAI-style ``cache_control``/``control`` dict.
``type`` is always ``"default"`` (the only value Bedrock's Converse API
accepts). ``ttl`` is only honored for models that support extended TTL
caching (Claude 4.5 family on Bedrock).
"""
cache_point = CachePointBlock(type="default")
if isinstance(control, dict) and "ttl" in control:
ttl = control["ttl"]
if ttl in ["5m", "1h"] and model is not None and is_claude_4_5_on_bedrock(model):
cache_point["ttl"] = ttl
return cache_point
def _transform_system_message(
self, messages: List[AllMessageValues], model: Optional[str] = None
) -> Tuple[List[AllMessageValues], List[SystemContentBlock]]:
@ -1241,7 +1252,7 @@ class AmazonConverseConfig(BaseConfig):
# Handle parallel_tool_calls configuration
parallel_tool_use_config = additional_request_params.pop("_parallel_tool_use_config", None)
if parallel_tool_use_config is not None and is_claude_4_5_on_bedrock(model):
if parallel_tool_use_config is not None and bedrock_converse_supports_parallel_tool_use_config(model):
for key, value in parallel_tool_use_config.items():
if (
key in additional_request_params
@ -1526,7 +1537,8 @@ class AmazonConverseConfig(BaseConfig):
if cache_injection_points and len(bedrock_tools) > 0:
for point in cache_injection_points:
if point.get("location") == "tool_config":
bedrock_tools.append({"cachePoint": {"type": "default"}})
cache_point = self._build_cache_point_block(point.get("control"), model)
bedrock_tools.append(ToolBlock(cachePoint=cache_point))
break
bedrock_tool_config: Optional[ToolConfigBlock] = None

View file

@ -4,9 +4,11 @@ from __future__ import annotations
Common utilities used across bedrock chat/embedding/image generation
"""
import contextlib
import functools
import json
import os
import re
from typing import (
TYPE_CHECKING,
Any,
@ -683,39 +685,72 @@ def get_bedrock_base_model(model: str) -> str:
return model
def bedrock_converse_supports_parallel_tool_use_config(model: str) -> bool:
return any(
(litellm.model_cost.get(candidate) or {}).get("supports_parallel_tool_use_config") is True
for candidate in (model, get_bedrock_base_model(model))
)
def is_claude_4_5_on_bedrock(model: str) -> bool:
"""
Check if the model is a Claude 4.5 model on Bedrock.
Claude 4.5 models support prompt caching with '5m' and '1h' TTL on Bedrock.
Check if the model supports Bedrock prompt caching with an extended '1h' TTL
(in addition to the default 5m TTL).
Backed by the ``cache_creation_input_token_cost_above_1hr`` field in
``model_prices_and_context_window.json`` instead of a hardcoded list of
model-name patterns, so newly released models pick up support as soon as
their pricing entry ships, with no code change required here.
"""
model_lower = model.lower()
claude_4_5_patterns = [
"sonnet-4.5",
"sonnet_4.5",
"sonnet-4-5",
"sonnet_4_5",
"haiku-4.5",
"haiku_4.5",
"haiku-4-5",
"haiku_4_5",
"opus-4.5",
"opus_4.5",
"opus-4-5",
"opus_4_5",
"sonnet-4.6",
"sonnet_4.6",
"sonnet-4-6",
"sonnet_4_6",
"opus-4.6",
"opus_4.6",
"opus-4-6",
"opus_4_6",
"opus-4.7",
"opus_4.7",
"opus-4-7",
"opus_4_7",
]
return any(pattern in model_lower for pattern in claude_4_5_patterns)
return any(
(litellm.model_cost.get(candidate) or {}).get("cache_creation_input_token_cost_above_1hr") is not None
for candidate in (model, get_bedrock_base_model(model))
)
_BEDROCK_MODEL_VERSION_SUFFIX_RE = re.compile(r"-v\d+(?::\d+)?$")
def bedrock_converse_supports_strict_tools(model: str) -> bool:
"""
Whether ``toolSpec.strict`` can be forwarded to Bedrock Converse for ``model``.
Non-Anthropic Bedrock families (Nova, Llama, GPT-OSS) reject the field
outright. Anthropic models forward it unless their entry in
``model_prices_and_context_window.json`` sets
``bedrock_converse_supports_strict_tools: false`` Bedrock routes those
(Opus 4.7/4.8, see #31582) through a stricter validator that rejects the
``strict`` key on ``toolSpec`` even though Anthropic's native API accepts
it as a top-level tool field.
"""
base = get_bedrock_base_model(model)
if not base.startswith("anthropic"):
return False
flag = _get_bedrock_converse_strict_tools_flag(base)
return flag if flag is not None else True
def _get_bedrock_converse_strict_tools_flag(base_model: str) -> Optional[bool]:
candidates = dict.fromkeys((base_model, _BEDROCK_MODEL_VERSION_SUFFIX_RE.sub("", base_model)))
for candidate in candidates:
with contextlib.suppress(Exception):
model_info = get_cached_model_info()(
model=candidate,
custom_llm_provider="bedrock",
)
flag = model_info.get("bedrock_converse_supports_strict_tools")
if isinstance(flag, bool):
return flag
model_cost_key = model_info.get("key")
if isinstance(model_cost_key, str):
local_flag = (
_get_local_model_cost_map().get(model_cost_key, {}).get("bedrock_converse_supports_strict_tools")
)
if isinstance(local_flag, bool):
return local_flag
return None
def normalize_bedrock_opus_output_config_effort(model: str, output_config: Any) -> None:

View file

@ -724,6 +724,7 @@
"supports_tool_choice": true
},
"anthropic.claude-haiku-4-5-20251001-v1:0": {
"supports_parallel_tool_use_config": true,
"cache_creation_input_token_cost": 1.25e-06,
"cache_creation_input_token_cost_above_1hr": 2e-06,
"cache_read_input_token_cost": 1e-07,
@ -747,6 +748,7 @@
"supports_native_structured_output": true
},
"anthropic.claude-haiku-4-5@20251001": {
"supports_parallel_tool_use_config": true,
"cache_creation_input_token_cost": 1.25e-06,
"cache_creation_input_token_cost_above_1hr": 2e-06,
"cache_read_input_token_cost": 1e-07,
@ -787,8 +789,6 @@
"output_cost_per_token_above_200k_tokens": 3e-05,
"cache_creation_input_token_cost_above_200k_tokens": 7.5e-06,
"cache_read_input_token_cost_above_200k_tokens": 6e-07,
"cache_creation_input_token_cost_above_1hr": 7.5e-06,
"cache_creation_input_token_cost_above_1hr_above_200k_tokens": 1.5e-05,
"cache_creation_input_token_cost": 3.75e-06,
"cache_read_input_token_cost": 3e-07
},
@ -813,9 +813,7 @@
"input_cost_per_token_above_200k_tokens": 6e-06,
"output_cost_per_token_above_200k_tokens": 3e-05,
"cache_creation_input_token_cost_above_200k_tokens": 7.5e-06,
"cache_read_input_token_cost_above_200k_tokens": 6e-07,
"cache_creation_input_token_cost_above_1hr": 7.5e-06,
"cache_creation_input_token_cost_above_1hr_above_200k_tokens": 1.5e-05
"cache_read_input_token_cost_above_200k_tokens": 6e-07
},
"anthropic.claude-3-7-sonnet-20240620-v1:0": {
"cache_creation_input_token_cost": 4.5e-06,
@ -965,6 +963,7 @@
"supports_vision": true
},
"anthropic.claude-opus-4-5-20251101-v1:0": {
"supports_parallel_tool_use_config": true,
"cache_creation_input_token_cost": 6.25e-06,
"cache_creation_input_token_cost_above_1hr": 1e-05,
"cache_read_input_token_cost": 5e-07,
@ -994,6 +993,7 @@
"bedrock_output_config_effort_ceiling": "high"
},
"anthropic.claude-opus-4-6-v1": {
"supports_parallel_tool_use_config": true,
"cache_creation_input_token_cost": 6.25e-06,
"cache_creation_input_token_cost_above_1hr": 1e-05,
"cache_read_input_token_cost": 5e-07,
@ -1025,6 +1025,7 @@
"bedrock_output_config_effort_ceiling": "max"
},
"global.anthropic.claude-opus-4-6-v1": {
"supports_parallel_tool_use_config": true,
"cache_creation_input_token_cost": 6.25e-06,
"cache_creation_input_token_cost_above_1hr": 1e-05,
"cache_read_input_token_cost": 5e-07,
@ -1056,6 +1057,7 @@
"bedrock_output_config_effort_ceiling": "max"
},
"us.anthropic.claude-opus-4-6-v1": {
"supports_parallel_tool_use_config": true,
"cache_creation_input_token_cost": 6.875e-06,
"cache_creation_input_token_cost_above_1hr": 1.1e-05,
"cache_read_input_token_cost": 5.5e-07,
@ -1087,6 +1089,7 @@
"bedrock_output_config_effort_ceiling": "max"
},
"eu.anthropic.claude-opus-4-6-v1": {
"supports_parallel_tool_use_config": true,
"cache_creation_input_token_cost": 6.875e-06,
"cache_creation_input_token_cost_above_1hr": 1.1e-05,
"cache_read_input_token_cost": 5.5e-07,
@ -1118,6 +1121,7 @@
"bedrock_output_config_effort_ceiling": "max"
},
"au.anthropic.claude-opus-4-6-v1": {
"supports_parallel_tool_use_config": true,
"cache_creation_input_token_cost": 6.875e-06,
"cache_creation_input_token_cost_above_1hr": 1.1e-05,
"cache_read_input_token_cost": 5.5e-07,
@ -1149,6 +1153,8 @@
"bedrock_output_config_effort_ceiling": "max"
},
"anthropic.claude-opus-4-7": {
"supports_parallel_tool_use_config": true,
"bedrock_converse_supports_strict_tools": false,
"cache_creation_input_token_cost": 6.25e-06,
"cache_creation_input_token_cost_above_1hr": 1e-05,
"cache_read_input_token_cost": 5e-07,
@ -1197,6 +1203,8 @@
"supports_output_config": true
},
"global.anthropic.claude-opus-4-7": {
"supports_parallel_tool_use_config": true,
"bedrock_converse_supports_strict_tools": false,
"cache_creation_input_token_cost": 6.25e-06,
"cache_creation_input_token_cost_above_1hr": 1e-05,
"cache_read_input_token_cost": 5e-07,
@ -1230,6 +1238,8 @@
"bedrock_output_config_effort_ceiling": "xhigh"
},
"us.anthropic.claude-opus-4-7": {
"supports_parallel_tool_use_config": true,
"bedrock_converse_supports_strict_tools": false,
"cache_creation_input_token_cost": 6.875e-06,
"cache_creation_input_token_cost_above_1hr": 1.1e-05,
"cache_read_input_token_cost": 5.5e-07,
@ -1263,6 +1273,8 @@
"bedrock_output_config_effort_ceiling": "xhigh"
},
"eu.anthropic.claude-opus-4-7": {
"supports_parallel_tool_use_config": true,
"bedrock_converse_supports_strict_tools": false,
"cache_creation_input_token_cost": 6.875e-06,
"cache_creation_input_token_cost_above_1hr": 1.1e-05,
"cache_read_input_token_cost": 5.5e-07,
@ -1296,6 +1308,8 @@
"bedrock_output_config_effort_ceiling": "xhigh"
},
"au.anthropic.claude-opus-4-7": {
"supports_parallel_tool_use_config": true,
"bedrock_converse_supports_strict_tools": false,
"cache_creation_input_token_cost": 6.875e-06,
"cache_creation_input_token_cost_above_1hr": 1.1e-05,
"cache_read_input_token_cost": 5.5e-07,
@ -1329,6 +1343,7 @@
"bedrock_output_config_effort_ceiling": "xhigh"
},
"anthropic.claude-fable-5": {
"supports_parallel_tool_use_config": true,
"cache_creation_input_token_cost": 1.25e-05,
"cache_creation_input_token_cost_above_1hr": 2e-05,
"cache_read_input_token_cost": 1e-06,
@ -1362,6 +1377,7 @@
"bedrock_output_config_effort_ceiling": "xhigh"
},
"global.anthropic.claude-fable-5": {
"supports_parallel_tool_use_config": true,
"cache_creation_input_token_cost": 1.25e-05,
"cache_creation_input_token_cost_above_1hr": 2e-05,
"cache_read_input_token_cost": 1e-06,
@ -1395,6 +1411,7 @@
"bedrock_output_config_effort_ceiling": "xhigh"
},
"us.anthropic.claude-fable-5": {
"supports_parallel_tool_use_config": true,
"cache_creation_input_token_cost": 1.375e-05,
"cache_creation_input_token_cost_above_1hr": 2.2e-05,
"cache_read_input_token_cost": 1.1e-06,
@ -1428,6 +1445,7 @@
"bedrock_output_config_effort_ceiling": "xhigh"
},
"eu.anthropic.claude-fable-5": {
"supports_parallel_tool_use_config": true,
"cache_creation_input_token_cost": 1.375e-05,
"cache_creation_input_token_cost_above_1hr": 2.2e-05,
"cache_read_input_token_cost": 1.1e-06,
@ -1461,6 +1479,8 @@
"bedrock_output_config_effort_ceiling": "xhigh"
},
"anthropic.claude-opus-4-8": {
"supports_parallel_tool_use_config": true,
"bedrock_converse_supports_strict_tools": false,
"cache_creation_input_token_cost": 6.25e-06,
"cache_creation_input_token_cost_above_1hr": 1e-05,
"cache_read_input_token_cost": 5e-07,
@ -1494,6 +1514,8 @@
"bedrock_output_config_effort_ceiling": "xhigh"
},
"global.anthropic.claude-opus-4-8": {
"supports_parallel_tool_use_config": true,
"bedrock_converse_supports_strict_tools": false,
"cache_creation_input_token_cost": 6.25e-06,
"cache_creation_input_token_cost_above_1hr": 1e-05,
"cache_read_input_token_cost": 5e-07,
@ -1527,6 +1549,8 @@
"bedrock_output_config_effort_ceiling": "xhigh"
},
"us.anthropic.claude-opus-4-8": {
"supports_parallel_tool_use_config": true,
"bedrock_converse_supports_strict_tools": false,
"cache_creation_input_token_cost": 6.875e-06,
"cache_creation_input_token_cost_above_1hr": 1.1e-05,
"cache_read_input_token_cost": 5.5e-07,
@ -1560,6 +1584,8 @@
"bedrock_output_config_effort_ceiling": "xhigh"
},
"eu.anthropic.claude-opus-4-8": {
"supports_parallel_tool_use_config": true,
"bedrock_converse_supports_strict_tools": false,
"cache_creation_input_token_cost": 6.875e-06,
"cache_creation_input_token_cost_above_1hr": 1.1e-05,
"cache_read_input_token_cost": 5.5e-07,
@ -1593,6 +1619,8 @@
"bedrock_output_config_effort_ceiling": "xhigh"
},
"au.anthropic.claude-opus-4-8": {
"supports_parallel_tool_use_config": true,
"bedrock_converse_supports_strict_tools": false,
"cache_creation_input_token_cost": 6.875e-06,
"cache_creation_input_token_cost_above_1hr": 1.1e-05,
"cache_read_input_token_cost": 5.5e-07,
@ -1626,6 +1654,8 @@
"bedrock_output_config_effort_ceiling": "xhigh"
},
"jp.anthropic.claude-opus-4-7": {
"supports_parallel_tool_use_config": true,
"bedrock_converse_supports_strict_tools": false,
"cache_creation_input_token_cost": 6.875e-06,
"cache_read_input_token_cost": 5.5e-07,
"input_cost_per_token": 5.5e-06,
@ -1657,6 +1687,7 @@
"supports_minimal_reasoning_effort": true
},
"anthropic.claude-sonnet-4-6": {
"supports_parallel_tool_use_config": true,
"cache_creation_input_token_cost": 3.75e-06,
"cache_creation_input_token_cost_above_1hr": 6e-06,
"cache_read_input_token_cost": 3e-07,
@ -1687,6 +1718,7 @@
"supports_output_config": true
},
"global.anthropic.claude-sonnet-4-6": {
"supports_parallel_tool_use_config": true,
"cache_creation_input_token_cost": 3.75e-06,
"cache_creation_input_token_cost_above_1hr": 6e-06,
"cache_read_input_token_cost": 3e-07,
@ -1717,6 +1749,7 @@
"supports_output_config": true
},
"us.anthropic.claude-sonnet-4-6": {
"supports_parallel_tool_use_config": true,
"cache_creation_input_token_cost": 4.125e-06,
"cache_creation_input_token_cost_above_1hr": 6.6e-06,
"cache_read_input_token_cost": 3.3e-07,
@ -1747,6 +1780,7 @@
"supports_output_config": true
},
"eu.anthropic.claude-sonnet-4-6": {
"supports_parallel_tool_use_config": true,
"cache_creation_input_token_cost": 4.125e-06,
"cache_creation_input_token_cost_above_1hr": 6.6e-06,
"cache_read_input_token_cost": 3.3e-07,
@ -1777,6 +1811,7 @@
"supports_output_config": true
},
"au.anthropic.claude-sonnet-4-6": {
"supports_parallel_tool_use_config": true,
"cache_creation_input_token_cost": 4.125e-06,
"cache_creation_input_token_cost_above_1hr": 6.6e-06,
"cache_read_input_token_cost": 3.3e-07,
@ -1807,6 +1842,7 @@
"supports_output_config": true
},
"jp.anthropic.claude-sonnet-4-6": {
"supports_parallel_tool_use_config": true,
"cache_creation_input_token_cost": 4.125e-06,
"cache_creation_input_token_cost_above_1hr": 6.6e-06,
"cache_read_input_token_cost": 3.3e-07,
@ -1866,6 +1902,7 @@
"supports_vision": true
},
"anthropic.claude-sonnet-4-5-20250929-v1:0": {
"supports_parallel_tool_use_config": true,
"cache_creation_input_token_cost": 3.75e-06,
"cache_creation_input_token_cost_above_1hr": 6e-06,
"cache_read_input_token_cost": 3e-07,
@ -2126,6 +2163,7 @@
"cache_creation_input_token_cost": 3.125e-07
},
"apac.anthropic.claude-haiku-4-5-20251001-v1:0": {
"supports_parallel_tool_use_config": true,
"cache_creation_input_token_cost": 1.375e-06,
"cache_read_input_token_cost": 1.1e-07,
"input_cost_per_token": 1.1e-06,
@ -2205,6 +2243,7 @@
"output_cost_per_second": 0.0
},
"au.anthropic.claude-sonnet-4-5-20250929-v1:0": {
"supports_parallel_tool_use_config": true,
"cache_creation_input_token_cost": 4.125e-06,
"cache_creation_input_token_cost_above_1hr": 6.6e-06,
"cache_read_input_token_cost": 3.3e-07,
@ -9327,6 +9366,7 @@
"cache_creation_input_token_cost": 3.75e-07
},
"bedrock/us-gov-east-1/anthropic.claude-sonnet-4-5-20250929-v1:0": {
"supports_parallel_tool_use_config": true,
"cache_creation_input_token_cost": 4.5e-06,
"cache_creation_input_token_cost_above_1hr": 7.2e-06,
"cache_read_input_token_cost": 3.6e-07,
@ -9349,6 +9389,7 @@
"supports_native_structured_output": true
},
"bedrock/us-gov-east-1/claude-sonnet-4-5-20250929-v1:0": {
"supports_parallel_tool_use_config": true,
"cache_creation_input_token_cost": 4.5e-06,
"cache_creation_input_token_cost_above_1hr": 7.2e-06,
"cache_read_input_token_cost": 3.6e-07,
@ -9502,6 +9543,7 @@
"cache_creation_input_token_cost": 3.75e-07
},
"bedrock/us-gov-west-1/anthropic.claude-sonnet-4-5-20250929-v1:0": {
"supports_parallel_tool_use_config": true,
"cache_creation_input_token_cost": 4.5e-06,
"cache_creation_input_token_cost_above_1hr": 7.2e-06,
"cache_read_input_token_cost": 3.6e-07,
@ -9524,6 +9566,7 @@
"supports_native_structured_output": true
},
"bedrock/us-gov-west-1/claude-sonnet-4-5-20250929-v1:0": {
"supports_parallel_tool_use_config": true,
"cache_creation_input_token_cost": 4.5e-06,
"cache_creation_input_token_cost_above_1hr": 7.2e-06,
"cache_read_input_token_cost": 3.6e-07,
@ -10250,6 +10293,7 @@
"supports_output_config": true
},
"claude-sonnet-4-5-20250929-v1:0": {
"supports_parallel_tool_use_config": true,
"cache_creation_input_token_cost": 3.75e-06,
"cache_creation_input_token_cost_above_1hr": 6e-06,
"cache_creation_input_token_cost_above_1hr_above_200k_tokens": 1.2e-05,
@ -14324,6 +14368,7 @@
"cache_creation_input_token_cost": 3.125e-07
},
"eu.anthropic.claude-haiku-4-5-20251001-v1:0": {
"supports_parallel_tool_use_config": true,
"cache_creation_input_token_cost": 1.375e-06,
"cache_creation_input_token_cost_above_1hr": 2.2e-06,
"cache_read_input_token_cost": 1.1e-07,
@ -14529,6 +14574,7 @@
"supports_vision": true
},
"eu.anthropic.claude-sonnet-4-5-20250929-v1:0": {
"supports_parallel_tool_use_config": true,
"cache_creation_input_token_cost": 4.125e-06,
"cache_creation_input_token_cost_above_1hr": 6.6e-06,
"cache_read_input_token_cost": 3.3e-07,
@ -19725,6 +19771,7 @@
"mode": "search"
},
"global.anthropic.claude-sonnet-4-5-20250929-v1:0": {
"supports_parallel_tool_use_config": true,
"cache_creation_input_token_cost": 3.75e-06,
"cache_creation_input_token_cost_above_1hr": 6e-06,
"cache_read_input_token_cost": 3e-07,
@ -19786,6 +19833,7 @@
"supports_vision": true
},
"global.anthropic.claude-haiku-4-5-20251001-v1:0": {
"supports_parallel_tool_use_config": true,
"cache_creation_input_token_cost": 1.25e-06,
"cache_creation_input_token_cost_above_1hr": 2e-06,
"cache_read_input_token_cost": 1e-07,
@ -23846,6 +23894,7 @@
"output_cost_per_token": 1.8e-08
},
"jp.anthropic.claude-sonnet-4-5-20250929-v1:0": {
"supports_parallel_tool_use_config": true,
"cache_creation_input_token_cost": 4.125e-06,
"cache_creation_input_token_cost_above_1hr": 6.6e-06,
"cache_read_input_token_cost": 3.3e-07,
@ -23878,6 +23927,7 @@
"supports_native_structured_output": true
},
"jp.anthropic.claude-haiku-4-5-20251001-v1:0": {
"supports_parallel_tool_use_config": true,
"cache_creation_input_token_cost": 1.375e-06,
"cache_creation_input_token_cost_above_1hr": 2.2e-06,
"cache_read_input_token_cost": 1.1e-07,
@ -32494,6 +32544,7 @@
"supports_tool_choice": true
},
"us.anthropic.claude-haiku-4-5-20251001-v1:0": {
"supports_parallel_tool_use_config": true,
"cache_creation_input_token_cost": 1.375e-06,
"cache_creation_input_token_cost_above_1hr": 2.2e-06,
"cache_read_input_token_cost": 1.1e-07,
@ -32644,6 +32695,7 @@
"supports_vision": true
},
"us.anthropic.claude-sonnet-4-5-20250929-v1:0": {
"supports_parallel_tool_use_config": true,
"cache_creation_input_token_cost": 4.125e-06,
"cache_creation_input_token_cost_above_1hr": 6.6e-06,
"cache_read_input_token_cost": 3.3e-07,
@ -32676,6 +32728,7 @@
"supports_native_structured_output": true
},
"us-gov.anthropic.claude-sonnet-4-5-20250929-v1:0": {
"supports_parallel_tool_use_config": true,
"cache_creation_input_token_cost": 4.5e-06,
"cache_creation_input_token_cost_above_1hr": 7.2e-06,
"cache_read_input_token_cost": 3.6e-07,
@ -32703,6 +32756,7 @@
"supports_native_structured_output": true
},
"au.anthropic.claude-haiku-4-5-20251001-v1:0": {
"supports_parallel_tool_use_config": true,
"cache_creation_input_token_cost": 1.375e-06,
"cache_creation_input_token_cost_above_1hr": 2.2e-06,
"cache_read_input_token_cost": 1.1e-07,
@ -32750,6 +32804,7 @@
"supports_vision": true
},
"us.anthropic.claude-opus-4-5-20251101-v1:0": {
"supports_parallel_tool_use_config": true,
"cache_creation_input_token_cost": 6.875e-06,
"cache_creation_input_token_cost_above_1hr": 1.1e-05,
"cache_read_input_token_cost": 5.5e-07,
@ -32779,6 +32834,7 @@
"bedrock_output_config_effort_ceiling": "high"
},
"global.anthropic.claude-opus-4-5-20251101-v1:0": {
"supports_parallel_tool_use_config": true,
"cache_creation_input_token_cost": 6.25e-06,
"cache_creation_input_token_cost_above_1hr": 1e-05,
"cache_read_input_token_cost": 5e-07,
@ -32808,6 +32864,7 @@
"bedrock_output_config_effort_ceiling": "high"
},
"eu.anthropic.claude-opus-4-5-20251101-v1:0": {
"supports_parallel_tool_use_config": true,
"cache_creation_input_token_cost": 6.25e-06,
"cache_read_input_token_cost": 5e-07,
"input_cost_per_token": 5e-06,
@ -42759,6 +42816,7 @@
"source": "https://aws.amazon.com/bedrock/pricing/"
},
"bedrock/us-gov-east-1/anthropic.claude-haiku-4-5-20251001-v1:0": {
"supports_parallel_tool_use_config": true,
"cache_creation_input_token_cost": 1.5e-06,
"cache_creation_input_token_cost_above_1hr": 2.4e-06,
"cache_read_input_token_cost": 1.2e-07,
@ -42782,6 +42840,7 @@
"supports_pdf_input": true
},
"bedrock/us-gov-west-1/anthropic.claude-haiku-4-5-20251001-v1:0": {
"supports_parallel_tool_use_config": true,
"cache_creation_input_token_cost": 1.5e-06,
"cache_creation_input_token_cost_above_1hr": 2.4e-06,
"cache_read_input_token_cost": 1.2e-07,

View file

@ -0,0 +1,78 @@
"""Client authentication for OAuth 2.0 token-endpoint requests (RFC 6749 section 2.3.1).
A confidential MCP upstream may require ``client_secret_basic`` (HTTP Basic, the OIDC
default) or ``client_secret_post`` (credentials in the form body). Every token-endpoint
POST in the MCP gateway builds its client authentication here so the two methods are
applied identically across the inbound exchange, the refresh grants, the M2M
client_credentials fetch, and RFC 8693 token exchange. The default is
``client_secret_post`` so servers that never set ``token_endpoint_auth_method`` keep
their current behavior.
"""
from __future__ import annotations
import base64
from dataclasses import dataclass
from urllib.parse import quote_plus
from litellm.types.mcp_server.mcp_server_manager import MCPTokenEndpointAuthMethod
@dataclass(frozen=True, slots=True)
class TokenEndpointClientAuth:
headers: dict[str, str]
body: dict[str, str]
class TokenEndpointAuthConfigError(ValueError):
"""``client_secret_basic`` is configured but the client credentials needed for it are missing.
Subclasses ``ValueError`` so existing call sites that already guard missing credentials with
``except ValueError`` / ``except Exception`` keep mapping it to their own failure contract.
"""
def normalize_token_endpoint_auth_method(
value: object,
) -> MCPTokenEndpointAuthMethod | None:
"""Narrow an untyped (DB/JSON-sourced) value to the auth-method literal, else ``None``."""
if value == "client_secret_basic":
return "client_secret_basic"
if value == "client_secret_post":
return "client_secret_post"
return None
def build_token_endpoint_client_auth(
*,
auth_method: MCPTokenEndpointAuthMethod | None,
client_id: str | None,
client_secret: str | None,
) -> TokenEndpointClientAuth:
"""Return the headers and body fields that authenticate the client to the token endpoint.
``client_secret_basic`` is a confidential-client method, so it requires both ``client_id`` and
``client_secret`` and raises ``TokenEndpointAuthConfigError`` when either is missing rather than
silently degrading to a weaker request (RFC 6749 section 2.3.1; matches the "absent credential
must surface, never fall sideways" rule). It sends an HTTP Basic ``Authorization`` header and
keeps the credentials out of the body. Any other method (including ``None``, the default) is the
``client_secret_post`` path: it places whichever of ``client_id`` / ``client_secret`` are present
into the body, so a secretless client_id (a public client authenticating with PKCE) stays valid.
"""
if auth_method == "client_secret_basic":
if not client_id or not client_secret:
raise TokenEndpointAuthConfigError(
"token_endpoint_auth_method=client_secret_basic requires both client_id and client_secret"
)
# RFC 6749 section 2.3.1: form-urlencode each value before joining with ':' so a
# client_id/secret containing reserved characters (':', '+', '%', ...) is transmitted intact.
userpass = f"{quote_plus(client_id)}:{quote_plus(client_secret)}"
encoded = base64.b64encode(userpass.encode()).decode()
return TokenEndpointClientAuth(headers={"Authorization": f"Basic {encoded}"}, body={})
return TokenEndpointClientAuth(
headers={},
body={
**({"client_id": client_id} if client_id else {}),
**({"client_secret": client_secret} if client_secret else {}),
},
)

View file

@ -24,6 +24,9 @@ from litellm.constants import (
MCP_TOKEN_EXCHANGE_CACHE_MAX_SIZE,
)
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
from litellm.proxy._experimental.mcp_server.auth.token_endpoint_auth import (
build_token_endpoint_client_auth,
)
from litellm.types.llms.custom_http import httpxSpecialProvider
if TYPE_CHECKING:
@ -113,12 +116,16 @@ class TokenExchangeHandler:
f"but missing client_id or client_secret"
)
client_auth = build_token_endpoint_client_auth(
auth_method=server.token_endpoint_auth_method,
client_id=server.client_id,
client_secret=server.client_secret,
)
data: Dict[str, str] = {
"grant_type": TOKEN_EXCHANGE_GRANT_TYPE,
"subject_token": subject_token,
"subject_token_type": server.subject_token_type or DEFAULT_SUBJECT_TOKEN_TYPE,
"client_id": server.client_id,
"client_secret": server.client_secret,
**client_auth.body,
}
if server.audience:
data["audience"] = server.audience
@ -133,8 +140,9 @@ class TokenExchangeHandler:
)
client = get_async_httpx_client(llm_provider=httpxSpecialProvider.MCP)
post_kwargs = {"data": data, **({"headers": client_auth.headers} if client_auth.headers else {})}
try:
response = await client.post(endpoint, data=data)
response = await client.post(endpoint, **post_kwargs)
response.raise_for_status()
except httpx.HTTPStatusError as exc:
verbose_logger.debug(

View file

@ -9,6 +9,10 @@ from litellm._logging import verbose_proxy_logger
from litellm._uuid import uuid
from litellm.constants import MCP_PER_USER_TOKEN_EXPIRY_BUFFER_SECONDS
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
from litellm.proxy._experimental.mcp_server.auth.token_endpoint_auth import (
build_token_endpoint_client_auth,
normalize_token_endpoint_auth_method,
)
from litellm.proxy._types import (
LiteLLM_MCPServerTable,
LiteLLM_ObjectPermissionTable,
@ -1030,20 +1034,21 @@ async def refresh_user_oauth_token(
)
return None
token_data: Dict[str, str] = {
"grant_type": "refresh_token",
"refresh_token": refresh_token,
}
if client_id:
token_data["client_id"] = client_id
if client_secret:
token_data["client_secret"] = client_secret
try:
client_auth = build_token_endpoint_client_auth(
auth_method=normalize_token_endpoint_auth_method(getattr(server, "token_endpoint_auth_method", None)),
client_id=client_id,
client_secret=client_secret,
)
token_data: Dict[str, str] = {
"grant_type": "refresh_token",
"refresh_token": refresh_token,
**client_auth.body,
}
async_client = get_async_httpx_client(llm_provider=httpxSpecialProvider.Oauth2Check)
response = await async_client.post(
token_url,
headers={"Accept": "application/json"},
headers={"Accept": "application/json", **client_auth.headers},
data=token_data,
)
response.raise_for_status()

View file

@ -2,18 +2,23 @@ import asyncio
import html as _html
import json
import time
from typing import Any, Dict, Optional, Tuple
from typing import TYPE_CHECKING, Any, Dict, Literal, Optional, Tuple
from urllib.parse import parse_qsl, urlencode, urlparse, urlunparse
import httpx
from fastapi import APIRouter, Form, HTTPException, Request
from fastapi.responses import HTMLResponse, JSONResponse, RedirectResponse
from pydantic import BaseModel, ValidationError
from litellm._logging import verbose_logger
from litellm.llms.custom_httpx.http_handler import (
get_async_httpx_client,
httpxSpecialProvider,
)
from litellm.proxy._experimental.mcp_server.auth.token_endpoint_auth import (
TokenEndpointAuthConfigError,
build_token_endpoint_client_auth,
)
from litellm.proxy._experimental.mcp_server.oauth_utils import (
TOKEN_NO_CACHE_HEADERS,
get_request_base_url,
@ -26,9 +31,12 @@ from litellm.proxy.common_utils.encrypt_decrypt_utils import (
)
from litellm.proxy.common_utils.http_parsing_utils import _read_request_body
from litellm.proxy.utils import get_server_root_path
from litellm.types.mcp import MCPAuth
from litellm.types.mcp import MCPAuth, MCPCredentials
from litellm.types.mcp_server.mcp_server_manager import MCPServer
if TYPE_CHECKING:
from litellm.proxy._types import LiteLLM_MCPServerTable
# TTL cache for upstream OAuth metadata fetched from pass-through MCP servers.
# Keeps us from hammering the upstream IdP on each discovery request.
# Keyed by (server_id, resource_url) → (expires_at_epoch, payload).
@ -398,6 +406,14 @@ async def exchange_token_with_server(
resolved_client_id = mcp_server.client_id if mcp_server.client_id else client_id
resolved_client_secret = mcp_server.client_secret if mcp_server.client_secret else client_secret
try:
client_auth = build_token_endpoint_client_auth(
auth_method=mcp_server.token_endpoint_auth_method,
client_id=resolved_client_id,
client_secret=resolved_client_secret,
)
except TokenEndpointAuthConfigError as exc:
raise HTTPException(status_code=400, detail=str(exc)) from exc
if grant_type == "refresh_token":
if not refresh_token:
@ -408,10 +424,8 @@ async def exchange_token_with_server(
token_data: dict = {
"grant_type": "refresh_token",
"refresh_token": refresh_token,
"client_id": resolved_client_id,
**client_auth.body,
}
if resolved_client_secret is not None:
token_data["client_secret"] = resolved_client_secret
if scope:
token_data["scope"] = scope
else:
@ -423,19 +437,17 @@ async def exchange_token_with_server(
proxy_base_url = get_request_base_url(request)
token_data = {
"grant_type": "authorization_code",
"client_id": resolved_client_id,
"code": code,
"redirect_uri": f"{proxy_base_url}/callback",
**client_auth.body,
}
if resolved_client_secret is not None:
token_data["client_secret"] = resolved_client_secret
if code_verifier:
token_data["code_verifier"] = code_verifier
async_client = get_async_httpx_client(llm_provider=httpxSpecialProvider.Oauth2Check)
response = await async_client.post(
mcp_server.token_url,
headers={"Accept": "application/json"},
headers={"Accept": "application/json", **client_auth.headers},
data=token_data,
)
if response is None:
@ -501,6 +513,178 @@ async def exchange_token_with_server(
return JSONResponse(result, headers=TOKEN_NO_CACHE_HEADERS)
class _DcrClientRegistration(BaseModel):
"""RFC 7591 dynamic client registration response, narrowed to the fields the gateway
must persist to authenticate later token-endpoint calls. Extra members are ignored."""
client_id: str
client_secret: Optional[str] = None
token_endpoint_auth_method: Optional[str] = None
class _PersistedDcrCredentials(BaseModel):
client_id: Optional[str] = None
client_secret: Optional[str] = None
token_endpoint_auth_method: Optional[str] = None
def _get_persisted_dcr_credentials(credentials: object) -> Optional[_PersistedDcrCredentials]:
if not credentials:
return None
try:
return (
_PersistedDcrCredentials.model_validate_json(credentials)
if isinstance(credentials, str)
else _PersistedDcrCredentials.model_validate(credentials)
)
except ValidationError:
return None
def _decrypt_persisted_dcr_credential(value: Optional[str], key: str) -> Optional[str]:
if value is None:
return None
return decrypt_value_helper(
value=value,
key=key,
exception_type="debug",
return_original_value=True,
)
def _apply_persisted_dcr_credentials(mcp_server: MCPServer, credentials: _PersistedDcrCredentials) -> bool:
client_id = _decrypt_persisted_dcr_credential(credentials.client_id, "client_id")
if not client_id:
return False
mcp_server.client_id = client_id
mcp_server.client_secret = _decrypt_persisted_dcr_credential(credentials.client_secret, "client_secret")
mcp_server.token_endpoint_auth_method = credentials.token_endpoint_auth_method
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
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,
)
except Exception as exc: # noqa: BLE001
verbose_logger.debug(
"register_client_with_server: failed to read persisted DCR client registration 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)
if credentials is None or not credentials.client_id:
return None
return persisted_mcp_server, credentials
async def _reuse_persisted_dcr_client_if_available(mcp_server: MCPServer) -> bool:
persisted = await _get_persisted_mcp_server_with_dcr_client_id(mcp_server)
if persisted is None:
return False
persisted_mcp_server, credentials = persisted
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,
)
return bool(mcp_server.client_id)
DcrRegistrationPersistenceResult = Literal["persisted", "reused", "failed"]
async def _persist_dcr_client_registration(
mcp_server: MCPServer, registration_response: object
) -> DcrRegistrationPersistenceResult:
"""Persist the dynamically registered OAuth client (RFC 7591) onto the MCP server row.
The interactive authorization_code flow mints a ``client_id`` via Dynamic Client
Registration that discovery cannot re-derive; without persisting it the autonomous
``refresh_token`` grant has no client identity, so an expired access token forces a
full re-authorization instead of a silent refresh. Mirrors the ``encrypt_credentials``
write that ``client_credentials`` and token exchange already use. Failures are logged,
never raised: registration still returns to the caller even when persistence fails.
"""
try:
registration = _DcrClientRegistration.model_validate(registration_response)
except ValidationError as exc:
verbose_logger.warning(
"register_client_with_server: DCR response has no usable client_id for server_id=%s; "
"client registration not persisted (%s)",
mcp_server.server_id,
exc,
)
return "failed"
if await _reuse_persisted_dcr_client_if_available(mcp_server):
return "reused"
credentials: MCPCredentials = {
"client_id": registration.client_id,
**({"client_secret": registration.client_secret} if registration.client_secret is not None else {}),
**(
{"token_endpoint_auth_method": "client_secret_basic"}
if registration.token_endpoint_auth_method == "client_secret_basic"
else {}
),
}
from litellm.proxy._experimental.mcp_server.db import update_mcp_server # noqa: PLC0415
from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( # noqa: PLC0415
global_mcp_server_manager,
)
from litellm.proxy._types import UpdateMCPServerRequest # noqa: PLC0415
from litellm.proxy.utils import get_prisma_client_or_throw # noqa: PLC0415
try:
prisma_client = get_prisma_client_or_throw(
"Database not connected. Cannot persist MCP OAuth client registration."
)
updated_row = await update_mcp_server(
prisma_client=prisma_client,
data=UpdateMCPServerRequest(
server_id=mcp_server.server_id,
credentials=credentials,
**({"token_url": mcp_server.token_url} if mcp_server.token_url else {}),
),
touched_by="mcp_oauth_dcr",
)
await global_mcp_server_manager.update_server(updated_row)
return "persisted"
except Exception as exc: # noqa: BLE001
verbose_logger.warning(
"register_client_with_server: failed to persist DCR client registration for server_id=%s: %s",
mcp_server.server_id,
exc,
)
return "failed"
async def register_client_with_server(
request: Request,
mcp_server: MCPServer,
@ -509,6 +693,7 @@ async def register_client_with_server(
response_types: Optional[list],
token_endpoint_auth_method: Optional[str],
fallback_client_id: Optional[str] = None,
persist_credentials: bool = False,
):
request_base_url = get_request_base_url(request)
dummy_return = {
@ -517,7 +702,10 @@ async def register_client_with_server(
"redirect_uris": [f"{request_base_url}/callback"],
}
if mcp_server.client_id and mcp_server.client_secret:
if mcp_server.client_id:
return dummy_return
if await _reuse_persisted_dcr_client_if_available(mcp_server):
return dummy_return
if mcp_server.authorization_url is None:
@ -553,6 +741,11 @@ async def register_client_with_server(
token_response = response.json()
if persist_credentials:
persistence_result = await _persist_dcr_client_registration(mcp_server, token_response)
if persistence_result == "reused":
return dummy_return
return JSONResponse(token_response)

View file

@ -754,6 +754,7 @@ class MCPServerManager:
authorization_url=resolved_authorization_url,
token_url=resolved_token_url,
registration_url=resolved_registration_url,
token_endpoint_auth_method=server_config.get("token_endpoint_auth_method", None),
# TODO: utility fn the default values
transport=server_config.get("transport", MCPTransport.http),
auth_type=auth_type,
@ -1127,6 +1128,9 @@ class MCPServerManager:
authorization_url=mcp_server.authorization_url or getattr(mcp_oauth_metadata, "authorization_url", None),
token_url=mcp_server.token_url or getattr(mcp_oauth_metadata, "token_url", None),
registration_url=mcp_server.registration_url or getattr(mcp_oauth_metadata, "registration_url", None),
token_endpoint_auth_method=(
credentials_dict.get("token_endpoint_auth_method") if credentials_dict else None
),
command=getattr(mcp_server, "command", None),
args=getattr(mcp_server, "args", None) or [],
env=env_dict,
@ -2051,7 +2055,7 @@ class MCPServerManager:
]
return tools
else:
tools = await self._fetch_tools_with_timeout(client, server.name, server=server)
tools = await self._fetch_tools_with_timeout(client, server.name)
self._remember_upstream_initialize_instructions(server, client)
prefixed_or_original_tools = self._create_prefixed_tools(tools, server, add_prefix=add_prefix)
@ -2063,6 +2067,17 @@ class MCPServerManager:
# client triggers the upstream OAuth flow. The multi-server
# aggregator catches this explicitly to keep absorbing.
raise
except HTTPException as e:
headers = e.headers or {}
www_authenticate = headers.get("WWW-Authenticate") or headers.get("www-authenticate")
if e.status_code == 401 and www_authenticate is not None:
raise MCPUpstreamAuthError(
status_code=401,
www_authenticate=www_authenticate,
server_name=server.name,
) from e
verbose_logger.warning(f"Failed to get tools from server {server.name}: {str(e)}")
return []
except Exception as e:
verbose_logger.warning(f"Failed to get tools from server {server.name}: {str(e)}")
return []
@ -2613,7 +2628,6 @@ class MCPServerManager:
self,
client: MCPClient,
server_name: str,
server: Optional[MCPServer] = None,
) -> List[MCPTool]:
"""
Fetch tools from MCP client with timeout and error handling.
@ -2621,38 +2635,27 @@ class MCPServerManager:
Uses anyio.fail_after() instead of asyncio.wait_for() to avoid conflicts
with the MCP SDK's anyio TaskGroup. See GitHub issue #20715 for details.
For OAuth pass-through and upstream-delegated OAuth2 MCP servers, an
upstream HTTP 401 is converted into :class:`MCPUpstreamAuthError`
instead of being swallowed to an empty tool list. That lets the
single-server HTTP routes surface a proper 401 + ``WWW-Authenticate``
challenge so standards-compliant MCP clients trigger the upstream
OAuth flow. Other servers keep today's swallow-and-log behaviour so
the multi-server ``/mcp`` aggregator doesn't get tainted by a single
bad server.
An upstream HTTP 401 is converted into :class:`MCPUpstreamAuthError`
instead of being swallowed to an empty tool list, regardless of the
server's auth_type. Callers route it by surface: the single-server HTTP
routes turn it into a 401 + ``WWW-Authenticate`` challenge so standards-
compliant MCP clients trigger the upstream OAuth flow, while the
multi-server ``/mcp`` aggregator absorbs it to an empty list so one
unauthenticated server doesn't fail the whole listing. Only a 401
(missing/invalid credential) drives the re-auth challenge; a 403
(authenticated but forbidden, e.g. insufficient scope) is not a re-auth
signal and, like other non-auth errors, returns an empty list.
Args:
client: MCP client instance
server_name: Name of the server for logging
server: Optional MCPServer; when upstream auth is delegated, auth
errors are re-raised as :class:`MCPUpstreamAuthError`.
Returns:
List of tools from the server
"""
should_surface_upstream_auth = bool(
server is not None
and (
server.is_oauth_passthrough
or (
server.auth_type == MCPAuth.oauth2
and getattr(server, "delegate_auth_to_upstream", False) is True
and not server.has_client_credentials
)
)
)
try:
with anyio.fail_after(MCP_TOOL_LISTING_TIMEOUT):
tools = await client.list_tools(raise_on_error=should_surface_upstream_auth)
tools = await client.list_tools(raise_on_error=True)
verbose_logger.debug(f"Tools from {server_name}: {tools}")
return tools
except TimeoutError:
@ -2665,16 +2668,15 @@ class MCPServerManager:
verbose_logger.warning(f"Connection error while listing tools from {server_name}: {str(e)}")
return []
except Exception as e:
if should_surface_upstream_auth:
auth_info = _extract_upstream_auth_failure(e)
if auth_info is not None:
status_code, www_authenticate = auth_info
verbose_logger.info(f"Upstream auth failure from MCP server {server_name}: HTTP {status_code}")
raise MCPUpstreamAuthError(
status_code=status_code,
www_authenticate=www_authenticate,
server_name=server_name,
) from e
auth_info = _extract_upstream_auth_failure(e)
if auth_info is not None and auth_info[0] == 401:
_, www_authenticate = auth_info
verbose_logger.info(f"Upstream auth failure from MCP server {server_name}: HTTP 401")
raise MCPUpstreamAuthError(
status_code=401,
www_authenticate=www_authenticate,
server_name=server_name,
) from e
verbose_logger.warning(f"Error listing tools from {server_name}: {str(e)}")
return []

View file

@ -27,6 +27,9 @@ from litellm.proxy.common_utils.encrypt_decrypt_utils import (
encrypt_value_helper,
)
from litellm.proxy._experimental.mcp_server.auth import token_exchange
from litellm.proxy._experimental.mcp_server.auth.token_endpoint_auth import (
build_token_endpoint_client_auth,
)
from litellm.types.llms.custom_http import httpxSpecialProvider
if TYPE_CHECKING:
@ -103,10 +106,14 @@ class MCPOAuth2TokenCache(InMemoryCache):
f"token_url={bool(server.token_url)}"
)
client_auth = build_token_endpoint_client_auth(
auth_method=server.token_endpoint_auth_method,
client_id=server.client_id,
client_secret=server.client_secret,
)
data: Dict[str, str] = {
"grant_type": "client_credentials",
"client_id": server.client_id,
"client_secret": server.client_secret,
**client_auth.body,
}
if server.scopes:
data["scope"] = " ".join(server.scopes)
@ -116,8 +123,9 @@ class MCPOAuth2TokenCache(InMemoryCache):
server.server_id,
)
post_kwargs = {"data": data, **({"headers": client_auth.headers} if client_auth.headers else {})}
try:
response = await client.post(server.token_url, data=data)
response = await client.post(server.token_url, **post_kwargs)
response.raise_for_status()
except httpx.HTTPStatusError as exc:
raise ValueError(

View file

@ -14,6 +14,11 @@ import time
from collections.abc import Awaitable, Callable
from typing import TYPE_CHECKING, Protocol
from litellm._logging import verbose_logger
from litellm.proxy._experimental.mcp_server.auth.token_endpoint_auth import (
TokenEndpointAuthConfigError,
build_token_endpoint_client_auth,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.oauth_token_store import (
OAuthToken,
)
@ -22,7 +27,7 @@ if TYPE_CHECKING:
from litellm.types.mcp_server.mcp_server_manager import MCPServer
ServerLookup = Callable[[str], "MCPServer | None"]
TokenEndpointPost = Callable[[str, dict[str, str]], Awaitable["dict[str, object] | None"]]
TokenEndpointPost = Callable[[str, dict[str, str], dict[str, str]], Awaitable["dict[str, object] | None"]]
class CredentialPersist(Protocol):
@ -86,13 +91,21 @@ class AuthorizationCodeRefresher:
if server is None or not server.token_url:
return None
try:
client_auth = build_token_endpoint_client_auth(
auth_method=server.token_endpoint_auth_method,
client_id=server.client_id,
client_secret=server.client_secret,
)
except TokenEndpointAuthConfigError as exc:
verbose_logger.warning("MCP OAuth refresh misconfigured for server %s: %s", server_id, exc)
return None
form = {
"grant_type": "refresh_token",
"refresh_token": token.refresh_token,
**({"client_id": server.client_id} if server.client_id else {}),
**({"client_secret": server.client_secret} if server.client_secret else {}),
**client_auth.body,
}
body = await self._token_endpoint(server.token_url, form)
body = await self._token_endpoint(server.token_url, form, client_auth.headers)
if body is None:
return None
access_token = body.get("access_token")

View file

@ -92,7 +92,7 @@ async def _persist_credential(
)
async def _post_token_endpoint(url: str, form: dict[str, str]) -> dict[str, object] | None:
async def _post_token_endpoint(url: str, form: dict[str, str], headers: dict[str, str]) -> dict[str, object] | None:
from litellm.llms.custom_httpx.http_handler import ( # noqa: PLC0415
get_async_httpx_client, # pyright: ignore
)
@ -101,11 +101,11 @@ async def _post_token_endpoint(url: str, form: dict[str, str]) -> dict[str, obje
# litellm's httpx handler and httpx.Response are only partially typed; the IdP returns a JSON
# object and the refresher validates each field, so the untyped boundary is contained here.
provider = httpxSpecialProvider.Oauth2Check
headers = {"Accept": "application/json"}
request_headers = {"Accept": "application/json", **headers}
# A failed refresh is a miss, not a 500 (matches v1), so any error becomes None.
try:
client = get_async_httpx_client(llm_provider=provider) # pyright: ignore
response = await client.post(url, headers=headers, data=form) # pyright: ignore
response = await client.post(url, headers=request_headers, data=form) # pyright: ignore
response.raise_for_status() # pyright: ignore
body: dict[str, object] = response.json() # pyright: ignore
except Exception as exc: # noqa: BLE001

View file

@ -1714,6 +1714,7 @@ if MCP_AVAILABLE:
response_types=data.get("response_types", []),
token_endpoint_auth_method=data.get("token_endpoint_auth_method", ""),
fallback_client_id=server_id,
persist_credentials=_user_is_full_admin(user_api_key_dict),
)
@router.delete(

View file

@ -18,6 +18,7 @@ class CacheControlToolConfigInjectionPoint(TypedDict):
"""Type for tool_config-level injection points (Bedrock)."""
location: Literal["tool_config"]
control: Optional[ChatCompletionCachedContent]
CacheControlInjectionPoint = Union[

View file

@ -73,6 +73,10 @@ class MCPPublicServer(BaseModel):
mcp_info: Optional[Dict[str, Any]] = None
# OAuth 2.0 token-endpoint client authentication method (RFC 6749 section 2.3.1).
MCPTokenEndpointAuthMethod = Literal["client_secret_basic", "client_secret_post"]
class MCPCredentials(TypedDict, total=False):
auth_value: Optional[str]
"""
@ -132,6 +136,12 @@ class MCPCredentials(TypedDict, total=False):
Default: urn:ietf:params:oauth:token-type:access_token
"""
token_endpoint_auth_method: Optional[MCPTokenEndpointAuthMethod]
"""
How the gateway authenticates to the upstream token endpoint. "client_secret_basic"
sends HTTP Basic; defaults to "client_secret_post" when unset.
"""
class MCPServerCostInfo(TypedDict, total=False):
default_cost_per_query: Optional[float]

View file

@ -3,7 +3,12 @@ from typing import Any, Dict, List, Literal, Optional
from pydantic import BaseModel, ConfigDict
from litellm.types.mcp import MCPAuth, MCPAuthType, MCPTransportType
from litellm.types.mcp import (
MCPAuth,
MCPAuthType,
MCPTokenEndpointAuthMethod,
MCPTransportType,
)
# MCPInfo now allows arbitrary additional fields for custom metadata
MCPInfo = Dict[str, Any]
@ -48,6 +53,10 @@ class MCPServer(BaseModel):
authorization_url: Optional[str] = None
token_url: Optional[str] = None
registration_url: Optional[str] = None
# How the gateway authenticates to the upstream token endpoint. When
# "client_secret_basic" the credentials go in an HTTP Basic Authorization
# header (omitted from the body); None defaults to "client_secret_post".
token_endpoint_auth_method: Optional[MCPTokenEndpointAuthMethod] = None
# AWS SigV4 fields
aws_access_key_id: Optional[str] = None
aws_secret_access_key: Optional[str] = None

View file

@ -152,6 +152,7 @@ class ProviderSpecificModelInfo(TypedDict, total=False):
supports_output_config: Optional[bool]
supports_image_size: Optional[bool]
bedrock_output_config_effort_ceiling: Optional[Literal["low", "medium", "high", "max", "xhigh"]]
bedrock_converse_supports_strict_tools: Optional[bool]
class SearchContextCostPerQuery(TypedDict, total=False):

View file

@ -5418,6 +5418,7 @@ def _get_model_info_helper(
supports_xhigh_reasoning_effort=_model_info.get("supports_xhigh_reasoning_effort", None),
supports_max_reasoning_effort=_model_info.get("supports_max_reasoning_effort", None),
bedrock_output_config_effort_ceiling=_model_info.get("bedrock_output_config_effort_ceiling", None),
bedrock_converse_supports_strict_tools=_model_info.get("bedrock_converse_supports_strict_tools", None),
supports_computer_use=_model_info.get("supports_computer_use", None),
search_context_cost_per_query=_model_info.get("search_context_cost_per_query", None),
web_search_billing_unit=_model_info.get("web_search_billing_unit", None),

View file

@ -724,6 +724,7 @@
"supports_tool_choice": true
},
"anthropic.claude-haiku-4-5-20251001-v1:0": {
"supports_parallel_tool_use_config": true,
"cache_creation_input_token_cost": 1.25e-06,
"cache_creation_input_token_cost_above_1hr": 2e-06,
"cache_read_input_token_cost": 1e-07,
@ -747,6 +748,7 @@
"supports_native_structured_output": true
},
"anthropic.claude-haiku-4-5@20251001": {
"supports_parallel_tool_use_config": true,
"cache_creation_input_token_cost": 1.25e-06,
"cache_creation_input_token_cost_above_1hr": 2e-06,
"cache_read_input_token_cost": 1e-07,
@ -787,8 +789,6 @@
"output_cost_per_token_above_200k_tokens": 3e-05,
"cache_creation_input_token_cost_above_200k_tokens": 7.5e-06,
"cache_read_input_token_cost_above_200k_tokens": 6e-07,
"cache_creation_input_token_cost_above_1hr": 7.5e-06,
"cache_creation_input_token_cost_above_1hr_above_200k_tokens": 1.5e-05,
"cache_creation_input_token_cost": 3.75e-06,
"cache_read_input_token_cost": 3e-07
},
@ -813,9 +813,7 @@
"input_cost_per_token_above_200k_tokens": 6e-06,
"output_cost_per_token_above_200k_tokens": 3e-05,
"cache_creation_input_token_cost_above_200k_tokens": 7.5e-06,
"cache_read_input_token_cost_above_200k_tokens": 6e-07,
"cache_creation_input_token_cost_above_1hr": 7.5e-06,
"cache_creation_input_token_cost_above_1hr_above_200k_tokens": 1.5e-05
"cache_read_input_token_cost_above_200k_tokens": 6e-07
},
"anthropic.claude-3-7-sonnet-20240620-v1:0": {
"cache_creation_input_token_cost": 4.5e-06,
@ -965,6 +963,7 @@
"supports_vision": true
},
"anthropic.claude-opus-4-5-20251101-v1:0": {
"supports_parallel_tool_use_config": true,
"cache_creation_input_token_cost": 6.25e-06,
"cache_creation_input_token_cost_above_1hr": 1e-05,
"cache_read_input_token_cost": 5e-07,
@ -994,6 +993,7 @@
"bedrock_output_config_effort_ceiling": "high"
},
"anthropic.claude-opus-4-6-v1": {
"supports_parallel_tool_use_config": true,
"cache_creation_input_token_cost": 6.25e-06,
"cache_creation_input_token_cost_above_1hr": 1e-05,
"cache_read_input_token_cost": 5e-07,
@ -1025,6 +1025,7 @@
"bedrock_output_config_effort_ceiling": "max"
},
"global.anthropic.claude-opus-4-6-v1": {
"supports_parallel_tool_use_config": true,
"cache_creation_input_token_cost": 6.25e-06,
"cache_creation_input_token_cost_above_1hr": 1e-05,
"cache_read_input_token_cost": 5e-07,
@ -1056,6 +1057,7 @@
"bedrock_output_config_effort_ceiling": "max"
},
"us.anthropic.claude-opus-4-6-v1": {
"supports_parallel_tool_use_config": true,
"cache_creation_input_token_cost": 6.875e-06,
"cache_creation_input_token_cost_above_1hr": 1.1e-05,
"cache_read_input_token_cost": 5.5e-07,
@ -1087,6 +1089,7 @@
"bedrock_output_config_effort_ceiling": "max"
},
"eu.anthropic.claude-opus-4-6-v1": {
"supports_parallel_tool_use_config": true,
"cache_creation_input_token_cost": 6.875e-06,
"cache_creation_input_token_cost_above_1hr": 1.1e-05,
"cache_read_input_token_cost": 5.5e-07,
@ -1118,6 +1121,7 @@
"bedrock_output_config_effort_ceiling": "max"
},
"au.anthropic.claude-opus-4-6-v1": {
"supports_parallel_tool_use_config": true,
"cache_creation_input_token_cost": 6.875e-06,
"cache_creation_input_token_cost_above_1hr": 1.1e-05,
"cache_read_input_token_cost": 5.5e-07,
@ -1149,6 +1153,8 @@
"bedrock_output_config_effort_ceiling": "max"
},
"anthropic.claude-opus-4-7": {
"supports_parallel_tool_use_config": true,
"bedrock_converse_supports_strict_tools": false,
"cache_creation_input_token_cost": 6.25e-06,
"cache_creation_input_token_cost_above_1hr": 1e-05,
"cache_read_input_token_cost": 5e-07,
@ -1197,6 +1203,8 @@
"supports_output_config": true
},
"global.anthropic.claude-opus-4-7": {
"supports_parallel_tool_use_config": true,
"bedrock_converse_supports_strict_tools": false,
"cache_creation_input_token_cost": 6.25e-06,
"cache_creation_input_token_cost_above_1hr": 1e-05,
"cache_read_input_token_cost": 5e-07,
@ -1230,6 +1238,8 @@
"bedrock_output_config_effort_ceiling": "xhigh"
},
"us.anthropic.claude-opus-4-7": {
"supports_parallel_tool_use_config": true,
"bedrock_converse_supports_strict_tools": false,
"cache_creation_input_token_cost": 6.875e-06,
"cache_creation_input_token_cost_above_1hr": 1.1e-05,
"cache_read_input_token_cost": 5.5e-07,
@ -1263,6 +1273,8 @@
"bedrock_output_config_effort_ceiling": "xhigh"
},
"eu.anthropic.claude-opus-4-7": {
"supports_parallel_tool_use_config": true,
"bedrock_converse_supports_strict_tools": false,
"cache_creation_input_token_cost": 6.875e-06,
"cache_creation_input_token_cost_above_1hr": 1.1e-05,
"cache_read_input_token_cost": 5.5e-07,
@ -1296,6 +1308,8 @@
"bedrock_output_config_effort_ceiling": "xhigh"
},
"au.anthropic.claude-opus-4-7": {
"supports_parallel_tool_use_config": true,
"bedrock_converse_supports_strict_tools": false,
"cache_creation_input_token_cost": 6.875e-06,
"cache_creation_input_token_cost_above_1hr": 1.1e-05,
"cache_read_input_token_cost": 5.5e-07,
@ -1329,6 +1343,7 @@
"bedrock_output_config_effort_ceiling": "xhigh"
},
"anthropic.claude-fable-5": {
"supports_parallel_tool_use_config": true,
"cache_creation_input_token_cost": 1.25e-05,
"cache_creation_input_token_cost_above_1hr": 2e-05,
"cache_read_input_token_cost": 1e-06,
@ -1362,6 +1377,7 @@
"bedrock_output_config_effort_ceiling": "xhigh"
},
"global.anthropic.claude-fable-5": {
"supports_parallel_tool_use_config": true,
"cache_creation_input_token_cost": 1.25e-05,
"cache_creation_input_token_cost_above_1hr": 2e-05,
"cache_read_input_token_cost": 1e-06,
@ -1395,6 +1411,7 @@
"bedrock_output_config_effort_ceiling": "xhigh"
},
"us.anthropic.claude-fable-5": {
"supports_parallel_tool_use_config": true,
"cache_creation_input_token_cost": 1.375e-05,
"cache_creation_input_token_cost_above_1hr": 2.2e-05,
"cache_read_input_token_cost": 1.1e-06,
@ -1428,6 +1445,7 @@
"bedrock_output_config_effort_ceiling": "xhigh"
},
"eu.anthropic.claude-fable-5": {
"supports_parallel_tool_use_config": true,
"cache_creation_input_token_cost": 1.375e-05,
"cache_creation_input_token_cost_above_1hr": 2.2e-05,
"cache_read_input_token_cost": 1.1e-06,
@ -1461,6 +1479,8 @@
"bedrock_output_config_effort_ceiling": "xhigh"
},
"anthropic.claude-opus-4-8": {
"supports_parallel_tool_use_config": true,
"bedrock_converse_supports_strict_tools": false,
"cache_creation_input_token_cost": 6.25e-06,
"cache_creation_input_token_cost_above_1hr": 1e-05,
"cache_read_input_token_cost": 5e-07,
@ -1494,6 +1514,8 @@
"bedrock_output_config_effort_ceiling": "xhigh"
},
"global.anthropic.claude-opus-4-8": {
"supports_parallel_tool_use_config": true,
"bedrock_converse_supports_strict_tools": false,
"cache_creation_input_token_cost": 6.25e-06,
"cache_creation_input_token_cost_above_1hr": 1e-05,
"cache_read_input_token_cost": 5e-07,
@ -1527,6 +1549,8 @@
"bedrock_output_config_effort_ceiling": "xhigh"
},
"us.anthropic.claude-opus-4-8": {
"supports_parallel_tool_use_config": true,
"bedrock_converse_supports_strict_tools": false,
"cache_creation_input_token_cost": 6.875e-06,
"cache_creation_input_token_cost_above_1hr": 1.1e-05,
"cache_read_input_token_cost": 5.5e-07,
@ -1560,6 +1584,8 @@
"bedrock_output_config_effort_ceiling": "xhigh"
},
"eu.anthropic.claude-opus-4-8": {
"supports_parallel_tool_use_config": true,
"bedrock_converse_supports_strict_tools": false,
"cache_creation_input_token_cost": 6.875e-06,
"cache_creation_input_token_cost_above_1hr": 1.1e-05,
"cache_read_input_token_cost": 5.5e-07,
@ -1593,6 +1619,8 @@
"bedrock_output_config_effort_ceiling": "xhigh"
},
"au.anthropic.claude-opus-4-8": {
"supports_parallel_tool_use_config": true,
"bedrock_converse_supports_strict_tools": false,
"cache_creation_input_token_cost": 6.875e-06,
"cache_creation_input_token_cost_above_1hr": 1.1e-05,
"cache_read_input_token_cost": 5.5e-07,
@ -1626,6 +1654,8 @@
"bedrock_output_config_effort_ceiling": "xhigh"
},
"jp.anthropic.claude-opus-4-7": {
"supports_parallel_tool_use_config": true,
"bedrock_converse_supports_strict_tools": false,
"cache_creation_input_token_cost": 6.875e-06,
"cache_read_input_token_cost": 5.5e-07,
"input_cost_per_token": 5.5e-06,
@ -1657,6 +1687,7 @@
"supports_minimal_reasoning_effort": true
},
"anthropic.claude-sonnet-4-6": {
"supports_parallel_tool_use_config": true,
"cache_creation_input_token_cost": 3.75e-06,
"cache_creation_input_token_cost_above_1hr": 6e-06,
"cache_read_input_token_cost": 3e-07,
@ -1687,6 +1718,7 @@
"supports_output_config": true
},
"global.anthropic.claude-sonnet-4-6": {
"supports_parallel_tool_use_config": true,
"cache_creation_input_token_cost": 3.75e-06,
"cache_creation_input_token_cost_above_1hr": 6e-06,
"cache_read_input_token_cost": 3e-07,
@ -1717,6 +1749,7 @@
"supports_output_config": true
},
"us.anthropic.claude-sonnet-4-6": {
"supports_parallel_tool_use_config": true,
"cache_creation_input_token_cost": 4.125e-06,
"cache_creation_input_token_cost_above_1hr": 6.6e-06,
"cache_read_input_token_cost": 3.3e-07,
@ -1747,6 +1780,7 @@
"supports_output_config": true
},
"eu.anthropic.claude-sonnet-4-6": {
"supports_parallel_tool_use_config": true,
"cache_creation_input_token_cost": 4.125e-06,
"cache_creation_input_token_cost_above_1hr": 6.6e-06,
"cache_read_input_token_cost": 3.3e-07,
@ -1777,6 +1811,7 @@
"supports_output_config": true
},
"au.anthropic.claude-sonnet-4-6": {
"supports_parallel_tool_use_config": true,
"cache_creation_input_token_cost": 4.125e-06,
"cache_creation_input_token_cost_above_1hr": 6.6e-06,
"cache_read_input_token_cost": 3.3e-07,
@ -1807,6 +1842,7 @@
"supports_output_config": true
},
"jp.anthropic.claude-sonnet-4-6": {
"supports_parallel_tool_use_config": true,
"cache_creation_input_token_cost": 4.125e-06,
"cache_creation_input_token_cost_above_1hr": 6.6e-06,
"cache_read_input_token_cost": 3.3e-07,
@ -1866,6 +1902,7 @@
"supports_vision": true
},
"anthropic.claude-sonnet-4-5-20250929-v1:0": {
"supports_parallel_tool_use_config": true,
"cache_creation_input_token_cost": 3.75e-06,
"cache_creation_input_token_cost_above_1hr": 6e-06,
"cache_read_input_token_cost": 3e-07,
@ -2126,6 +2163,7 @@
"cache_creation_input_token_cost": 3.125e-07
},
"apac.anthropic.claude-haiku-4-5-20251001-v1:0": {
"supports_parallel_tool_use_config": true,
"cache_creation_input_token_cost": 1.375e-06,
"cache_read_input_token_cost": 1.1e-07,
"input_cost_per_token": 1.1e-06,
@ -2205,6 +2243,7 @@
"output_cost_per_second": 0.0
},
"au.anthropic.claude-sonnet-4-5-20250929-v1:0": {
"supports_parallel_tool_use_config": true,
"cache_creation_input_token_cost": 4.125e-06,
"cache_creation_input_token_cost_above_1hr": 6.6e-06,
"cache_read_input_token_cost": 3.3e-07,
@ -9327,6 +9366,7 @@
"cache_creation_input_token_cost": 3.75e-07
},
"bedrock/us-gov-east-1/anthropic.claude-sonnet-4-5-20250929-v1:0": {
"supports_parallel_tool_use_config": true,
"cache_creation_input_token_cost": 4.5e-06,
"cache_creation_input_token_cost_above_1hr": 7.2e-06,
"cache_read_input_token_cost": 3.6e-07,
@ -9349,6 +9389,7 @@
"supports_native_structured_output": true
},
"bedrock/us-gov-east-1/claude-sonnet-4-5-20250929-v1:0": {
"supports_parallel_tool_use_config": true,
"cache_creation_input_token_cost": 4.5e-06,
"cache_creation_input_token_cost_above_1hr": 7.2e-06,
"cache_read_input_token_cost": 3.6e-07,
@ -9502,6 +9543,7 @@
"cache_creation_input_token_cost": 3.75e-07
},
"bedrock/us-gov-west-1/anthropic.claude-sonnet-4-5-20250929-v1:0": {
"supports_parallel_tool_use_config": true,
"cache_creation_input_token_cost": 4.5e-06,
"cache_creation_input_token_cost_above_1hr": 7.2e-06,
"cache_read_input_token_cost": 3.6e-07,
@ -9524,6 +9566,7 @@
"supports_native_structured_output": true
},
"bedrock/us-gov-west-1/claude-sonnet-4-5-20250929-v1:0": {
"supports_parallel_tool_use_config": true,
"cache_creation_input_token_cost": 4.5e-06,
"cache_creation_input_token_cost_above_1hr": 7.2e-06,
"cache_read_input_token_cost": 3.6e-07,
@ -10250,6 +10293,7 @@
"supports_output_config": true
},
"claude-sonnet-4-5-20250929-v1:0": {
"supports_parallel_tool_use_config": true,
"cache_creation_input_token_cost": 3.75e-06,
"cache_creation_input_token_cost_above_1hr": 6e-06,
"cache_creation_input_token_cost_above_1hr_above_200k_tokens": 1.2e-05,
@ -14324,6 +14368,7 @@
"cache_creation_input_token_cost": 3.125e-07
},
"eu.anthropic.claude-haiku-4-5-20251001-v1:0": {
"supports_parallel_tool_use_config": true,
"cache_creation_input_token_cost": 1.375e-06,
"cache_creation_input_token_cost_above_1hr": 2.2e-06,
"cache_read_input_token_cost": 1.1e-07,
@ -14529,6 +14574,7 @@
"supports_vision": true
},
"eu.anthropic.claude-sonnet-4-5-20250929-v1:0": {
"supports_parallel_tool_use_config": true,
"cache_creation_input_token_cost": 4.125e-06,
"cache_creation_input_token_cost_above_1hr": 6.6e-06,
"cache_read_input_token_cost": 3.3e-07,
@ -19886,6 +19932,7 @@
"mode": "search"
},
"global.anthropic.claude-sonnet-4-5-20250929-v1:0": {
"supports_parallel_tool_use_config": true,
"cache_creation_input_token_cost": 3.75e-06,
"cache_creation_input_token_cost_above_1hr": 6e-06,
"cache_read_input_token_cost": 3e-07,
@ -19947,6 +19994,7 @@
"supports_vision": true
},
"global.anthropic.claude-haiku-4-5-20251001-v1:0": {
"supports_parallel_tool_use_config": true,
"cache_creation_input_token_cost": 1.25e-06,
"cache_creation_input_token_cost_above_1hr": 2e-06,
"cache_read_input_token_cost": 1e-07,
@ -24007,6 +24055,7 @@
"output_cost_per_token": 1.8e-08
},
"jp.anthropic.claude-sonnet-4-5-20250929-v1:0": {
"supports_parallel_tool_use_config": true,
"cache_creation_input_token_cost": 4.125e-06,
"cache_creation_input_token_cost_above_1hr": 6.6e-06,
"cache_read_input_token_cost": 3.3e-07,
@ -24039,6 +24088,7 @@
"supports_native_structured_output": true
},
"jp.anthropic.claude-haiku-4-5-20251001-v1:0": {
"supports_parallel_tool_use_config": true,
"cache_creation_input_token_cost": 1.375e-06,
"cache_creation_input_token_cost_above_1hr": 2.2e-06,
"cache_read_input_token_cost": 1.1e-07,
@ -32671,6 +32721,7 @@
"supports_tool_choice": true
},
"us.anthropic.claude-haiku-4-5-20251001-v1:0": {
"supports_parallel_tool_use_config": true,
"cache_creation_input_token_cost": 1.375e-06,
"cache_creation_input_token_cost_above_1hr": 2.2e-06,
"cache_read_input_token_cost": 1.1e-07,
@ -32821,6 +32872,7 @@
"supports_vision": true
},
"us.anthropic.claude-sonnet-4-5-20250929-v1:0": {
"supports_parallel_tool_use_config": true,
"cache_creation_input_token_cost": 4.125e-06,
"cache_creation_input_token_cost_above_1hr": 6.6e-06,
"cache_read_input_token_cost": 3.3e-07,
@ -32853,6 +32905,7 @@
"supports_native_structured_output": true
},
"us-gov.anthropic.claude-sonnet-4-5-20250929-v1:0": {
"supports_parallel_tool_use_config": true,
"cache_creation_input_token_cost": 4.5e-06,
"cache_creation_input_token_cost_above_1hr": 7.2e-06,
"cache_read_input_token_cost": 3.6e-07,
@ -32880,6 +32933,7 @@
"supports_native_structured_output": true
},
"au.anthropic.claude-haiku-4-5-20251001-v1:0": {
"supports_parallel_tool_use_config": true,
"cache_creation_input_token_cost": 1.375e-06,
"cache_creation_input_token_cost_above_1hr": 2.2e-06,
"cache_read_input_token_cost": 1.1e-07,
@ -32927,6 +32981,7 @@
"supports_vision": true
},
"us.anthropic.claude-opus-4-5-20251101-v1:0": {
"supports_parallel_tool_use_config": true,
"cache_creation_input_token_cost": 6.875e-06,
"cache_creation_input_token_cost_above_1hr": 1.1e-05,
"cache_read_input_token_cost": 5.5e-07,
@ -32956,6 +33011,7 @@
"bedrock_output_config_effort_ceiling": "high"
},
"global.anthropic.claude-opus-4-5-20251101-v1:0": {
"supports_parallel_tool_use_config": true,
"cache_creation_input_token_cost": 6.25e-06,
"cache_creation_input_token_cost_above_1hr": 1e-05,
"cache_read_input_token_cost": 5e-07,
@ -32985,6 +33041,7 @@
"bedrock_output_config_effort_ceiling": "high"
},
"eu.anthropic.claude-opus-4-5-20251101-v1:0": {
"supports_parallel_tool_use_config": true,
"cache_creation_input_token_cost": 6.25e-06,
"cache_read_input_token_cost": 5e-07,
"input_cost_per_token": 5e-06,
@ -42994,6 +43051,7 @@
"source": "https://aws.amazon.com/bedrock/pricing/"
},
"bedrock/us-gov-east-1/anthropic.claude-haiku-4-5-20251001-v1:0": {
"supports_parallel_tool_use_config": true,
"cache_creation_input_token_cost": 1.5e-06,
"cache_creation_input_token_cost_above_1hr": 2.4e-06,
"cache_read_input_token_cost": 1.2e-07,
@ -43017,6 +43075,7 @@
"supports_pdf_input": true
},
"bedrock/us-gov-west-1/anthropic.claude-haiku-4-5-20251001-v1:0": {
"supports_parallel_tool_use_config": true,
"cache_creation_input_token_cost": 1.5e-06,
"cache_creation_input_token_cost_above_1hr": 2.4e-06,
"cache_read_input_token_cost": 1.2e-07,

View file

@ -0,0 +1,107 @@
"""Regression tests for Bedrock Converse ``toolSpec.strict`` forwarding.
Bedrock Converse routes Claude Opus 4.7/4.8 through an Anthropic-compatible
validator that rejects ``toolSpec.strict`` even though Anthropic's native API
accepts ``strict`` as a top-level tool field for the same models. See
BerriAI/litellm#31582.
"""
import pytest
from litellm.litellm_core_utils.prompt_templates.factory import _bedrock_tools_pt
from litellm.llms.bedrock.common_utils import bedrock_converse_supports_strict_tools
_STRICT_TOOL = [
{
"type": "function",
"function": {
"name": "get_weather",
"strict": True,
"description": "Get the weather for a city",
"parameters": {
"type": "object",
"properties": {
"city": {"type": "string"},
"unit": {"type": "string", "enum": ["celsius"]},
},
"required": ["city", "unit"],
"additionalProperties": False,
},
},
}
]
@pytest.mark.parametrize(
"model_id",
[
"bedrock/us.anthropic.claude-opus-4-7",
"bedrock/us.anthropic.claude-opus-4-8",
"anthropic.claude-opus-4-7",
"anthropic.claude-opus-4-8",
"anthropic.claude-opus-4-7-v1:0",
"bedrock/eu.anthropic.claude-opus-4-8-v1:0",
"bedrock/global.anthropic.claude-opus-4-7",
],
)
def test_bedrock_tools_pt_strict_dropped_for_opus_47_48(model_id: str) -> None:
"""Opus 4.7/4.8 on Bedrock Converse reject toolSpec.strict — must be dropped."""
result = _bedrock_tools_pt(_STRICT_TOOL, model=model_id)
assert "strict" not in result[0]["toolSpec"], f"strict leaked into toolSpec for {model_id}: {result[0]['toolSpec']}"
@pytest.mark.parametrize(
"model_id",
[
"anthropic.claude-sonnet-4-5-20250929-v1:0",
"bedrock/us.anthropic.claude-sonnet-4-6",
"bedrock/us.anthropic.claude-opus-4-6",
"bedrock/us.anthropic.claude-opus-4-5",
],
)
def test_bedrock_tools_pt_strict_kept_for_other_anthropic(model_id: str) -> None:
"""Sonnet 4.5/4.6 and Opus <=4.6 accept toolSpec.strict — keep forwarding it."""
result = _bedrock_tools_pt(_STRICT_TOOL, model=model_id)
assert result[0]["toolSpec"]["strict"] is True, f"strict missing for {model_id}: {result[0]['toolSpec']}"
@pytest.mark.parametrize(
"model_id",
[
"us.amazon.nova-micro-v1:0",
"meta.llama3-2-11b-instruct-v1:0",
],
)
def test_bedrock_tools_pt_strict_dropped_for_non_anthropic(model_id: str) -> None:
"""Non-Anthropic Bedrock families reject toolSpec.strict — must be dropped."""
result = _bedrock_tools_pt(_STRICT_TOOL, model=model_id)
assert "strict" not in result[0]["toolSpec"]
def test_bedrock_converse_supports_strict_tools_helper() -> None:
"""Direct check for the gate helper used by factory.py."""
assert bedrock_converse_supports_strict_tools("bedrock/us.anthropic.claude-opus-4-7") is False
assert bedrock_converse_supports_strict_tools("bedrock/us.anthropic.claude-opus-4-8") is False
assert bedrock_converse_supports_strict_tools("anthropic.claude-sonnet-4-5-20250929-v1:0") is True
assert bedrock_converse_supports_strict_tools("bedrock/us.anthropic.claude-opus-4-6") is True
assert bedrock_converse_supports_strict_tools("us.amazon.nova-micro-v1:0") is False
assert bedrock_converse_supports_strict_tools("") is False
@pytest.mark.parametrize(
"cost_map_key",
[
"anthropic.claude-opus-4-7",
"us.anthropic.claude-opus-4-7",
"anthropic.claude-opus-4-8",
"us.anthropic.claude-opus-4-8",
],
)
def test_strict_tools_flag_set_in_model_cost_map(cost_map_key: str) -> None:
"""The gate is driven by ``bedrock_converse_supports_strict_tools: false`` in
``model_prices_and_context_window.json``, not hardcoded model patterns."""
from litellm.litellm_core_utils.get_model_cost_map import GetModelCostMap
cost_map = GetModelCostMap.load_local_model_cost_map()
assert cost_map[cost_map_key]["bedrock_converse_supports_strict_tools"] is False

View file

@ -1,5 +1,6 @@
import base64
import json
import os
from unittest.mock import MagicMock, patch
import pytest
@ -2608,102 +2609,132 @@ def test_add_cache_point_tool_block_passes_ttl_for_claude_4_5():
TTL ordering constraint (tools -> system -> messages).
Ref: https://github.com/BerriAI/litellm/issues/XXXXX
Forces the bundled local cost map so ttl eligibility (driven by
`cache_creation_input_token_cost_above_1hr` in litellm.model_cost) reads
this branch's pricing data rather than the network-fetched `main` copy,
which lacks the fix until merge.
"""
from litellm.litellm_core_utils.prompt_templates.factory import (
add_cache_point_tool_block,
)
tool_with_1h = {
"type": "function",
"function": {"name": "get_weather", "parameters": {"type": "object"}},
"cache_control": {"type": "ephemeral", "ttl": "1h"},
}
old_env = os.environ.get("LITELLM_LOCAL_MODEL_COST_MAP")
old_cost = litellm.model_cost
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
litellm.model_cost = litellm.get_model_cost_map(url="")
try:
tool_with_1h = {
"type": "function",
"function": {"name": "get_weather", "parameters": {"type": "object"}},
"cache_control": {"type": "ephemeral", "ttl": "1h"},
}
# Claude 4.5 model: ttl should be preserved
result = add_cache_point_tool_block(
tool_with_1h, model="us.anthropic.claude-sonnet-4-5-20250514-v1:0"
)
assert result is not None
assert result["cachePoint"]["type"] == "default"
assert result["cachePoint"]["ttl"] == "1h"
# Claude 4.5 model: ttl should be preserved
result = add_cache_point_tool_block(
tool_with_1h, model="jp.anthropic.claude-opus-4-7"
)
assert result is not None
assert result["cachePoint"]["type"] == "default"
assert result["cachePoint"]["ttl"] == "1h"
# Claude 4.5 model with 5m ttl: also preserved
tool_with_5m = {
"cache_control": {"type": "ephemeral", "ttl": "5m"},
}
result_5m = add_cache_point_tool_block(
tool_with_5m, model="us.anthropic.claude-sonnet-4-5-20250514-v1:0"
)
assert result_5m is not None
assert result_5m["cachePoint"]["ttl"] == "5m"
# Claude 4.5 model with 5m ttl: also preserved
tool_with_5m = {
"cache_control": {"type": "ephemeral", "ttl": "5m"},
}
result_5m = add_cache_point_tool_block(
tool_with_5m, model="jp.anthropic.claude-opus-4-7"
)
assert result_5m is not None
assert result_5m["cachePoint"]["ttl"] == "5m"
# Older model: ttl should be stripped
result_old = add_cache_point_tool_block(
tool_with_1h, model="anthropic.claude-3-5-sonnet-20241022-v2:0"
)
assert result_old is not None
assert result_old["cachePoint"]["type"] == "default"
assert "ttl" not in result_old["cachePoint"]
# Older model: ttl should be stripped
result_old = add_cache_point_tool_block(
tool_with_1h, model="anthropic.claude-3-5-sonnet-20241022-v2:0"
)
assert result_old is not None
assert result_old["cachePoint"]["type"] == "default"
assert "ttl" not in result_old["cachePoint"]
# No model provided: ttl should be stripped (safe default)
result_no_model = add_cache_point_tool_block(tool_with_1h, model=None)
assert result_no_model is not None
assert "ttl" not in result_no_model["cachePoint"]
# No model provided: ttl should be stripped (safe default)
result_no_model = add_cache_point_tool_block(tool_with_1h, model=None)
assert result_no_model is not None
assert "ttl" not in result_no_model["cachePoint"]
# No cache_control: returns None (unchanged behavior)
tool_no_cache = {
"type": "function",
"function": {"name": "get_weather", "parameters": {"type": "object"}},
}
assert add_cache_point_tool_block(tool_no_cache) is None
# No cache_control: returns None (unchanged behavior)
tool_no_cache = {
"type": "function",
"function": {"name": "get_weather", "parameters": {"type": "object"}},
}
assert add_cache_point_tool_block(tool_no_cache) is None
# cache_control without ttl: returns default cachePoint (unchanged behavior)
tool_no_ttl = {"cache_control": {"type": "ephemeral"}}
result_no_ttl = add_cache_point_tool_block(
tool_no_ttl, model="us.anthropic.claude-sonnet-4-5-20250514-v1:0"
)
assert result_no_ttl is not None
assert result_no_ttl["cachePoint"]["type"] == "default"
assert "ttl" not in result_no_ttl["cachePoint"]
# cache_control without ttl: returns default cachePoint (unchanged behavior)
tool_no_ttl = {"cache_control": {"type": "ephemeral"}}
result_no_ttl = add_cache_point_tool_block(
tool_no_ttl, model="us.anthropic.claude-sonnet-4-5-20250929-v1:0"
)
assert result_no_ttl is not None
assert result_no_ttl["cachePoint"]["type"] == "default"
assert "ttl" not in result_no_ttl["cachePoint"]
finally:
litellm.model_cost = old_cost
if old_env is None:
os.environ.pop("LITELLM_LOCAL_MODEL_COST_MAP", None)
else:
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = old_env
def test_bedrock_tools_pt_passes_ttl_for_claude_4_5():
"""
End-to-end: _bedrock_tools_pt should produce cachePoint blocks with ttl
for Claude 4.5+ models when tools have cache_control with ttl.
Forces the bundled local cost map so ttl eligibility (driven by
`cache_creation_input_token_cost_above_1hr` in litellm.model_cost) reads
this branch's pricing data rather than the network-fetched `main` copy,
which lacks the fix until merge.
"""
from litellm.litellm_core_utils.prompt_templates.factory import _bedrock_tools_pt
tools = [
{
"type": "function",
"function": {
"name": "get_weather",
"description": "Get weather",
"parameters": {
"type": "object",
"properties": {"city": {"type": "string"}},
old_env = os.environ.get("LITELLM_LOCAL_MODEL_COST_MAP")
old_cost = litellm.model_cost
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
litellm.model_cost = litellm.get_model_cost_map(url="")
try:
tools = [
{
"type": "function",
"function": {
"name": "get_weather",
"description": "Get weather",
"parameters": {
"type": "object",
"properties": {"city": {"type": "string"}},
},
},
},
"cache_control": {"type": "ephemeral", "ttl": "1h"},
}
]
"cache_control": {"type": "ephemeral", "ttl": "1h"},
}
]
# Claude 4.5: cachePoint should have ttl
result = _bedrock_tools_pt(
tools, model="us.anthropic.claude-sonnet-4-5-20250514-v1:0"
)
cache_blocks = [b for b in result if "cachePoint" in b]
assert len(cache_blocks) == 1
assert cache_blocks[0]["cachePoint"]["ttl"] == "1h"
# Claude 4.5: cachePoint should have ttl
result = _bedrock_tools_pt(tools, model="jp.anthropic.claude-opus-4-7")
cache_blocks = [b for b in result if "cachePoint" in b]
assert len(cache_blocks) == 1
assert cache_blocks[0]["cachePoint"]["ttl"] == "1h"
# Older model: cachePoint should not have ttl
result_old = _bedrock_tools_pt(
tools, model="anthropic.claude-3-5-sonnet-20241022-v2:0"
)
cache_blocks_old = [b for b in result_old if "cachePoint" in b]
assert len(cache_blocks_old) == 1
assert "ttl" not in cache_blocks_old[0]["cachePoint"]
# Older model: cachePoint should not have ttl
result_old = _bedrock_tools_pt(
tools, model="anthropic.claude-3-5-sonnet-20241022-v2:0"
)
cache_blocks_old = [b for b in result_old if "cachePoint" in b]
assert len(cache_blocks_old) == 1
assert "ttl" not in cache_blocks_old[0]["cachePoint"]
finally:
litellm.model_cost = old_cost
if old_env is None:
os.environ.pop("LITELLM_LOCAL_MODEL_COST_MAP", None)
else:
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = old_env
def test_convert_to_anthropic_tool_result_openai_file_pdf_becomes_document():

View file

@ -611,6 +611,65 @@ def test_transform_request_helper_includes_anthropic_beta_and_tools():
assert fields["tools"][0]["type"] == "computer_20250124"
def test_parallel_tool_calls_config_kept_for_sonnet_4_6():
old_env = os.environ.get("LITELLM_LOCAL_MODEL_COST_MAP")
old_cost = litellm.model_cost
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
litellm.model_cost = litellm.get_model_cost_map(url="")
try:
config = AmazonConverseConfig()
optional_params = config.map_openai_params(
model="anthropic.claude-sonnet-4-6",
non_default_params={"parallel_tool_calls": False},
optional_params={},
drop_params=False,
)
data = config._transform_request_helper(
model="anthropic.claude-sonnet-4-6",
system_content_blocks=[],
optional_params=optional_params,
messages=None,
)
assert data["additionalModelRequestFields"]["tool_choice"] == {
"disable_parallel_tool_use": True
}
finally:
litellm.model_cost = old_cost
if old_env is None:
os.environ.pop("LITELLM_LOCAL_MODEL_COST_MAP", None)
else:
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = old_env
def test_parallel_tool_calls_config_dropped_for_ttl_only_model(
monkeypatch: pytest.MonkeyPatch,
):
model = "anthropic.claude-fable-5"
monkeypatch.setitem(
litellm.model_cost,
model,
{"cache_creation_input_token_cost_above_1hr": 2e-05},
)
config = AmazonConverseConfig()
optional_params = config.map_openai_params(
model=model,
non_default_params={"parallel_tool_calls": False},
optional_params={},
drop_params=False,
)
data = config._transform_request_helper(
model=model,
system_content_blocks=[],
optional_params=optional_params,
messages=None,
)
assert "tool_choice" not in data.get("additionalModelRequestFields", {})
def test_transform_response_with_computer_use_tool():
"""Test response transformation with computer use tool call."""
import httpx
@ -4130,6 +4189,42 @@ def test_parallel_tool_calls_newer_model_adds_disable_flag():
assert "parallel_tool_calls" not in request_data["additionalModelRequestFields"]
def test_parallel_tool_calls_flag_decoupled_from_ttl_pricing(monkeypatch):
"""
The disable_parallel_tool_use gate must read supports_parallel_tool_use_config,
not the 1h-TTL pricing field: a model carrying only the former still gets the flag.
"""
from litellm.llms.bedrock.common_utils import is_claude_4_5_on_bedrock
config = AmazonConverseConfig()
model = "anthropic.claude-parallel-tool-use-only"
monkeypatch.setitem(litellm.model_cost, model, {"supports_parallel_tool_use_config": True})
assert is_claude_4_5_on_bedrock(model) is False
messages = [{"role": "user", "content": "What's the weather in SF and NYC?"}]
optional_params = config.map_openai_params(
non_default_params={"parallel_tool_calls": False, "tools": _TOOL_PARAM},
optional_params={},
model=model,
drop_params=False,
)
request_data = config.transform_request(
model=model,
messages=messages,
optional_params=optional_params,
litellm_params={},
headers={},
)
assert (
request_data["additionalModelRequestFields"]["tool_choice"][
"disable_parallel_tool_use"
]
is True
)
def test_parallel_tool_calls_older_model_drops_disable_flag():
"""Older Claude models (pre-4.5) must NOT receive disable_parallel_tool_use — Bedrock rejects it."""
config = AmazonConverseConfig()
@ -4554,6 +4649,154 @@ def test_cache_control_injection_tool_config_not_added_without_injection_point()
assert all("cachePoint" not in tool for tool in tools)
def test_cache_control_injection_tool_config_honors_ttl_for_supported_model():
"""
Regression test: cache_control_injection_points with location=tool_config
must honor the requested `control.ttl`, mirroring the message/system
cache_control behavior, instead of always emitting a bare
{"type": "default"} cachePoint with no ttl.
Forces the bundled local cost map so `is_claude_4_5_on_bedrock` (which
reads `cache_creation_input_token_cost_above_1hr` from litellm.model_cost)
sees this branch's pricing data rather than the network-fetched `main`
copy, which lacks it until merge.
"""
old_env = os.environ.get("LITELLM_LOCAL_MODEL_COST_MAP")
old_cost = litellm.model_cost
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
litellm.model_cost = litellm.get_model_cost_map(url="")
try:
config = AmazonConverseConfig()
messages = [
{"role": "user", "content": "What is the weather?"},
]
optional_params = {
"tools": [
{
"type": "function",
"function": {
"name": "get_weather",
"description": "Get weather for a location",
"parameters": {
"type": "object",
"properties": {"location": {"type": "string"}},
"required": ["location"],
},
},
}
],
"cache_control_injection_points": [
{"location": "tool_config", "control": {"type": "ephemeral", "ttl": "1h"}},
],
}
result = config._transform_request(
model="us.anthropic.claude-sonnet-4-5-20250929-v1:0",
messages=messages,
optional_params=optional_params,
litellm_params={},
)
tools = result["toolConfig"]["tools"]
assert tools[-1] == {"cachePoint": {"type": "default", "ttl": "1h"}}
finally:
litellm.model_cost = old_cost
if old_env is None:
os.environ.pop("LITELLM_LOCAL_MODEL_COST_MAP", None)
else:
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = old_env
def test_cache_control_injection_tool_config_honors_ttl_for_regional_model_lacking_own_pricing():
"""
Regression test: a regional pricing entry that omits
`cache_creation_input_token_cost_above_1hr` (e.g. `jp.anthropic.claude-opus-4-7`)
must not shadow the base model entry that carries it; the requested ttl
survives through the base-model fallback.
"""
old_env = os.environ.get("LITELLM_LOCAL_MODEL_COST_MAP")
old_cost = litellm.model_cost
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
litellm.model_cost = litellm.get_model_cost_map(url="")
try:
assert "cache_creation_input_token_cost_above_1hr" not in litellm.model_cost["jp.anthropic.claude-opus-4-7"]
assert "cache_creation_input_token_cost_above_1hr" in litellm.model_cost["anthropic.claude-opus-4-7"]
config = AmazonConverseConfig()
messages = [
{"role": "user", "content": "What is the weather?"},
]
optional_params = {
"tools": [
{
"type": "function",
"function": {
"name": "get_weather",
"description": "Get weather for a location",
"parameters": {
"type": "object",
"properties": {"location": {"type": "string"}},
"required": ["location"],
},
},
}
],
"cache_control_injection_points": [
{"location": "tool_config", "control": {"type": "ephemeral", "ttl": "1h"}},
],
}
result = config._transform_request(
model="jp.anthropic.claude-opus-4-7",
messages=messages,
optional_params=optional_params,
litellm_params={},
)
tools = result["toolConfig"]["tools"]
assert tools[-1] == {"cachePoint": {"type": "default", "ttl": "1h"}}
finally:
litellm.model_cost = old_cost
if old_env is None:
os.environ.pop("LITELLM_LOCAL_MODEL_COST_MAP", None)
else:
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = old_env
def test_cache_control_injection_tool_config_drops_ttl_for_unsupported_model():
"""
Models that don't support extended TTL caching (only Claude 4.5+ on
Bedrock does) must fall back to the default cachePoint with no ttl,
even if the caller requested one, matching message/system behavior.
"""
config = AmazonConverseConfig()
messages = [
{"role": "user", "content": "What is the weather?"},
]
optional_params = {
"tools": [
{
"type": "function",
"function": {
"name": "get_weather",
"description": "Get weather for a location",
"parameters": {
"type": "object",
"properties": {"location": {"type": "string"}},
"required": ["location"],
},
},
}
],
"cache_control_injection_points": [
{"location": "tool_config", "control": {"type": "ephemeral", "ttl": "1h"}},
],
}
result = config._transform_request(
model="anthropic.claude-3-5-haiku-20241022-v1:0",
messages=messages,
optional_params=optional_params,
litellm_params={},
)
tools = result["toolConfig"]["tools"]
assert tools[-1] == {"cachePoint": {"type": "default"}}
def test_translate_response_format_json_schema_still_injects_tool():
"""
response_format with an explicit json_schema should still use the

View file

@ -492,7 +492,7 @@ def test_bedrock_invoke_messages_transform_converts_custom_tool_schema_type_to_o
assert result["tools"][0]["type"] == "custom"
def test_remove_ttl_from_cache_control_processes_tools():
def test_remove_ttl_from_cache_control_processes_tools(local_model_cost_map):
"""
Ensure _remove_ttl_from_cache_control also sanitizes cache_control on tools.
@ -538,7 +538,7 @@ def test_remove_ttl_from_cache_control_processes_tools():
assert "ttl" not in request["system"][0]["cache_control"]
def test_remove_ttl_from_cache_control_preserves_tools_ttl_for_claude_4_5():
def test_remove_ttl_from_cache_control_preserves_tools_ttl_for_claude_4_5(local_model_cost_map):
"""
For Claude 4.5+ models, ttl in ["5m", "1h"] should be preserved on tools,
just like it is for system and messages.
@ -564,7 +564,7 @@ def test_remove_ttl_from_cache_control_preserves_tools_ttl_for_claude_4_5():
}
cfg._remove_ttl_from_cache_control(
request, model="us.anthropic.claude-sonnet-4-5-20250514-v1:0"
request, model="us.anthropic.claude-sonnet-4-5-20250929-v1:0"
)
# Both tools and system should preserve ttl for Claude 4.5

View file

@ -445,3 +445,31 @@ def test_explicit_invoke_route_does_not_match_async_invoke():
BedrockModelInfo._explicit_async_invoke_route(f"bedrock/{async_invoke_model}")
is True
)
def test_capability_lookups_fall_back_to_base_model_when_regional_entry_lacks_field(monkeypatch):
"""
Regression test: a regional model_cost entry without the capability field
must not shadow a base entry that has it (`get(model) or get(base)` used to
short-circuit on the truthy regional dict and drop the capability).
"""
import litellm
from litellm.llms.bedrock.common_utils import (
bedrock_converse_supports_parallel_tool_use_config,
is_claude_4_5_on_bedrock,
)
base = "anthropic.claude-fallback-test"
regional = f"eu.{base}"
monkeypatch.setitem(litellm.model_cost, regional, {"input_cost_per_token": 1e-06})
monkeypatch.setitem(
litellm.model_cost,
base,
{
"cache_creation_input_token_cost_above_1hr": 1e-05,
"supports_parallel_tool_use_config": True,
},
)
assert is_claude_4_5_on_bedrock(regional) is True
assert bedrock_converse_supports_parallel_tool_use_config(regional) is True

View file

@ -0,0 +1,80 @@
"""Tests for token-endpoint client authentication (client_secret_basic vs client_secret_post)."""
import base64
import pytest
from litellm.proxy._experimental.mcp_server.auth.token_endpoint_auth import (
TokenEndpointAuthConfigError,
build_token_endpoint_client_auth,
normalize_token_endpoint_auth_method,
)
def _expected_basic(client_id: str, client_secret: str) -> str:
return "Basic " + base64.b64encode(f"{client_id}:{client_secret}".encode()).decode()
def test_basic_puts_credentials_in_header_and_not_body():
auth = build_token_endpoint_client_auth(auth_method="client_secret_basic", client_id="cid", client_secret="sec")
assert auth.headers == {"Authorization": _expected_basic("cid", "sec")}
assert "client_secret" not in auth.body
assert auth.body == {}
def test_basic_form_urlencodes_reserved_characters():
"""RFC 6749 2.3.1: client_id and client_secret are form-urlencoded before the ':' join, so reserved
characters survive base64 transport instead of corrupting the username/password split."""
auth = build_token_endpoint_client_auth(
auth_method="client_secret_basic", client_id="client:one", client_secret="sec+ret:two"
)
decoded = base64.b64decode(auth.headers["Authorization"].removeprefix("Basic ")).decode()
assert decoded == "client%3Aone:sec%2Bret%3Atwo"
def test_post_default_puts_credentials_in_body_and_no_auth_header():
auth = build_token_endpoint_client_auth(auth_method="client_secret_post", client_id="cid", client_secret="sec")
assert auth.headers == {}
assert auth.body == {"client_id": "cid", "client_secret": "sec"}
def test_none_method_defaults_to_post():
auth = build_token_endpoint_client_auth(auth_method=None, client_id="cid", client_secret="sec")
assert auth.headers == {}
assert auth.body == {"client_id": "cid", "client_secret": "sec"}
def test_explicit_basic_without_secret_raises():
"""client_secret_basic is a confidential-client method; a missing secret is a misconfiguration
that must surface, not silently downgrade to a body request (RFC 6749; the no-silent-fallback rule)."""
with pytest.raises(TokenEndpointAuthConfigError):
build_token_endpoint_client_auth(auth_method="client_secret_basic", client_id="cid", client_secret=None)
def test_explicit_basic_without_client_id_raises():
with pytest.raises(TokenEndpointAuthConfigError):
build_token_endpoint_client_auth(auth_method="client_secret_basic", client_id=None, client_secret="sec")
def test_default_method_without_secret_is_public_client_post():
"""A secretless client_id under the default method is the legitimate public-client / PKCE case:
client_id goes in the body, no secret, no error."""
auth = build_token_endpoint_client_auth(auth_method=None, client_id="cid", client_secret=None)
assert auth.headers == {}
assert auth.body == {"client_id": "cid"}
def test_explicit_post_without_secret_does_not_raise():
"""Unlike basic, explicit client_secret_post degrades to a valid public-client request, so it
does not error on a missing secret."""
auth = build_token_endpoint_client_auth(auth_method="client_secret_post", client_id="cid", client_secret=None)
assert auth.headers == {}
assert auth.body == {"client_id": "cid"}
def test_normalize_only_accepts_known_methods():
assert normalize_token_endpoint_auth_method("client_secret_basic") == "client_secret_basic"
assert normalize_token_endpoint_auth_method("client_secret_post") == "client_secret_post"
assert normalize_token_endpoint_auth_method("private_key_jwt") is None
assert normalize_token_endpoint_auth_method(None) is None
assert normalize_token_endpoint_auth_method(123) is None

View file

@ -509,3 +509,31 @@ async def test_database_loading_token_exchange_scopes_from_credentials():
assert server.token_exchange_endpoint == "https://idp.example.com/oauth2/token"
assert server.audience == "api://db-mcp"
assert server.scopes == ["db.read", "db.write"]
@pytest.mark.asyncio
async def test_exchange_token_uses_client_secret_basic_when_configured():
"""LIT-4091: token exchange with token_endpoint_auth_method=client_secret_basic sends the
client credentials as HTTP Basic and omits client_secret from the body."""
import base64
handler = TokenExchangeHandler()
server = _obo_server(
server_id="srv-obo-basic", token_endpoint_auth_method="client_secret_basic"
)
mock_client = AsyncMock()
mock_client.post.return_value = _exchange_response("scoped-basic")
with patch(
"litellm.proxy._experimental.mcp_server.auth.token_exchange.get_async_httpx_client",
return_value=mock_client,
):
result = await handler.exchange_token("user-jwt-basic", server)
assert result == "scoped-basic"
_, kwargs = mock_client.post.call_args
expected = "Basic " + base64.b64encode(b"litellm-client-id:litellm-client-secret").decode()
assert kwargs["headers"]["Authorization"] == expected
assert "client_secret" not in kwargs["data"]
assert "client_id" not in kwargs["data"]
assert kwargs["data"]["grant_type"] == TOKEN_EXCHANGE_GRANT_TYPE

View file

@ -16,10 +16,12 @@ class _Server:
token_url="https://idp.example.com/token",
client_id="cid",
client_secret="sec",
token_endpoint_auth_method=None,
):
self.token_url = token_url
self.client_id = client_id
self.client_secret = client_secret
self.token_endpoint_auth_method = token_endpoint_auth_method
def _lookup(server):
@ -27,9 +29,9 @@ def _lookup(server):
def _endpoint(body, sink=None):
async def post(url, form):
async def post(url, form, headers):
if sink is not None:
sink.append((url, form))
sink.append((url, form, headers))
return body
return post
@ -81,8 +83,8 @@ async def test_refreshes_persists_and_returns_typed_token():
assert token.expires_at == 1000.0 + 3600 # clock + expires_in -> epoch
# the rotated triple is persisted for (user, server) with parsed scopes
assert persisted == [("alice", "srv", "new-at", "new-rt", 3600, ("a", "b"))]
# the grant carried the refresh_token + client credentials
url, form = posted[0]
# the grant carried the refresh_token + client credentials in the body (client_secret_post default)
url, form, headers = posted[0]
assert url == "https://idp.example.com/token"
assert form == {
"grant_type": "refresh_token",
@ -90,6 +92,44 @@ async def test_refreshes_persists_and_returns_typed_token():
"client_id": "cid",
"client_secret": "sec",
}
assert "Authorization" not in headers
@pytest.mark.asyncio
async def test_client_secret_basic_sends_authorization_header_not_body():
"""A server with token_endpoint_auth_method=client_secret_basic authenticates via HTTP Basic;
the secret must not also leak into the form body."""
import base64
posted = []
server = _Server(token_endpoint_auth_method="client_secret_basic")
refresher = _refresher(
server=server,
body={"access_token": "new-at"},
post_sink=posted,
)
token = await refresher.refresh(
"alice", "srv", OAuthToken(access_token="old", refresh_token="old-rt")
)
assert token is not None
_url, form, headers = posted[0]
expected = "Basic " + base64.b64encode(b"cid:sec").decode()
assert headers["Authorization"] == expected
assert "client_secret" not in form
assert "client_id" not in form
assert form == {"grant_type": "refresh_token", "refresh_token": "old-rt"}
@pytest.mark.asyncio
async def test_client_secret_basic_without_secret_is_a_failed_refresh():
"""A server set to client_secret_basic but missing its secret cannot authenticate; the refresh
returns None (failed refresh -> needs reauth) and never posts a downgraded request to the IdP."""
posted = []
server = _Server(client_secret=None, token_endpoint_auth_method="client_secret_basic")
refresher = _refresher(server=server, body={"access_token": "x"}, post_sink=posted)
assert await refresher.refresh("a", "s", OAuthToken("old", refresh_token="rt")) is None
assert posted == [] # never hit the IdP
@pytest.mark.asyncio

View file

@ -643,3 +643,82 @@ async def test_rotate_user_env_vars_skips_undecryptable_rows():
assert prisma.db.litellm_mcpuserenvvars.update.call_count == 1
where = prisma.db.litellm_mcpuserenvvars.update.call_args.kwargs["where"]
assert where["user_id_server_id"]["server_id"] == "srv-ok"
@pytest.mark.asyncio
async def test_refresh_user_oauth_token_uses_client_secret_basic(monkeypatch):
"""LIT-4091: a per-user refresh against a server with token_endpoint_auth_method=client_secret_basic
sends HTTP Basic and keeps the secret out of the body."""
import litellm.proxy._experimental.mcp_server.db as db_mod
server = MagicMock()
server.token_url = "https://idp.example.com/oauth2/token"
server.server_id = "srv"
server.client_id = "cid"
server.client_secret = "sec"
server.token_endpoint_auth_method = "client_secret_basic"
mock_response = MagicMock()
mock_response.json.return_value = {"access_token": "new-at", "expires_in": 3600}
mock_response.raise_for_status = MagicMock()
mock_client = AsyncMock()
mock_client.post.return_value = mock_response
monkeypatch.setattr(db_mod, "get_async_httpx_client", lambda **kwargs: mock_client)
monkeypatch.setattr(db_mod, "store_user_oauth_credential", AsyncMock())
monkeypatch.setattr(
db_mod, "get_user_oauth_credential", AsyncMock(return_value={"access_token": "new-at"})
)
result = await db_mod.refresh_user_oauth_token(
prisma_client=MagicMock(),
user_id="alice",
server=server,
cred={"refresh_token": "rt"},
)
assert result is not None
_, kwargs = mock_client.post.call_args
assert kwargs["headers"]["Authorization"] == "Basic " + base64.b64encode(b"cid:sec").decode()
assert "client_secret" not in kwargs["data"]
assert "client_id" not in kwargs["data"]
assert kwargs["data"]["grant_type"] == "refresh_token"
assert kwargs["data"]["refresh_token"] == "rt"
@pytest.mark.asyncio
async def test_refresh_user_oauth_token_defaults_to_client_secret_post(monkeypatch):
"""Backward compatibility: with no token_endpoint_auth_method the refresh keeps credentials in
the body (client_secret_post) and sends no Authorization header."""
import litellm.proxy._experimental.mcp_server.db as db_mod
server = MagicMock()
server.token_url = "https://idp.example.com/oauth2/token"
server.server_id = "srv"
server.client_id = "cid"
server.client_secret = "sec"
server.token_endpoint_auth_method = None
mock_response = MagicMock()
mock_response.json.return_value = {"access_token": "new-at", "expires_in": 3600}
mock_response.raise_for_status = MagicMock()
mock_client = AsyncMock()
mock_client.post.return_value = mock_response
monkeypatch.setattr(db_mod, "get_async_httpx_client", lambda **kwargs: mock_client)
monkeypatch.setattr(db_mod, "store_user_oauth_credential", AsyncMock())
monkeypatch.setattr(
db_mod, "get_user_oauth_credential", AsyncMock(return_value={"access_token": "new-at"})
)
await db_mod.refresh_user_oauth_token(
prisma_client=MagicMock(),
user_id="alice",
server=server,
cred={"refresh_token": "rt"},
)
_, kwargs = mock_client.post.call_args
assert "Authorization" not in kwargs["headers"]
assert kwargs["data"]["client_id"] == "cid"
assert kwargs["data"]["client_secret"] == "sec"

View file

@ -534,6 +534,548 @@ async def test_register_client_remote_registration_success():
)
@pytest.mark.asyncio
async def test_register_client_persists_dcr_client_identity():
"""A dynamic client registration (RFC 7591) must persist the issued client_id /
client_secret / token_endpoint_auth_method and the token_url onto the server row so
autonomous refresh can authenticate as the registered client. Without persistence the
minted client_id is discarded and the refresh_token grant has no client identity."""
try:
from fastapi import Request
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
register_client_with_server,
)
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
global_mcp_server_manager,
)
from litellm.proxy._types import MCPTransport
from litellm.types.mcp import MCPAuth
from litellm.types.mcp_server.mcp_server_manager import MCPServer
except ImportError:
pytest.skip("MCP discoverable endpoints not available")
oauth2_server = MCPServer(
server_id="remote_server",
name="remote_server",
server_name="remote_server",
alias="remote_server",
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_request = MagicMock(spec=Request)
mock_request.base_url = "https://proxy.litellm.example/"
mock_request.headers = {}
mock_response = MagicMock()
mock_response.json.return_value = {
"client_id": "generated-client",
"client_secret": "generated-secret",
"token_endpoint_auth_method": "client_secret_basic",
}
mock_response.raise_for_status = MagicMock()
mock_async_client = MagicMock()
mock_async_client.post = AsyncMock(return_value=mock_response)
mock_update = AsyncMock(return_value=MagicMock())
mock_update_server = AsyncMock()
with (
patch(
"litellm.proxy._experimental.mcp_server.discoverable_endpoints.get_async_httpx_client",
return_value=mock_async_client,
),
patch("litellm.proxy.utils.get_prisma_client_or_throw", return_value=MagicMock()),
patch("litellm.proxy._experimental.mcp_server.db.update_mcp_server", new=mock_update),
patch.object(global_mcp_server_manager, "update_server", new=mock_update_server),
):
response = await register_client_with_server(
request=mock_request,
mcp_server=oauth2_server,
client_name="Litellm Proxy",
grant_types=["authorization_code", "refresh_token"],
response_types=["code"],
token_endpoint_auth_method="client_secret_basic",
persist_credentials=True,
)
import json
assert response.status_code == 200
assert json.loads(response.body.decode("utf-8")) == mock_response.json.return_value
mock_update.assert_called_once()
update_data = mock_update.call_args.kwargs["data"]
assert update_data.server_id == "remote_server"
assert update_data.token_url == "https://provider.example/oauth/token"
assert update_data.credentials["client_id"] == "generated-client"
assert update_data.credentials["client_secret"] == "generated-secret"
assert update_data.credentials["token_endpoint_auth_method"] == "client_secret_basic"
mock_update_server.assert_called_once()
@pytest.mark.asyncio
async def test_register_client_does_not_clobber_token_url_when_absent():
"""When the in-memory server has no token_url, the DCR persist must omit it from the
partial update rather than passing None, so exclude_unset leaves the token_url column
untouched instead of overwriting an existing value with NULL."""
try:
from fastapi import Request
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
register_client_with_server,
)
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
global_mcp_server_manager,
)
from litellm.proxy._types import MCPTransport
from litellm.types.mcp import MCPAuth
from litellm.types.mcp_server.mcp_server_manager import MCPServer
except ImportError:
pytest.skip("MCP discoverable endpoints not available")
oauth2_server = MCPServer(
server_id="remote_server",
name="remote_server",
server_name="remote_server",
alias="remote_server",
transport=MCPTransport.http,
auth_type=MCPAuth.oauth2,
client_id=None,
client_secret=None,
authorization_url="https://provider.example/oauth/authorize",
token_url=None,
registration_url="https://provider.example/oauth/register",
)
mock_request = MagicMock(spec=Request)
mock_request.base_url = "https://proxy.litellm.example/"
mock_request.headers = {}
mock_response = MagicMock()
mock_response.json.return_value = {"client_id": "generated-client"}
mock_response.raise_for_status = MagicMock()
mock_async_client = MagicMock()
mock_async_client.post = AsyncMock(return_value=mock_response)
mock_update = AsyncMock(return_value=MagicMock())
mock_update_server = AsyncMock()
with (
patch(
"litellm.proxy._experimental.mcp_server.discoverable_endpoints.get_async_httpx_client",
return_value=mock_async_client,
),
patch("litellm.proxy.utils.get_prisma_client_or_throw", return_value=MagicMock()),
patch("litellm.proxy._experimental.mcp_server.db.update_mcp_server", new=mock_update),
patch.object(global_mcp_server_manager, "update_server", new=mock_update_server),
):
await register_client_with_server(
request=mock_request,
mcp_server=oauth2_server,
client_name="Litellm Proxy",
grant_types=["authorization_code", "refresh_token"],
response_types=["code"],
token_endpoint_auth_method="none",
persist_credentials=True,
)
mock_update.assert_called_once()
update_data = mock_update.call_args.kwargs["data"]
assert update_data.credentials["client_id"] == "generated-client"
assert "token_url" not in update_data.model_fields_set
@pytest.mark.asyncio
async def test_register_client_reuses_persisted_client_id_for_non_admin_when_registry_is_stale():
try:
from fastapi import Request
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
register_client_with_server,
)
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
global_mcp_server_manager,
)
from litellm.proxy._types import MCPTransport
from litellm.types.mcp import MCPAuth
from litellm.types.mcp_server.mcp_server_manager import MCPServer
except ImportError:
pytest.skip("MCP discoverable endpoints not available")
oauth2_server = MCPServer(
server_id="remote_server",
name="remote_server",
server_name="remote_server",
alias="remote_server",
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_request = MagicMock(spec=Request)
mock_request.base_url = "https://proxy.litellm.example/"
mock_request.headers = {}
persisted_server = MagicMock()
persisted_server.credentials = {"client_id": "persisted-client"}
mock_get_mcp_server = AsyncMock(return_value=persisted_server)
mock_update_mcp_server = AsyncMock()
mock_update_server = AsyncMock()
mock_async_client = MagicMock()
mock_async_client.post = AsyncMock()
with (
patch(
"litellm.proxy._experimental.mcp_server.discoverable_endpoints.get_async_httpx_client",
return_value=mock_async_client,
),
patch("litellm.proxy.utils.get_prisma_client_or_throw", return_value=MagicMock()),
patch(
"litellm.proxy._experimental.mcp_server.db.get_mcp_server",
new=mock_get_mcp_server,
),
patch(
"litellm.proxy._experimental.mcp_server.db.update_mcp_server",
new=mock_update_mcp_server,
),
patch.object(global_mcp_server_manager, "update_server", new=mock_update_server),
):
response = await register_client_with_server(
request=mock_request,
mcp_server=oauth2_server,
client_name="Litellm Proxy",
grant_types=["authorization_code", "refresh_token"],
response_types=["code"],
token_endpoint_auth_method="none",
persist_credentials=False,
)
assert response["client_id"] == "remote_server"
assert oauth2_server.client_id == "persisted-client"
mock_async_client.post.assert_not_called()
mock_update_mcp_server.assert_not_called()
mock_update_server.assert_called_once_with(persisted_server)
@pytest.mark.asyncio
async def test_register_client_reuse_refreshes_request_server_when_manager_update_fails():
try:
from fastapi import Request
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
register_client_with_server,
)
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
global_mcp_server_manager,
)
from litellm.proxy._types import MCPTransport
from litellm.types.mcp import MCPAuth
from litellm.types.mcp_server.mcp_server_manager import MCPServer
except ImportError:
pytest.skip("MCP discoverable endpoints not available")
oauth2_server = MCPServer(
server_id="remote_server",
name="remote_server",
server_name="remote_server",
alias="remote_server",
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_request = MagicMock(spec=Request)
mock_request.base_url = "https://proxy.litellm.example/"
mock_request.headers = {}
persisted_server = MagicMock()
persisted_server.credentials = {
"client_id": "persisted-client",
"client_secret": "persisted-secret",
"token_endpoint_auth_method": "client_secret_basic",
}
mock_get_mcp_server = AsyncMock(return_value=persisted_server)
mock_async_client = MagicMock()
mock_async_client.post = AsyncMock()
mock_update_server = AsyncMock(side_effect=RuntimeError("registry update failed"))
with (
patch(
"litellm.proxy._experimental.mcp_server.discoverable_endpoints.get_async_httpx_client",
return_value=mock_async_client,
),
patch("litellm.proxy.utils.get_prisma_client_or_throw", return_value=MagicMock()),
patch(
"litellm.proxy._experimental.mcp_server.db.get_mcp_server",
new=mock_get_mcp_server,
),
patch.object(global_mcp_server_manager, "update_server", new=mock_update_server),
):
response = await register_client_with_server(
request=mock_request,
mcp_server=oauth2_server,
client_name="Litellm Proxy",
grant_types=["authorization_code", "refresh_token"],
response_types=["code"],
token_endpoint_auth_method="none",
persist_credentials=False,
)
assert response["client_id"] == "remote_server"
assert oauth2_server.client_id == "persisted-client"
assert oauth2_server.client_secret == "persisted-secret"
assert oauth2_server.token_endpoint_auth_method == "client_secret_basic"
mock_async_client.post.assert_not_called()
mock_update_server.assert_called_once_with(persisted_server)
@pytest.mark.asyncio
async def test_register_client_returns_reused_client_when_concurrent_persist_wins():
try:
from fastapi import Request
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
register_client_with_server,
)
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
global_mcp_server_manager,
)
from litellm.proxy._types import MCPTransport
from litellm.types.mcp import MCPAuth
from litellm.types.mcp_server.mcp_server_manager import MCPServer
except ImportError:
pytest.skip("MCP discoverable endpoints not available")
oauth2_server = MCPServer(
server_id="remote_server",
name="remote_server",
server_name="remote_server",
alias="remote_server",
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_request = MagicMock(spec=Request)
mock_request.base_url = "https://proxy.litellm.example/"
mock_request.headers = {}
mock_response = MagicMock()
mock_response.json.return_value = {"client_id": "generated-client"}
mock_response.raise_for_status = MagicMock()
mock_async_client = MagicMock()
mock_async_client.post = AsyncMock(return_value=mock_response)
persisted_server = MagicMock()
persisted_server.credentials = {"client_id": "persisted-client"}
mock_get_mcp_server = AsyncMock(side_effect=[None, persisted_server])
mock_update_mcp_server = AsyncMock()
mock_update_server = AsyncMock()
with (
patch(
"litellm.proxy._experimental.mcp_server.discoverable_endpoints.get_async_httpx_client",
return_value=mock_async_client,
),
patch("litellm.proxy.utils.get_prisma_client_or_throw", return_value=MagicMock()),
patch(
"litellm.proxy._experimental.mcp_server.db.get_mcp_server",
new=mock_get_mcp_server,
),
patch(
"litellm.proxy._experimental.mcp_server.db.update_mcp_server",
new=mock_update_mcp_server,
),
patch.object(global_mcp_server_manager, "update_server", new=mock_update_server),
):
response = await register_client_with_server(
request=mock_request,
mcp_server=oauth2_server,
client_name="Litellm Proxy",
grant_types=["authorization_code", "refresh_token"],
response_types=["code"],
token_endpoint_auth_method="none",
persist_credentials=True,
)
assert response["client_id"] == "remote_server"
assert oauth2_server.client_id == "persisted-client"
mock_async_client.post.assert_called_once()
mock_update_mcp_server.assert_not_called()
mock_update_server.assert_called_once_with(persisted_server)
@pytest.mark.asyncio
async def test_register_client_reuses_existing_client_id_without_re_dcr():
"""A server that already has a client_id (admin-configured or previously DCR'd) must be
reused, not re-registered, even without a client_secret. A client_id is one-per-application
in OAuth and shared across users; re-minting per authorize would orphan other users' refresh
tokens by overwriting the server's client_id."""
try:
from fastapi import Request
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
register_client,
)
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
global_mcp_server_manager,
)
from litellm.proxy._types import MCPTransport
from litellm.types.mcp import MCPAuth
from litellm.types.mcp_server.mcp_server_manager import MCPServer
except ImportError:
pytest.skip("MCP discoverable endpoints not available")
global_mcp_server_manager.registry.clear()
oauth2_server = MCPServer(
server_id="remote_server",
name="remote_server",
server_name="remote_server",
alias="remote_server",
transport=MCPTransport.http,
auth_type=MCPAuth.oauth2,
client_id="existing-shared-client",
client_secret=None,
authorization_url="https://provider.example/oauth/authorize",
token_url="https://provider.example/oauth/token",
registration_url="https://provider.example/oauth/register",
)
global_mcp_server_manager.registry[oauth2_server.server_id] = oauth2_server
mock_request = MagicMock(spec=Request)
mock_request.base_url = "https://proxy.litellm.example/"
mock_request.headers = {}
request_payload = {
"client_name": "Litellm Proxy",
"grant_types": ["authorization_code", "refresh_token"],
"response_types": ["code"],
"token_endpoint_auth_method": "none",
}
mock_async_client = MagicMock()
mock_async_client.post = AsyncMock()
try:
with (
patch(
"litellm.proxy._experimental.mcp_server.discoverable_endpoints._read_request_body",
new=AsyncMock(return_value=request_payload),
),
patch(
"litellm.proxy._experimental.mcp_server.discoverable_endpoints.get_async_httpx_client",
return_value=mock_async_client,
),
):
response = await register_client(request=mock_request, mcp_server_name=oauth2_server.server_name)
finally:
global_mcp_server_manager.registry.clear()
mock_async_client.post.assert_not_called()
body = response if isinstance(response, dict) else json.loads(response.body.decode("utf-8"))
assert body["client_secret"] == "dummy"
@pytest.mark.asyncio
async def test_public_register_route_does_not_persist_client_credentials():
"""The unauthenticated root /register route must not persist the DCR result onto the
server row; only the authenticated management path passes persist_credentials=True. An
external caller could otherwise bind a caller-controlled client (and leak its secret) to
a server that has no client yet."""
try:
from fastapi import Request
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
register_client,
)
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
global_mcp_server_manager,
)
from litellm.proxy._types import MCPTransport
from litellm.types.mcp import MCPAuth
from litellm.types.mcp_server.mcp_server_manager import MCPServer
except ImportError:
pytest.skip("MCP discoverable endpoints not available")
global_mcp_server_manager.registry.clear()
oauth2_server = MCPServer(
server_id="remote_server",
name="remote_server",
server_name="remote_server",
alias="remote_server",
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",
)
global_mcp_server_manager.registry[oauth2_server.server_id] = oauth2_server
mock_request = MagicMock(spec=Request)
mock_request.base_url = "https://proxy.litellm.example/"
mock_request.headers = {}
request_payload = {
"client_name": "attacker",
"grant_types": ["authorization_code"],
"response_types": ["code"],
"token_endpoint_auth_method": "none",
}
mock_response = MagicMock()
mock_response.json.return_value = {
"client_id": "attacker-client",
"client_secret": "attacker-secret",
}
mock_response.raise_for_status = MagicMock()
mock_async_client = MagicMock()
mock_async_client.post = AsyncMock(return_value=mock_response)
mock_update = AsyncMock(return_value=MagicMock())
try:
with (
patch(
"litellm.proxy._experimental.mcp_server.discoverable_endpoints._read_request_body",
new=AsyncMock(return_value=request_payload),
),
patch(
"litellm.proxy._experimental.mcp_server.discoverable_endpoints.get_async_httpx_client",
return_value=mock_async_client,
),
patch("litellm.proxy.utils.get_prisma_client_or_throw", return_value=MagicMock()),
patch.object(global_mcp_server_manager, "update_server", new=AsyncMock()),
patch("litellm.proxy._experimental.mcp_server.db.update_mcp_server", new=mock_update),
):
await register_client(request=mock_request, mcp_server_name=oauth2_server.server_name)
finally:
global_mcp_server_manager.registry.clear()
mock_update.assert_not_called()
@pytest.mark.asyncio
@pytest.mark.usefixtures("trust_xff")
async def test_authorize_endpoint_respects_x_forwarded_proto():
@ -2733,3 +3275,128 @@ async def test_token_exchange_passes_through_upstream_expires_in():
{"access_token": "tok", "token_type": "Bearer", "expires_in": 43200}
)
assert body["expires_in"] == 43200
@pytest.mark.asyncio
async def test_token_endpoint_uses_client_secret_basic_when_configured():
"""LIT-4091: a server with token_endpoint_auth_method=client_secret_basic must send the
credentials as an HTTP Basic Authorization header and omit client_secret from the body;
providers requiring Basic rejected body credentials with invalid_client."""
import base64
from unittest.mock import AsyncMock
from fastapi import Request
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
token_endpoint,
)
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
global_mcp_server_manager,
)
from litellm.proxy._types import MCPTransport
from litellm.types.mcp import MCPAuth
from litellm.types.mcp_server.mcp_server_manager import MCPServer
global_mcp_server_manager.registry.clear()
oauth2_server = MCPServer(
server_id="basic_mcp",
name="basic_mcp",
server_name="basic_mcp",
alias="basic_mcp",
transport=MCPTransport.http,
auth_type=MCPAuth.oauth2,
client_id="basic-client",
client_secret="basic-secret",
authorization_url="https://idp.example.com/authorize",
token_url="https://idp.example.com/oauth2/token",
token_endpoint_auth_method="client_secret_basic",
)
global_mcp_server_manager.registry[oauth2_server.server_id] = oauth2_server
mock_request = MagicMock(spec=Request)
mock_request.base_url = "https://litellm-proxy.example.com/"
mock_request.headers = {}
mock_response = MagicMock()
mock_response.json.return_value = {
"access_token": "at",
"token_type": "Bearer",
"expires_in": 3599,
}
mock_response.raise_for_status = MagicMock()
with patch(
"litellm.proxy._experimental.mcp_server.discoverable_endpoints.get_async_httpx_client"
) as mock_get_client:
mock_async_client = MagicMock()
mock_async_client.post = AsyncMock(return_value=mock_response)
mock_get_client.return_value = mock_async_client
await token_endpoint(
request=mock_request,
grant_type="authorization_code",
code="auth-code",
redirect_uri="http://localhost/callback",
client_id="basic-client",
mcp_server_name="basic_mcp",
client_secret="basic-secret",
code_verifier="verifier",
)
call_args = mock_async_client.post.call_args
expected = "Basic " + base64.b64encode(b"basic-client:basic-secret").decode()
assert call_args[1]["headers"]["Authorization"] == expected
assert "client_secret" not in call_args[1]["data"]
assert "client_id" not in call_args[1]["data"]
assert call_args[1]["data"]["grant_type"] == "authorization_code"
assert call_args[1]["data"]["code"] == "auth-code"
@pytest.mark.asyncio
async def test_token_endpoint_client_secret_basic_without_secret_returns_400():
"""A server configured client_secret_basic but missing its secret is a misconfiguration; the
inbound /token endpoint surfaces it as a 400 rather than silently posting a downgraded request."""
from fastapi import HTTPException, Request
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
token_endpoint,
)
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
global_mcp_server_manager,
)
from litellm.proxy._types import MCPTransport
from litellm.types.mcp import MCPAuth
from litellm.types.mcp_server.mcp_server_manager import MCPServer
global_mcp_server_manager.registry.clear()
oauth2_server = MCPServer(
server_id="basic_no_secret",
name="basic_no_secret",
server_name="basic_no_secret",
alias="basic_no_secret",
transport=MCPTransport.http,
auth_type=MCPAuth.oauth2,
client_id="basic-client",
client_secret=None,
authorization_url="https://idp.example.com/authorize",
token_url="https://idp.example.com/oauth2/token",
token_endpoint_auth_method="client_secret_basic",
)
global_mcp_server_manager.registry[oauth2_server.server_id] = oauth2_server
mock_request = MagicMock(spec=Request)
mock_request.base_url = "https://litellm-proxy.example.com/"
mock_request.headers = {}
with pytest.raises(HTTPException) as exc_info:
await token_endpoint(
request=mock_request,
grant_type="authorization_code",
code="auth-code",
redirect_uri="http://localhost/callback",
client_id="basic-client",
mcp_server_name="basic_no_secret",
client_secret=None,
code_verifier="verifier",
)
assert exc_info.value.status_code == 400

View file

@ -76,9 +76,7 @@ async def test_fetch_tools_from_passthrough_raises_on_upstream_401():
mock_client.list_tools = AsyncMock(side_effect=upstream_error)
with pytest.raises(MCPUpstreamAuthError) as exc_info:
await manager._fetch_tools_with_timeout(
mock_client, passthrough_server.name, server=passthrough_server
)
await manager._fetch_tools_with_timeout(mock_client, passthrough_server.name)
assert exc_info.value.status_code == 401
assert exc_info.value.www_authenticate == (
@ -113,9 +111,7 @@ async def test_fetch_tools_from_delegated_oauth2_raises_on_upstream_401():
mock_client.list_tools = AsyncMock(side_effect=upstream_error)
with pytest.raises(MCPUpstreamAuthError) as exc_info:
await manager._fetch_tools_with_timeout(
mock_client, delegated_server.name, server=delegated_server
)
await manager._fetch_tools_with_timeout(mock_client, delegated_server.name)
assert exc_info.value.status_code == 401
assert exc_info.value.www_authenticate == (
@ -126,7 +122,10 @@ async def test_fetch_tools_from_delegated_oauth2_raises_on_upstream_401():
@pytest.mark.asyncio
async def test_fetch_tools_from_client_credentials_oauth2_keeps_swallow_behavior():
async def test_fetch_tools_from_client_credentials_oauth2_surfaces_upstream_401():
"""The auth_type carve-out was removed: a client_credentials (M2M) server now
surfaces an upstream 401 as MCPUpstreamAuthError too, instead of swallowing it
to an empty list, so single-server routes can return a 401 challenge."""
manager = MCPServerManager()
m2m_server = MCPServer(
server_id="oauth-m2m",
@ -150,12 +149,12 @@ async def test_fetch_tools_from_client_credentials_oauth2_keeps_swallow_behavior
mock_client = MagicMock()
mock_client.list_tools = AsyncMock(side_effect=upstream_error)
tools = await manager._fetch_tools_with_timeout(
mock_client, m2m_server.name, server=m2m_server
)
with pytest.raises(MCPUpstreamAuthError) as exc_info:
await manager._fetch_tools_with_timeout(mock_client, m2m_server.name)
assert tools == []
mock_client.list_tools.assert_awaited_with(raise_on_error=False)
assert exc_info.value.status_code == 401
assert exc_info.value.server_name == "m2m_docs"
mock_client.list_tools.assert_awaited_with(raise_on_error=True)
@pytest.mark.asyncio
@ -176,9 +175,7 @@ async def test_fetch_tools_from_passthrough_returns_tools_on_success():
mock_client = MagicMock()
mock_client.list_tools = AsyncMock(return_value=[tool])
tools = await manager._fetch_tools_with_timeout(
mock_client, passthrough_server.name, server=passthrough_server
)
tools = await manager._fetch_tools_with_timeout(mock_client, passthrough_server.name)
assert tools == [tool]
@ -238,8 +235,11 @@ def test_to_http_exception_skips_challenge_for_non_401_status():
@pytest.mark.asyncio
async def test_fetch_tools_from_gateway_managed_swallows_errors():
"""Regression guard: non-pass-through servers keep returning [] on errors."""
async def test_fetch_tools_from_gateway_managed_surfaces_upstream_401():
"""An oauth2 server that is neither pass-through nor delegate now surfaces an
upstream 401 as MCPUpstreamAuthError as well; the auth_type carve-out that
swallowed it to [] was removed. A missing upstream WWW-Authenticate is carried
through as None (the single-server route fabricates one from the gateway URL)."""
manager = MCPServerManager()
oauth2_server = MCPServer(
server_id="o1",
@ -260,8 +260,12 @@ async def test_fetch_tools_from_gateway_managed_swallows_errors():
mock_client = MagicMock()
mock_client.list_tools = AsyncMock(side_effect=upstream_error)
tools = await manager._fetch_tools_with_timeout(
mock_client, oauth2_server.name, server=oauth2_server
)
assert tools == []
mock_client.list_tools.assert_awaited_with(raise_on_error=False)
with pytest.raises(MCPUpstreamAuthError) as exc_info:
await manager._fetch_tools_with_timeout(mock_client, oauth2_server.name)
assert exc_info.value.status_code == 401
assert exc_info.value.www_authenticate is None
assert exc_info.value.server_name == "keycloak_whoami"
mock_client.list_tools.assert_awaited_with(raise_on_error=True)

View file

@ -323,6 +323,28 @@ class TestMCPServerManager:
assert cost_info["tool_name_to_cost_per_query"]["geocode"] == 1e-3
assert isinstance(cost_info["tool_name_to_cost_per_query"]["geocode"], float)
@pytest.mark.asyncio
async def test_load_servers_from_config_sets_token_endpoint_auth_method(self):
"""token_endpoint_auth_method from config is carried onto the MCPServer (LIT-4091)."""
manager = MCPServerManager()
config = {
"basic_provider": {
"url": "https://example.com/mcp",
"transport": MCPTransport.http,
"token_endpoint_auth_method": "client_secret_basic",
},
"default_provider": {
"url": "https://example.com/mcp2",
"transport": MCPTransport.http,
},
}
await manager.load_servers_from_config(config)
by_name = {s.server_name: s for s in manager.config_mcp_servers.values()}
assert by_name["basic_provider"].token_endpoint_auth_method == "client_secret_basic"
assert by_name["default_provider"].token_endpoint_auth_method is None
def test_normalize_mcp_server_cost_info_preserves_float_values(self):
mcp_info = {
"server_name": "maps",
@ -3292,6 +3314,31 @@ class TestMCPServerTimestamps:
assert mcp_server.created_at == created
assert mcp_server.updated_at == updated
@pytest.mark.asyncio
async def test_build_mcp_server_from_table_reads_token_endpoint_auth_method(self):
"""token_endpoint_auth_method stored in the credentials JSON is loaded onto the MCPServer (LIT-4091)."""
manager = MCPServerManager()
basic_record = LiteLLM_MCPServerTable(
server_id="basic-db-1",
server_name="basic_db",
url="https://example.com/mcp",
transport=MCPTransport.http,
credentials={"token_endpoint_auth_method": "client_secret_basic"},
)
basic_server = await manager.build_mcp_server_from_table(basic_record, credentials_are_encrypted=False)
assert basic_server.token_endpoint_auth_method == "client_secret_basic"
default_record = LiteLLM_MCPServerTable(
server_id="default-db-1",
server_name="default_db",
url="https://example.com/mcp",
transport=MCPTransport.http,
credentials={},
)
default_server = await manager.build_mcp_server_from_table(default_record, credentials_are_encrypted=False)
assert default_server.token_endpoint_auth_method is None
def test_build_mcp_server_table_preserves_timestamps(self):
"""_build_mcp_server_table must use the MCPServer's stored timestamps, not datetime.now()."""
manager = MCPServerManager()
@ -4937,5 +4984,139 @@ class TestCreateMcpClientV2Graft:
assert client._get_auth_headers()["Authorization"] == "Bearer hook-jwt"
def _upstream_status_error(status_code: int, challenge: str) -> httpx.HTTPStatusError:
request = httpx.Request("POST", "https://upstream.example/mcp")
response = httpx.Response(
status_code,
headers={"WWW-Authenticate": challenge},
request=request,
)
return httpx.HTTPStatusError(
"upstream rejected token", request=request, response=response
)
class TestMCPToolsListAuthSurfacing:
"""Regression: MCP tools/list 401 auth failures must surface as MCPUpstreamAuthError.
Previously a missing/expired per-user OAuth token, or an upstream 401 for any
non-carveout auth_type, was swallowed to an empty tool list, so a single-server
client saw a 200 with no tools instead of a 401 challenge. The listing helpers
now raise MCPUpstreamAuthError on a 401 regardless of auth_type; the single-server
routes turn it into a 401 + WWW-Authenticate while the aggregator absorbs it to an
empty list. Only a 401 challenges; a 403 (forbidden) degrades like any other error.
"""
@pytest.mark.asyncio
async def test_fetch_tools_with_timeout_surfaces_upstream_401(self):
from litellm.proxy._experimental.mcp_server.exceptions import (
MCPUpstreamAuthError,
)
manager = MCPServerManager()
challenge = 'Bearer resource_metadata="https://upstream.example/.well-known/oauth-protected-resource"'
client = MagicMock()
client.list_tools = AsyncMock(side_effect=_upstream_status_error(401, challenge))
with pytest.raises(MCPUpstreamAuthError) as exc_info:
await manager._fetch_tools_with_timeout(client, "static-key-server")
assert exc_info.value.status_code == 401
assert exc_info.value.www_authenticate == challenge
assert exc_info.value.server_name == "static-key-server"
@pytest.mark.asyncio
async def test_fetch_tools_with_timeout_absorbs_upstream_403(self):
"""Only a 401 drives the re-auth challenge. A 403 (authenticated but
forbidden, e.g. insufficient scope) is not a re-auth signal, so even
with a WWW-Authenticate header it degrades to an empty list rather than
surfacing a challenge."""
manager = MCPServerManager()
challenge = 'Bearer error="insufficient_scope", scope="read:tools"'
client = MagicMock()
client.list_tools = AsyncMock(side_effect=_upstream_status_error(403, challenge))
assert await manager._fetch_tools_with_timeout(client, "forbidden-server") == []
@pytest.mark.asyncio
async def test_fetch_tools_with_timeout_returns_empty_on_non_auth_error(self):
manager = MCPServerManager()
client = MagicMock()
client.list_tools = AsyncMock(side_effect=RuntimeError("upstream 500"))
assert await manager._fetch_tools_with_timeout(client, "srv") == []
@pytest.mark.asyncio
async def test_get_tools_from_server_surfaces_unusable_user_token(self):
from litellm.proxy._experimental.mcp_server.exceptions import (
MCPUpstreamAuthError,
)
manager = MCPServerManager()
server = MCPServer(
server_id="oauth-srv", name="oauth-srv", transport=MCPTransport.http
)
challenge = 'Bearer resource_metadata="/.well-known/oauth-protected-resource/mcp/oauth-srv"'
manager._create_mcp_client = AsyncMock(
side_effect=HTTPException(
status_code=401,
detail="Unauthorized",
headers={"WWW-Authenticate": challenge},
)
)
with pytest.raises(MCPUpstreamAuthError) as exc_info:
await manager._get_tools_from_server(server)
assert exc_info.value.status_code == 401
assert exc_info.value.www_authenticate == challenge
assert exc_info.value.server_name == "oauth-srv"
@pytest.mark.asyncio
async def test_get_tools_from_server_absorbs_non_challenge_http_error(self):
manager = MCPServerManager()
server = MCPServer(
server_id="stdio-srv", name="stdio-srv", transport=MCPTransport.http
)
manager._create_mcp_client = AsyncMock(
side_effect=HTTPException(
status_code=403,
detail="MCP stdio command 'foo' is not in the allowlist",
)
)
assert await manager._get_tools_from_server(server) == []
@pytest.mark.asyncio
async def test_aggregate_list_tools_absorbs_unauthenticated_server(self):
from litellm.proxy._experimental.mcp_server.exceptions import (
MCPUpstreamAuthError,
)
manager = MCPServerManager()
good = MCPServer(server_id="good", name="good", transport=MCPTransport.http)
bad = MCPServer(server_id="bad", name="bad", transport=MCPTransport.http)
manager.get_allowed_mcp_servers = AsyncMock(return_value=["good", "bad"])
manager.get_mcp_server_by_id = MagicMock(
side_effect=lambda server_id: {"good": good, "bad": bad}.get(server_id)
)
good_tool = MCPTool(name="good-do_thing", description="do thing", inputSchema={})
async def fake_get_tools(server, **kwargs):
if server.server_id == "bad":
raise MCPUpstreamAuthError(
status_code=401,
www_authenticate='Bearer realm="x"',
server_name="bad",
)
return [good_tool]
manager._get_tools_from_server = fake_get_tools
result = await manager.list_tools()
assert [t.name for t in result] == ["good-do_thing"]
if __name__ == "__main__":
pytest.main([__file__])

View file

@ -173,3 +173,27 @@ async def test_non_dict_response_raises_value_error():
pytest.raises(ValueError, match="non-object JSON"),
):
await resolve_mcp_auth(server)
@pytest.mark.asyncio
async def test_client_credentials_uses_client_secret_basic_when_configured():
"""LIT-4091: a client_credentials server with token_endpoint_auth_method=client_secret_basic
authenticates via HTTP Basic and keeps the secret out of the form body."""
import base64
server = _server(server_id="srv-basic", token_endpoint_auth_method="client_secret_basic")
mock_client = AsyncMock()
mock_client.post.return_value = _token_response("m2m-basic")
with patch(
"litellm.proxy._experimental.mcp_server.oauth2_token_cache.get_async_httpx_client",
return_value=mock_client,
):
result = await resolve_mcp_auth(server)
assert result == "m2m-basic"
_, kwargs = mock_client.post.call_args
assert kwargs["headers"]["Authorization"] == "Basic " + base64.b64encode(b"cid:csec").decode()
assert "client_secret" not in kwargs["data"]
assert "client_id" not in kwargs["data"]
assert kwargs["data"]["grant_type"] == "client_credentials"

View file

@ -839,6 +839,76 @@ class TestListToolsRestAPI:
assert exc_info.value.status_code == upstream_status
assert exc_info.value.headers == {"www-authenticate": challenge}
async def test_aggregate_list_absorbs_one_server_auth_failure(self, monkeypatch):
"""The multi-server aggregate listing degrades a server whose upstream
rejects auth to an empty contribution and still returns the healthy
server's tools with a 200, rather than surfacing a 401."""
from litellm.proxy._experimental.mcp_server.exceptions import (
MCPUpstreamAuthError,
)
class StubServer:
def __init__(self, name):
self.alias = name
self.server_name = name
self.name = name
self.allowed_tools = None
self.mcp_info = {"server_name": name}
self.available_on_public_internet = True
good = StubServer("good")
bad = StubServer("bad")
async def fake_contexts(user_api_key_auth):
return [user_api_key_auth]
async def fake_get_allowed_mcp_servers(*args, **kwargs):
return ["good", "bad"]
async def fake_get_tools(server, *args, **kwargs):
if server.server_name == "bad":
raise MCPUpstreamAuthError(
status_code=401,
www_authenticate='Bearer realm="x"',
server_name="bad",
)
return ["good-tool"]
monkeypatch.setattr(
rest_endpoints,
"build_effective_auth_contexts",
fake_contexts,
raising=False,
)
monkeypatch.setattr(
rest_endpoints.global_mcp_server_manager,
"get_allowed_mcp_servers",
fake_get_allowed_mcp_servers,
raising=False,
)
monkeypatch.setattr(
rest_endpoints.global_mcp_server_manager,
"get_mcp_server_by_id",
lambda server_id: {"good": good, "bad": bad}.get(server_id),
raising=False,
)
monkeypatch.setattr(
rest_endpoints,
"_get_tools_for_single_server",
fake_get_tools,
raising=False,
)
request = _build_request(path="/mcp-rest/tools/list", method="GET")
result = await rest_endpoints.list_tool_rest_api(
request,
server_id=None,
user_api_key_dict=UserAPIKeyAuth(),
)
assert result["tools"] == ["good-tool"]
assert result["error"] is None
async def test_name_resolution_finds_server_by_uuid(self, monkeypatch):
"""When server_id is a name string, it should be resolved to its UUID
and used for the tools lookup when the UUID is in allowed_server_ids."""

View file

@ -2264,8 +2264,55 @@ class TestTemporaryMCPSessionEndpoints:
response_types=["code"],
token_endpoint_auth_method="client_secret_basic",
fallback_client_id="server-1",
persist_credentials=True,
)
@pytest.mark.asyncio
async def test_mcp_register_does_not_persist_for_non_admin(self):
"""A non-admin caller (who may have access to a real server) must not persist the DCR
result onto the shared server row. register_client_with_server is invoked with
persist_credentials=False, so user-side registration returns the DCR response without
writing shared client credentials. Only a full PROXY_ADMIN establishes the shared client."""
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
mcp_register,
)
request = MagicMock()
server = generate_mock_mcp_server_config_record(server_id="server-1")
register_response = {"client_id": "generated"}
request_body = {
"client_name": "LiteLLM",
"grant_types": ["authorization_code"],
"response_types": ["code"],
"token_endpoint_auth_method": "client_secret_basic",
}
non_admin_auth = generate_mock_user_api_key_auth(
user_role=LitellmUserRoles.INTERNAL_USER,
)
with (
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints._get_cached_temporary_mcp_server_or_404",
return_value=server,
),
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints._read_request_body",
AsyncMock(return_value=request_body),
),
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.register_client_with_server",
AsyncMock(return_value=register_response),
) as register_mock,
):
result = await mcp_register(
request=request,
server_id="server-1",
user_api_key_dict=non_admin_auth,
)
assert result is register_response
assert register_mock.await_args.kwargs["persist_credentials"] is False
@pytest.mark.asyncio
async def test_get_cached_temporary_mcp_server_falls_back_to_redis(self):
from litellm.proxy.management_endpoints.mcp_management_endpoints import (

View file

@ -817,6 +817,7 @@ def test_aaamodel_prices_and_context_window_json_is_valid():
"supports_image_input": {"type": "boolean"},
"supports_nova_canvas_image_edit": {"type": "boolean"},
"supports_parallel_function_calling": {"type": "boolean"},
"supports_parallel_tool_use_config": {"type": "boolean"},
"supports_pdf_input": {"type": "boolean"},
"supports_prompt_caching": {"type": "boolean"},
"supports_response_schema": {"type": "boolean"},
@ -842,6 +843,7 @@ def test_aaamodel_prices_and_context_window_json_is_valid():
"type": "string",
"enum": ["low", "medium", "high", "max", "xhigh"],
},
"bedrock_converse_supports_strict_tools": {"type": "boolean"},
"tpm": {"type": "number"},
"provider_specific_entry": {"type": "object"},
"supported_endpoints": {

View file

@ -7,6 +7,7 @@ import CreateMCPServer from "./create_mcp_server";
vi.mock("../networking", () => ({
createMCPServer: vi.fn(),
fetchOpenAPIRegistry: vi.fn().mockResolvedValue({ apis: [] }),
registerMCPServer: vi.fn(),
storeMCPOAuthUserCredential: vi.fn().mockResolvedValue({}),
testMCPToolsListRequest: vi.fn().mockResolvedValue({ tools: [], error: null }),
@ -16,15 +17,26 @@ vi.mock("@/utils/mcpTokenStore", () => ({
setToken: vi.fn(),
}));
vi.mock("./OpenAPIQuickPicker", () => ({
default: () => null,
}));
// Mutable holder so individual tests can simulate "Authorize & Fetch" having
// produced a token before submit, and inspect the reset wiring.
const oauthHook = vi.hoisted(() => ({
tokenResponse: null as Record<string, unknown> | null,
reset: vi.fn(),
onTokenReceived: null as ((token: Record<string, unknown> | null) => void) | null,
onTokenReceived: null as
| ((token: Record<string, unknown> | null, registeredClient?: { clientId?: string; clientSecret?: string }) => void)
| null,
}));
vi.mock("@/hooks/useMcpOAuthFlow", () => ({
useMcpOAuthFlow: (opts: { onTokenReceived: (token: Record<string, unknown> | null) => void }) => {
useMcpOAuthFlow: (opts: {
onTokenReceived: (
token: Record<string, unknown> | null,
registeredClient?: { clientId?: string; clientSecret?: string },
) => void;
}) => {
oauthHook.onTokenReceived = opts.onTokenReceived;
return {
startOAuthFlow: vi.fn(),
@ -495,6 +507,170 @@ describe("CreateMCPServer", () => {
expect(payload.token_validation).toEqual({ organization: "my-org", "team.id": "42" });
});
it("invalidates the DCR client and OAuth flow when the MCP URL changes after Authorize & Fetch", async () => {
await setupOAuthInteractive();
const nameInput = document.getElementById("server_name") as HTMLInputElement;
await act(async () => {
fireEvent.change(nameInput, { target: { value: "Url_Change_Server" } });
});
const urlInput = screen.getByPlaceholderText("https://your-mcp-server.com");
await act(async () => {
fireEvent.change(urlInput, { target: { value: "https://a.example.com/mcp" } });
});
act(() => {
oauthHook.onTokenReceived?.({ access_token: "tok-a" }, { clientId: "client-a", clientSecret: "secret-a" });
});
oauthHook.reset.mockClear();
await act(async () => {
fireEvent.change(urlInput, { target: { value: "https://b.example.com/mcp" } });
});
await waitFor(() => expect(oauthHook.reset).toHaveBeenCalled());
vi.mocked(networking.createMCPServer).mockResolvedValue({
server_id: "new-server-oauth",
server_name: "Url_Change_Server",
alias: "Url_Change_Server",
url: "https://b.example.com/mcp",
transport: "http",
auth_type: "oauth2",
created_at: "2024-01-01T00:00:00Z",
created_by: "user-1",
updated_at: "2024-01-01T00:00:00Z",
updated_by: "user-1",
});
const submitButton = screen.getByRole("button", { name: "Add MCP Server" });
await act(async () => {
fireEvent.click(submitButton);
});
await waitFor(() => expect(networking.createMCPServer).toHaveBeenCalledTimes(1));
const [, payload] = vi.mocked(networking.createMCPServer).mock.calls[0];
expect(payload.credentials?.client_id).toBeUndefined();
expect(payload.credentials?.client_secret).toBeUndefined();
});
it("invalidates the DCR client and OAuth flow when the OpenAPI spec URL changes after Authorize & Fetch", async () => {
render(<CreateMCPServer {...defaultProps} />);
await selectAntOption("Transport Type", "OpenAPI Spec");
await waitFor(() => {
expect(screen.getByPlaceholderText("https://petstore3.swagger.io/api/v3/openapi.json")).toBeInTheDocument();
});
await selectAntOption("Authentication", "OAuth");
await waitFor(() => {
expect(screen.getByText("OAuth Flow Type")).toBeInTheDocument();
});
const nameInput = document.getElementById("server_name") as HTMLInputElement;
await act(async () => {
fireEvent.change(nameInput, { target: { value: "OpenAPI_Server" } });
});
const specInput = screen.getByPlaceholderText("https://petstore3.swagger.io/api/v3/openapi.json");
await act(async () => {
fireEvent.change(specInput, { target: { value: "https://a.example.com/openapi.json" } });
});
act(() => {
oauthHook.onTokenReceived?.({ access_token: "tok-a" }, { clientId: "client-a", clientSecret: "secret-a" });
});
oauthHook.reset.mockClear();
await act(async () => {
fireEvent.change(specInput, { target: { value: "https://b.example.com/openapi.json" } });
});
await waitFor(() => expect(oauthHook.reset).toHaveBeenCalled());
vi.mocked(networking.createMCPServer).mockResolvedValue({
server_id: "new-openapi-server",
server_name: "OpenAPI_Server",
alias: "OpenAPI_Server",
url: "https://b.example.com/openapi.json",
transport: "http",
auth_type: "oauth2",
created_at: "2024-01-01T00:00:00Z",
created_by: "user-1",
updated_at: "2024-01-01T00:00:00Z",
updated_by: "user-1",
});
const submitButton = screen.getByRole("button", { name: "Add MCP Server" });
await act(async () => {
fireEvent.click(submitButton);
});
await waitFor(() => expect(networking.createMCPServer).toHaveBeenCalledTimes(1));
const [, payload] = vi.mocked(networking.createMCPServer).mock.calls[0];
expect(payload.spec_path).toBe("https://b.example.com/openapi.json");
expect(payload.credentials?.client_id).toBeUndefined();
expect(payload.credentials?.client_secret).toBeUndefined();
});
it("invalidates the DCR client and OAuth flow when the transport changes after Authorize & Fetch", async () => {
render(<CreateMCPServer {...defaultProps} />);
await selectAntOption("Transport Type", "OpenAPI Spec");
await waitFor(() => {
expect(screen.getByPlaceholderText("https://petstore3.swagger.io/api/v3/openapi.json")).toBeInTheDocument();
});
await selectAntOption("Authentication", "OAuth");
await waitFor(() => {
expect(screen.getByText("OAuth Flow Type")).toBeInTheDocument();
});
const nameInput = document.getElementById("server_name") as HTMLInputElement;
await act(async () => {
fireEvent.change(nameInput, { target: { value: "Transport_Change_Server" } });
});
const specInput = screen.getByPlaceholderText("https://petstore3.swagger.io/api/v3/openapi.json");
await act(async () => {
fireEvent.change(specInput, { target: { value: "https://same.example.com/spec-or-mcp" } });
});
act(() => {
oauthHook.onTokenReceived?.({ access_token: "tok-a" }, { clientId: "client-a", clientSecret: "secret-a" });
});
oauthHook.reset.mockClear();
await selectAntOption("Transport Type", "Streamable HTTP");
await waitFor(() => {
expect(screen.getByPlaceholderText("https://your-mcp-server.com")).toBeInTheDocument();
});
const urlInput = screen.getByPlaceholderText("https://your-mcp-server.com");
await act(async () => {
fireEvent.change(urlInput, { target: { value: "https://same.example.com/spec-or-mcp" } });
});
await waitFor(() => expect(oauthHook.reset).toHaveBeenCalled());
vi.mocked(networking.createMCPServer).mockResolvedValue({
server_id: "new-transport-server",
server_name: "Transport_Change_Server",
alias: "Transport_Change_Server",
url: "https://same.example.com/spec-or-mcp",
transport: "http",
auth_type: "oauth2",
created_at: "2024-01-01T00:00:00Z",
created_by: "user-1",
updated_at: "2024-01-01T00:00:00Z",
updated_by: "user-1",
});
const submitButton = screen.getByRole("button", { name: "Add MCP Server" });
await act(async () => {
fireEvent.click(submitButton);
});
await waitFor(() => expect(networking.createMCPServer).toHaveBeenCalledTimes(1));
const [, payload] = vi.mocked(networking.createMCPServer).mock.calls[0];
expect(payload.url).toBe("https://same.example.com/spec-or-mcp");
expect(payload.credentials?.client_id).toBeUndefined();
expect(payload.credentials?.client_secret).toBeUndefined();
});
it("omits token_validation from payload when token_validation_json is empty", async () => {
vi.mocked(networking.createMCPServer).mockResolvedValue({
server_id: "new-server-oauth",
@ -667,11 +843,13 @@ describe("CreateMCPServer", () => {
// Reopen for a brand-new server and enter a different URL without re-authorizing.
rerender(<CreateMCPServer {...defaultProps} isModalVisible={true} />);
const reopenedUrlInput = screen.getByPlaceholderText("https://your-mcp-server.com");
oauthHook.reset.mockClear();
await act(async () => {
fireEvent.change(reopenedUrlInput, { target: { value: "https://server-b.example.com/mcp" } });
});
// The previous server's token must never be replayed for the new session.
expect(oauthHook.reset).not.toHaveBeenCalled();
expect(usedToken("stale-token-A")).toBe(false);
});

View file

@ -89,6 +89,7 @@ const CreateMCPServer: React.FC<CreateMCPServerProps> = ({
const [oauthAccessToken, setOauthAccessToken] = useState<string | null>(null);
const [logoUrl, setLogoUrl] = useState<string | undefined>(undefined);
const [oauthDocsUrl, setOauthDocsUrl] = useState<string | null>(null);
const [authorizedUrl, setAuthorizedUrl] = useState<string | undefined>(undefined);
// Single hook call shared by MCPConnectionStatus and MCPToolConfiguration to avoid duplicate requests.
const { tools, isLoadingTools, toolsError, toolsErrorStackTrace, canFetchTools, fetchTools, clearTools } =
@ -105,6 +106,12 @@ const CreateMCPServer: React.FC<CreateMCPServerProps> = ({
const isAwsSigV4AuthType = authType === AUTH_TYPE.AWS_SIGV4;
const isM2MFlow = isOAuthAuthType && formValues.oauth_flow_type === OAUTH_FLOW.M2M;
const getOAuthAuthorizationTarget = (values: Record<string, unknown>): string | undefined => {
const transport = values.transport || transportType;
const target = transport === TRANSPORT.OPENAPI ? values.spec_path : values.url;
return typeof target === "string" ? target : undefined;
};
const persistCreateUiState = () => {
if (typeof window === "undefined") {
return;
@ -171,7 +178,7 @@ const CreateMCPServer: React.FC<CreateMCPServerProps> = ({
env: values.env,
};
},
onTokenReceived: (token) => {
onTokenReceived: (token, registeredClient) => {
setOauthAccessToken(token?.access_token ?? null);
if (token?.access_token) {
@ -180,9 +187,12 @@ const CreateMCPServer: React.FC<CreateMCPServerProps> = ({
...(token.refresh_token && { refresh_token: token.refresh_token }),
...(token.expires_in && { expires_in: token.expires_in }),
...(token.scope && { scope: token.scope }),
...(registeredClient?.clientId && { client_id: registeredClient.clientId }),
...(registeredClient?.clientSecret && { client_secret: registeredClient.clientSecret }),
};
form.setFieldsValue({ credentials });
setAuthorizedUrl(getOAuthAuthorizationTarget(form.getFieldsValue(true)));
NotificationsManager.success(
"OAuth authorization successful! Please click 'Create MCP Server' to save the configuration.",
@ -193,6 +203,15 @@ const CreateMCPServer: React.FC<CreateMCPServerProps> = ({
flowSource: "create",
});
const clearAuthorizedOAuthState = (values: Record<string, unknown>) => {
form.resetFields(["credentials", "authorization_url", "token_url", "registration_url"]);
form.setFieldsValue(values);
setOauthAccessToken(null);
clearTools();
resetOAuthFlow();
setAuthorizedUrl(undefined);
};
React.useEffect(() => {
if (typeof window === "undefined") {
return;
@ -506,12 +525,28 @@ const CreateMCPServer: React.FC<CreateMCPServerProps> = ({
const handleTransportChange = (value: string) => {
setTransportType(value);
// Clear fields that are not relevant for the selected transport
if (value === "stdio") {
form.setFieldsValue({ url: undefined, spec_path: undefined, auth_type: undefined, credentials: undefined });
} else if (value === TRANSPORT.OPENAPI) {
form.setFieldsValue({ url: undefined, command: undefined, args: undefined, env: undefined });
const transportValues =
value === "stdio"
? { url: undefined, spec_path: undefined, auth_type: undefined, credentials: undefined }
: value === TRANSPORT.OPENAPI
? { url: undefined, command: undefined, args: undefined, env: undefined }
: { spec_path: undefined, command: undefined, args: undefined, env: undefined };
const nextValues =
authorizedUrl === undefined
? transportValues
: {
...transportValues,
credentials: undefined,
authorization_url: undefined,
token_url: undefined,
registration_url: undefined,
};
if (authorizedUrl !== undefined) {
clearAuthorizedOAuthState(nextValues);
} else {
form.setFieldsValue({ spec_path: undefined, command: undefined, args: undefined, env: undefined });
form.setFieldsValue(nextValues);
}
};
@ -567,11 +602,32 @@ const CreateMCPServer: React.FC<CreateMCPServerProps> = ({
setOauthAccessToken(null);
clearTools();
resetOAuthFlow();
setAuthorizedUrl(undefined);
}
}, [isModalVisible, form, clearTools, resetOAuthFlow]);
const isAdmin = isAdminRole(userRole);
const handleFormValuesChange = (changedValues: Record<string, unknown>, allValues: Record<string, unknown>) => {
const changedAuthorizationTarget = "url" in changedValues || "spec_path" in changedValues;
if (
changedAuthorizationTarget &&
authorizedUrl !== undefined &&
getOAuthAuthorizationTarget(allValues) !== authorizedUrl
) {
const invalidated = {
credentials: undefined,
authorization_url: changedValues.authorization_url,
token_url: changedValues.token_url,
registration_url: changedValues.registration_url,
};
clearAuthorizedOAuthState(invalidated);
setFormValues({ ...allValues, ...invalidated });
return;
}
setFormValues(allValues);
};
// rendering
return (
<Modal
@ -616,7 +672,7 @@ const CreateMCPServer: React.FC<CreateMCPServerProps> = ({
<Form
form={form}
onFinish={handleCreate}
onValuesChange={(_, allValues) => setFormValues(allValues)}
onValuesChange={handleFormValuesChange}
layout="vertical"
className="space-y-6"
>
@ -736,7 +792,9 @@ const CreateMCPServer: React.FC<CreateMCPServerProps> = ({
<OpenAPIFormSection
form={form}
accessToken={isModalVisible ? accessToken : null}
onValuesChange={(updates) => setFormValues((prev) => ({ ...prev, ...updates }))}
onValuesChange={(updates) =>
handleFormValuesChange(updates, { ...form.getFieldsValue(true), ...updates })
}
onKeyToolsChange={setKeyTools}
onLogoUrlChange={setLogoUrl}
onOAuthDocsUrlChange={setOauthDocsUrl}

View file

@ -131,6 +131,25 @@ describe("MCPToolsViewer auth gate routing", () => {
expect(screen.queryByText(GATE_TEXT)).not.toBeInTheDocument();
});
it("gates an OBO server whose stored token is expired and the list call 401s (refresh could not mint a token)", async () => {
// has_credential=true but the list call 401s: the server-side refresh could not
// produce a valid token (e.g. expired with no usable refresh token), so the user
// must reauthorize instead of seeing a dead empty list.
vi.mocked(getMCPOAuthUserCredentialStatus).mockResolvedValue(
credStatus({ has_credential: true, is_expired: true }),
);
vi.mocked(listMCPTools).mockResolvedValue({
tools: [],
error: "unauthorized",
status: 401,
} as unknown as Awaited<ReturnType<typeof listMCPTools>>);
renderViewer({ oauth2_flow: null, delegate_auth_to_upstream: false });
expect(await screen.findByText(GATE_TEXT)).toBeInTheDocument();
expect(screen.getByRole("button", { name: "Authorize" })).toBeInTheDocument();
});
it("does not gate an M2M server; lists with the LiteLLM key", async () => {
renderViewer({ oauth2_flow: "client_credentials", delegate_auth_to_upstream: false });

View file

@ -238,9 +238,14 @@ const MCPToolsViewer = ({
const toolsData = mcpToolsResponse?.tools || [];
const oboToolsError = mcpToolsError as (Error & { status?: number; response?: { status?: number } }) | null;
const oboTokenRejected = isObo && (oboToolsError?.status ?? oboToolsError?.response?.status) === 401;
// An auth gate replaces the tool list when the user must authenticate first:
// passthrough needs a browser token, OBO needs a stored DB credential.
const authGateActive = (isPassthrough && !oauthToken) || oboNeedsAuth;
// passthrough needs a browser token; OBO needs a stored DB credential or a
// still-valid one — a 401 from the list call means the backend has none even
// after attempting a refresh, so re-authorization is required.
const authGateActive = (isPassthrough && !oauthToken) || oboNeedsAuth || oboTokenRejected;
// Treat OBO credential-status loading as "tools loading" so the empty state
// doesn't flash before we know whether the user needs to authorize.
const toolsAreaLoading = isLoadingTools || oboStatusLoading;
@ -364,10 +369,12 @@ const MCPToolsViewer = ({
</div>
)}
{/* OBO auth gate only when no credential row exists for this user.
An existing-but-expired token is refreshed server-side on the
list call, so the gate never appears for a stored credential. */}
{oboNeedsAuth && (
{/* OBO auth gate shown when there is no credential row for this
user, or when the list call returns 401 (no valid token and the
server-side refresh could not mint one, e.g. an expired token
with no usable refresh token). A refreshable token is refreshed
on the list call and never trips this gate. */}
{(oboNeedsAuth || oboTokenRejected) && (
<div className="p-4 text-center bg-white border border-gray-200 rounded-lg">
<LockOutlined className="text-2xl text-gray-400 mb-2" />
<p className="text-xs font-medium text-gray-700 mb-1">Authentication required</p>

View file

@ -77,7 +77,7 @@ describe("useMcpOAuthFlow reset", () => {
await waitFor(() => expect(result.current.status).toBe("success"));
expect(result.current.tokenResponse).toEqual(token);
expect(onTokenReceived).toHaveBeenCalledWith(token);
expect(onTokenReceived).toHaveBeenCalledWith(token, expect.objectContaining({ clientId: "client-1" }));
act(() => {
result.current.reset();
@ -88,6 +88,34 @@ describe("useMcpOAuthFlow reset", () => {
expect(result.current.error).toBeNull();
});
it("ignores an in-flight exchange result after reset", async () => {
const token = { access_token: "stale-token" };
let resolveExchange: (value: typeof token) => void = () => undefined;
const exchangePromise = new Promise<typeof token>((resolve) => {
resolveExchange = resolve;
});
vi.mocked(networking.exchangeMcpOAuthToken).mockReturnValueOnce(exchangePromise);
seedCompletedRedirect();
const onTokenReceived = vi.fn();
const { result } = renderFlow(onTokenReceived);
await waitFor(() => expect(result.current.status).toBe("exchanging"));
act(() => {
result.current.reset();
});
await act(async () => {
resolveExchange(token);
await exchangePromise;
});
expect(onTokenReceived).not.toHaveBeenCalled();
expect(result.current.status).toBe("idle");
expect(result.current.tokenResponse).toBeNull();
});
it("clears the in-flight guard so a callback after a mid-exchange close is not swallowed", async () => {
// First exchange hangs, mimicking the modal being closed while the token
// endpoint is still in flight. processingRef is left true at that point.
@ -112,6 +140,92 @@ describe("useMcpOAuthFlow reset", () => {
const onTokenReceived2 = vi.fn();
rerender({ onTokenReceived: onTokenReceived2 });
await waitFor(() => expect(onTokenReceived2).toHaveBeenCalledWith(token));
await waitFor(() =>
expect(onTokenReceived2).toHaveBeenCalledWith(token, expect.objectContaining({ clientId: "client-1" })),
);
});
it("passes the DCR-registered client_id and client_secret to onTokenReceived so the created server persists them", async () => {
const token = { access_token: "tok-xyz", refresh_token: "ref-xyz", expires_in: 3600 };
vi.mocked(networking.exchangeMcpOAuthToken).mockResolvedValue(token);
setSecureItem(RESULT_KEY, JSON.stringify({ state: "state-1", code: "code-1" }));
setSecureItem(
FLOW_STATE_KEY,
JSON.stringify({
state: "state-1",
codeVerifier: "verifier-1",
serverId: "server-1",
clientId: "dcr-client-xyz",
clientSecret: "dcr-secret-abc",
redirectUri: "https://app.example.com/ui/mcp/oauth/callback",
flowSource: "create",
}),
);
const onTokenReceived = vi.fn();
const { result } = renderFlow(onTokenReceived);
await waitFor(() => expect(result.current.status).toBe("success"));
expect(onTokenReceived).toHaveBeenCalledWith(token, {
clientId: "dcr-client-xyz",
clientSecret: "dcr-secret-abc",
});
});
it("reuses an existing client_id and does not register a new client (second Authorize & Fetch, same server)", async () => {
vi.mocked(networking.cacheTemporaryMcpServer).mockResolvedValue({ server_id: "server-1" });
vi.mocked(networking.buildMcpOAuthAuthorizeUrl).mockReturnValue("https://idp.example.com/authorize");
const { result } = renderHook(() =>
useMcpOAuthFlow({
accessToken: "admin-token",
getCredentials: () => ({ client_id: "existing-client" }),
getTemporaryPayload: () => ({
url: "https://server-1.example.com/mcp",
transport: "http",
credentials: { client_id: "existing-client" },
}),
onTokenReceived: vi.fn(),
flowSource: "create",
}),
);
await act(async () => {
await result.current.startOAuthFlow();
});
expect(networking.registerMcpOAuthClient).not.toHaveBeenCalled();
expect(networking.buildMcpOAuthAuthorizeUrl).toHaveBeenCalledWith(
expect.objectContaining({ clientId: "existing-client" }),
);
});
it("registers a fresh client when no client_id is present (new URL after the derived client is cleared)", async () => {
vi.mocked(networking.cacheTemporaryMcpServer).mockResolvedValue({ server_id: "server-2" });
vi.mocked(networking.registerMcpOAuthClient).mockResolvedValue({ client_id: "fresh-client" });
vi.mocked(networking.buildMcpOAuthAuthorizeUrl).mockReturnValue("https://idp.example.com/authorize");
const { result } = renderHook(() =>
useMcpOAuthFlow({
accessToken: "admin-token",
getCredentials: () => ({}),
getTemporaryPayload: () => ({
url: "https://server-2.example.com/mcp",
transport: "http",
credentials: {},
}),
onTokenReceived: vi.fn(),
flowSource: "create",
}),
);
await act(async () => {
await result.current.startOAuthFlow();
});
expect(networking.registerMcpOAuthClient).toHaveBeenCalledTimes(1);
expect(networking.buildMcpOAuthAuthorizeUrl).toHaveBeenCalledWith(
expect.objectContaining({ clientId: "fresh-client" }),
);
});
});

View file

@ -26,7 +26,10 @@ interface UseMcpOAuthFlowOptions {
}
| undefined;
getTemporaryPayload: () => Record<string, any> | null;
onTokenReceived: (tokenResponse: Record<string, any>) => void;
onTokenReceived: (
tokenResponse: Record<string, any>,
registeredClient?: { clientId?: string; clientSecret?: string },
) => void;
onBeforeRedirect?: () => void;
// Distinguishes which form started the flow (e.g. "create" vs "edit"). Both forms
// mount this hook with shared storage keys, so the return handler only processes a
@ -55,6 +58,7 @@ export const useMcpOAuthFlow = ({
const [error, setError] = useState<string | null>(null);
const [tokenResponse, setTokenResponse] = useState<Record<string, any> | null>(null);
const processingRef = useRef(false);
const resetVersionRef = useRef(0);
const FLOW_STATE_KEY = "litellm-mcp-oauth-flow-state";
const RESULT_KEY = "litellm-mcp-oauth-result";
@ -144,9 +148,7 @@ export const useMcpOAuthFlow = ({
}
let registeredClient: { clientId?: string; clientSecret?: string } = {};
const hasPreconfiguredCredentials = Boolean(
temporaryPayload.credentials?.client_id && temporaryPayload.credentials?.client_secret,
);
const hasPreconfiguredCredentials = Boolean(temporaryPayload.credentials?.client_id);
if (!hasPreconfiguredCredentials) {
const registration = await registerMcpOAuthClient(accessToken, serverId, {
@ -286,6 +288,8 @@ export const useMcpOAuthFlow = ({
}
}
const resetVersion = resetVersionRef.current;
try {
if (!flowState || !flowState.state || !flowState.codeVerifier || !flowState.serverId) {
throw new Error(
@ -314,22 +318,31 @@ export const useMcpOAuthFlow = ({
accessToken,
});
onTokenReceived(token);
if (resetVersion !== resetVersionRef.current) {
return;
}
onTokenReceived(token, { clientId: flowState.clientId, clientSecret: flowState.clientSecret });
setTokenResponse(token);
setStatus("success");
setError(null);
NotificationsManager.success("OAuth token retrieved successfully");
} catch (err) {
if (resetVersion !== resetVersionRef.current) {
return;
}
const message = extractErrorMessage(err);
setError(message);
setStatus("error");
NotificationsManager.error(message);
} finally {
clearStoredFlow();
// Reset processing flag after a delay to allow UI updates
setTimeout(() => {
processingRef.current = false;
}, 1000);
if (resetVersion === resetVersionRef.current) {
clearStoredFlow();
// Reset processing flag after a delay to allow UI updates
setTimeout(() => {
processingRef.current = false;
}, 1000);
}
}
}, [onTokenReceived]);
@ -338,6 +351,7 @@ export const useMcpOAuthFlow = ({
}, [resumeOAuthFlow]);
const reset = useCallback(() => {
resetVersionRef.current += 1;
setStatus("idle");
setError(null);
setTokenResponse(null);