mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
Merge 5b26eefdc9 into 6f5ad78a1f
This commit is contained in:
commit
fb9ed21a92
4 changed files with 232 additions and 6 deletions
|
|
@ -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": [
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue