Merge remote-tracking branch 'origin/litellm_internal_staging' into litellm_lit_5546_count_tokens_offload

This commit is contained in:
mateo-berri 2026-09-07 18:11:29 -07:00
commit 486af36963
29 changed files with 1677 additions and 134 deletions

View file

@ -16,6 +16,8 @@ from collections.abc import Iterable, Mapping, Sequence
from typing import TYPE_CHECKING, Any, Final, cast
from urllib.parse import urlparse
from pydantic import TypeAdapter, ValidationError
from litellm._logging import verbose_logger
from litellm.integrations.custom_logger import CustomLogger
from litellm.integrations.custom_prompt_management import CustomPromptManagement
@ -23,6 +25,7 @@ from litellm.integrations.prompt_management_base import PromptManagementClient
from litellm.litellm_core_utils.prompt_templates.common_utils import (
with_prompt_cache_breakpoint,
)
from litellm.llms.anthropic.common_utils import is_claude_code_one_shot_subagent_request
from litellm.types.integrations.anthropic_cache_control_hook import (
GATEWAY_INJECTED_CACHE_METADATA_KEY,
GATEWAY_INJECTED_FOR_EVERY_DEPLOYMENT,
@ -62,10 +65,26 @@ OPENAI_PROMPT_CACHE_BREAKPOINT_BLOCK_TYPES: Final = frozenset(
)
OPENAI_API_HOST: Final = "api.openai.com"
OPENAI_API_BASE_ENV_VARS: Final = ("OPENAI_BASE_URL", "OPENAI_API_BASE")
_OBJECT_MAPPING_ADAPTER: Final = TypeAdapter(dict[object, object])
_OBJECT_LIST_ADAPTER: Final = TypeAdapter(list[object])
AllToolParamValues = ChatCompletionToolParam | AllAnthropicToolsValues
def _validated_object_mapping(value: object) -> dict[object, object] | None:
try:
return _OBJECT_MAPPING_ADAPTER.validate_python(value)
except ValidationError:
return None
def _validated_object_list(value: object) -> list[object] | None:
try:
return _OBJECT_LIST_ADAPTER.validate_python(value)
except ValidationError:
return None
def supports_openai_prompt_cache_breakpoint(model: str) -> bool:
model_map_flag: Final = _model_map_prompt_cache_breakpoint_flag(model)
if model_map_flag is not None:
@ -114,6 +133,36 @@ CARRY_UNMATCHED_MESSAGE_POINTS: Final = "_litellm_carry_unmatched_cache_control_
class AnthropicCacheControlHook(CustomPromptManagement):
@staticmethod
def _request_value(request_kwargs: object, key: str) -> object:
request_mapping: Final = _validated_object_mapping(request_kwargs)
if request_mapping is None:
return None
return request_mapping.get(key)
@staticmethod
def _request_user_agent(request_kwargs: object) -> str | None:
proxy_server_request: Final = AnthropicCacheControlHook._request_value(request_kwargs, "proxy_server_request")
proxy_server_request_mapping: Final = _validated_object_mapping(proxy_server_request)
if proxy_server_request_mapping is None:
return None
headers: Final = proxy_server_request_mapping.get("headers")
headers_mapping: Final = _validated_object_mapping(headers)
if headers_mapping is None:
return None
user_agent: Final = next(
(value for key, value in headers_mapping.items() if isinstance(key, str) and key.lower() == "user-agent"),
None,
)
return user_agent if isinstance(user_agent, str) else None
@staticmethod
def _request_system(request_kwargs: object) -> str | list[object] | None:
system: Final = AnthropicCacheControlHook._request_value(request_kwargs, "system")
if isinstance(system, str):
return system
return _validated_object_list(system)
def get_chat_completion_prompt(
self,
model: str,
@ -520,12 +569,13 @@ class AnthropicCacheControlHook(CustomPromptManagement):
points: Sequence[CacheControlInjectionPoint],
messages: list[AllMessageValues],
tools: list[object] | None,
cache_control: object,
model: str,
custom_llm_provider: str | None,
api_base: object,
prompt_cache_options: object,
) -> Sequence[Mapping[str, object]] | None:
if AnthropicCacheControlHook._should_stand_down(points, messages, None, tools):
if AnthropicCacheControlHook._should_stand_down(points, messages, None, tools, cache_control):
return None
return AnthropicCacheControlHook._stamped_with_dialect(
points, model, custom_llm_provider, api_base, prompt_cache_options
@ -561,6 +611,7 @@ class AnthropicCacheControlHook(CustomPromptManagement):
messages: list[AllMessageValues],
system: str | list | None,
tools: list | None,
cache_control: object = None,
) -> bool:
"""Whether configured injection points must yield to client-set cache_control.
@ -573,13 +624,14 @@ class AnthropicCacheControlHook(CustomPromptManagement):
"""
if all(point.get("_litellm_judged") for point in points):
return False
return AnthropicCacheControlHook._request_has_cache_control(messages, system, tools)
return AnthropicCacheControlHook._request_has_cache_control(messages, system, tools, cache_control)
@staticmethod
def _request_has_cache_control(
messages: list[AllMessageValues],
system: str | list | None,
tools: list | None = None,
cache_control: object = None,
) -> bool:
"""Return True if the request already carries any client-supplied cache_control.
@ -591,6 +643,8 @@ class AnthropicCacheControlHook(CustomPromptManagement):
carry the mark either at the top level (Anthropic shape) or nested under
``function`` (OpenAI shape); the Anthropic chat transform accepts both.
"""
if cache_control is not None:
return True
if AnthropicCacheControlHook.count_request_cache_breakpoints(messages, system) > 0:
return True
if tools is not None:
@ -612,6 +666,8 @@ class AnthropicCacheControlHook(CustomPromptManagement):
custom_llm_provider: str | None,
tools: list | None = None,
enable_prompt_caching: bool | None = None,
cache_control: object = None,
request_kwargs: object = None,
) -> list[CacheControlInjectionPoint]:
"""Default breakpoints when ``litellm.enable_anthropic_prompt_caching`` is on.
@ -649,7 +705,12 @@ class AnthropicCacheControlHook(CustomPromptManagement):
if not supports_prompt_caching(model=model, custom_llm_provider=provider):
return []
if AnthropicCacheControlHook._request_has_cache_control(messages, system, tools):
if AnthropicCacheControlHook._request_has_cache_control(messages, system, tools, cache_control):
return []
if is_claude_code_one_shot_subagent_request(
messages, system, tools, AnthropicCacheControlHook._request_user_agent(request_kwargs)
):
return []
control: Final = AnthropicCacheControlHook._default_control()
@ -665,6 +726,7 @@ class AnthropicCacheControlHook(CustomPromptManagement):
models: Iterable[str],
tools: list[AllToolParamValues] | None = None,
enable_prompt_caching: bool | None = None,
request_kwargs: object = None,
) -> list[AllMessageValues]:
"""Return the messages auto prompt caching will send, default breakpoints included.
@ -681,11 +743,13 @@ class AnthropicCacheControlHook(CustomPromptManagement):
for candidate in (
AnthropicCacheControlHook.get_default_injection_points(
messages=messages,
system=None,
model=model,
custom_llm_provider=None,
tools=tools,
enable_prompt_caching=enable_prompt_caching,
system=AnthropicCacheControlHook._request_system(request_kwargs),
cache_control=AnthropicCacheControlHook._request_value(request_kwargs, "cache_control"),
request_kwargs=request_kwargs,
)
for model in models
)
@ -730,6 +794,7 @@ class AnthropicCacheControlHook(CustomPromptManagement):
non_default_params["cache_control_injection_points"],
messages,
tools,
non_default_params.get("cache_control"),
model,
custom_llm_provider,
api_base,
@ -747,6 +812,8 @@ class AnthropicCacheControlHook(CustomPromptManagement):
custom_llm_provider=custom_llm_provider,
tools=tools,
enable_prompt_caching=enable_prompt_caching,
cache_control=non_default_params.get("cache_control"),
request_kwargs=non_default_params,
)
if points:
non_default_params["cache_control_injection_points"] = points
@ -853,10 +920,13 @@ class AnthropicCacheControlHook(CustomPromptManagement):
enable_prompt_caching: Final = cast( # cast-ok: kwargs is untyped; key stamped as bool by the proxy
bool | None, kwargs.pop("enable_prompt_caching", None)
)
cache_control: Final = kwargs.get("cache_control")
configured: Final = cast( # cast-ok: kwargs is untyped; this key only holds the documented injection-point list
list[CacheControlInjectionPoint] | None, kwargs.pop("cache_control_injection_points", None)
)
if configured and AnthropicCacheControlHook._should_stand_down(configured, typed_messages, system, tools):
if configured and AnthropicCacheControlHook._should_stand_down(
configured, typed_messages, system, tools, cache_control
):
return messages, system
injection_points: list[CacheControlInjectionPoint] = configured or []
if not injection_points and model is not None:
@ -867,6 +937,8 @@ class AnthropicCacheControlHook(CustomPromptManagement):
model=model,
custom_llm_provider=custom_llm_provider,
enable_prompt_caching=enable_prompt_caching,
cache_control=cache_control,
request_kwargs=kwargs,
)
if not injection_points:
return messages, system

View file

@ -67,6 +67,96 @@ _BEDROCK_VERSION_SUFFIX_RE: Final = re.compile(r"-v\d+(?::\d+)?$")
_INFERENCE_PROFILE_MINOR_RE: Final = re.compile(r":\d+$")
_DATED_RELEASE_SUFFIX_RE: Final = re.compile(r"-\d{8}$")
_DOTTED_VERSION_RE: Final = re.compile(r"(\d)\.(\d)")
_CLAUDE_CODE_BILLING_HEADER_PREFIX: Final = "x-anthropic-billing-header:"
_CLAUDE_CODE_OBJECT_MAPPING_ADAPTER: Final = TypeAdapter(dict[object, object])
_CLAUDE_CODE_OBJECT_LIST_ADAPTER: Final = TypeAdapter(list[object])
def is_claude_code_user_agent(user_agent: str) -> bool:
return user_agent.startswith("claude-cli/")
def _validated_claude_code_mapping(value: object) -> dict[object, object] | None:
try:
return _CLAUDE_CODE_OBJECT_MAPPING_ADAPTER.validate_python(value)
except ValidationError:
return None
def _validated_claude_code_list(value: object) -> list[object] | None:
try:
return _CLAUDE_CODE_OBJECT_LIST_ADAPTER.validate_python(value)
except ValidationError:
return None
def _claude_code_billing_fields(text: str) -> tuple[tuple[str, str], ...] | None:
stripped: Final = text.strip()
if "\n" in stripped or "\r" in stripped or not stripped.startswith(_CLAUDE_CODE_BILLING_HEADER_PREFIX):
return None
fields: Final = tuple(
field
for raw_field in stripped.removeprefix(_CLAUDE_CODE_BILLING_HEADER_PREFIX).split(";")
if (field := raw_field.strip())
)
if not fields or any("=" not in field for field in fields):
return None
parsed_fields: Final = tuple(
(parts[0].strip(), parts[1].strip()) for field in fields for parts in (field.split("=", 1),)
)
if any(not key or not value for key, value in parsed_fields):
return None
return parsed_fields
def _claude_code_billing_texts(system: object) -> tuple[str, ...] | None:
if isinstance(system, str):
return (system,)
blocks: Final = _validated_claude_code_list(system)
if blocks is None:
return None
block_mappings: Final = tuple(_validated_claude_code_mapping(block) for block in blocks)
if any(block is None for block in block_mappings):
return None
text_values: Final = tuple(
block.get("text") for block in block_mappings if block is not None and block.get("type") == "text"
)
if len(text_values) != len(blocks) or any(not isinstance(text, str) for text in text_values):
return None
meaningful_text: Final = tuple(text for text in text_values if isinstance(text, str) and text.strip())
return meaningful_text or None
def _is_claude_code_subagent_billing_system(system: object) -> bool:
billing_texts: Final = _claude_code_billing_texts(system)
if billing_texts is None:
return False
billing_fields: Final = tuple(
fields for text in billing_texts if (fields := _claude_code_billing_fields(text)) is not None
)
if len(billing_fields) != len(billing_texts):
return False
subagent_values: Final = tuple(
value for fields in billing_fields for key, value in fields if key == "cc_is_subagent"
)
return subagent_values == ("true",)
def is_claude_code_one_shot_subagent_request(
messages: list[AllMessageValues],
system: object,
tools: object,
user_agent: str | None,
) -> bool:
only_message: Final = _validated_claude_code_mapping(messages[0]) if len(messages) == 1 else None
return (
user_agent is not None
and is_claude_code_user_agent(user_agent)
and not tools
and only_message is not None
and only_message.get("role") == "user"
and _is_claude_code_subagent_billing_system(system)
)
def _strip_bedrock_id_suffixes(model: str) -> str:

View file

