Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
22 commits
Select commit Hold shift + click to select a range
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
11 changes: 11 additions & 0 deletions airflow-core/src/airflow/assets/manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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)
Expand Down
3 changes: 3 additions & 0 deletions airflow-core/src/airflow/serialization/definitions/dag.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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",
Expand Down
1 change: 1 addition & 0 deletions airflow-core/src/airflow/serialization/schema.json
Original file line number Diff line number Diff line change
Expand Up @@ -168,6 +168,7 @@
"tasks": { "$ref": "#/definitions/tasks" },
"timezone": { "$ref": "#/definitions/timezone" },
"owner_links": { "type": "object" },
"outlets": { "type": "array", "default": [] },
"timetable": {
"type": "object",
"properties": {
Expand Down
119 changes: 105 additions & 14 deletions task-sdk/src/airflow/sdk/definitions/dag.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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").
Expand Down Expand Up @@ -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(
Expand Down Expand Up @@ -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):
"""
Expand Down Expand Up @@ -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

Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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,
Expand Down
63 changes: 62 additions & 1 deletion task-sdk/tests/task_sdk/definitions/test_dag.py
Original file line number Diff line number Diff line change
Expand Up @@ -16,6 +16,7 @@
# under the License.
from __future__ import annotations

import operator
import re
import warnings
import weakref
Expand All @@ -37,14 +38,74 @@
)
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

DEFAULT_DATE = datetime(2016, 1, 1, tzinfo=timezone.utc)


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"),
[
Expand Down
Loading