diff --git a/litellm/constants.py b/litellm/constants.py index 6237e67f6af..9c25cf77906 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -84,9 +84,18 @@ MAX_MCP_SEMANTIC_FILTER_TOOLS_HEADER_LENGTH = int( ) # MCP OAuth2 Client Credentials Defaults -MCP_OAUTH2_TOKEN_EXPIRY_BUFFER_SECONDS = 60 -MCP_OAUTH2_TOKEN_CACHE_MAX_SIZE = 200 -MCP_OAUTH2_TOKEN_CACHE_DEFAULT_TTL = 3600 +MCP_OAUTH2_TOKEN_EXPIRY_BUFFER_SECONDS = int( + os.getenv("MCP_OAUTH2_TOKEN_EXPIRY_BUFFER_SECONDS", "60") +) +MCP_OAUTH2_TOKEN_CACHE_MAX_SIZE = int( + os.getenv("MCP_OAUTH2_TOKEN_CACHE_MAX_SIZE", "200") +) +MCP_OAUTH2_TOKEN_CACHE_DEFAULT_TTL = int( + os.getenv("MCP_OAUTH2_TOKEN_CACHE_DEFAULT_TTL", "3600") +) +MCP_OAUTH2_TOKEN_CACHE_MIN_TTL = int( + os.getenv("MCP_OAUTH2_TOKEN_CACHE_MIN_TTL", "10") +) LITELLM_UI_ALLOW_HEADERS = [ "x-litellm-semantic-filter", diff --git a/litellm/proxy/_experimental/mcp_server/oauth2_token_cache.py b/litellm/proxy/_experimental/mcp_server/oauth2_token_cache.py index 10856267fea..ebe71270563 100644 --- a/litellm/proxy/_experimental/mcp_server/oauth2_token_cache.py +++ b/litellm/proxy/_experimental/mcp_server/oauth2_token_cache.py @@ -13,6 +13,7 @@ from litellm.caching.in_memory_cache import InMemoryCache from litellm.constants import ( MCP_OAUTH2_TOKEN_CACHE_DEFAULT_TTL, MCP_OAUTH2_TOKEN_CACHE_MAX_SIZE, + MCP_OAUTH2_TOKEN_CACHE_MIN_TTL, MCP_OAUTH2_TOKEN_EXPIRY_BUFFER_SECONDS, ) from litellm.llms.custom_httpx.http_handler import get_async_httpx_client @@ -101,7 +102,7 @@ class MCPOAuth2TokenCache(InMemoryCache): ) expires_in = int(body.get("expires_in", 3600)) - ttl = max(expires_in - MCP_OAUTH2_TOKEN_EXPIRY_BUFFER_SECONDS, 0) + ttl = max(expires_in - MCP_OAUTH2_TOKEN_EXPIRY_BUFFER_SECONDS, MCP_OAUTH2_TOKEN_CACHE_MIN_TTL) verbose_logger.info( "Fetched OAuth2 token for MCP server %s (expires in %ds)", diff --git a/tests/mcp_tests/test_oauth2_e2e.sh b/tests/mcp_tests/test_oauth2_e2e.sh new file mode 100755 index 00000000000..edeffd93fb7 --- /dev/null +++ b/tests/mcp_tests/test_oauth2_e2e.sh @@ -0,0 +1,67 @@ +#!/usr/bin/env bash +# E2E test for OAuth2 client_credentials MCP flow +# Usage: bash tests/mcp_tests/test_oauth2_e2e.sh +set -euo pipefail + +MOCK_PORT=8765 +PROXY_PORT=4000 +SCRIPT_DIR="$(cd "$(dirname "$0")" && pwd)" +CONFIG="$SCRIPT_DIR/test_oauth2_mcp_config.yaml" +MOCK_SERVER="$SCRIPT_DIR/mock_oauth2_mcp_server.py" + +cleanup() { + echo "" + echo "=== Cleaning up ===" + kill "$MOCK_PID" 2>/dev/null || true + kill "$PROXY_PID" 2>/dev/null || true + wait "$MOCK_PID" 2>/dev/null || true + wait "$PROXY_PID" 2>/dev/null || true + echo "Done." +} +trap cleanup EXIT + +# ── 1. Start mock OAuth2 MCP server ────────────────────────────────────────── +echo "=== Starting mock OAuth2 MCP server on :$MOCK_PORT ===" +python "$MOCK_SERVER" & +MOCK_PID=$! +sleep 2 + +# Quick smoke test on the token endpoint +TOKEN_RESP=$(curl -sf http://localhost:$MOCK_PORT/oauth/token \ + -d "grant_type=client_credentials&client_id=test-client&client_secret=test-secret") +echo "Token endpoint OK: $TOKEN_RESP" + + +# ── 3. List tools ──────────────────────────────────────────────────────────── +echo "" +echo "=== Request 1: List MCP tools ===" +curl -s http://localhost:$PROXY_PORT/mcp-rest/tools/list \ + -H "Authorization: Bearer sk-1234" | python3 -m json.tool + +# ── 4. Call the echo tool ──────────────────────────────────────────────────── +echo "" +echo "=== Request 2: Call echo tool ===" +# Get the server_id from health endpoint +SERVER_ID=$(curl -s http://localhost:$PROXY_PORT/v1/mcp/server/health \ + -H "Authorization: Bearer sk-1234" | python3 -c "import json,sys; print(json.load(sys.stdin)[0]['server_id'])") + +curl -s http://localhost:$PROXY_PORT/mcp-rest/tools/call \ + -H "Content-Type: application/json" \ + -H "Authorization: Bearer sk-1234" \ + -d "{\"name\": \"echo\", \"arguments\": {\"message\": \"Hello from OAuth2 client_credentials\"}, \"server_id\": \"$SERVER_ID\"}" | python3 -m json.tool + +# ── 5. Call again (uses cached token) ──────────────────────────────────────── +echo "" +echo "=== Request 3: Call echo again (cached token) ===" +curl -s http://localhost:$PROXY_PORT/mcp-rest/tools/call \ + -H "Content-Type: application/json" \ + -H "Authorization: Bearer sk-1234" \ + -d "{\"name\": \"echo\", \"arguments\": {\"message\": \"Second call - token should be cached\"}, \"server_id\": \"$SERVER_ID\"}" | python3 -m json.tool + +# ── 6. Show OAuth2-specific proxy logs ─────────────────────────────────────── +echo "" +echo "=== Proxy OAuth2 logs ===" +grep -E "(Fetching OAuth2|Fetched OAuth2)" /tmp/litellm_oauth2_test.log || echo "(no OAuth2 log lines found)" + +echo "" +echo "=== All requests succeeded ===" diff --git a/tests/mcp_tests/test_oauth2_mcp_config.yaml b/tests/mcp_tests/test_oauth2_mcp_config.yaml new file mode 100644 index 00000000000..c2704c5fe71 --- /dev/null +++ b/tests/mcp_tests/test_oauth2_mcp_config.yaml @@ -0,0 +1,14 @@ +model_list: + - model_name: fake-model + litellm_params: + model: openai/fake + api_key: fake-key + +mcp_servers: + test_oauth2_server: + url: "http://localhost:8765/mcp" + transport: "http" + auth_type: "oauth2" + client_id: "test-client" + client_secret: "test-secret" + token_url: "http://localhost:8765/oauth/token" diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_oauth2_token_cache.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_oauth2_token_cache.py new file mode 100644 index 00000000000..2e0f78df605 --- /dev/null +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_oauth2_token_cache.py @@ -0,0 +1,122 @@ +""" +Core tests for MCP OAuth2 machine-to-machine (client_credentials) token management. + +Covers the critical path: resolve_mcp_auth(), token caching, auth priority, +fallback to static token, and the skip-condition property. +""" + +import asyncio +import time +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest + +from litellm.proxy._experimental.mcp_server.oauth2_token_cache import ( + MCPOAuth2TokenCache, + resolve_mcp_auth, +) +from litellm.proxy._types import MCPTransport +from litellm.types.mcp import MCPAuth +from litellm.types.mcp_server.mcp_server_manager import MCPServer + + +def _server(**overrides) -> MCPServer: + defaults = dict( + server_id="srv-1", + name="test", + url="https://mcp.example.com/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + client_id="cid", + client_secret="csec", + token_url="https://auth.example.com/token", + ) + defaults.update(overrides) + return MCPServer(**defaults) + + +def _token_response(token="tok-abc", expires_in=3600): + resp = MagicMock() + resp.json.return_value = { + "access_token": token, + "token_type": "bearer", + "expires_in": expires_in, + } + resp.raise_for_status = MagicMock() + return resp + + +@pytest.mark.asyncio +async def test_resolve_mcp_auth_fetches_oauth2_token(): + """resolve_mcp_auth fetches a token via client_credentials when the server has OAuth2 config.""" + server = _server() + mock_client = AsyncMock() + mock_client.post.return_value = _token_response("m2m-token-1") + + with patch( + "litellm.proxy._experimental.mcp_server.oauth2_token_cache.get_async_httpx_client", + return_value=mock_client, + ): + result = await resolve_mcp_auth(server) + + assert result == "m2m-token-1" + mock_client.post.assert_called_once() + post_data = mock_client.post.call_args[1]["data"] + assert post_data["grant_type"] == "client_credentials" + assert post_data["client_id"] == "cid" + assert post_data["client_secret"] == "csec" + + +@pytest.mark.asyncio +async def test_token_cached_across_calls(): + """Second resolve_mcp_auth call reuses the cached token — only 1 HTTP POST.""" + cache = MCPOAuth2TokenCache() + server = _server() + mock_client = AsyncMock() + mock_client.post.return_value = _token_response("cached-tok") + + with patch( + "litellm.proxy._experimental.mcp_server.oauth2_token_cache.get_async_httpx_client", + return_value=mock_client, + ), patch( + "litellm.proxy._experimental.mcp_server.oauth2_token_cache.mcp_oauth2_token_cache", + cache, + ): + t1 = await resolve_mcp_auth(server) + t2 = await resolve_mcp_auth(server) + + assert t1 == t2 == "cached-tok" + assert mock_client.post.call_count == 1 + + +@pytest.mark.asyncio +async def test_per_request_header_beats_oauth2(): + """An explicit mcp_auth_header takes priority over the OAuth2 token.""" + server = _server() + result = await resolve_mcp_auth(server, mcp_auth_header="Bearer user-tok") + assert result == "Bearer user-tok" + + +@pytest.mark.asyncio +async def test_falls_back_to_static_token(): + """When no client_credentials config, resolve_mcp_auth returns the static authentication_token.""" + server = _server( + client_id=None, + client_secret=None, + token_url=None, + authentication_token="static-tok-xyz", + ) + result = await resolve_mcp_auth(server) + assert result == "static-tok-xyz" + + +def test_needs_user_oauth_token_property(): + """needs_user_oauth_token is True only for OAuth2 servers WITHOUT client_credentials.""" + # OAuth2 with credentials → M2M, no user token needed + assert _server().needs_user_oauth_token is False + + # OAuth2 without credentials → needs per-user token + assert _server(client_id=None, client_secret=None, token_url=None).needs_user_oauth_token is True + + # Non-OAuth2 → never needs user OAuth token + assert _server(auth_type=MCPAuth.bearer_token).needs_user_oauth_token is False