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:
Sameer Kankute 2026-02-24 19:36:08 +05:30 • committed by GitHub
commit 6531d01959
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
7 changed files with 382 additions and 10 deletions

View file

@ -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 (

View file

@ -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

View file

@ -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)."""

View file

@ -171,6 +171,7 @@ export const MCPServerView: React.FC<MCPServerViewProps> = ({
userRole={userRole}
userID={userID}
serverAlias={mcpServer.alias}
extraHeaders={mcpServer.extra_headers}
/>
</TabPanel>

View file

@ -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">

View file

@ -132,6 +132,7 @@ export interface MCPToolsViewerProps {
userRole: string | null;
userID: string | null;
serverAlias?: string | null;
extraHeaders?: string[] | null;
}
export interface MCPServer {

View file

@ -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> = {