This commit is contained in:
Onat Özmen 2026-09-28 12:57:19 -04:00 • committed by GitHub
commit fb9ed21a92
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
4 changed files with 232 additions and 6 deletions

View file

@ -178,6 +178,8 @@ __all__ = (
"_prefetch_oauth_creds_for_user",
"_prepare_mcp_server_headers",
"_raise_if_initialize_grants_no_mcp_servers",
"_request_tags_from_raw_headers",
"_request_tags_header",
"_resolve_display_name_to_original",
"_run_post_mcp_call_guardrails",
"_server_answers_to",
@ -214,6 +216,34 @@ def _mcp_session_id_from_headers(
return None
def _request_tags_header(
raw_headers: Mapping[str, str] | None,
) -> str | None:
"""The caller's ``x-litellm-tags`` value, read case-insensitively like the other header
lookups in this module. ``None`` when the caller sent no tags."""
if not raw_headers:
return None
for key, value in raw_headers.items():
if key.lower() == "x-litellm-tags":
return value or None
return None
def _request_tags_from_raw_headers(
raw_headers: Mapping[str, str] | None,
) -> Sequence[str] | None:
"""The caller's tags, parsed by the same helper the LLM routes use so an MCP operation and a
chat completion attribute an identical header identically."""
header_value: Final = _request_tags_header(raw_headers)
if header_value is None:
return None
return LiteLLMProxyRequestSetup.add_request_tag_to_metadata(
llm_router=None,
headers={"x-litellm-tags": header_value}, # mutable-ok: the shared parser reads a plain dict
data={}, # mutable-ok: no request body to read tags from on this path
)
class ListMCPToolsRestAPIResponseObject(MCPTool):
"""
Object returned by the /tools/list REST API route.
@ -989,6 +1019,11 @@ async def _get_tools_from_mcp_servers(
list_tools_call_id: Final = str(uuid.uuid4())
# Derive trace_id from raw_headers when not explicitly passed (same as A2A / MCP call_tool)
effective_litellm_trace_id: Final = litellm_trace_id or get_chain_id_from_headers(raw_headers)
# An explicit [] means the caller resolved to no tags; only fall back to the
# header when nothing was passed at all.
effective_request_tags: Final = (
request_tags if request_tags is not None else _request_tags_from_raw_headers(raw_headers)
)
spend_logs_metadata: Final[dict[str, object]] = {
"mcp_operation": "list_tools",
}
@ -1005,7 +1040,7 @@ async def _get_tools_from_mcp_servers(
"metadata": {
"spend_logs_metadata": spend_logs_metadata,
"headers": logging_safe_mcp_headers(raw_headers),
**({"tags": request_tags} if request_tags else {}),
**({"tags": effective_request_tags} if effective_request_tags else {}),
},
# Provide a small input payload for standard logging
"input": [

View file

@ -14,7 +14,6 @@ from pydantic import BaseModel, TypeAdapter, ValidationError
from litellm._logging import verbose_proxy_logger
from litellm.caching.redis_cache import RedisCache
from litellm.constants import DEFAULT_IN_MEMORY_TTL
from litellm.models.organization import LiteLLM_OrganizationTable
from litellm.models.team import LiteLLM_TeamTableCachedObj
from litellm.models.team_membership import LiteLLM_TeamMembership
@ -190,14 +189,17 @@ def _iter_entries(refs: AuthObjectRefs, management_ttl: float) -> Iterator[_Cach
None,
)
if refs.organization_id is not None:
# Organization entries use the management TTL like every other prefetched object: a
# 5s fuse expires before slow requests reach the getters, sending them back to
# the DB the prefetch was meant to spare.
yield _CacheEntry(
f"org_id:{refs.organization_id}", "organization_row", LiteLLM_OrganizationTable, DEFAULT_IN_MEMORY_TTL
f"org_id:{refs.organization_id}", "organization_row", LiteLLM_OrganizationTable, management_ttl
)
yield _CacheEntry(
f"org_id:{refs.organization_id}:with_budget",
"organization_row",
LiteLLM_OrganizationTable,
DEFAULT_IN_MEMORY_TTL,
management_ttl,
)
if refs.project_id is not None:
yield _CacheEntry(f"project_id:{refs.project_id}", "project_row", LiteLLM_ProjectTableCachedObj, management_ttl)

View file

@ -177,6 +177,7 @@ async def test_mcp_server_tool_call_forwards_client_headers_to_logging(_mcp_requ
raw_headers={
"x-nuid": "nuid-1",
"x-app-id": "app-1",
"x-litellm-tags": "application:orders, service:checkout",
"content-length": "42",
"x-forwarded-for": "9.9.9.9",
},
@ -205,6 +206,7 @@ async def test_mcp_server_tool_call_forwards_client_headers_to_logging(_mcp_requ
assert captured_headers.get("x-nuid") == "nuid-1"
assert captured_headers.get("x-app-id") == "app-1"
assert captured_headers.get("x-litellm-tags") == "application:orders, service:checkout"
assert "content-length" not in captured_headers
assert captured_headers.get("x-forwarded-for") == "1.2.3.4"
@ -5747,6 +5749,193 @@ async def test_get_tools_from_mcp_servers_logs_list_tools_to_spendlogs_when_enab
assert spend_meta["per_server_list_outcomes"] == {"server_a": {"status": "ok", "tool_count": 1}}
@pytest.mark.asyncio
async def test_get_tools_from_mcp_servers_takes_list_tools_tags_from_x_litellm_tags_header():
"""A gateway that stamps `x-litellm-tags` on proxied traffic gets per-application attribution on
LLM routes; list_tools must read the same header so MCP usage is not stuck under the shared key.
Nothing populated `request_tags`, so the header was the only source and it was being dropped."""
try:
from litellm.proxy._experimental.mcp_server.server import (
_get_tools_from_mcp_servers,
)
from litellm.proxy._types import UserAPIKeyAuth
from mcp.types import Tool as MCPTool
except ImportError:
pytest.skip("MCP server not available")
user_auth = UserAPIKeyAuth(api_key="test-key", user_id="test-user")
server_a = MagicMock(name="server_a_obj")
server_a.name = "server_a"
server_a.alias = "server_a"
server_a.server_name = "server_a"
server_a.server_id = "a"
server_a.auth_type = None
server_a.extra_headers = None
tool_1 = MCPTool(name="server_a-tool_1", description="test tool", inputSchema={"type": "object"})
dummy_logging_obj = MagicMock()
dummy_logging_obj.model_call_details = {"metadata": {"spend_logs_metadata": {}}}
dummy_logging_obj.async_success_handler = AsyncMock()
function_setup_kwargs = {}
def _capture_function_setup(*_args, **kwargs):
function_setup_kwargs.update(kwargs)
return dummy_logging_obj, None
with (
patch( # test-quality-ok: server allowlist is a module-level function; the suite has no injection seam
"litellm.proxy._experimental.mcp_server.operations._get_allowed_mcp_servers",
new=AsyncMock(return_value=[server_a]),
),
patch( # test-quality-ok: header prep is a module-level function; the suite has no injection seam
"litellm.proxy._experimental.mcp_server.operations._prepare_mcp_server_headers",
return_value=(None, None),
),
patch( # test-quality-ok: manager is a module-level singleton; patching it is the suite established seam
"litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager",
) as mock_manager,
patch( # test-quality-ok: tool filter is a module-level function; the suite has no injection seam
"litellm.proxy._experimental.mcp_server.operations.filter_tools_by_allowed_tools",
side_effect=lambda tools, _server: tools,
),
patch( # test-quality-ok: permission filter is a module-level async function; the suite has no injection seam
"litellm.proxy._experimental.mcp_server.operations.filter_tools_by_key_team_permissions",
new=AsyncMock(side_effect=lambda tools, **_: tools),
),
patch( # test-quality-ok: logging setup is a module-level function; patched to capture spend-log metadata kwargs
"litellm.proxy._experimental.mcp_server.operations.function_setup",
side_effect=_capture_function_setup,
),
):
mock_manager._get_tools_from_server = AsyncMock(return_value=[tool_1])
listing = await _get_tools_from_mcp_servers(
user_api_key_auth=user_auth,
mcp_auth_header=None,
mcp_servers=["server_a"],
mcp_server_auth_headers=None,
raw_headers={"X-LiteLLM-Tags": "application:orders, service:checkout"},
log_list_tools_to_spendlogs=True,
list_tools_log_source="mcp_protocol",
)
assert listing.tools == [tool_1]
assert function_setup_kwargs["metadata"]["tags"] == ["application:orders", "service:checkout"]
@pytest.mark.asyncio
async def test_get_tools_from_mcp_servers_prefers_explicit_request_tags_over_the_header():
"""`request_tags` is the resolved value a caller passes in; a header must not override it.
An explicit empty list resolves to no tags rather than falling back to the header."""
try:
from litellm.proxy._experimental.mcp_server.server import (
_get_tools_from_mcp_servers,
)
from litellm.proxy._types import UserAPIKeyAuth
from mcp.types import Tool as MCPTool
except ImportError:
pytest.skip("MCP server not available")
user_auth = UserAPIKeyAuth(api_key="test-key", user_id="test-user")
server_a = MagicMock(name="server_a_obj")
server_a.name = "server_a"
server_a.alias = "server_a"
server_a.server_name = "server_a"
server_a.server_id = "a"
server_a.auth_type = None
server_a.extra_headers = None
tool_1 = MCPTool(name="server_a-tool_1", description="test tool", inputSchema={"type": "object"})
dummy_logging_obj = MagicMock()
dummy_logging_obj.model_call_details = {"metadata": {"spend_logs_metadata": {}}}
dummy_logging_obj.async_success_handler = AsyncMock()
function_setup_kwargs = {}
def _capture_function_setup(*_args, **kwargs):
function_setup_kwargs.update(kwargs)
return dummy_logging_obj, None
with (
patch( # test-quality-ok: server allowlist is a module-level function; the suite has no injection seam
"litellm.proxy._experimental.mcp_server.operations._get_allowed_mcp_servers",
new=AsyncMock(return_value=[server_a]),
),
patch( # test-quality-ok: header prep is a module-level function; the suite has no injection seam
"litellm.proxy._experimental.mcp_server.operations._prepare_mcp_server_headers",
return_value=(None, None),
),
patch( # test-quality-ok: manager is a module-level singleton; patching it is the suite established seam
"litellm.proxy._experimental.mcp_server.operations.global_mcp_server_manager",
) as mock_manager,
patch( # test-quality-ok: tool filter is a module-level function; the suite has no injection seam
"litellm.proxy._experimental.mcp_server.operations.filter_tools_by_allowed_tools",
side_effect=lambda tools, _server: tools,
),
patch( # test-quality-ok: permission filter is a module-level async function; the suite has no injection seam
"litellm.proxy._experimental.mcp_server.operations.filter_tools_by_key_team_permissions",
new=AsyncMock(side_effect=lambda tools, **_: tools),
),
patch( # test-quality-ok: logging setup is a module-level function; patched to capture spend-log metadata kwargs
"litellm.proxy._experimental.mcp_server.operations.function_setup",
side_effect=_capture_function_setup,
),
):
mock_manager._get_tools_from_server = AsyncMock(return_value=[tool_1])
await _get_tools_from_mcp_servers(
user_api_key_auth=user_auth,
mcp_auth_header=None,
mcp_servers=["server_a"],
mcp_server_auth_headers=None,
raw_headers={"x-litellm-tags": "from-header"},
log_list_tools_to_spendlogs=True,
list_tools_log_source="mcp_protocol",
request_tags=["explicit"],
)
explicit_metadata = dict(function_setup_kwargs["metadata"])
await _get_tools_from_mcp_servers(
user_api_key_auth=user_auth,
mcp_auth_header=None,
mcp_servers=["server_a"],
mcp_server_auth_headers=None,
raw_headers={"x-litellm-tags": "from-header"},
log_list_tools_to_spendlogs=True,
list_tools_log_source="mcp_protocol",
request_tags=[],
)
empty_metadata = dict(function_setup_kwargs["metadata"])
assert explicit_metadata["tags"] == ["explicit"]
assert "tags" not in empty_metadata
@pytest.mark.parametrize(
"raw_headers, expected",
[
(None, None),
({"mcp-session-id": "abc"}, None),
({"x-litellm-tags": ""}, None),
({"X-LiteLLM-Tags": "application:orders, service:checkout"}, ["application:orders", "service:checkout"]),
],
)
def test_request_tags_from_raw_headers_only_reads_the_tag_header(raw_headers, expected):
"""Only `x-litellm-tags` carries tags, whatever its casing, and an empty value is not a tag.
A request whose headers are unrelated must attribute nothing rather than the first value seen."""
try:
from litellm.proxy._experimental.mcp_server.operations import (
_request_tags_from_raw_headers,
)
except ImportError:
pytest.skip("MCP server not available")
assert _request_tags_from_raw_headers(raw_headers) == expected
@pytest.mark.asyncio
async def test_get_tools_from_mcp_servers_returns_tools_when_success_logging_fails():
"""

View file

@ -181,8 +181,8 @@ async def test_cold_regime_is_one_mget_one_query_and_the_getters_never_touch_io_
assert sets == sorted(
[
f"SET {TEAM_ID}_{USER_ID} ttl=5",
f"SET org_id:{ORG_ID} ttl=5",
f"SET org_id:{ORG_ID}:with_budget ttl=5",
f"SET org_id:{ORG_ID} ttl=60",
f"SET org_id:{ORG_ID}:with_budget ttl=60",
f"SET {USER_ID} ttl=60",
f"SET team_id:{TEAM_ID} ttl=60",
f"SET team_membership:{USER_ID}:{TEAM_ID} ttl=None",