# Licensed to the Apache Software Foundation (ASF) under one
# or more contributor license agreements. See the NOTICE file
# distributed with this work for additional information
# regarding copyright ownership. The ASF licenses this file
# to you under the Apache License, Version 2.0 (the
# "License"); you may not use this file except in compliance
# with the License. You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing,
# software distributed under the License is distributed on an
# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
# KIND, either express or implied. See the License for the
# specific language governing permissions and limitations
# under the License.
from __future__ import annotations
import logging
from collections.abc import Sequence
from copy import deepcopy
from datetime import datetime, timedelta
from typing import TYPE_CHECKING, Any
from sqlalchemy import delete, select
from airflow.executors import workloads
from airflow.executors.base_executor import BaseExecutor
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, 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 (
CALLBACK_JOB_MAP_INDEX,
CALLBACK_JOB_TRY_NUMBER,
EXECUTE_CALLBACK_TAG,
build_callback_run_id,
is_callback_execute,
)
from airflow.providers.edge3.version_compat import AIRFLOW_V_3_4_PLUS
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
from airflow.utils.state import TaskInstanceState
if AIRFLOW_V_3_4_PLUS:
from airflow.executors.workloads.base import WorkloadType
if TYPE_CHECKING:
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
[docs]
CommandType = Sequence[str]
# Task tuple to send to be executed
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,
)
[docs]
class EdgeExecutor(BaseExecutor):
"""Implementation of the EdgeExecutor to distribute work to Edge Workers via HTTP."""
[docs]
supports_multi_team: bool = True
def __init__(self, *args, **kwargs):
super().__init__(*args, **kwargs)
[docs]
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.
# In such cases, fall back to the global configuration object.
# This allows the changes to be backwards compatible with older versions of Airflow.
# Can be removed when minimum supported provider version is equal to the version of core airflow
# which introduces multi-team configuration (3.2+).
if not hasattr(self, "conf") or not hasattr(self.conf, "getint"):
from airflow.configuration import conf as global_conf
self.conf = global_conf
# Also set team_name to None if it doesn't exist, since the Celery app creation expects it to be
# there (even if it's None)
if not hasattr(self, "team_name"):
self.team_name = None
@provide_session
[docs]
def start(self, *, session: Session = NEW_SESSION):
"""If EdgeExecutor provider is loaded first time, ensure table exists."""
check_db_manager_config()
edge_db_manager = EdgeDBManager(session)
if edge_db_manager.check_migration():
return
with create_global_lock(session=session, lock=DBLocks.MIGRATIONS):
edge_db_manager.initdb()
def _process_tasks(self, task_tuples: list[TaskTuple]) -> None:
"""
Temporary overwrite of _process_tasks function.
Idea is to not change the interface of the execute_async function in BaseExecutor as it will be changed in Airflow 3.
Edge worker needs task_instance in execute_async but BaseExecutor deletes this out of the self.queued_tasks.
Store queued_tasks in own var to be able to access this in execute_async function.
"""
self.edge_queued_tasks = deepcopy(self.queued_tasks)
super()._process_tasks(task_tuples) # type: ignore[misc]
[docs]
def queue_workload(
self,
workload: workloads.All,
session: Session,
) -> None:
"""Put new workload to queue. Airflow 3 entry point to execute a task."""
key: TaskInstanceKey | CallbackKey
if is_callback_execute(workload):
existing_job = session.scalars(
select(EdgeJobModel).where(
EdgeJobModel.dag_id == EXECUTE_CALLBACK_TAG,
EdgeJobModel.task_id == workload.callback.id,
EdgeJobModel.run_id == build_callback_run_id(workload.callback.id),
)
).first()
if existing_job:
existing_job.state = TaskInstanceState.QUEUED
existing_job.command = workload.model_dump_json()
else:
session.add(
EdgeJobModel(
dag_id=EXECUTE_CALLBACK_TAG,
task_id=str(workload.callback.id),
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
# Check if job already exists with same dag_id, task_id, run_id, map_index, try_number
existing_job = session.scalars(
select(EdgeJobModel).where(
EdgeJobModel.dag_id == key.dag_id,
EdgeJobModel.task_id == key.task_id,
EdgeJobModel.run_id == key.run_id,
EdgeJobModel.map_index == key.map_index,
EdgeJobModel.try_number == key.try_number,
)
).first()
if existing_job:
existing_job.state = TaskInstanceState.QUEUED
existing_job.queue = task_instance.queue
existing_job.concurrency_slots = task_instance.pool_slots
existing_job.command = workload.model_dump_json()
existing_job.team_name = self.team_name
else:
session.add(
EdgeJobModel(
dag_id=key.dag_id,
task_id=key.task_id,
run_id=key.run_id,
map_index=key.map_index,
try_number=key.try_number,
state=TaskInstanceState.QUEUED,
queue=task_instance.queue,
concurrency_slots=task_instance.pool_slots,
command=workload.model_dump_json(),
team_name=self.team_name,
)
)
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:
"""
No-op: EdgeExecutor does not use the BaseExecutor workload pipeline.
EdgeExecutor handles task queuing directly in queue_workload() by writing
to the EdgeJobModel database table, bypassing BaseExecutor's queued_tasks.
Therefore, trigger_tasks() never accumulates workloads to pass here.
"""
def _check_worker_liveness(self, session: Session) -> bool:
"""Reset worker state if heartbeat timed out."""
changed = False
heartbeat_interval: int = self.conf.getint("edge", "heartbeat_interval")
lifeless_workers = session.scalars(
select(EdgeWorkerModel)
.with_for_update(skip_locked=True)
.where(
EdgeWorkerModel.team_name == self.team_name,
EdgeWorkerModel.state.not_in(
[
EdgeWorkerState.UNKNOWN,
EdgeWorkerState.OFFLINE,
EdgeWorkerState.OFFLINE_MAINTENANCE,
]
),
EdgeWorkerModel.last_update < (timezone.utcnow() - timedelta(seconds=heartbeat_interval * 5)),
)
).all()
for worker in lifeless_workers:
changed = True
# If the worker dies in maintenance mode we want to remember it, so it can start in maintenance mode
worker.state = (
EdgeWorkerState.OFFLINE_MAINTENANCE
if worker.state
in (
EdgeWorkerState.MAINTENANCE_MODE,
EdgeWorkerState.MAINTENANCE_PENDING,
EdgeWorkerState.MAINTENANCE_REQUEST,
)
else EdgeWorkerState.UNKNOWN
)
# Reset presented status
sysinfo = dict(worker.sysinfo or {}) # copy needed to have alembic detect change in content
sysinfo["status"] = logging.NOTSET
sysinfo.pop("status_text", None) # Remove old status text if exists
worker.sysinfo = sysinfo
self.log.warning("Worker %s is lifeless. Setting state to %s", worker.worker_name, worker.state)
reset_metrics(worker.worker_name, team_name=worker.team_name)
return changed
def _update_orphaned_jobs(self, session: Session) -> bool:
"""Update status ob jobs when workers die and don't update anymore."""
heartbeat_interval: int = self.conf.getint("scheduler", "task_instance_heartbeat_timeout")
lifeless_jobs = session.scalars(
select(EdgeJobModel)
.with_for_update(skip_locked=True)
.where(
EdgeJobModel.team_name == self.team_name,
EdgeJobModel.state == TaskInstanceState.RUNNING,
EdgeJobModel.last_update < (timezone.utcnow() - timedelta(seconds=heartbeat_interval)),
)
).all()
for job in lifeless_jobs:
ti = TaskInstance.get_task_instance(
dag_id=job.dag_id,
run_id=job.run_id,
task_id=job.task_id,
map_index=job.map_index,
session=session,
)
job.state = ti.state if ti and ti.state else TaskInstanceState.REMOVED
if job.state != TaskInstanceState.RUNNING:
# Edge worker does not backport emitted Airflow metrics, so export some metrics
# Export metrics as failed as these jobs will be deleted in the future
tags = {
"dag_id": job.dag_id,
"task_id": job.task_id,
"queue": job.queue,
"state": str(TaskInstanceState.FAILED),
"team_name": job.team_name,
}
Stats.incr("edge_worker.ti.finish", tags=prune_dict(tags))
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
job_success_purge = self.conf.getint("edge", "job_success_purge")
job_fail_purge = self.conf.getint("edge", "job_fail_purge")
jobs = session.scalars(
select(EdgeJobModel)
.with_for_update(skip_locked=True)
.where(
EdgeJobModel.team_name == self.team_name,
EdgeJobModel.state.in_(_PURGE_HANDLED_STATES),
)
).all()
# 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:
if job.state == TaskInstanceState.RUNNING:
if (
job.key not in self.last_reported_state
or self.last_reported_state[job.key] != job.state
):
self.running_state(job.key)
self.last_reported_state[job.key] = job.state
elif job.state == TaskInstanceState.SUCCESS:
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.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
and job.last_update_t < (datetime.now() - timedelta(minutes=job_success_purge)).timestamp()
) or (
job.state
in (
TaskInstanceState.FAILED,
TaskInstanceState.REMOVED,
TaskInstanceState.RESTARTING,
TaskInstanceState.UP_FOR_RETRY,
)
and job.last_update_t < (datetime.now() - timedelta(minutes=job_fail_purge)).timestamp()
):
if job.key in self.last_reported_state:
del self.last_reported_state[job.key]
purged_marker = True
session.delete(job)
session.execute(
delete(EdgeLogsModel).where(
EdgeLogsModel.dag_id == job.dag_id,
EdgeLogsModel.run_id == job.run_id,
EdgeLogsModel.task_id == job.task_id,
EdgeLogsModel.map_index == job.map_index,
EdgeLogsModel.try_number == job.try_number,
)
)
return purged_marker
@provide_session
[docs]
def sync(self, *, session: Session = NEW_SESSION) -> None:
"""Sync will get called periodically by the heartbeat method."""
with Stats.timer("edge_executor.sync.duration", tags=prune_dict({"team_name": self.team_name})):
orphaned = self._update_orphaned_jobs(session)
purged = self._purge_jobs(session)
liveness = self._check_worker_liveness(session)
if purged or liveness or orphaned:
session.commit()
[docs]
def end(self) -> None:
"""End the executor."""
self.log.info("Shutting down EdgeExecutor")
[docs]
def terminate(self):
"""Terminate the executor is not doing anything."""
@provide_session
[docs]
def revoke_task(self, *, ti: TaskInstance, session: Session = NEW_SESSION):
"""
Revoke a task instance from the executor.
This method removes the task from the executor's internal state and deletes
the corresponding EdgeJobModel record to prevent edge workers from picking it up.
:param ti: Task instance to revoke
:param session: Database session
"""
# Remove from executor's internal state
self.running.discard(ti.key)
if AIRFLOW_V_3_4_PLUS:
self.executor_queues[WorkloadType.EXECUTE_TASK].pop(ti.key, None)
else:
self.queued_tasks.pop(ti.key, None)
if ti.key in self.last_reported_state:
del self.last_reported_state[ti.key]
# Delete the job from the database to prevent edge workers from picking it up
session.execute(
delete(EdgeJobModel).where(
EdgeJobModel.dag_id == ti.dag_id,
EdgeJobModel.task_id == ti.task_id,
EdgeJobModel.run_id == ti.run_id,
EdgeJobModel.map_index == ti.map_index,
EdgeJobModel.try_number == ti.try_number,
)
)
self.log.info("Revoked task instance %s from EdgeExecutor", ti.key)
@provide_session
[docs]
def try_adopt_task_instances(
self, tis: Sequence[TaskInstance], *, session: Session = NEW_SESSION
) -> Sequence[TaskInstance]:
"""
Adopt the task instances whose job is still in flight in the edge_job table.
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
"""
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
[docs]
def get_cli_commands() -> list[GroupCommand]:
from airflow.providers.edge3.cli.definition import get_edge_cli_commands
return get_edge_cli_commands()