From b98b080c985e655d8b97fbbcab50bbe210f2f3a5 Mon Sep 17 00:00:00 2001 From: Jeremiah Lowin <153965+jlowin@users.noreply.github.com> Date: Thu, 26 Jun 2025 15:08:31 -0400 Subject: [PATCH] Implement TokenVerifier protocol for mcp-python-sdk compatibility MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Fixes breaking changes from mcp-python-sdk PR #982 which updated BearerAuthBackend to use TokenVerifier protocol instead of OAuth providers. Changes: - Add verify_token() method to BearerAuthProvider implementing TokenVerifier protocol - Add verify_token() method to InMemoryOAuthProvider implementing TokenVerifier protocol - Update BearerAuthBackend usage in setup_auth_middleware_and_routes() to pass TokenVerifier - Add comprehensive unit tests for TokenVerifier implementations - Add integration tests for BearerAuthBackend with TokenVerifier - Add tests for HTTP auth setup functions All existing functionality preserved with full backwards compatibility. 🤖 Generated with [Claude Code](https://claude.ai/code) Co-Authored-By: Claude --- src/fastmcp/server/auth/providers/bearer.py | 15 ++ .../server/auth/providers/in_memory.py | 15 ++ src/fastmcp/server/http.py | 2 +- tests/auth/providers/test_token_verifier.py | 179 +++++++++++++++++ tests/server/http/test_auth_setup.py | 187 ++++++++++++++++++ tests/server/http/test_bearer_auth_backend.py | 178 +++++++++++++++++ 6 files changed, 575 insertions(+), 1 deletion(-) create mode 100644 tests/auth/providers/test_token_verifier.py create mode 100644 tests/server/http/test_auth_setup.py create mode 100644 tests/server/http/test_bearer_auth_backend.py diff --git a/src/fastmcp/server/auth/providers/bearer.py b/src/fastmcp/server/auth/providers/bearer.py index fd6402895..edb7abb3f 100644 --- a/src/fastmcp/server/auth/providers/bearer.py +++ b/src/fastmcp/server/auth/providers/bearer.py @@ -385,6 +385,21 @@ class BearerAuthProvider(OAuthProvider): return scope_claim return [] + async def verify_token(self, token: str) -> AccessToken | None: + """ + Verify a bearer token and return access info if valid. + + This method implements the TokenVerifier protocol by delegating + to our existing load_access_token method. + + Args: + token: The JWT token string to validate + + Returns: + AccessToken object if valid, None if invalid or expired + """ + return await self.load_access_token(token) + # --- Unused OAuth server methods --- async def get_client(self, client_id: str) -> OAuthClientInformationFull | None: raise NotImplementedError("Client management not supported") diff --git a/src/fastmcp/server/auth/providers/in_memory.py b/src/fastmcp/server/auth/providers/in_memory.py index d948fb620..92408b134 100644 --- a/src/fastmcp/server/auth/providers/in_memory.py +++ b/src/fastmcp/server/auth/providers/in_memory.py @@ -271,6 +271,21 @@ class InMemoryOAuthProvider(OAuthProvider): return token_obj return None + async def verify_token(self, token: str) -> AccessToken | None: + """ + Verify a bearer token and return access info if valid. + + This method implements the TokenVerifier protocol by delegating + to our existing load_access_token method. + + Args: + token: The token string to validate + + Returns: + AccessToken object if valid, None if invalid or expired + """ + return await self.load_access_token(token) + def _revoke_internal( self, access_token_str: str | None = None, refresh_token_str: str | None = None ): diff --git a/src/fastmcp/server/http.py b/src/fastmcp/server/http.py index 7041ba975..8a85edfca 100644 --- a/src/fastmcp/server/http.py +++ b/src/fastmcp/server/http.py @@ -87,7 +87,7 @@ def setup_auth_middleware_and_routes( middleware = [ Middleware( AuthenticationMiddleware, - backend=BearerAuthBackend(provider=auth), + backend=BearerAuthBackend(auth), ), Middleware(AuthContextMiddleware), ] diff --git a/tests/auth/providers/test_token_verifier.py b/tests/auth/providers/test_token_verifier.py new file mode 100644 index 000000000..f8bac52ef --- /dev/null +++ b/tests/auth/providers/test_token_verifier.py @@ -0,0 +1,179 @@ +"""Tests for TokenVerifier protocol implementation in auth providers.""" + +import pytest +from mcp.server.auth.provider import AccessToken + +from fastmcp.server.auth.providers.bearer import BearerAuthProvider, RSAKeyPair +from fastmcp.server.auth.providers.in_memory import InMemoryOAuthProvider + + +class TestBearerAuthProviderTokenVerifier: + """Test that BearerAuthProvider implements TokenVerifier protocol correctly.""" + + @pytest.fixture + def rsa_key_pair(self) -> RSAKeyPair: + """Generate RSA key pair for testing.""" + return RSAKeyPair.generate() + + @pytest.fixture + def bearer_provider(self, rsa_key_pair: RSAKeyPair) -> BearerAuthProvider: + """Create BearerAuthProvider for testing.""" + return BearerAuthProvider( + public_key=rsa_key_pair.public_key, + issuer="https://test.example.com", + audience="https://api.example.com", + ) + + @pytest.fixture + def valid_token(self, rsa_key_pair: RSAKeyPair) -> str: + """Create a valid test token.""" + return rsa_key_pair.create_token( + subject="test-user", + issuer="https://test.example.com", + audience="https://api.example.com", + scopes=["read", "write"], + ) + + @pytest.fixture + def expired_token(self, rsa_key_pair: RSAKeyPair) -> str: + """Create an expired test token.""" + return rsa_key_pair.create_token( + subject="test-user", + issuer="https://test.example.com", + audience="https://api.example.com", + expires_in_seconds=-3600, # Expired 1 hour ago + ) + + async def test_verify_token_with_valid_token( + self, bearer_provider: BearerAuthProvider, valid_token: str + ): + """Test that verify_token returns AccessToken for valid token.""" + result = await bearer_provider.verify_token(valid_token) + + assert result is not None + assert isinstance(result, AccessToken) + assert result.token == valid_token + assert result.client_id == "test-user" + assert "read" in result.scopes + assert "write" in result.scopes + + async def test_verify_token_with_expired_token( + self, bearer_provider: BearerAuthProvider, expired_token: str + ): + """Test that verify_token returns None for expired token.""" + result = await bearer_provider.verify_token(expired_token) + assert result is None + + async def test_verify_token_with_invalid_token( + self, bearer_provider: BearerAuthProvider + ): + """Test that verify_token returns None for invalid token.""" + result = await bearer_provider.verify_token("invalid.token.here") + assert result is None + + async def test_verify_token_with_malformed_token( + self, bearer_provider: BearerAuthProvider + ): + """Test that verify_token returns None for malformed token.""" + result = await bearer_provider.verify_token("not-a-jwt") + assert result is None + + async def test_verify_token_delegation_to_load_access_token( + self, bearer_provider: BearerAuthProvider, valid_token: str + ): + """Test that verify_token delegates to load_access_token.""" + # Both methods should return the same result + verify_result = await bearer_provider.verify_token(valid_token) + load_result = await bearer_provider.load_access_token(valid_token) + + assert verify_result == load_result + if verify_result is not None and load_result is not None: + assert verify_result.token == load_result.token + assert verify_result.client_id == load_result.client_id + assert verify_result.scopes == load_result.scopes + + +class TestInMemoryOAuthProviderTokenVerifier: + """Test that InMemoryOAuthProvider implements TokenVerifier protocol correctly.""" + + @pytest.fixture + def in_memory_provider(self) -> InMemoryOAuthProvider: + """Create InMemoryOAuthProvider for testing.""" + return InMemoryOAuthProvider( + issuer_url="https://test.example.com", + required_scopes=["user"], + ) + + async def test_verify_token_with_nonexistent_token( + self, in_memory_provider: InMemoryOAuthProvider + ): + """Test that verify_token returns None for nonexistent token.""" + result = await in_memory_provider.verify_token("nonexistent-token") + assert result is None + + async def test_verify_token_delegation_to_load_access_token( + self, in_memory_provider: InMemoryOAuthProvider + ): + """Test that verify_token delegates to load_access_token.""" + # Create a test token in the provider's storage + test_token = "test-access-token" + test_access_token = AccessToken( + token=test_token, + client_id="test-client", + scopes=["user"], + expires_at=None, # No expiry + ) + in_memory_provider.access_tokens[test_token] = test_access_token + + # Both methods should return the same result + verify_result = await in_memory_provider.verify_token(test_token) + load_result = await in_memory_provider.load_access_token(test_token) + + assert verify_result == load_result + assert verify_result is not None + assert verify_result.token == test_token + assert verify_result.client_id == "test-client" + assert verify_result.scopes == ["user"] + + async def test_verify_token_with_expired_token( + self, in_memory_provider: InMemoryOAuthProvider + ): + """Test that verify_token returns None for expired token.""" + import time + + # Create an expired token + expired_token = "expired-token" + expired_access_token = AccessToken( + token=expired_token, + client_id="test-client", + scopes=["user"], + expires_at=int(time.time()) - 3600, # Expired 1 hour ago + ) + in_memory_provider.access_tokens[expired_token] = expired_access_token + + result = await in_memory_provider.verify_token(expired_token) + assert result is None + + # Token should be cleaned up from storage + assert expired_token not in in_memory_provider.access_tokens + + +class TestTokenVerifierProtocolCompliance: + """Test that our providers properly implement the TokenVerifier protocol.""" + + async def test_bearer_provider_implements_protocol(self): + """Test that BearerAuthProvider can be used as TokenVerifier.""" + key_pair = RSAKeyPair.generate() + provider = BearerAuthProvider(public_key=key_pair.public_key) + + # Should have the required method for TokenVerifier protocol + assert hasattr(provider, "verify_token") + assert callable(provider.verify_token) + + async def test_in_memory_provider_implements_protocol(self): + """Test that InMemoryOAuthProvider can be used as TokenVerifier.""" + provider = InMemoryOAuthProvider() + + # Should have the required method for TokenVerifier protocol + assert hasattr(provider, "verify_token") + assert callable(provider.verify_token) diff --git a/tests/server/http/test_auth_setup.py b/tests/server/http/test_auth_setup.py new file mode 100644 index 000000000..212495a40 --- /dev/null +++ b/tests/server/http/test_auth_setup.py @@ -0,0 +1,187 @@ +"""Tests for authentication setup in HTTP apps.""" + +import pytest +from mcp.server.auth.middleware.bearer_auth import BearerAuthBackend +from mcp.server.auth.provider import AccessToken +from starlette.middleware import Middleware +from starlette.middleware.authentication import AuthenticationMiddleware + +from fastmcp.server.auth.providers.bearer import BearerAuthProvider, RSAKeyPair +from fastmcp.server.auth.providers.in_memory import InMemoryOAuthProvider +from fastmcp.server.http import setup_auth_middleware_and_routes + + +class TestSetupAuthMiddlewareAndRoutes: + """Test setup_auth_middleware_and_routes with TokenVerifier providers.""" + + @pytest.fixture + def bearer_provider(self) -> BearerAuthProvider: + """Create BearerAuthProvider for testing.""" + key_pair = RSAKeyPair.generate() + return BearerAuthProvider( + public_key=key_pair.public_key, + issuer="https://test.example.com", + audience="https://api.example.com", + required_scopes=["read", "write"], + ) + + @pytest.fixture + def in_memory_provider(self) -> InMemoryOAuthProvider: + """Create InMemoryOAuthProvider for testing.""" + return InMemoryOAuthProvider( + issuer_url="https://test.example.com", + required_scopes=["user"], + ) + + def test_setup_with_bearer_provider(self, bearer_provider: BearerAuthProvider): + """Test that setup works with BearerAuthProvider as TokenVerifier.""" + middleware, auth_routes, required_scopes = setup_auth_middleware_and_routes( + bearer_provider + ) + + # Should return middleware list + assert isinstance(middleware, list) + assert len(middleware) == 2 # AuthenticationMiddleware + AuthContextMiddleware + + # First middleware should be AuthenticationMiddleware with BearerAuthBackend + auth_middleware = middleware[0] + assert isinstance(auth_middleware, Middleware) + assert auth_middleware.cls == AuthenticationMiddleware + assert "backend" in auth_middleware.kwargs + + backend = auth_middleware.kwargs["backend"] + assert isinstance(backend, BearerAuthBackend) + assert backend.token_verifier is bearer_provider # type: ignore[attr-defined] + + # Should return auth routes + assert isinstance(auth_routes, list) + assert len(auth_routes) > 0 # Should have OAuth routes + + # Should return required scopes + assert required_scopes == ["read", "write"] + + def test_setup_with_in_memory_provider( + self, in_memory_provider: InMemoryOAuthProvider + ): + """Test that setup works with InMemoryOAuthProvider as TokenVerifier.""" + middleware, auth_routes, required_scopes = setup_auth_middleware_and_routes( + in_memory_provider + ) + + # Should return middleware list + assert isinstance(middleware, list) + assert len(middleware) == 2 + + # Backend should use the provider as token verifier + auth_middleware = middleware[0] + backend = auth_middleware.kwargs["backend"] + assert isinstance(backend, BearerAuthBackend) + assert backend.token_verifier is in_memory_provider # type: ignore[attr-defined] + + # Should return required scopes + assert required_scopes == ["user"] + + def test_setup_preserves_provider_functionality( + self, bearer_provider: BearerAuthProvider + ): + """Test that setup doesn't break the provider's functionality.""" + # Setup should not modify the provider + original_issuer = bearer_provider.issuer + original_scopes = bearer_provider.required_scopes + + middleware, auth_routes, required_scopes = setup_auth_middleware_and_routes( + bearer_provider + ) + + # Provider should be unchanged + assert bearer_provider.issuer == original_issuer + assert bearer_provider.required_scopes == original_scopes + + # Provider should still work as TokenVerifier + assert hasattr(bearer_provider, "verify_token") + assert callable(bearer_provider.verify_token) + + +class MockOAuthProvider: + """Mock OAuth provider that implements TokenVerifier.""" + + def __init__(self, required_scopes=None, issuer_url="http://localhost:8000"): + from pydantic import AnyHttpUrl + + from fastmcp.server.auth.auth import ( + ClientRegistrationOptions, + RevocationOptions, + ) + + self.required_scopes = required_scopes or [] + self.issuer_url = AnyHttpUrl(issuer_url) + self.service_documentation_url = None + self.client_registration_options = ClientRegistrationOptions(enabled=False) + self.revocation_options = RevocationOptions(enabled=False) + + async def verify_token(self, token: str) -> AccessToken | None: + """Mock verify_token implementation.""" + if token == "valid-token": + return AccessToken( + token=token, + client_id="mock-client", + scopes=self.required_scopes, + expires_at=None, + ) + return None + + +class TestSetupWithMockProvider: + """Test setup function with mock provider.""" + + def test_setup_with_mock_token_verifier(self): + """Test that setup works with any TokenVerifier implementation.""" + mock_provider = MockOAuthProvider(required_scopes=["mock-scope"]) + + middleware, auth_routes, required_scopes = setup_auth_middleware_and_routes( + mock_provider # type: ignore[arg-type] + ) + + # Should work with any TokenVerifier + assert len(middleware) == 2 + auth_middleware = middleware[0] + backend = auth_middleware.kwargs["backend"] + assert isinstance(backend, BearerAuthBackend) + assert backend.token_verifier is mock_provider # type: ignore[attr-defined] + + assert required_scopes == ["mock-scope"] + + async def test_setup_middleware_can_authenticate(self): + """Test that the setup middleware can actually authenticate requests.""" + mock_provider = MockOAuthProvider() + + middleware, _, _ = setup_auth_middleware_and_routes(mock_provider) # type: ignore[arg-type] + + # Extract the BearerAuthBackend + auth_middleware = middleware[0] + backend = auth_middleware.kwargs["backend"] + + # Test authentication with valid token + from starlette.requests import HTTPConnection + + scope = { + "type": "http", + "headers": [(b"authorization", b"Bearer valid-token")], + } + conn = HTTPConnection(scope) + + result = await backend.authenticate(conn) # type: ignore[attr-defined] + assert result is not None + + credentials, user = result + assert user.username == "mock-client" + + # Test authentication with invalid token + scope = { + "type": "http", + "headers": [(b"authorization", b"Bearer invalid-token")], + } + conn = HTTPConnection(scope) + + result = await backend.authenticate(conn) # type: ignore[attr-defined] + assert result is None diff --git a/tests/server/http/test_bearer_auth_backend.py b/tests/server/http/test_bearer_auth_backend.py new file mode 100644 index 000000000..10d1bc69b --- /dev/null +++ b/tests/server/http/test_bearer_auth_backend.py @@ -0,0 +1,178 @@ +"""Tests for BearerAuthBackend integration with TokenVerifier.""" + +import pytest +from mcp.server.auth.middleware.bearer_auth import BearerAuthBackend +from mcp.server.auth.provider import AccessToken +from starlette.requests import HTTPConnection + +from fastmcp.server.auth.providers.bearer import BearerAuthProvider, RSAKeyPair + + +class TestBearerAuthBackendTokenVerifierIntegration: + """Test BearerAuthBackend works with TokenVerifier protocol.""" + + @pytest.fixture + def rsa_key_pair(self) -> RSAKeyPair: + """Generate RSA key pair for testing.""" + return RSAKeyPair.generate() + + @pytest.fixture + def bearer_provider(self, rsa_key_pair: RSAKeyPair) -> BearerAuthProvider: + """Create BearerAuthProvider for testing.""" + return BearerAuthProvider( + public_key=rsa_key_pair.public_key, + issuer="https://test.example.com", + audience="https://api.example.com", + ) + + @pytest.fixture + def valid_token(self, rsa_key_pair: RSAKeyPair) -> str: + """Create a valid test token.""" + return rsa_key_pair.create_token( + subject="test-user", + issuer="https://test.example.com", + audience="https://api.example.com", + scopes=["read", "write"], + ) + + def test_bearer_auth_backend_constructor_accepts_token_verifier( + self, bearer_provider: BearerAuthProvider + ): + """Test that BearerAuthBackend constructor accepts TokenVerifier.""" + # This should not raise an error + backend = BearerAuthBackend(bearer_provider) + assert backend.token_verifier is bearer_provider # type: ignore[attr-defined] + + async def test_bearer_auth_backend_authenticate_with_valid_token( + self, bearer_provider: BearerAuthProvider, valid_token: str + ): + """Test BearerAuthBackend authentication with valid token.""" + backend = BearerAuthBackend(bearer_provider) + + # Create mock HTTPConnection with Authorization header + scope = { + "type": "http", + "headers": [(b"authorization", f"Bearer {valid_token}".encode())], + } + conn = HTTPConnection(scope) + + result = await backend.authenticate(conn) + + assert result is not None + credentials, user = result + assert credentials.scopes == ["read", "write"] + assert user.username == "test-user" + assert hasattr(user, "access_token") + assert user.access_token.token == valid_token + + async def test_bearer_auth_backend_authenticate_with_invalid_token( + self, bearer_provider: BearerAuthProvider + ): + """Test BearerAuthBackend authentication with invalid token.""" + backend = BearerAuthBackend(bearer_provider) + + # Create mock HTTPConnection with invalid Authorization header + scope = { + "type": "http", + "headers": [(b"authorization", b"Bearer invalid-token")], + } + conn = HTTPConnection(scope) + + result = await backend.authenticate(conn) + assert result is None + + async def test_bearer_auth_backend_authenticate_with_no_header( + self, bearer_provider: BearerAuthProvider + ): + """Test BearerAuthBackend authentication with no Authorization header.""" + backend = BearerAuthBackend(bearer_provider) + + # Create mock HTTPConnection without Authorization header + scope = { + "type": "http", + "headers": [], + } + conn = HTTPConnection(scope) + + result = await backend.authenticate(conn) + assert result is None + + async def test_bearer_auth_backend_authenticate_with_non_bearer_token( + self, bearer_provider: BearerAuthProvider + ): + """Test BearerAuthBackend authentication with non-Bearer token.""" + backend = BearerAuthBackend(bearer_provider) + + # Create mock HTTPConnection with Basic auth header + scope = { + "type": "http", + "headers": [(b"authorization", b"Basic dXNlcjpwYXNz")], + } + conn = HTTPConnection(scope) + + result = await backend.authenticate(conn) + assert result is None + + +class MockTokenVerifier: + """Mock TokenVerifier for testing backend integration.""" + + def __init__(self, return_value: AccessToken | None = None): + self.return_value = return_value + self.verify_token_calls = [] + + async def verify_token(self, token: str) -> AccessToken | None: + """Mock verify_token method.""" + self.verify_token_calls.append(token) + return self.return_value + + +class TestBearerAuthBackendWithMockVerifier: + """Test BearerAuthBackend with mock TokenVerifier.""" + + async def test_backend_calls_verify_token_method(self): + """Test that BearerAuthBackend calls verify_token on the verifier.""" + mock_access_token = AccessToken( + token="test-token", + client_id="test-client", + scopes=["read"], + expires_at=None, + ) + mock_verifier = MockTokenVerifier(return_value=mock_access_token) + backend = BearerAuthBackend(mock_verifier) # type: ignore[arg-type] + + scope = { + "type": "http", + "headers": [(b"authorization", b"Bearer test-token")], + } + conn = HTTPConnection(scope) + + result = await backend.authenticate(conn) + + # Should have called verify_token with the token + assert mock_verifier.verify_token_calls == ["test-token"] + + # Should return authentication result + assert result is not None + credentials, user = result + assert credentials.scopes == ["read"] + assert user.username == "test-client" + + async def test_backend_handles_verify_token_none_result(self): + """Test that BearerAuthBackend handles None result from verify_token.""" + mock_verifier = MockTokenVerifier(return_value=None) + backend = BearerAuthBackend(mock_verifier) # type: ignore[arg-type] + + scope = { + "type": "http", + "headers": [(b"authorization", b"Bearer invalid-token")], + } + conn = HTTPConnection(scope) + + result = await backend.authenticate(conn) + + # Should have called verify_token + assert mock_verifier.verify_token_calls == ["invalid-token"] + + # Should return None for authentication failure + assert result is None