Skip to content
Merged
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
1 change: 1 addition & 0 deletions airflow-core/docs/core-concepts/auth-manager/index.rst
Original file line number Diff line number Diff line change
Expand Up @@ -235,6 +235,7 @@ The following methods aren't required to override to have a functional Airflow a
* ``batch_is_authorized_dag``: Batch version of ``is_authorized_dag``. If not overridden, it calls ``is_authorized_dag`` for every single item.
* ``batch_is_authorized_pool``: Batch version of ``is_authorized_pool``. If not overridden, it calls ``is_authorized_pool`` for every single item.
* ``batch_is_authorized_variable``: Batch version of ``is_authorized_variable``. If not overridden, it calls ``is_authorized_variable`` for every single item.
* ``filter_authorized_assets``: Given a list of assets (each carrying its id, name and uri), return the ids of the assets the user has access to. If not overridden, it calls ``is_authorized_asset`` for every single asset passed as parameter.
* ``filter_authorized_connections``: Given a list of connection IDs (``conn_id``), return the list of connection IDs the user has access to. If not overridden, it calls ``is_authorized_connection`` for every single connection passed as parameter.
* ``filter_authorized_dag_ids``: Given a list of Dag IDs, return the list of Dag IDs the user has access to. If not overridden, it calls ``is_authorized_dag`` for every single Dag passes as parameter.
* ``filter_authorized_pools``: Given a list of pool names, return the list of pool names the user has access to. If not overridden, it calls ``is_authorized_pool`` for every single pool passed as parameter.
Expand Down
1 change: 1 addition & 0 deletions airflow-core/newsfragments/72682.feature.rst
Original file line number Diff line number Diff line change
@@ -0,0 +1 @@
Auth managers can now restrict which assets a user may see: ``BaseAuthManager`` gains ``get_authorized_assets`` and ``filter_authorized_assets``, ``AssetDetails`` now carries the asset's ``name`` and ``uri`` so an implementation can authorize on a URI prefix rather than an opaque id, and the asset list and asset events endpoints scope both their rows and their ``total_entries`` to the assets the caller may read. The default implementation calls ``is_authorized_asset`` once per asset, so a deployment whose auth manager ignores asset details is unaffected.
Original file line number Diff line number Diff line change
Expand Up @@ -32,6 +32,7 @@

