mirror of
https://github.com/PrefectHQ/fastmcp.git
synced 2026-08-20 12:34:17 +02:00
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
This commit is contained in:
parent
01ecc91807
commit
246a0adefd
2 changed files with 185 additions and 4 deletions
|
|
@ -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()
|
||||
|
|
|
|||
159
tests/server/http/test_stale_access_token.py
Normal file
159
tests/server/http/test_stale_access_token.py
Normal file
|
|
@ -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)
|
||||
Loading…
Add table
Add a link
Reference in a new issue