Skip to content
Open
Show file tree
Hide file tree
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
101 changes: 75 additions & 26 deletions airflow-core/src/airflow/jobs/scheduler_job_runner.py
Original file line number Diff line number Diff line change
Expand Up @@ -43,7 +43,9 @@
delete,
exists,
func,
insert,
inspect,
literal,
or_,
select,
text,
Expand Down Expand Up @@ -138,6 +140,7 @@
from sqlalchemy.engine import CursorResult
from sqlalchemy.orm import Session
from sqlalchemy.orm.interfaces import LoaderOption
from sqlalchemy.sql import Select
from sqlalchemy.sql.elements import ColumnElement
from sqlalchemy.sql.selectable import Subquery

Expand All @@ -163,6 +166,23 @@
MAX_PARTITION_DAG_RUNS_PER_LOOP = 500


def _associate_asset_events_with_dag_run(
*,
dag_run: DagRun,
asset_event_ids_select: Select[tuple[int]],
session: Session,
) -> None:
"""Associate selected asset events without materializing ORM objects."""
selected_event_ids = asset_event_ids_select.subquery()
session.execute(
insert(association_table).from_select(
["dag_run_id", "event_id"],
select(literal(dag_run.id), selected_event_ids.c.event_id),
)
)
session.expire(dag_run, ["consumed_asset_events"])

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Worth a comment explaining why manual expiration is needed.



def _eager_load_dag_run_for_validation() -> tuple[LoaderOption, LoaderOption]:
"""
Eager-load DagRun relations required for execution API datamodel validation.
Expand Down Expand Up @@ -2425,13 +2445,15 @@ def _create_dagruns_for_partitioned_asset_dags(self, session: Session) -> set[st
creating_job_id=self.job.id,
session=session,
)
asset_events = session.scalars(
select(AssetEvent).where(
PartitionedAssetKeyLog.asset_partition_dag_run_id == apdr.id,
PartitionedAssetKeyLog.asset_event_id == AssetEvent.id,
)
asset_event_ids_select = select(AssetEvent.id.label("event_id")).where(
PartitionedAssetKeyLog.asset_partition_dag_run_id == apdr.id,
PartitionedAssetKeyLog.asset_event_id == AssetEvent.id,
)
_associate_asset_events_with_dag_run(
dag_run=dag_run,
asset_event_ids_select=asset_event_ids_select,
session=session,
)
dag_run.consumed_asset_events.extend(asset_events)
session.flush()
apdr.created_dag_run_id = dag_run.id
session.flush()
Expand Down Expand Up @@ -2688,26 +2710,33 @@ def _create_dag_runs_asset_triggered(
)
),
)
asset_events = list(
session.scalars(
select(AssetEvent)
.where(
event_predicate,
~(
select(association_table.c.event_id)
.join(DagRun, DagRun.id == association_table.c.dag_run_id)
.where(
DagRun.dag_id == dag.dag_id,
association_table.c.event_id == AssetEvent.id,
)
.exists()
),
)
.order_by(AssetEvent.timestamp.asc(), AssetEvent.id.asc())
selected_asset_events = (
select(
AssetEvent.id.label("event_id"),
AssetEvent.timestamp.label("timestamp"),
)
.where(
event_predicate,
~(
select(association_table.c.event_id)
.join(DagRun, DagRun.id == association_table.c.dag_run_id)
.where(
DagRun.dag_id == dag.dag_id,
association_table.c.event_id == AssetEvent.id,
)
.exists()
),
)
.cte()
)
if asset_events:
triggered_date = timezone.coerce_datetime(max(event.timestamp for event in asset_events))
triggered_date, max_selected_event_id = session.execute(
select(
func.max(selected_asset_events.c.timestamp),
func.max(selected_asset_events.c.event_id),
)
).one()
if triggered_date is not None and max_selected_event_id is not None:
triggered_date = timezone.coerce_datetime(triggered_date)
self.log.debug(
"Creating asset-triggered DagRun for '%s': %d queued assets, triggered_date=%s",
dag.dag_id,
Expand All @@ -2733,12 +2762,32 @@ def _create_dag_runs_asset_triggered(
else None
)
stats.incr("asset.triggered_dagruns", tags=prune_dict({"team_name": team_name}))
dag_run.consumed_asset_events.extend(asset_events)
_associate_asset_events_with_dag_run(
dag_run=dag_run,
asset_event_ids_select=select(selected_asset_events.c.event_id)
.where(
selected_asset_events.c.timestamp <= triggered_date,
selected_asset_events.c.event_id <= max_selected_event_id,
Comment on lines +2769 to +2770

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

This looks like dead weight unless the reader understands the CTE is re-evaluated fresh in this second, independent query and could pick up events inserted after the aggregate read. Add a comment explaining this is a snapshot bound against a race, not a no-op.

)
.order_by(
selected_asset_events.c.timestamp.asc(),
selected_asset_events.c.event_id.asc(),
),
session=session,
)
consumed_asset_event_count = (
session.scalar(
select(func.count())
.select_from(association_table)
.where(association_table.c.dag_run_id == dag_run.id)
)
or 0
)
Comment on lines +2778 to +2785

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

We can probably use rowcount instead? If we must use a separate query, it’d be a good idea to at least add a guard on log level (with isEnabledFor)

self.log.info(
"Created asset-triggered DagRun for '%s': run_id=%s, consumed %d asset events",
dag.dag_id,
dag_run.run_id,
len(asset_events),
consumed_asset_event_count,
)
else:
self.log.info(
Expand Down
136 changes: 132 additions & 4 deletions airflow-core/tests/unit/jobs/test_scheduler_job.py
Original file line number Diff line number Diff line change
Expand Up @@ -38,7 +38,7 @@
import psutil
import pytest
import time_machine
from sqlalchemy import delete, func, inspect, select, update
from sqlalchemy import delete, event as sqlalchemy_event, func, inspect, select, update
from sqlalchemy.dialects import mysql
from sqlalchemy.orm import joinedload

Expand All @@ -64,7 +64,7 @@
from airflow.executors.executor_utils import ExecutorName
from airflow.executors.local_executor import LocalExecutor
from airflow.jobs.job import Job, run_job
from airflow.jobs.scheduler_job_runner import SchedulerJobRunner
from airflow.jobs.scheduler_job_runner import SchedulerJobRunner, _associate_asset_events_with_dag_run
from airflow.models.asset import (
AssetActive,
AssetAliasModel,
Expand Down Expand Up @@ -196,6 +196,29 @@
/ "standard"
/ "example_dags"
)


@contextmanager
def assert_no_asset_event_materialization(session: Session) -> Generator[None, None, None]:
"""Assert a scheduler operation does not select full AssetEvent entities."""
materializing_queries = []

def track_orm_execute(orm_execute_state):
if orm_execute_state.is_select and any(
description.get("expr") is AssetEvent
for description in orm_execute_state.statement.column_descriptions
):
materializing_queries.append(orm_execute_state.statement)

sqlalchemy_event.listen(session, "do_orm_execute", track_orm_execute)
try:
yield
finally:
sqlalchemy_event.remove(session, "do_orm_execute", track_orm_execute)

assert materializing_queries == []


DEFAULT_DATE = timezone.datetime(2016, 1, 1)
DEFAULT_LOGICAL_DATE = timezone.coerce_datetime(DEFAULT_DATE)
TRY_NUMBER = 1
Expand Down Expand Up @@ -5658,6 +5681,31 @@ def test_create_dag_runs(self, dag_maker):
assert dr.start_date is None
assert dr.creating_job_id == scheduler_job.id

def test_associate_asset_events_with_dag_run_keeps_relationship_coherent(self, session, dag_maker):
with dag_maker(dag_id="asset-event-association"):
pass
dag_run = dag_maker.create_dagrun()
asset = AssetModel(name="association-asset", uri="association-asset", group="asset")
session.add(asset)
session.flush()
asset_event = AssetEvent(asset_id=asset.id)
session.add(asset_event)
session.flush()
asset_event_id = asset_event.id
session.expunge(asset_event)
assert dag_run.consumed_asset_events == []

with assert_queries_count(1, session=session):
_associate_asset_events_with_dag_run(
dag_run=dag_run,
asset_event_ids_select=select(AssetEvent.id.label("event_id")).where(
AssetEvent.id == asset_event_id
),
session=session,
)

assert [event.id for event in dag_run.consumed_asset_events] == [asset_event_id]

@pytest.mark.need_serialized_dag
def test_create_dag_runs_assets(self, session, dag_maker):
"""
Expand Down Expand Up @@ -5731,7 +5779,8 @@ def test_create_dag_runs_assets(self, session, dag_maker):
self.job_runner = SchedulerJobRunner(job=scheduler_job, executors=[self.null_exec])

with create_session() as session:
self.job_runner._create_dagruns_for_dags(session, session)
with assert_no_asset_event_materialization(session):
self.job_runner._create_dagruns_for_dags(session, session)

def dict_from_obj(obj):
"""Get dict of column attrs from SqlAlchemy object."""
Expand Down Expand Up @@ -5939,6 +5988,84 @@ def test_create_dag_runs_asset_triggered_skips_stale_triggered_date(self, sessio
# We do not create a new DagRun since the ADRQ has already been consumed
assert session.scalars(select(DagRun).where(DagRun.dag_id == dag_model.dag_id)).one_or_none() is None

@pytest.mark.need_serialized_dag
@mock.patch("airflow.jobs.scheduler_job_runner._associate_asset_events_with_dag_run", autospec=True)
def test_asset_event_queued_during_association_waits_for_next_run(
self, mock_associate_asset_events, session, dag_maker
):
asset = Asset(name="event-queued-during-association")
with dag_maker(
dag_id="event-queued-during-association-consumer",
schedule=[asset],
catchup=True,
session=session,
):
pass
dag_model = dag_maker.dag_model
asset_id = session.scalar(select(AssetModel.id).where(AssetModel.uri == asset.uri))
first_timestamp = timezone.utcnow()
first_event = AssetEvent(asset_id=asset_id, timestamp=first_timestamp)
session.add(first_event)
session.flush()
session.add(
AssetDagRunQueue(
target_dag_id=dag_model.dag_id,
asset_id=asset_id,
asset_event_id=first_event.id,
)
)
session.flush()

second_timestamp = first_timestamp + timedelta(minutes=1)
second_event_ids = []

def queue_second_event_then_associate(*, dag_run, asset_event_ids_select, session):
if not second_event_ids:
second_event = AssetEvent(asset_id=asset_id, timestamp=second_timestamp)
session.add(second_event)
session.flush()
second_event_ids.append(second_event.id)
session.add(
AssetDagRunQueue(
target_dag_id=dag_model.dag_id,
asset_id=asset_id,
asset_event_id=second_event.id,
)
)
session.flush()
_associate_asset_events_with_dag_run(
dag_run=dag_run,
asset_event_ids_select=asset_event_ids_select,
session=session,
)

mock_associate_asset_events.side_effect = queue_second_event_then_associate
scheduler_job = Job()
self.job_runner = SchedulerJobRunner(job=scheduler_job, executors=[self.null_exec])

self.job_runner._create_dag_runs_asset_triggered(dag_models=[dag_model], session=session)

first_run = session.scalars(select(DagRun).where(DagRun.dag_id == dag_model.dag_id)).one()
assert [event.id for event in first_run.consumed_asset_events] == [first_event.id]
assert first_run.run_after == first_timestamp
assert (
session.scalars(
select(AssetDagRunQueue.asset_event_id).where(
AssetDagRunQueue.target_dag_id == dag_model.dag_id
)
).all()
== second_event_ids
)

self.job_runner._create_dag_runs_asset_triggered(dag_models=[dag_model], session=session)

runs = session.scalars(
select(DagRun).where(DagRun.dag_id == dag_model.dag_id).order_by(DagRun.run_after)
).all()
assert len(runs) == 2
assert [event.id for event in runs[1].consumed_asset_events] == second_event_ids
assert runs[1].run_after == second_timestamp

@pytest.mark.need_serialized_dag
def test_create_dag_runs_asset_triggered_deletes_only_selected_adrq_rows(
self, session: Session, dag_maker
Expand Down Expand Up @@ -11589,7 +11716,8 @@ def test_partitioned_dag_run_with_customized_mapper(
dag_maker=dag_maker,
expected_partition_key="key-1",
)
partition_dags = runner._create_dagruns_for_partitioned_asset_dags(session=session)
with assert_no_asset_event_materialization(session):
partition_dags = runner._create_dagruns_for_partitioned_asset_dags(session=session)
session.refresh(apdr)
# Since asset event for Asset(name="asset-2") with key "key-1" has not yet been created,
# no Dag run will be created
Expand Down