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 86b8a521837 Bind the AWS auth manager SAML response to the browser 
that started the login (#73698)
86b8a521837 is described below

commit 86b8a521837ee4abe56c5d85263ce0a2643d32f4
Author: Jarek Potiuk <[email protected]>
AuthorDate: Mon Oct 5 00:09:08 2026 +0200

    Bind the AWS auth manager SAML response to the browser that started the 
login (#73698)
    
    * Bind the AWS auth manager SAML response to the browser that started the 
login
    
    login() discarded the AuthnRequest id that saml_auth.login() generates, and
    login_callback() called process_response() with no request_id. python3-saml
    only validates InResponseTo when it is given that id, so the check was
    skipped and any assertion the IdP had signed was accepted, including one
    issued for a login that a different browser started.
    
    Each login now gets a nonce that travels to the IdP in RelayState and names,
    on the way back, which of this browser's pending logins the response 
answers.
    The pending set lives in a short-lived signed cookie (HttpOnly, 
SameSite=Lax,
    Secure behind TLS, 10 minutes, one entry per login). At the callback the
    response must match one of those entries and InResponseTo must match that
    entry's AuthnRequest id; the entry is consumed on use.
    
    IdP-initiated SSO (for example the AWS Identity Center portal tile) is kept
    available behind [aws_auth_manager] allow_idp_initiated_login, default 
False.
    With it enabled the assertion must still be signed by the configured IdP and
    a response carrying InResponseTo is refused.
    
    Also fix TestLoginRouter never being collected (mock_plugin_manager applied
    as a class decorator), register it in check_contextmanager_class_decorators,
    and fix the two assertions that had never run.
    
    Generated-by: Claude Opus 5
    
    * Add idp to the spelling wordlist
    
    Generated-by: Claude Opus 5
---
 docs/spelling_wordlist.txt                         |   1 +
 .../docs/auth-manager/setup/identity-center.rst    |  36 +++
 providers/amazon/provider.yaml                     |  15 +
 .../providers/amazon/aws/auth_manager/constants.py |   1 +
 .../amazon/aws/auth_manager/routes/login.py        | 200 ++++++++++--
 .../airflow/providers/amazon/get_provider_info.py  |   7 +
 .../amazon/aws/auth_manager/routes/test_login.py   | 359 +++++++++++++++++----
 .../prek/check_contextmanager_class_decorators.py  |   1 +
 .../test_check_contextmanager_class_decorators.py  |  80 +++++
 9 files changed, 627 insertions(+), 73 deletions(-)

diff --git a/docs/spelling_wordlist.txt b/docs/spelling_wordlist.txt
index 2719dde430d..6db426432a6 100644
--- a/docs/spelling_wordlist.txt
+++ b/docs/spelling_wordlist.txt
@@ -840,6 +840,7 @@ ideation
 idempotence
 idempotency
 IdP
+idp
 ie
 iframe
 iframes
diff --git a/providers/amazon/docs/auth-manager/setup/identity-center.rst 
b/providers/amazon/docs/auth-manager/setup/identity-center.rst
index 25fc80d807e..134fd121fee 100644
--- a/providers/amazon/docs/auth-manager/setup/identity-center.rst
+++ b/providers/amazon/docs/auth-manager/setup/identity-center.rst
@@ -123,3 +123,39 @@ or
 
     export AIRFLOW__API__BASE_URL='<base_url>'
     export 
AIRFLOW__AWS_AUTH_MANAGER__SAML_METADATA_URL='<saml_metadata_file_url>'
+
+.. _identity_center_idp_initiated_login:
+
+Logging in from the access portal (IdP-initiated SSO)
+=====================================================
+
+By default the AWS auth manager only accepts a SAML response that answers a 
login started from
+Airflow itself. When you open ``<base_url>/auth/login``, Airflow sends an 
AuthnRequest to Identity
+Center and remembers its id in a short-lived, signed, ``HttpOnly`` cookie; the 
response is accepted
+only if it echoes that id back in ``InResponseTo``.
+
+This binding matters because a SAML assertion is signed by the identity 
provider, which
+authenticates *the identity in the response* — it says nothing about *which 
browser asked for it*.
+Without the binding, an assertion issued for one login would be accepted by a 
browser that never
+started it, signing that browser in as the assertion's subject rather than as 
the person using it.
+
+The cost is that clicking the Airflow tile in the AWS IAM Identity Center 
access portal no longer
+works: that flow produces an assertion nobody asked for, and Airflow cannot 
tell it apart from a
+replayed one. If you need it, enable it explicitly:
+
+.. code-block:: ini
+
+    [aws_auth_manager]
+    allow_idp_initiated_login = True
+
+or
+
+.. code-block:: bash
+
+    export AIRFLOW__AWS_AUTH_MANAGER__ALLOW_IDP_INITIATED_LOGIN='True'
+
+With this enabled, Airflow still requires the assertion to be signed by the 
configured identity
+provider, and still refuses one that carries an ``InResponseTo`` value — an 
unsolicited assertion
+answers no request, so a response that names one is a solicited assertion 
being replayed. What it
+gives up is the guarantee that the browser receiving the response is the one 
that asked for it.
+Leave it disabled unless the access portal flow is required.
diff --git a/providers/amazon/provider.yaml b/providers/amazon/provider.yaml
index f6d91b78fcf..3e2d9c2c493 100644
--- a/providers/amazon/provider.yaml
+++ b/providers/amazon/provider.yaml
@@ -1558,6 +1558,21 @@ config:
         type: string
         example: ~
         default: ~
+      allow_idp_initiated_login:
+        description: |
+          Whether to accept SAML assertions that no login started from Airflow 
asked for, such as
+          the ones produced by clicking the Airflow tile in the AWS Identity 
Center access portal
+          (IdP-initiated SSO).
+
+          When this is disabled, a SAML response is only accepted if it 
answers an AuthnRequest that
+          the same browser started. That binding is what stops an assertion 
obtained elsewhere from
+          being accepted in another user's browser, which would sign that user 
in as the assertion's
+          subject. Enable it only where the access portal flow is required and 
that trade-off is
+          accepted.
+        version_added: 9.37.0
+        type: boolean
+        example: "True"
+        default: "False"
 
 executors:
   - 
airflow.providers.amazon.aws.executors.aws_lambda.lambda_executor.AwsLambdaExecutor
diff --git 
a/providers/amazon/src/airflow/providers/amazon/aws/auth_manager/constants.py 
b/providers/amazon/src/airflow/providers/amazon/aws/auth_manager/constants.py
index b05636fdaab..1841b82c6e4 100644
--- 
a/providers/amazon/src/airflow/providers/amazon/aws/auth_manager/constants.py
+++ 
b/providers/amazon/src/airflow/providers/amazon/aws/auth_manager/constants.py
@@ -23,3 +23,4 @@ CONF_CONN_ID_KEY = "conn_id"
 CONF_REGION_NAME_KEY = "region_name"
 CONF_SAML_METADATA_URL_KEY = "saml_metadata_url"
 CONF_AVP_POLICY_STORE_ID_KEY = "avp_policy_store_id"
+CONF_ALLOW_IDP_INITIATED_LOGIN_KEY = "allow_idp_initiated_login"
diff --git 
a/providers/amazon/src/airflow/providers/amazon/aws/auth_manager/routes/login.py
 
b/providers/amazon/src/airflow/providers/amazon/aws/auth_manager/routes/login.py
index d88c2b15252..4f7af3fed31 100644
--- 
a/providers/amazon/src/airflow/providers/amazon/aws/auth_manager/routes/login.py
+++ 
b/providers/amazon/src/airflow/providers/amazon/aws/auth_manager/routes/login.py
@@ -17,13 +17,20 @@
 
 from __future__ import annotations
 
+import base64
+import binascii
+import hmac
+import json
 import logging
+import secrets
+import time
+from hashlib import sha256
 from typing import TYPE_CHECKING, Any
 from urllib.parse import urlparse
 
 import anyio
 from fastapi import HTTPException, Request, status
-from fastapi.responses import RedirectResponse
+from fastapi.responses import JSONResponse, RedirectResponse
 
 from airflow.api_fastapi.app import (
     AUTH_MANAGER_FASTAPI_APP_PREFIX,
@@ -31,7 +38,11 @@ from airflow.api_fastapi.app import (
 )
 from airflow.api_fastapi.auth.managers.base_auth_manager import 
COOKIE_NAME_JWT_TOKEN
 from airflow.api_fastapi.common.router import AirflowRouter
-from airflow.providers.amazon.aws.auth_manager.constants import 
CONF_SAML_METADATA_URL_KEY, CONF_SECTION_NAME
+from airflow.providers.amazon.aws.auth_manager.constants import (
+    CONF_ALLOW_IDP_INITIATED_LOGIN_KEY,
+    CONF_SAML_METADATA_URL_KEY,
+    CONF_SECTION_NAME,
+)
 from airflow.providers.amazon.aws.auth_manager.datamodels.login import 
LoginResponse
 from airflow.providers.amazon.aws.auth_manager.user import AwsAuthManagerUser
 from airflow.providers.amazon.version_compat import AIRFLOW_V_3_1_1_PLUS, 
AIRFLOW_V_3_1_8_PLUS
@@ -59,6 +70,101 @@ except ImportError:
 log = logging.getLogger(__name__)
 login_router = AirflowRouter(tags=["AWSAuthManagerLogin"])
 
+# Name of the short-lived cookie recording which logins this browser has 
started.
+COOKIE_NAME_LOGIN_STATE = "_awsam_login_state"
+
+# The login flow is a redirect to the IdP and back. Ten minutes is generous 
for that and
+# keeps a stale request id from lingering.
+LOGIN_STATE_MAX_AGE = 600
+
+# One entry per login started and not yet completed, so opening a second tab 
does not
+# invalidate the first. The cap bounds the cookie; the oldest pending login is 
dropped.
+MAX_PENDING_LOGINS = 5
+
+LOGIN_MODE_REDIRECT = "login-redirect"
+LOGIN_MODE_TOKEN = "login-token"
+LOGIN_MODES = (LOGIN_MODE_REDIRECT, LOGIN_MODE_TOKEN)
+
+NO_LOGIN_IN_PROGRESS = "No login in progress for this browser. Start the login 
from Airflow and try again."
+
+
+def _is_secure_request(request: Request) -> bool:
+    return request.base_url.scheme == "https" or bool(conf.get("api", 
"ssl_cert", fallback=""))
+
+
+def _allows_idp_initiated_login() -> bool:
+    return conf.getboolean(CONF_SECTION_NAME, 
CONF_ALLOW_IDP_INITIATED_LOGIN_KEY, fallback=False)
+
+
+def _sign_login_state(payload: str) -> str:
+    # The API server secret key is already required to be identical across API 
servers, so a
+    # login may start on one instance and finish on another.
+    secret = conf.get("api", "secret_key", fallback="")
+    return hmac.new(secret.encode(), payload.encode(), sha256).hexdigest()
+
+
+def _read_pending_logins(request: Request) -> list[dict[str, Any]]:
+    """
+    Return the logins this browser started and has not yet completed.
+
+    This cookie is the browser's half of the binding: it states that *this* 
browser asked for
+    these AuthnRequests. A SAML assertion is signed by the identity provider, 
which
+    authenticates *the identity in the response* -- it says nothing about 
*which browser
+    asked*. Without this, an assertion issued for one login would be accepted 
in a browser
+    that never started it, signing that browser in as the assertion's subject.
+
+    Entries are signed so a response cannot contribute one of its own, and 
each carries its
+    own deadline so an abandoned tab expires without affecting the others.
+    """
+    raw = request.cookies.get(COOKIE_NAME_LOGIN_STATE)
+    if not raw:
+        return []
+    payload, _, signature = raw.rpartition(".")
+    if not payload or not hmac.compare_digest(signature, 
_sign_login_state(payload)):
+        log.warning("Ignoring a login state cookie that this deployment did 
not sign.")
+        return []
+    try:
+        entries = json.loads(base64.urlsafe_b64decode(payload))
+    except (ValueError, binascii.Error):
+        log.warning("Ignoring a login state cookie that could not be decoded.")
+        return []
+    if not isinstance(entries, list):
+        return []
+    now = time.time()
+    return [
+        entry
+        for entry in entries
+        if isinstance(entry, dict) and isinstance(entry.get("exp"), (int, 
float)) and entry["exp"] > now
+    ]
+
+
+def _write_pending_logins(request: Request, response: Any, entries: 
list[dict[str, Any]]) -> None:
+    cookie_path = get_cookie_path()
+    if not entries:
+        response.delete_cookie(COOKIE_NAME_LOGIN_STATE, path=cookie_path)
+        return
+    payload = base64.urlsafe_b64encode(json.dumps(entries, separators=(",", 
":")).encode()).decode()
+    response.set_cookie(
+        COOKIE_NAME_LOGIN_STATE,
+        f"{payload}.{_sign_login_state(payload)}",
+        max_age=LOGIN_STATE_MAX_AGE,
+        path=cookie_path,
+        secure=_is_secure_request(request),
+        httponly=True,
+        samesite="lax",
+    )
+
+
+def _match_pending_login(pending: list[dict[str, Any]], relay_state: str) -> 
dict[str, Any] | None:
+    """Find which of this browser's pending logins a response claims to 
answer."""
+    mode, _, nonce = relay_state.partition(":")
+    if not nonce or mode not in LOGIN_MODES:
+        return None
+    for entry in pending:
+        if entry.get("mode") == mode and 
hmac.compare_digest(str(entry.get("nonce", "")), nonce):
+            return entry
+    return None
+
 
 def _read_form(request: Request) -> FormData:
     """Read the request form from a synchronous context running in a worker 
thread."""
@@ -69,28 +175,73 @@ def _read_form(request: Request) -> FormData:
     return anyio.from_thread.run(_form)
 
 
+def _start_login(request: Request, mode: str) -> RedirectResponse:
+    """
+    Begin a login, remembering enough about it to recognise its response later.
+
+    The nonce travels to the IdP in ``RelayState`` and returns with the 
response, naming
+    which of this browser's pending logins that response answers. It is only 
an index into
+    the signed cookie; the binding itself is the AuthnRequest id, which the 
IdP echoes in
+    ``InResponseTo`` and which whoever posts the response cannot choose. The 
return mode is
+    read back from the matched entry rather than from the form, so a response 
cannot select
+    a mode the browser did not ask for.
+    """
+    saml_auth = _init_saml_auth(request)
+    nonce = secrets.token_urlsafe(16)
+    callback_url = saml_auth.login(f"{mode}:{nonce}")
+    response = RedirectResponse(url=callback_url)
+    pending = _read_pending_logins(request)[-(MAX_PENDING_LOGINS - 1) :]
+    pending.append(
+        {
+            "nonce": nonce,
+            "request_id": saml_auth.get_last_request_id(),
+            "mode": mode,
+            "exp": time.time() + LOGIN_STATE_MAX_AGE,
+        }
+    )
+    _write_pending_logins(request, response, pending)
+    return response
+
+
 @login_router.get("/login")
 def login(request: Request):
     """Initiate the authentication."""
-    saml_auth = _init_saml_auth(request)
-    callback_url = saml_auth.login("login-redirect")
-    return RedirectResponse(url=callback_url)
+    return _start_login(request, LOGIN_MODE_REDIRECT)
 
 
 @login_router.get("/login/token")
 def login_token(request: Request) -> RedirectResponse:
     """Initiate the authentication to create a token."""
-    saml_auth = _init_saml_auth(request)
-    callback_url = saml_auth.login("login-token")
-    return RedirectResponse(url=callback_url)
+    return _start_login(request, LOGIN_MODE_TOKEN)
 
 
 @login_router.post("/login_callback")
 def login_callback(request: Request):
     """Authenticate the user."""
+    form_data = _read_form(request)
+    pending = _read_pending_logins(request)
+    relay_state = form_data.get("RelayState")
+    # A multipart upload posted under this name is not a relay state, so it 
matches nothing.
+    matched = _match_pending_login(pending, relay_state) if 
isinstance(relay_state, str) else None
+
+    if matched is not None:
+        mode = matched["mode"]
+        expected_request_id = matched["request_id"]
+    elif _allows_idp_initiated_login():
+        # Opted in: this deployment accepts assertions that no login from this 
browser asked
+        # for, so the Identity Center access portal tile keeps working. 
Nothing ties such a
+        # response to the browser receiving it -- that is the trade the option 
names.
+        mode = LOGIN_MODE_REDIRECT
+        expected_request_id = None
+    else:
+        log.error("SAML response received that answers no login started by 
this browser.")
+        raise HTTPException(status.HTTP_401_UNAUTHORIZED, NO_LOGIN_IN_PROGRESS)
+
     saml_auth = _init_saml_auth(request)
     try:
-        saml_auth.process_response()
+        # Passing the request id makes python3-saml enforce InResponseTo. 
Without it the
+        # check is skipped entirely and any valid assertion is accepted.
+        saml_auth.process_response(request_id=expected_request_id)
     except OneLogin_Saml2_Error as e:
         log.exception(e)
         raise HTTPException(status.HTTP_500_INTERNAL_SERVER_ERROR, "Failed to 
authenticate")
@@ -103,6 +254,13 @@ def login_callback(request: Request):
         log.error("Error reason: %s", error_reason)
         raise HTTPException(status.HTTP_500_INTERNAL_SERVER_ERROR, f"Failed to 
authenticate: {error_reason}")
 
+    if expected_request_id is None and 
saml_auth.get_last_response_in_response_to() is not None:
+        # An unsolicited assertion answers no request, so one carrying 
InResponseTo is a
+        # solicited assertion being replayed here. python3-saml skips the 
comparison entirely
+        # when it is given no request id, so this is checked rather than 
assumed.
+        log.error("Unsolicited SAML response carries InResponseTo; refusing it 
as a replay.")
+        raise HTTPException(status.HTTP_401_UNAUTHORIZED, "Invalid SAML 
response")
+
     attributes = saml_auth.get_attributes()
     user = AwsAuthManagerUser(
         user_id=attributes["id"][0],
@@ -113,23 +271,27 @@ def login_callback(request: Request):
     url = conf.get("api", "base_url", fallback="/")
     token = get_auth_manager().generate_jwt(user)
 
-    form_data = _read_form(request)
-    relay_state = form_data["RelayState"]
-
-    if relay_state == "login-redirect":
+    response: Any
+    if mode == LOGIN_MODE_REDIRECT:
         response = RedirectResponse(url=url, status_code=303)
-        secure = request.base_url.scheme == "https" or bool(conf.get("api", 
"ssl_cert", fallback=""))
+        cookie_path = get_cookie_path()
+        secure = _is_secure_request(request)
         # In Airflow 3.1.1 authentication changes, front-end no longer handle 
the token
         # See https://github.com/apache/airflow/pull/55506
-        cookie_path = get_cookie_path()
         if AIRFLOW_V_3_1_1_PLUS:
             response.set_cookie(COOKIE_NAME_JWT_TOKEN, token, 
path=cookie_path, secure=secure, httponly=True)
         else:
             response.set_cookie(COOKIE_NAME_JWT_TOKEN, token, 
path=cookie_path, secure=secure)
-        return response
-    if relay_state == "login-token":
-        return LoginResponse(access_token=token)
-    raise HTTPException(status.HTTP_500_INTERNAL_SERVER_ERROR, f"Invalid relay 
state: {relay_state}")
+    else:
+        # Returned as a JSONResponse rather than the bare model so the 
consumed login state
+        # can be cleared on this path too. Left in place, the same assertion 
could be reposted
+        # to mint further tokens until the cookie expired.
+        response = 
JSONResponse(content=LoginResponse(access_token=token).model_dump())
+
+    if matched is not None:
+        # One response per request. Logins started in other tabs keep theirs.
+        _write_pending_logins(request, response, [entry for entry in pending 
if entry is not matched])
+    return response
 
 
 def _init_saml_auth(request: Request) -> OneLogin_Saml2_Auth:
diff --git a/providers/amazon/src/airflow/providers/amazon/get_provider_info.py 
b/providers/amazon/src/airflow/providers/amazon/get_provider_info.py
index 74a0f97fb6e..1eaed1193de 100644
--- a/providers/amazon/src/airflow/providers/amazon/get_provider_info.py
+++ b/providers/amazon/src/airflow/providers/amazon/get_provider_info.py
@@ -1597,6 +1597,13 @@ def get_provider_info():
                         "example": None,
                         "default": None,
                     },
+                    "allow_idp_initiated_login": {
+                        "description": "Whether to accept SAML assertions that 
no login started from Airflow asked for, such as\nthe ones produced by clicking 
the Airflow tile in the AWS Identity Center access portal\n(IdP-initiated 
SSO).\n\nWhen this is disabled, a SAML response is only accepted if it answers 
an AuthnRequest that\nthe same browser started. That binding is what stops an 
assertion obtained elsewhere from\nbeing accepted in another user's browser, 
which would sign that use [...]
+                        "version_added": "9.37.0",
+                        "type": "boolean",
+                        "example": "True",
+                        "default": "False",
+                    },
                 },
             },
         },
diff --git 
a/providers/amazon/tests/unit/amazon/aws/auth_manager/routes/test_login.py 
b/providers/amazon/tests/unit/amazon/aws/auth_manager/routes/test_login.py
index 350b1b9afe4..ba32814dbc3 100644
--- a/providers/amazon/tests/unit/amazon/aws/auth_manager/routes/test_login.py
+++ b/providers/amazon/tests/unit/amazon/aws/auth_manager/routes/test_login.py
@@ -16,6 +16,9 @@
 # under the License.
 from __future__ import annotations
 
+import base64
+import json
+import time
 from unittest.mock import Mock, patch
 
 import pytest
@@ -39,6 +42,15 @@ OneLogin_Saml2_IdPMetadataParser = pytest.importorskip(
     "onelogin.saml2.idp_metadata_parser"
 ).OneLogin_Saml2_IdPMetadataParser
 
+# Imported after the importorskip above: this module raises ImportError when 
python3-saml
+# is absent, and the lowest-dependency check deliberately runs without it.
+from airflow.providers.amazon.aws.auth_manager.routes.login import (  # noqa: 
E402
+    COOKIE_NAME_LOGIN_STATE,
+    LOGIN_MODE_REDIRECT,
+    LOGIN_MODE_TOKEN,
+    _sign_login_state,
+)
+
 SAML_METADATA_URL = "/saml/metadata"
 SAML_METADATA_PARSED = {
     "idp": {
@@ -58,17 +70,70 @@ SAML_METADATA_PARSED = {
 }
 
 
+EXPECTED_REQUEST_ID = "ONELOGIN_authn_request_id"
+TEST_NONCE = "nonce-for-this-login"
+TEST_SECRET_KEY = "login-state-signing-key"
+
+RELAY_REDIRECT = f"{LOGIN_MODE_REDIRECT}:{TEST_NONCE}"
+RELAY_TOKEN = f"{LOGIN_MODE_TOKEN}:{TEST_NONCE}"
+
+BASE_CONF = {
+    ("core", "auth_manager"): 
"airflow.providers.amazon.aws.auth_manager.aws_auth_manager.AwsAuthManager",
+    ("aws_auth_manager", "saml_metadata_url"): SAML_METADATA_URL,
+    ("api", "ssl_cert"): "",
+    ("api", "secret_key"): TEST_SECRET_KEY,
+    ("api", "base_url"): "http://localhost:8080/";,
+}
+
+
+def make_login_state(*entries: tuple[str, str, str], expires_in: float = 600) 
-> str:
+    """
+    Build the signed login-state cookie a browser would hold for ``entries``.
+
+    Each entry is ``(nonce, request_id, mode)``. Must be called inside the 
``conf_vars``
+    block that sets the signing key, so the signature matches what the route 
computes.
+    """
+    payload = base64.urlsafe_b64encode(
+        json.dumps(
+            [
+                {
+                    "nonce": nonce,
+                    "request_id": request_id,
+                    "mode": mode,
+                    "exp": time.time() + expires_in,
+                }
+                for nonce, request_id, mode in entries
+            ],
+            separators=(",", ":"),
+        ).encode()
+    ).decode()
+    return f"{payload}.{_sign_login_state(payload)}"
+
+
+def read_login_state_cookie(response) -> str | None:
+    """Return the login-state cookie value the response sets, if it sets 
one."""
+    for header in response.headers.get_list("set-cookie"):
+        if header.startswith(f"{COOKIE_NAME_LOGIN_STATE}="):
+            return header.split("=", 1)[1].split(";", 1)[0]
+    return None
+
+
[email protected](autouse=True)
+def no_plugins():
+    """
+    Keep the plugin manager out of these tests.
+
+    Applied as a fixture rather than as ``@mock_plugin_manager(...)`` on the 
class:
+    ``mock_plugin_manager`` is a ``contextmanager``, and a 
``ContextDecorator`` used on a class
+    replaces it with a function, which pytest then does not collect at all.
+    """
+    with mock_plugin_manager(plugins=[]):
+        yield
+
+
 @pytest.fixture
 def test_client():
-    with conf_vars(
-        {
-            (
-                "core",
-                "auth_manager",
-            ): 
"airflow.providers.amazon.aws.auth_manager.aws_auth_manager.AwsAuthManager",
-            ("aws_auth_manager", "saml_metadata_url"): SAML_METADATA_URL,
-        }
-    ):
+    with conf_vars(BASE_CONF):
         with (
             patch.object(OneLogin_Saml2_IdPMetadataParser, "parse_remote") as 
mock_parse_remote,
             patch(
@@ -80,17 +145,27 @@ def test_client():
             yield TestClient(create_app())
 
 
-def get_login_callback_response(relay_state: str, *, base_url: str = 
"http://testserver";):
-    with conf_vars(
-        {
-            (
-                "core",
-                "auth_manager",
-            ): 
"airflow.providers.amazon.aws.auth_manager.aws_auth_manager.AwsAuthManager",
-            ("aws_auth_manager", "saml_metadata_url"): SAML_METADATA_URL,
-            ("api", "ssl_cert"): "",
-        }
-    ):
+def get_login_callback_response(
+    relay_state: str,
+    *,
+    base_url: str = "http://testserver";,
+    pending: list[tuple[str, str, str]] | None = None,
+    raw_login_state: str | None = None,
+    expires_in: float = 600,
+    in_response_to: str | None = None,
+    is_authenticated: bool = True,
+    extra_conf: dict | None = None,
+    return_auth_mock: bool = False,
+):
+    """
+    Post a SAML response to the callback.
+
+    ``pending`` is what this browser has started, as ``(nonce, request_id, 
mode)`` tuples; it
+    defaults to a single entry answering ``relay_state``, i.e. a browser that 
really did start
+    this login. ``raw_login_state`` sets the cookie verbatim instead -- pass 
``""`` for a
+    browser that started nothing.
+    """
+    with conf_vars({**BASE_CONF, **(extra_conf or {})}):
         with (
             patch.object(OneLogin_Saml2_IdPMetadataParser, "parse_remote") as 
mock_parse_remote,
             patch(
@@ -104,23 +179,40 @@ def get_login_callback_response(relay_state: str, *, 
base_url: str = "http://tes
             mock_parse_remote.return_value = SAML_METADATA_PARSED
 
             auth = Mock()
-            auth.is_authenticated.return_value = True
+            auth.is_authenticated.return_value = is_authenticated
             auth.get_nameid.return_value = "user_id"
+            auth.get_last_response_in_response_to.return_value = in_response_to
             auth.get_attributes.return_value = {
                 "id": ["1"],
                 "groups": ["group_1", "group_2"],
                 "email": ["email"],
             }
             mock_init_saml_auth.return_value = auth
+
+            if raw_login_state is None:
+                mode, _, nonce = relay_state.partition(":")
+                entries = (
+                    pending
+                    if pending is not None
+                    else [(nonce or TEST_NONCE, EXPECTED_REQUEST_ID, mode or 
LOGIN_MODE_REDIRECT)]
+                )
+                cookie = make_login_state(*entries, expires_in=expires_in)
+            else:
+                cookie = raw_login_state
+
             client = TestClient(create_app(), base_url=base_url)
-            return client.post(
+            if cookie:
+                client.cookies.set(COOKIE_NAME_LOGIN_STATE, cookie)
+            response = client.post(
                 AUTH_MANAGER_FASTAPI_APP_PREFIX + "/login_callback",
                 follow_redirects=False,
                 data={"RelayState": relay_state},
             )
+            if return_auth_mock:
+                return response, auth
+            return response
 
 
-@mock_plugin_manager(plugins=[])
 class TestLoginRouter:
     @pytest.mark.parametrize(
         "url",
@@ -135,51 +227,210 @@ class TestLoginRouter:
         )
 
     def test_login_callback_successful_with_relay_state_redirect(self):
-        response = get_login_callback_response("login-redirect")
+        response = get_login_callback_response(RELAY_REDIRECT)
         assert response.status_code == 303
         assert "location" in response.headers
         assert "_token" in response.cookies
         assert 
response.headers["location"].startswith("http://localhost:8080/";)
 
     def test_login_callback_sets_secure_cookie_behind_tls_proxy(self):
-        response = get_login_callback_response("login-redirect", 
base_url="https://testserver";)
+        response = get_login_callback_response(RELAY_REDIRECT, 
base_url="https://testserver";)
 
         assert "Secure" in response.headers["set-cookie"]
 
     def test_login_callback_successful_with_relay_state_token(self):
-        response = get_login_callback_response("login-token")
+        response = get_login_callback_response(RELAY_TOKEN)
         assert response.status_code == 200
         assert "access_token" in response.json()
 
     def test_login_callback_with_invalid_relay_state(self):
         response = get_login_callback_response("dummy")
-        assert response.status_code == 500
+        assert response.status_code == 401
+
+    # ------------------------------------------------------------------
+    # Binding the SAML response to the browser that started the login
+    # ------------------------------------------------------------------
+
+    def test_login_sets_a_login_state_cookie(self, test_client):
+        """The AuthnRequest id must be remembered so the response can be tied 
to it."""
+        response = test_client.get(AUTH_MANAGER_FASTAPI_APP_PREFIX + "/login", 
follow_redirects=False)
+        assert COOKIE_NAME_LOGIN_STATE in response.cookies
+        set_cookie = response.headers["set-cookie"].lower()
+        assert "httponly" in set_cookie
+        assert "samesite=lax" in set_cookie
+
+    def test_login_sends_the_nonce_to_the_idp(self, test_client):
+        """The nonce has to survive the round trip, so it goes out in 
RelayState."""
+        response = test_client.get(AUTH_MANAGER_FASTAPI_APP_PREFIX + "/login", 
follow_redirects=False)
+        assert f"RelayState={LOGIN_MODE_REDIRECT}" in 
response.headers["location"]
+
+    def test_login_callback_rejects_a_browser_that_started_no_login(self):
+        """
+        The core case: an assertion posted to a browser that never started a 
login.
+
+        The assertion here is valid and authenticates successfully -- the mock 
returns an
+        authenticated user. What must stop it is the absence of any login this 
browser
+        began, so the caller is not logged in as the assertion's subject.
+        """
+        response = get_login_callback_response(RELAY_REDIRECT, 
raw_login_state="")
+        assert response.status_code == 401
+        assert "_token" not in response.cookies
+
+    def test_login_callback_enforces_in_response_to(self):
+        """python3-saml only validates InResponseTo when it is given the 
request id."""
+        response, auth = get_login_callback_response(RELAY_REDIRECT, 
return_auth_mock=True)
+        assert response.status_code == 303
+        
auth.process_response.assert_called_once_with(request_id=EXPECTED_REQUEST_ID)
+
+    def 
test_login_callback_rejects_a_relay_state_the_browser_did_not_ask_for(self):
+        """The return mode is fixed when the flow starts, not chosen by the 
response."""
+        response = get_login_callback_response(
+            RELAY_TOKEN, pending=[(TEST_NONCE, EXPECTED_REQUEST_ID, 
LOGIN_MODE_REDIRECT)]
+        )
+        assert response.status_code == 401
+
+    def test_login_callback_rejects_a_nonce_this_browser_never_had(self):
+        """The nonce names one of this browser's pending logins; an unknown 
one matches none."""
+        response = get_login_callback_response(
+            f"{LOGIN_MODE_REDIRECT}:some-other-nonce",
+            pending=[(TEST_NONCE, EXPECTED_REQUEST_ID, LOGIN_MODE_REDIRECT)],
+        )
+        assert response.status_code == 401
+
+    @pytest.mark.parametrize(
+        "mangle",
+        [
+            pytest.param(lambda state: state.split(".")[0], 
id="signature_stripped"),
+            pytest.param(lambda state: f"{state.split('.')[0]}.{'0' * 64}", 
id="signature_wrong"),
+            pytest.param(lambda state: f"x{state}", id="payload_altered"),
+            pytest.param(lambda state: "not-a-cookie", id="not_a_cookie"),
+        ],
+    )
+    def test_login_callback_rejects_a_login_state_it_did_not_sign(self, 
mangle):
+        """
+        Forging an entry would defeat the binding outright.
+
+        Anyone who could write this cookie could name an arbitrary 
AuthnRequest id, so an
+        unrelated response would match it. The signature is what makes the 
entry
+        something only this deployment can produce.
+        """
+        with conf_vars(BASE_CONF):
+            valid = make_login_state((TEST_NONCE, EXPECTED_REQUEST_ID, 
LOGIN_MODE_REDIRECT))
+        response = get_login_callback_response(RELAY_REDIRECT, 
raw_login_state=mangle(valid))
+        assert response.status_code == 401
+
+    def test_login_callback_rejects_an_expired_login_state(self):
+        """A login left unfinished stops being answerable once its deadline 
passes."""
+        response = get_login_callback_response(RELAY_REDIRECT, expires_in=-1)
+        assert response.status_code == 401
+
+    def test_login_callback_clears_the_login_state_on_the_token_path(self):
+        """
+        A consumed request id must not stay usable after a token login either.
+
+        Left set, the same assertion could be reposted within the cookie's 
lifetime to
+        mint further API tokens.
+        """
+        response = get_login_callback_response(RELAY_TOKEN)
+        assert response.status_code == 200
+        assert read_login_state_cookie(response) == '""'
+
+    def test_login_callback_clears_the_login_state_on_success(self):
+        """A consumed request id must not stay usable for a second response."""
+        response = get_login_callback_response(RELAY_REDIRECT)
+        assert response.status_code == 303
+        assert read_login_state_cookie(response) == '""'
+
+    def test_login_callback_rejects_the_same_response_twice(self):
+        """
+        Single use, end to end: the cookie the first response hands back no 
longer answers it.
+
+        This is the replay the clearing above exists to stop, exercised 
through the route
+        rather than by inspecting the header.
+        """
+        first = get_login_callback_response(RELAY_REDIRECT)
+        assert first.status_code == 303
+        replayed = get_login_callback_response(
+            RELAY_REDIRECT, 
raw_login_state=read_login_state_cookie(first).strip('"')
+        )
+        assert replayed.status_code == 401
+
+    def test_login_callback_leaves_logins_started_in_other_tabs_alone(self):
+        """
+        Two tabs, two pending logins. Finishing one must not strand the other.
+
+        A single-slot cookie made the second tab overwrite the first, so 
whichever login
+        finished second failed for a user who had done nothing wrong.
+        """
+        other = ("nonce-from-the-other-tab", "other_request_id", 
LOGIN_MODE_REDIRECT)
+        response = get_login_callback_response(
+            RELAY_REDIRECT,
+            pending=[other, (TEST_NONCE, EXPECTED_REQUEST_ID, 
LOGIN_MODE_REDIRECT)],
+        )
+        assert response.status_code == 303
+
+        remaining = read_login_state_cookie(response)
+        assert remaining is not None
+        assert remaining != '""'
+        entries = json.loads(base64.urlsafe_b64decode(remaining.split(".")[0]))
+        assert [entry["nonce"] for entry in entries] == [other[0]]
+
+    # ------------------------------------------------------------------
+    # IdP-initiated (unsolicited) logins, which are opt-in
+    # ------------------------------------------------------------------
+
+    def test_login_callback_rejects_an_unsolicited_response_by_default(self):
+        """The access portal flow is off unless a deployment asks for it."""
+        response = get_login_callback_response("", raw_login_state="")
+        assert response.status_code == 401
+
+    def 
test_login_callback_accepts_an_unsolicited_response_when_opted_in(self):
+        """With the option on, an assertion nobody asked for logs the caller 
in."""
+        response = get_login_callback_response(
+            "",
+            raw_login_state="",
+            extra_conf={("aws_auth_manager", "allow_idp_initiated_login"): 
"True"},
+        )
+        assert response.status_code == 303
+        assert "_token" in response.cookies
+
+    def 
test_login_callback_does_not_enforce_in_response_to_when_unsolicited(self):
+        """There is no request to bind to, so no request id is given to 
python3-saml."""
+        response, auth = get_login_callback_response(
+            "",
+            raw_login_state="",
+            extra_conf={("aws_auth_manager", "allow_idp_initiated_login"): 
"True"},
+            return_auth_mock=True,
+        )
+        assert response.status_code == 303
+        auth.process_response.assert_called_once_with(request_id=None)
+
+    def 
test_login_callback_refuses_an_unsolicited_response_carrying_in_response_to(self):
+        """
+        An unsolicited assertion answers no request, so one naming a request 
is a replay.
+
+        python3-saml skips the comparison entirely when it is handed no 
request id, so
+        opting in to the access portal flow would otherwise also accept a 
solicited
+        assertion reposted with the browser's cookie removed.
+        """
+        response = get_login_callback_response(
+            "",
+            raw_login_state="",
+            in_response_to=EXPECTED_REQUEST_ID,
+            extra_conf={("aws_auth_manager", "allow_idp_initiated_login"): 
"True"},
+        )
+        assert response.status_code == 401
+
+    def 
test_login_callback_still_binds_a_solicited_response_when_unsolicited_allowed(self):
+        """Opting in relaxes the no-state case only; a login that started here 
is still bound."""
+        response, auth = get_login_callback_response(
+            RELAY_REDIRECT,
+            extra_conf={("aws_auth_manager", "allow_idp_initiated_login"): 
"True"},
+            return_auth_mock=True,
+        )
+        assert response.status_code == 303
+        
auth.process_response.assert_called_once_with(request_id=EXPECTED_REQUEST_ID)
 
     def test_login_callback_unsuccessful(self):
-        with conf_vars(
-            {
-                (
-                    "core",
-                    "auth_manager",
-                ): 
"airflow.providers.amazon.aws.auth_manager.aws_auth_manager.AwsAuthManager",
-                ("aws_auth_manager", "saml_metadata_url"): SAML_METADATA_URL,
-            }
-        ):
-            with (
-                patch.object(OneLogin_Saml2_IdPMetadataParser, "parse_remote") 
as mock_parse_remote,
-                patch(
-                    
"airflow.providers.amazon.aws.auth_manager.routes.login._init_saml_auth"
-                ) as mock_init_saml_auth,
-                patch(
-                    
"airflow.providers.amazon.aws.auth_manager.avp.facade.AwsAuthManagerAmazonVerifiedPermissionsFacade.is_policy_store_schema_up_to_date"
-                ) as mock_is_policy_store_schema_up_to_date,
-            ):
-                mock_is_policy_store_schema_up_to_date.return_value = True
-                mock_parse_remote.return_value = SAML_METADATA_PARSED
-
-                auth = Mock()
-                auth.is_authenticated.return_value = False
-                mock_init_saml_auth.return_value = auth
-                client = TestClient(create_app())
-                response = client.post(AUTH_MANAGER_FASTAPI_APP_PREFIX + 
"/login_callback")
-                assert response.status_code == 500
+        response = get_login_callback_response(RELAY_REDIRECT, 
is_authenticated=False)
+        assert response.status_code == 500
diff --git a/scripts/ci/prek/check_contextmanager_class_decorators.py 
b/scripts/ci/prek/check_contextmanager_class_decorators.py
index 149d00bc6c2..e1921f373b5 100755
--- a/scripts/ci/prek/check_contextmanager_class_decorators.py
+++ b/scripts/ci/prek/check_contextmanager_class_decorators.py
@@ -80,6 +80,7 @@ class ContextManagerClassDecoratorChecker(ast.NodeVisitor):
         problematic_decorators = {
             "conf_vars",
             "env_vars",
+            "mock_plugin_manager",
             "contextlib.contextmanager",
             "contextmanager",
         }
diff --git 
a/scripts/tests/ci/prek/test_check_contextmanager_class_decorators.py 
b/scripts/tests/ci/prek/test_check_contextmanager_class_decorators.py
new file mode 100644
index 00000000000..7250ee9ad6f
--- /dev/null
+++ b/scripts/tests/ci/prek/test_check_contextmanager_class_decorators.py
@@ -0,0 +1,80 @@
+# 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
+
+import pytest
+from check_contextmanager_class_decorators import check_file
+
+
+class TestCheckFile:
+    @pytest.mark.parametrize(
+        "decorator",
+        [
+            pytest.param("@conf_vars({})", id="conf_vars"),
+            pytest.param("@env_vars({})", id="env_vars"),
+            pytest.param("@mock_plugin_manager(plugins=[])", 
id="mock_plugin_manager"),
+            pytest.param("@contextmanager", id="contextmanager"),
+            pytest.param("@contextlib.contextmanager", 
id="contextlib_contextmanager"),
+        ],
+    )
+    def test_context_manager_on_a_test_class_is_reported(self, 
write_python_file, decorator):
+        """Each of these turns the class into a function, so pytest silently 
collects nothing."""
+        path = write_python_file(
+            f"""
+            {decorator}
+            class TestSomething:
+                def test_one(self):
+                    pass
+            """
+        )
+        errors = check_file(path)
+        assert len(errors) == 1
+        assert "TestSomething" in errors[0]
+
+    @pytest.mark.parametrize(
+        "code",
+        [
+            pytest.param(
+                """
+                @pytest.mark.usefixtures("no_plugins")
+                class TestSomething:
+                    def test_one(self):
+                        pass
+                """,
+                id="usefixtures_is_the_supported_form",
+            ),
+            pytest.param(
+                """
+                class TestSomething:
+                    @mock_plugin_manager(plugins=[])
+                    def test_one(self):
+                        pass
+                """,
+                id="on_a_method_is_fine",
+            ),
+            pytest.param(
+                """
+                @mock_plugin_manager(plugins=[])
+                class HelperNotATestClass:
+                    pass
+                """,
+                id="only_test_classes_are_checked",
+            ),
+        ],
+    )
+    def test_accepted_usages(self, write_python_file, code):
+        assert check_file(write_python_file(code)) == []

Reply via email to