From 246a0adefd03d033f78ce2964fe102728019ed94 Mon Sep 17 00:00:00 2001 From: Jeremiah Lowin <153965+jlowin@users.noreply.github.com> Date: Mon, 1 Dec 2025 10:51:02 -0500 Subject: [PATCH] Fix get_access_token() returning stale token after OAuth refresh (#2505) * Fix get_access_token() returning stale token after OAuth refresh Fixes #1863 * Update dependencies.py --- src/fastmcp/server/dependencies.py | 30 +++- tests/server/http/test_stale_access_token.py | 159 +++++++++++++++++++ 2 files changed, 185 insertions(+), 4 deletions(-) create mode 100644 tests/server/http/test_stale_access_token.py diff --git a/src/fastmcp/server/dependencies.py b/src/fastmcp/server/dependencies.py index 7b3386b42..68f22e606 100644 --- a/src/fastmcp/server/dependencies.py +++ b/src/fastmcp/server/dependencies.py @@ -6,6 +6,7 @@ from typing import TYPE_CHECKING from mcp.server.auth.middleware.auth_context import ( get_access_token as _sdk_get_access_token, ) +from mcp.server.auth.middleware.bearer_auth import AuthenticatedUser from mcp.server.auth.provider import ( AccessToken as _SDKAccessToken, ) @@ -111,17 +112,38 @@ def get_access_token() -> AccessToken | None: """ Get the FastMCP access token from the current context. + This function first tries to get the token from the current HTTP request's scope, + which is more reliable for long-lived connections where the SDK's auth_context_var + may become stale after token refresh. Falls back to the SDK's context var if no + request is available. + Returns: The access token if an authenticated user is available, None otherwise. """ - # - access_token: _SDKAccessToken | None = _sdk_get_access_token() + access_token: _SDKAccessToken | None = None + + # First, try to get from current HTTP request's scope (issue #1863) + # This is more reliable than auth_context_var for Streamable HTTP sessions + # where tokens may be refreshed between MCP messages + try: + request = get_http_request() + user = request.scope.get("user") + if isinstance(user, AuthenticatedUser): + access_token = user.access_token + except RuntimeError: + # No HTTP request available, fall back to context var + pass + + # Fall back to SDK's context var if we didn't get a token from the request + if access_token is None: + access_token = _sdk_get_access_token() if access_token is None or isinstance(access_token, AccessToken): return access_token - # If the object is not a FastMCP AccessToken, convert it to one if the fields are compatible - # This is a workaround for the case where the SDK returns a different type + # If the object is not a FastMCP AccessToken, convert it to one if the + # fields are compatible (e.g. `claims` is not present in the SDK's AccessToken). + # This is a workaround for the case where the SDK or auth provider returns a different type # If it fails, it will raise a TypeError try: access_token_as_dict = access_token.model_dump() diff --git a/tests/server/http/test_stale_access_token.py b/tests/server/http/test_stale_access_token.py new file mode 100644 index 000000000..34f271e79 --- /dev/null +++ b/tests/server/http/test_stale_access_token.py @@ -0,0 +1,159 @@ +""" +Test for issue #1863: get_access_token() returns stale token after OAuth refresh. + +This test demonstrates the bug where auth_context_var holds a stale token, +but the current HTTP request (via request_ctx) has a fresh token. + +The test should FAIL with the current implementation and PASS after the fix. +""" + +from unittest.mock import MagicMock + +from mcp.server.auth.middleware.auth_context import auth_context_var +from mcp.server.auth.middleware.bearer_auth import AuthenticatedUser +from mcp.server.lowlevel.server import request_ctx +from mcp.shared.context import RequestContext +from starlette.requests import Request + +from fastmcp.server.auth import AccessToken +from fastmcp.server.dependencies import get_access_token + + +class TestStaleAccessToken: + """Test that get_access_token returns fresh token from request scope.""" + + def test_get_access_token_prefers_request_scope_over_stale_context_var(self): + """ + Regression test for issue #1863. + + Scenario: + - auth_context_var has a STALE token (set at HTTP middleware level) + - request_ctx.request.scope["user"] has a FRESH token (per MCP message) + - get_access_token() should return the FRESH token + + This simulates the case where: + 1. A Streamable HTTP session was established with token A + 2. auth_context_var was set to token A during session setup + 3. Token expired, client refreshed, got token B + 4. New MCP message arrives with token B in the request + 5. get_access_token() should return token B, not stale token A + """ + # Create STALE token (in auth_context_var) + # Using FastMCP's AccessToken to avoid conversion issues + stale_token = AccessToken( + token="stale-token-from-initial-auth", + client_id="test-client", + scopes=["read"], + ) + stale_user = AuthenticatedUser(stale_token) + + # Create FRESH token (in request.scope["user"]) + fresh_token = AccessToken( + token="fresh-token-after-refresh", + client_id="test-client", + scopes=["read"], + ) + fresh_user = AuthenticatedUser(fresh_token) + + # Create a mock request with fresh token in scope + scope = { + "type": "http", + "user": fresh_user, + "auth": MagicMock(), + } + mock_request = Request(scope) + + # Create a mock RequestContext with the request + mock_request_context = MagicMock(spec=RequestContext) + mock_request_context.request = mock_request + + # Set up the context vars: + # - auth_context_var has STALE token + # - request_ctx has request with FRESH token + auth_token = auth_context_var.set(stale_user) + request_token = request_ctx.set(mock_request_context) + + try: + # Call get_access_token - should return FRESH token + result = get_access_token() + + # Assert we get the FRESH token, not the stale one + assert result is not None, "Expected an access token but got None" + assert result.token == "fresh-token-after-refresh", ( + f"Expected fresh token 'fresh-token-after-refresh' but got '{result.token}'. " + "get_access_token() is returning the stale token from auth_context_var " + "instead of the fresh token from request.scope['user']." + ) + finally: + # Clean up context vars + auth_context_var.reset(auth_token) + request_ctx.reset(request_token) + + def test_get_access_token_falls_back_to_context_var_when_no_request(self): + """ + Verify that get_access_token falls back to auth_context_var + when there's no HTTP request available. + """ + # Create token in auth_context_var using FastMCP's AccessToken + token = AccessToken( + token="context-var-token", + client_id="test-client", + scopes=["read"], + ) + user = AuthenticatedUser(token) + + # Set up auth_context_var but NOT request_ctx + auth_token = auth_context_var.set(user) + + try: + result = get_access_token() + + assert result is not None + assert result.token == "context-var-token" + finally: + auth_context_var.reset(auth_token) + + def test_get_access_token_returns_none_when_no_auth(self): + """ + Verify that get_access_token returns None when there's no + authenticated user anywhere. + """ + result = get_access_token() + assert result is None + + def test_get_access_token_falls_back_when_scope_user_is_not_authenticated(self): + """ + Verify that get_access_token falls back to auth_context_var when + scope["user"] exists but is not an AuthenticatedUser (e.g., UnauthenticatedUser). + """ + from starlette.authentication import UnauthenticatedUser + + # Create token in auth_context_var + token = AccessToken( + token="context-var-token", + client_id="test-client", + scopes=["read"], + ) + user = AuthenticatedUser(token) + + # Create request with UnauthenticatedUser in scope + scope = { + "type": "http", + "user": UnauthenticatedUser(), + } + mock_request = Request(scope) + mock_request_context = MagicMock(spec=RequestContext) + mock_request_context.request = mock_request + + auth_token = auth_context_var.set(user) + request_token = request_ctx.set(mock_request_context) + + try: + result = get_access_token() + + # Should fall back to auth_context_var since scope user is unauthenticated + assert result is not None + assert result.token == "context-var-token" + finally: + auth_context_var.reset(auth_token) + request_ctx.reset(request_token)