Skip to content
Closed
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
341 changes: 340 additions & 1 deletion airflow-core/tests/unit/jobs/test_scheduler_job.py
Original file line number Diff line number Diff line change
Expand Up @@ -172,7 +172,7 @@
)
from tests_common.test_utils.mock_executor import MockExecutor
from tests_common.test_utils.mock_operators import CustomOperator
from tests_common.test_utils.taskinstance import create_task_instance, run_task_instance
from tests_common.test_utils.taskinstance import create_task_instance, get_template_context, run_task_instance
from unit.listeners import dag_listener
from unit.models import TEST_DAGS_FOLDER

Expand Down Expand Up @@ -5922,6 +5922,345 @@ def _make_event(timestamp):
== []
)

def _tick_asset_dagruns(self, session):
scheduler_job = Job()
self.job_runner = SchedulerJobRunner(job=scheduler_job, executors=[self.null_exec])
self.job_runner._create_dagruns_for_dags(session, session)

@staticmethod
def _consumed_by_timestamp_then_id(dag_run):
# Association table has no order_by; consume selects timestamp then id.
return sorted(dag_run.consumed_asset_events, key=lambda e: (e.timestamp, e.id))

@staticmethod
def _triggering_asset_events(dag_run, task, asset, session):
# Access the relationship first so model_validate includes consumed events instead of [].
_ = dag_run.consumed_asset_events
ti = dag_run.get_task_instance(task.task_id, session=session)
context = get_template_context(ti, task, session=session)
return list(context["triggering_asset_events"][asset])

@pytest.mark.need_serialized_dag
def test_mapped_outlet_asset_events_consumed_in_one_tick(self, session, dag_maker):
"""Every mapped outlet event is consumed in the next asset-triggered run.

When N mapped TIs of the same producer succeed with the same unpartitioned
outlet, one scheduler tick must attach all N events (timestamp then id
order) to a single consumer DagRun and clear the queue.
"""
asset = Asset(uri="test://mapped-outlet-one-tick", name="mapped_outlet_one_tick")

with dag_maker(
dag_id="mapped-outlet-one-tick-consumer",
schedule=[asset],
max_active_runs=16,
session=session,
):
EmptyOperator(task_id="consume")
consumer_dag = dag_maker.dag
consumer_dag_id = consumer_dag.dag_id
consume_task = consumer_dag.get_task("consume")

with dag_maker(dag_id="mapped-outlet-one-tick-producer", schedule=None, session=session):

@task(outlets=[asset])
def produce(item):
return item

produce.expand(item=[1, 2, 3])

producer_run = dag_maker.create_dagrun(run_id="producer-run")
for map_index in range(3):
dag_maker.run_ti("produce", dag_run=producer_run, map_index=map_index, session=session)
session.commit()

events = session.scalars(
select(AssetEvent).where(AssetEvent.source_dag_id == producer_run.dag_id)
).all()
assert {e.source_map_index for e in events} == {0, 1, 2}
assert len(events) == 3
assert (
len(
session.scalars(
select(AssetDagRunQueue).where(AssetDagRunQueue.target_dag_id == consumer_dag_id)
).all()
)
== 3
)

self._tick_asset_dagruns(session)

created_run = session.scalars(select(DagRun).where(DagRun.dag_id == consumer_dag_id)).one()
assert created_run.state == State.QUEUED
consumed = self._consumed_by_timestamp_then_id(created_run)
assert [e.source_map_index for e in consumed] == [0, 1, 2]
trig = self._triggering_asset_events(created_run, consume_task, asset, session)
assert {e.source_map_index for e in trig} == {0, 1, 2}
assert (
session.scalars(
select(AssetDagRunQueue).where(AssetDagRunQueue.target_dag_id == consumer_dag_id)
).all()
== []
)

@pytest.mark.need_serialized_dag
def test_mapped_outlet_asset_events_consumed_across_staggered_ticks(self, session, dag_maker):
"""Events that arrive after a tick are consumed by a later run.

map_index 0 is consumed alone on the first tick; 1 and 2 batch into the
next run. max_active_runs is high enough that the second run is created
while the first is still queued.
"""
asset = Asset(uri="test://mapped-outlet-staggered", name="mapped_outlet_staggered")

with dag_maker(
dag_id="mapped-outlet-staggered-consumer",
schedule=[asset],
max_active_runs=16,
session=session,
):
EmptyOperator(task_id="consume")
consumer_dag = dag_maker.dag
consumer_dag_id = consumer_dag.dag_id
consume_task = consumer_dag.get_task("consume")

with dag_maker(dag_id="mapped-outlet-staggered-producer", schedule=None, session=session):

