diff --git a/litellm/proxy/_experimental/mcp_server/operations.py b/litellm/proxy/_experimental/mcp_server/operations.py index a19246b6e90..4850a327db4 100644 --- a/litellm/proxy/_experimental/mcp_server/operations.py +++ b/litellm/proxy/_experimental/mcp_server/operations.py @@ -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": [ diff --git a/litellm/proxy/auth/auth_object_prefetch.py b/litellm/proxy/auth/auth_object_prefetch.py index 52e26e885c9..503223e7d87 100644 --- a/litellm/proxy/auth/auth_object_prefetch.py +++ b/litellm/proxy/auth/auth_object_prefetch.py @@ -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) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py index ab00ec4da1e..e039e274a92 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py @@ -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(): """ diff --git a/tests/test_litellm/proxy/auth/test_auth_object_prefetch.py b/tests/test_litellm/proxy/auth/test_auth_object_prefetch.py index 0fd0dda3017..58761e270a0 100644 --- a/tests/test_litellm/proxy/auth/test_auth_object_prefetch.py +++ b/tests/test_litellm/proxy/auth/test_auth_object_prefetch.py @@ -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",