diff --git a/airflow-core/docs/core-concepts/auth-manager/index.rst b/airflow-core/docs/core-concepts/auth-manager/index.rst index b5e11baeb8db9..f54bad49f9b99 100644 --- a/airflow-core/docs/core-concepts/auth-manager/index.rst +++ b/airflow-core/docs/core-concepts/auth-manager/index.rst @@ -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. diff --git a/airflow-core/newsfragments/72682.feature.rst b/airflow-core/newsfragments/72682.feature.rst new file mode 100644 index 0000000000000..b005272f4f879 --- /dev/null +++ b/airflow-core/newsfragments/72682.feature.rst @@ -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. diff --git a/airflow-core/src/airflow/api_fastapi/auth/managers/base_auth_manager.py b/airflow-core/src/airflow/api_fastapi/auth/managers/base_auth_manager.py index baa8b0174b092..e26ed408dd14a 100644 --- a/airflow-core/src/airflow/api_fastapi/auth/managers/base_auth_manager.py +++ b/airflow-core/src/airflow/api_fastapi/auth/managers/base_auth_manager.py @@ -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, @@ -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 @@ -72,7 +74,6 @@ from airflow.api_fastapi.auth.managers.models.resource_details import ( AccessView, AssetAliasDetails, - AssetDetails, ConfigurationDetails, DagAccessEntity, ) @@ -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, diff --git a/airflow-core/src/airflow/api_fastapi/auth/managers/models/resource_details.py b/airflow-core/src/airflow/api_fastapi/auth/managers/models/resource_details.py index 89319b9237bdf..cd8b0a7f55cba 100644 --- a/airflow-core/src/airflow/api_fastapi/auth/managers/models/resource_details.py +++ b/airflow-core/src/airflow/api_fastapi/auth/managers/models/resource_details.py @@ -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 diff --git a/airflow-core/src/airflow/api_fastapi/core_api/routes/public/assets.py b/airflow-core/src/airflow/api_fastapi/core_api/routes/public/assets.py index 0a30d719d1a4e..e9fdcdfa737f0 100644 --- a/airflow-core/src/airflow/api_fastapi/core_api/routes/public/assets.py +++ b/airflow-core/src/airflow/api_fastapi/core_api/routes/public/assets.py @@ -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, @@ -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.""" @@ -202,6 +205,7 @@ def get_assets( uri_pattern, uri_prefix_pattern, dag_ids, + readable_assets_filter, ], order_by=order_by, offset=offset, @@ -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.""" @@ -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, diff --git a/airflow-core/src/airflow/api_fastapi/core_api/routes/ui/assets.py b/airflow-core/src/airflow/api_fastapi/core_api/routes/ui/assets.py index 61384896625e1..dd5feb55783e2 100644 --- a/airflow-core/src/airflow/api_fastapi/core_api/routes/ui/assets.py +++ b/airflow-core/src/airflow/api_fastapi/core_api/routes/ui/assets.py @@ -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, @@ -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.""" @@ -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, diff --git a/airflow-core/src/airflow/api_fastapi/core_api/security.py b/airflow-core/src/airflow/api_fastapi/core_api/security.py index 04472a8d303ec..a1a907ed02d2e 100644 --- a/airflow-core/src/airflow/api_fastapi/core_api/security.py +++ b/airflow-core/src/airflow/api_fastapi/core_api/security.py @@ -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 @@ -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 ), ) diff --git a/airflow-core/src/airflow/models/asset.py b/airflow-core/src/airflow/models/asset.py index d980c2d69acec..aee173590cc12 100644 --- a/airflow-core/src/airflow/models/asset.py +++ b/airflow-core/src/airflow/models/asset.py @@ -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: @@ -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): """ diff --git a/airflow-core/tests/unit/api_fastapi/auth/managers/test_base_auth_manager.py b/airflow-core/tests/unit/api_fastapi/auth/managers/test_base_auth_manager.py index 790d1ca506da9..f25f4dbc9fdfa 100644 --- a/airflow-core/tests/unit/api_fastapi/auth/managers/test_base_auth_manager.py +++ b/airflow-core/tests/unit/api_fastapi/auth/managers/test_base_auth_manager.py @@ -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, @@ -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, ) @@ -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"), [ diff --git a/airflow-core/tests/unit/api_fastapi/core_api/routes/public/test_assets.py b/airflow-core/tests/unit/api_fastapi/core_api/routes/public/test_assets.py index d2a607522811d..33cc8f9fbd159 100644 --- a/airflow-core/tests/unit/api_fastapi/core_api/routes/public/test_assets.py +++ b/airflow-core/tests/unit/api_fastapi/core_api/routes/public/test_assets.py @@ -26,7 +26,11 @@ from airflow._shared.timezones import timezone from airflow.api_fastapi.auth.managers.base_auth_manager import BaseAuthManager -from airflow.api_fastapi.auth.managers.models.resource_details import DagAccessEntity, DagDetails +from airflow.api_fastapi.auth.managers.models.resource_details import ( + AssetDetails, + DagAccessEntity, + DagDetails, +) from airflow.api_fastapi.core_api.security import PermittedAssetEventFilter from airflow.models import DagModel from airflow.models.asset import ( @@ -369,7 +373,9 @@ def test_should_respond_200(self, test_client, session): assert len(session.scalars(select(AssetModel)).all()) == 3 assert len(session.scalars(select(AssetActive)).all()) == 2 - with assert_queries_count(7): + # 8 rather than 7: resolving the caller's readable assets, so the list can be scoped + # to them, costs one additional query. + with assert_queries_count(8): response = test_client.get("/assets") assert response.status_code == 200 @@ -412,6 +418,22 @@ def test_should_respond_200(self, test_client, session): "total_entries": 2, } + @mock.patch("airflow.api_fastapi.auth.managers.base_auth_manager.BaseAuthManager.get_authorized_assets") + def test_should_return_only_assets_the_caller_may_read( + self, mock_get_authorized_assets, test_client, session + ): + self.create_assets(session=session, num=3) + mock_get_authorized_assets.return_value = {1, 3} + + response = test_client.get("/assets") + + mock_get_authorized_assets.assert_called_once_with(user=mock.ANY, method="GET") + assert response.status_code == 200 + body = response.json() + assert [asset["id"] for asset in body["assets"]] == [1, 3] + # The count must be scoped too, so the existence of hidden assets does not leak. + assert body["total_entries"] == 2 + def test_should_respond_200_with_watchers(self, test_client, session): """Test that assets with watchers return the watcher information in the API response.""" asset1, asset2 = self.create_assets_with_watchers(session=session, num=2) @@ -587,7 +609,7 @@ def test_query_count_with_multi_team(self, test_client, session): """ _create_assets_with_team_references(session, num=5) - with assert_queries_count(9): + with assert_queries_count(10): response = test_client.get("/assets") assert response.status_code == 200 @@ -996,6 +1018,37 @@ def test_endpoint_returns_only_events_the_caller_may_read( # The count must be scoped too, so the existence of hidden events does not leak. assert body["total_entries"] == len(expected_ids) + @mock.patch("airflow.api_fastapi.auth.managers.base_auth_manager.BaseAuthManager.get_authorized_assets") + def test_endpoint_returns_only_events_of_assets_the_caller_may_read( + self, mock_get_authorized_assets, test_client, session + ): + """An event of an asset the caller cannot read is hidden even when its source Dag is readable. + + An event whose asset has since been deleted has no name or uri left to authorize on and + stays visible, like an event with no source Dag. + """ + mock_get_authorized_assets.return_value = {2} + self.create_assets(session=session, num=2) + session.add_all( + [ + AssetEvent(id=1, asset_id=1, extra={}, source_dag_id="source_dag_id", timestamp=DEFAULT_DATE), + AssetEvent(id=2, asset_id=2, extra={}, source_dag_id="source_dag_id", timestamp=DEFAULT_DATE), + AssetEvent(id=3, asset_id=1, extra={}, timestamp=DEFAULT_DATE), + # Asset 99 does not exist any more. + AssetEvent( + id=4, asset_id=99, extra={}, source_dag_id="source_dag_id", timestamp=DEFAULT_DATE + ), + ] + ) + session.commit() + + response = test_client.get("/assets/events") + + assert response.status_code == 200 + body = response.json() + assert [event["id"] for event in body["asset_events"]] == [2, 4] + assert body["total_entries"] == 2 + class TestGetAssetEvents(TestAssets): def test_should_respond_200(self, test_client, session): @@ -1007,9 +1060,9 @@ def test_should_respond_200(self, test_client, session): session.commit() assert len(assets) == 2 - # 5 rather than 4: resolving the caller's readable Dags, so events can be scoped - # to them, costs one additional query — the same cost the queued-events routes pay. - with assert_queries_count(5): + # 6 rather than 4: resolving the caller's readable Dags and readable assets, so events + # can be scoped to them, costs one additional query each. + with assert_queries_count(6): response = test_client.get("/assets/events") assert response.status_code == 200 @@ -1632,7 +1685,9 @@ def test_should_respond_200(self, test_client, *, session): self.create_assets(num=1) assert session.scalars(select(func.count(AssetModel.id))).one() == 1 tz_datetime_format = from_datetime_to_zulu_without_ms(DEFAULT_DATE) - with assert_queries_count(6): + # 7 rather than 6: resolving the asset's name and uri for the authorization check costs + # one additional query. + with assert_queries_count(7): response = test_client.get("/assets/1") assert response.status_code == 200 assert response.json() == { @@ -1651,6 +1706,23 @@ def test_should_respond_200(self, test_client, *, session): "last_asset_event": {"id": None, "timestamp": None}, } + @mock.patch( + "airflow.api_fastapi.auth.managers.simple.simple_auth_manager.SimpleAuthManager.is_authorized_asset" + ) + def test_should_authorize_with_asset_name_and_uri(self, mock_is_authorized_asset, test_client): + """The auth manager receives the name and uri so it can authorize on more than the id.""" + self.create_assets(num=1) + mock_is_authorized_asset.return_value = True + + response = test_client.get("/assets/1") + + assert response.status_code == 200 + mock_is_authorized_asset.assert_called_once_with( + method="GET", + details=AssetDetails(id="1", name="simple1", uri="s3://bucket/key/1"), + user=mock.ANY, + ) + @provide_session def test_should_respond_200_with_watchers(self, test_client, *, session): """Test that single asset endpoint returns watcher information.""" @@ -2512,7 +2584,9 @@ def test_should_respond_200(self, test_client, session, create_dummy_dag): (asset,) = self.create_assets(session=session, num=1) self._create_asset_dag_run_queues(dag_id, asset.id, session) - with assert_queries_count(3): + # 4 rather than 3: resolving the asset's name and uri for the authorization check costs + # one additional query. + with assert_queries_count(4): response = test_client.get(f"/assets/{asset.id}/queuedEvents") assert response.status_code == 200 diff --git a/airflow-core/tests/unit/api_fastapi/core_api/routes/ui/test_assets.py b/airflow-core/tests/unit/api_fastapi/core_api/routes/ui/test_assets.py index 36b200405cc59..ef4ca97fb2bcf 100644 --- a/airflow-core/tests/unit/api_fastapi/core_api/routes/ui/test_assets.py +++ b/airflow-core/tests/unit/api_fastapi/core_api/routes/ui/test_assets.py @@ -531,6 +531,25 @@ def test_should_respond_200(self, test_client, session): assert body["total_entries"] == 1 assert body["assets"][0]["name"] == "ui_asset" + @mock.patch("airflow.api_fastapi.auth.managers.base_auth_manager.BaseAuthManager.get_authorized_assets") + def test_should_return_only_assets_the_caller_may_read( + self, mock_get_authorized_assets, test_client, session + ): + assets = [AssetModel(name=f"asset{i}", uri=f"s3://bucket/asset{i}", group="asset") for i in range(3)] + session.add_all(assets) + session.add_all(AssetActive.for_asset(asset) for asset in assets) + session.commit() + mock_get_authorized_assets.return_value = {assets[1].id} + + response = test_client.get("/assets") + + mock_get_authorized_assets.assert_called_once_with(user=mock.ANY, method="GET") + assert response.status_code == 200 + body = response.json() + assert [asset["name"] for asset in body["assets"]] == ["asset1"] + # The count must be scoped too, so the existence of hidden assets does not leak. + assert body["total_entries"] == 1 + def test_sort_by_last_asset_event_timestamp(self, test_client, session): older = AssetModel(name="older", uri="s3://bucket/older", group="asset") newer = AssetModel(name="newer", uri="s3://bucket/newer", group="asset") @@ -828,7 +847,8 @@ def test_query_count(self, test_client, session): assets[i].aliases.append(AssetAliasModel(name=f"alias{i}", group="")) session.commit() - with assert_queries_count(8): + # One of these queries resolves the caller's readable assets so the list can be scoped to them. + with assert_queries_count(9): assert test_client.get("/assets").status_code == 200 @conf_vars({("core", "multi_team"): "True"}) @@ -856,7 +876,7 @@ def test_query_count_with_multi_team(self, test_client, session): session.add(TaskOutletAssetReference(dag_id=f"producing_dag{i}", task_id="task", asset=asset)) session.commit() - with assert_queries_count(12): + with assert_queries_count(13): response = test_client.get("/assets") assert response.status_code == 200 diff --git a/airflow-core/tests/unit/api_fastapi/core_api/test_security.py b/airflow-core/tests/unit/api_fastapi/core_api/test_security.py index 168a8edafdd1e..151859f203c6e 100644 --- a/airflow-core/tests/unit/api_fastapi/core_api/test_security.py +++ b/airflow-core/tests/unit/api_fastapi/core_api/test_security.py @@ -22,6 +22,7 @@ import pytest from fastapi import HTTPException, Request from jwt import ExpiredSignatureError, InvalidTokenError +from sqlalchemy import select from sqlalchemy.orm import Session from airflow import settings @@ -30,6 +31,7 @@ from airflow.api_fastapi.auth.managers.models.base_user import BaseUser from airflow.api_fastapi.auth.managers.models.resource_details import ( AccessView, + AssetDetails, ConnectionDetails, DagAccessEntity, DagDetails, @@ -42,9 +44,13 @@ 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 ( + PermittedAssetEventByAssetFilter, + PermittedAssetFilter, _build_dag_run_access_requests, get_user, is_safe_url, + permitted_asset_filter_factory, + requires_access_asset, requires_access_backfill, requires_access_connection, requires_access_connection_bulk, @@ -57,6 +63,7 @@ resolve_user_from_token, ) from airflow.models import Connection, Pool, Variable +from airflow.models.asset import AssetEvent, AssetModel from airflow.models.dag import DagModel from airflow.models.dagbundle import DagBundleModel from airflow.models.team import Team @@ -1486,6 +1493,98 @@ def test_requires_access_variable_bulk_checks_destination_team( user=user, ) + @pytest.mark.parametrize( + ("path_params", "name_and_uri", "expected_details"), + [ + pytest.param( + {"asset_id": "1"}, + ("simple1", "s3://bucket/key/1"), + AssetDetails(id="1", name="simple1", uri="s3://bucket/key/1"), + id="existing-asset", + ), + pytest.param({"asset_id": "1"}, None, AssetDetails(id="1"), id="missing-asset"), + pytest.param({}, None, AssetDetails(id=None), id="no-asset-id"), + ], + ) + @patch.object(AssetModel, "get_name_and_uri") + @patch("airflow.api_fastapi.core_api.security.get_auth_manager") + def test_requires_access_asset_resolves_name_and_uri( + self, mock_get_auth_manager, mock_get_name_and_uri, path_params, name_and_uri, expected_details + ): + auth_manager = Mock(spec=BaseAuthManager) + auth_manager.is_authorized_asset.return_value = True + mock_get_auth_manager.return_value = auth_manager + mock_get_name_and_uri.return_value = name_and_uri + + fastapi_request = Mock(spec=Request) + fastapi_request.path_params = path_params + user = Mock(spec=BaseUser) + + requires_access_asset("GET")(fastapi_request, user) + + auth_manager.is_authorized_asset.assert_called_once_with( + method="GET", details=expected_details, user=user + ) + if path_params: + mock_get_name_and_uri.assert_called_once_with(1) + else: + mock_get_name_and_uri.assert_not_called() + + @patch.object(AssetModel, "get_name_and_uri") + @patch("airflow.api_fastapi.core_api.security.get_auth_manager") + def test_requires_access_asset_skips_lookup_for_non_numeric_id( + self, mock_get_auth_manager, mock_get_name_and_uri + ): + auth_manager = Mock(spec=BaseAuthManager) + auth_manager.is_authorized_asset.return_value = True + mock_get_auth_manager.return_value = auth_manager + + fastapi_request = Mock(spec=Request) + fastapi_request.path_params = {"asset_id": "not-a-number"} + user = Mock(spec=BaseUser) + + requires_access_asset("GET")(fastapi_request, user) + + mock_get_name_and_uri.assert_not_called() + auth_manager.is_authorized_asset.assert_called_once_with( + method="GET", details=AssetDetails(id="not-a-number"), user=user + ) + + @pytest.mark.parametrize( + ("filter_class", "model", "expected_column"), + [ + pytest.param(PermittedAssetFilter, AssetModel, "asset.id IN", id="assets"), + pytest.param( + PermittedAssetEventByAssetFilter, AssetEvent, "asset_event.asset_id IN", id="asset-events" + ), + ], + ) + def test_permitted_asset_filters_scope_on_asset_id(self, filter_class, model, expected_column): + rendered = str(filter_class({1, 2}).to_orm(select(model))) + assert expected_column in rendered + + def test_permitted_asset_event_filter_keeps_events_of_deleted_assets(self): + rendered = str(PermittedAssetEventByAssetFilter({1}).to_orm(select(AssetEvent))) + assert "asset_event.asset_id NOT IN (SELECT asset.id" in rendered + + @pytest.mark.parametrize( + "filter_class", + [ + pytest.param(PermittedAssetFilter, id="default"), + pytest.param(PermittedAssetEventByAssetFilter, id="events"), + ], + ) + def test_permitted_asset_filter_factory(self, filter_class): + auth_manager = Mock(spec=BaseAuthManager) + auth_manager.get_authorized_assets.return_value = {1, 3} + user = Mock(spec=BaseUser) + + permitted_filter = permitted_asset_filter_factory("GET", filter_class)(user, auth_manager) + + assert isinstance(permitted_filter, filter_class) + assert permitted_filter.value == {1, 3} + auth_manager.get_authorized_assets.assert_called_once_with(user=user, method="GET") + class TestAuthManagerDependency: """Test the auth_manager_from_app dependency function.""" diff --git a/providers/amazon/src/airflow/providers/amazon/aws/auth_manager/aws_auth_manager.py b/providers/amazon/src/airflow/providers/amazon/aws/auth_manager/aws_auth_manager.py index c2733d902468f..b29c21971481b 100644 --- a/providers/amazon/src/airflow/providers/amazon/aws/auth_manager/aws_auth_manager.py +++ b/providers/amazon/src/airflow/providers/amazon/aws/auth_manager/aws_auth_manager.py @@ -379,6 +379,36 @@ def batch_is_authorized_variable( ] return self.avp_facade.batch_is_authorized(requests=facade_requests, user=user) + def filter_authorized_assets( + self, + *, + assets: Sequence[AssetDetails], + user: AwsAuthManagerUser, + method: ResourceMethod = "GET", + ) -> set[str]: + requests: dict[str, IsAuthorizedRequest] = {} + requests_list: list[IsAuthorizedRequest] = [] + for details in assets: + if details.id is None: + continue + request: IsAuthorizedRequest = { + "method": method, + "entity_type": AvpEntities.ASSET, + "entity_id": details.id, + } + requests[details.id] = request + requests_list.append(request) + + batch_is_authorized_results = self.avp_facade.get_batch_is_authorized_results( + requests=requests_list, user=user + ) + + return { + asset_id + for asset_id, request in requests.items() + if self._is_authorized_from_batch_response(batch_is_authorized_results, request, user) + } + def filter_authorized_connections( self, *, diff --git a/providers/amazon/tests/unit/amazon/aws/auth_manager/test_aws_auth_manager.py b/providers/amazon/tests/unit/amazon/aws/auth_manager/test_aws_auth_manager.py index 91d69c5bf192a..a4e33e31ab486 100644 --- a/providers/amazon/tests/unit/amazon/aws/auth_manager/test_aws_auth_manager.py +++ b/providers/amazon/tests/unit/amazon/aws/auth_manager/test_aws_auth_manager.py @@ -24,7 +24,11 @@ from airflow.exceptions import AirflowProviderDeprecationWarning -from tests_common.test_utils.version_compat import AIRFLOW_V_3_0_PLUS, AIRFLOW_V_3_2_PLUS +from tests_common.test_utils.version_compat import ( + AIRFLOW_V_3_0_PLUS, + AIRFLOW_V_3_2_PLUS, + AIRFLOW_V_3_4_PLUS, +) if not AIRFLOW_V_3_0_PLUS: pytest.skip("AWS auth manager is only compatible with Airflow >= 3.0.0", allow_module_level=True) @@ -32,6 +36,7 @@ from airflow.api_fastapi.app import AUTH_MANAGER_FASTAPI_APP_PREFIX from airflow.api_fastapi.auth.managers.models.resource_details import ( AccessView, + AssetDetails, BackfillDetails, ConfigurationDetails, ConnectionDetails, @@ -894,6 +899,51 @@ def test_filter_authorized( auth_manager.avp_facade.get_batch_is_authorized_results.assert_called() assert result == expected_result + @pytest.mark.skipif( + not AIRFLOW_V_3_4_PLUS, reason="AssetDetails name and uri not available before Airflow 3.4.0" + ) + def test_filter_authorized_assets(self, auth_manager): + user = AwsAuthManagerUser(user_id="test_user_id1", groups=[]) + assets = [ + AssetDetails(id="1", name="sales", uri="s3://team-a/sales.csv"), + AssetDetails(id="2", name="salary", uri="s3://team-b/salary.csv"), + ] + avp_entity = AvpEntities.ASSET.value + batch_is_authorized_output = [ + { + "request": { + "principal": {"entityType": "Airflow::User", "entityId": "test_user_id1"}, + "action": {"actionType": "Airflow::Action", "actionId": f"{avp_entity}.GET"}, + "resource": {"entityType": f"Airflow::{avp_entity}", "entityId": "1"}, + }, + "decision": "ALLOW", + }, + { + "request": { + "principal": {"entityType": "Airflow::User", "entityId": "test_user_id1"}, + "action": {"actionType": "Airflow::Action", "actionId": f"{avp_entity}.GET"}, + "resource": {"entityType": f"Airflow::{avp_entity}", "entityId": "2"}, + }, + "decision": "DENY", + }, + ] + auth_manager.avp_facade.get_batch_is_authorized_results = Mock( + return_value=batch_is_authorized_output + ) + + result = auth_manager.filter_authorized_assets(assets=assets, user=user, method="GET") + + assert result == {"1"} + auth_manager.avp_facade.get_batch_is_authorized_results.assert_called_once() + sent = auth_manager.avp_facade.get_batch_is_authorized_results.call_args.kwargs["requests"] + assert [request["entity_type"] for request in sent] == [AvpEntities.ASSET, AvpEntities.ASSET] + assert {request["entity_id"] for request in sent} == {"1", "2"} + + def test_filter_authorized_assets_empty(self, auth_manager, test_user): + auth_manager.avp_facade.get_batch_is_authorized_results = Mock(return_value=[]) + + assert auth_manager.filter_authorized_assets(assets=[], user=test_user, method="GET") == set() + def test_get_url_login(self, auth_manager): result = auth_manager.get_url_login() assert result == f"{AUTH_MANAGER_FASTAPI_APP_PREFIX}/login" diff --git a/providers/fab/src/airflow/providers/fab/auth_manager/fab_auth_manager.py b/providers/fab/src/airflow/providers/fab/auth_manager/fab_auth_manager.py index 4bb6d8abe4f71..7e7241e48f03d 100644 --- a/providers/fab/src/airflow/providers/fab/auth_manager/fab_auth_manager.py +++ b/providers/fab/src/airflow/providers/fab/auth_manager/fab_auth_manager.py @@ -55,6 +55,7 @@ from airflow.api_fastapi.common.types import ExtraMenuItem, MenuItem from airflow.exceptions import AirflowConfigException, AirflowProviderDeprecationWarning from airflow.models import Connection, DagModel, Pool, Variable +from airflow.models.asset import AssetModel from airflow.providers.common.compat.sdk import AirflowException, conf from airflow.providers.common.compat.security.access_view import ( AUDIT_LOGS_ALL_ACCESS_VIEW, @@ -563,6 +564,26 @@ def filter_authorized_menu_items(self, menu_items: list[MenuItem], user: User) - ) ] + @provide_session + def get_authorized_assets( + self, + *, + user: User, + method: ResourceMethod = "GET", + session: Session = NEW_SESSION, + ) -> set[int]: + """ + Get the ids of the assets the user has access to. + + Fab auth manager does not allow fine-grained access with assets. Thus, return all the asset ids. + + :param user: the user + :param method: the method to filter on + :param session: the session + """ + rows = session.execute(select(AssetModel.id)).scalars().all() + return set(rows) + @provide_session def get_authorized_connections( self, diff --git a/providers/fab/tests/unit/fab/auth_manager/test_fab_auth_manager.py b/providers/fab/tests/unit/fab/auth_manager/test_fab_auth_manager.py index 43841959cd105..5ddad4a1ae870 100644 --- a/providers/fab/tests/unit/fab/auth_manager/test_fab_auth_manager.py +++ b/providers/fab/tests/unit/fab/auth_manager/test_fab_auth_manager.py @@ -904,6 +904,12 @@ def test_filter_authorized_menu_items( result = auth_manager.filter_authorized_menu_items(menu_items, user=user) assert result == expected_result + def test_get_authorized_assets(self, auth_manager): + session = Mock() + session.execute.return_value.scalars.return_value.all.return_value = [1, 2] + result = auth_manager.get_authorized_assets(user=Mock(), method="GET", session=session) + assert result == {1, 2} + def test_get_authorized_connections(self, auth_manager): session = Mock() session.execute.return_value.scalars.return_value.all.return_value = ["conn1", "conn2"] 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 3fe177bc6cc04..d945b4da47090 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 @@ -664,6 +664,38 @@ def check(request: IsAuthorizedVariableRequest) -> bool: results = executor.map(check, requests) return all(results) + def filter_authorized_assets( + self, + *, + assets: Sequence[AssetDetails], + user: KeycloakAuthManagerUser, + method: ResourceMethod = "GET", + ) -> set[str]: + candidates = [details for details in assets if details.id is not None] + cache_key = ( + user.get_id(), + method, + frozenset(cast("str", details.id) for details in candidates), + ) + + def query_keycloak() -> set[str]: + if not candidates: + return set() + max_workers = min( + len(candidates), conf.getint(CONF_SECTION_NAME, CONF_REQUESTS_POOL_SIZE_KEY, fallback=10) + ) + + def check(details: AssetDetails) -> tuple[str, bool]: + return cast("str", details.id), self.is_authorized_asset( + method=method, user=user, details=details + ) + + with ThreadPoolExecutor(max_workers=max_workers) as executor: + results = executor.map(check, candidates) + return {asset_id for asset_id, authorized in results if authorized} + + return single_flight(cache_key, query_keycloak) + def filter_authorized_connections( self, *, 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 3e05e01304ac5..e52b0c24019df 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 @@ -45,6 +45,7 @@ AIRFLOW_V_3_1_7_PLUS, AIRFLOW_V_3_2_PLUS, AIRFLOW_V_3_3_PLUS, + AIRFLOW_V_3_4_PLUS, ) if AIRFLOW_V_3_1_7_PLUS: @@ -1395,6 +1396,43 @@ def test_filter_authorized_dag_ids_cache_hit(self, mock_is_authorized, auth_mana # is_authorized_dag should only be called for the first invocation (2 dag_ids × 1 call) assert mock_is_authorized.call_count == 2 + @pytest.mark.skipif( + not AIRFLOW_V_3_4_PLUS, reason="AssetDetails name and uri not available before Airflow 3.4.0" + ) + @patch.object( + KeycloakAuthManager, + "is_authorized_asset", + side_effect=lambda *, details, **kw: details.uri.startswith("s3://team-a/"), + ) + def test_filter_authorized_assets(self, mock_is_authorized, auth_manager, user): + assets = [ + AssetDetails(id="1", name="sales", uri="s3://team-a/sales.csv"), + AssetDetails(id="2", name="salary", uri="s3://team-b/salary.csv"), + AssetDetails(id="3", name="logs", uri="s3://team-a/logs.csv"), + ] + + result = auth_manager.filter_authorized_assets(assets=assets, user=user, method="GET") + + assert result == {"1", "3"} + assert mock_is_authorized.call_count == 3 + + def test_filter_authorized_assets_empty(self, auth_manager, user): + assert auth_manager.filter_authorized_assets(assets=[], user=user, method="GET") == set() + + @pytest.mark.skipif( + not AIRFLOW_V_3_4_PLUS, reason="AssetDetails name and uri not available before Airflow 3.4.0" + ) + @patch.object(KeycloakAuthManager, "is_authorized_asset", return_value=True) + def test_filter_authorized_assets_cache_hit(self, mock_is_authorized, auth_manager, user): + """A second identical call is served from the cache without asking Keycloak again.""" + assets = [AssetDetails(id="1", name="sales", uri="s3://team-a/sales.csv")] + + first = auth_manager.filter_authorized_assets(assets=assets, user=user, method="GET") + second = auth_manager.filter_authorized_assets(assets=assets, user=user, method="GET") + + assert first == second == {"1"} + assert mock_is_authorized.call_count == 1 + @patch.object( KeycloakAuthManager, "is_authorized_connection",