mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-12 23:01:41 +00:00
test(mcp): add regression tests for oauth2_flow M2M opt-in behavior
This commit is contained in:
parent
013882e6c8
commit
6f5b217b98
1 changed files with 410 additions and 0 deletions
|
|
@ -1036,6 +1036,241 @@ 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"""
|
||||
# M2M must be opted in explicitly with oauth2_flow="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",
|
||||
oauth2_flow="client_credentials",
|
||||
)
|
||||
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)."""
|
||||
|
|
@ -2073,5 +2308,180 @@ class TestMCPServerManager:
|
|||
assert resolved_server.server_name == "test_server" # server_name matches
|
||||
|
||||
|
||||
class TestMCPServerTimestamps:
|
||||
"""Regression tests: created_at/updated_at must be preserved, not overwritten with datetime.now()."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_build_mcp_server_from_table_preserves_timestamps(self):
|
||||
"""build_mcp_server_from_table must carry created_at and updated_at into MCPServer."""
|
||||
manager = MCPServerManager()
|
||||
|
||||
created = datetime(2024, 1, 15, 10, 0, 0)
|
||||
updated = datetime(2024, 6, 20, 12, 30, 0)
|
||||
|
||||
table_record = LiteLLM_MCPServerTable(
|
||||
server_id="ts-server-1",
|
||||
server_name="ts_server",
|
||||
url="https://example.com/mcp",
|
||||
transport=MCPTransport.http,
|
||||
created_at=created,
|
||||
updated_at=updated,
|
||||
)
|
||||
|
||||
mcp_server = await manager.build_mcp_server_from_table(table_record)
|
||||
|
||||
assert mcp_server.created_at == created
|
||||
assert mcp_server.updated_at == updated
|
||||
|
||||
def test_build_mcp_server_table_preserves_timestamps(self):
|
||||
"""_build_mcp_server_table must use the MCPServer's stored timestamps, not datetime.now()."""
|
||||
manager = MCPServerManager()
|
||||
|
||||
created = datetime(2024, 1, 15, 10, 0, 0)
|
||||
updated = datetime(2024, 6, 20, 12, 30, 0)
|
||||
|
||||
server = MCPServer(
|
||||
server_id="ts-server-2",
|
||||
name="ts_server",
|
||||
url="https://example.com/mcp",
|
||||
transport=MCPTransport.http,
|
||||
created_at=created,
|
||||
updated_at=updated,
|
||||
)
|
||||
|
||||
table = manager._build_mcp_server_table(server)
|
||||
|
||||
assert table.created_at == created
|
||||
assert table.updated_at == updated
|
||||
|
||||
def test_build_mcp_server_table_none_timestamps_when_not_set(self):
|
||||
"""_build_mcp_server_table must return None timestamps when not set on MCPServer."""
|
||||
manager = MCPServerManager()
|
||||
|
||||
server = MCPServer(
|
||||
server_id="ts-server-3",
|
||||
name="ts_server",
|
||||
url="https://example.com/mcp",
|
||||
transport=MCPTransport.http,
|
||||
)
|
||||
|
||||
table = manager._build_mcp_server_table(server)
|
||||
|
||||
assert table.created_at is None
|
||||
assert table.updated_at is None
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_round_trip_timestamps_preserved(self):
|
||||
"""Timestamps survive the full round-trip: LiteLLM_MCPServerTable -> MCPServer -> LiteLLM_MCPServerTable."""
|
||||
manager = MCPServerManager()
|
||||
|
||||
created = datetime(2023, 3, 10, 8, 0, 0)
|
||||
updated = datetime(2023, 9, 5, 16, 45, 0)
|
||||
|
||||
table_record = LiteLLM_MCPServerTable(
|
||||
server_id="ts-server-4",
|
||||
server_name="ts_server_rt",
|
||||
url="https://example.com/mcp",
|
||||
transport=MCPTransport.http,
|
||||
created_at=created,
|
||||
updated_at=updated,
|
||||
)
|
||||
|
||||
mcp_server = await manager.build_mcp_server_from_table(table_record)
|
||||
rebuilt_table = manager._build_mcp_server_table(mcp_server)
|
||||
|
||||
assert rebuilt_table.created_at == created
|
||||
assert rebuilt_table.updated_at == updated
|
||||
|
||||
|
||||
class TestHasClientCredentialsOAuth2Flow:
|
||||
"""
|
||||
Regression tests for the M2M auto-detection bug.
|
||||
|
||||
Before the fix, has_client_credentials returned True whenever
|
||||
client_id + client_secret + token_url were all set, even for
|
||||
interactive OAuth setups (e.g. GitHub Enterprise). This silently
|
||||
dropped user tokens and fetched M2M tokens instead.
|
||||
|
||||
The fix: M2M must be opted in explicitly via oauth2_flow="client_credentials".
|
||||
"""
|
||||
|
||||
def _make_server(self, **kwargs) -> MCPServer:
|
||||
return MCPServer(
|
||||
server_id="test-server",
|
||||
name="test-server",
|
||||
transport=MCPTransport.http,
|
||||
auth_type=MCPAuth.oauth2,
|
||||
url="https://github.example.com/mcp",
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
def test_all_three_fields_set_without_oauth2_flow_is_not_m2m(self):
|
||||
"""
|
||||
GitHub Enterprise regression: client_id + client_secret + token_url
|
||||
should NOT trigger M2M flow unless oauth2_flow is explicitly set.
|
||||
"""
|
||||
server = self._make_server(
|
||||
client_id="gh-client-id",
|
||||
client_secret="gh-client-secret",
|
||||
token_url="https://github.example.com/login/oauth/access_token",
|
||||
)
|
||||
assert server.has_client_credentials is False
|
||||
|
||||
def test_explicit_client_credentials_flow_enables_m2m(self):
|
||||
"""oauth2_flow='client_credentials' opts in to M2M."""
|
||||
server = self._make_server(
|
||||
client_id="svc-client-id",
|
||||
client_secret="svc-client-secret",
|
||||
token_url="https://idp.example.com/token",
|
||||
oauth2_flow="client_credentials",
|
||||
)
|
||||
assert server.has_client_credentials is True
|
||||
|
||||
def test_explicit_authorization_code_flow_disables_m2m(self):
|
||||
"""oauth2_flow='authorization_code' always returns False."""
|
||||
server = self._make_server(
|
||||
client_id="gh-client-id",
|
||||
client_secret="gh-client-secret",
|
||||
token_url="https://github.example.com/login/oauth/access_token",
|
||||
oauth2_flow="authorization_code",
|
||||
)
|
||||
assert server.has_client_credentials is False
|
||||
|
||||
def test_no_fields_no_flow_is_not_m2m(self):
|
||||
"""No credentials configured — not M2M."""
|
||||
server = self._make_server()
|
||||
assert server.has_client_credentials is False
|
||||
|
||||
def test_partial_fields_without_flow_is_not_m2m(self):
|
||||
"""Partial credential fields without explicit flow — not M2M."""
|
||||
server = self._make_server(
|
||||
client_id="only-client-id",
|
||||
)
|
||||
assert server.has_client_credentials is False
|
||||
|
||||
def test_needs_user_oauth_token_true_without_explicit_m2m(self):
|
||||
"""
|
||||
Without oauth2_flow='client_credentials', an oauth2 server with
|
||||
client fields set still needs a user OAuth token (interactive flow).
|
||||
"""
|
||||
server = self._make_server(
|
||||
client_id="gh-client-id",
|
||||
client_secret="gh-client-secret",
|
||||
token_url="https://github.example.com/login/oauth/access_token",
|
||||
)
|
||||
assert server.needs_user_oauth_token is True
|
||||
|
||||
def test_needs_user_oauth_token_false_with_explicit_m2m(self):
|
||||
"""With oauth2_flow='client_credentials', no per-user token needed."""
|
||||
server = self._make_server(
|
||||
client_id="svc-client-id",
|
||||
client_secret="svc-client-secret",
|
||||
token_url="https://idp.example.com/token",
|
||||
oauth2_flow="client_credentials",
|
||||
)
|
||||
assert server.needs_user_oauth_token is False
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
pytest.main([__file__])
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue