mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
Merge 1c92361284 into e340e546e2
This commit is contained in:
commit
05dd72e7f7
5 changed files with 311 additions and 7 deletions
|
|
@ -244,6 +244,8 @@ mcp.<operation>.<auth_family>.<assertion>
|
|||
operation : list_tools | call_tool | list_resources | read_resource | list_prompts | get_prompt
|
||||
auth_family : none | api_key | bearer | oauth
|
||||
assertion : succeeds | denied_without_permission | persists_across_processes
|
||||
| access_group_scoped | toolset_scoped | scoped | denied_invalid_signature
|
||||
| denied_expired | denied_inactive_user | explicit_header_precedence
|
||||
e.g. mcp.call_tool.oauth.succeeds
|
||||
```
|
||||
|
||||
|
|
|
|||
|
|
@ -135,3 +135,93 @@
|
|||
assertions: [toolset_scoped]
|
||||
source: "user_api_key_auth_mcp.py:2137"
|
||||
rationale: "A key granted a toolset lists exactly the toolset's tools: the rest of the server's catalog stays hidden and every stored name resolves"
|
||||
|
||||
- id: mcp.list_tools.bearer.scoped
|
||||
module: mcp
|
||||
tier: P0
|
||||
operation: list_tools
|
||||
auth_family: bearer
|
||||
assertions: [scoped]
|
||||
source: "user_api_key_auth.py"
|
||||
rationale: "LIT-4506 preserves scoped gateway JWT admission independently of upstream credentials"
|
||||
|
||||
- id: mcp.call_tool.bearer.scoped
|
||||
module: mcp
|
||||
tier: P0
|
||||
operation: call_tool
|
||||
auth_family: bearer
|
||||
assertions: [scoped]
|
||||
source: "user_api_key_auth.py"
|
||||
rationale: "LIT-4506 preserves scoped gateway JWT admission independently of upstream credentials"
|
||||
|
||||
- id: mcp.list_tools.bearer.denied_invalid_signature
|
||||
module: mcp
|
||||
tier: P0
|
||||
operation: list_tools
|
||||
auth_family: bearer
|
||||
assertions: [denied_invalid_signature]
|
||||
source: "user_api_key_auth.py"
|
||||
rationale: "LIT-4506 preserves scoped gateway JWT admission independently of upstream credentials"
|
||||
|
||||
- id: mcp.call_tool.bearer.denied_invalid_signature
|
||||
module: mcp
|
||||
tier: P0
|
||||
operation: call_tool
|
||||
auth_family: bearer
|
||||
assertions: [denied_invalid_signature]
|
||||
source: "user_api_key_auth.py"
|
||||
rationale: "LIT-4506 preserves scoped gateway JWT admission independently of upstream credentials"
|
||||
|
||||
- id: mcp.list_tools.bearer.denied_expired
|
||||
module: mcp
|
||||
tier: P0
|
||||
operation: list_tools
|
||||
auth_family: bearer
|
||||
assertions: [denied_expired]
|
||||
source: "user_api_key_auth.py"
|
||||
rationale: "LIT-4506 preserves scoped gateway JWT admission independently of upstream credentials"
|
||||
|
||||
- id: mcp.call_tool.bearer.denied_expired
|
||||
module: mcp
|
||||
tier: P0
|
||||
operation: call_tool
|
||||
auth_family: bearer
|
||||
assertions: [denied_expired]
|
||||
source: "user_api_key_auth.py"
|
||||
rationale: "LIT-4506 preserves scoped gateway JWT admission independently of upstream credentials"
|
||||
|
||||
- id: mcp.list_tools.bearer.denied_inactive_user
|
||||
module: mcp
|
||||
tier: P0
|
||||
operation: list_tools
|
||||
auth_family: bearer
|
||||
assertions: [denied_inactive_user]
|
||||
source: "user_api_key_auth.py:1722"
|
||||
rationale: "LIT-4506 preserves scoped gateway JWT admission independently of upstream credentials"
|
||||
|
||||
- id: mcp.call_tool.bearer.denied_inactive_user
|
||||
module: mcp
|
||||
tier: P0
|
||||
operation: call_tool
|
||||
auth_family: bearer
|
||||
assertions: [denied_inactive_user]
|
||||
source: "user_api_key_auth.py:1722"
|
||||
rationale: "LIT-4506 preserves scoped gateway JWT admission independently of upstream credentials"
|
||||
|
||||
- id: mcp.list_tools.bearer.explicit_header_precedence
|
||||
module: mcp
|
||||
tier: P0
|
||||
operation: list_tools
|
||||
auth_family: bearer
|
||||
assertions: [explicit_header_precedence]
|
||||
source: "user_api_key_auth.py"
|
||||
rationale: "LIT-4506 preserves scoped gateway JWT admission independently of upstream credentials"
|
||||
|
||||
- id: mcp.call_tool.bearer.explicit_header_precedence
|
||||
module: mcp
|
||||
tier: P0
|
||||
operation: call_tool
|
||||
auth_family: bearer
|
||||
assertions: [explicit_header_precedence]
|
||||
source: "user_api_key_auth.py"
|
||||
rationale: "LIT-4506 preserves scoped gateway JWT admission independently of upstream credentials"
|
||||
|
|
|
|||
|
|
@ -17,12 +17,11 @@ from collections.abc import Mapping
|
|||
from dataclasses import dataclass
|
||||
from typing import Literal
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, Field, RootModel
|
||||
|
||||
from e2e_config import settle_propagation
|
||||
from e2e_http import Headers, NoBody, Result, Success, UnknownApiError, unwrap
|
||||
from e2e_http import AuthHeaders, Headers, NoBody, Result, Success, UnknownApiError, unwrap
|
||||
from models import KeyGenerateBody, McpServerListResponse, McpServerRow, ObjectPermission
|
||||
from proxy_client import ProxyClient
|
||||
from pydantic import BaseModel, ConfigDict, Field, RootModel
|
||||
|
||||
McpToolArg = str | int | float | bool | list[str] | dict[str, str]
|
||||
McpToolArguments = Mapping[str, McpToolArg]
|
||||
|
|
@ -254,10 +253,10 @@ class McpClient:
|
|||
)
|
||||
)
|
||||
|
||||
def list_tools(self, key: str) -> Result[McpToolsListResponse]:
|
||||
def list_tools(self, key: str, *, headers: AuthHeaders | None = None) -> Result[McpToolsListResponse]:
|
||||
return self.proxy.transport.get(
|
||||
"/mcp-rest/tools/list",
|
||||
headers=ApiKeyHeaders(x_litellm_api_key=key),
|
||||
headers=headers if headers is not None else ApiKeyHeaders(x_litellm_api_key=key),
|
||||
params=NoBody(),
|
||||
response_type=McpToolsListResponse,
|
||||
)
|
||||
|
|
@ -308,6 +307,7 @@ class McpClient:
|
|||
server_id: str,
|
||||
name: str,
|
||||
arguments: McpToolArguments,
|
||||
headers: AuthHeaders | None = None,
|
||||
) -> McpCallToolResponse:
|
||||
"""Poll tools/call until the result is not a multi-worker registry miss.
|
||||
|
||||
|
|
@ -318,7 +318,7 @@ class McpClient:
|
|||
deadline = time.monotonic() + self.proxy.poll_timeout
|
||||
last: Result[McpCallToolResponse] | None = None
|
||||
while True:
|
||||
last = self.call_tool(key, server_id=server_id, name=name, arguments=arguments)
|
||||
last = self.call_tool(key, server_id=server_id, name=name, arguments=arguments, headers=headers)
|
||||
if not _is_mcp_not_synced(last, tool_name=name):
|
||||
return unwrap(last)
|
||||
if time.monotonic() >= deadline:
|
||||
|
|
@ -393,10 +393,11 @@ class McpClient:
|
|||
server_id: str,
|
||||
name: str,
|
||||
arguments: McpToolArguments,
|
||||
headers: AuthHeaders | None = None,
|
||||
) -> Result[McpCallToolResponse]:
|
||||
return self.proxy.transport.post(
|
||||
"/mcp-rest/tools/call",
|
||||
headers=ApiKeyHeaders(x_litellm_api_key=key),
|
||||
headers=headers if headers is not None else ApiKeyHeaders(x_litellm_api_key=key),
|
||||
json=McpCallToolBody(
|
||||
name=name, arguments=dict(arguments), server_id=server_id
|
||||
),
|
||||
|
|
|
|||
204
tests/e2e/mcp/test_mcp_jwt_auth_e2e.py
Normal file
204
tests/e2e/mcp/test_mcp_jwt_auth_e2e.py
Normal file
|
|
@ -0,0 +1,204 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import time
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Final, Literal
|
||||
|
||||
import pytest
|
||||
from datadog_mcp import SEARCH_LOGS_TOOL, register_datadog_mcp
|
||||
from e2e_config import DD_SEARCH_FROM, unique_marker
|
||||
from e2e_http import AuthHeaders, UnauthorizedError, unwrap
|
||||
from idp import SHORT_LIVED_CLIENT_ID, Identity, Keycloak, token_claims
|
||||
from lifecycle import ResourceManager
|
||||
from management.management_client import build_client as build_management_client
|
||||
from mcp_client import McpClient
|
||||
from models import KeyGenerateBody, ObjectPermission, TeamUpdateBody, UserScimMetadata, UserUpdateBody
|
||||
|
||||
pytestmark = pytest.mark.e2e
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class GrantedMcp:
|
||||
server_id: str
|
||||
other_server_id: str
|
||||
tool: str
|
||||
token: str = field(repr=False)
|
||||
|
||||
|
||||
def _allowed(client: McpClient, access: GrantedMcp, headers: AuthHeaders) -> None:
|
||||
listing: Final = unwrap(client.list_tools(access.token, headers=headers))
|
||||
assert listing.tool_names_for_server(access.server_id) == frozenset((access.tool,))
|
||||
assert listing.tool_names_for_server(access.other_server_id) == frozenset()
|
||||
result: Final = client.await_call_tool(
|
||||
access.token,
|
||||
server_id=access.server_id,
|
||||
name=access.tool,
|
||||
arguments={"query": f"lit4506-{unique_marker()}", "from": DD_SEARCH_FROM, "to": "now", "max_tokens": 1000},
|
||||
headers=headers,
|
||||
)
|
||||
assert result.is_error is False, result
|
||||
assert result.all_text.strip(), "The successful control must return a real tool result"
|
||||
|
||||
|
||||
def _denied(client: McpClient, access: GrantedMcp, headers: AuthHeaders, reason: str) -> None:
|
||||
listing: Final = client.list_tools(access.token, headers=headers)
|
||||
assert isinstance(listing, UnauthorizedError), listing
|
||||
assert reason in listing.body.lower(), listing.body
|
||||
execution: Final = client.call_tool(
|
||||
access.token,
|
||||
server_id=access.server_id,
|
||||
name=access.tool,
|
||||
arguments={"query": f"lit4506-{unique_marker()}", "from": DD_SEARCH_FROM, "to": "now", "max_tokens": 1000},
|
||||
headers=headers,
|
||||
)
|
||||
assert isinstance(execution, UnauthorizedError), execution
|
||||
assert reason in execution.body.lower(), execution.body
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def granted_mcp(client: McpClient, resources: ResourceManager, jwt_identity: Identity, idp: Keycloak) -> GrantedMcp:
|
||||
server_id: Final = register_datadog_mcp(client, resources)
|
||||
other_server_id: Final = register_datadog_mcp(client, resources)
|
||||
client.await_registered(server_id)
|
||||
client.await_registered(other_server_id)
|
||||
management: Final = build_management_client(client.proxy)
|
||||
management.update_team(
|
||||
TeamUpdateBody(
|
||||
team_id=jwt_identity.group,
|
||||
team_alias=jwt_identity.group,
|
||||
object_permission=ObjectPermission(mcp_servers=[server_id]),
|
||||
)
|
||||
)
|
||||
management.update_user(
|
||||
UserUpdateBody(
|
||||
user_id=jwt_identity.user_id,
|
||||
user_role="internal_user",
|
||||
object_permission=ObjectPermission(mcp_servers=[server_id]),
|
||||
)
|
||||
)
|
||||
stored: Final = management.user_info(jwt_identity.user_id)
|
||||
assert stored.user_info.user_id == jwt_identity.user_id
|
||||
assert stored.user_info.user_role == "internal_user"
|
||||
token: Final = idp.access_token(jwt_identity)
|
||||
assert token_claims(token).sub == jwt_identity.user_id
|
||||
tool: Final = client.await_tool(token, server_id, SEARCH_LOGS_TOOL)
|
||||
return GrantedMcp(server_id, other_server_id, tool, token)
|
||||
|
||||
|
||||
class TestMcpGatewayJwt:
|
||||
@pytest.mark.covers("mcp.list_tools.bearer.scoped", "mcp.call_tool.bearer.scoped")
|
||||
@pytest.mark.parametrize("header", ("authorization", "x_litellm_api_key"))
|
||||
def test_valid_scoped_jwt_lists_and_calls(
|
||||
self, client: McpClient, granted_mcp: GrantedMcp, header: Literal["authorization", "x_litellm_api_key"]
|
||||
) -> None:
|
||||
headers: Final = (
|
||||
AuthHeaders(authorization=f"Bearer {granted_mcp.token}")
|
||||
if header == "authorization"
|
||||
else AuthHeaders.model_validate({"x-litellm-api-key": granted_mcp.token})
|
||||
)
|
||||
_allowed(client, granted_mcp, headers)
|
||||
|
||||
@pytest.mark.covers(
|
||||
"mcp.list_tools.bearer.denied_invalid_signature", "mcp.call_tool.bearer.denied_invalid_signature"
|
||||
)
|
||||
def test_tampered_signature_denies_both_operations(self, client: McpClient, granted_mcp: GrantedMcp) -> None:
|
||||
valid: Final = AuthHeaders(authorization=f"Bearer {granted_mcp.token}")
|
||||
_allowed(client, granted_mcp, valid)
|
||||
header, payload, signature = granted_mcp.token.split(".")
|
||||
flipped: Final = "A" if signature[10] != "A" else "B"
|
||||
tampered: Final = f"{header}.{payload}.{signature[:10]}{flipped}{signature[11:]}"
|
||||
_denied(client, granted_mcp, AuthHeaders(authorization=f"Bearer {tampered}"), "signature")
|
||||
_allowed(client, granted_mcp, valid)
|
||||
|
||||
@pytest.mark.covers("mcp.list_tools.bearer.denied_expired", "mcp.call_tool.bearer.denied_expired")
|
||||
def test_expired_signed_jwt_denies_both_operations(
|
||||
self, client: McpClient, granted_mcp: GrantedMcp, jwt_identity: Identity, idp: Keycloak
|
||||
) -> None:
|
||||
valid: Final = AuthHeaders(authorization=f"Bearer {granted_mcp.token}")
|
||||
_allowed(client, granted_mcp, valid)
|
||||
expiring: Final = idp.access_token(jwt_identity, client_id=SHORT_LIVED_CLIENT_ID)
|
||||
claims: Final = token_claims(expiring)
|
||||
assert claims.sub == jwt_identity.user_id
|
||||
delay: Final = claims.exp - time.time() + 1
|
||||
assert delay <= 5, "Short-lived token configuration or IdP clock drifted"
|
||||
time.sleep(max(0, delay))
|
||||
_denied(client, granted_mcp, AuthHeaders(authorization=f"Bearer {expiring}"), "expired")
|
||||
_allowed(client, granted_mcp, valid)
|
||||
|
||||
@pytest.mark.covers("mcp.list_tools.bearer.denied_inactive_user", "mcp.call_tool.bearer.denied_inactive_user")
|
||||
def test_deactivated_user_cannot_reuse_warm_jwt(
|
||||
self, client: McpClient, granted_mcp: GrantedMcp, jwt_identity: Identity
|
||||
) -> None:
|
||||
headers: Final = AuthHeaders(authorization=f"Bearer {granted_mcp.token}")
|
||||
_allowed(client, granted_mcp, headers)
|
||||
management: Final = build_management_client(client.proxy)
|
||||
management.update_user(
|
||||
UserUpdateBody(
|
||||
user_id=jwt_identity.user_id, user_role="internal_user", metadata=UserScimMetadata(scim_active=False)
|
||||
)
|
||||
)
|
||||
stored: Final = management.user_info(jwt_identity.user_id)
|
||||
assert stored.user_info.metadata is not None and stored.user_info.metadata.scim_active is False
|
||||
_denied(client, granted_mcp, headers, "deactivated")
|
||||
management.update_user(
|
||||
UserUpdateBody(
|
||||
user_id=jwt_identity.user_id, user_role="internal_user", metadata=UserScimMetadata(scim_active=True)
|
||||
)
|
||||
)
|
||||
_allowed(client, granted_mcp, headers)
|
||||
|
||||
@pytest.mark.covers("mcp.list_tools.bearer.denied_inactive_user", "mcp.call_tool.bearer.denied_inactive_user")
|
||||
def test_deactivated_user_is_denied_on_a_cold_jwt(
|
||||
self, client: McpClient, granted_mcp: GrantedMcp, jwt_identity: Identity, idp: Keycloak
|
||||
) -> None:
|
||||
management: Final = build_management_client(client.proxy)
|
||||
management.update_user(
|
||||
UserUpdateBody(
|
||||
user_id=jwt_identity.user_id, user_role="internal_user", metadata=UserScimMetadata(scim_active=False)
|
||||
)
|
||||
)
|
||||
cold: Final = idp.access_token(jwt_identity)
|
||||
assert cold != granted_mcp.token
|
||||
access: Final = GrantedMcp(granted_mcp.server_id, granted_mcp.other_server_id, granted_mcp.tool, cold)
|
||||
_denied(client, access, AuthHeaders(authorization=f"Bearer {cold}"), "deactivated")
|
||||
management.update_user(
|
||||
UserUpdateBody(
|
||||
user_id=jwt_identity.user_id, user_role="internal_user", metadata=UserScimMetadata(scim_active=True)
|
||||
)
|
||||
)
|
||||
_allowed(client, access, AuthHeaders(authorization=f"Bearer {cold}"))
|
||||
|
||||
@pytest.mark.covers(
|
||||
"mcp.list_tools.bearer.explicit_header_precedence", "mcp.call_tool.bearer.explicit_header_precedence"
|
||||
)
|
||||
def test_explicit_gateway_header_wins_without_fallback(
|
||||
self, client: McpClient, granted_mcp: GrantedMcp, resources: ResourceManager
|
||||
) -> None:
|
||||
_allowed(
|
||||
client,
|
||||
granted_mcp,
|
||||
AuthHeaders.model_validate(
|
||||
{"x-litellm-api-key": granted_mcp.token, "authorization": "Bearer invalid-secondary-token"}
|
||||
),
|
||||
)
|
||||
_denied(
|
||||
client,
|
||||
granted_mcp,
|
||||
AuthHeaders.model_validate(
|
||||
{"x-litellm-api-key": "sk-invalid-primary-token", "authorization": f"Bearer {granted_mcp.token}"}
|
||||
),
|
||||
"key",
|
||||
)
|
||||
sibling_key: Final = client.proxy.generate_key(
|
||||
KeyGenerateBody(object_permission=ObjectPermission(mcp_servers=[granted_mcp.other_server_id]))
|
||||
)
|
||||
resources.defer(lambda: client.proxy.delete_key(sibling_key))
|
||||
sibling_tool: Final = client.await_tool(sibling_key, granted_mcp.other_server_id, SEARCH_LOGS_TOOL)
|
||||
sibling: Final = GrantedMcp(granted_mcp.other_server_id, granted_mcp.server_id, sibling_tool, sibling_key)
|
||||
_allowed(
|
||||
client,
|
||||
sibling,
|
||||
AuthHeaders.model_validate(
|
||||
{"x-litellm-api-key": sibling_key, "authorization": f"Bearer {granted_mcp.token}"}
|
||||
),
|
||||
)
|
||||
|
|
@ -1540,9 +1540,15 @@ class UserNewResponse(BaseModel):
|
|||
key: str | None = None
|
||||
|
||||
|
||||
class UserScimMetadata(BaseModel):
|
||||
scim_active: bool | None = None
|
||||
|
||||
|
||||
class UserUpdateBody(BaseModel):
|
||||
user_id: str
|
||||
user_role: UserRole
|
||||
object_permission: ObjectPermission | None = None
|
||||
metadata: UserScimMetadata | None = None
|
||||
|
||||
|
||||
class UserInfoParams(BaseModel):
|
||||
|
|
@ -1553,6 +1559,7 @@ class UserData(BaseModel):
|
|||
user_id: str | None = None
|
||||
user_email: str | None = None
|
||||
user_role: str | None = None
|
||||
metadata: UserScimMetadata | None = None
|
||||
|
||||
|
||||
class UserInfoResponse(BaseModel):
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue