mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
Merge pull request #21982 from BerriAI/litellm_fix_pat_token_mcp
Fix: skip health check for MCP integration with passthrough token auth
This commit is contained in:
commit
6531d01959
7 changed files with 382 additions and 10 deletions
|
|
@ -2463,8 +2463,8 @@ class MCPServerManager:
|
|||
# Check if we should skip health check based on auth configuration
|
||||
should_skip_health_check = False
|
||||
|
||||
# Skip if auth_type is oauth2
|
||||
if server.needs_user_oauth_token:
|
||||
# Skip if server requires per-user authentication (OAuth2 or passthrough auth)
|
||||
if server.requires_per_user_auth:
|
||||
should_skip_health_check = True
|
||||
# Skip if auth_type is not none and authentication_token is missing
|
||||
elif (
|
||||
|
|
|
|||
|
|
@ -65,3 +65,25 @@ class MCPServer(BaseModel):
|
|||
def needs_user_oauth_token(self) -> bool:
|
||||
"""True if this is an OAuth2 server that relies on per-user tokens (no client_credentials)."""
|
||||
return self.auth_type == MCPAuth.oauth2 and not self.has_client_credentials
|
||||
|
||||
@property
|
||||
def requires_per_user_auth(self) -> bool:
|
||||
"""
|
||||
True if this server requires per-user/per-request authentication.
|
||||
This includes:
|
||||
- OAuth2 servers without client credentials
|
||||
- Servers with auth_type=none but extra_headers configured for auth passthrough
|
||||
|
||||
Health checks should be skipped for these servers since they cannot
|
||||
authenticate without user-provided credentials.
|
||||
"""
|
||||
# OAuth2 without client credentials
|
||||
if self.needs_user_oauth_token:
|
||||
return True
|
||||
|
||||
# PAT passthrough: auth_type is none but extra_headers includes auth headers
|
||||
if self.auth_type == MCPAuth.none and self.extra_headers:
|
||||
auth_header_names = {"authorization", "x-api-key", "api-key", "apikey"}
|
||||
return any(h.lower() in auth_header_names for h in self.extra_headers)
|
||||
|
||||
return False
|
||||
|
|
|
|||
|
|
@ -1036,6 +1036,240 @@ class TestMCPServerManager:
|
|||
assert result.status == "healthy"
|
||||
assert result.health_check_error is None
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_health_check_skips_passthrough_auth_with_authorization_header(self):
|
||||
"""Test that health check is skipped for servers with passthrough Authorization header"""
|
||||
manager = MCPServerManager()
|
||||
|
||||
# Mock server with auth_type=none and Authorization in extra_headers (passthrough auth)
|
||||
server = MCPServer(
|
||||
server_id="github-server",
|
||||
name="github-server",
|
||||
transport=MCPTransport.http,
|
||||
auth_type=MCPAuth.none,
|
||||
authentication_token=None,
|
||||
url="http://github-server.com",
|
||||
extra_headers=["Authorization"], # Passthrough auth configured
|
||||
)
|
||||
|
||||
manager.get_mcp_server_by_id = MagicMock(return_value=server)
|
||||
|
||||
# _create_mcp_client should not be called (health check should be skipped)
|
||||
manager._create_mcp_client = AsyncMock()
|
||||
|
||||
# Perform health check
|
||||
result = await manager.health_check_server("github-server")
|
||||
|
||||
# Verify that client was not created (health check was skipped)
|
||||
manager._create_mcp_client.assert_not_called()
|
||||
|
||||
# Verify results
|
||||
assert isinstance(result, LiteLLM_MCPServerTable)
|
||||
assert result.server_id == "github-server"
|
||||
assert result.status == "unknown"
|
||||
assert result.health_check_error is None
|
||||
assert result.last_health_check is not None
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_health_check_skips_passthrough_auth_with_api_key_header(self):
|
||||
"""Test that health check is skipped for servers with passthrough x-api-key header"""
|
||||
manager = MCPServerManager()
|
||||
|
||||
# Mock server with auth_type=none and x-api-key in extra_headers
|
||||
server = MCPServer(
|
||||
server_id="sourcegraph-server",
|
||||
name="sourcegraph-server",
|
||||
transport=MCPTransport.http,
|
||||
auth_type=MCPAuth.none,
|
||||
authentication_token=None,
|
||||
url="http://sourcegraph-server.com",
|
||||
extra_headers=["x-api-key"], # Passthrough auth configured
|
||||
)
|
||||
|
||||
manager.get_mcp_server_by_id = MagicMock(return_value=server)
|
||||
|
||||
# _create_mcp_client should not be called
|
||||
manager._create_mcp_client = AsyncMock()
|
||||
|
||||
# Perform health check
|
||||
result = await manager.health_check_server("sourcegraph-server")
|
||||
|
||||
# Verify that client was not created (health check was skipped)
|
||||
manager._create_mcp_client.assert_not_called()
|
||||
|
||||
# Verify results
|
||||
assert isinstance(result, LiteLLM_MCPServerTable)
|
||||
assert result.server_id == "sourcegraph-server"
|
||||
assert result.status == "unknown"
|
||||
assert result.health_check_error is None
|
||||
assert result.last_health_check is not None
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_health_check_runs_when_no_passthrough_auth(self):
|
||||
"""Test that health check runs normally for servers with auth_type=none but no passthrough headers"""
|
||||
manager = MCPServerManager()
|
||||
|
||||
# Mock server with auth_type=none but no extra_headers (no passthrough auth)
|
||||
server = MCPServer(
|
||||
server_id="public-server",
|
||||
name="public-server",
|
||||
transport=MCPTransport.http,
|
||||
auth_type=MCPAuth.none,
|
||||
authentication_token=None,
|
||||
url="http://public-server.com",
|
||||
extra_headers=None, # No passthrough auth
|
||||
)
|
||||
|
||||
manager.get_mcp_server_by_id = MagicMock(return_value=server)
|
||||
|
||||
# Mock successful client
|
||||
mock_client = AsyncMock()
|
||||
mock_client.run_with_session = AsyncMock(return_value="ok")
|
||||
manager._create_mcp_client = AsyncMock(return_value=mock_client)
|
||||
|
||||
# Perform health check
|
||||
result = await manager.health_check_server("public-server")
|
||||
|
||||
# Verify that client WAS created (health check should run)
|
||||
manager._create_mcp_client.assert_called_once()
|
||||
|
||||
# Verify results
|
||||
assert isinstance(result, LiteLLM_MCPServerTable)
|
||||
assert result.server_id == "public-server"
|
||||
assert result.status == "healthy"
|
||||
assert result.health_check_error is None
|
||||
assert result.last_health_check is not None
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_health_check_runs_when_extra_headers_no_auth(self):
|
||||
"""Test that health check runs when extra_headers exist but don't include auth headers"""
|
||||
manager = MCPServerManager()
|
||||
|
||||
# Mock server with extra_headers but no auth-related headers
|
||||
server = MCPServer(
|
||||
server_id="custom-server",
|
||||
name="custom-server",
|
||||
transport=MCPTransport.http,
|
||||
auth_type=MCPAuth.none,
|
||||
authentication_token=None,
|
||||
url="http://custom-server.com",
|
||||
extra_headers=["X-Custom-Header", "X-Request-ID"], # Non-auth headers
|
||||
)
|
||||
|
||||
manager.get_mcp_server_by_id = MagicMock(return_value=server)
|
||||
|
||||
# Mock successful client
|
||||
mock_client = AsyncMock()
|
||||
mock_client.run_with_session = AsyncMock(return_value="ok")
|
||||
manager._create_mcp_client = AsyncMock(return_value=mock_client)
|
||||
|
||||
# Perform health check
|
||||
result = await manager.health_check_server("custom-server")
|
||||
|
||||
# Verify that client WAS created (health check should run)
|
||||
manager._create_mcp_client.assert_called_once()
|
||||
|
||||
# Verify results
|
||||
assert isinstance(result, LiteLLM_MCPServerTable)
|
||||
assert result.server_id == "custom-server"
|
||||
assert result.status == "healthy"
|
||||
assert result.health_check_error is None
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_requires_per_user_auth_property_oauth2(self):
|
||||
"""Test that requires_per_user_auth returns True for OAuth2 without client credentials"""
|
||||
# OAuth2 without client credentials
|
||||
server = MCPServer(
|
||||
server_id="oauth-server",
|
||||
name="oauth-server",
|
||||
transport=MCPTransport.http,
|
||||
auth_type=MCPAuth.oauth2,
|
||||
url="http://oauth-server.com",
|
||||
client_id=None,
|
||||
client_secret=None,
|
||||
token_url=None,
|
||||
)
|
||||
assert server.requires_per_user_auth is True
|
||||
assert server.needs_user_oauth_token is True
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_requires_per_user_auth_property_oauth2_with_client_creds(self):
|
||||
"""Test that requires_per_user_auth returns False for OAuth2 with client credentials"""
|
||||
# OAuth2 with client credentials
|
||||
server = MCPServer(
|
||||
server_id="oauth-server",
|
||||
name="oauth-server",
|
||||
transport=MCPTransport.http,
|
||||
auth_type=MCPAuth.oauth2,
|
||||
url="http://oauth-server.com",
|
||||
client_id="client-id",
|
||||
client_secret="client-secret",
|
||||
token_url="http://oauth-server.com/token",
|
||||
)
|
||||
assert server.requires_per_user_auth is False
|
||||
assert server.has_client_credentials is True
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_requires_per_user_auth_property_passthrough_auth(self):
|
||||
"""Test that requires_per_user_auth returns True for passthrough auth (auth_type=none + Authorization header)"""
|
||||
# Passthrough auth with Authorization header
|
||||
server = MCPServer(
|
||||
server_id="github-server",
|
||||
name="github-server",
|
||||
transport=MCPTransport.http,
|
||||
auth_type=MCPAuth.none,
|
||||
url="http://github-server.com",
|
||||
extra_headers=["Authorization"],
|
||||
)
|
||||
assert server.requires_per_user_auth is True
|
||||
|
||||
# Passthrough auth with x-api-key header
|
||||
server2 = MCPServer(
|
||||
server_id="sourcegraph-server",
|
||||
name="sourcegraph-server",
|
||||
transport=MCPTransport.http,
|
||||
auth_type=MCPAuth.none,
|
||||
url="http://sourcegraph-server.com",
|
||||
extra_headers=["x-api-key"],
|
||||
)
|
||||
assert server2.requires_per_user_auth is True
|
||||
|
||||
# Passthrough auth with api-key header (case insensitive)
|
||||
server3 = MCPServer(
|
||||
server_id="api-server",
|
||||
name="api-server",
|
||||
transport=MCPTransport.http,
|
||||
auth_type=MCPAuth.none,
|
||||
url="http://api-server.com",
|
||||
extra_headers=["API-Key"],
|
||||
)
|
||||
assert server3.requires_per_user_auth is True
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_requires_per_user_auth_property_no_passthrough(self):
|
||||
"""Test that requires_per_user_auth returns False when no passthrough auth is configured"""
|
||||
# auth_type=none but no extra_headers
|
||||
server = MCPServer(
|
||||
server_id="public-server",
|
||||
name="public-server",
|
||||
transport=MCPTransport.http,
|
||||
auth_type=MCPAuth.none,
|
||||
url="http://public-server.com",
|
||||
extra_headers=None,
|
||||
)
|
||||
assert server.requires_per_user_auth is False
|
||||
|
||||
# auth_type=none with non-auth extra_headers
|
||||
server2 = MCPServer(
|
||||
server_id="custom-server",
|
||||
name="custom-server",
|
||||
transport=MCPTransport.http,
|
||||
auth_type=MCPAuth.none,
|
||||
url="http://custom-server.com",
|
||||
extra_headers=["X-Custom-Header", "X-Request-ID"],
|
||||
)
|
||||
assert server2.requires_per_user_auth is False
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_register_openapi_tools_includes_static_headers(self, tmp_path):
|
||||
"""Ensure OpenAPI-to-MCP tool calls include server.static_headers (Issue #19341)."""
|
||||
|
|
|
|||
|
|
@ -171,6 +171,7 @@ export const MCPServerView: React.FC<MCPServerViewProps> = ({
|
|||
userRole={userRole}
|
||||
userID={userID}
|
||||
serverAlias={mcpServer.alias}
|
||||
extraHeaders={mcpServer.extra_headers}
|
||||
/>
|
||||
</TabPanel>
|
||||
|
||||
|
|
|
|||
|
|
@ -1,12 +1,12 @@
|
|||
import React, { useState } from "react";
|
||||
import { useQuery, useMutation } from "@tanstack/react-query";
|
||||
import { ToolTestPanel } from "./ToolTestPanel";
|
||||
import { MCPTool, MCPToolsViewerProps, MCPContent, CallMCPToolResponse } from "./types";
|
||||
import { MCPTool, MCPToolsViewerProps, MCPContent, CallMCPToolResponse, AUTH_TYPE } from "./types";
|
||||
import { listMCPTools, callMCPTool } from "../networking";
|
||||
|
||||
import { Card, Title, Text } from "@tremor/react";
|
||||
import { RobotOutlined, ToolOutlined, SearchOutlined } from "@ant-design/icons";
|
||||
import { Input } from "antd";
|
||||
import { RobotOutlined, ToolOutlined, SearchOutlined, LockOutlined, KeyOutlined } from "@ant-design/icons";
|
||||
import { Input, Alert, Button as AntdButton } from "antd";
|
||||
|
||||
const MCPToolsViewer = ({
|
||||
serverId,
|
||||
|
|
@ -14,23 +14,50 @@ const MCPToolsViewer = ({
|
|||
auth_type,
|
||||
userRole,
|
||||
userID,
|
||||
serverAlias, // Add serverAlias prop
|
||||
serverAlias,
|
||||
extraHeaders,
|
||||
}: MCPToolsViewerProps) => {
|
||||
const [selectedTool, setSelectedTool] = useState<MCPTool | null>(null);
|
||||
const [toolResult, setToolResult] = useState<MCPContent[] | null>(null);
|
||||
const [toolError, setToolError] = useState<Error | null>(null);
|
||||
const [toolSearchTerm, setToolSearchTerm] = useState("");
|
||||
|
||||
// State for passthrough headers
|
||||
const [passthroughHeaders, setPassthroughHeaders] = useState<Record<string, string>>({});
|
||||
const [showHeaderInput, setShowHeaderInput] = useState(false);
|
||||
|
||||
// Check if this server has extra headers configured
|
||||
const hasExtraHeaders = extraHeaders && extraHeaders.length > 0;
|
||||
|
||||
// Build custom headers for MCP server requests
|
||||
const buildCustomHeaders = () => {
|
||||
if (!serverAlias || !hasExtraHeaders) return undefined;
|
||||
|
||||
const customHeaders: Record<string, string> = {};
|
||||
|
||||
// Add passthrough headers with server-specific prefix
|
||||
Object.entries(passthroughHeaders).forEach(([headerName, headerValue]) => {
|
||||
if (headerValue && headerValue.trim()) {
|
||||
// Format: x-mcp-{alias}-{header_name}
|
||||
const mcpHeaderName = `x-mcp-${serverAlias}-${headerName.toLowerCase()}`;
|
||||
customHeaders[mcpHeaderName] = headerValue;
|
||||
}
|
||||
});
|
||||
|
||||
return Object.keys(customHeaders).length > 0 ? customHeaders : undefined;
|
||||
};
|
||||
|
||||
// Query to fetch MCP tools
|
||||
const {
|
||||
data: mcpToolsResponse,
|
||||
isLoading: isLoadingTools,
|
||||
error: mcpToolsError,
|
||||
refetch: refetchTools,
|
||||
} = useQuery({
|
||||
queryKey: ["mcpTools", serverId],
|
||||
queryKey: ["mcpTools", serverId, passthroughHeaders],
|
||||
queryFn: () => {
|
||||
if (!accessToken) throw new Error("Access Token required");
|
||||
return listMCPTools(accessToken, serverId);
|
||||
return listMCPTools(accessToken, serverId, buildCustomHeaders());
|
||||
},
|
||||
enabled: !!accessToken,
|
||||
staleTime: 30000, // Consider data fresh for 30 seconds
|
||||
|
|
@ -42,7 +69,13 @@ const MCPToolsViewer = ({
|
|||
if (!accessToken) throw new Error("Access Token required");
|
||||
|
||||
try {
|
||||
const result: CallMCPToolResponse = await callMCPTool(accessToken, serverId, args.tool.name, args.arguments);
|
||||
const result: CallMCPToolResponse = await callMCPTool(
|
||||
accessToken,
|
||||
serverId,
|
||||
args.tool.name,
|
||||
args.arguments,
|
||||
{ customHeaders: buildCustomHeaders() }
|
||||
);
|
||||
return result;
|
||||
} catch (error) {
|
||||
throw error;
|
||||
|
|
@ -79,6 +112,80 @@ const MCPToolsViewer = ({
|
|||
<Title className="text-xl font-semibold mb-6 mt-2">MCP Tools</Title>
|
||||
|
||||
<div className="flex flex-col flex-1">
|
||||
{/* Extra Headers Input Section */}
|
||||
{hasExtraHeaders && (
|
||||
<div className="mb-4 p-3 bg-blue-50 border border-blue-200 rounded-lg">
|
||||
<div className="flex items-center justify-between mb-2">
|
||||
<div className="flex items-center">
|
||||
<KeyOutlined className="text-blue-600 mr-2" />
|
||||
<Text className="text-sm font-medium text-blue-800">
|
||||
Additional Headers
|
||||
</Text>
|
||||
</div>
|
||||
<AntdButton
|
||||
size="small"
|
||||
type="link"
|
||||
onClick={() => setShowHeaderInput(!showHeaderInput)}
|
||||
className="text-blue-700 p-0 h-auto"
|
||||
>
|
||||
{showHeaderInput ? "Hide" : "Configure"}
|
||||
</AntdButton>
|
||||
</div>
|
||||
|
||||
{!showHeaderInput && Object.keys(passthroughHeaders).length === 0 && (
|
||||
<Text className="text-xs text-blue-700">
|
||||
This server requires additional headers. Click "Configure" to provide values.
|
||||
</Text>
|
||||
)}
|
||||
|
||||
{showHeaderInput && (
|
||||
<div className="mt-3 space-y-2">
|
||||
{extraHeaders?.map((headerName) => (
|
||||
<div key={headerName}>
|
||||
<label className="block text-xs font-medium text-gray-700 mb-1">
|
||||
{headerName}
|
||||
</label>
|
||||
<Input
|
||||
size="small"
|
||||
placeholder={`Enter ${headerName}`}
|
||||
value={passthroughHeaders[headerName] || ""}
|
||||
onChange={(e) => {
|
||||
setPassthroughHeaders({
|
||||
...passthroughHeaders,
|
||||
[headerName]: e.target.value,
|
||||
});
|
||||
}}
|
||||
prefix={<KeyOutlined className="text-gray-400" />}
|
||||
className="rounded"
|
||||
/>
|
||||
</div>
|
||||
))}
|
||||
<AntdButton
|
||||
size="small"
|
||||
type="primary"
|
||||
onClick={() => {
|
||||
refetchTools();
|
||||
setShowHeaderInput(false);
|
||||
}}
|
||||
disabled={Object.values(passthroughHeaders).every(v => !v || !v.trim())}
|
||||
className="w-full mt-2"
|
||||
>
|
||||
Load Tools
|
||||
</AntdButton>
|
||||
</div>
|
||||
)}
|
||||
|
||||
{!showHeaderInput && Object.keys(passthroughHeaders).length > 0 && (
|
||||
<div className="mt-2">
|
||||
<Text className="text-xs text-green-700 flex items-center">
|
||||
<span className="inline-block w-2 h-2 bg-green-500 rounded-full mr-2"></span>
|
||||
{Object.keys(passthroughHeaders).length} header(s) configured
|
||||
</Text>
|
||||
</div>
|
||||
)}
|
||||
</div>
|
||||
)}
|
||||
|
||||
{/* Tool Selection - Show tools first */}
|
||||
<div className="flex flex-col flex-1 min-h-0">
|
||||
<Text className="font-medium block mb-3 text-gray-700 flex items-center">
|
||||
|
|
|
|||
|
|
@ -132,6 +132,7 @@ export interface MCPToolsViewerProps {
|
|||
userRole: string | null;
|
||||
userID: string | null;
|
||||
serverAlias?: string | null;
|
||||
extraHeaders?: string[] | null;
|
||||
}
|
||||
|
||||
export interface MCPServer {
|
||||
|
|
|
|||
|
|
@ -7147,7 +7147,11 @@ export const testSearchToolConnection = async (accessToken: string, litellmParam
|
|||
}
|
||||
};
|
||||
|
||||
export const listMCPTools = async (accessToken: string, serverId: string) => {
|
||||
export const listMCPTools = async (
|
||||
accessToken: string,
|
||||
serverId: string,
|
||||
customHeaders?: Record<string, string>
|
||||
) => {
|
||||
try {
|
||||
// Construct base URL
|
||||
let url = proxyBaseUrl
|
||||
|
|
@ -7159,6 +7163,7 @@ export const listMCPTools = async (accessToken: string, serverId: string) => {
|
|||
const headers: Record<string, string> = {
|
||||
[globalLitellmHeaderName]: `Bearer ${accessToken}`,
|
||||
"Content-Type": "application/json",
|
||||
...customHeaders, // Merge custom headers for passthrough auth
|
||||
};
|
||||
|
||||
const response = await fetch(url, {
|
||||
|
|
@ -7194,6 +7199,7 @@ export const listMCPTools = async (accessToken: string, serverId: string) => {
|
|||
|
||||
export interface CallMCPToolOptions {
|
||||
guardrails?: string[];
|
||||
customHeaders?: Record<string, string>;
|
||||
}
|
||||
|
||||
export const callMCPTool = async (
|
||||
|
|
@ -7212,6 +7218,7 @@ export const callMCPTool = async (
|
|||
const headers: Record<string, string> = {
|
||||
[globalLitellmHeaderName]: `Bearer ${accessToken}`,
|
||||
"Content-Type": "application/json",
|
||||
...(options?.customHeaders || {}), // Merge custom headers for passthrough auth
|
||||
};
|
||||
|
||||
const body: Record<string, any> = {
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue