mirror of
https://github.com/PrefectHQ/fastmcp.git
synced 2026-08-19 03:54:18 +02:00
Implement TokenVerifier protocol for mcp-python-sdk compatibility
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 <noreply@anthropic.com>
This commit is contained in:
parent
5ca7a118a6
commit
b98b080c98
6 changed files with 575 additions and 1 deletions
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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
|
||||
):
|
||||
|
|
|
|||
|
|
@ -87,7 +87,7 @@ def setup_auth_middleware_and_routes(
|
|||
middleware = [
|
||||
Middleware(
|
||||
AuthenticationMiddleware,
|
||||
backend=BearerAuthBackend(provider=auth),
|
||||
backend=BearerAuthBackend(auth),
|
||||
),
|
||||
Middleware(AuthContextMiddleware),
|
||||
]
|
||||
|
|
|
|||
179
tests/auth/providers/test_token_verifier.py
Normal file
179
tests/auth/providers/test_token_verifier.py
Normal file
|
|
@ -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)
|
||||
187
tests/server/http/test_auth_setup.py
Normal file
187
tests/server/http/test_auth_setup.py
Normal file
|
|
@ -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
|
||||
178
tests/server/http/test_bearer_auth_backend.py
Normal file
178
tests/server/http/test_bearer_auth_backend.py
Normal file
|
|
@ -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
|
||||
Loading…
Add table
Add a link
Reference in a new issue