From b9b73f91b9fb7969fa312f2fd2ddb2ffbfcb7122 Mon Sep 17 00:00:00 2001 From: Leondon9 Date: Wed, 3 Jun 2026 16:43:32 +0900 Subject: [PATCH 1/2] Add asset partition sensor --- .../example_dags/example_asset_partition.py | 34 ++++ .../src/airflow/jobs/triggerer_job_runner.py | 8 + providers/standard/provider.yaml | 2 + .../airflow/providers/standard/exceptions.py | 4 + .../providers/standard/get_provider_info.py | 2 + .../providers/standard/sensors/asset.py | 123 ++++++++++++++ .../providers/standard/triggers/asset.py | 107 ++++++++++++ .../providers/standard/version_compat.py | 2 + .../tests/unit/standard/sensors/test_asset.py | 157 ++++++++++++++++++ .../unit/standard/triggers/test_asset.py | 104 ++++++++++++ .../sdk/execution_time/request_handlers.py | 38 +++++ .../airflow/sdk/execution_time/supervisor.py | 32 +--- .../execution_time/test_request_handlers.py | 60 ++++++- 13 files changed, 644 insertions(+), 29 deletions(-) create mode 100644 providers/standard/src/airflow/providers/standard/sensors/asset.py create mode 100644 providers/standard/src/airflow/providers/standard/triggers/asset.py create mode 100644 providers/standard/tests/unit/standard/sensors/test_asset.py create mode 100644 providers/standard/tests/unit/standard/triggers/test_asset.py diff --git a/airflow-core/src/airflow/example_dags/example_asset_partition.py b/airflow-core/src/airflow/example_dags/example_asset_partition.py index e7151680404ed..97d0f1e00e130 100644 --- a/airflow-core/src/airflow/example_dags/example_asset_partition.py +++ b/airflow-core/src/airflow/example_dags/example_asset_partition.py @@ -19,6 +19,7 @@ from typing import TYPE_CHECKING +from airflow.providers.standard.sensors.asset import AssetPartitionSensor from airflow.sdk import ( DAG, AllowedKeyMapper, @@ -112,6 +113,39 @@ def combine_player_stats(dag_run=None): combine_player_stats() +with DAG( + dag_id="wait_for_combined_player_stats_partition", + schedule="@hourly", + catchup=False, + tags=["example", "player-stats", "sensor"], +): + """ + Wait for a specific ``combined_player_stats`` partition from a time-scheduled Dag. + + ``AssetPartitionSensor`` bridges time-based scheduling and asset partitioning: this hourly + Dag blocks until the combined-stats partition for its own data interval has an asset event. + ``after`` bounds the lookup to the current interval so a stale event with the same partition + key from an earlier run does not satisfy the wait. + """ + + wait_for_combined_partition = AssetPartitionSensor( + task_id="wait_for_combined_partition", + asset=combined_player_stats, + partition_key="{{ data_interval_start.strftime('%Y-%m-%dT%H') }}", + after="{{ data_interval_start }}", + deferrable=True, + ) + + @task + def report_partition_ready(dag_run=None): + """Run once the awaited combined-stats partition is available.""" + if TYPE_CHECKING: + assert dag_run + print(f"combined_player_stats partition ready for {dag_run.logical_date}") + + wait_for_combined_partition >> report_partition_ready() + + @asset( uri="file://analytics/player-stats/computed-player-odds.csv", # Fallback to IdentityMapper if no partition_mapper is specified. diff --git a/airflow-core/src/airflow/jobs/triggerer_job_runner.py b/airflow-core/src/airflow/jobs/triggerer_job_runner.py index 50a899c2cf256..fedf78396b3d0 100644 --- a/airflow-core/src/airflow/jobs/triggerer_job_runner.py +++ b/airflow-core/src/airflow/jobs/triggerer_job_runner.py @@ -73,6 +73,8 @@ DeleteXCom, DRCount, ErrorResponse, + GetAssetEventByAsset, + GetAssetEventByAssetAlias, GetAssetStateStoreByName, GetAssetStateStoreByUri, GetConnection, @@ -108,6 +110,8 @@ handle_delete_asset_state_store_by_uri, handle_delete_variable, handle_delete_xcom, + handle_get_asset_event_by_asset, + handle_get_asset_event_by_asset_alias, handle_get_asset_state_store_by_name, handle_get_asset_state_store_by_uri, handle_get_connection, @@ -641,6 +645,10 @@ def _handle_request(self, msg: ToTriggerSupervisor, log: FilteringBoundLogger, r resp, dump_opts = handle_get_dr_count(self.client, msg) elif isinstance(msg, GetDagRunState): resp, dump_opts = handle_get_dag_run_state(self.client, msg) + elif isinstance(msg, GetAssetEventByAsset): + resp, dump_opts = handle_get_asset_event_by_asset(self.client, msg) + elif isinstance(msg, GetAssetEventByAssetAlias): + resp, dump_opts = handle_get_asset_event_by_asset_alias(self.client, msg) elif isinstance(msg, GetTICount): resp, dump_opts = handle_get_ti_count(self.client, msg) diff --git a/providers/standard/provider.yaml b/providers/standard/provider.yaml index cb493caa69d09..0643478d903cc 100644 --- a/providers/standard/provider.yaml +++ b/providers/standard/provider.yaml @@ -106,6 +106,7 @@ sensors: - airflow.providers.standard.sensors.python - airflow.providers.standard.sensors.filesystem - airflow.providers.standard.sensors.external_task + - airflow.providers.standard.sensors.asset hooks: - integration-name: Standard python-modules: @@ -120,6 +121,7 @@ triggers: - airflow.providers.standard.triggers.file - airflow.providers.standard.triggers.temporal - airflow.providers.standard.triggers.hitl + - airflow.providers.standard.triggers.asset extra-links: - airflow.providers.standard.operators.trigger_dagrun.TriggerDagRunLink diff --git a/providers/standard/src/airflow/providers/standard/exceptions.py b/providers/standard/src/airflow/providers/standard/exceptions.py index df13fd5040b2d..15ca686427816 100644 --- a/providers/standard/src/airflow/providers/standard/exceptions.py +++ b/providers/standard/src/airflow/providers/standard/exceptions.py @@ -67,3 +67,7 @@ class HITLTimeoutError(HITLTriggerEventError): class HITLRejectException(AirflowException): """Raised when an ApprovalOperator receives a "Reject" response when fail_on_reject is set to True.""" + + +class AssetPartitionTriggerEventError(AirflowException): + """Raised when AssetPartitionTrigger reports an error event.""" diff --git a/providers/standard/src/airflow/providers/standard/get_provider_info.py b/providers/standard/src/airflow/providers/standard/get_provider_info.py index 1f7b2049454d1..d3a2f063fd890 100644 --- a/providers/standard/src/airflow/providers/standard/get_provider_info.py +++ b/providers/standard/src/airflow/providers/standard/get_provider_info.py @@ -75,6 +75,7 @@ def get_provider_info(): "airflow.providers.standard.sensors.python", "airflow.providers.standard.sensors.filesystem", "airflow.providers.standard.sensors.external_task", + "airflow.providers.standard.sensors.asset", ], } ], @@ -96,6 +97,7 @@ def get_provider_info(): "airflow.providers.standard.triggers.file", "airflow.providers.standard.triggers.temporal", "airflow.providers.standard.triggers.hitl", + "airflow.providers.standard.triggers.asset", ], } ], diff --git a/providers/standard/src/airflow/providers/standard/sensors/asset.py b/providers/standard/src/airflow/providers/standard/sensors/asset.py new file mode 100644 index 0000000000000..2ed41efec5fcd --- /dev/null +++ b/providers/standard/src/airflow/providers/standard/sensors/asset.py @@ -0,0 +1,123 @@ +# +# 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 datetime +from collections.abc import Sequence +from typing import TYPE_CHECKING, Any + +from airflow.providers.common.compat.sdk import ( + AirflowOptionalProviderFeatureException, + BaseSensorOperator, + conf, + timezone, +) +from airflow.providers.standard.version_compat import AIRFLOW_V_3_4_PLUS + +if not AIRFLOW_V_3_4_PLUS: + raise AirflowOptionalProviderFeatureException("Asset partition sensor needs Airflow 3.4+.") + +from airflow.providers.standard.exceptions import AssetPartitionTriggerEventError +from airflow.providers.standard.triggers.asset import AssetPartitionTrigger + +if TYPE_CHECKING: + from airflow.providers.common.compat.sdk import Asset, Context + + +class AssetPartitionSensor(BaseSensorOperator): + """ + Wait for an asset event with the given partition key. + + :param asset: asset to wait for. + :param partition_key: partition key for the asset event to wait for. + :param after: only match events whose timestamp is at or after this (timezone-aware) + datetime. When unset, any historical event with the partition key satisfies the + wait. Bound it (e.g. ``after="{{ data_interval_start }}"``) when the partition key + can be reused across events, so a stale event from an earlier run is not matched. + :param deferrable: If waiting for completion, whether to defer the task until done. + """ + + template_fields: Sequence[str] = ("partition_key", "after") + ui_color = "#e6f1f2" + + def __init__( + self, + *, + asset: Asset, + partition_key: str, + after: datetime.datetime | str | None = None, + deferrable: bool = conf.getboolean("operators", "default_deferrable", fallback=False), + **kwargs, + ) -> None: + super().__init__(**kwargs) + self.asset = asset + self.partition_key = partition_key + self.after = after + self.deferrable = deferrable + + def poke(self, context: Context) -> bool: + from airflow.sdk.exceptions import AirflowRuntimeError + from airflow.sdk.execution_time.comms import AssetEventsResult, ErrorResponse, GetAssetEventByAsset + from airflow.sdk.execution_time.task_runner import SUPERVISOR_COMMS + + self.log.info("Poking for asset event: asset=%s, partition_key=%s", self.asset, self.partition_key) + after = timezone.parse(self.after) if isinstance(self.after, str) else self.after + response = SUPERVISOR_COMMS.send( + GetAssetEventByAsset( + name=self.asset.name, + uri=self.asset.uri, + partition_key=self.partition_key, + after=after, + ascending=False, + limit=1, + ) + ) + if isinstance(response, ErrorResponse): + raise AirflowRuntimeError(response) + if TYPE_CHECKING: + assert isinstance(response, AssetEventsResult) + return bool(response and response.asset_events) + + def execute(self, context: Context) -> None: + if not self.deferrable: + super().execute(context=context) + return + + if not self.poke(context=context): + self.defer( + timeout=datetime.timedelta(seconds=self.timeout), + trigger=AssetPartitionTrigger( + asset_name=self.asset.name, + asset_uri=self.asset.uri, + partition_key=self.partition_key, + after=self.after, + poke_interval=self.poke_interval, + ), + method_name="execute_complete", + ) + + def execute_complete(self, context: Context, event: dict[str, Any] | None = None) -> None: + if event and event.get("status") == "success": + self.log.info( + "Asset partition event found: asset=%s, partition_key=%s", + self.asset, + self.partition_key, + ) + return + message = event.get("message") if event else "Trigger completed without an event" + raise AssetPartitionTriggerEventError(message) diff --git a/providers/standard/src/airflow/providers/standard/triggers/asset.py b/providers/standard/src/airflow/providers/standard/triggers/asset.py new file mode 100644 index 0000000000000..08383e2298e03 --- /dev/null +++ b/providers/standard/src/airflow/providers/standard/triggers/asset.py @@ -0,0 +1,107 @@ +# +# 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 asyncio +from collections.abc import AsyncIterator +from typing import TYPE_CHECKING, Any + +from airflow.providers.common.compat.sdk import AirflowOptionalProviderFeatureException, timezone +from airflow.providers.standard.version_compat import AIRFLOW_V_3_4_PLUS + +if not AIRFLOW_V_3_4_PLUS: + raise AirflowOptionalProviderFeatureException("Asset partition sensor needs Airflow 3.4+.") + +from airflow.triggers.base import BaseTrigger, TriggerEvent + +if TYPE_CHECKING: + import datetime + + +class AssetPartitionTrigger(BaseTrigger): + """ + Trigger when an asset event exists for the given partition key. + + :param asset_name: name of the asset to wait for. + :param asset_uri: URI of the asset to wait for. + :param partition_key: partition key for the asset event to wait for. + :param after: only match events whose timestamp is at or after this (timezone-aware) + datetime. Leave unset to match any event with the partition key. + :param poke_interval: polling interval in seconds. + """ + + def __init__( + self, + *, + asset_name: str | None, + asset_uri: str | None, + partition_key: str, + after: datetime.datetime | str | None = None, + poke_interval: float = 5.0, + ) -> None: + super().__init__() + self.asset_name = asset_name + self.asset_uri = asset_uri + self.partition_key = partition_key + self.after = after + self.poke_interval = poke_interval + + def serialize(self) -> tuple[str, dict[str, Any]]: + """Serialize AssetPartitionTrigger arguments and classpath.""" + return ( + "airflow.providers.standard.triggers.asset.AssetPartitionTrigger", + { + "asset_name": self.asset_name, + "asset_uri": self.asset_uri, + "partition_key": self.partition_key, + "after": self.after, + "poke_interval": self.poke_interval, + }, + ) + + async def run(self) -> AsyncIterator[TriggerEvent]: + """Poll until the requested asset partition event exists.""" + from airflow.sdk.execution_time.comms import AssetEventsResult, ErrorResponse, GetAssetEventByAsset + from airflow.sdk.execution_time.task_runner import SUPERVISOR_COMMS + + after = timezone.parse(self.after) if isinstance(self.after, str) else self.after + while True: + response = await SUPERVISOR_COMMS.asend( + GetAssetEventByAsset( + name=self.asset_name, + uri=self.asset_uri, + partition_key=self.partition_key, + after=after, + ascending=False, + limit=1, + ) + ) + if isinstance(response, ErrorResponse): + yield TriggerEvent( + { + "status": "error", + "message": f"{response.error.value}: {response.detail}", + } + ) + return + if TYPE_CHECKING: + assert isinstance(response, AssetEventsResult) + if response and response.asset_events: + yield TriggerEvent({"status": "success"}) + return + await asyncio.sleep(self.poke_interval) diff --git a/providers/standard/src/airflow/providers/standard/version_compat.py b/providers/standard/src/airflow/providers/standard/version_compat.py index f99a558937807..146a5828fa4bc 100644 --- a/providers/standard/src/airflow/providers/standard/version_compat.py +++ b/providers/standard/src/airflow/providers/standard/version_compat.py @@ -37,6 +37,7 @@ def get_base_airflow_version_tuple() -> tuple[int, int, int]: AIRFLOW_V_3_1_3_PLUS: bool = get_base_airflow_version_tuple() >= (3, 1, 3) AIRFLOW_V_3_2_PLUS: bool = get_base_airflow_version_tuple() >= (3, 2, 0) AIRFLOW_V_3_3_PLUS: bool = get_base_airflow_version_tuple() >= (3, 3, 0) +AIRFLOW_V_3_4_PLUS: bool = get_base_airflow_version_tuple() >= (3, 4, 0) # BaseOperator: Use 3.1+ due to xcom_push method missing in SDK BaseOperator 3.0.x # This is needed for DecoratedOperator compatibility @@ -60,6 +61,7 @@ def is_arg_set(value): # type: ignore[misc,no-redef] "AIRFLOW_V_3_1_PLUS", "AIRFLOW_V_3_2_PLUS", "AIRFLOW_V_3_3_PLUS", + "AIRFLOW_V_3_4_PLUS", "ArgNotSet", "BaseOperator", "is_arg_set", diff --git a/providers/standard/tests/unit/standard/sensors/test_asset.py b/providers/standard/tests/unit/standard/sensors/test_asset.py new file mode 100644 index 0000000000000..bb30f96d27369 --- /dev/null +++ b/providers/standard/tests/unit/standard/sensors/test_asset.py @@ -0,0 +1,157 @@ +# +# 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 + +from unittest import mock + +import pytest + +from airflow.providers.common.compat.sdk import AirflowException, Asset, TaskDeferred +from airflow.providers.standard.sensors.asset import AssetPartitionSensor +from airflow.providers.standard.triggers.asset import AssetPartitionTrigger +from airflow.sdk import timezone +from airflow.sdk.api.datamodels._generated import AssetEventResponse, AssetResponse +from airflow.sdk.exceptions import AirflowRuntimeError, ErrorType +from airflow.sdk.execution_time import task_runner +from airflow.sdk.execution_time.comms import ( + AssetEventsResult, + ErrorResponse, + GetAssetEventByAsset, +) + + +class TestAssetPartitionSensor: + def test_template_fields(self): + assert AssetPartitionSensor.template_fields == ("partition_key", "after") + + def test_poke_returns_true_when_partition_event_exists(self, monkeypatch): + comms = mock.Mock() + comms.send.return_value = AssetEventsResult( + asset_events=[ + AssetEventResponse( + id=1, + timestamp=timezone.utcnow(), + asset=AssetResponse(name="orders", uri="s3://warehouse/orders", group="asset"), + partition_key="2024-01-01", + created_dagruns=[], + ) + ], + ) + monkeypatch.setattr(task_runner, "SUPERVISOR_COMMS", comms, raising=False) + + sensor = AssetPartitionSensor( + task_id="wait_orders", + asset=Asset(name="orders", uri="s3://warehouse/orders"), + partition_key="2024-01-01", + ) + + assert sensor.poke({}) is True + comms.send.assert_called_once_with( + GetAssetEventByAsset( + name="orders", + uri="s3://warehouse/orders", + partition_key="2024-01-01", + ascending=False, + limit=1, + ) + ) + + def test_poke_forwards_after_bound(self, monkeypatch): + comms = mock.Mock() + comms.send.return_value = AssetEventsResult(asset_events=[]) + monkeypatch.setattr(task_runner, "SUPERVISOR_COMMS", comms, raising=False) + + after = timezone.datetime(2024, 1, 1) + sensor = AssetPartitionSensor( + task_id="wait_orders", + asset=Asset(name="orders", uri="s3://warehouse/orders"), + partition_key="2024-01-01", + after=after, + ) + + assert sensor.poke({}) is False + comms.send.assert_called_once_with( + GetAssetEventByAsset( + name="orders", + uri="s3://warehouse/orders", + partition_key="2024-01-01", + after=after, + ascending=False, + limit=1, + ) + ) + + def test_poke_returns_false_when_partition_event_is_missing(self, monkeypatch): + comms = mock.Mock() + comms.send.return_value = AssetEventsResult(asset_events=[]) + monkeypatch.setattr(task_runner, "SUPERVISOR_COMMS", comms, raising=False) + + sensor = AssetPartitionSensor( + task_id="wait_orders", + asset=Asset(name="orders", uri="s3://warehouse/orders"), + partition_key="2024-01-01", + ) + + assert sensor.poke({}) is False + + def test_poke_raises_runtime_error_for_supervisor_error(self, monkeypatch): + comms = mock.Mock() + comms.send.return_value = ErrorResponse(error=ErrorType.ASSET_NOT_FOUND) + monkeypatch.setattr(task_runner, "SUPERVISOR_COMMS", comms, raising=False) + + sensor = AssetPartitionSensor( + task_id="wait_orders", + asset=Asset(name="orders", uri="s3://warehouse/orders"), + partition_key="2024-01-01", + ) + + with pytest.raises(AirflowRuntimeError): + sensor.poke({}) + + def test_execute_defers_when_partition_event_is_missing(self, monkeypatch): + comms = mock.Mock() + comms.send.return_value = AssetEventsResult(asset_events=[]) + monkeypatch.setattr(task_runner, "SUPERVISOR_COMMS", comms, raising=False) + + after = timezone.datetime(2024, 1, 1) + sensor = AssetPartitionSensor( + task_id="wait_orders", + asset=Asset(name="orders", uri="s3://warehouse/orders"), + partition_key="2024-01-01", + after=after, + deferrable=True, + ) + + with pytest.raises(TaskDeferred) as exc: + sensor.execute({}) + + assert isinstance(exc.value.trigger, AssetPartitionTrigger) + assert exc.value.trigger.asset_name == "orders" + assert exc.value.trigger.asset_uri == "s3://warehouse/orders" + assert exc.value.trigger.partition_key == "2024-01-01" + assert exc.value.trigger.after == after + + def test_execute_complete_raises_for_trigger_error(self): + sensor = AssetPartitionSensor( + task_id="wait_orders", + asset=Asset(name="orders", uri="s3://warehouse/orders"), + partition_key="2024-01-01", + ) + + with pytest.raises(AirflowException, match="failed"): + sensor.execute_complete({}, {"status": "error", "message": "failed"}) diff --git a/providers/standard/tests/unit/standard/triggers/test_asset.py b/providers/standard/tests/unit/standard/triggers/test_asset.py new file mode 100644 index 0000000000000..e79e6d01d6a18 --- /dev/null +++ b/providers/standard/tests/unit/standard/triggers/test_asset.py @@ -0,0 +1,104 @@ +# +# 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 + +from unittest import mock + +import pytest + +from airflow.providers.standard.triggers.asset import AssetPartitionTrigger +from airflow.sdk import timezone +from airflow.sdk.api.datamodels._generated import AssetEventResponse, AssetResponse +from airflow.sdk.exceptions import ErrorType +from airflow.sdk.execution_time import task_runner +from airflow.sdk.execution_time.comms import AssetEventsResult, ErrorResponse, GetAssetEventByAsset +from airflow.triggers.base import TriggerEvent + + +class TestAssetPartitionTrigger: + def test_serialization(self): + after = timezone.datetime(2024, 1, 1) + trigger = AssetPartitionTrigger( + asset_name="orders", + asset_uri="s3://warehouse/orders", + partition_key="2024-01-01", + after=after, + poke_interval=10, + ) + + classpath, kwargs = trigger.serialize() + + assert classpath == "airflow.providers.standard.triggers.asset.AssetPartitionTrigger" + assert kwargs == { + "asset_name": "orders", + "asset_uri": "s3://warehouse/orders", + "partition_key": "2024-01-01", + "after": after, + "poke_interval": 10, + } + + @pytest.mark.asyncio + async def test_run_yields_success_when_partition_event_exists(self, monkeypatch): + comms = mock.Mock() + comms.asend = mock.AsyncMock( + return_value=AssetEventsResult( + asset_events=[ + AssetEventResponse( + id=1, + timestamp=timezone.utcnow(), + asset=AssetResponse(name="orders", uri="s3://warehouse/orders", group="asset"), + partition_key="2024-01-01", + created_dagruns=[], + ) + ], + ) + ) + monkeypatch.setattr(task_runner, "SUPERVISOR_COMMS", comms, raising=False) + trigger = AssetPartitionTrigger( + asset_name="orders", + asset_uri="s3://warehouse/orders", + partition_key="2024-01-01", + poke_interval=0, + ) + + assert await trigger.run().__anext__() == TriggerEvent({"status": "success"}) + comms.asend.assert_awaited_once_with( + GetAssetEventByAsset( + name="orders", + uri="s3://warehouse/orders", + partition_key="2024-01-01", + ascending=False, + limit=1, + ) + ) + + @pytest.mark.asyncio + async def test_run_yields_error_for_supervisor_error(self, monkeypatch): + comms = mock.Mock() + comms.asend = mock.AsyncMock(return_value=ErrorResponse(error=ErrorType.ASSET_NOT_FOUND)) + monkeypatch.setattr(task_runner, "SUPERVISOR_COMMS", comms, raising=False) + trigger = AssetPartitionTrigger( + asset_name="orders", + asset_uri="s3://warehouse/orders", + partition_key="2024-01-01", + poke_interval=0, + ) + + assert await trigger.run().__anext__() == TriggerEvent( + {"status": "error", "message": "ASSET_NOT_FOUND: None"} + ) diff --git a/task-sdk/src/airflow/sdk/execution_time/request_handlers.py b/task-sdk/src/airflow/sdk/execution_time/request_handlers.py index a31596b3d2b82..e8ecac71ee9d1 100644 --- a/task-sdk/src/airflow/sdk/execution_time/request_handlers.py +++ b/task-sdk/src/airflow/sdk/execution_time/request_handlers.py @@ -41,6 +41,7 @@ XComSequenceSliceResponse, ) from airflow.sdk.execution_time.comms import ( + AssetEventsResult, AssetStateStoreResult, ClearAssetStateStoreByName, ClearAssetStateStoreByUri, @@ -50,6 +51,8 @@ DeleteAssetStateStoreByUri, DeleteVariable, DeleteXCom, + GetAssetEventByAsset, + GetAssetEventByAssetAlias, GetAssetStateStoreByName, GetAssetStateStoreByUri, GetConnection, @@ -217,6 +220,41 @@ def handle_get_dag_run_state(client: Client, msg: GetDagRunState) -> tuple[BaseM return dr_resp, {} +def handle_get_asset_event_by_asset( + client: Client, msg: GetAssetEventByAsset +) -> tuple[BaseModel | None, dict[str, bool]]: + """Fetch asset events for an asset, optionally filtered by partition key.""" + resp = client.asset_events.get( + uri=msg.uri, + name=msg.name, + after=msg.after, + before=msg.before, + ascending=msg.ascending, + limit=msg.limit, + partition_key=msg.partition_key, + partition_key_regexp_pattern=msg.partition_key_regexp_pattern, + extra=msg.extra, + ) + return AssetEventsResult.from_asset_events_response(resp), {"exclude_unset": True} + + +def handle_get_asset_event_by_asset_alias( + client: Client, msg: GetAssetEventByAssetAlias +) -> tuple[BaseModel | None, dict[str, bool]]: + """Fetch asset events for an asset alias, optionally filtered by partition key.""" + resp = client.asset_events.get( + alias_name=msg.alias_name, + after=msg.after, + before=msg.before, + ascending=msg.ascending, + limit=msg.limit, + partition_key=msg.partition_key, + partition_key_regexp_pattern=msg.partition_key_regexp_pattern, + extra=msg.extra, + ) + return AssetEventsResult.from_asset_events_response(resp), {"exclude_unset": True} + + def handle_get_previous_dag_run( client: Client, msg: GetPreviousDagRun ) -> tuple[BaseModel | None, dict[str, bool]]: diff --git a/task-sdk/src/airflow/sdk/execution_time/supervisor.py b/task-sdk/src/airflow/sdk/execution_time/supervisor.py index 87311f02da7a1..c132afdbeb5b3 100644 --- a/task-sdk/src/airflow/sdk/execution_time/supervisor.py +++ b/task-sdk/src/airflow/sdk/execution_time/supervisor.py @@ -61,7 +61,6 @@ from airflow.sdk.exceptions import ErrorType from airflow.sdk.execution_time import comms from airflow.sdk.execution_time.comms import ( - AssetEventsResult, AssetResult, AssetStateStoreResult, AwaitInputTask, @@ -136,6 +135,8 @@ from airflow.sdk.execution_time.request_handlers import ( handle_delete_variable, handle_delete_xcom, + handle_get_asset_event_by_asset, + handle_get_asset_event_by_asset_alias, handle_get_connection, handle_get_dag_run_state, handle_get_dr_count, @@ -1767,34 +1768,9 @@ def _handle_request(self, msg: ToSupervisor, log: FilteringBoundLogger, req_id: elif isinstance(msg, GetAssetsByAlias): resp = self.client.assets.get_by_alias(alias_name=msg.alias_name) elif isinstance(msg, GetAssetEventByAsset): - asset_event_resp = self.client.asset_events.get( - uri=msg.uri, - name=msg.name, - after=msg.after, - before=msg.before, - ascending=msg.ascending, - limit=msg.limit, - partition_key=msg.partition_key, - partition_key_regexp_pattern=msg.partition_key_regexp_pattern, - extra=msg.extra, - ) - asset_event_result = AssetEventsResult.from_asset_events_response(asset_event_resp) - resp = asset_event_result - dump_opts = {"exclude_unset": True} + resp, dump_opts = handle_get_asset_event_by_asset(self.client, msg) elif isinstance(msg, GetAssetEventByAssetAlias): - asset_event_resp = self.client.asset_events.get( - alias_name=msg.alias_name, - after=msg.after, - before=msg.before, - ascending=msg.ascending, - limit=msg.limit, - partition_key=msg.partition_key, - partition_key_regexp_pattern=msg.partition_key_regexp_pattern, - extra=msg.extra, - ) - asset_event_result = AssetEventsResult.from_asset_events_response(asset_event_resp) - resp = asset_event_result - dump_opts = {"exclude_unset": True} + resp, dump_opts = handle_get_asset_event_by_asset_alias(self.client, msg) elif isinstance(msg, GetPrevSuccessfulDagRun): resp, dump_opts = handle_get_prev_successful_dag_run(self.client, self.id) elif isinstance(msg, GetXComCount): diff --git a/task-sdk/tests/task_sdk/execution_time/test_request_handlers.py b/task-sdk/tests/task_sdk/execution_time/test_request_handlers.py index c3999f9baca9a..a4ec5d35c2086 100644 --- a/task-sdk/tests/task_sdk/execution_time/test_request_handlers.py +++ b/task-sdk/tests/task_sdk/execution_time/test_request_handlers.py @@ -20,16 +20,20 @@ import pytest +from airflow.sdk import timezone from airflow.sdk.api import client as sdk_client -from airflow.sdk.api.datamodels._generated import AssetStateStoreResponse +from airflow.sdk.api.datamodels._generated import AssetEventsResponse, AssetStateStoreResponse from airflow.sdk.exceptions import ErrorType from airflow.sdk.execution_time.comms import ( + AssetEventsResult, AssetStateStoreResult, ClearAssetStateStoreByName, ClearAssetStateStoreByUri, DeleteAssetStateStoreByName, DeleteAssetStateStoreByUri, ErrorResponse, + GetAssetEventByAsset, + GetAssetEventByAssetAlias, GetAssetStateStoreByName, GetAssetStateStoreByUri, SetAssetStateStoreByName, @@ -40,6 +44,8 @@ handle_clear_asset_state_store_by_uri, handle_delete_asset_state_store_by_name, handle_delete_asset_state_store_by_uri, + handle_get_asset_event_by_asset, + handle_get_asset_event_by_asset_alias, handle_get_asset_state_store_by_name, handle_get_asset_state_store_by_uri, handle_set_asset_state_store_by_name, @@ -100,6 +106,58 @@ def test_get_asset_state_store_by_uri_passes_through_error_response(client): assert dump_opts == {} +def test_get_asset_event_by_asset_delegates_and_wraps(client): + client.asset_events.get.return_value = AssetEventsResponse(asset_events=[]) + + result, dump_opts = handle_get_asset_event_by_asset( + client, + GetAssetEventByAsset( + name="orders", + uri="s3://warehouse/orders", + partition_key="2024-01-01", + after=timezone.datetime(2024, 1, 1), + limit=1, + ascending=False, + ), + ) + + client.asset_events.get.assert_called_once_with( + uri="s3://warehouse/orders", + name="orders", + after=timezone.datetime(2024, 1, 1), + before=None, + ascending=False, + limit=1, + partition_key="2024-01-01", + partition_key_regexp_pattern=None, + extra=None, + ) + assert isinstance(result, AssetEventsResult) + assert dump_opts == {"exclude_unset": True} + + +def test_get_asset_event_by_asset_alias_delegates_and_wraps(client): + client.asset_events.get.return_value = AssetEventsResponse(asset_events=[]) + + result, dump_opts = handle_get_asset_event_by_asset_alias( + client, + GetAssetEventByAssetAlias(alias_name="orders_alias", partition_key="2024-01-01"), + ) + + client.asset_events.get.assert_called_once_with( + alias_name="orders_alias", + after=None, + before=None, + ascending=True, + limit=None, + partition_key="2024-01-01", + partition_key_regexp_pattern=None, + extra=None, + ) + assert isinstance(result, AssetEventsResult) + assert dump_opts == {"exclude_unset": True} + + @pytest.mark.parametrize( ("handler", "msg", "call_kwargs", "method"), [ From 03b434db16bf48975c87a83db8572ae10deaff51 Mon Sep 17 00:00:00 2001 From: Leondon9 Date: Wed, 3 Jun 2026 16:45:25 +0900 Subject: [PATCH 2/2] Add newsfragment for asset partition sensor --- airflow-core/newsfragments/67941.feature.rst | 1 + 1 file changed, 1 insertion(+) create mode 100644 airflow-core/newsfragments/67941.feature.rst diff --git a/airflow-core/newsfragments/67941.feature.rst b/airflow-core/newsfragments/67941.feature.rst new file mode 100644 index 0000000000000..37aa553a4c353 --- /dev/null +++ b/airflow-core/newsfragments/67941.feature.rst @@ -0,0 +1 @@ +Add an asset partition sensor for waiting on a specific asset event partition.