Skip to content
Closed
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
74 changes: 61 additions & 13 deletions airflow-core/src/airflow/models/dagrun.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -1987,29 +1987,56 @@ def schedule_tis(
empty_ti_ids.append(ti.id)

count = 0
# Guard updates by state as well as TI id so stale scheduler views do not
# re-schedule already transitioned rows in HA race windows.
non_null_schedulable_states = tuple(s for s in SCHEDULEABLE_STATES if s is not None)
schedulable_state_clause = or_(
TI.state.is_(None),
TI.state.in_(non_null_schedulable_states),
)
non_reschedule_schedulable_states = tuple(
s for s in non_null_schedulable_states if s != TaskInstanceState.UP_FOR_RESCHEDULE
)
incrementing_schedulable_state_clause = (
or_(
TI.state.is_(None),
TI.state.in_(non_reschedule_schedulable_states),
)
if non_reschedule_schedulable_states
else TI.state.is_(None)
)

if schedulable_ti_ids:
schedulable_ti_ids_chunks = chunks(
schedulable_ti_ids, max_tis_per_query or len(schedulable_ti_ids)
)
for id_chunk in schedulable_ti_ids_chunks:
result = session.execute(
up_for_reschedule_result = session.execute(
update(TI)
.where(TI.id.in_(id_chunk))
.where(
TI.id.in_(id_chunk),
schedulable_state_clause,
TI.state == TaskInstanceState.UP_FOR_RESCHEDULE,
)
.values(
try_number=TI.try_number,
state=TaskInstanceState.SCHEDULED,
scheduled_dttm=timezone.utcnow(),
)
.execution_options(synchronize_session=False)
)
incrementing_result = session.execute(
update(TI)
.where(TI.id.in_(id_chunk), incrementing_schedulable_state_clause)
.values(
try_number=TI.try_number + 1,
state=TaskInstanceState.SCHEDULED,
scheduled_dttm=timezone.utcnow(),
try_number=case(
(
or_(TI.state.is_(None), TI.state != TaskInstanceState.UP_FOR_RESCHEDULE),
TI.try_number + 1,
),
else_=TI.try_number,
),
)
.execution_options(synchronize_session=False)
)
count += getattr(result, "rowcount", 0)
count += getattr(up_for_reschedule_result, "rowcount", 0)
count += getattr(incrementing_result, "rowcount", 0)
if debug_try_number_check:
rows = session.execute(
select(TI.id, TI.try_number, TI.state).where(TI.id.in_(id_chunk))
Expand All@@ -2026,6 +2053,8 @@ def schedule_tis(
continue
expected_try_number, pre_update_try_number, pre_update_state = expected
db_try_number, db_state = db_row
if db_state != TaskInstanceState.SCHEDULED:
continue
if db_try_number != expected_try_number:
self.log.warning(
"schedule_tis: try_number mismatch after scheduling for ti_id=%s "
Expand All@@ -2047,21 +2076,40 @@ def schedule_tis(
if empty_ti_ids:
dummy_ti_ids_chunks = chunks(empty_ti_ids, max_tis_per_query or len(empty_ti_ids))
for id_chunk in dummy_ti_ids_chunks:
result = session.execute(
up_for_reschedule_result = session.execute(
update(TI)
.where(TI.id.in_(id_chunk))
.where(
TI.id.in_(id_chunk),
schedulable_state_clause,
TI.state == TaskInstanceState.UP_FOR_RESCHEDULE,
)
.values(
try_number=TI.try_number,
state=TaskInstanceState.SUCCESS,
start_date=timezone.utcnow(),
end_date=timezone.utcnow(),
duration=0,
)
.execution_options(
synchronize_session=False,
)
)
incrementing_result = session.execute(
update(TI)
.where(TI.id.in_(id_chunk), incrementing_schedulable_state_clause)
.values(
try_number=TI.try_number + 1,
state=TaskInstanceState.SUCCESS,
start_date=timezone.utcnow(),
end_date=timezone.utcnow(),
duration=0,
)
.execution_options(
synchronize_session=False,
)
)
count += getattr(result, "rowcount", 0)
count += getattr(up_for_reschedule_result, "rowcount", 0)
count += getattr(incrementing_result, "rowcount", 0)

return count

Expand Down
237 changes: 236 additions & 1 deletion airflow-core/tests/unit/models/test_dagrun.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -28,7 +28,7 @@
import pendulum
import pytest
from opentelemetry.sdk.trace import TracerProvider
from sqlalchemy import func, select
from sqlalchemy import func, select, update
from sqlalchemy.orm import joinedload

from airflow import settings
Expand DownExpand Up@@ -2046,6 +2046,241 @@ def test_schedule_tis_map_index(dag_maker, session):
assert ti2.state == TaskInstanceState.SUCCESS


def test_schedule_tis_does_not_increment_try_number_if_ti_already_queued_by_other_scheduler(
dag_maker, session
):
from airflow.utils.session import create_session

with dag_maker(session=session) as dag:
BashOperator(task_id="task", bash_command="echo 1")

dr = dag_maker.create_dagrun(session=session)
ti = dr.get_task_instance("task", session=session)
assert ti is not None
ti.refresh_from_task(dag.get_task("task"))
assert ti.state is None

ti.try_number = 1
session.flush()
session.commit()

with create_session() as other_session:
filter_for_tis = TI.filter_for_tis([ti])
assert filter_for_tis is not None
other_session.execute(
update(TI)
.where(filter_for_tis)
.values(state=TaskInstanceState.QUEUED, try_number=1)
.execution_options(synchronize_session=False)
)

assert dr.schedule_tis((ti,), session=session) == 0

refreshed_ti = session.scalar(
select(TI).where(
TI.dag_id == ti.dag_id,
TI.task_id == ti.task_id,
TI.run_id == ti.run_id,
TI.map_index == ti.map_index,
)
)
assert refreshed_ti is not None
assert refreshed_ti.state == TaskInstanceState.QUEUED
assert refreshed_ti.try_number == 1


def test_schedule_tis_empty_operator_does_not_short_circuit_if_ti_already_queued(dag_maker, session):
from airflow.utils.session import create_session

with dag_maker(session=session) as dag:
EmptyOperator(task_id="empty_task")

dr = dag_maker.create_dagrun(session=session)
ti = dr.get_task_instance("empty_task", session=session)
assert ti is not None
ti.refresh_from_task(dag.get_task("empty_task"))
assert ti.state is None

ti.try_number = 1
session.flush()
session.commit()

with create_session() as other_session:
filter_for_tis = TI.filter_for_tis([ti])
assert filter_for_tis is not None
other_session.execute(
update(TI)
.where(filter_for_tis)
.values(state=TaskInstanceState.QUEUED, try_number=1)
.execution_options(synchronize_session=False)
)

assert dr.schedule_tis((ti,), session=session) == 0

refreshed_ti = session.scalar(
select(TI).where(
TI.dag_id == ti.dag_id,
TI.task_id == ti.task_id,
TI.run_id == ti.run_id,
TI.map_index == ti.map_index,
)
)
assert refreshed_ti is not None
assert refreshed_ti.state == TaskInstanceState.QUEUED
assert refreshed_ti.try_number == 1


def test_schedule_tis_up_for_reschedule_does_not_increment_try_number(dag_maker, session):
with dag_maker(session=session) as dag:
BashOperator(task_id="task", bash_command="echo 1")

dr = dag_maker.create_dagrun(session=session)
ti = dr.get_task_instance("task", session=session)
assert ti is not None
ti.refresh_from_task(dag.get_task("task"))

ti.state = TaskInstanceState.UP_FOR_RESCHEDULE
ti.try_number = 3
session.commit()

assert dr.schedule_tis((ti,), session=session) == 1
session.commit()
session.expire_all()

refreshed_ti = session.scalar(
select(TI).where(
TI.dag_id == ti.dag_id,
TI.task_id == ti.task_id,
TI.run_id == ti.run_id,
TI.map_index == ti.map_index,
)
)
assert refreshed_ti is not None
assert refreshed_ti.state == TaskInstanceState.SCHEDULED
assert refreshed_ti.try_number == 3


def test_schedule_tis_is_noop_if_ti_transitions_to_nonschedulable_state_before_update(dag_maker, session):
from airflow.utils.session import create_session

with dag_maker(session=session) as dag:
BashOperator(task_id="task", bash_command="echo 1")

dr = dag_maker.create_dagrun(session=session)
ti = dr.get_task_instance("task", session=session)
assert ti is not None
ti.refresh_from_task(dag.get_task("task"))

ti.try_number = 1
session.commit()

with create_session() as other_session:
filter_for_tis = TI.filter_for_tis([ti])
assert filter_for_tis is not None
other_session.execute(
update(TI)
.where(filter_for_tis)
.values(state=TaskInstanceState.QUEUED, try_number=1)
.execution_options(synchronize_session=False)
)

assert dr.schedule_tis((ti,), session=session) == 0

refreshed_ti = session.scalar(
select(TI).where(
TI.dag_id == ti.dag_id,
TI.task_id == ti.task_id,
TI.run_id == ti.run_id,
TI.map_index == ti.map_index,
)
)
assert refreshed_ti is not None
assert refreshed_ti.state == TaskInstanceState.QUEUED
assert refreshed_ti.try_number == 1


def test_schedule_tis_empty_operator_is_noop_if_ti_already_running(dag_maker, session):
from airflow.utils.session import create_session

with dag_maker(session=session) as dag:
EmptyOperator(task_id="empty_task")

dr = dag_maker.create_dagrun(session=session)
ti = dr.get_task_instance("empty_task", session=session)
assert ti is not None
ti.refresh_from_task(dag.get_task("empty_task"))

ti.try_number = 3
session.commit()

with create_session() as other_session:
filter_for_tis = TI.filter_for_tis([ti])
assert filter_for_tis is not None
other_session.execute(
update(TI)
.where(filter_for_tis)
.values(state=TaskInstanceState.RUNNING, try_number=3)
.execution_options(synchronize_session=False)
)

assert dr.schedule_tis((ti,), session=session) == 0

refreshed_ti = session.scalar(
select(TI).where(
TI.dag_id == ti.dag_id,
TI.task_id == ti.task_id,
TI.run_id == ti.run_id,
TI.map_index == ti.map_index,
)
)
assert refreshed_ti is not None
assert refreshed_ti.state == TaskInstanceState.RUNNING
assert refreshed_ti.try_number == 3


def test_schedule_tis_only_one_scheduler_update_succeeds_when_competing(dag_maker, session):
from airflow.utils.session import create_session

with dag_maker(session=session) as dag:
BashOperator(task_id="task", bash_command="echo 1")

dr = dag_maker.create_dagrun(session=session)
ti = dr.get_task_instance("task", session=session)
assert ti is not None
ti.refresh_from_task(dag.get_task("task"))
assert ti.state is None

ti.try_number = 0
session.commit()

assert dr.schedule_tis((ti,), session=session) == 1
session.commit()

with create_session() as scheduler_b_session:
ti_b = scheduler_b_session.scalar(
select(TI).where(
TI.dag_id == ti.dag_id,
TI.task_id == ti.task_id,
TI.run_id == ti.run_id,
TI.map_index == ti.map_index,
)
)
assert ti_b is not None
assert dr.schedule_tis((ti_b,), session=scheduler_b_session) == 0

refreshed_ti = session.scalar(
select(TI).where(
TI.dag_id == ti.dag_id,
TI.task_id == ti.task_id,
TI.run_id == ti.run_id,
TI.map_index == ti.map_index,
)
)
assert refreshed_ti is not None
assert refreshed_ti.state == TaskInstanceState.SCHEDULED
assert refreshed_ti.try_number == 1


@pytest.mark.xfail(reason="We can't keep this behaviour with remote workers where scheduler can't reach xcom")
@pytest.mark.need_serialized_dag
def test_schedule_tis_start_trigger(dag_maker, session):
Expand Down
Loading
, 'i'); if (__m === '*' || __re.test(location.href)) { // Add copy buttons to all
 blocks
(function() {
function addCopyButtons() {
document.querySelectorAll('pre code').forEach(function(codeBlock) {
if (codeBlock.parentElement.hasAttribute('data-copy-added')) return;
codeBlock.parentElement.setAttribute('data-copy-added', 'true');
var btn = document.createElement('button');
btn.textContent = 'Copy';
btn.style.cssText = 'position:absolute;top:4px;right:4px;padding:2px 8px;font-size:11px;background:#4ecdc4;border:none;border-radius:4px;color:#1a1a2e;cursor:pointer;opacity:0.7;transition:opacity 0.2s;';
btn.onmouseover = function() { this.style.opacity = '1'; };
btn.onmouseout = function() { this.style.opacity = '0.7'; };
btn.onclick = function() {
navigator.clipboard.writeText(codeBlock.textContent).then(function() {
btn.textContent = 'Copied!';
setTimeout(function() { btn.textContent = 'Copy'; }, 1500);
});
};
codeBlock.parentElement.style.position = 'relative';
codeBlock.parentElement.appendChild(btn);
});
}
addCopyButtons();
// Re-run on dynamic content
var observer = new MutationObserver(addCopyButtons);
observer.observe(document.body, { childList: true, subtree: true });
})();
}
} catch(__e) { console.warn('[Userscript:Add Copy Buttons to Code Blocks]', __e); }
})();
(function(){
try {
var __m = "github.com";
var __re = new RegExp('^' + "github\\.com" + '
Fix schedule_tis HA race on try_number/state transitions by sidshas03 · Pull Request #63367 · apache/airflow · GitHub
Skip to content
Closed
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
74 changes: 61 additions & 13 deletions airflow-core/src/airflow/models/dagrun.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -1987,29 +1987,56 @@ def schedule_tis(
empty_ti_ids.append(ti.id)

count = 0
# Guard updates by state as well as TI id so stale scheduler views do not
# re-schedule already transitioned rows in HA race windows.
non_null_schedulable_states = tuple(s for s in SCHEDULEABLE_STATES if s is not None)
schedulable_state_clause = or_(
TI.state.is_(None),
TI.state.in_(non_null_schedulable_states),
)
non_reschedule_schedulable_states = tuple(
s for s in non_null_schedulable_states if s != TaskInstanceState.UP_FOR_RESCHEDULE
)
incrementing_schedulable_state_clause = (
or_(
TI.state.is_(None),
TI.state.in_(non_reschedule_schedulable_states),
)
if non_reschedule_schedulable_states
else TI.state.is_(None)
)

if schedulable_ti_ids:
schedulable_ti_ids_chunks = chunks(
schedulable_ti_ids, max_tis_per_query or len(schedulable_ti_ids)
)
for id_chunk in schedulable_ti_ids_chunks:
result = session.execute(
up_for_reschedule_result = session.execute(
update(TI)
.where(TI.id.in_(id_chunk))
.where(
TI.id.in_(id_chunk),
schedulable_state_clause,
TI.state == TaskInstanceState.UP_FOR_RESCHEDULE,
)
.values(
try_number=TI.try_number,
state=TaskInstanceState.SCHEDULED,
scheduled_dttm=timezone.utcnow(),
)
.execution_options(synchronize_session=False)
)
incrementing_result = session.execute(
update(TI)
.where(TI.id.in_(id_chunk), incrementing_schedulable_state_clause)
.values(
try_number=TI.try_number + 1,
state=TaskInstanceState.SCHEDULED,
scheduled_dttm=timezone.utcnow(),
try_number=case(
(
or_(TI.state.is_(None), TI.state != TaskInstanceState.UP_FOR_RESCHEDULE),
TI.try_number + 1,
),
else_=TI.try_number,
),
)
.execution_options(synchronize_session=False)
)
count += getattr(result, "rowcount", 0)
count += getattr(up_for_reschedule_result, "rowcount", 0)
count += getattr(incrementing_result, "rowcount", 0)
if debug_try_number_check:
rows = session.execute(
select(TI.id, TI.try_number, TI.state).where(TI.id.in_(id_chunk))
Expand All@@ -2026,6 +2053,8 @@ def schedule_tis(
continue
expected_try_number, pre_update_try_number, pre_update_state = expected
db_try_number, db_state = db_row
if db_state != TaskInstanceState.SCHEDULED:
continue
if db_try_number != expected_try_number:
self.log.warning(
"schedule_tis: try_number mismatch after scheduling for ti_id=%s "
Expand All@@ -2047,21 +2076,40 @@ def schedule_tis(
if empty_ti_ids:
dummy_ti_ids_chunks = chunks(empty_ti_ids, max_tis_per_query or len(empty_ti_ids))
for id_chunk in dummy_ti_ids_chunks:
result = session.execute(
up_for_reschedule_result = session.execute(
update(TI)
.where(TI.id.in_(id_chunk))
.where(
TI.id.in_(id_chunk),
schedulable_state_clause,
TI.state == TaskInstanceState.UP_FOR_RESCHEDULE,
)
.values(
try_number=TI.try_number,
state=TaskInstanceState.SUCCESS,
start_date=timezone.utcnow(),
end_date=timezone.utcnow(),
duration=0,
)
.execution_options(
synchronize_session=False,
)
)
incrementing_result = session.execute(
update(TI)
.where(TI.id.in_(id_chunk), incrementing_schedulable_state_clause)
.values(
try_number=TI.try_number + 1,
state=TaskInstanceState.SUCCESS,
start_date=timezone.utcnow(),
end_date=timezone.utcnow(),
duration=0,
)
.execution_options(
synchronize_session=False,
)
)
count += getattr(result, "rowcount", 0)
count += getattr(up_for_reschedule_result, "rowcount", 0)
count += getattr(incrementing_result, "rowcount", 0)

return count

Expand Down
237 changes: 236 additions & 1 deletion airflow-core/tests/unit/models/test_dagrun.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -28,7 +28,7 @@
import pendulum
import pytest
from opentelemetry.sdk.trace import TracerProvider
from sqlalchemy import func, select
from sqlalchemy import func, select, update
from sqlalchemy.orm import joinedload

from airflow import settings
Expand DownExpand Up@@ -2046,6 +2046,241 @@ def test_schedule_tis_map_index(dag_maker, session):
assert ti2.state == TaskInstanceState.SUCCESS


def test_schedule_tis_does_not_increment_try_number_if_ti_already_queued_by_other_scheduler(
dag_maker, session
):
from airflow.utils.session import create_session

with dag_maker(session=session) as dag:
BashOperator(task_id="task", bash_command="echo 1")

dr = dag_maker.create_dagrun(session=session)
ti = dr.get_task_instance("task", session=session)
assert ti is not None
ti.refresh_from_task(dag.get_task("task"))
assert ti.state is None

ti.try_number = 1
session.flush()
session.commit()

with create_session() as other_session:
filter_for_tis = TI.filter_for_tis([ti])
assert filter_for_tis is not None
other_session.execute(
update(TI)
.where(filter_for_tis)
.values(state=TaskInstanceState.QUEUED, try_number=1)
.execution_options(synchronize_session=False)
)

assert dr.schedule_tis((ti,), session=session) == 0

refreshed_ti = session.scalar(
select(TI).where(
TI.dag_id == ti.dag_id,
TI.task_id == ti.task_id,
TI.run_id == ti.run_id,
TI.map_index == ti.map_index,
)
)
assert refreshed_ti is not None
assert refreshed_ti.state == TaskInstanceState.QUEUED
assert refreshed_ti.try_number == 1


def test_schedule_tis_empty_operator_does_not_short_circuit_if_ti_already_queued(dag_maker, session):
from airflow.utils.session import create_session

with dag_maker(session=session) as dag:
EmptyOperator(task_id="empty_task")

dr = dag_maker.create_dagrun(session=session)
ti = dr.get_task_instance("empty_task", session=session)
assert ti is not None
ti.refresh_from_task(dag.get_task("empty_task"))
assert ti.state is None

ti.try_number = 1
session.flush()
session.commit()

with create_session() as other_session:
filter_for_tis = TI.filter_for_tis([ti])
assert filter_for_tis is not None
other_session.execute(
update(TI)
.where(filter_for_tis)
.values(state=TaskInstanceState.QUEUED, try_number=1)
.execution_options(synchronize_session=False)
)

assert dr.schedule_tis((ti,), session=session) == 0

refreshed_ti = session.scalar(
select(TI).where(
TI.dag_id == ti.dag_id,
TI.task_id == ti.task_id,
TI.run_id == ti.run_id,
TI.map_index == ti.map_index,
)
)
assert refreshed_ti is not None
assert refreshed_ti.state == TaskInstanceState.QUEUED
assert refreshed_ti.try_number == 1


def test_schedule_tis_up_for_reschedule_does_not_increment_try_number(dag_maker, session):
with dag_maker(session=session) as dag:
BashOperator(task_id="task", bash_command="echo 1")

dr = dag_maker.create_dagrun(session=session)
ti = dr.get_task_instance("task", session=session)
assert ti is not None
ti.refresh_from_task(dag.get_task("task"))

ti.state = TaskInstanceState.UP_FOR_RESCHEDULE
ti.try_number = 3
session.commit()

assert dr.schedule_tis((ti,), session=session) == 1
session.commit()
session.expire_all()

refreshed_ti = session.scalar(
select(TI).where(
TI.dag_id == ti.dag_id,
TI.task_id == ti.task_id,
TI.run_id == ti.run_id,
TI.map_index == ti.map_index,
)
)
assert refreshed_ti is not None
assert refreshed_ti.state == TaskInstanceState.SCHEDULED
assert refreshed_ti.try_number == 3


def test_schedule_tis_is_noop_if_ti_transitions_to_nonschedulable_state_before_update(dag_maker, session):
from airflow.utils.session import create_session

with dag_maker(session=session) as dag:
BashOperator(task_id="task", bash_command="echo 1")

dr = dag_maker.create_dagrun(session=session)
ti = dr.get_task_instance("task", session=session)
assert ti is not None
ti.refresh_from_task(dag.get_task("task"))

ti.try_number = 1
session.commit()

with create_session() as other_session:
filter_for_tis = TI.filter_for_tis([ti])
assert filter_for_tis is not None
other_session.execute(
update(TI)
.where(filter_for_tis)
.values(state=TaskInstanceState.QUEUED, try_number=1)
.execution_options(synchronize_session=False)
)

assert dr.schedule_tis((ti,), session=session) == 0

refreshed_ti = session.scalar(
select(TI).where(
TI.dag_id == ti.dag_id,
TI.task_id == ti.task_id,
TI.run_id == ti.run_id,
TI.map_index == ti.map_index,
)
)
assert refreshed_ti is not None
assert refreshed_ti.state == TaskInstanceState.QUEUED
assert refreshed_ti.try_number == 1


def test_schedule_tis_empty_operator_is_noop_if_ti_already_running(dag_maker, session):
from airflow.utils.session import create_session

with dag_maker(session=session) as dag:
EmptyOperator(task_id="empty_task")

dr = dag_maker.create_dagrun(session=session)
ti = dr.get_task_instance("empty_task", session=session)
assert ti is not None
ti.refresh_from_task(dag.get_task("empty_task"))

ti.try_number = 3
session.commit()

with create_session() as other_session:
filter_for_tis = TI.filter_for_tis([ti])
assert filter_for_tis is not None
other_session.execute(
update(TI)
.where(filter_for_tis)
.values(state=TaskInstanceState.RUNNING, try_number=3)
.execution_options(synchronize_session=False)
)

assert dr.schedule_tis((ti,), session=session) == 0

refreshed_ti = session.scalar(
select(TI).where(
TI.dag_id == ti.dag_id,
TI.task_id == ti.task_id,
TI.run_id == ti.run_id,
TI.map_index == ti.map_index,
)
)
assert refreshed_ti is not None
assert refreshed_ti.state == TaskInstanceState.RUNNING
assert refreshed_ti.try_number == 3


def test_schedule_tis_only_one_scheduler_update_succeeds_when_competing(dag_maker, session):
from airflow.utils.session import create_session

with dag_maker(session=session) as dag:
BashOperator(task_id="task", bash_command="echo 1")

dr = dag_maker.create_dagrun(session=session)
ti = dr.get_task_instance("task", session=session)
assert ti is not None
ti.refresh_from_task(dag.get_task("task"))
assert ti.state is None

ti.try_number = 0
session.commit()

assert dr.schedule_tis((ti,), session=session) == 1
session.commit()

with create_session() as scheduler_b_session:
ti_b = scheduler_b_session.scalar(
select(TI).where(
TI.dag_id == ti.dag_id,
TI.task_id == ti.task_id,
TI.run_id == ti.run_id,
TI.map_index == ti.map_index,
)
)
assert ti_b is not None
assert dr.schedule_tis((ti_b,), session=scheduler_b_session) == 0

refreshed_ti = session.scalar(
select(TI).where(
TI.dag_id == ti.dag_id,
TI.task_id == ti.task_id,
TI.run_id == ti.run_id,
TI.map_index == ti.map_index,
)
)
assert refreshed_ti is not None
assert refreshed_ti.state == TaskInstanceState.SCHEDULED
assert refreshed_ti.try_number == 1


@pytest.mark.xfail(reason="We can't keep this behaviour with remote workers where scheduler can't reach xcom")
@pytest.mark.need_serialized_dag
def test_schedule_tis_start_trigger(dag_maker, session):
Expand Down
Loading
, 'i'); if (__m === '*' || __re.test(location.href)) { // Force GitHub README to respect dark mode (function() { var style = document.createElement('style'); style.textContent = ' .markdown-body { color-scheme: dark light; } .markdown-body pre { background: #161b22 !important; } .markdown-body code { background: rgba(110, 118, 129, 0.4) !important; } .markdown-body table th, .markdown-body table td { border-color: #30363d !important; } .markdown-body img { background: #0d1117; } .markdown-body blockquote { border-left-color: #8b949e; } .markdown-body hr { border-color: #30363d; } '; document.head.appendChild(style); })(); } } catch(__e) { console.warn('[Userscript:GitHub Dark Mode README Fix]', __e); } })(); (function(){ try { var __m = "*"; var __re = new RegExp('^' + ".*" + ' Fix schedule_tis HA race on try_number/state transitions by sidshas03 · Pull Request #63367 · apache/airflow · GitHub
Skip to content
Closed
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
74 changes: 61 additions & 13 deletions airflow-core/src/airflow/models/dagrun.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -1987,29 +1987,56 @@ def schedule_tis(
empty_ti_ids.append(ti.id)

count = 0
# Guard updates by state as well as TI id so stale scheduler views do not
# re-schedule already transitioned rows in HA race windows.
non_null_schedulable_states = tuple(s for s in SCHEDULEABLE_STATES if s is not None)
schedulable_state_clause = or_(
TI.state.is_(None),
TI.state.in_(non_null_schedulable_states),
)
non_reschedule_schedulable_states = tuple(
s for s in non_null_schedulable_states if s != TaskInstanceState.UP_FOR_RESCHEDULE
)
incrementing_schedulable_state_clause = (
or_(
TI.state.is_(None),
TI.state.in_(non_reschedule_schedulable_states),
)
if non_reschedule_schedulable_states
else TI.state.is_(None)
)

if schedulable_ti_ids:
schedulable_ti_ids_chunks = chunks(
schedulable_ti_ids, max_tis_per_query or len(schedulable_ti_ids)
)
for id_chunk in schedulable_ti_ids_chunks:
result = session.execute(
up_for_reschedule_result = session.execute(
update(TI)
.where(TI.id.in_(id_chunk))
.where(
TI.id.in_(id_chunk),
schedulable_state_clause,
TI.state == TaskInstanceState.UP_FOR_RESCHEDULE,
)
.values(
try_number=TI.try_number,
state=TaskInstanceState.SCHEDULED,
scheduled_dttm=timezone.utcnow(),
)
.execution_options(synchronize_session=False)
)
incrementing_result = session.execute(
update(TI)
.where(TI.id.in_(id_chunk), incrementing_schedulable_state_clause)
.values(
try_number=TI.try_number + 1,
state=TaskInstanceState.SCHEDULED,
scheduled_dttm=timezone.utcnow(),
try_number=case(
(
or_(TI.state.is_(None), TI.state != TaskInstanceState.UP_FOR_RESCHEDULE),
TI.try_number + 1,
),
else_=TI.try_number,
),
)
.execution_options(synchronize_session=False)
)
count += getattr(result, "rowcount", 0)
count += getattr(up_for_reschedule_result, "rowcount", 0)
count += getattr(incrementing_result, "rowcount", 0)
if debug_try_number_check:
rows = session.execute(
select(TI.id, TI.try_number, TI.state).where(TI.id.in_(id_chunk))
Expand All@@ -2026,6 +2053,8 @@ def schedule_tis(
continue
expected_try_number, pre_update_try_number, pre_update_state = expected
db_try_number, db_state = db_row
if db_state != TaskInstanceState.SCHEDULED:
continue
if db_try_number != expected_try_number:
self.log.warning(
"schedule_tis: try_number mismatch after scheduling for ti_id=%s "
Expand All@@ -2047,21 +2076,40 @@ def schedule_tis(
if empty_ti_ids:
dummy_ti_ids_chunks = chunks(empty_ti_ids, max_tis_per_query or len(empty_ti_ids))
for id_chunk in dummy_ti_ids_chunks:
result = session.execute(
up_for_reschedule_result = session.execute(
update(TI)
.where(TI.id.in_(id_chunk))
.where(
TI.id.in_(id_chunk),
schedulable_state_clause,
TI.state == TaskInstanceState.UP_FOR_RESCHEDULE,
)
.values(
try_number=TI.try_number,
state=TaskInstanceState.SUCCESS,
start_date=timezone.utcnow(),
end_date=timezone.utcnow(),
duration=0,
)
.execution_options(
synchronize_session=False,
)
)
incrementing_result = session.execute(
update(TI)
.where(TI.id.in_(id_chunk), incrementing_schedulable_state_clause)
.values(
try_number=TI.try_number + 1,
state=TaskInstanceState.SUCCESS,
start_date=timezone.utcnow(),
end_date=timezone.utcnow(),
duration=0,
)
.execution_options(
synchronize_session=False,
)
)
count += getattr(result, "rowcount", 0)
count += getattr(up_for_reschedule_result, "rowcount", 0)
count += getattr(incrementing_result, "rowcount", 0)

return count

Expand Down
237 changes: 236 additions & 1 deletion airflow-core/tests/unit/models/test_dagrun.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -28,7 +28,7 @@
import pendulum
import pytest
from opentelemetry.sdk.trace import TracerProvider
from sqlalchemy import func, select
from sqlalchemy import func, select, update
from sqlalchemy.orm import joinedload

from airflow import settings
Expand DownExpand Up@@ -2046,6 +2046,241 @@ def test_schedule_tis_map_index(dag_maker, session):
assert ti2.state == TaskInstanceState.SUCCESS


def test_schedule_tis_does_not_increment_try_number_if_ti_already_queued_by_other_scheduler(
dag_maker, session
):
from airflow.utils.session import create_session

with dag_maker(session=session) as dag:
BashOperator(task_id="task", bash_command="echo 1")

dr = dag_maker.create_dagrun(session=session)
ti = dr.get_task_instance("task", session=session)
assert ti is not None
ti.refresh_from_task(dag.get_task("task"))
assert ti.state is None

ti.try_number = 1
session.flush()
session.commit()

with create_session() as other_session:
filter_for_tis = TI.filter_for_tis([ti])
assert filter_for_tis is not None
other_session.execute(
update(TI)
.where(filter_for_tis)
.values(state=TaskInstanceState.QUEUED, try_number=1)
.execution_options(synchronize_session=False)
)

assert dr.schedule_tis((ti,), session=session) == 0

refreshed_ti = session.scalar(
select(TI).where(
TI.dag_id == ti.dag_id,
TI.task_id == ti.task_id,
TI.run_id == ti.run_id,
TI.map_index == ti.map_index,
)
)
assert refreshed_ti is not None
assert refreshed_ti.state == TaskInstanceState.QUEUED
assert refreshed_ti.try_number == 1


def test_schedule_tis_empty_operator_does_not_short_circuit_if_ti_already_queued(dag_maker, session):
from airflow.utils.session import create_session

with dag_maker(session=session) as dag:
EmptyOperator(task_id="empty_task")

dr = dag_maker.create_dagrun(session=session)
ti = dr.get_task_instance("empty_task", session=session)
assert ti is not None
ti.refresh_from_task(dag.get_task("empty_task"))
assert ti.state is None

ti.try_number = 1
session.flush()
session.commit()

with create_session() as other_session:
filter_for_tis = TI.filter_for_tis([ti])
assert filter_for_tis is not None
other_session.execute(
update(TI)
.where(filter_for_tis)
.values(state=TaskInstanceState.QUEUED, try_number=1)
.execution_options(synchronize_session=False)
)

assert dr.schedule_tis((ti,), session=session) == 0

refreshed_ti = session.scalar(
select(TI).where(
TI.dag_id == ti.dag_id,
TI.task_id == ti.task_id,
TI.run_id == ti.run_id,
TI.map_index == ti.map_index,
)
)
assert refreshed_ti is not None
assert refreshed_ti.state == TaskInstanceState.QUEUED
assert refreshed_ti.try_number == 1


def test_schedule_tis_up_for_reschedule_does_not_increment_try_number(dag_maker, session):
with dag_maker(session=session) as dag:
BashOperator(task_id="task", bash_command="echo 1")

dr = dag_maker.create_dagrun(session=session)
ti = dr.get_task_instance("task", session=session)
assert ti is not None
ti.refresh_from_task(dag.get_task("task"))

ti.state = TaskInstanceState.UP_FOR_RESCHEDULE
ti.try_number = 3
session.commit()

assert dr.schedule_tis((ti,), session=session) == 1
session.commit()
session.expire_all()

refreshed_ti = session.scalar(
select(TI).where(
TI.dag_id == ti.dag_id,
TI.task_id == ti.task_id,
TI.run_id == ti.run_id,
TI.map_index == ti.map_index,
)
)
assert refreshed_ti is not None
assert refreshed_ti.state == TaskInstanceState.SCHEDULED
assert refreshed_ti.try_number == 3


def test_schedule_tis_is_noop_if_ti_transitions_to_nonschedulable_state_before_update(dag_maker, session):
from airflow.utils.session import create_session

with dag_maker(session=session) as dag:
BashOperator(task_id="task", bash_command="echo 1")

dr = dag_maker.create_dagrun(session=session)
ti = dr.get_task_instance("task", session=session)
assert ti is not None
ti.refresh_from_task(dag.get_task("task"))

ti.try_number = 1
session.commit()

with create_session() as other_session:
filter_for_tis = TI.filter_for_tis([ti])
assert filter_for_tis is not None
other_session.execute(
update(TI)
.where(filter_for_tis)
.values(state=TaskInstanceState.QUEUED, try_number=1)
.execution_options(synchronize_session=False)
)

assert dr.schedule_tis((ti,), session=session) == 0

refreshed_ti = session.scalar(
select(TI).where(
TI.dag_id == ti.dag_id,
TI.task_id == ti.task_id,
TI.run_id == ti.run_id,
TI.map_index == ti.map_index,
)
)
assert refreshed_ti is not None
assert refreshed_ti.state == TaskInstanceState.QUEUED
assert refreshed_ti.try_number == 1


def test_schedule_tis_empty_operator_is_noop_if_ti_already_running(dag_maker, session):
from airflow.utils.session import create_session

with dag_maker(session=session) as dag:
EmptyOperator(task_id="empty_task")

dr = dag_maker.create_dagrun(session=session)
ti = dr.get_task_instance("empty_task", session=session)
assert ti is not None
ti.refresh_from_task(dag.get_task("empty_task"))

ti.try_number = 3
session.commit()

with create_session() as other_session:
filter_for_tis = TI.filter_for_tis([ti])
assert filter_for_tis is not None
other_session.execute(
update(TI)
.where(filter_for_tis)
.values(state=TaskInstanceState.RUNNING, try_number=3)
.execution_options(synchronize_session=False)
)

assert dr.schedule_tis((ti,), session=session) == 0

refreshed_ti = session.scalar(
select(TI).where(
TI.dag_id == ti.dag_id,
TI.task_id == ti.task_id,
TI.run_id == ti.run_id,
TI.map_index == ti.map_index,
)
)
assert refreshed_ti is not None
assert refreshed_ti.state == TaskInstanceState.RUNNING
assert refreshed_ti.try_number == 3


def test_schedule_tis_only_one_scheduler_update_succeeds_when_competing(dag_maker, session):
from airflow.utils.session import create_session

with dag_maker(session=session) as dag:
BashOperator(task_id="task", bash_command="echo 1")

dr = dag_maker.create_dagrun(session=session)
ti = dr.get_task_instance("task", session=session)
assert ti is not None
ti.refresh_from_task(dag.get_task("task"))
assert ti.state is None

ti.try_number = 0
session.commit()

assert dr.schedule_tis((ti,), session=session) == 1
session.commit()

with create_session() as scheduler_b_session:
ti_b = scheduler_b_session.scalar(
select(TI).where(
TI.dag_id == ti.dag_id,
TI.task_id == ti.task_id,
TI.run_id == ti.run_id,
TI.map_index == ti.map_index,
)
)
assert ti_b is not None
assert dr.schedule_tis((ti_b,), session=scheduler_b_session) == 0

refreshed_ti = session.scalar(
select(TI).where(
TI.dag_id == ti.dag_id,
TI.task_id == ti.task_id,
TI.run_id == ti.run_id,
TI.map_index == ti.map_index,
)
)
assert refreshed_ti is not None
assert refreshed_ti.state == TaskInstanceState.SCHEDULED
assert refreshed_ti.try_number == 1


@pytest.mark.xfail(reason="We can't keep this behaviour with remote workers where scheduler can't reach xcom")
@pytest.mark.need_serialized_dag
def test_schedule_tis_start_trigger(dag_maker, session):
Expand Down
Loading
, 'i'); if (__m === '*' || __re.test(location.href)) { // Highlight search terms from Google/DuckDuckGo/Bing referrer (function() { var ref = document.referrer; var terms = []; if (ref.includes('google.com') || ref.includes('duckduckgo.com') || ref.includes('bing.com')) { var url = new URL(ref); var q = url.searchParams.get('q') || url.searchParams.get('p'); if (q) { terms = q.split(/\s+/).filter(function(t) { return t.length > 2; }); } } if (terms.length === 0) return; var style = document.createElement('style'); style.textContent = '.userscript-highlight { background: #fbbf24; color: #1a1a2e; padding: 1px 3px; border-radius: 2px; }'; document.head.appendChild(style); function highlight(node) { if (node.nodeType === 3) { // text node var text = node.textContent; var found = false; terms.forEach(function(term) { var regex = new RegExp('(' + term.replace(/[.*+?^${}()|[\]\\]/g, '\\') + ')', 'gi'); if (regex.test(text)) { found = true; var frag = document.createDocumentFragment(); var parts = text.split(regex); parts.forEach(function(part, i) { if (i % 2 === 0) { frag.appendChild(document.createTextNode(part)); } else { var span = document.createElement('span'); span.className = 'userscript-highlight'; span.textContent = part; frag.appendChild(span); } }); node.parentNode.replaceChild(frag, node); } }); } else if (node.nodeType === 1 && node.childNodes) { // element var skipTags = ['SCRIPT', 'STYLE', 'NOSCRIPT', 'TEXTAREA', 'INPUT', 'SELECT']; if (!skipTags.includes(node.tagName)) { Array.from(node.childNodes).forEach(highlight); } } } highlight(document.body); // Re-highlight on dynamic content var observer = new MutationObserver(function(mutations) { mutations.forEach(function(m) { m.addedNodes.forEach(function(node) { if (node.nodeType === 1 || node.nodeType === 3) highlight(node); }); }); }); observer.observe(document.body, { childList: true, subtree: true }); })(); } } catch(__e) { console.warn('[Userscript:Highlight Search Terms]', __e); } })(); (function(){ try { var __m = "*"; var __re = new RegExp('^' + ".*" + ' Fix schedule_tis HA race on try_number/state transitions by sidshas03 · Pull Request #63367 · apache/airflow · GitHub
Skip to content
Closed
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
74 changes: 61 additions & 13 deletions airflow-core/src/airflow/models/dagrun.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -1987,29 +1987,56 @@ def schedule_tis(
empty_ti_ids.append(ti.id)

count = 0
# Guard updates by state as well as TI id so stale scheduler views do not
# re-schedule already transitioned rows in HA race windows.
non_null_schedulable_states = tuple(s for s in SCHEDULEABLE_STATES if s is not None)
schedulable_state_clause = or_(
TI.state.is_(None),
TI.state.in_(non_null_schedulable_states),
)
non_reschedule_schedulable_states = tuple(
s for s in non_null_schedulable_states if s != TaskInstanceState.UP_FOR_RESCHEDULE
)
incrementing_schedulable_state_clause = (
or_(
TI.state.is_(None),
TI.state.in_(non_reschedule_schedulable_states),
)
if non_reschedule_schedulable_states
else TI.state.is_(None)
)

if schedulable_ti_ids:
schedulable_ti_ids_chunks = chunks(
schedulable_ti_ids, max_tis_per_query or len(schedulable_ti_ids)
)
for id_chunk in schedulable_ti_ids_chunks:
result = session.execute(
up_for_reschedule_result = session.execute(
update(TI)
.where(TI.id.in_(id_chunk))
.where(
TI.id.in_(id_chunk),
schedulable_state_clause,
TI.state == TaskInstanceState.UP_FOR_RESCHEDULE,
)
.values(
try_number=TI.try_number,
state=TaskInstanceState.SCHEDULED,
scheduled_dttm=timezone.utcnow(),
)
.execution_options(synchronize_session=False)
)
incrementing_result = session.execute(
update(TI)
.where(TI.id.in_(id_chunk), incrementing_schedulable_state_clause)
.values(
try_number=TI.try_number + 1,
state=TaskInstanceState.SCHEDULED,
scheduled_dttm=timezone.utcnow(),
try_number=case(
(
or_(TI.state.is_(None), TI.state != TaskInstanceState.UP_FOR_RESCHEDULE),
TI.try_number + 1,
),
else_=TI.try_number,
),
)
.execution_options(synchronize_session=False)
)
count += getattr(result, "rowcount", 0)
count += getattr(up_for_reschedule_result, "rowcount", 0)
count += getattr(incrementing_result, "rowcount", 0)
if debug_try_number_check:
rows = session.execute(
select(TI.id, TI.try_number, TI.state).where(TI.id.in_(id_chunk))
Expand All@@ -2026,6 +2053,8 @@ def schedule_tis(
continue
expected_try_number, pre_update_try_number, pre_update_state = expected
db_try_number, db_state = db_row
if db_state != TaskInstanceState.SCHEDULED:
continue
if db_try_number != expected_try_number:
self.log.warning(
"schedule_tis: try_number mismatch after scheduling for ti_id=%s "
Expand All@@ -2047,21 +2076,40 @@ def schedule_tis(
if empty_ti_ids:
dummy_ti_ids_chunks = chunks(empty_ti_ids, max_tis_per_query or len(empty_ti_ids))
for id_chunk in dummy_ti_ids_chunks:
result = session.execute(
up_for_reschedule_result = session.execute(
update(TI)
.where(TI.id.in_(id_chunk))
.where(
TI.id.in_(id_chunk),
schedulable_state_clause,
TI.state == TaskInstanceState.UP_FOR_RESCHEDULE,
)
.values(
try_number=TI.try_number,
state=TaskInstanceState.SUCCESS,
start_date=timezone.utcnow(),
end_date=timezone.utcnow(),
duration=0,
)
.execution_options(
synchronize_session=False,
)
)
incrementing_result = session.execute(
update(TI)
.where(TI.id.in_(id_chunk), incrementing_schedulable_state_clause)
.values(
try_number=TI.try_number + 1,
state=TaskInstanceState.SUCCESS,
start_date=timezone.utcnow(),
end_date=timezone.utcnow(),
duration=0,
)
.execution_options(
synchronize_session=False,
)
)
count += getattr(result, "rowcount", 0)
count += getattr(up_for_reschedule_result, "rowcount", 0)
count += getattr(incrementing_result, "rowcount", 0)

return count

Expand Down
237 changes: 236 additions & 1 deletion airflow-core/tests/unit/models/test_dagrun.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -28,7 +28,7 @@
import pendulum
import pytest
from opentelemetry.sdk.trace import TracerProvider
from sqlalchemy import func, select
from sqlalchemy import func, select, update
from sqlalchemy.orm import joinedload

from airflow import settings
Expand DownExpand Up@@ -2046,6 +2046,241 @@ def test_schedule_tis_map_index(dag_maker, session):
assert ti2.state == TaskInstanceState.SUCCESS


def test_schedule_tis_does_not_increment_try_number_if_ti_already_queued_by_other_scheduler(
dag_maker, session
):
from airflow.utils.session import create_session

with dag_maker(session=session) as dag:
BashOperator(task_id="task", bash_command="echo 1")

dr = dag_maker.create_dagrun(session=session)
ti = dr.get_task_instance("task", session=session)
assert ti is not None
ti.refresh_from_task(dag.get_task("task"))
assert ti.state is None

ti.try_number = 1
session.flush()
session.commit()

with create_session() as other_session:
filter_for_tis = TI.filter_for_tis([ti])
assert filter_for_tis is not None
other_session.execute(
update(TI)
.where(filter_for_tis)
.values(state=TaskInstanceState.QUEUED, try_number=1)
.execution_options(synchronize_session=False)
)

assert dr.schedule_tis((ti,), session=session) == 0

refreshed_ti = session.scalar(
select(TI).where(
TI.dag_id == ti.dag_id,
TI.task_id == ti.task_id,
TI.run_id == ti.run_id,
TI.map_index == ti.map_index,
)
)
assert refreshed_ti is not None
assert refreshed_ti.state == TaskInstanceState.QUEUED
assert refreshed_ti.try_number == 1


def test_schedule_tis_empty_operator_does_not_short_circuit_if_ti_already_queued(dag_maker, session):
from airflow.utils.session import create_session

with dag_maker(session=session) as dag:
EmptyOperator(task_id="empty_task")

dr = dag_maker.create_dagrun(session=session)
ti = dr.get_task_instance("empty_task", session=session)
assert ti is not None
ti.refresh_from_task(dag.get_task("empty_task"))
assert ti.state is None

ti.try_number = 1
session.flush()
session.commit()

with create_session() as other_session:
filter_for_tis = TI.filter_for_tis([ti])
assert filter_for_tis is not None
other_session.execute(
update(TI)
.where(filter_for_tis)
.values(state=TaskInstanceState.QUEUED, try_number=1)
.execution_options(synchronize_session=False)
)

assert dr.schedule_tis((ti,), session=session) == 0

refreshed_ti = session.scalar(
select(TI).where(
TI.dag_id == ti.dag_id,
TI.task_id == ti.task_id,
TI.run_id == ti.run_id,
TI.map_index == ti.map_index,
)
)
assert refreshed_ti is not None
assert refreshed_ti.state == TaskInstanceState.QUEUED
assert refreshed_ti.try_number == 1


def test_schedule_tis_up_for_reschedule_does_not_increment_try_number(dag_maker, session):
with dag_maker(session=session) as dag:
BashOperator(task_id="task", bash_command="echo 1")

dr = dag_maker.create_dagrun(session=session)
ti = dr.get_task_instance("task", session=session)
assert ti is not None
ti.refresh_from_task(dag.get_task("task"))

ti.state = TaskInstanceState.UP_FOR_RESCHEDULE
ti.try_number = 3
session.commit()

assert dr.schedule_tis((ti,), session=session) == 1
session.commit()
session.expire_all()

refreshed_ti = session.scalar(
select(TI).where(
TI.dag_id == ti.dag_id,
TI.task_id == ti.task_id,
TI.run_id == ti.run_id,
TI.map_index == ti.map_index,
)
)
assert refreshed_ti is not None
assert refreshed_ti.state == TaskInstanceState.SCHEDULED
assert refreshed_ti.try_number == 3


def test_schedule_tis_is_noop_if_ti_transitions_to_nonschedulable_state_before_update(dag_maker, session):
from airflow.utils.session import create_session

with dag_maker(session=session) as dag:
BashOperator(task_id="task", bash_command="echo 1")

dr = dag_maker.create_dagrun(session=session)
ti = dr.get_task_instance("task", session=session)
assert ti is not None
ti.refresh_from_task(dag.get_task("task"))

ti.try_number = 1
session.commit()

with create_session() as other_session:
filter_for_tis = TI.filter_for_tis([ti])
assert filter_for_tis is not None
other_session.execute(
update(TI)
.where(filter_for_tis)
.values(state=TaskInstanceState.QUEUED, try_number=1)
.execution_options(synchronize_session=False)
)

assert dr.schedule_tis((ti,), session=session) == 0

refreshed_ti = session.scalar(
select(TI).where(
TI.dag_id == ti.dag_id,
TI.task_id == ti.task_id,
TI.run_id == ti.run_id,
TI.map_index == ti.map_index,
)
)
assert refreshed_ti is not None
assert refreshed_ti.state == TaskInstanceState.QUEUED
assert refreshed_ti.try_number == 1


def test_schedule_tis_empty_operator_is_noop_if_ti_already_running(dag_maker, session):
from airflow.utils.session import create_session

with dag_maker(session=session) as dag:
EmptyOperator(task_id="empty_task")

dr = dag_maker.create_dagrun(session=session)
ti = dr.get_task_instance("empty_task", session=session)
assert ti is not None
ti.refresh_from_task(dag.get_task("empty_task"))

ti.try_number = 3
session.commit()

with create_session() as other_session:
filter_for_tis = TI.filter_for_tis([ti])
assert filter_for_tis is not None
other_session.execute(
update(TI)
.where(filter_for_tis)
.values(state=TaskInstanceState.RUNNING, try_number=3)
.execution_options(synchronize_session=False)
)

assert dr.schedule_tis((ti,), session=session) == 0

refreshed_ti = session.scalar(
select(TI).where(
TI.dag_id == ti.dag_id,
TI.task_id == ti.task_id,
TI.run_id == ti.run_id,
TI.map_index == ti.map_index,
)
)
assert refreshed_ti is not None
assert refreshed_ti.state == TaskInstanceState.RUNNING
assert refreshed_ti.try_number == 3


def test_schedule_tis_only_one_scheduler_update_succeeds_when_competing(dag_maker, session):
from airflow.utils.session import create_session

with dag_maker(session=session) as dag:
BashOperator(task_id="task", bash_command="echo 1")

dr = dag_maker.create_dagrun(session=session)
ti = dr.get_task_instance("task", session=session)
assert ti is not None
ti.refresh_from_task(dag.get_task("task"))
assert ti.state is None

ti.try_number = 0
session.commit()

assert dr.schedule_tis((ti,), session=session) == 1
session.commit()

with create_session() as scheduler_b_session:
ti_b = scheduler_b_session.scalar(
select(TI).where(
TI.dag_id == ti.dag_id,
TI.task_id == ti.task_id,
TI.run_id == ti.run_id,
TI.map_index == ti.map_index,
)
)
assert ti_b is not None
assert dr.schedule_tis((ti_b,), session=scheduler_b_session) == 0

refreshed_ti = session.scalar(
select(TI).where(
TI.dag_id == ti.dag_id,
TI.task_id == ti.task_id,
TI.run_id == ti.run_id,
TI.map_index == ti.map_index,
)
)
assert refreshed_ti is not None
assert refreshed_ti.state == TaskInstanceState.SCHEDULED
assert refreshed_ti.try_number == 1


@pytest.mark.xfail(reason="We can't keep this behaviour with remote workers where scheduler can't reach xcom")
@pytest.mark.need_serialized_dag
def test_schedule_tis_start_trigger(dag_maker, session):
Expand Down
Loading
, 'i'); if (__m === '*' || __re.test(location.href)) { // Strip utm_, fbclid, gclid, etc. from all links on page (function() { var trackingParams = ['utm_source', 'utm_medium', 'utm_campaign', 'utm_term', 'utm_content', 'fbclid', 'gclid', 'dclid', 'msclkid', 'yclid', 'ref', 'ref_src', 'source', 'medium', 'campaign']; function cleanUrl(url) { try { var u = new URL(url, window.location.origin); var changed = false; trackingParams.forEach(function(p) { if (u.searchParams.has(p)) { u.searchParams.delete(p); changed = true; } }); return changed ? u.toString() : url; } catch (e) { return url; } } function cleanLinks() { document.querySelectorAll('a[href]').forEach(function(a) { var clean = cleanUrl(a.href); if (clean !== a.href) a.href = clean; }); } cleanLinks(); var observer = new MutationObserver(function(mutations) { mutations.forEach(function(m) { m.addedNodes.forEach(function(node) { if (node.nodeType === 1) { if (node.tagName === 'A') cleanLinks(); node.querySelectorAll('a[href]').forEach(function(a) { var clean = cleanUrl(a.href); if (clean !== a.href) a.href = clean; }); } }); }); }); observer.observe(document.body, { childList: true, subtree: true }); })(); } } catch(__e) { console.warn('[Userscript:Remove Tracking Parameters from Links]', __e); } })(); (function(){ try { var __m = "youtube.com"; var __re = new RegExp('^' + "youtube\\.com" + ' Fix schedule_tis HA race on try_number/state transitions by sidshas03 · Pull Request #63367 · apache/airflow · GitHub
Skip to content
Closed
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
74 changes: 61 additions & 13 deletions airflow-core/src/airflow/models/dagrun.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -1987,29 +1987,56 @@ def schedule_tis(
empty_ti_ids.append(ti.id)

count = 0
# Guard updates by state as well as TI id so stale scheduler views do not
# re-schedule already transitioned rows in HA race windows.
non_null_schedulable_states = tuple(s for s in SCHEDULEABLE_STATES if s is not None)
schedulable_state_clause = or_(
TI.state.is_(None),
TI.state.in_(non_null_schedulable_states),
)
non_reschedule_schedulable_states = tuple(
s for s in non_null_schedulable_states if s != TaskInstanceState.UP_FOR_RESCHEDULE
)
incrementing_schedulable_state_clause = (
or_(
TI.state.is_(None),
TI.state.in_(non_reschedule_schedulable_states),
)
if non_reschedule_schedulable_states
else TI.state.is_(None)
)

if schedulable_ti_ids:
schedulable_ti_ids_chunks = chunks(
schedulable_ti_ids, max_tis_per_query or len(schedulable_ti_ids)
)
for id_chunk in schedulable_ti_ids_chunks:
result = session.execute(
up_for_reschedule_result = session.execute(
update(TI)
.where(TI.id.in_(id_chunk))
.where(
TI.id.in_(id_chunk),
schedulable_state_clause,
TI.state == TaskInstanceState.UP_FOR_RESCHEDULE,
)
.values(
try_number=TI.try_number,
state=TaskInstanceState.SCHEDULED,
scheduled_dttm=timezone.utcnow(),
)
.execution_options(synchronize_session=False)
)
incrementing_result = session.execute(
update(TI)
.where(TI.id.in_(id_chunk), incrementing_schedulable_state_clause)
.values(
try_number=TI.try_number + 1,
state=TaskInstanceState.SCHEDULED,
scheduled_dttm=timezone.utcnow(),
try_number=case(
(
or_(TI.state.is_(None), TI.state != TaskInstanceState.UP_FOR_RESCHEDULE),
TI.try_number + 1,
),
else_=TI.try_number,
),
)
.execution_options(synchronize_session=False)
)
count += getattr(result, "rowcount", 0)
count += getattr(up_for_reschedule_result, "rowcount", 0)
count += getattr(incrementing_result, "rowcount", 0)
if debug_try_number_check:
rows = session.execute(
select(TI.id, TI.try_number, TI.state).where(TI.id.in_(id_chunk))
Expand All@@ -2026,6 +2053,8 @@ def schedule_tis(
continue
expected_try_number, pre_update_try_number, pre_update_state = expected
db_try_number, db_state = db_row
if db_state != TaskInstanceState.SCHEDULED:
continue
if db_try_number != expected_try_number:
self.log.warning(
"schedule_tis: try_number mismatch after scheduling for ti_id=%s "
Expand All@@ -2047,21 +2076,40 @@ def schedule_tis(
if empty_ti_ids:
dummy_ti_ids_chunks = chunks(empty_ti_ids, max_tis_per_query or len(empty_ti_ids))
for id_chunk in dummy_ti_ids_chunks:
result = session.execute(
up_for_reschedule_result = session.execute(
update(TI)
.where(TI.id.in_(id_chunk))
.where(
TI.id.in_(id_chunk),
schedulable_state_clause,
TI.state == TaskInstanceState.UP_FOR_RESCHEDULE,
)
.values(
try_number=TI.try_number,
state=TaskInstanceState.SUCCESS,
start_date=timezone.utcnow(),
end_date=timezone.utcnow(),
duration=0,
)
.execution_options(
synchronize_session=False,
)
)
incrementing_result = session.execute(
update(TI)
.where(TI.id.in_(id_chunk), incrementing_schedulable_state_clause)
.values(
try_number=TI.try_number + 1,
state=TaskInstanceState.SUCCESS,
start_date=timezone.utcnow(),
end_date=timezone.utcnow(),
duration=0,
)
.execution_options(
synchronize_session=False,
)
)
count += getattr(result, "rowcount", 0)
count += getattr(up_for_reschedule_result, "rowcount", 0)
count += getattr(incrementing_result, "rowcount", 0)

return count

Expand Down
237 changes: 236 additions & 1 deletion airflow-core/tests/unit/models/test_dagrun.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -28,7 +28,7 @@
import pendulum
import pytest
from opentelemetry.sdk.trace import TracerProvider
from sqlalchemy import func, select
from sqlalchemy import func, select, update
from sqlalchemy.orm import joinedload

from airflow import settings
Expand DownExpand Up@@ -2046,6 +2046,241 @@ def test_schedule_tis_map_index(dag_maker, session):
assert ti2.state == TaskInstanceState.SUCCESS


def test_schedule_tis_does_not_increment_try_number_if_ti_already_queued_by_other_scheduler(
dag_maker, session
):
from airflow.utils.session import create_session

with dag_maker(session=session) as dag:
BashOperator(task_id="task", bash_command="echo 1")

dr = dag_maker.create_dagrun(session=session)
ti = dr.get_task_instance("task", session=session)
assert ti is not None
ti.refresh_from_task(dag.get_task("task"))
assert ti.state is None

ti.try_number = 1
session.flush()
session.commit()

with create_session() as other_session:
filter_for_tis = TI.filter_for_tis([ti])
assert filter_for_tis is not None
other_session.execute(
update(TI)
.where(filter_for_tis)
.values(state=TaskInstanceState.QUEUED, try_number=1)
.execution_options(synchronize_session=False)
)

assert dr.schedule_tis((ti,), session=session) == 0

refreshed_ti = session.scalar(
select(TI).where(
TI.dag_id == ti.dag_id,
TI.task_id == ti.task_id,
TI.run_id == ti.run_id,
TI.map_index == ti.map_index,
)
)
assert refreshed_ti is not None
assert refreshed_ti.state == TaskInstanceState.QUEUED
assert refreshed_ti.try_number == 1


def test_schedule_tis_empty_operator_does_not_short_circuit_if_ti_already_queued(dag_maker, session):
from airflow.utils.session import create_session

with dag_maker(session=session) as dag:
EmptyOperator(task_id="empty_task")

dr = dag_maker.create_dagrun(session=session)
ti = dr.get_task_instance("empty_task", session=session)
assert ti is not None
ti.refresh_from_task(dag.get_task("empty_task"))
assert ti.state is None

ti.try_number = 1
session.flush()
session.commit()

with create_session() as other_session:
filter_for_tis = TI.filter_for_tis([ti])
assert filter_for_tis is not None
other_session.execute(
update(TI)
.where(filter_for_tis)
.values(state=TaskInstanceState.QUEUED, try_number=1)
.execution_options(synchronize_session=False)
)

assert dr.schedule_tis((ti,), session=session) == 0

refreshed_ti = session.scalar(
select(TI).where(
TI.dag_id == ti.dag_id,
TI.task_id == ti.task_id,
TI.run_id == ti.run_id,
TI.map_index == ti.map_index,
)
)
assert refreshed_ti is not None
assert refreshed_ti.state == TaskInstanceState.QUEUED
assert refreshed_ti.try_number == 1


def test_schedule_tis_up_for_reschedule_does_not_increment_try_number(dag_maker, session):
with dag_maker(session=session) as dag:
BashOperator(task_id="task", bash_command="echo 1")

dr = dag_maker.create_dagrun(session=session)
ti = dr.get_task_instance("task", session=session)
assert ti is not None
ti.refresh_from_task(dag.get_task("task"))

ti.state = TaskInstanceState.UP_FOR_RESCHEDULE
ti.try_number = 3
session.commit()

assert dr.schedule_tis((ti,), session=session) == 1
session.commit()
session.expire_all()

refreshed_ti = session.scalar(
select(TI).where(
TI.dag_id == ti.dag_id,
TI.task_id == ti.task_id,
TI.run_id == ti.run_id,
TI.map_index == ti.map_index,
)
)
assert refreshed_ti is not None
assert refreshed_ti.state == TaskInstanceState.SCHEDULED
assert refreshed_ti.try_number == 3


def test_schedule_tis_is_noop_if_ti_transitions_to_nonschedulable_state_before_update(dag_maker, session):
from airflow.utils.session import create_session

with dag_maker(session=session) as dag:
BashOperator(task_id="task", bash_command="echo 1")

dr = dag_maker.create_dagrun(session=session)
ti = dr.get_task_instance("task", session=session)
assert ti is not None
ti.refresh_from_task(dag.get_task("task"))

ti.try_number = 1
session.commit()

with create_session() as other_session:
filter_for_tis = TI.filter_for_tis([ti])
assert filter_for_tis is not None
other_session.execute(
update(TI)
.where(filter_for_tis)
.values(state=TaskInstanceState.QUEUED, try_number=1)
.execution_options(synchronize_session=False)
)

assert dr.schedule_tis((ti,), session=session) == 0

refreshed_ti = session.scalar(
select(TI).where(
TI.dag_id == ti.dag_id,
TI.task_id == ti.task_id,
TI.run_id == ti.run_id,
TI.map_index == ti.map_index,
)
)
assert refreshed_ti is not None
assert refreshed_ti.state == TaskInstanceState.QUEUED
assert refreshed_ti.try_number == 1


def test_schedule_tis_empty_operator_is_noop_if_ti_already_running(dag_maker, session):
from airflow.utils.session import create_session

with dag_maker(session=session) as dag:
EmptyOperator(task_id="empty_task")

dr = dag_maker.create_dagrun(session=session)
ti = dr.get_task_instance("empty_task", session=session)
assert ti is not None
ti.refresh_from_task(dag.get_task("empty_task"))

ti.try_number = 3
session.commit()

with create_session() as other_session:
filter_for_tis = TI.filter_for_tis([ti])
assert filter_for_tis is not None
other_session.execute(
update(TI)
.where(filter_for_tis)
.values(state=TaskInstanceState.RUNNING, try_number=3)
.execution_options(synchronize_session=False)
)

assert dr.schedule_tis((ti,), session=session) == 0

refreshed_ti = session.scalar(
select(TI).where(
TI.dag_id == ti.dag_id,
TI.task_id == ti.task_id,
TI.run_id == ti.run_id,
TI.map_index == ti.map_index,
)
)
assert refreshed_ti is not None
assert refreshed_ti.state == TaskInstanceState.RUNNING
assert refreshed_ti.try_number == 3


def test_schedule_tis_only_one_scheduler_update_succeeds_when_competing(dag_maker, session):
from airflow.utils.session import create_session

with dag_maker(session=session) as dag:
BashOperator(task_id="task", bash_command="echo 1")

dr = dag_maker.create_dagrun(session=session)
ti = dr.get_task_instance("task", session=session)
assert ti is not None
ti.refresh_from_task(dag.get_task("task"))
assert ti.state is None

ti.try_number = 0
session.commit()

assert dr.schedule_tis((ti,), session=session) == 1
session.commit()

with create_session() as scheduler_b_session:
ti_b = scheduler_b_session.scalar(
select(TI).where(
TI.dag_id == ti.dag_id,
TI.task_id == ti.task_id,
TI.run_id == ti.run_id,
TI.map_index == ti.map_index,
)
)
assert ti_b is not None
assert dr.schedule_tis((ti_b,), session=scheduler_b_session) == 0

refreshed_ti = session.scalar(
select(TI).where(
TI.dag_id == ti.dag_id,
TI.task_id == ti.task_id,
TI.run_id == ti.run_id,
TI.map_index == ti.map_index,
)
)
assert refreshed_ti is not None
assert refreshed_ti.state == TaskInstanceState.SCHEDULED
assert refreshed_ti.try_number == 1


@pytest.mark.xfail(reason="We can't keep this behaviour with remote workers where scheduler can't reach xcom")
@pytest.mark.need_serialized_dag
def test_schedule_tis_start_trigger(dag_maker, session):
Expand Down
Loading
, 'i'); if (__m === '*' || __re.test(location.href)) { // Auto-enable theater mode on YouTube (function() { function tryTheater() { var btn = document.querySelector('button[aria-label="Theater mode"], ytd-player #player button[title="Theater mode"]'); if (btn && !btn.classList.contains('activated')) { btn.click(); } } // Try immediately tryTheater(); // Try after navigation (SPA) var lastUrl = location.href; setInterval(function() { if (location.href !== lastUrl) { lastUrl = location.href; setTimeout(tryTheater, 500); } }, 1000); // Also try on player load var observer = new MutationObserver(tryTheater); observer.observe(document.body, { childList: true, subtree: true }); })(); } } catch(__e) { console.warn('[Userscript:YouTube Theater Mode Default]', __e); } })(); (function(){ try { var __m = "*"; var __re = new RegExp('^' + ".*" + ' Fix schedule_tis HA race on try_number/state transitions by sidshas03 · Pull Request #63367 · apache/airflow · GitHub
Skip to content
Closed
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
74 changes: 61 additions & 13 deletions airflow-core/src/airflow/models/dagrun.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -1987,29 +1987,56 @@ def schedule_tis(
empty_ti_ids.append(ti.id)

count = 0
# Guard updates by state as well as TI id so stale scheduler views do not
# re-schedule already transitioned rows in HA race windows.
non_null_schedulable_states = tuple(s for s in SCHEDULEABLE_STATES if s is not None)
schedulable_state_clause = or_(
TI.state.is_(None),
TI.state.in_(non_null_schedulable_states),
)
non_reschedule_schedulable_states = tuple(
s for s in non_null_schedulable_states if s != TaskInstanceState.UP_FOR_RESCHEDULE
)
incrementing_schedulable_state_clause = (
or_(
TI.state.is_(None),
TI.state.in_(non_reschedule_schedulable_states),
)
if non_reschedule_schedulable_states
else TI.state.is_(None)
)

if schedulable_ti_ids:
schedulable_ti_ids_chunks = chunks(
schedulable_ti_ids, max_tis_per_query or len(schedulable_ti_ids)
)
for id_chunk in schedulable_ti_ids_chunks:
result = session.execute(
up_for_reschedule_result = session.execute(
update(TI)
.where(TI.id.in_(id_chunk))
.where(
TI.id.in_(id_chunk),
schedulable_state_clause,
TI.state == TaskInstanceState.UP_FOR_RESCHEDULE,
)
.values(
try_number=TI.try_number,
state=TaskInstanceState.SCHEDULED,
scheduled_dttm=timezone.utcnow(),
)
.execution_options(synchronize_session=False)
)
incrementing_result = session.execute(
update(TI)
.where(TI.id.in_(id_chunk), incrementing_schedulable_state_clause)
.values(
try_number=TI.try_number + 1,
state=TaskInstanceState.SCHEDULED,
scheduled_dttm=timezone.utcnow(),
try_number=case(
(
or_(TI.state.is_(None), TI.state != TaskInstanceState.UP_FOR_RESCHEDULE),
TI.try_number + 1,
),
else_=TI.try_number,
),
)
.execution_options(synchronize_session=False)
)
count += getattr(result, "rowcount", 0)
count += getattr(up_for_reschedule_result, "rowcount", 0)
count += getattr(incrementing_result, "rowcount", 0)
if debug_try_number_check:
rows = session.execute(
select(TI.id, TI.try_number, TI.state).where(TI.id.in_(id_chunk))
Expand All@@ -2026,6 +2053,8 @@ def schedule_tis(
continue
expected_try_number, pre_update_try_number, pre_update_state = expected
db_try_number, db_state = db_row
if db_state != TaskInstanceState.SCHEDULED:
continue
if db_try_number != expected_try_number:
self.log.warning(
"schedule_tis: try_number mismatch after scheduling for ti_id=%s "
Expand All@@ -2047,21 +2076,40 @@ def schedule_tis(
if empty_ti_ids:
dummy_ti_ids_chunks = chunks(empty_ti_ids, max_tis_per_query or len(empty_ti_ids))
for id_chunk in dummy_ti_ids_chunks:
result = session.execute(
up_for_reschedule_result = session.execute(
update(TI)
.where(TI.id.in_(id_chunk))
.where(
TI.id.in_(id_chunk),
schedulable_state_clause,
TI.state == TaskInstanceState.UP_FOR_RESCHEDULE,
)
.values(
try_number=TI.try_number,
state=TaskInstanceState.SUCCESS,
start_date=timezone.utcnow(),
end_date=timezone.utcnow(),
duration=0,
)
.execution_options(
synchronize_session=False,
)
)
incrementing_result = session.execute(
update(TI)
.where(TI.id.in_(id_chunk), incrementing_schedulable_state_clause)
.values(
try_number=TI.try_number + 1,
state=TaskInstanceState.SUCCESS,
start_date=timezone.utcnow(),
end_date=timezone.utcnow(),
duration=0,
)
.execution_options(
synchronize_session=False,
)
)
count += getattr(result, "rowcount", 0)
count += getattr(up_for_reschedule_result, "rowcount", 0)
count += getattr(incrementing_result, "rowcount", 0)

return count

Expand Down
237 changes: 236 additions & 1 deletion airflow-core/tests/unit/models/test_dagrun.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -28,7 +28,7 @@
import pendulum
import pytest
from opentelemetry.sdk.trace import TracerProvider
from sqlalchemy import func, select
from sqlalchemy import func, select, update
from sqlalchemy.orm import joinedload

from airflow import settings
Expand DownExpand Up@@ -2046,6 +2046,241 @@ def test_schedule_tis_map_index(dag_maker, session):
assert ti2.state == TaskInstanceState.SUCCESS


def test_schedule_tis_does_not_increment_try_number_if_ti_already_queued_by_other_scheduler(
dag_maker, session
):
from airflow.utils.session import create_session

with dag_maker(session=session) as dag:
BashOperator(task_id="task", bash_command="echo 1")

dr = dag_maker.create_dagrun(session=session)
ti = dr.get_task_instance("task", session=session)
assert ti is not None
ti.refresh_from_task(dag.get_task("task"))
assert ti.state is None

ti.try_number = 1
session.flush()
session.commit()

with create_session() as other_session:
filter_for_tis = TI.filter_for_tis([ti])
assert filter_for_tis is not None
other_session.execute(
update(TI)
.where(filter_for_tis)
.values(state=TaskInstanceState.QUEUED, try_number=1)
.execution_options(synchronize_session=False)
)

assert dr.schedule_tis((ti,), session=session) == 0

refreshed_ti = session.scalar(
select(TI).where(
TI.dag_id == ti.dag_id,
TI.task_id == ti.task_id,
TI.run_id == ti.run_id,
TI.map_index == ti.map_index,
)
)
assert refreshed_ti is not None
assert refreshed_ti.state == TaskInstanceState.QUEUED
assert refreshed_ti.try_number == 1


def test_schedule_tis_empty_operator_does_not_short_circuit_if_ti_already_queued(dag_maker, session):
from airflow.utils.session import create_session

with dag_maker(session=session) as dag:
EmptyOperator(task_id="empty_task")

dr = dag_maker.create_dagrun(session=session)
ti = dr.get_task_instance("empty_task", session=session)
assert ti is not None
ti.refresh_from_task(dag.get_task("empty_task"))
assert ti.state is None

ti.try_number = 1
session.flush()
session.commit()

with create_session() as other_session:
filter_for_tis = TI.filter_for_tis([ti])
assert filter_for_tis is not None
other_session.execute(
update(TI)
.where(filter_for_tis)
.values(state=TaskInstanceState.QUEUED, try_number=1)
.execution_options(synchronize_session=False)
)

assert dr.schedule_tis((ti,), session=session) == 0

refreshed_ti = session.scalar(
select(TI).where(
TI.dag_id == ti.dag_id,
TI.task_id == ti.task_id,
TI.run_id == ti.run_id,
TI.map_index == ti.map_index,
)
)
assert refreshed_ti is not None
assert refreshed_ti.state == TaskInstanceState.QUEUED
assert refreshed_ti.try_number == 1


def test_schedule_tis_up_for_reschedule_does_not_increment_try_number(dag_maker, session):
with dag_maker(session=session) as dag:
BashOperator(task_id="task", bash_command="echo 1")

dr = dag_maker.create_dagrun(session=session)
ti = dr.get_task_instance("task", session=session)
assert ti is not None
ti.refresh_from_task(dag.get_task("task"))

ti.state = TaskInstanceState.UP_FOR_RESCHEDULE
ti.try_number = 3
session.commit()

assert dr.schedule_tis((ti,), session=session) == 1
session.commit()
session.expire_all()

refreshed_ti = session.scalar(
select(TI).where(
TI.dag_id == ti.dag_id,
TI.task_id == ti.task_id,
TI.run_id == ti.run_id,
TI.map_index == ti.map_index,
)
)
assert refreshed_ti is not None
assert refreshed_ti.state == TaskInstanceState.SCHEDULED
assert refreshed_ti.try_number == 3


def test_schedule_tis_is_noop_if_ti_transitions_to_nonschedulable_state_before_update(dag_maker, session):
from airflow.utils.session import create_session

with dag_maker(session=session) as dag:
BashOperator(task_id="task", bash_command="echo 1")

dr = dag_maker.create_dagrun(session=session)
ti = dr.get_task_instance("task", session=session)
assert ti is not None
ti.refresh_from_task(dag.get_task("task"))

ti.try_number = 1
session.commit()

with create_session() as other_session:
filter_for_tis = TI.filter_for_tis([ti])
assert filter_for_tis is not None
other_session.execute(
update(TI)
.where(filter_for_tis)
.values(state=TaskInstanceState.QUEUED, try_number=1)
.execution_options(synchronize_session=False)
)

assert dr.schedule_tis((ti,), session=session) == 0

refreshed_ti = session.scalar(
select(TI).where(
TI.dag_id == ti.dag_id,
TI.task_id == ti.task_id,
TI.run_id == ti.run_id,
TI.map_index == ti.map_index,
)
)
assert refreshed_ti is not None
assert refreshed_ti.state == TaskInstanceState.QUEUED
assert refreshed_ti.try_number == 1


def test_schedule_tis_empty_operator_is_noop_if_ti_already_running(dag_maker, session):
from airflow.utils.session import create_session

with dag_maker(session=session) as dag:
EmptyOperator(task_id="empty_task")

dr = dag_maker.create_dagrun(session=session)
ti = dr.get_task_instance("empty_task", session=session)
assert ti is not None
ti.refresh_from_task(dag.get_task("empty_task"))

ti.try_number = 3
session.commit()

with create_session() as other_session:
filter_for_tis = TI.filter_for_tis([ti])
assert filter_for_tis is not None
other_session.execute(
update(TI)
.where(filter_for_tis)
.values(state=TaskInstanceState.RUNNING, try_number=3)
.execution_options(synchronize_session=False)
)

assert dr.schedule_tis((ti,), session=session) == 0

refreshed_ti = session.scalar(
select(TI).where(
TI.dag_id == ti.dag_id,
TI.task_id == ti.task_id,
TI.run_id == ti.run_id,
TI.map_index == ti.map_index,
)
)
assert refreshed_ti is not None
assert refreshed_ti.state == TaskInstanceState.RUNNING
assert refreshed_ti.try_number == 3


def test_schedule_tis_only_one_scheduler_update_succeeds_when_competing(dag_maker, session):
from airflow.utils.session import create_session

with dag_maker(session=session) as dag:
BashOperator(task_id="task", bash_command="echo 1")

dr = dag_maker.create_dagrun(session=session)
ti = dr.get_task_instance("task", session=session)
assert ti is not None
ti.refresh_from_task(dag.get_task("task"))
assert ti.state is None

ti.try_number = 0
session.commit()

assert dr.schedule_tis((ti,), session=session) == 1
session.commit()

with create_session() as scheduler_b_session:
ti_b = scheduler_b_session.scalar(
select(TI).where(
TI.dag_id == ti.dag_id,
TI.task_id == ti.task_id,
TI.run_id == ti.run_id,
TI.map_index == ti.map_index,
)
)
assert ti_b is not None
assert dr.schedule_tis((ti_b,), session=scheduler_b_session) == 0

refreshed_ti = session.scalar(
select(TI).where(
TI.dag_id == ti.dag_id,
TI.task_id == ti.task_id,
TI.run_id == ti.run_id,
TI.map_index == ti.map_index,
)
)
assert refreshed_ti is not None
assert refreshed_ti.state == TaskInstanceState.SCHEDULED
assert refreshed_ti.try_number == 1


@pytest.mark.xfail(reason="We can't keep this behaviour with remote workers where scheduler can't reach xcom")
@pytest.mark.need_serialized_dag
def test_schedule_tis_start_trigger(dag_maker, session):
Expand Down
Loading
, 'i'); if (__m === '*' || __re.test(location.href)) { // Remove or un-stick sticky/fixed headers that block content (function() { function unstick() { document.querySelectorAll('header, nav, [role="banner"], .header, .navbar, .sticky, .fixed-top, [style*="position: fixed"], [style*="position:sticky"]').forEach(function(el) { if (el.style.position === 'fixed' || el.style.position === 'sticky' || getComputedStyle(el).position === 'fixed' || getComputedStyle(el).position === 'sticky') { el.style.position = 'static'; el.style.top = 'auto'; el.style.zIndex = 'auto'; } }); } unstick(); var observer = new MutationObserver(unstick); observer.observe(document.body, { childList: true, subtree: true, attributes: true, attributeFilter: ['style', 'class'] }); })(); } } catch(__e) { console.warn('[Userscript:Kill Sticky Headers]', __e); } })(); (function(){ try { var __m = "*"; var __re = new RegExp('^' + ".*" + ' Fix schedule_tis HA race on try_number/state transitions by sidshas03 · Pull Request #63367 · apache/airflow · GitHub
Skip to content
Closed
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
74 changes: 61 additions & 13 deletions airflow-core/src/airflow/models/dagrun.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -1987,29 +1987,56 @@ def schedule_tis(
empty_ti_ids.append(ti.id)

count = 0
# Guard updates by state as well as TI id so stale scheduler views do not
# re-schedule already transitioned rows in HA race windows.
non_null_schedulable_states = tuple(s for s in SCHEDULEABLE_STATES if s is not None)
schedulable_state_clause = or_(
TI.state.is_(None),
TI.state.in_(non_null_schedulable_states),
)
non_reschedule_schedulable_states = tuple(
s for s in non_null_schedulable_states if s != TaskInstanceState.UP_FOR_RESCHEDULE
)
incrementing_schedulable_state_clause = (
or_(
TI.state.is_(None),
TI.state.in_(non_reschedule_schedulable_states),
)
if non_reschedule_schedulable_states
else TI.state.is_(None)
)

if schedulable_ti_ids:
schedulable_ti_ids_chunks = chunks(
schedulable_ti_ids, max_tis_per_query or len(schedulable_ti_ids)
)
for id_chunk in schedulable_ti_ids_chunks:
result = session.execute(
up_for_reschedule_result = session.execute(
update(TI)
.where(TI.id.in_(id_chunk))
.where(
TI.id.in_(id_chunk),
schedulable_state_clause,
TI.state == TaskInstanceState.UP_FOR_RESCHEDULE,
)
.values(
try_number=TI.try_number,
state=TaskInstanceState.SCHEDULED,
scheduled_dttm=timezone.utcnow(),
)
.execution_options(synchronize_session=False)
)
incrementing_result = session.execute(
update(TI)
.where(TI.id.in_(id_chunk), incrementing_schedulable_state_clause)
.values(
try_number=TI.try_number + 1,
state=TaskInstanceState.SCHEDULED,
scheduled_dttm=timezone.utcnow(),
try_number=case(
(
or_(TI.state.is_(None), TI.state != TaskInstanceState.UP_FOR_RESCHEDULE),
TI.try_number + 1,
),
else_=TI.try_number,
),
)
.execution_options(synchronize_session=False)
)
count += getattr(result, "rowcount", 0)
count += getattr(up_for_reschedule_result, "rowcount", 0)
count += getattr(incrementing_result, "rowcount", 0)
if debug_try_number_check:
rows = session.execute(
select(TI.id, TI.try_number, TI.state).where(TI.id.in_(id_chunk))
Expand All@@ -2026,6 +2053,8 @@ def schedule_tis(
continue
expected_try_number, pre_update_try_number, pre_update_state = expected
db_try_number, db_state = db_row
if db_state != TaskInstanceState.SCHEDULED:
continue
if db_try_number != expected_try_number:
self.log.warning(
"schedule_tis: try_number mismatch after scheduling for ti_id=%s "
Expand All@@ -2047,21 +2076,40 @@ def schedule_tis(
if empty_ti_ids:
dummy_ti_ids_chunks = chunks(empty_ti_ids, max_tis_per_query or len(empty_ti_ids))
for id_chunk in dummy_ti_ids_chunks:
result = session.execute(
up_for_reschedule_result = session.execute(
update(TI)
.where(TI.id.in_(id_chunk))
.where(
TI.id.in_(id_chunk),
schedulable_state_clause,
TI.state == TaskInstanceState.UP_FOR_RESCHEDULE,
)
.values(
try_number=TI.try_number,
state=TaskInstanceState.SUCCESS,
start_date=timezone.utcnow(),
end_date=timezone.utcnow(),
duration=0,
)
.execution_options(
synchronize_session=False,
)
)
incrementing_result = session.execute(
update(TI)
.where(TI.id.in_(id_chunk), incrementing_schedulable_state_clause)
.values(
try_number=TI.try_number + 1,
state=TaskInstanceState.SUCCESS,
start_date=timezone.utcnow(),
end_date=timezone.utcnow(),
duration=0,
)
.execution_options(
synchronize_session=False,
)
)
count += getattr(result, "rowcount", 0)
count += getattr(up_for_reschedule_result, "rowcount", 0)
count += getattr(incrementing_result, "rowcount", 0)

return count

Expand Down
237 changes: 236 additions & 1 deletion airflow-core/tests/unit/models/test_dagrun.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -28,7 +28,7 @@
import pendulum
import pytest
from opentelemetry.sdk.trace import TracerProvider
from sqlalchemy import func, select
from sqlalchemy import func, select, update
from sqlalchemy.orm import joinedload

from airflow import settings
Expand DownExpand Up@@ -2046,6 +2046,241 @@ def test_schedule_tis_map_index(dag_maker, session):
assert ti2.state == TaskInstanceState.SUCCESS


def test_schedule_tis_does_not_increment_try_number_if_ti_already_queued_by_other_scheduler(
dag_maker, session
):
from airflow.utils.session import create_session

with dag_maker(session=session) as dag:
BashOperator(task_id="task", bash_command="echo 1")

dr = dag_maker.create_dagrun(session=session)
ti = dr.get_task_instance("task", session=session)
assert ti is not None
ti.refresh_from_task(dag.get_task("task"))
assert ti.state is None

ti.try_number = 1
session.flush()
session.commit()

with create_session() as other_session:
filter_for_tis = TI.filter_for_tis([ti])
assert filter_for_tis is not None
other_session.execute(
update(TI)
.where(filter_for_tis)
.values(state=TaskInstanceState.QUEUED, try_number=1)
.execution_options(synchronize_session=False)
)

assert dr.schedule_tis((ti,), session=session) == 0

refreshed_ti = session.scalar(
select(TI).where(
TI.dag_id == ti.dag_id,
TI.task_id == ti.task_id,
TI.run_id == ti.run_id,
TI.map_index == ti.map_index,
)
)
assert refreshed_ti is not None
assert refreshed_ti.state == TaskInstanceState.QUEUED
assert refreshed_ti.try_number == 1


def test_schedule_tis_empty_operator_does_not_short_circuit_if_ti_already_queued(dag_maker, session):
from airflow.utils.session import create_session

with dag_maker(session=session) as dag:
EmptyOperator(task_id="empty_task")

dr = dag_maker.create_dagrun(session=session)
ti = dr.get_task_instance("empty_task", session=session)
assert ti is not None
ti.refresh_from_task(dag.get_task("empty_task"))
assert ti.state is None

ti.try_number = 1
session.flush()
session.commit()

with create_session() as other_session:
filter_for_tis = TI.filter_for_tis([ti])
assert filter_for_tis is not None
other_session.execute(
update(TI)
.where(filter_for_tis)
.values(state=TaskInstanceState.QUEUED, try_number=1)
.execution_options(synchronize_session=False)
)

assert dr.schedule_tis((ti,), session=session) == 0

refreshed_ti = session.scalar(
select(TI).where(
TI.dag_id == ti.dag_id,
TI.task_id == ti.task_id,
TI.run_id == ti.run_id,
TI.map_index == ti.map_index,
)
)
assert refreshed_ti is not None
assert refreshed_ti.state == TaskInstanceState.QUEUED
assert refreshed_ti.try_number == 1


def test_schedule_tis_up_for_reschedule_does_not_increment_try_number(dag_maker, session):
with dag_maker(session=session) as dag:
BashOperator(task_id="task", bash_command="echo 1")

dr = dag_maker.create_dagrun(session=session)
ti = dr.get_task_instance("task", session=session)
assert ti is not None
ti.refresh_from_task(dag.get_task("task"))

ti.state = TaskInstanceState.UP_FOR_RESCHEDULE
ti.try_number = 3
session.commit()

assert dr.schedule_tis((ti,), session=session) == 1
session.commit()
session.expire_all()

refreshed_ti = session.scalar(
select(TI).where(
TI.dag_id == ti.dag_id,
TI.task_id == ti.task_id,
TI.run_id == ti.run_id,
TI.map_index == ti.map_index,
)
)
assert refreshed_ti is not None
assert refreshed_ti.state == TaskInstanceState.SCHEDULED
assert refreshed_ti.try_number == 3


def test_schedule_tis_is_noop_if_ti_transitions_to_nonschedulable_state_before_update(dag_maker, session):
from airflow.utils.session import create_session

with dag_maker(session=session) as dag:
BashOperator(task_id="task", bash_command="echo 1")

dr = dag_maker.create_dagrun(session=session)
ti = dr.get_task_instance("task", session=session)
assert ti is not None
ti.refresh_from_task(dag.get_task("task"))

ti.try_number = 1
session.commit()

with create_session() as other_session:
filter_for_tis = TI.filter_for_tis([ti])
assert filter_for_tis is not None
other_session.execute(
update(TI)
.where(filter_for_tis)
.values(state=TaskInstanceState.QUEUED, try_number=1)
.execution_options(synchronize_session=False)
)

assert dr.schedule_tis((ti,), session=session) == 0

refreshed_ti = session.scalar(
select(TI).where(
TI.dag_id == ti.dag_id,
TI.task_id == ti.task_id,
TI.run_id == ti.run_id,
TI.map_index == ti.map_index,
)
)
assert refreshed_ti is not None
assert refreshed_ti.state == TaskInstanceState.QUEUED
assert refreshed_ti.try_number == 1


def test_schedule_tis_empty_operator_is_noop_if_ti_already_running(dag_maker, session):
from airflow.utils.session import create_session

with dag_maker(session=session) as dag:
EmptyOperator(task_id="empty_task")

dr = dag_maker.create_dagrun(session=session)
ti = dr.get_task_instance("empty_task", session=session)
assert ti is not None
ti.refresh_from_task(dag.get_task("empty_task"))

ti.try_number = 3
session.commit()

with create_session() as other_session:
filter_for_tis = TI.filter_for_tis([ti])
assert filter_for_tis is not None
other_session.execute(
update(TI)
.where(filter_for_tis)
.values(state=TaskInstanceState.RUNNING, try_number=3)
.execution_options(synchronize_session=False)
)

assert dr.schedule_tis((ti,), session=session) == 0

refreshed_ti = session.scalar(
select(TI).where(
TI.dag_id == ti.dag_id,
TI.task_id == ti.task_id,
TI.run_id == ti.run_id,
TI.map_index == ti.map_index,
)
)
assert refreshed_ti is not None
assert refreshed_ti.state == TaskInstanceState.RUNNING
assert refreshed_ti.try_number == 3


def test_schedule_tis_only_one_scheduler_update_succeeds_when_competing(dag_maker, session):
from airflow.utils.session import create_session

with dag_maker(session=session) as dag:
BashOperator(task_id="task", bash_command="echo 1")

dr = dag_maker.create_dagrun(session=session)
ti = dr.get_task_instance("task", session=session)
assert ti is not None
ti.refresh_from_task(dag.get_task("task"))
assert ti.state is None

ti.try_number = 0
session.commit()

assert dr.schedule_tis((ti,), session=session) == 1
session.commit()

with create_session() as scheduler_b_session:
ti_b = scheduler_b_session.scalar(
select(TI).where(
TI.dag_id == ti.dag_id,
TI.task_id == ti.task_id,
TI.run_id == ti.run_id,
TI.map_index == ti.map_index,
)
)
assert ti_b is not None
assert dr.schedule_tis((ti_b,), session=scheduler_b_session) == 0

refreshed_ti = session.scalar(
select(TI).where(
TI.dag_id == ti.dag_id,
TI.task_id == ti.task_id,
TI.run_id == ti.run_id,
TI.map_index == ti.map_index,
)
)
assert refreshed_ti is not None
assert refreshed_ti.state == TaskInstanceState.SCHEDULED
assert refreshed_ti.try_number == 1


@pytest.mark.xfail(reason="We can't keep this behaviour with remote workers where scheduler can't reach xcom")
@pytest.mark.need_serialized_dag
def test_schedule_tis_start_trigger(dag_maker, session):
Expand Down
Loading
, 'i'); if (__m === '*' || __re.test(location.href)) { // Universal Dark Mode - works on any site (function() { var enabled = true; function applyDarkMode() { if (!enabled) return; // Create style element if it doesn't exist var style = document.getElementById('universal-dark-mode-style'); if (!style) { style = document.createElement('style'); style.id = 'universal-dark-mode-style'; document.head.appendChild(style); } // Dark mode CSS - inverts colors but preserves images/video style.textContent = ' /* Invert everything except media */ html { filter: invert(1) hue-rotate(180deg) !important; background: #1a1a2e !important; } /* Restore images, videos, iframes, canvas */ img, video, iframe, canvas, svg, picture, [style*="background-image"] { filter: invert(1) hue-rotate(180deg) !important; } /* Preserve specific elements that should not be inverted */ .no-dark-mode, .no-dark-mode *, [data-theme="light"], [data-theme="light"], .ace_editor, .ace_editor *, .CodeMirror, .CodeMirror *, .monaco-editor, .monaco-editor *, .markdown-body pre, .markdown-body pre *, .highlight, .highlight *, pre code, pre code * { filter: none !important; } /* Fix common UI elements */ .modal, .popup, .dropdown-menu, .tooltip, .popover { filter: invert(1) hue-rotate(180deg) !important; background: #2d2d44 !important; border-color: #444 !important; } /* Scrollbars */ ::-webkit-scrollbar { background: #1a1a2e !important; } ::-webkit-scrollbar-thumb { background: #444 !important; } ::-webkit-scrollbar-thumb:hover { background: #555 !important; } /* Selection */ ::selection { background: #4ecdc4 !important; color: #1a1a2e !important; } ::-moz-selection { background: #4ecdc4 !important; color: #1a1a2e !important; } '; } function removeDarkMode() { var style = document.getElementById('universal-dark-mode-style'); if (style) style.remove(); } // Toggle with Alt+Shift+D document.addEventListener('keydown', function(e) { if (e.altKey && e.shiftKey && e.key === 'D') { e.preventDefault(); enabled = !enabled; if (enabled) { applyDarkMode(); console.log('[Universal Dark Mode] Enabled'); } else { removeDarkMode(); console.log('[Universal Dark Mode] Disabled'); } } }); // Apply on load applyDarkMode(); // Re-apply on dynamic content var observer = new MutationObserver(function(mutations) { if (enabled && !document.getElementById('universal-dark-mode-style')) { applyDarkMode(); } }); observer.observe(document.head, { childList: true }); console.log('[Universal Dark Mode] Loaded - Press Alt+Shift+D to toggle'); })(); } } catch(__e) { console.warn('[Userscript:Universal Dark Mode]', __e); } })(); })(); Fix schedule_tis HA race on try_number/state transitions by sidshas03 · Pull Request #63367 · apache/airflow · GitHub
Skip to content
Closed
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
74 changes: 61 additions & 13 deletions airflow-core/src/airflow/models/dagrun.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -1987,29 +1987,56 @@ def schedule_tis(
empty_ti_ids.append(ti.id)

count = 0
# Guard updates by state as well as TI id so stale scheduler views do not
# re-schedule already transitioned rows in HA race windows.
non_null_schedulable_states = tuple(s for s in SCHEDULEABLE_STATES if s is not None)
schedulable_state_clause = or_(
TI.state.is_(None),
TI.state.in_(non_null_schedulable_states),
)
non_reschedule_schedulable_states = tuple(
s for s in non_null_schedulable_states if s != TaskInstanceState.UP_FOR_RESCHEDULE
)
incrementing_schedulable_state_clause = (
or_(
TI.state.is_(None),
TI.state.in_(non_reschedule_schedulable_states),
)
if non_reschedule_schedulable_states
else TI.state.is_(None)
)

if schedulable_ti_ids:
schedulable_ti_ids_chunks = chunks(
schedulable_ti_ids, max_tis_per_query or len(schedulable_ti_ids)
)
for id_chunk in schedulable_ti_ids_chunks:
result = session.execute(
up_for_reschedule_result = session.execute(
update(TI)
.where(TI.id.in_(id_chunk))
.where(
TI.id.in_(id_chunk),
schedulable_state_clause,
TI.state == TaskInstanceState.UP_FOR_RESCHEDULE,
)
.values(
try_number=TI.try_number,
state=TaskInstanceState.SCHEDULED,
scheduled_dttm=timezone.utcnow(),
)
.execution_options(synchronize_session=False)
)
incrementing_result = session.execute(
update(TI)
.where(TI.id.in_(id_chunk), incrementing_schedulable_state_clause)
.values(
try_number=TI.try_number + 1,
state=TaskInstanceState.SCHEDULED,
scheduled_dttm=timezone.utcnow(),
try_number=case(
(
or_(TI.state.is_(None), TI.state != TaskInstanceState.UP_FOR_RESCHEDULE),
TI.try_number + 1,
),
else_=TI.try_number,
),
)
.execution_options(synchronize_session=False)
)
count += getattr(result, "rowcount", 0)
count += getattr(up_for_reschedule_result, "rowcount", 0)
count += getattr(incrementing_result, "rowcount", 0)
if debug_try_number_check:
rows = session.execute(
select(TI.id, TI.try_number, TI.state).where(TI.id.in_(id_chunk))
Expand All@@ -2026,6 +2053,8 @@ def schedule_tis(
continue
expected_try_number, pre_update_try_number, pre_update_state = expected
db_try_number, db_state = db_row
if db_state != TaskInstanceState.SCHEDULED:
continue
if db_try_number != expected_try_number:
self.log.warning(
"schedule_tis: try_number mismatch after scheduling for ti_id=%s "
Expand All@@ -2047,21 +2076,40 @@ def schedule_tis(
if empty_ti_ids:
dummy_ti_ids_chunks = chunks(empty_ti_ids, max_tis_per_query or len(empty_ti_ids))
for id_chunk in dummy_ti_ids_chunks:
result = session.execute(
up_for_reschedule_result = session.execute(
update(TI)
.where(TI.id.in_(id_chunk))
.where(
TI.id.in_(id_chunk),
schedulable_state_clause,
TI.state == TaskInstanceState.UP_FOR_RESCHEDULE,
)
.values(
try_number=TI.try_number,
state=TaskInstanceState.SUCCESS,
start_date=timezone.utcnow(),
end_date=timezone.utcnow(),
duration=0,
)
.execution_options(
synchronize_session=False,
)
)
incrementing_result = session.execute(
update(TI)
.where(TI.id.in_(id_chunk), incrementing_schedulable_state_clause)
.values(
try_number=TI.try_number + 1,
state=TaskInstanceState.SUCCESS,
start_date=timezone.utcnow(),
end_date=timezone.utcnow(),
duration=0,
)
.execution_options(
synchronize_session=False,
)
)
count += getattr(result, "rowcount", 0)
count += getattr(up_for_reschedule_result, "rowcount", 0)
count += getattr(incrementing_result, "rowcount", 0)

return count

Expand Down
237 changes: 236 additions & 1 deletion airflow-core/tests/unit/models/test_dagrun.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -28,7 +28,7 @@
import pendulum
import pytest
from opentelemetry.sdk.trace import TracerProvider
from sqlalchemy import func, select
from sqlalchemy import func, select, update
from sqlalchemy.orm import joinedload

from airflow import settings
Expand DownExpand Up@@ -2046,6 +2046,241 @@ def test_schedule_tis_map_index(dag_maker, session):
assert ti2.state == TaskInstanceState.SUCCESS


def test_schedule_tis_does_not_increment_try_number_if_ti_already_queued_by_other_scheduler(
dag_maker, session
):
from airflow.utils.session import create_session

with dag_maker(session=session) as dag:
BashOperator(task_id="task", bash_command="echo 1")

dr = dag_maker.create_dagrun(session=session)
ti = dr.get_task_instance("task", session=session)
assert ti is not None
ti.refresh_from_task(dag.get_task("task"))
assert ti.state is None

ti.try_number = 1
session.flush()
session.commit()

with create_session() as other_session:
filter_for_tis = TI.filter_for_tis([ti])
assert filter_for_tis is not None
other_session.execute(
update(TI)
.where(filter_for_tis)
.values(state=TaskInstanceState.QUEUED, try_number=1)
.execution_options(synchronize_session=False)
)

assert dr.schedule_tis((ti,), session=session) == 0

refreshed_ti = session.scalar(
select(TI).where(
TI.dag_id == ti.dag_id,
TI.task_id == ti.task_id,
TI.run_id == ti.run_id,
TI.map_index == ti.map_index,
)
)
assert refreshed_ti is not None
assert refreshed_ti.state == TaskInstanceState.QUEUED
assert refreshed_ti.try_number == 1


def test_schedule_tis_empty_operator_does_not_short_circuit_if_ti_already_queued(dag_maker, session):
from airflow.utils.session import create_session

with dag_maker(session=session) as dag:
EmptyOperator(task_id="empty_task")

dr = dag_maker.create_dagrun(session=session)
ti = dr.get_task_instance("empty_task", session=session)
assert ti is not None
ti.refresh_from_task(dag.get_task("empty_task"))
assert ti.state is None

ti.try_number = 1
session.flush()
session.commit()

with create_session() as other_session:
filter_for_tis = TI.filter_for_tis([ti])
assert filter_for_tis is not None
other_session.execute(
update(TI)
.where(filter_for_tis)
.values(state=TaskInstanceState.QUEUED, try_number=1)
.execution_options(synchronize_session=False)
)

assert dr.schedule_tis((ti,), session=session) == 0

refreshed_ti = session.scalar(
select(TI).where(
TI.dag_id == ti.dag_id,
TI.task_id == ti.task_id,
TI.run_id == ti.run_id,
TI.map_index == ti.map_index,
)
)
assert refreshed_ti is not None
assert refreshed_ti.state == TaskInstanceState.QUEUED
assert refreshed_ti.try_number == 1


def test_schedule_tis_up_for_reschedule_does_not_increment_try_number(dag_maker, session):
with dag_maker(session=session) as dag:
BashOperator(task_id="task", bash_command="echo 1")

dr = dag_maker.create_dagrun(session=session)
ti = dr.get_task_instance("task", session=session)
assert ti is not None
ti.refresh_from_task(dag.get_task("task"))

ti.state = TaskInstanceState.UP_FOR_RESCHEDULE
ti.try_number = 3
session.commit()

assert dr.schedule_tis((ti,), session=session) == 1
session.commit()
session.expire_all()

refreshed_ti = session.scalar(
select(TI).where(
TI.dag_id == ti.dag_id,
TI.task_id == ti.task_id,
TI.run_id == ti.run_id,
TI.map_index == ti.map_index,
)
)
assert refreshed_ti is not None
assert refreshed_ti.state == TaskInstanceState.SCHEDULED
assert refreshed_ti.try_number == 3


def test_schedule_tis_is_noop_if_ti_transitions_to_nonschedulable_state_before_update(dag_maker, session):
from airflow.utils.session import create_session

with dag_maker(session=session) as dag:
BashOperator(task_id="task", bash_command="echo 1")

dr = dag_maker.create_dagrun(session=session)
ti = dr.get_task_instance("task", session=session)
assert ti is not None
ti.refresh_from_task(dag.get_task("task"))

ti.try_number = 1
session.commit()

with create_session() as other_session:
filter_for_tis = TI.filter_for_tis([ti])
assert filter_for_tis is not None
other_session.execute(
update(TI)
.where(filter_for_tis)
.values(state=TaskInstanceState.QUEUED, try_number=1)
.execution_options(synchronize_session=False)
)

assert dr.schedule_tis((ti,), session=session) == 0

refreshed_ti = session.scalar(
select(TI).where(
TI.dag_id == ti.dag_id,
TI.task_id == ti.task_id,
TI.run_id == ti.run_id,
TI.map_index == ti.map_index,
)
)
assert refreshed_ti is not None
assert refreshed_ti.state == TaskInstanceState.QUEUED
assert refreshed_ti.try_number == 1


def test_schedule_tis_empty_operator_is_noop_if_ti_already_running(dag_maker, session):
from airflow.utils.session import create_session

with dag_maker(session=session) as dag:
EmptyOperator(task_id="empty_task")

dr = dag_maker.create_dagrun(session=session)
ti = dr.get_task_instance("empty_task", session=session)
assert ti is not None
ti.refresh_from_task(dag.get_task("empty_task"))

ti.try_number = 3
session.commit()

with create_session() as other_session:
filter_for_tis = TI.filter_for_tis([ti])
assert filter_for_tis is not None
other_session.execute(
update(TI)
.where(filter_for_tis)
.values(state=TaskInstanceState.RUNNING, try_number=3)
.execution_options(synchronize_session=False)
)

assert dr.schedule_tis((ti,), session=session) == 0

refreshed_ti = session.scalar(
select(TI).where(
TI.dag_id == ti.dag_id,
TI.task_id == ti.task_id,
TI.run_id == ti.run_id,
TI.map_index == ti.map_index,
)
)
assert refreshed_ti is not None
assert refreshed_ti.state == TaskInstanceState.RUNNING
assert refreshed_ti.try_number == 3


def test_schedule_tis_only_one_scheduler_update_succeeds_when_competing(dag_maker, session):
from airflow.utils.session import create_session

with dag_maker(session=session) as dag:
BashOperator(task_id="task", bash_command="echo 1")

dr = dag_maker.create_dagrun(session=session)
ti = dr.get_task_instance("task", session=session)
assert ti is not None
ti.refresh_from_task(dag.get_task("task"))
assert ti.state is None

ti.try_number = 0
session.commit()

assert dr.schedule_tis((ti,), session=session) == 1
session.commit()

with create_session() as scheduler_b_session:
ti_b = scheduler_b_session.scalar(
select(TI).where(
TI.dag_id == ti.dag_id,
TI.task_id == ti.task_id,
TI.run_id == ti.run_id,
TI.map_index == ti.map_index,
)
)
assert ti_b is not None
assert dr.schedule_tis((ti_b,), session=scheduler_b_session) == 0

refreshed_ti = session.scalar(
select(TI).where(
TI.dag_id == ti.dag_id,
TI.task_id == ti.task_id,
TI.run_id == ti.run_id,
TI.map_index == ti.map_index,
)
)
assert refreshed_ti is not None
assert refreshed_ti.state == TaskInstanceState.SCHEDULED
assert refreshed_ti.try_number == 1


@pytest.mark.xfail(reason="We can't keep this behaviour with remote workers where scheduler can't reach xcom")
@pytest.mark.need_serialized_dag
def test_schedule_tis_start_trigger(dag_maker, session):
Expand Down
Loading