from airflow.api_fastapi.auth.managers.models.base_user import BaseUser
from airflow.api_fastapi.auth.managers.models.resource_details import (
AssetDetails,
ConnectionDetails,
DagDetails,
PoolDetails,
Expand All @@ -48,6 +49,7 @@
from airflow.configuration import conf
from airflow.exceptions import RemovedInAirflow4Warning
from airflow.models import Connection, DagModel, Pool, Variable
from airflow.models.asset import AssetModel
from airflow.models.dagbundle import DagBundleModel
from airflow.models.revoked_token import RevokedToken
from airflow.models.team import Team, dag_bundle_team_association_table
Expand All @@ -72,7 +74,6 @@
from airflow.api_fastapi.auth.managers.models.resource_details import (
AccessView,
AssetAliasDetails,
AssetDetails,
ConfigurationDetails,
DagAccessEntity,
)
Expand Down Expand Up @@ -582,6 +583,52 @@ def batch_is_authorized_variable(
for request in requests
)

@provide_session
def get_authorized_assets(
self,
*,
user: T,
method: ResourceMethod = "GET",
session: Session = NEW_SESSION,
) -> set[int]:
"""
Get the ids of the assets the user has access to.

:param user: the user
:param method: the method to filter on
:param session: the session
"""
rows = session.execute(select(AssetModel.id, AssetModel.name, AssetModel.uri)).all()
assets = [AssetDetails(id=str(asset_id), name=name, uri=uri) for asset_id, name, uri in rows]
authorized_ids = self.filter_authorized_assets(assets=assets, user=user, method=method)
return {asset_id for asset_id, _, _ in rows if str(asset_id) in authorized_ids}

def filter_authorized_assets(
self,
*,
assets: Sequence[AssetDetails],
user: T,
method: ResourceMethod = "GET",
) -> set[str]:
"""
Filter assets the user has access to, returning the ids of the authorized ones.

By default, check individually if the user has permissions to access the asset. An auth manager
whose ``is_authorized_asset`` performs a remote call must override this method: a deployment can
hold far more assets than connections or pools, and the default costs one round trip per asset on
every asset listing.

:param assets: the assets to filter. Each item carries the asset id, name and uri, so an auth
manager can authorize on any of them (e.g. restrict by uri prefix).
:param user: the user
:param method: the method to filter on
"""
return {
details.id
for details in assets
if details.id is not None and self.is_authorized_asset(method=method, details=details, user=user)
}

@provide_session
def get_authorized_connections(
self,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -63,6 +63,8 @@ class AssetDetails:
"""Represents the details of an asset."""

id: str | None = None
name: str | None = None
uri: str | None = None


@dataclass
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -70,7 +70,9 @@
from airflow.api_fastapi.core_api.openapi.exceptions import create_openapi_http_exception_doc
from airflow.api_fastapi.core_api.security import (
GetUserDep,
ReadableAssetEventsByAssetFilterDep,
ReadableAssetEventsFilterDep,
ReadableAssetsFilterDep,
ReadableDagsFilterDep,
requires_access_asset,
requires_access_asset_alias,
Expand Down Expand Up @@ -157,6 +159,7 @@ def get_assets(
SortParam,
Depends(SortParam(["id", "name", "uri", "created_at", "updated_at"], AssetModel).dynamic_depends()),
],
readable_assets_filter: ReadableAssetsFilterDep,
session: SessionDep,
) -> AssetCollectionResponse:
"""Get assets."""
Expand Down Expand Up @@ -202,6 +205,7 @@ def get_assets(
uri_pattern,
uri_prefix_pattern,
dag_ids,
readable_assets_filter,
],
order_by=order_by,
offset=offset,
Expand Down Expand Up @@ -343,6 +347,7 @@ def get_asset_events(
extra_filter: QueryAssetEventExtraFilter,
timestamp_range: Annotated[RangeFilter, Depends(datetime_range_filter_factory("timestamp", AssetEvent))],
readable_asset_events_filter: ReadableAssetEventsFilterDep,
readable_asset_events_by_asset_filter: ReadableAssetEventsByAssetFilterDep,
session: SessionDep,
) -> AssetEventCollectionResponse:
"""Get asset events."""
Expand All @@ -367,6 +372,7 @@ def get_asset_events(
extra_filter,
timestamp_range,
readable_asset_events_filter,
readable_asset_events_by_asset_filter,
],
order_by=order_by,
offset=offset,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -48,6 +48,7 @@
)
from airflow.api_fastapi.core_api.routes.public.assets import OnlyActiveFilter
from airflow.api_fastapi.core_api.security import (
ReadableAssetsFilterDep,
requires_access_asset,
requires_access_asset_alias,
requires_access_dag,
Expand Down Expand Up @@ -109,6 +110,7 @@ def get_assets(
).dynamic_depends(default="-last_asset_event_timestamp")
),
],
readable_assets_filter: ReadableAssetsFilterDep,
session: SessionDep,
) -> AssetCollectionResponse:
"""Get assets. Like the public endpoint, but also supports sorting by group and last asset event timestamp."""
Expand All @@ -125,6 +127,7 @@ def get_assets(
group_prefix_pattern,
dag_ids,
last_asset_event_timestamp_range,
readable_assets_filter,
],
order_by=order_by,
offset=offset,
Expand Down
73 changes: 70 additions & 3 deletions airflow-core/src/airflow/api_fastapi/core_api/security.py
Original file line number Diff line number Diff line change
Expand Up @@ -70,7 +70,7 @@
from airflow.api_fastapi.core_api.datamodels.variables import VariableBody
from airflow.configuration import conf
from airflow.models import Connection, Pool, Variable
from airflow.models.asset import AssetEvent
from airflow.models.asset import AssetEvent, AssetModel
from airflow.models.backfill import Backfill
from airflow.models.dag import DagModel, DagRun, DagTag
from airflow.models.dag_version import DagVersion
Expand Down Expand Up @@ -991,16 +991,83 @@ def inner(
return inner


class PermittedAssetFilter(OrmClause[set[int]]):
"""A parameter that filters the permitted assets for the user."""

def to_orm(self, select: Select) -> Select:
return select.where(AssetModel.id.in_(self.value or set()))


# Uncorrelated on purpose: a correlated EXISTS would bind to the outer AssetModel join some
# asset event queries add (e.g. for name filters) and produce an invalid statement.
_existing_asset_ids = select(AssetModel.id)


class PermittedAssetEventByAssetFilter(PermittedAssetFilter):
"""A parameter that filters asset events to those of the assets the user may read."""

def to_orm(self, select: Select) -> Select:
# Events outlive their asset by design. Once the asset row is gone there is no name or
# uri left to authorize on, so such events stay visible to any caller who may read
# assets, the same way events with no source Dag do.
return select.where(
or_(
AssetEvent.asset_id.in_(self.value or set()),
AssetEvent.asset_id.not_in(_existing_asset_ids),
)
)


def permitted_asset_filter_factory(
method: ResourceMethod,
filter_class: type[PermittedAssetFilter] = PermittedAssetFilter,
) -> Callable[[BaseUser, BaseAuthManager], PermittedAssetFilter]:
"""
Create a callable for Depends in FastAPI that returns a filter of the permitted assets for the user.

:param method: whether filter readable or writable.
:param filter_class: the filter class to instantiate, defaulting to ``PermittedAssetFilter``.
"""

def depends_permitted_assets_filter(
user: GetUserDep,
auth_manager: AuthManagerDep,
) -> PermittedAssetFilter:
authorized_assets: set[int] = auth_manager.get_authorized_assets(user=user, method=method)
return filter_class(authorized_assets)

return depends_permitted_assets_filter


ReadableAssetsFilterDep = Annotated[PermittedAssetFilter, Depends(permitted_asset_filter_factory("GET"))]
ReadableAssetEventsByAssetFilterDep = Annotated[
PermittedAssetEventByAssetFilter,
Depends(permitted_asset_filter_factory("GET", PermittedAssetEventByAssetFilter)),
]


def _build_asset_details(asset_id: str | None) -> AssetDetails:
"""Resolve the name and uri of the asset so an auth manager can authorize on more than the id."""
if asset_id is None or not asset_id.isdigit():
# A non-numeric id fails the route's own path validation; there is nothing to look up.
return AssetDetails(id=asset_id)
name_and_uri = AssetModel.get_name_and_uri(int(asset_id))
if name_and_uri is None:
return AssetDetails(id=asset_id)
name, uri = name_and_uri
return AssetDetails(id=asset_id, name=name, uri=uri)


def requires_access_asset(method: ResourceMethod) -> Callable[[Request, BaseUser], None]:
def inner(
request: Request,
user: GetUserDep,
) -> None:
asset_id = request.path_params.get("asset_id")
details = _build_asset_details(request.path_params.get("asset_id"))

_requires_access(
is_authorized_callback=lambda: get_auth_manager().is_authorized_asset(
method=method, details=AssetDetails(id=asset_id), user=user
method=method, details=details, user=user
),
)

Expand Down
8 changes: 8 additions & 0 deletions airflow-core/src/airflow/models/asset.py
Original file line number Diff line number Diff line change
Expand Up @@ -41,6 +41,7 @@
from airflow._shared.timezones import timezone
from airflow.configuration import conf as airflow_conf
from airflow.models.base import Base, StringID
from airflow.utils.session import NEW_SESSION, provide_session
from airflow.utils.sqlalchemy import UtcDateTime

if TYPE_CHECKING:
Expand Down Expand Up @@ -395,6 +396,13 @@ def to_serialized(self) -> SerializedAsset:
def add_trigger(self, trigger: Trigger, watcher_name: str):
self.watchers.append(AssetWatcherModel(name=watcher_name, trigger_id=trigger.id))

@staticmethod
@provide_session
def get_name_and_uri(asset_id: int, *, session: Session = NEW_SESSION) -> tuple[str, str] | None:
stmt = select(AssetModel.name, AssetModel.uri).where(AssetModel.id == asset_id)
row = session.execute(stmt).one_or_none()
return (row.name, row.uri) if row is not None else None


class AssetActive(Base):
"""
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -18,15 +18,17 @@

import warnings
from typing import TYPE_CHECKING, Any
from unittest.mock import AsyncMock, MagicMock, Mock, patch
from unittest.mock import AsyncMock, MagicMock, Mock, create_autospec, patch

import pytest
from jwt import InvalidTokenError
from sqlalchemy.orm import Session

from airflow.api_fastapi.auth.managers.base_auth_manager import BaseAuthManager, T
from airflow.api_fastapi.auth.managers.models.base_user import BaseUser
from airflow.api_fastapi.auth.managers.models.resource_details import (
AccessView,
AssetDetails,
ConnectionDetails,
DagDetails,
PoolDetails,
Expand All @@ -44,7 +46,6 @@
from airflow.api_fastapi.auth.managers.base_auth_manager import ResourceMethod
from airflow.api_fastapi.auth.managers.models.resource_details import (
AssetAliasDetails,
AssetDetails,
ConfigurationDetails,
DagAccessEntity,
)
Expand Down Expand Up @@ -798,6 +799,75 @@ def side_effect_func(
result = auth_manager.get_authorized_pools(user=user, session=session)
assert result == expected

@pytest.mark.parametrize(
("authorized_uri_prefix", "rows", "expected"),
[
pytest.param(None, [(1, "a", "s3://team-a/a"), (2, "b", "s3://team-b/b")], set(), id="no-access"),
pytest.param(
"s3://team-a/",
[(1, "a", "s3://team-a/a"), (2, "b", "s3://team-b/b"), (3, "c", "s3://team-a/c")],
{1, 3},
id="access-by-uri-prefix",
),
],
)
def test_get_authorized_assets(self, auth_manager, authorized_uri_prefix, rows: list, expected: set):
def side_effect_func(
*,
method: ResourceMethod,
user: BaseAuthManagerUserTest,
details: AssetDetails | None = None,
):
if not details or not details.uri or authorized_uri_prefix is None:
return False
return details.uri.startswith(authorized_uri_prefix)

auth_manager.is_authorized_asset = create_autospec(
auth_manager.is_authorized_asset, side_effect=side_effect_func
)
user = Mock(spec=BaseAuthManagerUserTest)
session = Mock(spec=Session)
session.execute.return_value.all.return_value = rows
result = auth_manager.get_authorized_assets(user=user, session=session)
assert result == expected

def test_get_authorized_assets_passes_id_name_and_uri_to_filter(self, auth_manager):
auth_manager.filter_authorized_assets = create_autospec(
auth_manager.filter_authorized_assets, return_value={"2"}
)
user = Mock(spec=BaseAuthManagerUserTest)
session = Mock(spec=Session)
session.execute.return_value.all.return_value = [(1, "a", "s3://a"), (2, "b", "s3://b")]

result = auth_manager.get_authorized_assets(user=user, method="PUT", session=session)

auth_manager.filter_authorized_assets.assert_called_once_with(
assets=[
AssetDetails(id="1", name="a", uri="s3://a"),
AssetDetails(id="2", name="b", uri="s3://b"),
],
user=user,
method="PUT",
)
assert result == {2}

def test_filter_authorized_assets(self, auth_manager):
assets = [
AssetDetails(id="1", name="a", uri="s3://a"),
AssetDetails(id="2", name="b", uri="s3://b"),
AssetDetails(name="no-id", uri="s3://no-id"),
]
auth_manager.is_authorized_asset = create_autospec(
auth_manager.is_authorized_asset, side_effect=lambda *, method, user, details: details.id != "2"
)
user = Mock(spec=BaseAuthManagerUserTest)

result = auth_manager.filter_authorized_assets(assets=assets, user=user, method="DELETE")

assert result == {"1"}
auth_manager.is_authorized_asset.assert_any_call(method="DELETE", details=assets[0], user=user)
auth_manager.is_authorized_asset.assert_any_call(method="DELETE", details=assets[1], user=user)

@pytest.mark.parametrize(
("user_id", "assigned_users", "expected"),
[
Expand Down
Loading