Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -22,13 +22,16 @@
import warnings
from abc import ABCMeta, abstractmethod
from collections import defaultdict
from contextlib import suppress
from enum import Enum
from functools import cache, cached_property
from typing import TYPE_CHECKING, Any, Generic, Literal, TypeVar

from jwt import InvalidTokenError
from sqlalchemy import select
from sqlalchemy.exc import SQLAlchemyError

from airflow import settings
from airflow.api_fastapi.auth.managers.models.base_user import BaseUser
from airflow.api_fastapi.auth.managers.models.resource_details import (
ConnectionDetails,
Expand Down Expand Up @@ -161,7 +164,7 @@ async def get_user_from_token(self, token: str) -> BaseUser:
log.error("JWT token is not valid: %s", e)
raise e

if (jti := payload.get("jti")) and RevokedToken.is_revoked(jti):
if (jti := payload.get("jti")) and self._is_token_revoked(jti):
raise InvalidTokenError("Token has been revoked")

try:
Expand All @@ -170,6 +173,35 @@ async def get_user_from_token(self, token: str) -> BaseUser:
log.error("Couldn't deserialize user from token, JWT token is not valid: %s", e)
raise InvalidTokenError(str(e))

@staticmethod
def _is_token_revoked(jti: str) -> bool:
"""
Return whether the token ``jti`` has been revoked, tolerating a stale DB connection.

This is the first database access in the auth path and runs on the shared scoped
session. A prior request (for example FAB's ``deserialize_user``) can leave that session
bound to a connection the database later drops on idle timeout (MySQL error 4031,
"disconnected ... because of inactivity"). The connection is never re-checked-out, so
``pool_pre_ping`` / ``pool_recycle`` cannot detect it and the first request after an idle
period fails with a 500. On any ``SQLAlchemyError`` we discard the poisoned scoped session
and retry once on a fresh connection, mirroring the recovery
``FabAuthManager.deserialize_user`` gained in #62919. See #71395.
"""
try:
return RevokedToken.is_revoked(jti)
except SQLAlchemyError:
log.warning(
"Revoked-token check failed on a stale DB session; discarding the scoped "
"session and retrying once on a fresh connection.",
exc_info=True,
)
# settings.Session is Optional and only set once the DB is configured; guard so the
# discard is a no-op (rather than a crash) if it is missing.
if (session_registry := settings.Session) is not None:
with suppress(Exception):
session_registry.remove()
return RevokedToken.is_revoked(jti)

def get_fastapi_middlewares(self) -> list[tuple[_MiddlewareFactory[Any], dict[str, Any]]]:
"""
Return middlewares the auth manager wants registered on the main FastAPI app.
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -22,6 +22,7 @@

import pytest
from jwt import InvalidTokenError
from sqlalchemy.exc import OperationalError

from airflow.api_fastapi.auth.managers.base_auth_manager import BaseAuthManager, T
from airflow.api_fastapi.auth.managers.models.base_user import BaseUser
Expand Down Expand Up @@ -333,6 +334,73 @@ async def test_get_user_from_token_revoked(

mock_is_revoked.assert_called_once_with("some-jti")

@staticmethod
def _stale_connection_error() -> OperationalError:
"""Mimic MySQL error 4031 raised when a pooled connection was dropped on idle timeout."""
return OperationalError(
statement="SELECT EXISTS (SELECT 1 FROM revoked_token WHERE jti = %s)",
params=("some-jti",),
orig=Exception("(4031, 'The client was disconnected by the server because of inactivity.')"),
)

@patch("airflow.api_fastapi.auth.managers.base_auth_manager.settings.Session")
@patch("airflow.models.revoked_token.RevokedToken.is_revoked")
@patch(
"airflow.api_fastapi.auth.managers.base_auth_manager.BaseAuthManager._get_token_validator",
autospec=True,
)
@patch.object(EmptyAuthManager, "deserialize_user")
@pytest.mark.asyncio
async def test_get_user_from_token_recovers_from_stale_session_on_revoked_check(
self, mock_deserialize_user, mock__get_token_validator, mock_is_revoked, mock_session, auth_manager
):
"""
The revoked-token check must survive a stale pooled DB connection: on SQLAlchemyError it
discards the scoped session and retries once, so the first request after an idle period
succeeds instead of returning 500. Regression test for issue #71395.
"""
token = "token"
payload = {"jti": "some-jti"}
user = BaseAuthManagerUserTest(name="test")
signer = AsyncMock(spec=JWTValidator)
signer.avalidated_claims.return_value = payload
mock__get_token_validator.return_value = signer
mock_deserialize_user.return_value = user
# First check hits the poisoned connection, the retry runs on a fresh one.
mock_is_revoked.side_effect = [self._stale_connection_error(), False]

result = await auth_manager.get_user_from_token(token)

assert result == user
assert mock_is_revoked.call_count == 2
# The poisoned scoped session was discarded before the retry.
mock_session.remove.assert_called_once_with()
mock_deserialize_user.assert_called_once_with(payload)

@patch("airflow.api_fastapi.auth.managers.base_auth_manager.settings.Session")
@patch("airflow.models.revoked_token.RevokedToken.is_revoked")
@patch(
"airflow.api_fastapi.auth.managers.base_auth_manager.BaseAuthManager._get_token_validator",
autospec=True,
)
@pytest.mark.asyncio
async def test_get_user_from_token_revoked_after_stale_session_retry(
self, mock__get_token_validator, mock_is_revoked, mock_session, auth_manager
):
"""A revocation confirmed by the retry is still honored after recovering the session."""
token = "token"
payload = {"jti": "some-jti"}
signer = AsyncMock(spec=JWTValidator)
signer.avalidated_claims.return_value = payload
mock__get_token_validator.return_value = signer
mock_is_revoked.side_effect = [self._stale_connection_error(), True]

with pytest.raises(InvalidTokenError, match="Token has been revoked"):
await auth_manager.get_user_from_token(token)

assert mock_is_revoked.call_count == 2
mock_session.remove.assert_called_once_with()

@patch(
"airflow.api_fastapi.auth.managers.base_auth_manager.BaseAuthManager._get_token_validator",
autospec=True,
Expand Down