@ -82,7 +82,7 @@ async def anthropic_messages_with_mcp(
LiteLLM_Proxy_MCP_Handler,
)
mcp_references, other_tools = LiteLLM_Proxy_MCP_Handler._parse_mcp_tools(tools)
mcp_references, other_tools = await LiteLLM_Proxy_MCP_Handler._split_mcp_tools(tools)
if not mcp_references:
return await _AnthropicMessagesCall(fn=litellm.anthropic_messages).fn(

View file

@ -3108,15 +3108,17 @@ class MCPRequestHandler:
@staticmethod
async def _get_allowed_mcp_servers_for_agent(
user_api_key_auth: UserAPIKeyAuth | None = None,
agent_object_permission=None,
agent_object_permission: LiteLLM_ObjectPermissionTable | None = None,
) -> list[str]:
"""
Get allowed MCP servers for an agent (from the agent's object_permission).
Returns the MCP servers from the agent's object_permission.
If agent has no object_permission, returns [] (no extra restriction). An entitlement the
agent LINKS but that cannot be read raises ``UnloadableEntitlementError`` out of here so the
resolver denies.
Returns the agent's direct servers, the servers in its access groups, and the servers reached
through its toolsets, exactly as the key, team, and org levels count theirs. If agent has no
object_permission, returns [] (no extra restriction). An entitlement the agent LINKS but that
cannot be read, or a declared toolset that resolves to no grants, raises
``UnloadableEntitlementError`` out of here so the resolver denies instead of reading the
agent as unrestricted.
Args:
user_api_key_auth: User auth with agent_id
@ -3126,31 +3128,30 @@ class MCPRequestHandler:
if not user_api_key_auth or not user_api_key_auth.agent_id:
return []
obj_perm = agent_object_permission
if obj_perm is None:
obj_perm = await MCPRequestHandler._get_agent_object_permission(user_api_key_auth)
obj_perm: Final = (
agent_object_permission
if agent_object_permission is not None
else await MCPRequestHandler._get_agent_object_permission(user_api_key_auth)
)
if obj_perm is None:
return []
try:
direct_mcp_servers = getattr(obj_perm, "mcp_servers", None) or []
if isinstance(direct_mcp_servers, str):
direct_mcp_servers = []
mcp_access_groups = getattr(obj_perm, "mcp_access_groups", None) or []
if isinstance(mcp_access_groups, str):
mcp_access_groups = []
# Permission entries may be server_ids OR names/aliases — expand to ids.
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
global_mcp_server_manager,
)
expanded_direct_servers: Final = global_mcp_server_manager.expand_permission_list(list(direct_mcp_servers))
access_group_servers: Final = await MCPRequestHandler._get_mcp_servers_from_access_groups(mcp_access_groups)
all_servers: Final = expanded_direct_servers + access_group_servers
return list(set(all_servers))
expanded_direct_servers: Final = global_mcp_server_manager.expand_permission_list(
obj_perm.mcp_servers or []
)
access_group_servers: Final = await MCPRequestHandler._get_mcp_servers_from_access_groups(
obj_perm.mcp_access_groups or []
)
toolset_grants: Final = await MCPRequestHandler._toolset_tool_permissions(obj_perm)
return list({*expanded_direct_servers, *access_group_servers, *toolset_grants})
except Exception as e:
if isinstance(e, UnloadableEntitlementError):
raise
verbose_logger.warning("Failed to get allowed MCP servers for agent: %s", e)
return []
@ -3158,13 +3159,15 @@ class MCPRequestHandler:
async def _get_agent_tool_permissions_for_server(
server_id: str,
user_api_key_auth: UserAPIKeyAuth | None = None,
agent_object_permission=None,
agent_object_permission: LiteLLM_ObjectPermissionTable | None = None,
) -> list[str] | None:
"""
Get allowed tool names for a server from the agent's object_permission.
Returns None if agent has no tool restrictions for this server. An entitlement the agent
LINKS but that cannot be read raises ``UnloadableEntitlementError`` out of here, which the
tool resolver turns into deny-all for the server rather than an unrestricted tool list.
Get allowed tool names for a server from the agent's object_permission: the union of its
direct tool permissions and the tools its toolsets grant on that server, mirroring the key and
team levels. Returns None if agent has no tool restrictions for this server. An entitlement the
agent LINKS but that cannot be read, or a declared toolset that resolves to no grants, raises
``UnloadableEntitlementError`` out of here, which the tool resolver turns into deny-all for the
server rather than an unrestricted tool list.
Args:
server_id: Server ID to check permissions for
@ -3175,24 +3178,30 @@ class MCPRequestHandler:
if not user_api_key_auth or not user_api_key_auth.agent_id:
return None
obj_perm = agent_object_permission
if obj_perm is None:
obj_perm = await MCPRequestHandler._get_agent_object_permission(user_api_key_auth)
obj_perm: Final = (
agent_object_permission
if agent_object_permission is not None
else await MCPRequestHandler._get_agent_object_permission(user_api_key_auth)
)
if obj_perm is None:
return None
try:
mcp_tool_permissions: Final = getattr(obj_perm, "mcp_tool_permissions", None)
if not mcp_tool_permissions or not isinstance(mcp_tool_permissions, dict):
return None
# Dict keys may be server_ids OR names/aliases; normalize before lookup.
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
global_mcp_server_manager,
)
tools: Final = global_mcp_server_manager.expand_tool_permissions(mcp_tool_permissions).get(server_id)
return list(tools) if tools else None
direct_tools: Final = (
global_mcp_server_manager.expand_tool_permissions(obj_perm.mcp_tool_permissions).get(server_id)
if obj_perm.mcp_tool_permissions
else None
)
toolset_tools: Final = await MCPRequestHandler._toolset_tools_for_server(obj_perm, server_id)
agent_tools: Final = MCPRequestHandler._union_tool_grants(direct_tools, toolset_tools)
return list(agent_tools) if agent_tools else None
except Exception as e:
if isinstance(e, UnloadableEntitlementError):
raise
verbose_logger.warning("Failed to get agent tool permissions for server: %s", e)
return None

View file