@task(outlets=[asset])
def produce(item):
return item

produce.expand(item=[1, 2, 3])

producer_run = dag_maker.create_dagrun(run_id="producer-run")
dag_maker.run_ti("produce", dag_run=producer_run, map_index=0, session=session)
session.commit()

self._tick_asset_dagruns(session)

first_run = session.scalars(select(DagRun).where(DagRun.dag_id == consumer_dag_id)).one()
assert [e.source_map_index for e in self._consumed_by_timestamp_then_id(first_run)] == [0]
trig = self._triggering_asset_events(first_run, consume_task, asset, session)
assert {e.source_map_index for e in trig} == {0}
assert (
session.scalars(
select(AssetDagRunQueue).where(AssetDagRunQueue.target_dag_id == consumer_dag_id)
).all()
== []
)

for map_index in (1, 2):
dag_maker.run_ti("produce", dag_run=producer_run, map_index=map_index, session=session)
session.commit()

self._tick_asset_dagruns(session)

runs = session.scalars(
select(DagRun).where(DagRun.dag_id == consumer_dag_id).order_by(DagRun.id)
).all()
assert len(runs) == 2
second_run = runs[1]
consumed = self._consumed_by_timestamp_then_id(second_run)
assert [e.source_map_index for e in consumed] == [1, 2]
trig = self._triggering_asset_events(second_run, consume_task, asset, session)
assert {e.source_map_index for e in trig} == {1, 2}
assert (
session.scalars(
select(AssetDagRunQueue).where(AssetDagRunQueue.target_dag_id == consumer_dag_id)
).all()
== []
)

@pytest.mark.need_serialized_dag
def test_mapped_outlet_asset_events_same_timestamp_are_all_consumed(self, session, dag_maker):
"""Events that share a timestamp are still distinct by id and all consumed."""
asset = Asset(uri="test://mapped-outlet-same-ts", name="mapped_outlet_same_ts")

with dag_maker(
dag_id="mapped-outlet-same-ts-consumer",
schedule=[asset],
max_active_runs=16,
session=session,
):
EmptyOperator(task_id="consume")
consumer_dag = dag_maker.dag
consumer_dag_id = consumer_dag.dag_id
consume_task = consumer_dag.get_task("consume")

with dag_maker(dag_id="mapped-outlet-same-ts-producer", schedule=None, session=session):

@task(outlets=[asset])
def produce(item):
return item

produce.expand(item=[1, 2, 3])

producer_run = dag_maker.create_dagrun(run_id="producer-run")
with time_machine.travel(DEFAULT_DATE, tick=False):
for map_index in range(3):
dag_maker.run_ti("produce", dag_run=producer_run, map_index=map_index, session=session)
session.commit()

events = session.scalars(
select(AssetEvent).where(AssetEvent.source_dag_id == producer_run.dag_id)
).all()
assert {e.source_map_index for e in events} == {0, 1, 2}
assert len(events) == 3
assert len({e.timestamp for e in events}) == 1
assert (
len(
session.scalars(
select(AssetDagRunQueue).where(AssetDagRunQueue.target_dag_id == consumer_dag_id)
).all()
)
== 3
)

self._tick_asset_dagruns(session)

created_run = session.scalars(select(DagRun).where(DagRun.dag_id == consumer_dag_id)).one()
consumed = self._consumed_by_timestamp_then_id(created_run)
assert [e.source_map_index for e in consumed] == [0, 1, 2]
assert [e.id for e in consumed] == sorted(e.id for e in events)
assert len({e.timestamp for e in consumed}) == 1
trig = self._triggering_asset_events(created_run, consume_task, asset, session)
assert {e.source_map_index for e in trig} == {0, 1, 2}
assert (
session.scalars(
select(AssetDagRunQueue).where(AssetDagRunQueue.target_dag_id == consumer_dag_id)
).all()
== []
)

@pytest.mark.need_serialized_dag
def test_mapped_outlet_asset_events_and_condition_waits_for_all_assets(self, session, dag_maker):
"""Asset AND holds the mapped queue until every required asset has an event."""
asset_a = Asset(uri="test://mapped-outlet-and-a", name="mapped_outlet_and_a")
asset_b = Asset(uri="test://mapped-outlet-and-b", name="mapped_outlet_and_b")

with dag_maker(
dag_id="mapped-outlet-and-consumer",
schedule=[asset_a, asset_b],
max_active_runs=16,
session=session,
):
EmptyOperator(task_id="consume")
consumer_dag = dag_maker.dag
consumer_dag_id = consumer_dag.dag_id
consume_task = consumer_dag.get_task("consume")

with dag_maker(dag_id="mapped-outlet-and-producer-a", schedule=None, session=session):

@task(outlets=[asset_a])
def produce_a(item):
return item

produce_a.expand(item=[1, 2, 3])

producer_a_run = dag_maker.create_dagrun(run_id="producer-a-run")
for map_index in range(3):
dag_maker.run_ti("produce_a", dag_run=producer_a_run, map_index=map_index, session=session)
session.commit()

self._tick_asset_dagruns(session)

assert session.scalars(select(DagRun).where(DagRun.dag_id == consumer_dag_id)).one_or_none() is None
assert (
len(
session.scalars(
select(AssetDagRunQueue).where(AssetDagRunQueue.target_dag_id == consumer_dag_id)
).all()
)
== 3
)

with dag_maker(dag_id="mapped-outlet-and-producer-b", schedule=None, session=session):

@task(outlets=[asset_b])
def produce_b():
return 1

produce_b()

producer_b_run = dag_maker.create_dagrun(run_id="producer-b-run")
dag_maker.run_ti("produce_b", dag_run=producer_b_run, session=session)
session.commit()

self._tick_asset_dagruns(session)

created_run = session.scalars(select(DagRun).where(DagRun.dag_id == consumer_dag_id)).one()
consumed = list(created_run.consumed_asset_events)
asset_a_id = session.scalar(select(AssetModel.id).where(AssetModel.uri == asset_a.uri))
asset_b_id = session.scalar(select(AssetModel.id).where(AssetModel.uri == asset_b.uri))
assert len(consumed) == 4
assert {e.source_map_index for e in consumed if e.asset_id == asset_a_id} == {0, 1, 2}
assert sum(1 for e in consumed if e.asset_id == asset_b_id) == 1
trig_a = self._triggering_asset_events(created_run, consume_task, asset_a, session)
assert {e.source_map_index for e in trig_a} == {0, 1, 2}
trig_b = self._triggering_asset_events(created_run, consume_task, asset_b, session)
assert len(trig_b) == 1
assert (
session.scalars(
select(AssetDagRunQueue).where(AssetDagRunQueue.target_dag_id == consumer_dag_id)
).all()
== []
)

@pytest.mark.need_serialized_dag
def test_mapped_outlet_asset_alias_events_are_all_consumed(self, session, dag_maker):
"""A mapped alias outlet yields three events that the alias consumer consumes together."""
asset = Asset(uri="test://mapped-outlet-alias-asset", name="mapped_outlet_alias_asset")
alias = AssetAlias(name="mapped_outlet_alias")
asm = AssetModel(name=asset.name, uri=asset.uri, group=asset.group)
session.add_all([asm, AssetActive.for_asset(asm)])
session.commit()

with dag_maker(
dag_id="mapped-outlet-alias-consumer",
schedule=[alias],
max_active_runs=16,
session=session,
):
EmptyOperator(task_id="consume")
consumer_dag = dag_maker.dag
consumer_dag_id = consumer_dag.dag_id
consume_task = consumer_dag.get_task("consume")

with dag_maker(dag_id="mapped-outlet-alias-producer", schedule=None, session=session):

@task(outlets=[alias])
def produce(item, *, outlet_events):
outlet_events[alias].add(asset)
return item

produce.expand(item=[1, 2, 3])

producer_run = dag_maker.create_dagrun(run_id="producer-run")
for map_index in range(3):
dag_maker.run_ti("produce", dag_run=producer_run, map_index=map_index, session=session)
session.commit()

events = session.scalars(
select(AssetEvent).where(AssetEvent.source_dag_id == producer_run.dag_id)
).all()
assert {e.source_map_index for e in events} == {0, 1, 2}
assert len(events) == 3

self._tick_asset_dagruns(session)

created_run = session.scalars(select(DagRun).where(DagRun.dag_id == consumer_dag_id)).one()
consumed = list(created_run.consumed_asset_events)
assert {e.source_map_index for e in consumed} == {0, 1, 2}
assert len(consumed) == 3
trig = self._triggering_asset_events(created_run, consume_task, alias, session)
assert {e.source_map_index for e in trig} == {0, 1, 2}
assert (
session.scalars(
select(AssetDagRunQueue).where(AssetDagRunQueue.target_dag_id == consumer_dag_id)
).all()
== []
)

@pytest.mark.need_serialized_dag
def test_create_dag_runs_asset_triggered_skips_stale_triggered_date(self, session, dag_maker):
asset = Asset(uri="test://asset-for-stale-trigger-date", name="asset-for-stale-trigger-date")
Expand Down
Loading