Skip to content
Closed
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 @@ -58,7 +58,7 @@
from collections.abc import Sequence

from fastapi import FastAPI
from sqlalchemy import Row
from sqlalchemy import Row, Select
from sqlalchemy.orm import Session
from starlette.middleware import _MiddlewareFactory

Expand Down Expand Up @@ -625,6 +625,33 @@ def _is_authorized_connection(conn_id: str):

return {conn_id for conn_id in conn_ids if _is_authorized_connection(conn_id)}

def get_authorized_dag_ids_select(
self,
*,
user: T,
method: ResourceMethod = "GET",
) -> Select | None:
"""
Get a select of the Dag ids the user has access to, to be used as a subquery.

Returning ``None``, the default, makes the API materialize the whole authorized set
through :meth:`get_authorized_dag_ids` instead. Auth managers that keep their grants in
the Airflow metadata database can return a select here, which the API applies as
``dag_id IN (subquery)`` so that filtering and pagination happen in one statement rather
than after every authorized Dag id has been loaded into memory.

The select must produce a single column of Dag ids.

Returning a select replaces :meth:`get_authorized_dag_ids` entirely, including its
per-team grouping, so a multi-team deployment has to scope the select itself. Join
``DagBundleModel`` and ``dag_bundle_team_association_table`` the way
:meth:`get_authorized_dag_ids` does, or return ``None`` and keep the default fan-out.

:param user: the user
:param method: the method to filter on
"""
return None

@provide_session
def get_authorized_dag_ids(
self,
Expand Down
15 changes: 9 additions & 6 deletions airflow-core/src/airflow/api_fastapi/common/db/dags.py
Original file line number Diff line number Diff line change
Expand Up @@ -37,28 +37,31 @@


def generate_dag_with_latest_run_query(
max_run_filters: list[BaseParam], order_by: SortParam, *, dag_ids: set[str] | None = None
max_run_filters: list[BaseParam],
order_by: SortParam,
*,
dag_ids: set[str] | Select | None = None,
) -> Select:
"""
Generate a query to fetch Dags with their latest run.

:param max_run_filters: List of filters to apply to the latest run
:param order_by: Sort parameter for ordering results
:param dag_ids: Optional set of Dag IDs to limit the query to. When provided, both the main
Dag query and the subquery for finding the latest runs will be filtered to
only these Dag IDs, improving performance when users have limited Dag access.
:param dag_ids: Optional set of Dag IDs, or a select producing them, to limit the query
to. When provided, both the main Dag query and the subquery for finding the latest
runs are filtered to those Dag IDs.
:return: SQLAlchemy Select statement
"""
query = select(DagModel).options(selectinload(DagModel.tags))

# Filter main query by dag_ids if provided
if dag_ids is not None:
query = query.where(DagModel.dag_id.in_(dag_ids or set()))
query = query.where(DagModel.dag_id.in_(dag_ids))

# Also filter the subquery for finding latest runs
max_run_id_query_stmt = select(DagRun.dag_id, func.max(DagRun.id).label("max_dag_run_id"))
if dag_ids is not None:
max_run_id_query_stmt = max_run_id_query_stmt.where(DagRun.dag_id.in_(dag_ids or set()))
max_run_id_query_stmt = max_run_id_query_stmt.where(DagRun.dag_id.in_(dag_ids))
max_run_id_query = max_run_id_query_stmt.group_by(DagRun.dag_id).subquery(name="mrq")

has_max_run_filter = False
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -146,7 +146,7 @@ def get_dags(
last_dag_run_state,
],
order_by=order_by,
dag_ids=readable_dags_filter.value,
dag_ids=readable_dags_filter.permitted,
)

dags_select, total_entries = paginated_select(
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -152,7 +152,7 @@ def get_dags(
last_dag_run_state,
],
order_by=order_by,
dag_ids=readable_dags_filter.value,
dag_ids=readable_dags_filter.permitted,
)

dags_select, total_entries = paginated_select(
Expand Down
70 changes: 56 additions & 14 deletions airflow-core/src/airflow/api_fastapi/core_api/security.py
Original file line number Diff line number Diff line change
Expand Up @@ -30,6 +30,7 @@
from pydantic import NonNegativeInt, TypeAdapter, ValidationError
from sqlalchemy import or_, select
from sqlalchemy.orm import Session
from sqlalchemy.sql import Select

from airflow.api_fastapi.app import get_auth_manager
from airflow.api_fastapi.auth.managers.base_auth_manager import (
Expand Down Expand Up @@ -80,8 +81,6 @@
from airflow.models.xcom import XComModel

if TYPE_CHECKING:
from sqlalchemy.sql import Select

from airflow.api_fastapi.auth.managers.base_auth_manager import ResourceMethod


Expand Down Expand Up @@ -246,23 +245,61 @@ def inner(
class PermittedDagFilter(OrmClause[set[str]]):
"""A parameter that filters the permitted dags for the user."""

def __init__(
self,
value: set[str] | None = None,
*,
subquery: Select | None = None,
materialize: Callable[[], set[str]] | None = None,
):
self._materialized = value
self._subquery = subquery
self._materialize = materialize

@property
def value(self) -> set[str] | None:
"""
The permitted dag ids.

Reading this materializes the whole authorized set when the auth manager answered
with a subquery, since callers that ask for the ids need them in memory.
"""
if self._materialized is None and self._materialize is not None:
self._materialized = self._materialize()
return self._materialized

@value.setter
def value(self, value: set[str] | None) -> None:
self._materialized = value

@property
def permitted(self) -> set[str] | Select:
"""
The operand for ``in_()``, preferring the subquery so nothing is materialized.

Not ``self.value or set()``: a select has no truth value, and an empty set is a
legitimate answer meaning nothing is permitted.
"""
if self._subquery is not None:
return self._subquery
return self.value if self.value is not None else set()

def to_orm(self, select: Select) -> Select:
# self.value may be None (OrmClause holds Optional), ensure we pass an Iterable to in_
return select.where(DagModel.dag_id.in_(self.value or set()))
return select.where(DagModel.dag_id.in_(self.permitted))


class PermittedDagRunFilter(PermittedDagFilter):
"""A parameter that filters the permitted dag runs for the user."""

def to_orm(self, select: Select) -> Select:
return select.where(DagRun.dag_id.in_(self.value or set()))
return select.where(DagRun.dag_id.in_(self.permitted))


class PermittedDagWarningFilter(PermittedDagFilter):
"""A parameter that filters the permitted dag warnings for the user."""

def to_orm(self, select: Select) -> Select:
return select.where(DagWarning.dag_id.in_(self.value or set()))
return select.where(DagWarning.dag_id.in_(self.permitted))


class PermittedEventLogFilter(PermittedDagFilter):
Expand All @@ -273,7 +310,7 @@ def __init__(self, value: set[str] | None = None, *, include_non_dag_logs: bool
self.include_non_dag_logs = include_non_dag_logs

def to_orm(self, select: Select) -> Select:
permitted_dag_logs = Log.dag_id.in_(self.value or set())
permitted_dag_logs = Log.dag_id.in_(self.permitted)
if not self.include_non_dag_logs:
return select.where(permitted_dag_logs)
# Event logs not related to a Dag have dag_id as None. They record Connection,
Expand All @@ -288,35 +325,35 @@ class PermittedTIFilter(PermittedDagFilter):
"""A parameter that filters the permitted task instances for the user."""

def to_orm(self, select: Select) -> Select:
return select.where(TI.dag_id.in_(self.value or set()))
return select.where(TI.dag_id.in_(self.permitted))


class PermittedXComFilter(PermittedDagFilter):
"""A parameter that filters the permitted XComs for the user."""

def to_orm(self, select: Select) -> Select:
return select.where(XComModel.dag_id.in_(self.value or set()))
return select.where(XComModel.dag_id.in_(self.permitted))


class PermittedTagFilter(PermittedDagFilter):
"""A parameter that filters the permitted dag tags for the user."""

def to_orm(self, select: Select) -> Select:
return select.where(DagTag.dag_id.in_(self.value or set()))
return select.where(DagTag.dag_id.in_(self.permitted))


class PermittedDagVersionFilter(PermittedDagFilter):
"""A parameter that filters the permitted dag versions for the user."""

def to_orm(self, select: Select) -> Select:
return select.where(DagVersion.dag_id.in_(self.value or set()))
return select.where(DagVersion.dag_id.in_(self.permitted))


class PermittedBackfillFilter(PermittedDagFilter):
"""A parameter that filters the permitted backfills for the user."""

def to_orm(self, select: Select) -> Select:
return select.where(Backfill.dag_id.in_(self.value or set()))
return select.where(Backfill.dag_id.in_(self.permitted))


def permitted_dag_filter_factory(
Expand All @@ -333,8 +370,13 @@ def depends_permitted_dags_filter(
user: GetUserDep,
auth_manager: AuthManagerDep,
) -> PermittedDagFilter:
authorized_dags: set[str] = auth_manager.get_authorized_dag_ids(user=user, method=method)
return filter_class(authorized_dags)
subquery = auth_manager.get_authorized_dag_ids_select(user=user, method=method)
if subquery is not None:
return filter_class(
subquery=subquery,
materialize=lambda: auth_manager.get_authorized_dag_ids(user=user, method=method),
)
return filter_class(auth_manager.get_authorized_dag_ids(user=user, method=method))

return depends_permitted_dags_filter

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -594,6 +594,10 @@ def side_effect_func(
result = auth_manager.get_authorized_dag_ids(user=user, session=session)
assert result == expected

def test_get_authorized_dag_ids_select_defaults_to_none(self, auth_manager):
"""Returning None is what keeps every existing manager on the materializing path."""
assert auth_manager.get_authorized_dag_ids_select(user=Mock()) is None

@pytest.mark.parametrize(
("access_per_connection", "access_per_team", "rows", "expected"),
[
Expand Down
94 changes: 91 additions & 3 deletions airflow-core/tests/unit/api_fastapi/core_api/test_security.py
Original file line number Diff line number Diff line change
Expand Up @@ -22,6 +22,7 @@
import pytest
from fastapi import HTTPException, Request
from jwt import ExpiredSignatureError, InvalidTokenError
from sqlalchemy import false, select
from sqlalchemy.orm import Session

from airflow import settings
Expand All @@ -42,6 +43,8 @@
from airflow.api_fastapi.core_api.datamodels.pools import PoolBody
from airflow.api_fastapi.core_api.datamodels.variables import VariableBody
from airflow.api_fastapi.core_api.security import (
PermittedDagFilter,
PermittedDagRunFilter,
_build_dag_run_access_requests,
get_user,
is_safe_url,
Expand All @@ -56,10 +59,11 @@
requires_access_variable_bulk,
resolve_user_from_token,
)
from airflow.models import Connection, Pool, Variable
from airflow.models import Connection, DagModel, Pool, Variable
from airflow.models.backfill import Backfill
from airflow.models.dag import DagModel
from airflow.models.dagbundle import DagBundleModel
from airflow.models.dag import DagTag
from airflow.models.dagbundle import DagBundleModel, dag_bundle_team_association_table
from airflow.models.dagrun import DagRun
from airflow.models.team import Team

from tests_common.test_utils.asserts import assert_queries_count
Expand Down Expand Up @@ -1505,3 +1509,87 @@ def test_auth_manager_from_app_integration_with_test_client(self, test_client):
assert auth_manager is not None
assert hasattr(auth_manager, "get_url_login")
assert hasattr(auth_manager, "get_url_logout")


class TestPermittedDagFilterSubquery:
"""A manager keeping grants in the metadata database can hand back a select instead of a set.

Without this the whole authorized set is loaded into memory before pagination is applied,
which is what makes a list view cost O(all dags) even when the grants are already in SQL.
"""

@staticmethod
def _compiled(orm_clause) -> str:
stmt = orm_clause.to_orm(select(DagModel.dag_id))
return str(stmt.compile(compile_kwargs={"literal_binds": True}))

def test_a_set_still_becomes_an_in_list(self):
sql = self._compiled(PermittedDagFilter({"dag_a", "dag_b"}))
assert "IN (" in sql
assert "SELECT" in sql.split("IN (", 1)[1] or "dag_a" in sql

def test_a_select_becomes_a_subquery_rather_than_a_materialized_set(self):
sql = self._compiled(PermittedDagFilter(subquery=select(DagModel.dag_id).where(DagModel.is_paused)))
after_in = sql.split("IN (", 1)[1]
assert after_in.lstrip().upper().startswith("SELECT"), sql

def test_none_permits_nothing(self):
"""None has to stay deny-all; an empty select is a legitimate value, so `or` would be wrong."""
sql = self._compiled(PermittedDagFilter(None))
assert "IN (" in sql

def test_an_empty_select_is_not_treated_as_no_value(self):
empty = select(DagModel.dag_id).where(false())
sql = self._compiled(PermittedDagFilter(subquery=empty))
after_in = sql.split("IN (", 1)[1]
assert after_in.lstrip().upper().startswith("SELECT"), sql

def test_the_subquery_reaches_the_dag_run_filter_too(self):
"""Every permitted-* filter inherits the behaviour, so they all get it at once."""
stmt = PermittedDagRunFilter(subquery=select(DagModel.dag_id)).to_orm(select(DagRun.dag_id))
sql = str(stmt.compile(compile_kwargs={"literal_binds": True}))
after_in = sql.split("IN (", 1)[1]
assert after_in.lstrip().upper().startswith("SELECT"), sql

def test_a_tag_based_manager_needs_no_fab(self):
"""The docs point custom managers at Dag attributes like tags, which live in the database.

Without a select such a manager has to load every dag id carrying the tag into memory
before the API can page over them.
"""
by_tag = select(DagTag.dag_id).where(DagTag.name.in_({"team-a", "team-b"}))
sql = self._compiled(PermittedDagFilter(subquery=by_tag))
after_in = sql.split("IN (", 1)[1]
assert after_in.lstrip().upper().startswith("SELECT"), sql
assert "dag_tag" in sql

def test_a_team_scoped_select_stays_one_statement(self):
"""Returning a select replaces the per-team fan-out, so the select carries the scoping."""
by_team = (
select(DagModel.dag_id)
.join(DagBundleModel, DagModel.bundle_name == DagBundleModel.name)
.join(
dag_bundle_team_association_table,
DagBundleModel.name == dag_bundle_team_association_table.c.dag_bundle_name,
)
.where(dag_bundle_team_association_table.c.team_name == "team-a")
)
sql = self._compiled(PermittedDagFilter(subquery=by_team))
after_in = sql.split("IN (", 1)[1]
assert after_in.lstrip().upper().startswith("SELECT"), sql
assert "team_name" in sql

def test_reading_value_materializes_only_when_asked(self):
"""Callers that need the ids still get them, and nothing is loaded until one asks."""
calls = []

def materialize() -> set[str]:
calls.append(1)
return {"dag_a"}

f = PermittedDagFilter(subquery=select(DagModel.dag_id), materialize=materialize)
self._compiled(f)
assert calls == []
assert f.value == {"dag_a"}
assert f.value == {"dag_a"}
assert calls == [1]
Loading