Skip to content
Merged
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
6 changes: 6 additions & 0 deletions providers/edge3/docs/changelog.rst
Original file line number Diff line number Diff line change
Expand Up @@ -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
.....

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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
Expand All @@ -51,14 +58,25 @@
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."""

supports_multi_team: bool = True

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.
Expand Down Expand Up @@ -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()

Expand All @@ -122,16 +139,17 @@ 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,
command=workload.model_dump_json(),
team_name=self.team_name,
)
)
key = workload.key
elif isinstance(workload, workloads.ExecuteTask):
task_instance = workload.ti
key = task_instance.key
Expand Down Expand Up @@ -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:
"""
Expand Down Expand Up @@ -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
Expand All @@ -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:
Expand All @@ -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
Expand Down Expand Up @@ -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]:
Expand Down
31 changes: 28 additions & 3 deletions providers/edge3/src/airflow/providers/edge3/models/edge_job.py
Original file line number Diff line number Diff line change
Expand Up @@ -17,6 +17,7 @@
from __future__ import annotations

from datetime import datetime
from typing import TYPE_CHECKING

from sqlalchemy import (
Index,
Expand All @@ -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):
"""
Expand Down Expand Up @@ -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:
Expand Down
19 changes: 19 additions & 0 deletions providers/edge3/src/airflow/providers/edge3/models/types.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
)
Loading