mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-08 22:21:35 +00:00
fix(mcp): never forward an Authorization header that satisfied admission on the tools preview
Authorization doubles as the admission fallback when x-litellm-api-key is absent, so a caller who authenticated the preview request that way had their LiteLLM key forwarded to the upstream as the oauth2/client-forwarded token. The preview now forwards Authorization only when the primary admission header is present, which is how the dashboard has always sent it; with no primary header there is no upstream token on the request at all. Applies to oauth2 and both client-forwarded modes; parametrized regression test plus the admission header added to the existing extraction tests to mirror the real UI request shape
This commit is contained in:
parent
43726f2d0b
commit
65d0dcfb82
2 changed files with 71 additions and 75 deletions
|
|
@ -1321,12 +1321,15 @@ if MCP_AVAILABLE:
|
|||
if isinstance(credentials, dict):
|
||||
mcp_auth_header = credentials.get("auth_value")
|
||||
|
||||
# Authorization doubles as the admission fallback (LITELLM_API_KEY_HEADER_NAME_SECONDARY):
|
||||
# when the primary x-litellm-api-key header is absent, the Authorization value is the
|
||||
# caller's LiteLLM key, not an upstream token, and must never be forwarded upstream.
|
||||
oauth2_headers: Optional[Dict[str, str]] = None
|
||||
if new_mcp_server_request.auth_type in {
|
||||
MCPAuth.oauth2,
|
||||
MCPAuth.true_passthrough,
|
||||
MCPAuth.oauth_delegate,
|
||||
}:
|
||||
} and headers.get(MCPRequestHandler.LITELLM_API_KEY_HEADER_NAME_PRIMARY):
|
||||
oauth2_headers = MCPRequestHandler._get_oauth2_headers_from_headers(headers)
|
||||
|
||||
async def _list_tools_operation(client):
|
||||
|
|
|
|||
|
|
@ -37,10 +37,7 @@ def _build_request(
|
|||
body_bytes = body
|
||||
else:
|
||||
body_bytes = b""
|
||||
raw_headers = [
|
||||
(key.lower().encode("latin-1"), value.encode("latin-1"))
|
||||
for key, value in headers.items()
|
||||
]
|
||||
raw_headers = [(key.lower().encode("latin-1"), value.encode("latin-1")) for key, value in headers.items()]
|
||||
scope = {
|
||||
"type": "http",
|
||||
"http_version": "1.1",
|
||||
|
|
@ -62,25 +59,18 @@ def _build_request(
|
|||
|
||||
def _get_route(path: str, method: str):
|
||||
for route in rest_endpoints.router.routes:
|
||||
if getattr(route, "path", None) == path and method in getattr(
|
||||
route, "methods", set()
|
||||
):
|
||||
if getattr(route, "path", None) == path and method in getattr(route, "methods", set()):
|
||||
return route
|
||||
raise AssertionError(f"Route {method} {path} not found")
|
||||
|
||||
|
||||
def _route_has_dependency(route, dependency) -> bool:
|
||||
if any(
|
||||
getattr(dep, "dependency", None) == dependency
|
||||
for dep in getattr(route, "dependencies", [])
|
||||
):
|
||||
if any(getattr(dep, "dependency", None) == dependency for dep in getattr(route, "dependencies", [])):
|
||||
return True
|
||||
dependant = getattr(route, "dependant", None)
|
||||
if dependant is None:
|
||||
return False
|
||||
return any(
|
||||
getattr(dep, "call", None) == dependency for dep in dependant.dependencies
|
||||
)
|
||||
return any(getattr(dep, "call", None) == dependency for dep in dependant.dependencies)
|
||||
|
||||
|
||||
class TestExecuteWithMcpClient:
|
||||
|
|
@ -104,9 +94,7 @@ class TestExecuteWithMcpClient:
|
|||
auth_type=MCPAuth.none,
|
||||
)
|
||||
|
||||
result = await rest_endpoints._execute_with_mcp_client(
|
||||
payload, failing_operation
|
||||
)
|
||||
result = await rest_endpoints._execute_with_mcp_client(payload, failing_operation)
|
||||
|
||||
assert result["status"] == "error"
|
||||
assert "stack_trace" not in result
|
||||
|
|
@ -267,15 +255,10 @@ class TestExecuteWithMcpClient:
|
|||
assert result["status"] == "ok"
|
||||
# The incoming Authorization must be dropped — extra_headers should
|
||||
# contain no oauth2 headers (only static_headers, which are None here).
|
||||
assert (
|
||||
captured["extra_headers"] is None
|
||||
or "Authorization" not in captured["extra_headers"]
|
||||
)
|
||||
assert captured["extra_headers"] is None or "Authorization" not in captured["extra_headers"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_interactive_oauth_resolves_forwarded_token_via_presented_store(
|
||||
self, monkeypatch
|
||||
):
|
||||
async def test_interactive_oauth_resolves_forwarded_token_via_presented_store(self, monkeypatch):
|
||||
"""Interactive authorization_code preview (oauth2, no client credentials): the forwarded
|
||||
just-authorized token is resolved THROUGH the v2 resolver via a one-shot presented store
|
||||
(cred_provider), not the caller-override path. The bare token (Bearer stripped) is the
|
||||
|
|
@ -433,9 +416,7 @@ class TestExecuteWithMcpClient:
|
|||
return None
|
||||
|
||||
async def fake_create_client(*args, **kwargs):
|
||||
raise BaseExceptionGroup(
|
||||
"test group", [RuntimeError("Cancelled via cancel scope")]
|
||||
)
|
||||
raise BaseExceptionGroup("test group", [RuntimeError("Cancelled via cancel scope")])
|
||||
|
||||
monkeypatch.setattr(
|
||||
rest_endpoints.global_mcp_server_manager,
|
||||
|
|
@ -497,9 +478,7 @@ class TestTestToolsList:
|
|||
"message": "Successfully retrieved tools",
|
||||
}
|
||||
|
||||
monkeypatch.setattr(
|
||||
rest_endpoints, "_execute_with_mcp_client", fake_execute, raising=False
|
||||
)
|
||||
monkeypatch.setattr(rest_endpoints, "_execute_with_mcp_client", fake_execute, raising=False)
|
||||
|
||||
oauth_call_counter = {"count": 0}
|
||||
|
||||
|
|
@ -555,9 +534,7 @@ class TestTestToolsList:
|
|||
"message": "Successfully retrieved tools",
|
||||
}
|
||||
|
||||
monkeypatch.setattr(
|
||||
rest_endpoints, "_execute_with_mcp_client", fake_execute, raising=False
|
||||
)
|
||||
monkeypatch.setattr(rest_endpoints, "_execute_with_mcp_client", fake_execute, raising=False)
|
||||
|
||||
oauth_headers = {"Authorization": "Bearer oauth"}
|
||||
oauth_call_counter = {"count": 0}
|
||||
|
|
@ -573,7 +550,7 @@ class TestTestToolsList:
|
|||
raising=False,
|
||||
)
|
||||
|
||||
request = _build_request({"authorization": "Bearer incoming"})
|
||||
request = _build_request({"authorization": "Bearer incoming", "x-litellm-api-key": "sk-admission"})
|
||||
payload = NewMCPServerRequest(
|
||||
server_name="example",
|
||||
url="https://example.com",
|
||||
|
|
@ -627,7 +604,7 @@ class TestTestToolsList:
|
|||
raising=False,
|
||||
)
|
||||
|
||||
request = _build_request({"authorization": "Bearer upstream-token"})
|
||||
request = _build_request({"authorization": "Bearer upstream-token", "x-litellm-api-key": "sk-admission"})
|
||||
payload = NewMCPServerRequest(
|
||||
server_name="example",
|
||||
url="https://example.com",
|
||||
|
|
@ -646,6 +623,48 @@ class TestTestToolsList:
|
|||
assert captured["mcp_auth_header"] is None
|
||||
assert captured["oauth2_headers"] == oauth_headers
|
||||
|
||||
@pytest.mark.parametrize("auth_type", [MCPAuth.oauth2, MCPAuth.true_passthrough, MCPAuth.oauth_delegate])
|
||||
async def test_does_not_forward_authorization_that_satisfied_admission(self, monkeypatch, auth_type):
|
||||
"""Authorization is also the admission fallback: with no x-litellm-api-key on the request,
|
||||
the Authorization value is the caller's LiteLLM key, so forwarding it would send the
|
||||
admission credential to the upstream."""
|
||||
|
||||
captured: dict = {}
|
||||
|
||||
async def fake_execute(
|
||||
request,
|
||||
operation,
|
||||
mcp_auth_header=None,
|
||||
oauth2_headers=None,
|
||||
raw_headers=None,
|
||||
):
|
||||
captured["oauth2_headers"] = oauth2_headers
|
||||
return {
|
||||
"tools": [],
|
||||
"error": None,
|
||||
"message": "Successfully retrieved tools",
|
||||
}
|
||||
|
||||
monkeypatch.setattr(rest_endpoints, "_execute_with_mcp_client", fake_execute, raising=False)
|
||||
|
||||
request = _build_request({"authorization": "Bearer sk-litellm-admission-key"})
|
||||
payload = NewMCPServerRequest(
|
||||
server_name="example",
|
||||
url="https://example.com",
|
||||
auth_type=auth_type,
|
||||
)
|
||||
|
||||
from litellm.proxy._types import LitellmUserRoles
|
||||
|
||||
result = await rest_endpoints.test_tools_list(
|
||||
request,
|
||||
payload,
|
||||
user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN),
|
||||
)
|
||||
|
||||
assert result["message"] == "Successfully retrieved tools"
|
||||
assert captured["oauth2_headers"] is None
|
||||
|
||||
|
||||
class TestListToolsRestAPI:
|
||||
pytestmark = pytest.mark.asyncio
|
||||
|
|
@ -775,9 +794,7 @@ class TestListToolsRestAPI:
|
|||
stub_server = StubServer()
|
||||
captured = {}
|
||||
|
||||
async def fake_get_tools(
|
||||
server, server_auth_header, *args, apply_tool_filters=True, **kwargs
|
||||
):
|
||||
async def fake_get_tools(server, server_auth_header, *args, apply_tool_filters=True, **kwargs):
|
||||
captured["apply_tool_filters"] = apply_tool_filters
|
||||
return ["tool-1"]
|
||||
|
||||
|
|
@ -825,9 +842,7 @@ class TestListToolsRestAPI:
|
|||
assert captured["apply_tool_filters"] is True
|
||||
|
||||
@pytest.mark.parametrize("upstream_status", [401, 403])
|
||||
async def test_upstream_auth_failure_surfaces_status_and_challenge(
|
||||
self, monkeypatch, upstream_status
|
||||
):
|
||||
async def test_upstream_auth_failure_surfaces_status_and_challenge(self, monkeypatch, upstream_status):
|
||||
"""A single-server pass-through request whose upstream rejects the token
|
||||
must surface the upstream status (401 or 403) plus its WWW-Authenticate
|
||||
challenge, not collapse into a 200 ``unexpected_error`` body."""
|
||||
|
|
@ -1415,9 +1430,7 @@ class TestListToolsRestAPI:
|
|||
|
||||
oauth_headers = {"Authorization": "Bearer user-oauth-token"}
|
||||
|
||||
async def fake_get_user_oauth_extra_headers(
|
||||
server, user_api_key_dict, prefetched_creds=None
|
||||
):
|
||||
async def fake_get_user_oauth_extra_headers(server, user_api_key_dict, prefetched_creds=None):
|
||||
return oauth_headers
|
||||
|
||||
captured = {}
|
||||
|
|
@ -1661,9 +1674,7 @@ class TestGetToolsForSingleServer:
|
|||
|
||||
pytestmark = pytest.mark.asyncio
|
||||
|
||||
async def test_filters_tools_by_object_permission_mcp_tool_permissions(
|
||||
self, monkeypatch
|
||||
):
|
||||
async def test_filters_tools_by_object_permission_mcp_tool_permissions(self, monkeypatch):
|
||||
"""Test that tools are filtered by user_api_key_auth.object_permission.mcp_tool_permissions"""
|
||||
from litellm.proxy._experimental.mcp_server.server import MCPServer
|
||||
from litellm.proxy._types import LiteLLM_ObjectPermissionTable
|
||||
|
|
@ -1826,9 +1837,7 @@ class TestGetToolsForSingleServer:
|
|||
# All tools should be returned
|
||||
assert len(result) == 2
|
||||
|
||||
async def test_no_filtering_when_server_not_in_mcp_tool_permissions(
|
||||
self, monkeypatch
|
||||
):
|
||||
async def test_no_filtering_when_server_not_in_mcp_tool_permissions(self, monkeypatch):
|
||||
"""Test that all tools are returned when server is not in mcp_tool_permissions"""
|
||||
from litellm.proxy._experimental.mcp_server.server import MCPServer
|
||||
from litellm.proxy._types import LiteLLM_ObjectPermissionTable
|
||||
|
|
@ -1881,9 +1890,7 @@ class TestGetToolsForSingleServer:
|
|||
# All tools should be returned since server is not in permissions
|
||||
assert len(result) == 2
|
||||
|
||||
async def test_combines_server_allowed_tools_and_object_permission_filters(
|
||||
self, monkeypatch
|
||||
):
|
||||
async def test_combines_server_allowed_tools_and_object_permission_filters(self, monkeypatch):
|
||||
"""Test that both server.allowed_tools and object_permission.mcp_tool_permissions filters are applied"""
|
||||
from litellm.proxy._experimental.mcp_server.server import MCPServer
|
||||
from litellm.proxy._types import LiteLLM_ObjectPermissionTable
|
||||
|
|
@ -2201,9 +2208,7 @@ class TestPreviewOpenAPITools:
|
|||
"paths": {
|
||||
"/repos/{owner}/{repo}/actions/jobs/{job_id}/logs": {
|
||||
"get": {
|
||||
"operationId": (
|
||||
"actions/download-job-logs-for-workflow-run"
|
||||
),
|
||||
"operationId": ("actions/download-job-logs-for-workflow-run"),
|
||||
"summary": "Download job logs",
|
||||
}
|
||||
},
|
||||
|
|
@ -2246,9 +2251,7 @@ class TestPreviewOpenAPITools:
|
|||
names = [t["name"] for t in result["tools"]]
|
||||
anthropic_re = re.compile(r"^[a-zA-Z0-9_-]{1,128}$")
|
||||
for name in names:
|
||||
assert anthropic_re.match(
|
||||
name
|
||||
), f"preview tool name {name!r} violates ^[a-zA-Z0-9_-]+$"
|
||||
assert anthropic_re.match(name), f"preview tool name {name!r} violates ^[a-zA-Z0-9_-]+$"
|
||||
assert "actions_download-job-logs-for-workflow-run" in names
|
||||
assert "pulls_list-files" in names
|
||||
|
||||
|
|
@ -2305,9 +2308,7 @@ class TestPreviewOpenAPITools:
|
|||
|
||||
registered_summary_to_name: dict = {}
|
||||
|
||||
def fake_create_tool_function(
|
||||
path, method, operation, base_url
|
||||
): # noqa: ANN001
|
||||
def fake_create_tool_function(path, method, operation, base_url): # noqa: ANN001
|
||||
def _f():
|
||||
return None
|
||||
|
||||
|
|
@ -2320,9 +2321,7 @@ class TestPreviewOpenAPITools:
|
|||
)
|
||||
|
||||
class _StubRegistry:
|
||||
def register_tool(
|
||||
self, name, description, input_schema, handler
|
||||
): # noqa: ANN001
|
||||
def register_tool(self, name, description, input_schema, handler): # noqa: ANN001
|
||||
registered_summary_to_name[description] = name
|
||||
|
||||
monkeypatch.setattr(
|
||||
|
|
@ -2331,9 +2330,7 @@ class TestPreviewOpenAPITools:
|
|||
_StubRegistry(),
|
||||
)
|
||||
|
||||
openapi_to_mcp_generator.register_tools_from_openapi(
|
||||
spec, base_url="https://example.invalid"
|
||||
)
|
||||
openapi_to_mcp_generator.register_tools_from_openapi(spec, base_url="https://example.invalid")
|
||||
|
||||
assert preview_summary_to_name == registered_summary_to_name, (
|
||||
f"preview {preview_summary_to_name} != "
|
||||
|
|
@ -2361,15 +2358,11 @@ class TestConnectionErrorMessage:
|
|||
assert secret not in message
|
||||
|
||||
def test_connect_error_points_at_reachability(self):
|
||||
message = rest_endpoints._connection_error_message(
|
||||
httpx.ConnectError("All connection attempts failed")
|
||||
)
|
||||
message = rest_endpoints._connection_error_message(httpx.ConnectError("All connection attempts failed"))
|
||||
assert "unreachable" in message.lower()
|
||||
|
||||
def test_timeout_error_message(self):
|
||||
message = rest_endpoints._connection_error_message(
|
||||
httpx.ConnectTimeout("timed out")
|
||||
)
|
||||
message = rest_endpoints._connection_error_message(httpx.ConnectTimeout("timed out"))
|
||||
assert "unreachable" in message.lower()
|
||||
|
||||
def test_http_status_error_includes_status_code(self):
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue