Skip to content

Commit 6c06640

Browse files
committed
Let an auth manager push dag authorization into SQL
get_authorized_dag_ids returns a set, so every authorized dag id is loaded into memory before pagination is applied. FabAuthManager keeps its grants in the metadata database and still does this: a user authorized on all dags gets select(DagModel.dag_id) materialized on every list request. get_authorized_dag_ids_select lets a manager return a select instead, which the permitted-dag filters apply as a subquery. Returning None, the default, keeps the existing behaviour, and every permitted-* filter inherits it because they all build the clause with in_(). Signed-off-by: 1fanwang <1fannnw@gmail.com>
1 parent f9b72d8 commit 6c06640

8 files changed

Lines changed: 216 additions & 27 deletions

File tree

airflow-core/src/airflow/api_fastapi/auth/managers/base_auth_manager.py

Lines changed: 23 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -58,7 +58,7 @@
5858
from collections.abc import Sequence
5959

6060
from fastapi import FastAPI
61-
from sqlalchemy import Row
61+
from sqlalchemy import Row, Select
6262
from sqlalchemy.orm import Session
6363
from starlette.middleware import _MiddlewareFactory
6464

@@ -625,6 +625,28 @@ def _is_authorized_connection(conn_id: str):
625625

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

628+
def get_authorized_dag_ids_select(
629+
self,
630+
*,
631+
user: T,
632+
method: ResourceMethod = "GET",
633+
) -> Select | None:
634+
"""
635+
Get a select of the Dag ids the user has access to, to be used as a subquery.
636+
637+
Returning ``None``, the default, makes the API materialize the whole authorized set
638+
through :meth:`get_authorized_dag_ids` instead. Auth managers that keep their grants in
639+
the Airflow metadata database can return a select here, which the API applies as
640+
``dag_id IN (subquery)`` so that filtering and pagination happen in one statement rather
641+
than after every authorized Dag id has been loaded into memory.
642+
643+
The select must produce a single column of Dag ids.
644+
645+
:param user: the user
646+
:param method: the method to filter on
647+
"""
648+
return None
649+
628650
@provide_session
629651
def get_authorized_dag_ids(
630652
self,

airflow-core/src/airflow/api_fastapi/common/db/dags.py

Lines changed: 9 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -37,28 +37,31 @@
3737

3838

3939
def generate_dag_with_latest_run_query(
40-
max_run_filters: list[BaseParam], order_by: SortParam, *, dag_ids: set[str] | None = None
40+
max_run_filters: list[BaseParam],
41+
order_by: SortParam,
42+
*,
43+
dag_ids: set[str] | Select | None = None,
4144
) -> Select:
4245
"""
4346
Generate a query to fetch Dags with their latest run.
4447
4548
:param max_run_filters: List of filters to apply to the latest run
4649
:param order_by: Sort parameter for ordering results
47-
:param dag_ids: Optional set of Dag IDs to limit the query to. When provided, both the main
48-
Dag query and the subquery for finding the latest runs will be filtered to
49-
only these Dag IDs, improving performance when users have limited Dag access.
50+
:param dag_ids: Optional set of Dag IDs, or a select producing them, to limit the query
51+
to. When provided, both the main Dag query and the subquery for finding the latest
52+
runs are filtered to those Dag IDs.
5053
:return: SQLAlchemy Select statement
5154
"""
5255
query = select(DagModel).options(selectinload(DagModel.tags))
5356

5457
# Filter main query by dag_ids if provided
5558
if dag_ids is not None:
56-
query = query.where(DagModel.dag_id.in_(dag_ids or set()))
59+
query = query.where(DagModel.dag_id.in_(dag_ids))
5760

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

6467
has_max_run_filter = False

airflow-core/src/airflow/api_fastapi/core_api/routes/public/dags.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -146,7 +146,7 @@ def get_dags(
146146
last_dag_run_state,
147147
],
148148
order_by=order_by,
149-
dag_ids=readable_dags_filter.value,
149+
dag_ids=readable_dags_filter.permitted,
150150
)
151151

152152
dags_select, total_entries = paginated_select(

airflow-core/src/airflow/api_fastapi/core_api/routes/ui/dags.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -152,7 +152,7 @@ def get_dags(
152152
last_dag_run_state,
153153
],
154154
order_by=order_by,
155-
dag_ids=readable_dags_filter.value,
155+
dag_ids=readable_dags_filter.permitted,
156156
)
157157

158158
dags_select, total_entries = paginated_select(

airflow-core/src/airflow/api_fastapi/core_api/security.py

Lines changed: 56 additions & 14 deletions
Original file line numberDiff line numberDiff line change
@@ -30,6 +30,7 @@
3030
from pydantic import NonNegativeInt, TypeAdapter, ValidationError
3131
from sqlalchemy import or_, select
3232
from sqlalchemy.orm import Session
33+
from sqlalchemy.sql import Select
3334

3435
from airflow.api_fastapi.app import get_auth_manager
3536
from airflow.api_fastapi.auth.managers.base_auth_manager import (
@@ -80,8 +81,6 @@
8081
from airflow.models.xcom import XComModel
8182

8283
if TYPE_CHECKING:
83-
from sqlalchemy.sql import Select
84-
8584
from airflow.api_fastapi.auth.managers.base_auth_manager import ResourceMethod
8685

8786

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

248+
def __init__(
249+
self,
250+
value: set[str] | None = None,
251+
*,
252+
subquery: Select | None = None,
253+
materialize: Callable[[], set[str]] | None = None,
254+
):
255+
self._materialized = value
256+
self._subquery = subquery
257+
self._materialize = materialize
258+
259+
@property
260+
def value(self) -> set[str] | None:
261+
"""
262+
The permitted dag ids.
263+
264+
Reading this materializes the whole authorized set when the auth manager answered
265+
with a subquery, since callers that ask for the ids need them in memory.
266+
"""
267+
if self._materialized is None and self._materialize is not None:
268+
self._materialized = self._materialize()
269+
return self._materialized
270+
271+
@value.setter
272+
def value(self, value: set[str] | None) -> None:
273+
self._materialized = value
274+
275+
@property
276+
def permitted(self) -> set[str] | Select:
277+
"""
278+
The operand for ``in_()``, preferring the subquery so nothing is materialized.
279+
280+
Not ``self.value or set()``: a select has no truth value, and an empty set is a
281+
legitimate answer meaning nothing is permitted.
282+
"""
283+
if self._subquery is not None:
284+
return self._subquery
285+
return self.value if self.value is not None else set()
286+
249287
def to_orm(self, select: Select) -> Select:
250-
# self.value may be None (OrmClause holds Optional), ensure we pass an Iterable to in_
251-
return select.where(DagModel.dag_id.in_(self.value or set()))
288+
return select.where(DagModel.dag_id.in_(self.permitted))
252289

253290

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

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

260297

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

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

267304

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

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

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

293330

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

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

300337

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

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

307344

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

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

314351

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

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

321358

322359
def permitted_dag_filter_factory(
@@ -333,8 +370,13 @@ def depends_permitted_dags_filter(
333370
user: GetUserDep,
334371
auth_manager: AuthManagerDep,
335372
) -> PermittedDagFilter:
336-
authorized_dags: set[str] = auth_manager.get_authorized_dag_ids(user=user, method=method)
337-
return filter_class(authorized_dags)
373+
subquery = auth_manager.get_authorized_dag_ids_select(user=user, method=method)
374+
if subquery is not None:
375+
return filter_class(
376+
subquery=subquery,
377+
materialize=lambda: auth_manager.get_authorized_dag_ids(user=user, method=method),
378+
)
379+
return filter_class(auth_manager.get_authorized_dag_ids(user=user, method=method))
338380

339381
return depends_permitted_dags_filter
340382

airflow-core/tests/unit/api_fastapi/auth/managers/test_base_auth_manager.py

Lines changed: 4 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -594,6 +594,10 @@ def side_effect_func(
594594
result = auth_manager.get_authorized_dag_ids(user=user, session=session)
595595
assert result == expected
596596

597+
def test_get_authorized_dag_ids_select_defaults_to_none(self, auth_manager):
598+
"""Returning None is what keeps every existing manager on the materializing path."""
599+
assert auth_manager.get_authorized_dag_ids_select(user=Mock()) is None
600+
597601
@pytest.mark.parametrize(
598602
("access_per_connection", "access_per_team", "rows", "expected"),
599603
[

airflow-core/tests/unit/api_fastapi/core_api/test_security.py

Lines changed: 74 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -22,6 +22,7 @@
2222
import pytest
2323
from fastapi import HTTPException, Request
2424
from jwt import ExpiredSignatureError, InvalidTokenError
25+
from sqlalchemy import false, select
2526
from sqlalchemy.orm import Session
2627

2728
from airflow import settings
@@ -42,6 +43,8 @@
4243
from airflow.api_fastapi.core_api.datamodels.pools import PoolBody
4344
from airflow.api_fastapi.core_api.datamodels.variables import VariableBody
4445
from airflow.api_fastapi.core_api.security import (
46+
PermittedDagFilter,
47+
PermittedDagRunFilter,
4548
_build_dag_run_access_requests,
4649
get_user,
4750
is_safe_url,
@@ -56,10 +59,11 @@
5659
requires_access_variable_bulk,
5760
resolve_user_from_token,
5861
)
59-
from airflow.models import Connection, Pool, Variable
62+
from airflow.models import Connection, DagModel, Pool, Variable
6063
from airflow.models.backfill import Backfill
61-
from airflow.models.dag import DagModel
64+
from airflow.models.dag import DagTag
6265
from airflow.models.dagbundle import DagBundleModel
66+
from airflow.models.dagrun import DagRun
6367
from airflow.models.team import Team
6468

6569
from tests_common.test_utils.asserts import assert_queries_count
@@ -1505,3 +1509,71 @@ def test_auth_manager_from_app_integration_with_test_client(self, test_client):
15051509
assert auth_manager is not None
15061510
assert hasattr(auth_manager, "get_url_login")
15071511
assert hasattr(auth_manager, "get_url_logout")
1512+
1513+
1514+
class TestPermittedDagFilterSubquery:
1515+
"""A manager keeping grants in the metadata database can hand back a select instead of a set.
1516+
1517+
Without this the whole authorized set is loaded into memory before pagination is applied,
1518+
which is what makes a list view cost O(all dags) even when the grants are already in SQL.
1519+
"""
1520+
1521+
@staticmethod
1522+
def _compiled(orm_clause) -> str:
1523+
stmt = orm_clause.to_orm(select(DagModel.dag_id))
1524+
return str(stmt.compile(compile_kwargs={"literal_binds": True}))
1525+
1526+
def test_a_set_still_becomes_an_in_list(self):
1527+
sql = self._compiled(PermittedDagFilter({"dag_a", "dag_b"}))
1528+
assert "IN (" in sql
1529+
assert "SELECT" in sql.split("IN (", 1)[1] or "dag_a" in sql
1530+
1531+
def test_a_select_becomes_a_subquery_rather_than_a_materialized_set(self):
1532+
sql = self._compiled(PermittedDagFilter(subquery=select(DagModel.dag_id).where(DagModel.is_paused)))
1533+
after_in = sql.split("IN (", 1)[1]
1534+
assert after_in.lstrip().upper().startswith("SELECT"), sql
1535+
1536+
def test_none_permits_nothing(self):
1537+
"""None has to stay deny-all; an empty select is a legitimate value, so `or` would be wrong."""
1538+
sql = self._compiled(PermittedDagFilter(None))
1539+
assert "IN (" in sql
1540+
1541+
def test_an_empty_select_is_not_treated_as_no_value(self):
1542+
empty = select(DagModel.dag_id).where(false())
1543+
sql = self._compiled(PermittedDagFilter(subquery=empty))
1544+
after_in = sql.split("IN (", 1)[1]
1545+
assert after_in.lstrip().upper().startswith("SELECT"), sql
1546+
1547+
def test_the_subquery_reaches_the_dag_run_filter_too(self):
1548+
"""Every permitted-* filter inherits the behaviour, so they all get it at once."""
1549+
stmt = PermittedDagRunFilter(subquery=select(DagModel.dag_id)).to_orm(select(DagRun.dag_id))
1550+
sql = str(stmt.compile(compile_kwargs={"literal_binds": True}))
1551+
after_in = sql.split("IN (", 1)[1]
1552+
assert after_in.lstrip().upper().startswith("SELECT"), sql
1553+
1554+
def test_a_tag_based_manager_needs_no_fab(self):
1555+
"""The docs point custom managers at Dag attributes like tags, which live in the database.
1556+
1557+
Without a select such a manager has to load every dag id carrying the tag into memory
1558+
before the API can page over them.
1559+
"""
1560+
by_tag = select(DagTag.dag_id).where(DagTag.name.in_({"team-a", "team-b"}))
1561+
sql = self._compiled(PermittedDagFilter(subquery=by_tag))
1562+
after_in = sql.split("IN (", 1)[1]
1563+
assert after_in.lstrip().upper().startswith("SELECT"), sql
1564+
assert "dag_tag" in sql
1565+
1566+
def test_reading_value_materializes_only_when_asked(self):
1567+
"""Callers that need the ids still get them, and nothing is loaded until one asks."""
1568+
calls = []
1569+
1570+
def materialize() -> set[str]:
1571+
calls.append(1)
1572+
return {"dag_a"}
1573+
1574+
f = PermittedDagFilter(subquery=select(DagModel.dag_id), materialize=materialize)
1575+
self._compiled(f)
1576+
assert calls == []
1577+
assert f.value == {"dag_a"}
1578+
assert f.value == {"dag_a"}
1579+
assert calls == [1]

0 commit comments

Comments
 (0)