diff --git a/providers/edge3/docs/changelog.rst b/providers/edge3/docs/changelog.rst index 5e513a51d58f6..990b2886eaf8d 100644 --- a/providers/edge3/docs/changelog.rst +++ b/providers/edge3/docs/changelog.rst @@ -27,6 +27,12 @@ Changelog --------- +.. warning:: + ``EdgeExecutor`` now counts the tasks and callbacks it has queued against ``[core] parallelism``, as the + other executors do. Until now that limit had no effect on Edge. If a scheduler keeps more than + ``parallelism`` (default 32) workloads in flight on Edge, raise ``[core] parallelism``. Otherwise the + scheduler leaves the rest in ``scheduled`` state until slots free up. + 4.3.2 ..... diff --git a/providers/edge3/src/airflow/providers/edge3/executors/edge_executor.py b/providers/edge3/src/airflow/providers/edge3/executors/edge_executor.py index e6f7af33eb06c..0b5b9c9e4dd50 100644 --- a/providers/edge3/src/airflow/providers/edge3/executors/edge_executor.py +++ b/providers/edge3/src/airflow/providers/edge3/executors/edge_executor.py @@ -30,10 +30,16 @@ from airflow.models.taskinstance import TaskInstance from airflow.providers.common.compat.sdk import Stats, timezone from airflow.providers.edge3.models.db import EdgeDBManager, check_db_manager_config -from airflow.providers.edge3.models.edge_job import EdgeJobModel +from airflow.providers.edge3.models.edge_job import EdgeJobModel, build_job_key from airflow.providers.edge3.models.edge_logs import EdgeLogsModel from airflow.providers.edge3.models.edge_worker import EdgeWorkerModel, EdgeWorkerState, reset_metrics -from airflow.providers.edge3.models.types import is_callback_execute +from airflow.providers.edge3.models.types import ( + CALLBACK_JOB_MAP_INDEX, + CALLBACK_JOB_TRY_NUMBER, + EXECUTE_CALLBACK_TAG, + build_callback_run_id, + is_callback_execute, +) from airflow.utils.db import DBLocks, create_global_lock from airflow.utils.helpers import prune_dict from airflow.utils.session import NEW_SESSION, provide_session @@ -43,6 +49,7 @@ from sqlalchemy.orm import Session from airflow.cli.cli_config import GroupCommand + from airflow.models.callback import CallbackKey from airflow.models.taskinstancekey import TaskInstanceKey # TODO: Airflow 2 type hints; remove when Airflow 2 support is removed @@ -51,6 +58,17 @@ TaskTuple = tuple[TaskInstanceKey, CommandType, str | None, Any | None] +# _purge_jobs() reports on or deletes a job only while it is in one of these states. +_PURGE_HANDLED_STATES = ( + TaskInstanceState.RUNNING, + TaskInstanceState.SUCCESS, + TaskInstanceState.FAILED, + TaskInstanceState.REMOVED, + TaskInstanceState.RESTARTING, + TaskInstanceState.UP_FOR_RETRY, +) + + class EdgeExecutor(BaseExecutor): """Implementation of the EdgeExecutor to distribute work to Edge Workers via HTTP.""" @@ -58,7 +76,7 @@ class EdgeExecutor(BaseExecutor): def __init__(self, *args, **kwargs): super().__init__(*args, **kwargs) - self.last_reported_state: dict[TaskInstanceKey, TaskInstanceState | str] = {} + self.last_reported_state: dict[TaskInstanceKey | CallbackKey, TaskInstanceState | str] = {} # Check if self has the ExecutorConf set on the self.conf attribute with all required methods. # In Airflow 2.x, ExecutorConf exists but lacks methods like getint, getboolean, getsection, etc. @@ -103,14 +121,13 @@ def queue_workload( session: Session, ) -> None: """Put new workload to queue. Airflow 3 entry point to execute a task.""" + key: TaskInstanceKey | CallbackKey if is_callback_execute(workload): - from airflow.providers.edge3.models.types import EXECUTE_CALLBACK_TAG - existing_job = session.scalars( select(EdgeJobModel).where( EdgeJobModel.dag_id == EXECUTE_CALLBACK_TAG, EdgeJobModel.task_id == workload.callback.id, - EdgeJobModel.run_id == f"{EXECUTE_CALLBACK_TAG}-{workload.callback.id}", + EdgeJobModel.run_id == build_callback_run_id(workload.callback.id), ) ).first() @@ -122,9 +139,9 @@ def queue_workload( EdgeJobModel( dag_id=EXECUTE_CALLBACK_TAG, task_id=str(workload.callback.id), - run_id=f"{EXECUTE_CALLBACK_TAG}-{workload.callback.id}", - map_index=-1, - try_number=0, + run_id=build_callback_run_id(workload.callback.id), + map_index=CALLBACK_JOB_MAP_INDEX, + try_number=CALLBACK_JOB_TRY_NUMBER, queue=self.conf.get_mandatory_value("operators", "default_queue"), concurrency_slots=1, state=TaskInstanceState.QUEUED, @@ -132,6 +149,7 @@ def queue_workload( team_name=self.team_name, ) ) + key = workload.key elif isinstance(workload, workloads.ExecuteTask): task_instance = workload.ti key = task_instance.key @@ -170,6 +188,8 @@ def queue_workload( ) else: raise TypeError(f"Don't know how to queue workload of type {type(workload).__name__}") + # Added before the caller commits. On rollback, the reconciliation in _purge_jobs() drops the key. + self.running.add(key) def _process_workloads(self, workloads: Sequence[workloads.All]) -> None: """ @@ -260,6 +280,24 @@ def _update_orphaned_jobs(self, session: Session) -> bool: return bool(lifeless_jobs) + def _get_tracked_job_keys( + self, session: Session, states: Sequence[TaskInstanceState] + ) -> set[TaskInstanceKey | CallbackKey]: + """ + Read the keys of this team's jobs that are in one of ``states``. + + Rows are read without locking on purpose: an edge worker fetches its next job with + ``FOR UPDATE SKIP LOCKED``, so locking the queued rows here would make it come back empty. + """ + query = select( + EdgeJobModel.dag_id, + EdgeJobModel.task_id, + EdgeJobModel.run_id, + EdgeJobModel.try_number, + EdgeJobModel.map_index, + ).where(EdgeJobModel.team_name == self.team_name, EdgeJobModel.state.in_(states)) + return {build_job_key(*row) for row in session.execute(query)} + def _purge_jobs(self, session: Session) -> bool: """Clean finished jobs.""" purged_marker = False @@ -270,22 +308,16 @@ def _purge_jobs(self, session: Session) -> bool: .with_for_update(skip_locked=True) .where( EdgeJobModel.team_name == self.team_name, - EdgeJobModel.state.in_( - [ - TaskInstanceState.RUNNING, - TaskInstanceState.SUCCESS, - TaskInstanceState.FAILED, - TaskInstanceState.REMOVED, - TaskInstanceState.RESTARTING, - TaskInstanceState.UP_FOR_RETRY, - ] - ), + EdgeJobModel.state.in_(_PURGE_HANDLED_STATES), ) ).all() - # Sync DB with executor otherwise runs out of sync in multi scheduler deployment - already_removed = self.running - set(job.key for job in jobs) - self.running = self.running - already_removed + # Sync DB with executor otherwise runs out of sync in multi scheduler deployment. Only a queued job + # or one handled below keeps its slot. _update_orphaned_jobs() can leave a job in any task instance + # state, and a row this method never reads again would hold its slot until the scheduler restarts. + self.running &= self._get_tracked_job_keys( + session, states=(TaskInstanceState.QUEUED, *_PURGE_HANDLED_STATES) + ) for job in jobs: if job.key in self.running: @@ -300,15 +332,13 @@ def _purge_jobs(self, session: Session) -> bool: if job.key in self.last_reported_state: del self.last_reported_state[job.key] self.success(job.key) - elif job.state in [ - TaskInstanceState.FAILED, - TaskInstanceState.RESTARTING, - TaskInstanceState.UP_FOR_RETRY, - ]: + elif job.state in [TaskInstanceState.FAILED, TaskInstanceState.UP_FOR_RETRY]: if job.key in self.last_reported_state: del self.last_reported_state[job.key] self.fail(job.key) else: + # RESTARTING is not a failure here: the fetch endpoint parks a claimed job in that + # state until the worker reports RUNNING. self.last_reported_state[job.key] = TaskInstanceState(job.state) if ( job.state == TaskInstanceState.SUCCESS @@ -385,17 +415,25 @@ def revoke_task(self, *, ti: TaskInstance, session: Session = NEW_SESSION): ) self.log.info("Revoked task instance %s from EdgeExecutor", ti.key) - def try_adopt_task_instances(self, tis: Sequence[TaskInstance]) -> Sequence[TaskInstance]: + @provide_session + def try_adopt_task_instances( + self, tis: Sequence[TaskInstance], *, session: Session = NEW_SESSION + ) -> Sequence[TaskInstance]: """ - Try to adopt running task instances that have been abandoned by a SchedulerJob dying. + Adopt the task instances whose job is still in flight in the edge_job table. - Anything that is not adopted will be cleared by the scheduler (and then become eligible for - re-scheduling) + The ``running`` set is empty after a scheduler restart, so the adopted keys go back into it + to keep slot accounting accurate. Task instances whose job is finished or missing are + returned so the scheduler clears and re-schedules them. :return: any TaskInstances that were unable to be adopted """ - # We handle all running tasks from the DB in sync, no adoption logic needed. - return [] + tracked_keys = self._get_tracked_job_keys( + session, + states=(TaskInstanceState.QUEUED, TaskInstanceState.RESTARTING, TaskInstanceState.RUNNING), + ) + self.running.update(ti.key for ti in tis if ti.key in tracked_keys) + return [ti for ti in tis if ti.key not in tracked_keys] @staticmethod def get_cli_commands() -> list[GroupCommand]: diff --git a/providers/edge3/src/airflow/providers/edge3/models/edge_job.py b/providers/edge3/src/airflow/providers/edge3/models/edge_job.py index 79576031112ea..0309a57d6d434 100644 --- a/providers/edge3/src/airflow/providers/edge3/models/edge_job.py +++ b/providers/edge3/src/airflow/providers/edge3/models/edge_job.py @@ -17,6 +17,7 @@ from __future__ import annotations from datetime import datetime +from typing import TYPE_CHECKING from sqlalchemy import ( Index, @@ -27,12 +28,35 @@ from sqlalchemy.orm import Mapped from airflow.models.base import StringID -from airflow.providers.common.compat.sdk import TaskInstanceKey, timezone +from airflow.models.taskinstancekey import TaskInstanceKey +from airflow.providers.common.compat.sdk import timezone from airflow.providers.common.compat.sqlalchemy.orm import mapped_column from airflow.providers.edge3.models.edge_base import Base +from airflow.providers.edge3.models.types import is_callback_job +from airflow.providers.edge3.version_compat import AIRFLOW_V_3_3_PLUS from airflow.utils.log.logging_mixin import LoggingMixin from airflow.utils.sqlalchemy import UtcDateTime +if TYPE_CHECKING: + from airflow.models.callback import CallbackKey + + +def build_job_key( + dag_id: str, task_id: str, run_id: str, try_number: int, map_index: int +) -> TaskInstanceKey | CallbackKey: + """ + Build the key the executor layer uses for a job row. + + A row is a callback only if it has the full identity ``queue_workload()`` writes for callbacks, since + ``ExecuteCallback`` is a valid Dag id. A task row maps to the ``airflow.models`` ``TaskInstanceKey``, + not the ``airflow.sdk`` one, because ``BaseExecutor`` dispatches on it with ``isinstance``. + """ + if AIRFLOW_V_3_3_PLUS and is_callback_job(dag_id, task_id, run_id, try_number, map_index): + from airflow.models.callback import CallbackKey + + return CallbackKey(id=task_id) + return TaskInstanceKey(dag_id, task_id, run_id, try_number, map_index) + class EdgeJobModel(Base, LoggingMixin): """ @@ -92,8 +116,9 @@ def __init__( __table_args__ = (Index("rj_order", state, queued_dttm, queue),) @property - def key(self): - return TaskInstanceKey(self.dag_id, self.task_id, self.run_id, self.try_number, self.map_index) + def key(self) -> TaskInstanceKey | CallbackKey: + """Key of the job as the executor layer knows it.""" + return build_job_key(self.dag_id, self.task_id, self.run_id, self.try_number, self.map_index) @property def last_update_t(self) -> float: diff --git a/providers/edge3/src/airflow/providers/edge3/models/types.py b/providers/edge3/src/airflow/providers/edge3/models/types.py index 19cea39539dc9..dce93f52f827b 100644 --- a/providers/edge3/src/airflow/providers/edge3/models/types.py +++ b/providers/edge3/src/airflow/providers/edge3/models/types.py @@ -44,3 +44,22 @@ def is_callback_execute(workload: workloads.All) -> TypeGuard[ExecuteCallback]: # This is the key used to identify execute_callback jobs. # Changing this value may break compatibility with existing data in the edge_job table. EXECUTE_CALLBACK_TAG = "ExecuteCallback" + +# The rest of the identity queue_workload() writes for a callback row. "ExecuteCallback" is a valid +# Dag id, so a row is a callback only when all four fields match. +CALLBACK_JOB_TRY_NUMBER = 0 +CALLBACK_JOB_MAP_INDEX = -1 + + +def build_callback_run_id(callback_id: str) -> str: + return f"{EXECUTE_CALLBACK_TAG}-{callback_id}" + + +def is_callback_job(dag_id: str, task_id: str, run_id: str, try_number: int, map_index: int) -> bool: + """Return whether a job row matches the identity ``queue_workload()`` writes for a callback.""" + return ( + dag_id == EXECUTE_CALLBACK_TAG + and run_id == build_callback_run_id(task_id) + and try_number == CALLBACK_JOB_TRY_NUMBER + and map_index == CALLBACK_JOB_MAP_INDEX + ) diff --git a/providers/edge3/tests/unit/edge3/executors/test_edge_executor.py b/providers/edge3/tests/unit/edge3/executors/test_edge_executor.py index 2f135957c4743..fd08a503e554d 100644 --- a/providers/edge3/tests/unit/edge3/executors/test_edge_executor.py +++ b/providers/edge3/tests/unit/edge3/executors/test_edge_executor.py @@ -29,6 +29,7 @@ from sqlalchemy import delete, select from airflow.executors.workloads import BundleInfo, ExecuteTask +from airflow.models.taskinstance import TaskInstance from airflow.providers.common.compat.sdk import Stats, TaskInstanceKey, conf, timezone from airflow.providers.edge3.executors.edge_executor import EdgeExecutor from airflow.providers.edge3.models.edge_job import EdgeJobModel @@ -43,6 +44,7 @@ if AIRFLOW_V_3_3_PLUS: from airflow.executors.workloads import CallbackFetchMethod, ExecuteCallback, TaskInstanceDTO from airflow.executors.workloads.callback import CallbackDTO + from airflow.utils.state import CallbackState pytestmark = pytest.mark.db_test @@ -625,12 +627,12 @@ def setup(self): session.execute(delete(EdgeJobModel)) session.commit() - def _make_execute_task(self) -> ExecuteTask: + def _make_execute_task(self, task_id: str = "test_task", dag_id: str = "test_dag") -> ExecuteTask: ti = TaskInstanceDTO( id=uuid4(), dag_version_id=uuid4(), - task_id="test_task", - dag_id="test_dag", + task_id=task_id, + dag_id=dag_id, run_id="test_run", try_number=1, map_index=-1, @@ -646,6 +648,20 @@ def _make_execute_task(self) -> ExecuteTask: log_path="test.log", ) + def _make_execute_callback(self) -> ExecuteCallback: + callback = CallbackDTO( + id=str(uuid4()), + fetch_method=CallbackFetchMethod.IMPORT_PATH, + data={"path": "builtins.dict", "kwargs": {"a": 1, "b": 2, "c": 3}}, + ) + return ExecuteCallback( + callback=callback, + dag_rel_path=Path("test.py"), + bundle_info=BundleInfo(name="test_bundle", version="1.0"), + token="test_token", + log_path="test.log", + ) + def test_queue_workload_execute_task(self): executor = EdgeExecutor() workload = self._make_execute_task() @@ -662,6 +678,165 @@ def test_queue_workload_execute_task(self): assert job.state == TaskInstanceState.QUEUED assert '"type":"ExecuteTask"' in job.command or '"type": "ExecuteTask"' in job.command + @pytest.mark.parametrize("make_workload", ["_make_execute_task", "_make_execute_callback"]) + def test_queue_workload_occupies_an_executor_slot(self, make_workload): + executor = EdgeExecutor() + workload = getattr(self, make_workload)() + + with create_session() as session: + executor.queue_workload(workload, session=session) + session.commit() + + assert workload.key in executor.running + assert executor.slots_available == executor.parallelism - 1 + + # The slot stays taken while the job waits in the queue for a worker to pick it up. + executor.sync() + + assert workload.key in executor.running + assert executor.slots_available == executor.parallelism - 1 + + @pytest.mark.parametrize("make_workload", ["_make_execute_task", "_make_execute_callback"]) + @pytest.mark.parametrize( + ("job_state", "reported_state"), + [ + (TaskInstanceState.RUNNING, "running"), + (TaskInstanceState.SUCCESS, "success"), + (TaskInstanceState.FAILED, "failed"), + (TaskInstanceState.UP_FOR_RETRY, "failed"), + ], + ) + def test_sync_reports_state_of_queued_workload(self, make_workload, job_state, reported_state): + executor = EdgeExecutor() + workload = getattr(self, make_workload)() + + with create_session() as session: + executor.queue_workload(workload, session=session) + session.commit() + executor.sync() + + with create_session() as session: + session.scalar(select(EdgeJobModel)).state = job_state + session.commit() + + executor.sync() + + reported_states = TaskInstanceState if isinstance(workload, ExecuteTask) else CallbackState + assert executor.get_event_buffer() == {workload.key: (reported_states(reported_state), None)} + + def test_sync_keeps_slot_while_worker_claims_job(self): + executor = EdgeExecutor() + workload = self._make_execute_task() + + with create_session() as session: + executor.queue_workload(workload, session=session) + session.commit() + + # The fetch endpoint parks a claimed job in RESTARTING until the worker reports RUNNING. + with create_session() as session: + session.scalar(select(EdgeJobModel)).state = TaskInstanceState.RESTARTING + session.commit() + executor.sync() + + assert executor.get_event_buffer() == {} + assert workload.ti.key in executor.running + assert executor.slots_available == executor.parallelism - 1 + + with create_session() as session: + session.scalar(select(EdgeJobModel)).state = TaskInstanceState.RUNNING + session.commit() + executor.sync() + + assert executor.get_event_buffer() == {workload.ti.key: (TaskInstanceState.RUNNING, None)} + + def test_sync_reports_job_that_finishes_after_being_marked_removed(self): + executor = EdgeExecutor() + workload = self._make_execute_callback() + + with create_session() as session: + executor.queue_workload(workload, session=session) + session.commit() + + # When a callback job runs past the heartbeat timeout, _update_orphaned_jobs() marks it REMOVED. + for job_state in (TaskInstanceState.REMOVED, TaskInstanceState.SUCCESS): + with create_session() as session: + session.scalar(select(EdgeJobModel)).state = job_state + session.commit() + executor.sync() + + assert executor.get_event_buffer() == {workload.key: (CallbackState.SUCCESS, None)} + + @pytest.mark.parametrize( + "unhandled_state", + [TaskInstanceState.SCHEDULED, TaskInstanceState.DEFERRED, TaskInstanceState.UP_FOR_RESCHEDULE], + ) + def test_sync_frees_slot_of_job_in_state_purge_never_handles(self, unhandled_state): + executor = EdgeExecutor() + workload = self._make_execute_task() + + with create_session() as session: + executor.queue_workload(workload, session=session) + session.commit() + + # _update_orphaned_jobs() copies the task instance state into the job, whatever that state is. + with create_session() as session: + session.scalar(select(EdgeJobModel)).state = unhandled_state + session.commit() + executor.sync() + + assert workload.key not in executor.running + assert executor.slots_available == executor.parallelism + + def test_task_in_dag_named_after_callback_tag_keeps_its_task_key(self): + executor = EdgeExecutor() + workload = self._make_execute_task(dag_id=EXECUTE_CALLBACK_TAG) + + with create_session() as session: + executor.queue_workload(workload, session=session) + session.commit() + executor.sync() + + assert workload.key in executor.running + + with create_session() as session: + session.scalar(select(EdgeJobModel)).state = TaskInstanceState.RUNNING + session.commit() + executor.sync() + + assert executor.get_event_buffer() == {workload.key: (TaskInstanceState.RUNNING, None)} + + @pytest.mark.parametrize( + "finished_state", [TaskInstanceState.SUCCESS, TaskInstanceState.FAILED, TaskInstanceState.REMOVED] + ) + def test_try_adopt_task_instances_restores_slots_from_edge_job(self, finished_state): + executor = EdgeExecutor() + queued = self._make_execute_task() + finished = self._make_execute_task(task_id="finished") + with create_session() as session: + executor.queue_workload(queued, session=session) + executor.queue_workload(finished, session=session) + session.commit() + with create_session() as session: + finished_job = session.scalar(select(EdgeJobModel).where(EdgeJobModel.task_id == "finished")) + finished_job.state = finished_state + session.commit() + + restarted_executor = EdgeExecutor() + queued_ti = mock.Mock(spec=TaskInstance, key=queued.key) + finished_ti = mock.Mock(spec=TaskInstance, key=finished.key) + orphaned_ti = mock.Mock( + spec=TaskInstance, + key=TaskInstanceKey( + dag_id="test_dag", task_id="orphan", run_id="test_run", try_number=1, map_index=-1 + ), + ) + + not_adopted = restarted_executor.try_adopt_task_instances([queued_ti, finished_ti, orphaned_ti]) + + assert not_adopted == [finished_ti, orphaned_ti] + assert restarted_executor.running == {queued.key} + assert restarted_executor.slots_available == restarted_executor.parallelism - 1 + def test_queue_workload_execute_task_existing_job(self): executor = EdgeExecutor() workload = self._make_execute_task() @@ -678,22 +853,7 @@ def test_queue_workload_execute_task_existing_job(self): def test_queue_workload_execute_callback(self): executor = EdgeExecutor() - id = str(uuid4()) - callback_data = CallbackDTO( - id=id, - fetch_method=CallbackFetchMethod.IMPORT_PATH, - data={ - "path": "builtins.dict", - "kwargs": {"a": 1, "b": 2, "c": 3}, - }, - ) - workload = ExecuteCallback( - callback=callback_data, - dag_rel_path=Path("test.py"), - bundle_info=BundleInfo(name="test_bundle", version="1.0"), - token="test_token", - log_path="test.log", - ) + workload = self._make_execute_callback() with create_session() as session: executor.queue_workload(workload, session=session) @@ -702,28 +862,14 @@ def test_queue_workload_execute_callback(self): job = session.scalar(select(EdgeJobModel)) assert job is not None assert job.dag_id == EXECUTE_CALLBACK_TAG - assert job.task_id == id - assert job.run_id == f"{EXECUTE_CALLBACK_TAG}-{id}" + assert job.task_id == workload.callback.id + assert job.run_id == f"{EXECUTE_CALLBACK_TAG}-{workload.callback.id}" assert job.state == TaskInstanceState.QUEUED assert '"type":"ExecuteCallback"' in job.command or '"type": "ExecuteCallback"' in job.command def test_queue_workload_execute_callback_existing_job(self): executor = EdgeExecutor() - callback_data = CallbackDTO( - id=str(uuid4()), - fetch_method=CallbackFetchMethod.IMPORT_PATH, - data={ - "path": "builtins.dict", - "kwargs": {"a": 1, "b": 2, "c": 3}, - }, - ) - workload = ExecuteCallback( - callback=callback_data, - dag_rel_path=Path("test.py"), - bundle_info=BundleInfo(name="test_bundle", version="1.0"), - token="test_token", - log_path="test.log", - ) + workload = self._make_execute_callback() with create_session() as session: executor.queue_workload(workload, session=session) diff --git a/providers/edge3/tests/unit/edge3/models/test_edge_job.py b/providers/edge3/tests/unit/edge3/models/test_edge_job.py index e0f75c25964c5..dc4b04135b98d 100644 --- a/providers/edge3/tests/unit/edge3/models/test_edge_job.py +++ b/providers/edge3/tests/unit/edge3/models/test_edge_job.py @@ -24,9 +24,12 @@ from sqlalchemy import delete, select from airflow.providers.common.compat.sdk import TaskInstanceKey -from airflow.providers.edge3.models.edge_job import EdgeJobModel +from airflow.providers.edge3.models.edge_job import EdgeJobModel, build_job_key +from airflow.providers.edge3.models.types import EXECUTE_CALLBACK_TAG from airflow.utils.state import TaskInstanceState +from tests_common.test_utils.version_compat import AIRFLOW_V_3_3_PLUS + if TYPE_CHECKING: from sqlalchemy.orm import Session @@ -63,6 +66,30 @@ def test_key_builds_task_instance_key(): assert job.key == TaskInstanceKey("test_dag", "test_task", "test_run", 2, 3) +@pytest.mark.skipif(not AIRFLOW_V_3_3_PLUS, reason="Callback workloads need Airflow 3.3+") +def test_build_job_key_maps_callback_row_to_callback_key(): + from airflow.models.callback import CallbackKey + + key = build_job_key(EXECUTE_CALLBACK_TAG, "abc", f"{EXECUTE_CALLBACK_TAG}-abc", 0, -1) + + assert key == CallbackKey(id="abc") + + +@pytest.mark.parametrize( + ("dag_id", "run_id", "try_number", "map_index"), + [ + pytest.param("test_dag", f"{EXECUTE_CALLBACK_TAG}-abc", 0, -1, id="other_dag"), + pytest.param(EXECUTE_CALLBACK_TAG, "manual__2026-01-01T00:00:00+00:00", 0, -1, id="task_run_id"), + pytest.param(EXECUTE_CALLBACK_TAG, f"{EXECUTE_CALLBACK_TAG}-abc", 1, -1, id="try_number"), + pytest.param(EXECUTE_CALLBACK_TAG, f"{EXECUTE_CALLBACK_TAG}-abc", 0, 2, id="map_index"), + ], +) +def test_build_job_key_keeps_task_key_unless_full_callback_identity(dag_id, run_id, try_number, map_index): + key = build_job_key(dag_id, "abc", run_id, try_number, map_index) + + assert key == TaskInstanceKey(dag_id, "abc", run_id, try_number, map_index) + + @time_machine.travel(datetime(2026, 1, 1, 12, 0, 0, tzinfo=dt_timezone.utc), tick=False) def test_queued_dttm_defaults_to_now(): job = _make_job()