@ -23,6 +23,7 @@ from pydantic import AnyUrl, ConfigDict
from starlette.requests import Request as StarletteRequest
from starlette.responses import JSONResponse
from starlette.types import Message, Receive, Scope, Send
from typing_extensions import ReadOnly, TypedDict
from litellm._logging import verbose_logger
from litellm.constants import MAXIMUM_TRACEBACK_LINES_TO_LOG
@ -816,6 +817,11 @@ if MCP_AVAILABLE:
}
}
return ListToolsResult.model_validate({"tools": listing.tools, "_meta": outcome_meta})
except HTTPException as e:
from mcp.shared.exceptions import McpError
from mcp.types import INVALID_REQUEST, ErrorData
raise McpError(ErrorData(code=INVALID_REQUEST, message=_http_detail_message(e.detail))) from e
except Exception as e:
verbose_logger.exception("Error in list_tools endpoint: %s", e)
# Return empty list instead of failing completely
@ -1095,6 +1101,7 @@ if MCP_AVAILABLE:
mcp_server_auth_headers=mcp_server_auth_headers,
oauth2_headers=oauth2_headers,
raw_headers=raw_headers,
client_ip=_client_ip,
host_progress_callback=host_progress_callback,
**data, # for logging
)
@ -1128,7 +1135,7 @@ if MCP_AVAILABLE:
except HTTPException as e:
verbose_logger.error("HTTPException in MCP tool call: %s", e)
return CallToolResult(
content=[TextContent(text=f"Error: {e.detail}", type="text")],
content=[TextContent(text=f"Error: {_http_detail_message(e.detail)}", type="text")],
isError=True,
)
except MCPUpstreamAuthError as e:
@ -1392,7 +1399,7 @@ if MCP_AVAILABLE:
########################################################
async def _get_allowed_mcp_servers_from_mcp_server_names(
mcp_servers: list[str] | None,
mcp_servers: Sequence[str] | None,
allowed_mcp_servers: list[MCPServer],
) -> list[MCPServer]:
"""
@ -1413,13 +1420,10 @@ if MCP_AVAILABLE:
server_name_matched = False
for server in allowed_mcp_servers:
if server:
match_list = [s.lower() for s in iter_known_server_prefixes(server) if s]
if server_or_group.lower() in match_list:
filtered_server[server.server_id] = server
server_name_matched = True
break
if server and _server_answers_to(server, server_or_group):
filtered_server[server.server_id] = server
server_name_matched = True
break
if not server_name_matched:
try:
@ -1449,6 +1453,72 @@ if MCP_AVAILABLE:
return allowed_mcp_servers
def _http_detail_message(detail: object) -> str:
return str(detail.get("error")) if isinstance(detail, dict) and detail.get("error") else str(detail)
def _server_answers_to(server: MCPServer, name: str) -> bool:
requested: Final = name.lower()
return any(requested == known.lower() for known in iter_known_server_prefixes(server) if known)
class _McpDeniedDetail(TypedDict):
error: ReadOnly[str]
async def raise_denied_scoped_mcp_access(
requested_names: Sequence[str],
user_api_key_auth: UserAPIKeyAuth | None,
client_ip: str | None = None,
) -> None:
"""A scoped request (``/mcp/<name>`` path or ``x-mcp-servers`` header) resolved to zero
allowed servers, so the denial must be loud: a silent 200 with no tools reads as a healthy
server with no tools. Unknown, unauthorized, and access-group names all share one generic
error so scoping cannot probe which servers exist; the agent variant fires only when the
same request resolves once the agent binding is stripped, proving the binding caused the veto."""
agent_id: Final = user_api_key_auth.agent_id if user_api_key_auth else None
if user_api_key_auth is not None and agent_id:
resolved_without_agent: Final = await _get_allowed_mcp_servers(
user_api_key_auth=user_api_key_auth.model_copy(update=types.MappingProxyType({"agent_id": None})),
mcp_servers=requested_names,
client_ip=client_ip,
)
def _resolved_to_server(name: str) -> bool:
return any(_server_answers_to(server, name) for server in resolved_without_agent)
vetoed_server: Final = next((name for name in requested_names if _resolved_to_server(name)), None)
if vetoed_server is not None:
agent_denial: Final[_McpDeniedDetail] = {
"error": (
f"MCP server '{vetoed_server}' is not available to this key: the key is bound to "
f"agent '{agent_id}', whose MCP grants do not include this server. Add the server "
f"to the agent's object_permission.mcp_servers (edit the agent in the Admin UI or "
f"PATCH /v1/agents/{agent_id}), or use a key that is not bound to the agent."
)
}
raise HTTPException(status_code=403, detail=agent_denial)
vetoed_group: Final = next(
(
name
for name in requested_names
if not _resolved_to_server(name)
and any(name in (server.access_groups or ()) for server in resolved_without_agent)
),
None,
)
if vetoed_group is not None:
group_denial: Final[_McpDeniedDetail] = {
"error": (
f"MCP access group '{vetoed_group}' is not available to this key: the key is bound to "
f"agent '{agent_id}', whose MCP grants do not include it. Add the group to the "
f"agent's object_permission.mcp_access_groups (edit the agent in the Admin UI or "
f"PATCH /v1/agents/{agent_id}), or use a key that is not bound to the agent."
)
}
raise HTTPException(status_code=403, detail=group_denial)
generic_denial: Final[_McpDeniedDetail] = {
"error": f"The key is not allowed to access the requested MCP servers: {', '.join(requested_names)}"
}
raise HTTPException(status_code=403, detail=generic_denial)
def _tool_name_matches(tool_name: str, filter_list: list[str], mcp_server: MCPServer) -> bool:
"""
Check if a tool name matches any name in the filter list.
@ -1541,7 +1611,7 @@ if MCP_AVAILABLE:
async def _get_allowed_mcp_servers(
user_api_key_auth: UserAPIKeyAuth | None,
mcp_servers: list[str] | None,
mcp_servers: Sequence[str] | None,
client_ip: str | None = None,
) -> list[MCPServer]:
"""Return allowed MCP servers for a request after applying filters.
@ -1977,6 +2047,12 @@ if MCP_AVAILABLE:
mcp_servers=mcp_servers,
client_ip=client_ip,
)
if mcp_servers and not allowed_mcp_servers:
await raise_denied_scoped_mcp_access(
requested_names=mcp_servers,
user_api_key_auth=user_api_key_auth,
client_ip=client_ip,
)
# Pre-fetch OAuth credentials only when at least one server uses OAuth2,
# to avoid an unnecessary DB round-trip on requests with no OAuth2 MCP servers.
@ -2404,6 +2480,8 @@ if MCP_AVAILABLE:
)
verbose_logger.debug("Successfully fetched %s tools from managed MCP servers", len(listing.tools))
return listing
except HTTPException:
raise
except Exception as e:
verbose_logger.exception("Error getting tools from managed MCP servers: %s", e)
# Continue with an empty listing instead of failing completely
@ -3086,6 +3164,7 @@ if MCP_AVAILABLE:
mcp_server_auth_headers: dict[str, dict[str, str]] | None = None,
oauth2_headers: dict[str, str] | None = None,
raw_headers: dict[str, str] | None = None,
client_ip: str | None = None,
**kwargs: Any,
) -> CallToolResult:
"""
@ -3116,6 +3195,12 @@ if MCP_AVAILABLE:
mcp_servers=mcp_servers,
allowed_mcp_servers=allowed_mcp_servers,
)
if mcp_servers and not allowed_mcp_servers:
await raise_denied_scoped_mcp_access(
requested_names=mcp_servers,
user_api_key_auth=user_api_key_auth,
client_ip=client_ip,
)
if not allowed_mcp_servers:
raise HTTPException(
status_code=403,

View file

@ -366,6 +366,7 @@ async def handle_mcp_tool_call(
from litellm.proxy._experimental.mcp_server.server import (
_get_allowed_mcp_servers,
execute_mcp_tool,
raise_denied_scoped_mcp_access,
)
allowed_mcp_servers: Final = await _get_allowed_mcp_servers(
@ -373,6 +374,12 @@ async def handle_mcp_tool_call(
mcp_servers=mcp_servers,
client_ip=client_ip,
)
if mcp_servers and not allowed_mcp_servers:
await raise_denied_scoped_mcp_access(
requested_names=mcp_servers,
user_api_key_auth=user_api_key_dict,
client_ip=client_ip,
)
# Reject before dispatch when the key has no accessible servers; otherwise an
# unprefixed local tool name would fall through to the local registry in

View file

@ -2634,6 +2634,20 @@
],
"title": "Mcp Tool Permissions"
},
"mcp_toolsets": {
"anyOf": [
{
"items": {
"type": "string"
},
"type": "array"
},
{
"type": "null"
}
],
"title": "Mcp Toolsets"
},
"models": {
"anyOf": [
{

View file

@ -5,7 +5,7 @@ Pre-call hook that filters MCP tools semantically before LLM inference.
Reduces context window size and improves tool selection accuracy.
"""
from collections.abc import Iterable, Mapping, Sequence
from collections.abc import Awaitable, Callable, Collection, Iterable, Mapping, Sequence
from typing import TYPE_CHECKING, Final, Optional
from fastapi import HTTPException
@ -104,7 +104,7 @@ class SemanticToolFilterHook(CustomLogger):
)
# Parse to separate MCP tools from other tools
mcp_tools, _ = LiteLLM_Proxy_MCP_Handler._parse_mcp_tools(tools)
mcp_tools, _ = await LiteLLM_Proxy_MCP_Handler._split_mcp_tools(tools)
if not mcp_tools:
return []
@ -173,7 +173,11 @@ class SemanticToolFilterHook(CustomLogger):
return [name for name in names if name]
@staticmethod
def _narrow_mcp_references(tools: Sequence[Mapping[str, object]], selected_tool_names: list[str]) -> list[object]:
async def _narrow_mcp_references(
tools: Sequence[Mapping[str, object]],
selected_tool_names: list[str],
served_names: Callable[[Collection[str]], Awaitable[frozenset[str]]] | None = None,
) -> list[object]:
"""
Restrict each litellm_proxy MCP reference to the semantically selected tools.
@ -192,13 +196,14 @@ class SemanticToolFilterHook(CustomLogger):
LiteLLM_Proxy_MCP_Handler,
)
via_gateway: Final = await (
LiteLLM_Proxy_MCP_Handler.routes_through_gateway(tools, served_names)
if served_names is not None
else LiteLLM_Proxy_MCP_Handler.routes_through_gateway(tools)
)
return [
(
{**tool, "allowed_tools": selected_tool_names}
if isinstance(tool, dict) and LiteLLM_Proxy_MCP_Handler._should_use_litellm_mcp_gateway([tool])
else tool
)
for tool in tools
{**tool, "allowed_tools": selected_tool_names} if isinstance(tool, dict) and routed else tool
for tool, routed in zip(tools, via_gateway, strict=True)
]
def _is_mcp_tool(self, tool: object) -> bool:
@ -325,7 +330,7 @@ class SemanticToolFilterHook(CustomLogger):
filtered_expanded_tools = await self._filter_expanded_tools(data=data, expanded_tools=expanded_tools)
selected_tool_names: Final = self._selected_tool_names(filtered_expanded_tools)
narrowed_tools: Final = self._narrow_mcp_references(tools, selected_tool_names)
narrowed_tools: Final = await self._narrow_mcp_references(tools, selected_tool_names)
data["tools"] = narrowed_tools
self._emit_filter_metadata_safe(
data=data,

View file

@ -789,12 +789,6 @@ def apply_missing_session_id_policy(
)
def is_claude_code_user_agent(user_agent: str) -> bool:
"""Claude Code identifies itself as ``claude-cli/<version> ...``; the IDE
extensions and the Agent SDK run through the same CLI and share that prefix."""
return user_agent.startswith("claude-cli/")
def is_codex_user_agent(user_agent: str) -> bool:
"""Codex builds its user agent as ``<originator>/<version> ...`` and ships
several first-party originators: ``codex-tui``, ``codex_cli_rs``,
@ -811,6 +805,8 @@ def should_auto_drop_params_for_agentic_cli(user_agent: str, data: dict, proxy_c
requests routed to providers that reject them. An explicit drop_params
from the caller or in the operator's ``litellm_settings`` always wins
over this default."""
from litellm.llms.anthropic.common_utils import is_claude_code_user_agent
if not (is_claude_code_user_agent(user_agent) or is_codex_user_agent(user_agent)):
return False
if "drop_params" in data:

View file

@ -177,7 +177,7 @@ async def aresponses_api_with_mcp(
(
mcp_tools_with_litellm_proxy,
other_tools,
) = LiteLLM_Proxy_MCP_Handler._parse_mcp_tools(tools)
) = await LiteLLM_Proxy_MCP_Handler._split_mcp_tools(tools)
# Process MCP tools through the complete pipeline (fetch + filter + deduplicate + transform)
# Extract user_api_key_auth from litellm_metadata (where it's added by add_user_api_key_auth_to_request_metadata)
@ -236,6 +236,7 @@ async def aresponses_api_with_mcp(
"timeout": timeout,
"custom_llm_provider": custom_llm_provider,
**kwargs,
"_skip_mcp_handler": True,
}
# Handle MCP streaming if requested
@ -898,13 +899,14 @@ def _responses_try_dispatch_mcp_gateway(
custom_llm_provider: str | None,
kwargs: dict[str, object],
_is_async: bool,
skip_mcp_handler: bool,
) -> Any | None:
"""Return a response when MCP gateway handles the call; otherwise None."""
from litellm.responses.mcp.litellm_proxy_mcp_handler import (
LiteLLM_Proxy_MCP_Handler,
)
if not LiteLLM_Proxy_MCP_Handler._should_use_litellm_mcp_gateway(tools=tools):
if skip_mcp_handler or not LiteLLM_Proxy_MCP_Handler._should_use_litellm_mcp_gateway(tools=tools):
return None
mcp_call_kwargs: Final = {
"input": input,
@ -1074,6 +1076,7 @@ def responses(
litellm_logging_obj: Final[LiteLLMLoggingObj] = kwargs.get("litellm_logging_obj")
litellm_call_id: Final[str | None] = kwargs.get("litellm_call_id", None)
_is_async: Final = kwargs.pop("aresponses", False) is True
skip_mcp_handler: Final = kwargs.pop("_skip_mcp_handler", False)
use_chat_completions_api = _pop_use_chat_completions_api_kw(kwargs)
client_headers: Final = kwargs.get("headers")
@ -1168,6 +1171,7 @@ def responses(
custom_llm_provider=custom_llm_provider,
kwargs=kwargs,
_is_async=_is_async,
skip_mcp_handler=skip_mcp_handler,
)
if _mcp_dispatch is not None:
return _mcp_dispatch

View file

@ -106,7 +106,7 @@ async def acompletion_with_mcp(
(
mcp_tools_with_litellm_proxy,
other_tools,
) = LiteLLM_Proxy_MCP_Handler._parse_mcp_tools(tools)
) = await LiteLLM_Proxy_MCP_Handler._split_mcp_tools(tools)
if not mcp_tools_with_litellm_proxy:
# No MCP tools, proceed with regular completion
@ -114,6 +114,7 @@ async def acompletion_with_mcp(
model=model,
messages=messages,
tools=tools,
_skip_mcp_handler=True,
**kwargs,
)

View file

@ -1,6 +1,6 @@
import re
import traceback
from collections.abc import Iterable, Mapping, Sequence
from collections.abc import Awaitable, Callable, Collection, Iterable, Mapping, Sequence
from datetime import datetime
from typing import TYPE_CHECKING, Any, Final, Literal, Optional, TypeAlias, TypedDict, overload
@ -11,6 +11,7 @@ from litellm._logging import verbose_logger
from litellm.constants import MAXIMUM_TRACEBACK_LINES_TO_LOG
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.proxy._experimental.mcp_server.utils import (
iter_known_server_prefixes,
logging_safe_mcp_headers,
split_server_prefix_from_name,
strip_known_server_prefix,
@ -23,6 +24,7 @@ from litellm.types.llms.openai import (
ResponsesAPIStreamingResponse,
)
from litellm.types.llms.openai import ToolParam as ResponsesToolParam
from litellm.types.mcp_server.mcp_server_manager import MCPServer
from litellm.types.utils import (
CallTypes,
ChatCompletionMessageCustomToolCall,
@ -45,6 +47,7 @@ else:
# NOTE: We intentionally keep ToolParam as a broad type here to avoid tight coupling
ToolParam: TypeAlias = Mapping[str, object]
SplitTools: TypeAlias = tuple[list[ToolParam], list[Any]]
class MCPToolResult(TypedDict):
@ -56,14 +59,74 @@ class MCPToolResult(TypedDict):
LITELLM_PROXY_MCP_SERVER_URL: Final = "litellm_proxy"
LITELLM_PROXY_MCP_SERVER_URL_PREFIX: Final = f"{LITELLM_PROXY_MCP_SERVER_URL}/mcp/"
# Matches any URL whose path ends with /mcp/<server_name> — covers both root-path
# (http://host:port/mcp/name) and sub-path (http://host/base/mcp/name) proxy deployments.
# A false-positive match (e.g. an external URL that happens to end with /mcp/<name>) results
# in a "server not found" error from the internal gateway, not a silent failure or data leak,
# so this broad pattern is intentional and preferred over anchoring to localhost only.
_PROXY_MCP_PATH_RE: Final = re.compile(r"^https?://.+/mcp/([^/]+)$")
def _mcp_server_url(tool: ToolParam) -> str | None:
if not isinstance(tool, dict) or tool.get("type") != "mcp":
return None
server_url: Final = tool.get("server_url")
return server_url if isinstance(server_url, str) else None
def _names_gateway_explicitly(tool: ToolParam) -> bool:
return (_mcp_server_url(tool) or "").startswith(LITELLM_PROXY_MCP_SERVER_URL)
def _proxy_path_mcp_name(tool: ToolParam) -> str | None:
server_url: Final = _mcp_server_url(tool)
match: Final = None if server_url is None else _PROXY_MCP_PATH_RE.match(server_url)
return None if match is None else match.group(1)
def _registered_mcp_servers() -> Collection[MCPServer]:
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
global_mcp_server_manager,
)
return global_mcp_server_manager.get_registry().values()
def _registry_serves(name: str, servers: Collection[MCPServer]) -> bool:
requested: Final = name.lower()
return any(
requested in (known.lower() for known in (*iter_known_server_prefixes(server), server.name))
or name in (server.access_groups or ())
for server in servers
)
async def _toolset_exists(name: str) -> bool:
try:
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
global_mcp_server_manager,
)
from litellm.proxy.proxy_server import prisma_client
if prisma_client is None:
return False
return await global_mcp_server_manager.get_toolset_by_name_cached(prisma_client, name) is not None
except Exception as e:
verbose_logger.debug("Could not resolve '%s' as toolset: %s", name, e)
return False
async def _gateway_served_names(
names: Collection[str],
servers: Callable[[], Collection[MCPServer]] = _registered_mcp_servers,
toolset_exists: Callable[[str], Awaitable[bool]] = _toolset_exists,
) -> frozenset[str]:
registered: Final = tuple(servers()) if names else ()
return frozenset([name for name in names if _registry_serves(name, registered) or await toolset_exists(name)])
async def _served_mcp_path_names(
tools: Collection[ToolParam], served_names: Callable[[Collection[str]], Awaitable[frozenset[str]]]
) -> frozenset[str]:
names: Final = frozenset(name for name in map(_proxy_path_mcp_name, tools) if name is not None)
return await served_names(names) if names else frozenset[str]()
class LiteLLM_Proxy_MCP_Handler:
"""
Helper class with static methods for MCP integration with Responses API.
@ -87,57 +150,41 @@ class LiteLLM_Proxy_MCP_Handler:
@staticmethod
def _should_use_litellm_mcp_gateway(tools: Iterable[ToolParam] | None) -> bool:
"""
Returns True if any MCP tool should be handled via the litellm proxy MCP gateway.
This includes tools with server_url="litellm_proxy" as well as URLs ending in /mcp/<name>.
"""
if tools:
for tool in tools:
if isinstance(tool, dict) and tool.get("type") == "mcp":
server_url = tool.get("server_url", "")
if isinstance(server_url, str) and server_url.startswith(LITELLM_PROXY_MCP_SERVER_URL):
return True
if isinstance(server_url, str) and _PROXY_MCP_PATH_RE.match(server_url):
return True
return False
"""True when a tool may name this gateway: server_url "litellm_proxy..." or an http(s) URL ending in
/mcp/<name>. `_split_mcp_tools` then settles which of the latter the gateway actually serves."""
return any(_names_gateway_explicitly(tool) or _proxy_path_mcp_name(tool) is not None for tool in tools or ())
@staticmethod
def _parse_mcp_tools(
def _parse_mcp_tools(tools: Iterable[Mapping[str, object]] | None) -> SplitTools:
items: Final = tuple(tools or ())
gateway_tools: Final[list[ToolParam]] = [tool for tool in items if _names_gateway_explicitly(tool)]
other_tools: Final[list[Any]] = [tool for tool in items if not _names_gateway_explicitly(tool)]
return gateway_tools, other_tools
@staticmethod
async def _split_mcp_tools(
tools: Iterable[Mapping[str, object]] | None,
) -> tuple[list[ToolParam], list[Any]]:
"""
Parse tools and separate MCP tools with litellm_proxy from other tools.
served_names: Callable[[Collection[str]], Awaitable[frozenset[str]]] = _gateway_served_names,
) -> SplitTools:
items: Final = tuple(tools or ())
served: Final = await _served_mcp_path_names(items, served_names)
return LiteLLM_Proxy_MCP_Handler._parse_mcp_tools(
[
{**tool, "server_url": f"{LITELLM_PROXY_MCP_SERVER_URL_PREFIX}{name}"}
if (name := _proxy_path_mcp_name(tool)) in served
else tool
for tool in items
]
)
Returns:
Tuple of (mcp_tools_with_litellm_proxy, other_tools)
"""
mcp_tools_with_litellm_proxy: Final[list[ToolParam]] = []
other_tools: Final[list[Any]] = []
if tools:
for tool in tools:
if isinstance(tool, dict) and tool.get("type") == "mcp":
server_url = tool.get("server_url", "")
if isinstance(server_url, str) and server_url.startswith(LITELLM_PROXY_MCP_SERVER_URL):
mcp_tools_with_litellm_proxy.append(tool)
elif isinstance(server_url, str):
# Also intercept URLs like http://localhost:4000/mcp/atlassian_test
# by rewriting them to the internal litellm_proxy format.
m = _PROXY_MCP_PATH_RE.match(server_url)
if m:
rewritten = {
**tool,
"server_url": f"{LITELLM_PROXY_MCP_SERVER_URL_PREFIX}{m.group(1)}",
}
mcp_tools_with_litellm_proxy.append(rewritten)
else:
other_tools.append(tool)
else:
other_tools.append(tool)
else:
other_tools.append(tool)
return mcp_tools_with_litellm_proxy, other_tools
@staticmethod
async def routes_through_gateway(
tools: Iterable[Mapping[str, object]] | None,
served_names: Callable[[Collection[str]], Awaitable[frozenset[str]]] = _gateway_served_names,
) -> tuple[bool, ...]:
items: Final = tuple(tools or ())
served: Final = await _served_mcp_path_names(items, served_names)
return tuple(_names_gateway_explicitly(tool) or _proxy_path_mcp_name(tool) in served for tool in items)
@staticmethod
async def _apply_toolset_permissions(

View file

@ -90,6 +90,7 @@ class PromptCachingDeploymentCheck(CustomLogger):
enable_prompt_caching=(
request_kwargs.get("enable_prompt_caching") is True if request_kwargs is not None else None
),
request_kwargs=request_kwargs,
)
model_id_dict: Final = await prompt_cache.async_get_model_id(

View file

@ -1,9 +1,9 @@
from collections.abc import Mapping
from collections.abc import Mapping, Sequence
from datetime import datetime
from typing import TYPE_CHECKING, Any, Final, Literal
from pydantic import BaseModel, PrivateAttr, StrictInt
from typing_extensions import Required, TypedDict
from typing_extensions import ReadOnly, Required, TypedDict
from litellm.types.llms.base import LiteLLMPydanticObjectBase
@ -172,6 +172,7 @@ class AugmentedAgentCard(AgentCard):
class AgentObjectPermission(TypedDict, total=False):
mcp_servers: list[str] | None
mcp_access_groups: list[str] | None
mcp_toolsets: ReadOnly[Sequence[str] | None]
mcp_tool_permissions: dict[str, list[str]] | None
models: list[str] | None
agents: list[str] | None

View file

@ -1777,6 +1777,224 @@ class TestEnableAnthropicPromptCaching:
assert messages == before
class TestClaudeCodeOneShotAutoCaching:
BILLING_TEXT = "x-anthropic-billing-header: cc_version=2.1.263; cc_entrypoint=cli; cc_is_subagent=true;"
BILLING_SYSTEM = [{"type": "text", "text": BILLING_TEXT}]
MESSAGES = [{"role": "user", "content": [{"type": "text", "text": "unique fetched document"}]}]
@staticmethod
def _kwargs(configured=None):
kwargs = {
"litellm_metadata": {},
"proxy_server_request": {
"headers": {
"user-agent": "claude-cli/2.1.263 (external, cli)",
"x-app": "cli-bg",
}
},
}
if configured is not None:
kwargs["cache_control_injection_points"] = configured
return kwargs
@pytest.mark.parametrize(
"system",
[
BILLING_TEXT,
BILLING_SYSTEM,
[*BILLING_SYSTEM, {"type": "text", "text": " "}],
[
*BILLING_SYSTEM,
{"type": "text", "text": "x-anthropic-billing-header: cc_version=2.1.263; cc_entrypoint=cli;"},
],
],
ids=["string", "text_block", "whitespace_block", "multiple_billing_blocks"],
)
@pytest.mark.parametrize("tools", [None, []], ids=["absent_tools", "empty_tools"])
def test_skips_defaults_and_attribution_for_one_shot_subagent(self, monkeypatch, system, tools):
monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", True)
messages = copy.deepcopy(self.MESSAGES)
kwargs = self._kwargs()
result_messages, result_system = AnthropicCacheControlHook.maybe_inject_cache_control(
messages,
copy.deepcopy(system),
kwargs,
model="claude-sonnet-4-5",
custom_llm_provider="anthropic",
tools=tools,
)
assert result_messages == self.MESSAGES
assert result_system == system
assert "litellm_gateway_injected_cache" not in kwargs["litellm_metadata"]
def test_user_agent_header_lookup_is_case_insensitive(self, monkeypatch):
monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", True)
kwargs = self._kwargs()
user_agent = kwargs["proxy_server_request"]["headers"].pop("user-agent")
kwargs["proxy_server_request"]["headers"]["User-Agent"] = user_agent
result_messages, result_system = AnthropicCacheControlHook.maybe_inject_cache_control(
copy.deepcopy(self.MESSAGES),
copy.deepcopy(self.BILLING_SYSTEM),
kwargs,
model="claude-sonnet-4-5",
custom_llm_provider="anthropic",
)
assert result_messages == self.MESSAGES
assert result_system == self.BILLING_SYSTEM
def test_router_affinity_skips_string_billing_system(self, monkeypatch):
monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", True)
messages = copy.deepcopy(self.MESSAGES)
kwargs = self._kwargs()
kwargs["system"] = self.BILLING_TEXT
result = AnthropicCacheControlHook.messages_with_default_injections(
messages=messages,
models=("claude-sonnet-4-5",),
request_kwargs=kwargs,
)
assert result == messages
@pytest.mark.parametrize(
"headers,system",
[
("not-a-mapping", BILLING_SYSTEM),
(
{"user-agent": "claude-cli/2.1.263 (external, cli)"},
[{"type": "text", "text": "x-anthropic-billing-header: malformed"}],
),
({"user-agent": "claude-cli/2.1.263 (external, cli)"}, None),
({"user-agent": "claude-cli/2.1.263 (external, cli)"}, ["not-a-mapping"]),
(
{"user-agent": "claude-cli/2.1.263 (external, cli)"},
[{"type": "image", "text": BILLING_TEXT}],
),
],
ids=["malformed_headers", "malformed_billing", "missing_system", "malformed_block", "non_text_block"],
)
def test_malformed_untrusted_context_keeps_defaults(self, monkeypatch, headers, system):
monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", True)
points = AnthropicCacheControlHook.get_default_injection_points(
messages=copy.deepcopy(self.MESSAGES),
system=copy.deepcopy(system),
model="claude-sonnet-4-5",
custom_llm_provider="anthropic",
request_kwargs={"proxy_server_request": {"headers": headers}},
)
assert len(points) == 2
def test_message_without_role_keeps_defaults(self, monkeypatch):
monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", True)
points = AnthropicCacheControlHook.get_default_injection_points(
messages=[{"content": "missing role"}],
system=copy.deepcopy(self.BILLING_SYSTEM),
model="claude-sonnet-4-5",
custom_llm_provider="anthropic",
request_kwargs=self._kwargs(),
)
assert len(points) == 2
@pytest.mark.parametrize(
"messages,system,tools",
[
(
MESSAGES,
BILLING_SYSTEM,
[{"name": "WebFetch", "description": "fetch", "input_schema": {"type": "object"}}],
),
(MESSAGES, [*BILLING_SYSTEM, {"type": "text", "text": "Explore the repository"}], None),
(
[
{"role": "user", "content": "first turn"},
{"role": "assistant", "content": "reply"},
*MESSAGES,
],
BILLING_SYSTEM,
None,
),
],
ids=["tools", "real_system", "history"],
)
def test_keeps_defaults_for_reusable_subagents(self, monkeypatch, messages, system, tools):
monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", True)
kwargs = self._kwargs()
result_messages, result_system = AnthropicCacheControlHook.maybe_inject_cache_control(
copy.deepcopy(messages),
copy.deepcopy(system),
kwargs,
model="claude-sonnet-4-5",
custom_llm_provider="anthropic",
tools=copy.deepcopy(tools),
)
assert AnthropicCacheControlHook.count_request_cache_breakpoints(result_messages, result_system) == 2
assert kwargs["litellm_metadata"]["litellm_gateway_injected_cache"] == ""
@pytest.mark.parametrize(
"user_agent,system",
[
("anthropic-sdk-python/0.75.0", BILLING_SYSTEM),
(
"claude-cli/2.1.263 (external, cli)",
[
{
"type": "text",
"text": f"{BILLING_TEXT}\nadditional system instructions",
}
],
),
(
"claude-cli/2.1.263 (external, cli)",
[
{
"type": "text",
"text": "x-anthropic-billing-header: cc_version=2.1.263; cc_is_subagent=false;",
}
],
),
],
ids=["different_client", "appended_instructions", "not_a_subagent"],
)
def test_ambiguous_or_unmatched_signals_fail_open(self, monkeypatch, user_agent, system):
monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", True)
kwargs = self._kwargs()
kwargs["proxy_server_request"]["headers"]["user-agent"] = user_agent
result_messages, result_system = AnthropicCacheControlHook.maybe_inject_cache_control(
copy.deepcopy(self.MESSAGES),
copy.deepcopy(system),
kwargs,
model="claude-sonnet-4-5",
custom_llm_provider="anthropic",
)
assert AnthropicCacheControlHook.count_request_cache_breakpoints(result_messages, result_system) == 2
def test_explicit_injection_points_remain_authoritative(self, monkeypatch):
monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", True)
kwargs = self._kwargs([{"location": "message", "role": "user"}])
result_messages, _ = AnthropicCacheControlHook.maybe_inject_cache_control(
copy.deepcopy(self.MESSAGES),
copy.deepcopy(self.BILLING_SYSTEM),
kwargs,
model="claude-sonnet-4-5",
custom_llm_provider="anthropic",
)
assert result_messages[0]["content"][-1]["cache_control"] == {"type": "ephemeral"}
class TestPerKeyEnablePromptCaching:
"""Per-request enable_prompt_caching override (stamped from key metadata) with the global flag off."""
@ -1977,6 +2195,25 @@ class TestConfiguredInjectionPointsStandDown:
_, result_sys = self._inject(copy.deepcopy(self.V1_MESSAGES), kwargs)
assert result_sys == [{"type": "text", "text": "sys", "cache_control": {"type": "ephemeral"}}]
@pytest.mark.parametrize(
"configured",
[None, CONFIGURED],
ids=["automatic_defaults", "configured_points"],
)
def test_v1_messages_stands_down_for_root_cache_control(self, monkeypatch, configured):
monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", True)
root_cache_control = {"type": "ephemeral"}
kwargs = {"cache_control": root_cache_control, "litellm_metadata": {}}
if configured is not None:
kwargs["cache_control_injection_points"] = copy.deepcopy(configured)
result_messages, result_system = self._inject(copy.deepcopy(self.V1_MESSAGES), kwargs)
assert result_messages == self.V1_MESSAGES
assert result_system == "sys"
assert kwargs["cache_control"] is root_cache_control
assert "litellm_gateway_injected_cache" not in kwargs["litellm_metadata"]
def test_v1_messages_reentry_flow_preserves_tool_config_remainder(self):
"""The advisor interceptor re-enters anthropic_messages() with the outer
request's kwargs and post-injection messages. The first pass applies the

View file

@ -26,6 +26,30 @@ FAKE_REGULAR_KEY = "sk-ant-api03-regular-key-for-testing-123456789"
FAKE_AUTH_TOKEN = "sk-ant-aut01-fake-auth-token-for-testing-123456789"
@pytest.mark.parametrize(
"messages,system,expected",
[
([{"role": "user", "content": "hi"}], "x-anthropic-billing-header: cc_is_subagent=true;", True),
([{"role": "user", "content": "hi"}], "x-anthropic-billing-header: =junk; cc_is_subagent=true;", False),
([{"role": "user", "content": "hi"}], "x-anthropic-billing-header: cc_version=; cc_is_subagent=true;", False),
([{"role": "user", "content": "hi"}], "x-anthropic-billing-header: malformed", False),
([{"content": "missing role"}], "x-anthropic-billing-header: cc_is_subagent=true;", False),
(["not-a-mapping"], "x-anthropic-billing-header: cc_is_subagent=true;", False),
([{"role": "user", "content": "hi"}], ["not-a-mapping"], False),
([{"role": "user", "content": "hi"}], None, False),
],
)
def test_is_claude_code_one_shot_subagent_request(messages, system, expected):
from litellm.llms.anthropic.common_utils import is_claude_code_one_shot_subagent_request
assert is_claude_code_one_shot_subagent_request(
messages=messages,
system=system,
tools=None,
user_agent="claude-cli/2.1.263 (external, cli)",
) is expected
class TestOptionallyHandleAnthropicOAuth:
"""Tests for optionally_handle_anthropic_oauth function."""

View file

@ -11,6 +11,7 @@ from starlette.datastructures import Headers
from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import (
MCPRequestHandler,
UnloadableEntitlementError,
_is_mcp_admitted_user_subject,
)
from litellm.proxy._types import (
@ -4396,6 +4397,147 @@ class TestAgentMCPPermissions:
)
assert sorted(result) == ["tool_a", "tool_b"]
def _agent_object_permission(self, *, toolset_ids, servers=(), tool_permissions=None):
agent_object_permission = MagicMock()
agent_object_permission.mcp_servers = list(servers)
agent_object_permission.mcp_access_groups = []
agent_object_permission.mcp_tool_permissions = tool_permissions
agent_object_permission.mcp_toolsets = list(toolset_ids)
return agent_object_permission
def _mock_manager_with_toolsets(self, toolset_perms):
mock_manager = MagicMock()
mock_manager.expand_permission_list = MagicMock(side_effect=lambda servers: list(servers))
mock_manager.expand_tool_permissions = MagicMock(side_effect=lambda perms: perms or {})
mock_manager.resolve_toolset_tool_permissions = AsyncMock(return_value=toolset_perms)
return mock_manager
def _agent_toolset_patches(self, agent_object_permission, mock_manager):
return (
patch.object( # test-quality-ok: stub the agent perm loader; the resolver reads module globals with no injection seam
MCPRequestHandler, "_get_agent_object_permission", AsyncMock(return_value=agent_object_permission)
),
patch( # test-quality-ok: isolate the MCP registry, same seam as the sibling toolset tests
"litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager",
mock_manager,
),
patch.object( # test-quality-ok: access-group lookup hits the DB, not under test here
MCPRequestHandler, "_get_mcp_servers_from_access_groups", AsyncMock(return_value=[])
),
)
async def test_get_allowed_mcp_servers_for_agent_includes_toolset_servers(self):
"""An agent granted only mcp_toolsets reaches the toolset's servers, exactly as a
key, team, or org granted only toolsets does"""
user_api_key_auth = UserAPIKeyAuth(api_key="test-key", agent_id="agent-toolsets")
agent_object_permission = self._agent_object_permission(toolset_ids=["toolset-1"], servers=["server-direct"])
mock_manager = self._mock_manager_with_toolsets({"server-a": ["lookup_status"]})
with contextlib.ExitStack() as stack:
for patcher in self._agent_toolset_patches(agent_object_permission, mock_manager):
stack.enter_context(patcher)
result = await MCPRequestHandler._get_allowed_mcp_servers_for_agent(user_api_key_auth)
assert sorted(result) == ["server-a", "server-direct"]
mock_manager.resolve_toolset_tool_permissions.assert_awaited_once_with(toolset_ids=["toolset-1"])
async def test_get_allowed_mcp_servers_toolset_only_agent_caps_key_servers(self):
"""Regression: an agent whose only grant is a toolset used to resolve to [] and place
no ceiling at all, so a key bound to it kept every server the key itself granted"""
user_api_key_auth = UserAPIKeyAuth(api_key="test-key", agent_id="agent-toolsets")
agent_object_permission = self._agent_object_permission(toolset_ids=["toolset-1"])
mock_manager = self._mock_manager_with_toolsets({"server-a": ["lookup_status"]})
with contextlib.ExitStack() as stack:
for patcher in self._agent_toolset_patches(agent_object_permission, mock_manager):
stack.enter_context(patcher)
stack.enter_context(
patch.object( # test-quality-ok: key resolution has its own tests; pin its grants here
MCPRequestHandler, "_get_allowed_mcp_servers_for_key", AsyncMock(return_value=["server-a", "server-b"])
)
)
stack.enter_context(
patch.object( # test-quality-ok: team resolution has its own tests; pin it empty here
MCPRequestHandler, "_get_allowed_mcp_servers_for_team", AsyncMock(return_value=[])
)
)
result = await MCPRequestHandler.get_allowed_mcp_servers(user_api_key_auth)
assert result == ["server-a"]
async def test_get_allowed_mcp_servers_agent_dangling_toolset_denies(self):
"""An agent toolset that resolves to nothing is a known restriction with unknown
contents: deny, never fall through to the key's own servers"""
user_api_key_auth = UserAPIKeyAuth(api_key="test-key", agent_id="agent-toolsets")
agent_object_permission = self._agent_object_permission(toolset_ids=["toolset-gone"])
mock_manager = self._mock_manager_with_toolsets({})
with contextlib.ExitStack() as stack:
for patcher in self._agent_toolset_patches(agent_object_permission, mock_manager):
stack.enter_context(patcher)
with pytest.raises(UnloadableEntitlementError):
await MCPRequestHandler._get_allowed_mcp_servers_for_agent(user_api_key_auth)
stack.enter_context(
patch.object( # test-quality-ok: key resolution has its own tests; pin its grants here
MCPRequestHandler, "_get_allowed_mcp_servers_for_key", AsyncMock(return_value=["server-a", "server-b"])
)
)
stack.enter_context(
patch.object( # test-quality-ok: team resolution has its own tests; pin it empty here
MCPRequestHandler, "_get_allowed_mcp_servers_for_team", AsyncMock(return_value=[])
)
)
result = await MCPRequestHandler.get_allowed_mcp_servers(user_api_key_auth)
assert result == []
async def test_get_agent_tool_permissions_for_server_unions_direct_and_toolset_tools(self):
"""The agent's tool ceiling on a server is its direct tool grants plus the tools its
toolsets grant there, and None only when neither names the server"""
user_api_key_auth = UserAPIKeyAuth(api_key="test-key", agent_id="agent-toolsets")
agent_object_permission = self._agent_object_permission(
toolset_ids=["toolset-1"], tool_permissions={"server-a": ["tool_direct"]}
)
mock_manager = self._mock_manager_with_toolsets({"server-a": ["tool_via_toolset"], "server-b": ["tool_b"]})
with contextlib.ExitStack() as stack:
for patcher in self._agent_toolset_patches(agent_object_permission, mock_manager):
stack.enter_context(patcher)
server_a_tools = await MCPRequestHandler._get_agent_tool_permissions_for_server("server-a", user_api_key_auth)
server_b_tools = await MCPRequestHandler._get_agent_tool_permissions_for_server("server-b", user_api_key_auth)
server_c_tools = await MCPRequestHandler._get_agent_tool_permissions_for_server("server-c", user_api_key_auth)
assert sorted(server_a_tools) == ["tool_direct", "tool_via_toolset"]
assert server_b_tools == ["tool_b"]
assert server_c_tools is None
async def test_get_allowed_tools_for_server_toolset_only_agent_caps_key_tools(self):
"""Regression: a key allowing [tool_a, tool_b] bound to an agent whose toolset grants
only tool_a on the server ends with [tool_a]; the toolset used to be ignored"""
user_api_key_auth = UserAPIKeyAuth(api_key="test-key", agent_id="agent-toolsets")
agent_object_permission = self._agent_object_permission(toolset_ids=["toolset-1"])
mock_manager = self._mock_manager_with_toolsets({"server-a": ["tool_a"]})
key_perm = MagicMock()
key_perm.mcp_tool_permissions = {"server-a": ["tool_a", "tool_b"]}
key_perm.mcp_toolsets = []
with contextlib.ExitStack() as stack:
for patcher in self._agent_toolset_patches(agent_object_permission, mock_manager):
stack.enter_context(patcher)
stack.enter_context(
patch.object( # test-quality-ok: stub the key perm loader; the resolver reads module globals with no injection seam
MCPRequestHandler, "_get_key_object_permission", return_value=key_perm
)
)
stack.enter_context(
patch.object( # test-quality-ok: team resolution has its own tests; pin it absent here
MCPRequestHandler, "_get_team_object_permission", AsyncMock(return_value=None)
)
)
result = await MCPRequestHandler.get_allowed_tools_for_server("server-a", user_api_key_auth)
assert result == ["tool_a"]
async def test_get_agent_object_permission_uses_shared_helper(self):
"""``_get_agent_object_permission`` must resolve the agent's
``object_permission_id`` and then defer to the shared

View file

@ -2,6 +2,7 @@ import asyncio
import contextvars
import os
from datetime import datetime, timedelta
from types import SimpleNamespace
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
@ -1368,6 +1369,295 @@ async def test_get_tools_from_mcp_servers_handles_all_servers_failing():
mock_logger.info.assert_any_call("Successfully fetched %s tools total from all MCP servers", 0)
def _denied_scope_manager(known_server_names_to_ids: dict[str, str]) -> MagicMock:
servers = {name: MagicMock(server_id=server_id) for name, server_id in known_server_names_to_ids.items()}
manager = MagicMock()
manager.get_mcp_server_by_name = lambda name, client_ip=None: servers.get(name)
return manager
def _scope_resolver(resolved_without_agent: dict[str, str], access_groups: tuple[str, ...] = ()) -> AsyncMock:
async def resolve(user_api_key_auth, mcp_servers, client_ip=None):
if user_api_key_auth is not None and user_api_key_auth.agent_id:
return []
return [
SimpleNamespace(
server_id=server_id,
server_name=server_name,
alias=None,
short_prefix=None,
access_groups=list(access_groups),
)
for server_name, server_id in resolved_without_agent.items()
]
return AsyncMock(side_effect=resolve)
async def _denied_scoped_list(
user_api_key_auth: UserAPIKeyAuth,
mcp_servers: list[str],
mock_manager: MagicMock,
resolver: AsyncMock,
) -> HTTPException:
from litellm.proxy._experimental.mcp_server.server import _get_tools_from_mcp_servers
with (
patch( # test-quality-ok: the permission resolver is a module-level function; the suite's only seam
"litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers",
resolver,
),
patch( # test-quality-ok: the server registry is a module-level singleton; the suite's only seam
"litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager",
mock_manager,
),
):
with pytest.raises(HTTPException) as exc_info:
await _get_tools_from_mcp_servers(
user_api_key_auth=user_api_key_auth,
mcp_auth_header=None,
mcp_servers=mcp_servers,
)
return exc_info.value
@pytest.mark.asyncio
async def test_scoped_list_denied_by_agent_binding_raises_403_naming_agent():
"""The agent-binding veto must raise a 403 naming the agent, never a silent 200 with no tools."""
pytest.importorskip("litellm.proxy._experimental.mcp_server.server")
user_api_key_auth = UserAPIKeyAuth(api_key="test_key", user_id="test_user", agent_id="agent-123")
resolver = _scope_resolver(resolved_without_agent={"github": "srv-github"})
denial = await _denied_scoped_list(
user_api_key_auth, ["github"], _denied_scope_manager({"github": "srv-github"}), resolver
)
assert denial.status_code == 403
message = denial.detail["error"]
assert "MCP server 'github'" in message
assert "agent 'agent-123'" in message
assert "mcp_servers" in message
rerun_auth = resolver.await_args_list[1].kwargs["user_api_key_auth"]
assert rerun_auth.agent_id is None
assert rerun_auth.user_id == "test_user"
assert resolver.await_args_list[1].kwargs["mcp_servers"] == ["github"]
@pytest.mark.asyncio
async def test_empty_scope_lists_nothing_instead_of_raising_a_nameless_denial():
"""An empty ``x-mcp-servers`` header scopes to no servers; that is an empty listing, not a 403."""
pytest.importorskip("litellm.proxy._experimental.mcp_server.server")
from litellm.proxy._experimental.mcp_server.server import _get_tools_from_mcp_servers
user_api_key_auth = UserAPIKeyAuth(api_key="test_key", user_id="test_user")
resolver = AsyncMock(return_value=[])
with (
patch( # test-quality-ok: the permission resolver is a module-level function; the suite's only seam
"litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers",
resolver,
),
patch( # test-quality-ok: the server registry is a module-level singleton; the suite's only seam
"litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager",
_denied_scope_manager({"github": "srv-github"}),
),
):
listing = await _get_tools_from_mcp_servers(
user_api_key_auth=user_api_key_auth, mcp_auth_header=None, mcp_servers=[]
)
assert listing.tools == []
resolver.assert_awaited_once()
@pytest.mark.asyncio
async def test_scoped_list_denied_for_non_agent_key_raises_generic_403():
"""A denial for a key with no agent binding stays generic and skips the agent-stripped rerun."""
pytest.importorskip("litellm.proxy._experimental.mcp_server.server")
user_api_key_auth = UserAPIKeyAuth(api_key="test_key", user_id="test_user")
resolver = AsyncMock(return_value=[])
denial = await _denied_scoped_list(
user_api_key_auth, ["github"], _denied_scope_manager({"github": "srv-github"}), resolver
)
assert denial.status_code == 403
message = denial.detail["error"]
assert "github" in message
assert "agent" not in message
resolver.assert_awaited_once()
@pytest.mark.asyncio
async def test_scoped_list_unknown_name_raises_same_generic_403_as_unauthorized():
"""Unknown and registered-but-unauthorized names raise byte-identical generic 403s, so a
caller cannot probe which server names exist."""
pytest.importorskip("litellm.proxy._experimental.mcp_server.server")
user_api_key_auth = UserAPIKeyAuth(api_key="test_key", user_id="test_user", agent_id="agent-123")
unknown = await _denied_scoped_list(
user_api_key_auth, ["github"], _denied_scope_manager({}), _scope_resolver(resolved_without_agent={})
)
unauthorized = await _denied_scoped_list(
user_api_key_auth,
["github"],
_denied_scope_manager({"github": "srv-github"}),
_scope_resolver(resolved_without_agent={}),
)
assert unknown.status_code == unauthorized.status_code == 403
assert unknown.detail["error"] == unauthorized.detail["error"]
assert "github" in unknown.detail["error"]
assert "agent" not in unknown.detail["error"]
@pytest.mark.asyncio
async def test_scoped_list_access_group_vetoed_by_agent_names_agent_and_group():
"""An access-group scope vetoed by the agent binding raises the 403 naming the agent and the
group instead of the silent empty list."""
pytest.importorskip("litellm.proxy._experimental.mcp_server.server")
user_api_key_auth = UserAPIKeyAuth(api_key="test_key", user_id="test_user", agent_id="agent-123")
denial = await _denied_scoped_list(
user_api_key_auth,
["prod-group"],
_denied_scope_manager({}),
_scope_resolver(resolved_without_agent={"github": "srv-github"}, access_groups=("prod-group",)),
)
assert denial.status_code == 403
message = denial.detail["error"]
assert "access group 'prod-group'" in message
assert "agent 'agent-123'" in message
assert "mcp_access_groups" in message
@pytest.mark.asyncio
async def test_scoped_list_mixed_unknown_and_vetoed_group_names_the_group_that_resolved():
"""With an unknown name ahead of the agent-vetoed group in the scope, the 403 must name the group
whose servers the key can reach, never the unknown name, or the admin is told to grant a group
that does not exist."""
pytest.importorskip("litellm.proxy._experimental.mcp_server.server")
user_api_key_auth = UserAPIKeyAuth(api_key="test_key", user_id="test_user", agent_id="agent-123")
denial = await _denied_scoped_list(
user_api_key_auth,
["no-such-group", "prod-group"],
_denied_scope_manager({}),
_scope_resolver(resolved_without_agent={"github": "srv-github"}, access_groups=("prod-group",)),
)
assert denial.status_code == 403
message = denial.detail["error"]
assert "access group 'prod-group'" in message
assert "no-such-group" not in message
assert "agent 'agent-123'" in message
@pytest.mark.asyncio
async def test_scoped_list_agent_key_denied_by_key_grants_raises_generic_403():
"""When the agent-stripped rerun still resolves nothing, the 403 stays generic instead of
blaming the agent binding."""
pytest.importorskip("litellm.proxy._experimental.mcp_server.server")
user_api_key_auth = UserAPIKeyAuth(api_key="test_key", user_id="test_user", agent_id="agent-123")
resolver = _scope_resolver(resolved_without_agent={})
denial = await _denied_scoped_list(
user_api_key_auth, ["github"], _denied_scope_manager({"github": "srv-github"}), resolver
)
assert denial.status_code == 403
message = denial.detail["error"]
assert "github" in message
assert "agent" not in message
assert resolver.await_count == 2
@pytest.mark.asyncio
async def test_scoped_list_agent_veto_attributed_for_differently_cased_server_name():
"""The scope filter matches `/mcp/GitHub` to a server named `github` case-insensitively, so the
agent-attributed 403 must match the same way instead of falling back to the generic denial."""
pytest.importorskip("litellm.proxy._experimental.mcp_server.server")
user_api_key_auth = UserAPIKeyAuth(api_key="test_key", user_id="test_user", agent_id="agent-123")
denial = await _denied_scoped_list(
user_api_key_auth,
["GitHub"],
_denied_scope_manager({"github": "srv-github"}),
_scope_resolver(resolved_without_agent={"github": "srv-github"}),
)
assert denial.status_code == 403
message = denial.detail["error"]
assert "MCP server 'GitHub'" in message
assert "agent 'agent-123'" in message
@pytest.mark.asyncio
async def test_handle_list_tools_converts_permission_httpexception_to_mcp_error():
"""The MCP protocol handler surfaces a permission HTTPException as a clean JSON-RPC error
(McpError, INVALID_REQUEST) carrying the denial message, instead of a raw 500."""
try:
from litellm.proxy._experimental.mcp_server.server import handle_list_tools
except ImportError:
pytest.skip("MCP server not available")
from mcp.shared.exceptions import McpError
from mcp.types import INVALID_REQUEST
denial_message = "MCP server 'github' is not available to this key: the key is bound to agent 'agent-123'"
denial = HTTPException(status_code=403, detail={"error": denial_message})
with (
patch( # test-quality-ok: the protocol handler reads auth from module context; no injection seam
"litellm.proxy._experimental.mcp_server.server.get_or_extract_auth_context",
new=AsyncMock(return_value=(None, None, None, None, None, None, None)),
),
patch( # test-quality-ok: the listing helper is the handler's only collaborator; the suite's seam
"litellm.proxy._experimental.mcp_server.server._list_mcp_tools",
new=AsyncMock(side_effect=denial),
),
):
with pytest.raises(McpError) as exc_info:
await handle_list_tools()
assert exc_info.value.error.code == INVALID_REQUEST
assert exc_info.value.error.message == denial_message
@pytest.mark.asyncio
async def test_mcp_server_tool_call_renders_denial_message_not_detail_dict():
try:
from litellm.proxy._experimental.mcp_server.server import mcp_server_tool_call
except ImportError:
pytest.skip("MCP server not available")
denial_message = "MCP server 'github' is not available to this key: the key is bound to agent 'agent-123'"
denial = HTTPException(status_code=403, detail={"error": denial_message})
with (
patch( # test-quality-ok: the protocol handler reads auth from module context; no injection seam
"litellm.proxy._experimental.mcp_server.server.get_or_extract_auth_context",
new=AsyncMock(return_value=(None, None, None, None, None, None, None)),
),
patch( # test-quality-ok: the tool-call helper is the handler's only collaborator; the suite's seam
"litellm.proxy._experimental.mcp_server.server.call_mcp_tool",
new=AsyncMock(side_effect=denial),
),
):
result = await mcp_server_tool_call("github-search_issues", {})
assert result.isError is True
assert result.content[0].text == f"Error: {denial_message}"
@pytest.mark.asyncio
async def test_mcp_server_tool_call_body_with_none_arguments():
"""Test that proxy_server_request body handles None arguments correctly"""
@ -3518,6 +3808,35 @@ async def test_call_mcp_tool_user_unauthorized_access():
assert "User not allowed to call this tool" in exc_info.value.detail
@pytest.mark.asyncio
async def test_call_mcp_tool_scoped_denial_names_the_binding_agent():
from litellm.proxy._experimental.mcp_server.server import call_mcp_tool
agent_bound_key = UserAPIKeyAuth(api_key="test-key", user_id="test-user", agent_id="agent-123")
with (
patch( # test-quality-ok: the server registry is a module-level singleton; the suite's only seam
"litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager.get_allowed_mcp_servers",
AsyncMock(return_value=[]),
),
patch( # test-quality-ok: the permission resolver is a module-level function; the suite's only seam
"litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers",
_scope_resolver({"github": "srv-github"}),
),
):
with pytest.raises(HTTPException) as exc_info:
await call_mcp_tool(
name="github-search_issues",
arguments={},
user_api_key_auth=agent_bound_key,
mcp_servers=["github"],
)
assert exc_info.value.status_code == 403
assert "MCP server 'github'" in exc_info.value.detail["error"]
assert "agent 'agent-123'" in exc_info.value.detail["error"]
@pytest.mark.asyncio
async def test_call_mcp_tool_unauthorized_403_does_not_leak_server_credentials():
"""Regression for LIT-4703 / GH #29936.

View file

@ -11,6 +11,7 @@ Covers:
import json
from collections.abc import Sequence
from types import SimpleNamespace
from typing import Any
from unittest.mock import AsyncMock, MagicMock, patch
@ -1271,3 +1272,37 @@ class TestMcpServerToolCallErrorHandling:
assert result.isError is True
assert "User not allowed to call this tool" in result.content[0].text
@pytest.mark.asyncio
async def test_handle_mcp_tool_call_scoped_denial_names_the_binding_agent() -> None:
from fastapi import HTTPException
from litellm.proxy._experimental.mcp_server.tool_search import handle_mcp_tool_call
agent_bound_key = UserAPIKeyAuth(api_key="test_key", agent_id="agent-123")
async def resolve(user_api_key_auth, mcp_servers, client_ip=None):
if user_api_key_auth.agent_id:
return []
return [
SimpleNamespace(
server_id="srv-github", server_name="github", alias=None, short_prefix=None, access_groups=[]
)
]
with patch( # test-quality-ok: the permission resolver is a module-level function; the suite's only seam
"litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers",
new=AsyncMock(side_effect=resolve),
):
with pytest.raises(HTTPException) as exc_info:
await handle_mcp_tool_call(
tool_name="github-create_issue",
arguments={},
user_api_key_dict=agent_bound_key,
mcp_servers=["github"],
)
assert exc_info.value.status_code == 403
assert "MCP server 'github'" in exc_info.value.detail["error"]
assert "agent 'agent-123'" in exc_info.value.detail["error"]

View file

@ -2216,3 +2216,26 @@ async def test_top_k_above_router_default_is_respected():
assert len(filtered) == 6
print("✅ Configured top_k above the semantic-router default of 5 is honored")
@pytest.mark.asyncio
async def test_semantic_filter_hook_narrows_only_references_the_gateway_serves():
from litellm.proxy.hooks.mcp_semantic_filter import SemanticToolFilterHook
gateway_reference = {"type": "mcp", "server_url": "litellm_proxy", "require_approval": "never"}
external_tool = {
"type": "mcp",
"server_label": "zapier",
"server_url": "https://mcp.zapier.com/api/mcp/mcp",
"allowed_tools": ["zapier_send_email"],
}
async def served_names(names):
assert names == {"mcp"}
return frozenset()
narrowed = await SemanticToolFilterHook._narrow_mcp_references(
[gateway_reference, external_tool], ["srv-tool_1"], served_names=served_names
)
assert narrowed == [{**gateway_reference, "allowed_tools": ["srv-tool_1"]}, external_tool]

View file

@ -1,6 +1,13 @@
import json
import sys
import types
import pytest
import respx
from httpx import Response
from unittest.mock import AsyncMock, patch
import litellm
from litellm.types.utils import ModelResponse
from litellm.responses.mcp import chat_completions_handler
@ -1344,3 +1351,39 @@ async def test_acompletion_with_mcp_streaming_drains_inner_stream_after_exhausti
assert len(all_chunks) == 3
assert initial_stream.drained_after_exhaustion is True
@pytest.mark.asyncio
@respx.mock
async def test_acompletion_with_mcp_forwards_unserved_external_mcp_tool_to_the_provider(monkeypatch):
from litellm.proxy._experimental.mcp_server.mcp_server_manager import global_mcp_server_manager
zapier_tool = {"type": "mcp", "server_label": "zapier", "server_url": "https://mcp.zapier.com/api/mcp/mcp"}
monkeypatch.setitem(sys.modules, "litellm.proxy.proxy_server", types.SimpleNamespace(prisma_client=None))
monkeypatch.setattr(global_mcp_server_manager, "get_registry", lambda: {})
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
provider = respx.post("https://api.openai.com/v1/chat/completions").mock(
return_value=Response(
200,
json={
"id": "chatcmpl-zapier",
"object": "chat.completion",
"created": 0,
"model": "gpt-4.1",
"choices": [{"index": 0, "message": {"role": "assistant", "content": "ok"}, "finish_reason": "stop"}],
"usage": {"prompt_tokens": 1, "completion_tokens": 1, "total_tokens": 2},
},
)
)
result = await acompletion_with_mcp(
model="openai/gpt-4.1",
messages=[{"role": "user", "content": "hello"}],
tools=[zapier_tool],
api_key="sk-test",
acompletion=True,
)
assert isinstance(result, ModelResponse)
assert result.id == "chatcmpl-zapier"
assert json.loads(provider.calls.last.request.content)["tools"] == [zapier_tool]

View file

@ -6,6 +6,7 @@ from unittest.mock import AsyncMock, MagicMock
import pytest
from fastapi import HTTPException
from openai.types.responses.tool_param import Mcp
import importlib
from litellm.proxy._experimental.mcp_server.faults.list_outcomes import AggregateToolListing
@ -724,6 +725,178 @@ def test_extract_tool_call_details_still_prefers_openai_arguments():
assert arguments == '{"city": "Paris"}'
def _registered(
server_id: str,
name: str,
alias: str | None = None,
server_name: str | None = None,
access_groups: list[str] | None = None,
):
from litellm.types.mcp import MCPTransport
from litellm.types.mcp_server.mcp_server_manager import MCPServer
return MCPServer(
server_id=server_id,
name=name,
alias=alias,
server_name=server_name,
transport=MCPTransport.http,
access_groups=access_groups,
)
async def _no_toolset(_: str) -> bool:
return False
ZAPIER_TOOL: Mcp = {
"type": "mcp",
"server_label": "zapier",
"server_url": "https://mcp.zapier.com/api/mcp/mcp",
"require_approval": "never",
}
EXPLICIT_GATEWAY_TOOL = {"type": "mcp", "server_label": "github", "server_url": "litellm_proxy/mcp/github"}
FUNCTION_TOOL = {"type": "function", "name": "get_weather", "parameters": {}}
@pytest.mark.asyncio
async def test_gateway_served_names_matches_alias_server_name_name_access_group_and_toolset():
from litellm.responses.mcp.litellm_proxy_mcp_handler import _gateway_served_names
servers = (
_registered("id-1", "github-name", alias="github", server_name="github-server", access_groups=["prod-group"]),
_registered("id-2", "deepwiki"),
)
async def toolset_exists(name: str) -> bool:
return name == "my-toolset"
served = await _gateway_served_names(
{"github", "github-server", "github-name", "deepwiki", "prod-group", "my-toolset", "mcp", "nope"},
servers=lambda: servers,
toolset_exists=toolset_exists,
)
assert served == {"github", "github-server", "github-name", "deepwiki", "prod-group", "my-toolset"}
@pytest.mark.asyncio
async def test_gateway_served_names_matches_server_id_short_prefix_and_alias_case_like_the_gateway():
from litellm.proxy._experimental.mcp_server.utils import compute_short_server_prefix
from litellm.responses.mcp.litellm_proxy_mcp_handler import _gateway_served_names
server_id = "0b9ae4ca-1bd2-4faa-b183-7dd812597e3b"
short_prefix = compute_short_server_prefix(server_id)
servers = (_registered(server_id, "github-name", alias="github", access_groups=["prod-group"]),)
served = await _gateway_served_names(
{server_id, short_prefix, "GitHub", "PROD-GROUP", "nope"}, servers=lambda: servers, toolset_exists=_no_toolset
)
assert served == {server_id, short_prefix, "GitHub"}
@pytest.mark.asyncio
async def test_routes_through_gateway_flags_explicit_and_served_tools_only():
served_tool = {"type": "mcp", "server_label": "github", "server_url": "http://localhost:4000/mcp/github"}
async def served_names(names):
assert names == {"github", "mcp"}
return frozenset({"github"})
flags = await LiteLLM_Proxy_MCP_Handler.routes_through_gateway(
[ZAPIER_TOOL, EXPLICIT_GATEWAY_TOOL, served_tool, FUNCTION_TOOL], served_names=served_names
)
assert flags == (False, True, True, False)
@pytest.mark.asyncio
async def test_split_mcp_tools_leaves_external_mcp_path_urls_for_the_provider():
async def served_names(names):
assert names == {"mcp"}
return frozenset()
gateway_tools, other_tools = await LiteLLM_Proxy_MCP_Handler._split_mcp_tools(
[ZAPIER_TOOL, EXPLICIT_GATEWAY_TOOL, FUNCTION_TOOL], served_names=served_names
)
assert gateway_tools == [EXPLICIT_GATEWAY_TOOL]
assert other_tools == [ZAPIER_TOOL, FUNCTION_TOOL]
@pytest.mark.asyncio
async def test_split_mcp_tools_repoints_served_proxy_urls_at_the_gateway():
served_tool = {
"type": "mcp",
"server_label": "toolset",
"server_url": "http://localhost:4000/mcp/my-toolset",
"require_approval": "never",
"allowed_tools": ["get_me"],
}
unserved_tool = {"type": "mcp", "server_label": "typo", "server_url": "http://localhost:4000/mcp/githb"}
async def served_names(names):
return frozenset({"my-toolset"})
gateway_tools, other_tools = await LiteLLM_Proxy_MCP_Handler._split_mcp_tools(
[served_tool, unserved_tool], served_names=served_names
)
assert gateway_tools == [{**served_tool, "server_url": "litellm_proxy/mcp/my-toolset"}]
assert other_tools == [unserved_tool]
@pytest.mark.asyncio
async def test_split_mcp_tools_skips_resolution_when_nothing_points_at_the_proxy():
async def served_names(names):
raise AssertionError("no lookup expected")
gateway_tools, other_tools = await LiteLLM_Proxy_MCP_Handler._split_mcp_tools(
[EXPLICIT_GATEWAY_TOOL, FUNCTION_TOOL], served_names=served_names
)
assert gateway_tools == [EXPLICIT_GATEWAY_TOOL]
assert other_tools == [FUNCTION_TOOL]
def test_should_use_gateway_still_triggers_on_http_mcp_path():
assert LiteLLM_Proxy_MCP_Handler._should_use_litellm_mcp_gateway([ZAPIER_TOOL]) is True
assert LiteLLM_Proxy_MCP_Handler._should_use_litellm_mcp_gateway([EXPLICIT_GATEWAY_TOOL]) is True
assert LiteLLM_Proxy_MCP_Handler._should_use_litellm_mcp_gateway([FUNCTION_TOOL]) is False
assert LiteLLM_Proxy_MCP_Handler._should_use_litellm_mcp_gateway(None) is False
@pytest.mark.asyncio
async def test_aresponses_api_with_mcp_forwards_unserved_external_mcp_tool_to_the_provider(monkeypatch):
from litellm.proxy._experimental.mcp_server.mcp_server_manager import global_mcp_server_manager
from litellm.responses import main as responses_main
from litellm.types.llms.openai import ResponsesAPIResponse
monkeypatch.setitem(sys.modules, "litellm.proxy.proxy_server", types.SimpleNamespace(prisma_client=None))
monkeypatch.setattr(global_mcp_server_manager, "get_registry", lambda: {})
provider_tools: list[object] = []
def fake_provider(**kwargs: object) -> object:
request_params = cast(dict[str, object], kwargs["response_api_optional_request_params"])
provider_tools.append(request_params.get("tools"))
async def respond() -> ResponsesAPIResponse:
return ResponsesAPIResponse(id="resp_zapier", created_at=0, output=[])
return respond()
monkeypatch.setattr(responses_main.base_llm_http_handler, "response_api_handler", fake_provider)
response = await responses_main.aresponses_api_with_mcp(
input="Reply with the single word ok.", model="openai/gpt-4.1", tools=[ZAPIER_TOOL]
)
assert isinstance(response, ResponsesAPIResponse)
assert provider_tools == [[ZAPIER_TOOL]]
def _response_with_reasoning_and_tool_call() -> Any:
"""A first-turn response as a reasoning model returns it: reasoning item, then a function call."""
return ResponsesAPIResponse(

View file

@ -332,6 +332,67 @@ async def test_per_request_enable_prompt_caching_reaches_the_affinity_key(monkey
assert filtered == [deployments[1]]
@pytest.mark.asyncio
async def test_claude_code_one_shot_subagent_does_not_reuse_an_auto_injected_affinity_key(monkeypatch):
monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", True)
cache = DualCache()
check = PromptCachingDeploymentCheck(cache=cache)
deployments = _deployments(AUTO_CACHING_MODEL, AUTO_CACHING_MODEL)
messages = cast(List[AllMessageValues], [{"role": "user", "content": "unique " * 3000}])
request_kwargs = {
"system": [
{
"type": "text",
"text": "x-anthropic-billing-header: cc_version=2.1.263; cc_is_subagent=true;",
}
],
"proxy_server_request": {"headers": {"user-agent": "claude-cli/2.1.263 (external, cli)"}},
}
auto_injected_messages = AnthropicCacheControlHook.messages_with_default_injections(
messages=messages,
models=(AUTO_CACHING_MODEL,),
)
assert auto_injected_messages != messages
await PromptCachingCache(cache=cache).async_add_model_id(
model_id="dep-2", messages=auto_injected_messages, tools=None
)
filtered = await check.async_filter_deployments(
model=MODEL_GROUP_ALIAS,
healthy_deployments=deployments,
messages=messages,
request_kwargs=request_kwargs,
)
assert filtered == deployments
@pytest.mark.asyncio
async def test_root_cache_control_does_not_reuse_an_auto_injected_affinity_key(monkeypatch):
monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", True)
cache = DualCache()
check = PromptCachingDeploymentCheck(cache=cache)
deployments = _deployments(AUTO_CACHING_MODEL, AUTO_CACHING_MODEL)
messages = _auto_caching_messages()
auto_injected_messages = AnthropicCacheControlHook.messages_with_default_injections(
messages=messages,
models=(AUTO_CACHING_MODEL,),
)
assert auto_injected_messages != messages
await PromptCachingCache(cache=cache).async_add_model_id(
model_id="dep-2", messages=auto_injected_messages, tools=None
)
filtered = await check.async_filter_deployments(
model=MODEL_GROUP_ALIAS,
healthy_deployments=deployments,
messages=messages,
request_kwargs={"cache_control": {"type": "ephemeral"}},
)
assert filtered == deployments
@pytest.mark.asyncio
async def test_tool_marked_cache_control_keeps_routing_off_another_requests_prefix(monkeypatch, local_model_cost_map):
"""

View file

@ -313,6 +313,26 @@ export const buildAgentDataFromForm = (values: any, existingAgent?: any) => {
return agentData;
};
export const parseMcpPermissionsForForm = (agent: any) => ({
allowed_mcp_servers_and_groups: {
servers: agent.object_permission?.mcp_servers ?? [],
accessGroups: agent.object_permission?.mcp_access_groups ?? [],
toolsets: agent.object_permission?.mcp_toolsets ?? [],
},
mcp_tool_permissions: agent.object_permission?.mcp_tool_permissions ?? {},
});
/**
* Always includes every MCP key (empty when cleared) so removals persist;
* the proxy merges object_permission per key, leaving non-MCP grants untouched.
*/
export const buildMcpObjectPermission = (values: any) => ({
mcp_servers: values.allowed_mcp_servers_and_groups?.servers ?? [],
mcp_access_groups: values.allowed_mcp_servers_and_groups?.accessGroups ?? [],
mcp_toolsets: values.allowed_mcp_servers_and_groups?.toolsets ?? [],
mcp_tool_permissions: values.mcp_tool_permissions ?? {},
});
/**
* Parse agent data for form fields
*/
@ -356,5 +376,6 @@ export const parseAgentForForm = (agent: any) => {
: [],
// extra_headers: already an array of strings
extra_headers: agent.extra_headers ?? [],
...parseMcpPermissionsForForm(agent),
};
};

View file

@ -1,4 +1,5 @@
import React from "react";
import { QueryClient, QueryClientProvider } from "@tanstack/react-query";
import { fireEvent, render, screen, waitFor } from "@testing-library/react";
import userEvent, { PointerEventsCheckLevel } from "@testing-library/user-event";
import { describe, it, expect, vi, beforeEach } from "vitest";
@ -10,6 +11,12 @@ vi.mock("@/components/networking", () => ({
getAgentInfo: vi.fn(),
patchAgentCall: vi.fn(),
getAgentCreateMetadata: vi.fn(),
getProxyBaseUrl: vi.fn(() => ""),
getUiConfig: vi.fn(async () => ({})),
fetchMCPServers: vi.fn(async () => []),
fetchMCPAccessGroups: vi.fn(async () => []),
fetchMCPToolsets: vi.fn(async () => []),
listMCPTools: vi.fn(async () => ({ tools: [] })),
}));
vi.mock("@/app/(dashboard)/hooks/keys/useKeys", () => ({
@ -111,7 +118,14 @@ const bedrockAgentcoreInfo: AgentCreateInfo = {
const setup = () => userEvent.setup({ pointerEventsCheck: PointerEventsCheckLevel.Never });
const renderView = () => render(<AgentInfoView agentId="agent-1" onClose={vi.fn()} accessToken="tok" isAdmin={true} />);
const renderView = () => {
const queryClient = new QueryClient({ defaultOptions: { queries: { retry: false } } });
return render(
<QueryClientProvider client={queryClient}>
<AgentInfoView agentId="agent-1" onClose={vi.fn()} accessToken="tok" isAdmin={true} />
</QueryClientProvider>,
);
};
const openEditor = async (user: ReturnType<typeof setup>) => {
await user.click(await screen.findByRole("tab", { name: "Settings" }));
@ -161,6 +175,7 @@ describe("AgentInfoView update payload", () => {
rpm_limit: 222,
session_tpm_limit: 333,
session_rpm_limit: 444,
object_permission: { mcp_servers: [], mcp_access_groups: [], mcp_toolsets: [], mcp_tool_permissions: {} },
});
});
@ -201,6 +216,7 @@ describe("AgentInfoView update payload", () => {
rpm_limit: 222,
session_tpm_limit: 333,
session_rpm_limit: 444,
object_permission: { mcp_servers: [], mcp_access_groups: [], mcp_toolsets: [], mcp_tool_permissions: {} },
});
});
@ -278,9 +294,30 @@ describe("AgentInfoView update payload", () => {
api_base: "https://other.example.com",
model: "langgraph/asst_1",
},
object_permission: { mcp_servers: [], mcp_access_groups: [], mcp_toolsets: [], mcp_tool_permissions: {} },
});
});
it("keeps the agent's existing MCP grants in the update payload", async () => {
const existingMcpGrants = {
mcp_servers: ["srv-1"],
mcp_access_groups: ["grp-a"],
mcp_toolsets: ["toolset-1"],
mcp_tool_permissions: { "srv-1": ["tool_x"] },
};
vi.mocked(networking.getAgentInfo).mockResolvedValue({
...A2A_AGENT,
object_permission: existingMcpGrants,
} as never);
const user = setup();
renderView();
await openEditor(user);
await save(user);
expect(patchedPayload().object_permission).toEqual(existingMcpGrants);
});
it("preserves the full AgentCore runtime ARN (including the resource id after runtime/) across an unedited save", async () => {
vi.mocked(networking.getAgentCreateMetadata).mockResolvedValue([bedrockAgentcoreInfo]);
vi.mocked(networking.getAgentInfo).mockResolvedValue(BEDROCK_AGENTCORE_AGENT as never);

View file

@ -24,6 +24,18 @@ vi.mock("./agent_form_fields", () => ({
unmountedA2AFieldNames: () => [],
}));
vi.mock("@/app/(dashboard)/hooks/mcpServers/useMCPServers", () => ({
useMCPServers: () => ({ data: [{ server_id: "srv-1", server_name: "github" }] }),
}));
vi.mock("@/components/mcp_server_management/MCPServerSelector", () => ({
default: () => <div data-testid="mcp-server-selector" />,
}));
vi.mock("@/components/mcp_server_management/MCPToolPermissions", () => ({
default: () => <div data-testid="mcp-tool-permissions" />,
}));
const agent = {
agent_id: "agent-1",
agent_name: "support-agent",
@ -62,5 +74,18 @@ describe("AgentInfoView settings", () => {
expect(token).toBe("sk-test");
expect(agentId).toBe("agent-1");
expect(payload.tpm_limit).toBe(42);
const clearedMcpGrants = { mcp_servers: [], mcp_access_groups: [], mcp_toolsets: [], mcp_tool_permissions: {} };
expect(payload.object_permission).toEqual(clearedMcpGrants);
});
it("shows MCP grants with server names on the overview tab", async () => {
vi.mocked(networking.getAgentInfo).mockResolvedValue({
...agent,
object_permission: { mcp_servers: ["srv-1"] },
} as unknown as Agent);
render(<AgentInfoView agentId="agent-1" onClose={vi.fn()} accessToken="sk-test" isAdmin={true} />);
expect(await screen.findByText("github (srv-1)")).toBeInTheDocument();
});
});

View file

@ -15,16 +15,27 @@ import { getAgentInfo, patchAgentCall, getAgentCreateMetadata, AgentCreateInfo }
import { Agent } from "@/components/agents/types";
import { KeyResponse } from "@/components/key_team_helpers/key_list";
import { useKeys } from "@/app/(dashboard)/hooks/keys/useKeys";
import { useMCPServers } from "@/app/(dashboard)/hooks/mcpServers/useMCPServers";
import KeyInfoView from "@/components/templates/key_info_view";
import MCPServerSelector from "@/components/mcp_server_management/MCPServerSelector";
import MCPToolPermissions from "@/components/mcp_server_management/MCPToolPermissions";
import AgentVirtualKeys from "./agent_virtual_keys";
import AgentFormFields, { unmountedA2AFieldNames } from "./agent_form_fields";
import DynamicAgentFormFields, { buildDynamicAgentData, unmountedDynamicFieldNames } from "./dynamic_agent_form_fields";
import { AGENT_FORM_CONFIG, buildAgentDataFromForm, parseAgentForForm } from "./agent_config";
import {
AGENT_FORM_CONFIG,
buildAgentDataFromForm,
buildMcpObjectPermission,
parseAgentForForm,
parseMcpPermissionsForForm,
} from "./agent_config";
import {
AgentFormField,
AgentFormValues,
AgentNumberInput,
AgentRequestPayload,
McpServerSelection,
labelWithHint,
omitFieldValues,
useCollapsiblePanels,
} from "./AgentFormKit";
@ -111,7 +122,7 @@ const AgentInfoView: React.FC<AgentInfoViewProps> = ({ agentId, onClose, accessT
} else {
const typeInfo = agentTypeMetadata.find((t) => t.agent_type === agentType);
if (typeInfo) {
form.reset(parseDynamicAgentForForm(data, typeInfo));
form.reset({ ...parseDynamicAgentForForm(data, typeInfo), ...parseMcpPermissionsForForm(data) });
} else {
form.reset(parseAgentForForm(data));
}
@ -131,7 +142,7 @@ const AgentInfoView: React.FC<AgentInfoViewProps> = ({ agentId, onClose, accessT
if (agentType !== "a2a") {
const typeInfo = agentTypeMetadata.find((t) => t.agent_type === agentType);
if (typeInfo) {
form.reset(parseDynamicAgentForForm(agent, typeInfo));
form.reset({ ...parseDynamicAgentForForm(agent, typeInfo), ...parseMcpPermissionsForForm(agent) });
}
}
}
@ -139,6 +150,14 @@ const AgentInfoView: React.FC<AgentInfoViewProps> = ({ agentId, onClose, accessT
const selectedAgentTypeInfo = agentTypeMetadata.find((t) => t.agent_type === detectedAgentType);
const watchedFormValues = useWatch({ control: form.control });
const mcpSelection = useWatch({ control: form.control, name: "allowed_mcp_servers_and_groups" });
const mcpToolPermissions = useWatch({ control: form.control, name: "mcp_tool_permissions" });
const { data: mcpServers = [] } = useMCPServers();
const mcpServerLabel = (serverId: string) => {
const server = mcpServers.find((s) => s.server_id === serverId);
return server?.server_name ? `${server.server_name} (${serverId})` : serverId;
};
const discoveryRequest = useMemo(
() => buildDiscoveryRequest(detectedAgentType, watchedFormValues || {}, selectedAgentTypeInfo),
@ -199,7 +218,10 @@ const AgentInfoView: React.FC<AgentInfoViewProps> = ({ agentId, onClose, accessT
? overlayDiscoveredCardParams(built, appliedDiscoveredSelection.selected_card)
: built;
await patchAgentCall(accessToken, agentId, updateData);
await patchAgentCall(accessToken, agentId, {
...updateData,
object_permission: buildMcpObjectPermission(values),
});
toast.success("Agent updated successfully");
setIsEditing(false);
fetchAgentInfo();
@ -337,13 +359,20 @@ const AgentInfoView: React.FC<AgentInfoViewProps> = ({ agentId, onClose, accessT
{agent.object_permission &&
(agent.object_permission.mcp_servers?.length ||
agent.object_permission.mcp_access_groups?.length ||
agent.object_permission.mcp_toolsets?.length ||
(agent.object_permission.mcp_tool_permissions &&
Object.keys(agent.object_permission.mcp_tool_permissions).length > 0)) && (
<div style={{ marginTop: 24 }}>
<h3 className="text-lg font-medium">MCP Tool Permissions</h3>
<DetailList className="mt-4">
{agent.object_permission.mcp_servers && agent.object_permission.mcp_servers.length > 0 && (
<DetailItem label="MCP Servers">{agent.object_permission.mcp_servers.join(", ")}</DetailItem>
<DetailItem label="MCP Servers">
<div className="space-y-1">
{agent.object_permission.mcp_servers.map((serverId) => (
<div key={serverId}>{mcpServerLabel(serverId)}</div>
))}
</div>
</DetailItem>
)}
{agent.object_permission.mcp_access_groups &&
agent.object_permission.mcp_access_groups.length > 0 && (
@ -351,13 +380,16 @@ const AgentInfoView: React.FC<AgentInfoViewProps> = ({ agentId, onClose, accessT
{agent.object_permission.mcp_access_groups.join(", ")}
</DetailItem>
)}
{agent.object_permission.mcp_toolsets && agent.object_permission.mcp_toolsets.length > 0 && (
<DetailItem label="MCP Toolsets">{agent.object_permission.mcp_toolsets.join(", ")}</DetailItem>
)}
{agent.object_permission.mcp_tool_permissions &&
Object.keys(agent.object_permission.mcp_tool_permissions).length > 0 && (
<DetailItem label="Tool permissions per server">
<div className="space-y-1">
{Object.entries(agent.object_permission.mcp_tool_permissions).map(([serverId, tools]) => (
<div key={serverId}>
<span className="font-medium">{serverId}:</span>{" "}
<span className="font-medium">{mcpServerLabel(serverId)}:</span>{" "}
{Array.isArray(tools) ? tools.join(", ") : String(tools)}
</div>
))}
@ -457,6 +489,41 @@ const AgentInfoView: React.FC<AgentInfoViewProps> = ({ agentId, onClose, accessT
{rateLimitField("session_rpm_limit", "Session RPM Limit")}
</div>
<Separator className="my-6" />
<h3 className="text-lg font-medium mb-4">MCP Servers</h3>
<FieldGroup>
<AgentFormField
name="allowed_mcp_servers_and_groups"
label={labelWithHint(
"Allowed MCP Servers",
"Select which MCP servers or access groups this agent can access. Keys bound to this agent can only reach servers granted here.",
)}
>
{({ value, onChange }) => (
<MCPServerSelector
onChange={onChange}
value={{
servers: (value as McpServerSelection | undefined)?.servers ?? [],
accessGroups: (value as McpServerSelection | undefined)?.accessGroups ?? [],
toolsets: (value as McpServerSelection | undefined)?.toolsets ?? [],
}}
accessToken={accessToken ?? ""}
placeholder="Select MCP servers or access groups (optional)"
/>
)}
</AgentFormField>
</FieldGroup>
<div className="mt-4">
<MCPToolPermissions
accessToken={accessToken ?? ""}
selectedServers={mcpSelection?.servers ?? []}
toolPermissions={mcpToolPermissions ?? {}}
onChange={(toolPerms: Record<string, string[]>) =>
form.setValue("mcp_tool_permissions", toolPerms)
}
/>
</div>
<div className="mt-6 flex justify-end gap-2">
<Button
type="button"

View file

@ -6243,6 +6243,7 @@ export const patchAgentCall = async (
agent_name?: string;
litellm_params?: Record<string, any>;
agent_card_params?: Record<string, any>;
object_permission?: Record<string, any>;
tpm_limit?: number | null;
rpm_limit?: number | null;
session_tpm_limit?: number | null;
@ -6878,13 +6879,14 @@ export const buildMcpOAuthAuthorizeUrl = ({
const base = getProxyBaseUrl();
const normalizedServerId = encodeURIComponent(serverId.trim());
const url = `${base}/v1/mcp/server/oauth/${normalizedServerId}/authorize`;
const params = new URLSearchParams({
const authorizeParams = {
redirect_uri: redirectUri,
state,
response_type: "code",
code_challenge: codeChallenge,
code_challenge_method: "S256",
});
};
const params = new URLSearchParams(authorizeParams);
if (clientId && clientId.trim().length > 0) {
params.set("client_id", clientId);
}

View file

@ -23089,6 +23089,8 @@ export interface components {
mcp_tool_permissions?: {
[key: string]: string[];
} | null;
/** Mcp Toolsets */
mcp_toolsets?: string[] | null;
/** Models */
models?: string[] | null;
};