diff --git a/airflow-core/src/airflow/assets/manager.py b/airflow-core/src/airflow/assets/manager.py index c3491574e9148..d7b0fda97e055 100644 --- a/airflow-core/src/airflow/assets/manager.py +++ b/airflow-core/src/airflow/assets/manager.py @@ -274,6 +274,10 @@ def register_asset_change( asset: SerializedAsset | AssetModel | SerializedAssetUniqueKey, extra=None, source_alias_names: Collection[str] = (), + source_dag_id: str | None = None, + source_task_id: str | None = None, + source_run_id: str | None = None, + source_map_index: int | None = None, session: Session, partition_key: str | None = None, partition_date: datetime | None = None, @@ -341,6 +345,13 @@ def register_asset_change( source_run_id=task_instance.run_id, source_map_index=task_instance.map_index, ) + elif source_dag_id is not None or source_run_id is not None or source_task_id is not None: + event_kwargs.update( + source_task_id=source_task_id, + source_dag_id=source_dag_id, + source_run_id=source_run_id, + source_map_index=source_map_index, + ) asset_event = AssetEvent(**event_kwargs) session.add(asset_event) diff --git a/airflow-core/src/airflow/serialization/definitions/dag.py b/airflow-core/src/airflow/serialization/definitions/dag.py index 8ed0fee2ccabd..c5eedd7bce2a2 100644 --- a/airflow-core/src/airflow/serialization/definitions/dag.py +++ b/airflow-core/src/airflow/serialization/definitions/dag.py @@ -126,7 +126,9 @@ class SerializedDAG: max_active_runs: int = 16 max_active_tasks: int = 16 max_consecutive_failed_dag_runs: int = 0 + inlets: Sequence[Any] = attrs.field(factory=list) owner_links: dict[str, str] = attrs.field(factory=dict) + outlets: Sequence[Any] = attrs.field(factory=list) params: SerializedParamsDict = attrs.field(factory=SerializedParamsDict) partial: bool = False render_template_as_native_obj: bool = False @@ -175,6 +177,7 @@ def get_serialized_fields(cls) -> frozenset[str]: "max_active_runs", "max_active_tasks", "max_consecutive_failed_dag_runs", + "outlets", "owner_links", "relative_fileloc", "render_template_as_native_obj", diff --git a/airflow-core/src/airflow/serialization/schema.json b/airflow-core/src/airflow/serialization/schema.json index 872c3a1331ee3..288b01dd8f03b 100644 --- a/airflow-core/src/airflow/serialization/schema.json +++ b/airflow-core/src/airflow/serialization/schema.json @@ -168,6 +168,7 @@ "tasks": { "$ref": "#/definitions/tasks" }, "timezone": { "$ref": "#/definitions/timezone" }, "owner_links": { "type": "object" }, + "outlets": { "type": "array", "default": [] }, "timetable": { "type": "object", "properties": { diff --git a/task-sdk/src/airflow/sdk/definitions/dag.py b/task-sdk/src/airflow/sdk/definitions/dag.py index 5ba0254b208b0..726ca5ed443ed 100644 --- a/task-sdk/src/airflow/sdk/definitions/dag.py +++ b/task-sdk/src/airflow/sdk/definitions/dag.py @@ -45,7 +45,7 @@ from airflow.sdk.bases.timetable import BaseTimetable from airflow.sdk.definitions._internal.node import DAGNode, validate_key from airflow.sdk.definitions._internal.types import NOTSET, ArgNotSet, is_arg_set -from airflow.sdk.definitions.asset import AssetAll, BaseAsset +from airflow.sdk.definitions.asset import Asset, AssetAll, BaseAsset from airflow.sdk.definitions.context import Context from airflow.sdk.definitions.deadline import DeadlineAlert from airflow.sdk.definitions.param import DagParam, ParamsDict @@ -234,6 +234,60 @@ def _convert_deadline(deadline: list[DeadlineAlert] | DeadlineAlert | None) -> l return list(deadline) +def _collect_from_input(value_or_values: Any | Collection[Any] | None) -> list[Any]: + if not value_or_values: + return [] + if isinstance(value_or_values, Collection) and not isinstance(value_or_values, str): + return list(value_or_values) + return [value_or_values] + + +def _emit_dag_asset_events(*, context: Context) -> None: + """Emit asset events for Dag-level outlets on successful Dag run.""" + dag = context.get("dag") + if dag is None: + return + + from airflow.assets.manager import asset_manager + from airflow.models.asset import AssetModel + from airflow.serialization.definitions.assets import SerializedAsset + from airflow.utils.session import create_session + + run_id = context.get("run_id") + with create_session() as session: + for outlet in dag.outlets: + if not isinstance(outlet, BaseAsset): + continue + if not isinstance(outlet, Asset): + # Dag-level outlets only emit concrete assets as events. + continue + + serialized_asset = SerializedAsset( + name=outlet.name, + uri=outlet.uri, + group=outlet.group, + extra=outlet.extra, + watchers=[], + ) + event = asset_manager.register_asset_change( + asset=serialized_asset, + source_dag_id=dag.dag_id, + source_run_id=run_id, + partition_key=getattr(context.get("dag_run"), "partition_key", None), + session=session, + ) + if event is None: + session.add(AssetModel.from_serialized(serialized_asset)) + session.flush() + asset_manager.register_asset_change( + asset=serialized_asset, + source_dag_id=dag.dag_id, + source_run_id=run_id, + partition_key=getattr(context.get("dag_run"), "partition_key", None), + session=session, + ) + + def _convert_doc_md(doc_md: str | None) -> str | None: if doc_md is None: return doc_md @@ -418,6 +472,7 @@ class DAG: :param owner_links: Dict of owners and their links, that will be clickable on the Dags view UI. Can be used as an HTTP link (for example the link to your Slack channel), or a mailto link. e.g: ``{"dag_owner": "https://airflow.apache.org/"}`` + :param outlets: List of outlets that the Dag should emit when the Dag run is successful. :param auto_register: Automatically register this DAG when it is used in a ``with`` block :param fail_fast: Fails currently running tasks when task in Dag fails. **Warning**: A fail stop dag can only have tasks with the default trigger rule ("all_success"). @@ -532,6 +587,7 @@ def __rich_repr__(self): render_template_as_native_obj: bool = attrs.field(default=False, converter=bool) tags: MutableSet[str] = attrs.field(factory=set, converter=_convert_tags) owner_links: dict[str, str] = attrs.field(factory=dict) + outlets: list[Any] = attrs.field(factory=list, converter=_collect_from_input) auto_register: bool = attrs.field(default=True, converter=bool) fail_fast: bool = attrs.field(default=False, converter=bool) allowed_run_types: DagRunType | Collection[DagRunType] | None = attrs.field( @@ -600,6 +656,12 @@ def __attrs_post_init__(self): f"requires max_active_runs <= {active_runs_limit}" ) + if self.outlets: + callbacks = _collect_from_input(self.on_success_callback) + callbacks.append(_emit_dag_asset_events) + self.on_success_callback = callbacks + self.has_on_success_callback = True + @params.validator def _validate_params(self, _, params: ParamsDict): """ @@ -745,6 +807,28 @@ def __hash__(self): hash_components.append(repr(val)) return hash(tuple(hash_components)) + def __gt__(self, other): + """ + Return [Dag] > [Outlet]. + + If other is an attr annotated object it is set as an outlet of this Dag. + """ + if isinstance(other, str) or not isinstance(other, Iterable): + other = [other] + else: + other = list(other) + + for obj in other: + if not attrs.has(obj): + raise TypeError(f"Left hand side ({obj}) is not an outlet") + self.add_outlets(other) + + return self + + def add_outlets(self, outlets: Iterable[Any]) -> None: + """Define the outlets of this Dag.""" + self.outlets.extend(outlets) + def __enter__(self) -> Self: from airflow.sdk.definitions._internal.contextmanager import DagContext @@ -1282,25 +1366,31 @@ def test( scheduler_dag = DagSerialization.deserialize_dag(DagSerialization.serialize_dag(self)) # Allow users to explicitly pass None. If it isn't set, we default to current time. - logical_date = logical_date if is_arg_set(logical_date) else timezone.utcnow() + logical_date_val: datetime | None = ( + logical_date if is_arg_set(logical_date) else timezone.utcnow() + ) - log.debug("Clearing existing task instances for logical date %s", logical_date) + log.debug("Clearing existing task instances for logical date %s", logical_date_val) # TODO: Replace with calling client.dag_run.clear in Execution API at some point SerializedDAG.clear_dags( dags=[scheduler_dag], - start_date=logical_date, - end_date=logical_date, + start_date=logical_date_val, + end_date=logical_date_val, dag_run_state=False, ) log.debug("Getting dagrun for dag %s", self.dag_id) - logical_date = timezone.coerce_datetime(logical_date) - run_after = timezone.coerce_datetime(run_after) or timezone.coerce_datetime(timezone.utcnow()) - if logical_date is None: + logical_date_val = timezone.coerce_datetime(logical_date_val) + run_after_val: datetime = timezone.coerce_datetime(run_after) or timezone.coerce_datetime( + timezone.utcnow() + ) + if logical_date_val is None: data_interval: DataInterval | None = None else: timetable = coerce_to_core_timetable(self.timetable) - data_interval = timetable.infer_manual_data_interval(run_after=logical_date) + # logical_date_val is not None here, but mypy might still be unsure about its type + # We cast to Any because the core Timetable expects pendulum.DateTime which is a subclass of datetime + data_interval = timetable.infer_manual_data_interval(run_after=cast(Any, logical_date_val)) # These imports are intentionally lazy: this Task SDK module must not # pull in airflow-core at import time (worker isolation). from airflow.dag_processing.bundles.manager import DagBundlesManager @@ -1363,14 +1453,14 @@ def test( dr: DagRun = get_or_create_dagrun( dag=scheduler_dag, - start_date=logical_date or run_after, - logical_date=logical_date, + start_date=logical_date_val or run_after_val, + logical_date=logical_date_val, data_interval=data_interval, - run_after=run_after, + run_after=run_after_val, run_id=DagRun.generate_run_id( run_type=DagRunType.MANUAL, - logical_date=logical_date, - run_after=run_after, + logical_date=logical_date_val, + run_after=run_after_val, ), session=session, conf=run_conf, @@ -1631,6 +1721,7 @@ def dag( render_template_as_native_obj: bool = False, tags: Collection[str] | None = None, owner_links: dict[str, str] | None = None, + outlets: Any | None = None, auto_register: bool = True, fail_fast: bool = False, allowed_run_types: DagRunType | Collection[DagRunType] | None = None, diff --git a/task-sdk/tests/task_sdk/definitions/test_dag.py b/task-sdk/tests/task_sdk/definitions/test_dag.py index 9b76816886c76..d46f92dc55c62 100644 --- a/task-sdk/tests/task_sdk/definitions/test_dag.py +++ b/task-sdk/tests/task_sdk/definitions/test_dag.py @@ -16,6 +16,7 @@ # under the License. from __future__ import annotations +import operator import re import warnings import weakref @@ -37,7 +38,10 @@ ) from airflow.sdk.bases.operator import BaseOperator from airflow.sdk.bases.timetable import BaseTimetable -from airflow.sdk.definitions.param import DagParam, ParamsDict +from airflow.sdk.definitions.asset import Asset +from airflow.sdk.definitions.dag import DAG, dag as dag_decorator +from airflow.sdk.definitions.param import DagParam, Param, ParamsDict +from airflow.sdk.definitions.timetables import assets, events, interval, simple, trigger # noqa: F401 from airflow.sdk.exceptions import AirflowDagCycleException, DuplicateTaskIdFound, RemovedInAirflow4Warning from airflow.utils.types import DagRunType @@ -45,6 +49,63 @@ class TestDag: + def test_dag_outlets_init_from_single_value(self): + outlet = Asset("asset://dag_outlet") + dag = DAG("dag-with-outlet", schedule=None, outlets=outlet) + + assert dag.outlets == [outlet] + + def test_dag_outlets_init_from_collection(self): + outlet_1 = Asset("asset://dag_outlet_1") + outlet_2 = Asset("asset://dag_outlet_2") + dag = DAG("dag-with-outlets", schedule=None, outlets=[outlet_1, outlet_2]) + + assert dag.outlets == [outlet_1, outlet_2] + + def test_dag_gt_sets_outlets(self): + dag = DAG("dag-gt-outlet", schedule=None) + outlet = Asset("asset://dag_outlet") + + result = dag > outlet + + assert result is dag + assert dag.outlets == [outlet] + + def test_dag_gt_rejects_non_outlets(self): + dag = DAG("dag-invalid-outlet", schedule=None) + + with pytest.raises(TypeError, match=r"Left hand side \(not-an-outlet\) is not an outlet"): + operator.gt(dag, "not-an-outlet") + + def test_dag_outlets_register_success_callback(self): + dag = DAG("dag-callback-outlet", schedule=None, outlets=Asset("asset://dag_outlet")) + + assert dag.has_on_success_callback + callbacks = ( + dag.on_success_callback + if isinstance(dag.on_success_callback, list) + else [dag.on_success_callback] + ) + assert any(getattr(cb, "__name__", "") == "_emit_dag_asset_events" for cb in callbacks) + + def test_dag_outlets_preserve_existing_success_callback(self): + def callback(context): + pass + + dag = DAG( + "dag-callback-merge", + schedule=None, + outlets=Asset("asset://dag_outlet"), + on_success_callback=callback, + ) + callbacks = ( + dag.on_success_callback + if isinstance(dag.on_success_callback, list) + else [dag.on_success_callback] + ) + assert callback in callbacks + assert any(getattr(cb, "__name__", "") == "_emit_dag_asset_events" for cb in callbacks) + @pytest.mark.parametrize( ("dag_id", "exc_type", "exc_value"), [