|
6 | 6 | from typing import Any, cast |
7 | 7 |
|
8 | 8 | import pytest |
| 9 | +from pydantic import AnyHttpUrl |
9 | 10 | from starlette.authentication import AuthCredentials |
10 | 11 | from starlette.datastructures import Headers |
11 | 12 | from starlette.requests import Request |
@@ -265,6 +266,56 @@ async def test_mixed_case_authorization_header( |
265 | 266 | assert user.access_token == valid_access_token |
266 | 267 |
|
267 | 268 |
|
| 269 | +class SingleTokenVerifier: |
| 270 | + """A `TokenVerifier` that knows exactly one token.""" |
| 271 | + |
| 272 | + def __init__(self, access_token: AccessToken) -> None: |
| 273 | + self.access_token = access_token |
| 274 | + |
| 275 | + async def verify_token(self, token: str) -> AccessToken | None: |
| 276 | + return self.access_token if token == self.access_token.token else None |
| 277 | + |
| 278 | + |
| 279 | +RS = "https://api.example.com/mcp" |
| 280 | + |
| 281 | + |
| 282 | +@pytest.mark.anyio |
| 283 | +@pytest.mark.parametrize( |
| 284 | + ("resource_server_url", "token_resource", "accepted"), |
| 285 | + [ |
| 286 | + (None, "https://other.example.com/mcp", True), # nothing configured to compare against |
| 287 | + (None, None, True), |
| 288 | + (RS, None, False), # the verifier did not report what the token was issued for |
| 289 | + (RS, RS, True), |
| 290 | + (RS, RS + "/", True), |
| 291 | + (RS, "https://API.EXAMPLE.COM:443/mcp", True), # same URL, different spelling |
| 292 | + (RS, "https://api.example.com", False), |
| 293 | + (RS, RS + "/child", False), |
| 294 | + (RS, "https://api.example.com/other", False), |
| 295 | + (RS, "https://other.example.com/mcp", False), |
| 296 | + (RS, "api.example.com", False), # not a URL |
| 297 | + ], |
| 298 | +) |
| 299 | +async def test_backend_accepts_only_tokens_issued_for_its_resource( |
| 300 | + resource_server_url: str | None, token_resource: str | None, accepted: bool |
| 301 | +): |
| 302 | + """With `resource_server_url` set, only a token whose `resource` (RFC 8707) is that URL is |
| 303 | + accepted and anything else is treated like an unrecognized token (spec-mandated audience |
| 304 | + check); without it the verifier's answer stands (SDK-defined, the default wiring).""" |
| 305 | + token = AccessToken(token="t", client_id="c", scopes=["read"], resource=token_resource) |
| 306 | + backend = BearerAuthBackend( |
| 307 | + SingleTokenVerifier(token), |
| 308 | + resource_server_url=AnyHttpUrl(resource_server_url) if resource_server_url else None, |
| 309 | + ) |
| 310 | + |
| 311 | + result = await backend.authenticate(Request({"type": "http", "headers": [(b"authorization", b"Bearer t")]})) |
| 312 | + |
| 313 | + if accepted: |
| 314 | + assert result is not None and result[1].access_token == token |
| 315 | + else: |
| 316 | + assert result is None |
| 317 | + |
| 318 | + |
268 | 319 | @pytest.mark.anyio |
269 | 320 | class TestRequireAuthMiddleware: |
270 | 321 | """Tests for the RequireAuthMiddleware class.""" |
|
0 commit comments