This is an automated email from the ASF dual-hosted git repository.
potiuk pushed a commit to branch main
in repository https://gitbox.apache.org/repos/asf/airflow.git
The following commit(s) were added to refs/heads/main by this push:
new 607df554985 Add KeycloakJWTMiddleware to KeycloakAuthManager (#70800)
607df554985 is described below
commit 607df5549859105be835000c3505db6d0a547363
Author: stephen-bracken <[email protected]>
AuthorDate: Sun Aug 2 22:56:37 2026 +0100
Add KeycloakJWTMiddleware to KeycloakAuthManager (#70800)
Co-authored-by: Stephen Bracken <email-protected>
---
.../providers/keycloak/auth_manager/constants.py | 6 +-
.../keycloak/auth_manager/keycloak_auth_manager.py | 46 ++-
.../providers/keycloak/auth_manager/middleware.py | 246 ++++++++++++++
.../keycloak/auth_manager/routes/login.py | 15 -
.../airflow/providers/keycloak/version_compat.py | 1 +
.../keycloak/auth_manager/routes/test_login.py | 2 +
.../auth_manager/test_keycloak_auth_manager.py | 123 +++++--
.../unit/keycloak/auth_manager/test_middleware.py | 372 +++++++++++++++++++++
8 files changed, 746 insertions(+), 65 deletions(-)
diff --git
a/providers/keycloak/src/airflow/providers/keycloak/auth_manager/constants.py
b/providers/keycloak/src/airflow/providers/keycloak/auth_manager/constants.py
index 85749f26500..7a6cd2e0ece 100644
---
a/providers/keycloak/src/airflow/providers/keycloak/auth_manager/constants.py
+++
b/providers/keycloak/src/airflow/providers/keycloak/auth_manager/constants.py
@@ -27,9 +27,7 @@ CONF_REQUESTS_POOL_SIZE_KEY = "requests_pool_size"
CONF_REQUESTS_RETRIES_KEY = "requests_retries"
# Extra Cookie names
-COOKIE_NAME_ACCESS_TOKEN = "access_token"
+COOKIE_NAME_ACCESS_TOKEN = "_access_token"
COOKIE_NAME_ID_TOKEN = "_id_token"
-COOKIE_NAME_NAME = "name"
COOKIE_NAME_OAUTH_STATE = "_oauth_state"
-COOKIE_NAME_REFRESH_TOKEN = "refresh_token"
-COOKIE_NAME_USER_ID = "user_id"
+COOKIE_NAME_REFRESH_TOKEN = "_refresh_token"
diff --git
a/providers/keycloak/src/airflow/providers/keycloak/auth_manager/keycloak_auth_manager.py
b/providers/keycloak/src/airflow/providers/keycloak/auth_manager/keycloak_auth_manager.py
index a680be488f7..1bb0c9bda1b 100644
---
a/providers/keycloak/src/airflow/providers/keycloak/auth_manager/keycloak_auth_manager.py
+++
b/providers/keycloak/src/airflow/providers/keycloak/auth_manager/keycloak_auth_manager.py
@@ -24,11 +24,11 @@ import warnings
from base64 import urlsafe_b64decode
from collections.abc import Sequence
from concurrent.futures import ThreadPoolExecutor
-from typing import TYPE_CHECKING, Annotated, Any
+from typing import TYPE_CHECKING, Any, cast
from urllib.parse import urljoin
import requests
-from fastapi import Cookie, FastAPI
+from fastapi import FastAPI
from keycloak import KeycloakOpenID
from keycloak.exceptions import KeycloakPostError
from requests.adapters import HTTPAdapter
@@ -101,11 +101,6 @@ TEAM_SCOPED_RESOURCES = frozenset(
)
-def _get_keycloak_jwt(user: Annotated[KeycloakAuthManagerUser | None,
Cookie(default=None)] = None):
- """Populate Keycloak user from cookies."""
- return user
-
-
class KeycloakAuthManager(BaseAuthManager[KeycloakAuthManagerUser]):
"""
Keycloak auth manager.
@@ -143,13 +138,9 @@ class
KeycloakAuthManager(BaseAuthManager[KeycloakAuthManagerUser]):
return self._http_session
def deserialize_user(self, token: dict[str, Any]) ->
KeycloakAuthManagerUser:
- user = _get_keycloak_jwt()
- if user is None:
- raise ValueError("Couldn't deserialise user from Cookies.")
- if user_id := token.pop("user_id"):
- if user.get_id() != user_id:
- raise ValueError("Keycloak user in Cookies does not match
Airflow JWT.")
- return user
+ return KeycloakAuthManagerUser(
+ user_id=token["user_id"], name=token["name"], access_token="",
refresh_token=None
+ )
def serialize_user(self, user: KeycloakAuthManagerUser) -> dict[str, Any]:
return {
@@ -157,6 +148,24 @@ class
KeycloakAuthManager(BaseAuthManager[KeycloakAuthManagerUser]):
"name": user.get_name(),
}
+ async def get_user_from_token(
+ self, token: str, access_token: str | None = None, refresh_token: str
| None = None
+ ):
+ """
+ Get the user from the Airflow and Keycloak Tokens.
+
+ :param token: Airflow JWT
+ :param access_token: Keycloak access JWT
+ :param refresh_token: Keycloak refresh JWT
+ """
+ user = cast("KeycloakAuthManagerUser", await
super().get_user_from_token(token))
+ if access_token:
+ user.access_token = access_token
+ user.refresh_token = refresh_token
+ return user
+ # Skip refreshing JWT if Keycloak JWTs are not included.
+ return None
+
def get_url_login(self, **kwargs) -> str:
base_url = conf.get("api", "base_url", fallback="/")
return urljoin(base_url, f"{AUTH_MANAGER_FASTAPI_APP_PREFIX}/login")
@@ -165,12 +174,12 @@ class
KeycloakAuthManager(BaseAuthManager[KeycloakAuthManagerUser]):
base_url = conf.get("api", "base_url", fallback="/")
return urljoin(base_url, f"{AUTH_MANAGER_FASTAPI_APP_PREFIX}/logout")
- def refresh_user(self, *, user: KeycloakAuthManagerUser) ->
KeycloakAuthManagerUser | None:
+ def refresh_user(self, *, user: KeycloakAuthManagerUser | None) ->
KeycloakAuthManagerUser | None:
# According to RFC6749 section 4.4.3, a refresh token should not be
included when using
# the Service accounts/client_credentials flow.
# We check whether the user has a refresh token; if not, we assume
it's a service account
# and return None.
- if not user.refresh_token:
+ if not user or not user.refresh_token:
return None
if self._token_expired(user.access_token):
@@ -370,6 +379,11 @@ class
KeycloakAuthManager(BaseAuthManager[KeycloakAuthManagerUser]):
return app
+ def get_fastapi_middlewares(self):
+ from airflow.providers.keycloak.auth_manager.middleware import
KeycloakJWTMiddleware
+
+ return [(KeycloakJWTMiddleware, {})]
+
@staticmethod
def get_cli_commands() -> list[CLICommand]:
"""Vends CLI commands to be included in Airflow CLI."""
diff --git
a/providers/keycloak/src/airflow/providers/keycloak/auth_manager/middleware.py
b/providers/keycloak/src/airflow/providers/keycloak/auth_manager/middleware.py
new file mode 100644
index 00000000000..fcbfc56f044
--- /dev/null
+++
b/providers/keycloak/src/airflow/providers/keycloak/auth_manager/middleware.py
@@ -0,0 +1,246 @@
+# Licensed to the Apache Software Foundation (ASF) under one
+# or more contributor license agreements. See the NOTICE file
+# distributed with this work for additional information
+# regarding copyright ownership. The ASF licenses this file
+# to you under the Apache License, Version 2.0 (the
+# "License"); you may not use this file except in compliance
+# with the License. You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing,
+# software distributed under the License is distributed on an
+# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
+# KIND, either express or implied. See the License for the
+# specific language governing permissions and limitations
+# under the License.
+
+from __future__ import annotations
+
+from typing import TYPE_CHECKING, cast
+
+from fastapi import HTTPException, status
+from fastapi.responses import JSONResponse
+from jwt import ExpiredSignatureError, InvalidTokenError
+from starlette.middleware.base import BaseHTTPMiddleware
+
+from airflow.api_fastapi.app import get_auth_manager
+from airflow.api_fastapi.auth.managers.base_auth_manager import
COOKIE_NAME_JWT_TOKEN
+from airflow.api_fastapi.core_api import security as core_api_security
+from airflow.providers.common.compat.sdk import conf
+from airflow.providers.keycloak.auth_manager.constants import (
+ COOKIE_NAME_ACCESS_TOKEN,
+ COOKIE_NAME_REFRESH_TOKEN,
+)
+from airflow.providers.keycloak.version_compat import AIRFLOW_V_3_1_8_PLUS
+
+try:
+ from airflow.api_fastapi.auth.managers.exceptions import
AuthManagerRefreshTokenExpiredException
+except ImportError:
+
+ class AuthManagerRefreshTokenExpiredException(Exception): # type:
ignore[no-redef]
+ """In case it is using a version of Airflow without
``AuthManagerRefreshTokenExpiredException``."""
+
+ pass
+
+
+if AIRFLOW_V_3_1_8_PLUS:
+ from airflow.api_fastapi.app import get_cookie_path
+else:
+
+ def get_cookie_path() -> str:
+ return "/"
+
+
+if TYPE_CHECKING:
+ from fastapi import Request, Response
+
+ from airflow.providers.keycloak.auth_manager.keycloak_auth_manager import
KeycloakAuthManager
+ from airflow.providers.keycloak.auth_manager.user import
KeycloakAuthManagerUser
+
+
+class KeycloakJWTMiddleware(BaseHTTPMiddleware):
+ """
+ Attach the Keycloak JWT tokens to the user.
+
+ Gets the Keycloak JWT tokens from the request cookies
+ and attaches them to the user. If the token is expired,
+ attempt to refresh it using the refresh token.
+ """
+
+ async def dispatch(self, request: Request, call_next):
+ user = None
+ new_token = None
+ new_user = None
+ try:
+ try:
+ new_user, current_user = await self._refresh_user(request)
+ user = new_user or current_user
+ except (
+ AuthManagerRefreshTokenExpiredException,
+ ExpiredSignatureError,
+ InvalidTokenError,
+ HTTPException,
+ ):
+ new_token = ""
+
+ if user is not None:
+ request.state.user = user
+
+ user_injected = getattr(
+ core_api_security,
+ "USER_INJECTED_BY_TRUSTED_MIDDLEWARE",
+ None,
+ )
+ if user_injected is not None:
+ request.state.user_authenticated_via = user_injected
+
+ response = await call_next(request)
+
+ if new_user or new_token is not None:
+ secure = request.base_url.scheme == "https" or
bool(conf.get("api", "ssl_cert", fallback=""))
+ cookie_path = get_cookie_path()
+ if new_token == "":
+ response.set_cookie(
+ COOKIE_NAME_JWT_TOKEN,
+ new_token,
+ path=cookie_path,
+ httponly=True,
+ secure=secure,
+ samesite="lax",
+ max_age=0,
+ )
+ if cookie_path != "/":
+ response.set_cookie(
+ COOKIE_NAME_JWT_TOKEN,
+ "",
+ path="/",
+ httponly=True,
+ secure=secure,
+ samesite="lax",
+ max_age=0,
+ )
+ else:
+ response = await self._set_new_token(new_user, secure,
response, cookie_path)
+
+ except HTTPException as exc:
+ # If any HTTPException is raised during user resolution or
refresh, return it as response
+ return JSONResponse(status_code=exc.status_code,
content={"detail": exc.detail})
+
+ return response
+
+ @classmethod
+ async def _set_new_token(
+ cls,
+ new_user: KeycloakAuthManagerUser | None,
+ secure: bool,
+ response: Response,
+ cookie_path: str | None = None,
+ ) -> Response:
+ """
+ Set Cookies in the response based on a new JWT token and a new user
model.
+
+ :param new_user: User model for the JWT token
+ :param secure: HTTP secure property for cookies
+ :param response: FastAPI response object to set the cookies on
+ :param cookie_path: Path for cookies in the response
+ """
+ if cookie_path is None:
+ cookie_path = get_cookie_path()
+ if new_user:
+ # If we created a new user, serialize it and set it as a cookie
+ new_token = get_auth_manager().generate_jwt(new_user)
+ else:
+ new_token = ""
+ response.set_cookie(
+ COOKIE_NAME_JWT_TOKEN,
+ new_token,
+ path=cookie_path,
+ httponly=True,
+ secure=secure,
+ samesite="lax",
+ max_age=0 if new_token == "" else None,
+ )
+ if new_user:
+ # Update keycloak token cookies
+ response.set_cookie(
+ COOKIE_NAME_ACCESS_TOKEN,
+ new_user.access_token,
+ path=cookie_path,
+ secure=secure,
+ samesite="lax",
+ httponly=True,
+ )
+ if new_user.refresh_token:
+ response.set_cookie(
+ COOKIE_NAME_REFRESH_TOKEN,
+ new_user.refresh_token,
+ path=cookie_path,
+ secure=secure,
+ samesite="lax",
+ httponly=True,
+ )
+ else:
+ # No refresh token
+ response.set_cookie(
+ COOKIE_NAME_REFRESH_TOKEN,
+ "",
+ path=cookie_path,
+ secure=secure,
+ samesite="lax",
+ httponly=True,
+ max_age=0,
+ )
+ else:
+ # User is not populated, clear cookies for login
+ response.set_cookie(
+ COOKIE_NAME_ACCESS_TOKEN,
+ "",
+ path=cookie_path,
+ secure=secure,
+ samesite="lax",
+ httponly=True,
+ max_age=0,
+ )
+ response.set_cookie(
+ COOKIE_NAME_REFRESH_TOKEN,
+ "",
+ path=cookie_path,
+ secure=secure,
+ samesite="lax",
+ httponly=True,
+ max_age=0,
+ )
+ # Clear any stale _token cookie at root path "/".
+ # Older Airflow instances may have set the cookie there;
+ # without this, the root-path cookie keeps being sent on
+ # every request, causing an infinite redirect loop.
+ if cookie_path != "/":
+ response.set_cookie(
+ key=COOKIE_NAME_JWT_TOKEN,
+ path="/",
+ httponly=True,
+ secure=secure,
+ samesite="lax",
+ max_age=0,
+ )
+ return response
+
+ @staticmethod
+ async def _refresh_user(
+ request: Request,
+ ) -> tuple[KeycloakAuthManagerUser | None, KeycloakAuthManagerUser | None]:
+ jwt_token = request.cookies.get(COOKIE_NAME_JWT_TOKEN)
+ access_token = request.cookies.get(COOKIE_NAME_ACCESS_TOKEN)
+ refresh_token = request.cookies.get(COOKIE_NAME_REFRESH_TOKEN)
+ if not jwt_token:
+ raise HTTPException(
+ status_code=status.HTTP_401_UNAUTHORIZED, detail="User is not
logged into Airflow."
+ )
+ if not access_token:
+ raise HTTPException(
+ status_code=status.HTTP_401_UNAUTHORIZED, detail="User is not
logged into Keycloak."
+ )
+ auth_manager = cast("KeycloakAuthManager", get_auth_manager())
+ user = await auth_manager.get_user_from_token(jwt_token, access_token,
refresh_token)
+ return get_auth_manager().refresh_user(user=user), user
diff --git
a/providers/keycloak/src/airflow/providers/keycloak/auth_manager/routes/login.py
b/providers/keycloak/src/airflow/providers/keycloak/auth_manager/routes/login.py
index 95758f6122f..804733757c4 100644
---
a/providers/keycloak/src/airflow/providers/keycloak/auth_manager/routes/login.py
+++
b/providers/keycloak/src/airflow/providers/keycloak/auth_manager/routes/login.py
@@ -50,10 +50,8 @@ from airflow.providers.common.compat.sdk import conf
from airflow.providers.keycloak.auth_manager.constants import (
COOKIE_NAME_ACCESS_TOKEN,
COOKIE_NAME_ID_TOKEN,
- COOKIE_NAME_NAME,
COOKIE_NAME_OAUTH_STATE,
COOKIE_NAME_REFRESH_TOKEN,
- COOKIE_NAME_USER_ID,
)
from airflow.providers.keycloak.auth_manager.keycloak_auth_manager import
KeycloakAuthManager
from airflow.providers.keycloak.auth_manager.user import
KeycloakAuthManagerUser
@@ -141,12 +139,6 @@ def login_callback(request: Request):
COOKIE_NAME_ID_TOKEN, tokens["id_token"], path=cookie_path,
secure=secure, httponly=True
)
- response.set_cookie(COOKIE_NAME_USER_ID, userinfo["sub"],
path=cookie_path, secure=secure, httponly=True)
-
- response.set_cookie(
- COOKIE_NAME_NAME, userinfo["preferred_username"], path=cookie_path,
secure=secure, httponly=True
- )
-
response.set_cookie(
COOKIE_NAME_ACCESS_TOKEN, tokens["access_token"], path=cookie_path,
secure=secure, httponly=True
)
@@ -205,13 +197,6 @@ def logout_callback(request: Request):
secure=secure,
httponly=True,
)
- response.delete_cookie(key=COOKIE_NAME_USER_ID, path=cookie_path,
secure=secure, httponly=True)
- response.delete_cookie(
- key=COOKIE_NAME_NAME,
- path=cookie_path,
- secure=secure,
- httponly=True,
- )
response.delete_cookie(
key=COOKIE_NAME_ACCESS_TOKEN,
path=cookie_path,
diff --git
a/providers/keycloak/src/airflow/providers/keycloak/version_compat.py
b/providers/keycloak/src/airflow/providers/keycloak/version_compat.py
index 917adca937c..e97df79b887 100644
--- a/providers/keycloak/src/airflow/providers/keycloak/version_compat.py
+++ b/providers/keycloak/src/airflow/providers/keycloak/version_compat.py
@@ -33,4 +33,5 @@ def get_base_airflow_version_tuple() -> tuple[int, int, int]:
AIRFLOW_V_3_1_1_PLUS = get_base_airflow_version_tuple() >= (3, 1, 1)
+AIRFLOW_V_3_1_7_PLUS = get_base_airflow_version_tuple() >= (3, 1, 7)
AIRFLOW_V_3_1_8_PLUS = get_base_airflow_version_tuple() >= (3, 1, 8)
diff --git
a/providers/keycloak/tests/unit/keycloak/auth_manager/routes/test_login.py
b/providers/keycloak/tests/unit/keycloak/auth_manager/routes/test_login.py
index 28323f184ef..058bda6590d 100644
--- a/providers/keycloak/tests/unit/keycloak/auth_manager/routes/test_login.py
+++ b/providers/keycloak/tests/unit/keycloak/auth_manager/routes/test_login.py
@@ -79,6 +79,8 @@ class TestLoginRouter:
assert "_token" in response.cookies
assert response.cookies["_token"] == token
assert response.cookies["_id_token"] == "id_token"
+ assert response.cookies["_access_token"] == "access_token"
+ assert response.cookies["_refresh_token"] == "refresh_token"
@patch("airflow.providers.keycloak.auth_manager.routes.login.KeycloakAuthManager.get_keycloak_client")
def test_login_sets_secure_state_cookie_behind_tls_proxy(self,
mock_get_keycloak_client, client):
diff --git
a/providers/keycloak/tests/unit/keycloak/auth_manager/test_keycloak_auth_manager.py
b/providers/keycloak/tests/unit/keycloak/auth_manager/test_keycloak_auth_manager.py
index d3f075d2a46..44ab60304bd 100644
---
a/providers/keycloak/tests/unit/keycloak/auth_manager/test_keycloak_auth_manager.py
+++
b/providers/keycloak/tests/unit/keycloak/auth_manager/test_keycloak_auth_manager.py
@@ -20,7 +20,7 @@ import base64
import json
from concurrent.futures import ThreadPoolExecutor
from contextlib import ExitStack
-from unittest.mock import Mock, patch
+from unittest.mock import AsyncMock, Mock, patch
import pytest
from keycloak import KeycloakPostError
@@ -42,10 +42,16 @@ from
airflow.api_fastapi.auth.managers.models.resource_details import (
from tests_common.test_utils.config import conf_vars
from tests_common.test_utils.version_compat import AIRFLOW_V_3_1_7_PLUS,
AIRFLOW_V_3_2_PLUS
+if AIRFLOW_V_3_1_7_PLUS:
+ from airflow.api_fastapi.auth.managers.exceptions import
AuthManagerRefreshTokenExpiredException
+else:
+ AuthManagerRefreshTokenExpiredException = None # type:
ignore[assignment,misc]
+
if AIRFLOW_V_3_2_PLUS:
from airflow.api_fastapi.auth.managers.models.resource_details import
TeamDetails
else:
TeamDetails = None # type: ignore[assignment,misc]
+from airflow.api_fastapi.auth.managers.base_auth_manager import BaseAuthManager
from airflow.api_fastapi.common.types import MenuItem
from airflow.exceptions import AirflowProviderDeprecationWarning
@@ -66,6 +72,7 @@ from
airflow.providers.keycloak.auth_manager.keycloak_auth_manager import (
RESOURCE_ID_ATTRIBUTE_NAME,
KeycloakAuthManager,
)
+from airflow.providers.keycloak.auth_manager.middleware import
KeycloakJWTMiddleware
from airflow.providers.keycloak.auth_manager.user import
KeycloakAuthManagerUser
@@ -122,33 +129,27 @@ def _clear_filter_cache():
class TestKeycloakAuthManager:
-
@patch("airflow.providers.keycloak.auth_manager.keycloak_auth_manager._get_keycloak_jwt")
- def test_deserialize_user(self, mock_get_keycloak_jwt, auth_manager):
- mock_get_keycloak_jwt.return_value = KeycloakAuthManagerUser(
- user_id="user_id", name="name", access_token="access_token",
refresh_token="refresh_token"
- )
- result = auth_manager.deserialize_user({"user_id": "user_id", "name":
"name"})
+ @pytest.mark.parametrize(
+ "token_data",
+ [
+ {
+ "user_id": "user_id",
+ "name": "name",
+ },
+ {
+ "user_id": "user_id",
+ "name": "name",
+ "access_token": "access_token",
+ "refresh_token": "refresh_token",
+ },
+ ],
+ )
+ def test_deserialize_user(self, auth_manager, token_data):
+ result = auth_manager.deserialize_user(token_data)
assert result.user_id == "user_id"
assert result.name == "name"
- assert result.access_token == "access_token"
- assert result.refresh_token == "refresh_token"
-
-
@patch("airflow.providers.keycloak.auth_manager.keycloak_auth_manager._get_keycloak_jwt")
- def test_deserialize_user_missing(self, mock_get_keycloak_jwt,
auth_manager):
- mock_get_keycloak_jwt.return_value = None
- with pytest.raises(ValueError, match="Couldn't deserialise user from
Cookies."):
- auth_manager.deserialize_user({"user_id": "user_id", "name":
"name"})
-
-
@patch("airflow.providers.keycloak.auth_manager.keycloak_auth_manager._get_keycloak_jwt")
- def test_deserialize_user_doesnt_match(self, mock_get_keycloak_jwt,
auth_manager):
- mock_get_keycloak_jwt.return_value = KeycloakAuthManagerUser(
- user_id="user_2",
- name="name",
- access_token="access_token",
- refresh_token="refresh_token",
- )
- with pytest.raises(ValueError, match="Keycloak user in Cookies does
not match Airflow JWT."):
- auth_manager.deserialize_user({"user_id": "user_id", "name":
"name"})
+ assert result.access_token == ""
+ assert result.refresh_token is None
def test_serialize_user(self, auth_manager):
result = auth_manager.serialize_user(
@@ -156,10 +157,65 @@ class TestKeycloakAuthManager:
user_id="user_id", name="name", access_token="access_token",
refresh_token="refresh_token"
)
)
- assert result == {
- "user_id": "user_id",
- "name": "name",
- }
+ assert result == {"user_id": "user_id", "name": "name"}
+
+ @pytest.mark.asyncio
+ async def test_get_user_from_token(self, auth_manager):
+ mock_get_user_from_token = AsyncMock(
+ return_value=KeycloakAuthManagerUser(
+ user_id="user_id", name="name", access_token="",
refresh_token=None
+ )
+ )
+ with (
+ patch.object(
+ BaseAuthManager,
+ "get_user_from_token",
+ mock_get_user_from_token,
+ ),
+ ):
+ user = await auth_manager.get_user_from_token("token",
"access_token", "refresh_token")
+ mock_get_user_from_token.assert_called_with("token")
+ assert user.get_id() == "user_id"
+ assert user.get_name() == "name"
+ assert user.access_token == "access_token"
+ assert user.refresh_token == "refresh_token"
+
+ @pytest.mark.asyncio
+ async def test_get_user_from_token_keycloak_jwts_missing(self,
auth_manager):
+ mock_get_user_from_token = AsyncMock(
+ return_value=KeycloakAuthManagerUser(
+ user_id="user_id", name="name", access_token="",
refresh_token=None
+ )
+ )
+ with (
+ patch.object(
+ BaseAuthManager,
+ "get_user_from_token",
+ mock_get_user_from_token,
+ ),
+ ):
+ assert await auth_manager.get_user_from_token("token") is None
+
+ @pytest.mark.asyncio
+ async def test_get_user_from_token_keycloak_jwt(self, auth_manager):
+ mock_get_user_from_token = AsyncMock(
+ return_value=KeycloakAuthManagerUser(
+ user_id="user_id", name="name", access_token="",
refresh_token=None
+ )
+ )
+ with (
+ patch.object(
+ BaseAuthManager,
+ "get_user_from_token",
+ mock_get_user_from_token,
+ ),
+ ):
+ user = await auth_manager.get_user_from_token("token",
"access_token", "refresh_token")
+ mock_get_user_from_token.assert_called_with("token")
+ assert user.get_id() == "user_id"
+ assert user.get_name() == "name"
+ assert user.access_token == "access_token"
+ assert user.refresh_token == "refresh_token"
def test_get_url_login(self, auth_manager):
result = auth_manager.get_url_login()
@@ -177,6 +233,10 @@ class TestKeycloakAuthManager:
assert result is None
+ def test_refresh_user_not_user(self, auth_manager):
+ """When called from JWTRefreshMiddleware, ensure a None user can be
passed through."""
+ assert auth_manager.refresh_user(user=None) is None
+
def test_refresh_user_no_refresh_token(self, auth_manager):
"""Test that refresh_user returns None when refresh_token is empty
(client_credentials case)."""
user_without_refresh = Mock()
@@ -1311,3 +1371,6 @@ class TestKeycloakAuthManager:
auth_manager.filter_authorized_dag_ids(dag_ids=dag_ids, user=user)
mock_executor.assert_called_once_with(max_workers=expected_max_workers)
+
+ def test_get_fastapi_middleware(self, auth_manager):
+ assert auth_manager.get_fastapi_middlewares() ==
[(KeycloakJWTMiddleware, {})]
diff --git
a/providers/keycloak/tests/unit/keycloak/auth_manager/test_middleware.py
b/providers/keycloak/tests/unit/keycloak/auth_manager/test_middleware.py
new file mode 100644
index 00000000000..29a72988823
--- /dev/null
+++ b/providers/keycloak/tests/unit/keycloak/auth_manager/test_middleware.py
@@ -0,0 +1,372 @@
+# Licensed to the Apache Software Foundation (ASF) under one
+# or more contributor license agreements. See the NOTICE file
+# distributed with this work for additional information
+# regarding copyright ownership. The ASF licenses this file
+# to you under the Apache License, Version 2.0 (the
+# "License"); you may not use this file except in compliance
+# with the License. You may obtain a copy of the License at
+#
+# http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing,
+# software distributed under the License is distributed on an
+# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
+# KIND, either express or implied. See the License for the
+# specific language governing permissions and limitations
+# under the License.
+from __future__ import annotations
+
+from unittest.mock import AsyncMock, MagicMock, Mock, patch
+
+import pytest
+from fastapi import Request
+from jwt import InvalidTokenError
+
+from airflow.api_fastapi.auth.managers.base_auth_manager import
COOKIE_NAME_JWT_TOKEN
+from airflow.api_fastapi.core_api import security as core_api_security
+from airflow.providers.keycloak.auth_manager.constants import (
+ COOKIE_NAME_ACCESS_TOKEN,
+ COOKIE_NAME_REFRESH_TOKEN,
+)
+from airflow.providers.keycloak.auth_manager.middleware import
KeycloakJWTMiddleware
+from airflow.providers.keycloak.auth_manager.user import
KeycloakAuthManagerUser
+
+from tests_common.test_utils.version_compat import AIRFLOW_V_3_1_7_PLUS
+
+if AIRFLOW_V_3_1_7_PLUS:
+ from airflow.api_fastapi.auth.managers.exceptions import
AuthManagerRefreshTokenExpiredException
+else:
+ AuthManagerRefreshTokenExpiredException = None # type:
ignore[assignment,misc]
+
+
+def pytest_generate_tests(metafunc):
+ if "secure" in metafunc.fixturenames:
+ metafunc.parametrize("secure", [True, False], indirect=True)
+
+
[email protected]
+class TestKeycloakJWTMiddleware:
+ @pytest.fixture
+ def middleware(self):
+ return KeycloakJWTMiddleware(app=Mock(name="app"))
+
+ @pytest.fixture
+ def mock_request(self, secure):
+ request = MagicMock(spec=Request, name="request")
+ request.base_url.scheme = "https" if secure else "http"
+ request.cookies = {}
+ request.headers = {}
+ request.state = MagicMock(name="state", spec=[])
+ request.state.user = None
+ del request.state.user_authenticated_via
+ return request
+
+ @pytest.fixture
+ def mock_user(self):
+ user = Mock(name="user", spec=KeycloakAuthManagerUser)
+ user.user_id = "user_id"
+ user.name = "name"
+ user.access_token = "access_token"
+ user.refresh_token = "refresh_token"
+ return user
+
+ @pytest.fixture
+ def call_next(self):
+ return AsyncMock(return_value=Mock(name="response"), name="call_next")
+
+ @pytest.fixture
+ def auth_manager(self):
+ return Mock(name="auth_manager")
+
+ @pytest.fixture
+ def secure(self, request):
+ return request.param
+
+
@patch("airflow.providers.keycloak.auth_manager.middleware.get_auth_manager")
+ async def test_get_keycloak_tokens_from_cookies(
+ self, mock_get_auth_manager, auth_manager, call_next, mock_request,
middleware, mock_user
+ ):
+ auth_manager.get_user_from_token = AsyncMock(return_value=mock_user)
+ auth_manager.refresh_user.return_value = None
+ mock_get_auth_manager.return_value = auth_manager
+
+ mock_request.cookies = {
+ COOKIE_NAME_JWT_TOKEN: "token",
+ COOKIE_NAME_ACCESS_TOKEN: "access_token",
+ COOKIE_NAME_REFRESH_TOKEN: "refresh_token",
+ }
+
+ await middleware.dispatch(mock_request, call_next)
+
+ assert mock_request.state.user is mock_user
+ assert mock_request.state.user.access_token == "access_token"
+ assert mock_request.state.user.refresh_token == "refresh_token"
+
+ trusted_marker = getattr(
+ core_api_security,
+ "USER_INJECTED_BY_TRUSTED_MIDDLEWARE",
+ None,
+ )
+
+ if trusted_marker is not None:
+ assert mock_request.state.user_authenticated_via is trusted_marker
+ else:
+ assert not hasattr(mock_request.state, "user_authenticated_via")
+
+ auth_manager.get_user_from_token.assert_called_once_with("token",
"access_token", "refresh_token")
+ auth_manager.refresh_user.assert_called_once_with(user=mock_user)
+ call_next.assert_awaited_once_with(mock_request)
+
+
@patch("airflow.providers.keycloak.auth_manager.middleware.get_auth_manager")
+ async def test_refresh_keycloak_token(
+ self,
+ mock_get_auth_manager,
+ auth_manager,
+ call_next,
+ mock_request,
+ middleware,
+ mock_user,
+ secure,
+ ):
+ new_user = Mock(name="user", spec=KeycloakAuthManagerUser)
+ new_user.access_token = "new_access_token"
+ new_user.refresh_token = "new_refresh_token"
+ auth_manager.get_user_from_token = AsyncMock(return_value=mock_user)
+ auth_manager.refresh_user = Mock(return_value=new_user)
+ auth_manager.generate_jwt.return_value = "new_token"
+ mock_get_auth_manager.return_value = auth_manager
+
+ mock_request.cookies = {
+ COOKIE_NAME_JWT_TOKEN: "token",
+ COOKIE_NAME_ACCESS_TOKEN: "access_token",
+ COOKIE_NAME_REFRESH_TOKEN: "refresh_token",
+ }
+
+ response = await middleware.dispatch(mock_request, call_next)
+
+ assert mock_request.state.user is new_user
+ assert mock_request.state.user.access_token == "new_access_token"
+ assert mock_request.state.user.refresh_token == "new_refresh_token"
+
+ response.set_cookie.assert_any_call(
+ COOKIE_NAME_JWT_TOKEN,
+ "new_token",
+ path="/",
+ secure=secure,
+ samesite="lax",
+ httponly=True,
+ max_age=None,
+ )
+ response.set_cookie.assert_any_call(
+ COOKIE_NAME_ACCESS_TOKEN,
+ "new_access_token",
+ path="/",
+ samesite="lax",
+ secure=secure,
+ httponly=True,
+ )
+ response.set_cookie.assert_any_call(
+ COOKIE_NAME_REFRESH_TOKEN,
+ "new_refresh_token",
+ path="/",
+ samesite="lax",
+ secure=secure,
+ httponly=True,
+ )
+
+ trusted_marker = getattr(
+ core_api_security,
+ "USER_INJECTED_BY_TRUSTED_MIDDLEWARE",
+ None,
+ )
+
+ if trusted_marker is not None:
+ assert mock_request.state.user_authenticated_via is trusted_marker
+ else:
+ assert not hasattr(mock_request.state, "user_authenticated_via")
+
+ auth_manager.get_user_from_token.assert_called_once_with("token",
"access_token", "refresh_token")
+ auth_manager.refresh_user.assert_called_once_with(user=mock_user)
+ auth_manager.generate_jwt.assert_called_once_with(new_user)
+ call_next.assert_awaited_once_with(mock_request)
+
+
@patch("airflow.providers.keycloak.auth_manager.middleware.get_auth_manager")
+ async def test_no_keycloak_token(
+ self, mock_get_auth_manager, auth_manager, call_next, middleware,
mock_request, secure
+ ):
+ mock_get_auth_manager.return_value = auth_manager
+
+ mock_request.cookies = {COOKIE_NAME_JWT_TOKEN: "token"}
+
+ response = await middleware.dispatch(mock_request, call_next)
+
+ auth_manager.get_user_from_token.assert_not_called()
+ auth_manager.refresh_user.assert_not_called()
+
+ assert mock_request.state.user is None
+
+ trusted_marker = getattr(
+ core_api_security,
+ "USER_INJECTED_BY_TRUSTED_MIDDLEWARE",
+ None,
+ )
+
+ if trusted_marker is not None:
+ assert getattr(mock_request.state, "user_authenticated_via", None)
is not trusted_marker
+ else:
+ assert not hasattr(mock_request.state, "user_authenticated_via")
+
+ call_next.assert_awaited_with(mock_request)
+
+ response.set_cookie.assert_any_call(
+ COOKIE_NAME_JWT_TOKEN,
+ "",
+ path="/",
+ secure=secure,
+ httponly=True,
+ samesite="lax",
+ max_age=0,
+ )
+
+
@patch("airflow.providers.keycloak.auth_manager.middleware.get_auth_manager")
+ async def test_no_airflow_jwt_token(
+ self, mock_get_auth_manager, auth_manager, call_next, middleware,
mock_request
+ ):
+ mock_get_auth_manager.return_value = auth_manager
+
+ mock_request.cookies = {
+ COOKIE_NAME_ACCESS_TOKEN: "access_token",
+ COOKIE_NAME_REFRESH_TOKEN: "refresh_token",
+ }
+
+ await middleware.dispatch(mock_request, call_next)
+
+ auth_manager.get_user_from_token.assert_not_called()
+ auth_manager.refresh_user.assert_not_called()
+
+ assert mock_request.state.user is None
+
+ trusted_marker = getattr(
+ core_api_security,
+ "USER_INJECTED_BY_TRUSTED_MIDDLEWARE",
+ None,
+ )
+
+ if trusted_marker is not None:
+ assert getattr(mock_request.state, "user_authenticated_via", None)
is not trusted_marker
+ else:
+ assert not hasattr(mock_request.state, "user_authenticated_via")
+
+ call_next.assert_awaited_once_with(mock_request)
+
+
@patch("airflow.providers.keycloak.auth_manager.middleware.get_auth_manager")
+ @pytest.mark.asyncio
+ async def test_dispatch_expired_token(
+ self,
+ mock_get_auth_manager,
+ auth_manager,
+ call_next,
+ middleware,
+ mock_request,
+ secure,
+ ):
+ mock_get_auth_manager.return_value = auth_manager
+ mock_request.cookies = {
+ COOKIE_NAME_JWT_TOKEN: "invalid_token",
+ COOKIE_NAME_ACCESS_TOKEN: "access_token",
+ COOKIE_NAME_REFRESH_TOKEN: "refresh_token",
+ }
+ auth_manager.get_user_from_token.side_effect = InvalidTokenError()
+
+ response = await middleware.dispatch(mock_request, call_next)
+
+ call_next.assert_called_once_with(mock_request)
+ auth_manager.get_user_from_token.assert_called_once_with(
+ "invalid_token", "access_token", "refresh_token"
+ )
+
+ response.set_cookie.assert_any_call(
+ COOKIE_NAME_JWT_TOKEN,
+ "",
+ path="/",
+ secure=secure,
+ httponly=True,
+ samesite="lax",
+ max_age=0,
+ )
+
+
@patch("airflow.providers.keycloak.auth_manager.middleware.get_auth_manager")
+ @pytest.mark.asyncio
+ async def test_dispatch_expired_keycloak_token(
+ self,
+ mock_get_auth_manager,
+ auth_manager,
+ call_next,
+ middleware,
+ mock_request,
+ mock_user,
+ secure,
+ ):
+ mock_get_auth_manager.return_value = auth_manager
+ mock_request.cookies = {
+ COOKIE_NAME_JWT_TOKEN: "token",
+ COOKIE_NAME_ACCESS_TOKEN: "expired_token",
+ COOKIE_NAME_REFRESH_TOKEN: "refresh_token",
+ }
+ mock_user.access_token = "expired_token"
+ mock_user.refresh_token = "refresh_token"
+ auth_manager.get_user_from_token = AsyncMock(return_value=mock_user)
+ if AIRFLOW_V_3_1_7_PLUS:
+ auth_manager.refresh_user.side_effect =
AuthManagerRefreshTokenExpiredException()
+ else:
+ auth_manager.refresh_user.return_value = None
+
+ response = await middleware.dispatch(mock_request, call_next)
+
+ call_next.assert_called_once_with(mock_request)
+ auth_manager.get_user_from_token.assert_called_once_with("token",
"expired_token", "refresh_token")
+ auth_manager.refresh_user.assert_called_once_with(user=mock_user)
+
+ if AIRFLOW_V_3_1_7_PLUS:
+ response.set_cookie.assert_any_call(
+ COOKIE_NAME_JWT_TOKEN,
+ "",
+ path="/",
+ secure=secure,
+ httponly=True,
+ samesite="lax",
+ max_age=0,
+ )
+ auth_manager.generate_jwt.assert_not_called()
+
+
@patch("airflow.providers.keycloak.auth_manager.middleware.get_cookie_path")
+
@patch("airflow.providers.keycloak.auth_manager.middleware.get_auth_manager")
+ @pytest.mark.asyncio
+ async def test_dispatch_invalid_token_clears_root_cookie(
+ self,
+ mock_get_auth_manager,
+ mock_get_cookie_path,
+ auth_manager,
+ call_next,
+ middleware,
+ mock_request,
+ secure,
+ ):
+ mock_get_cookie_path.return_value = "/foo/"
+ mock_get_auth_manager.return_value = auth_manager
+ auth_manager.get_user_from_token.side_effect = InvalidTokenError()
+ """When a stale _token exists at root path, clearing must target both
the subpath and root."""
+ mock_request.cookies = {
+ COOKIE_NAME_JWT_TOKEN: "stale_root_token",
+ COOKIE_NAME_ACCESS_TOKEN: "access_token",
+ COOKIE_NAME_REFRESH_TOKEN: "refresh_token",
+ }
+ response = await middleware.dispatch(mock_request, call_next)
+
+ # Expect two delete cookies: one at the subpath and one at root "/"
+ response.set_cookie.assert_any_call(
+ "_token", "", path="/foo/", secure=secure, samesite="lax",
httponly=True, max_age=0
+ )
+ response.set_cookie.assert_any_call(
+ "_token", "", path="/", secure=secure, samesite="lax",
httponly=True, max_age=0
+ )