Source code for airflow_nomad.airflow

from __future__ import annotations

import logging
import subprocess
from pathlib import Path
from typing import TYPE_CHECKING, Any, Literal

from airflow_pydantic import Pool, fail, skip
from nomad_pydantic import NomadClient

from airflow_nomad.config import NomadAirflowConfiguration

if TYPE_CHECKING:
    from airflow_ha import CheckResult, HighAvailabilityOperator
    from nomad_pydantic import CommandResult, JobStatus

    DAG = Any
    Operator = Any

__all__ = ("Nomad", "check_nomad_health")

Step = Literal["check-job", "cleanup-nomad", "configure-nomad", "force-kill", "register-job", "restart-job", "stop-job"]


def _python_operator() -> Any:
    from airflow_pydantic.airflow import PythonOperator

    return PythonOperator


def _remove_jobspec(path: Path | None, working_dir: Path | None) -> bool:
    if path is not None and path.exists():
        path.unlink()
    if working_dir is not None and working_dir.exists():
        try:
            working_dir.rmdir()
        except OSError:
            pass
    return True


def _forward_logs(
    cfg: NomadAirflowConfiguration,
    client: NomadClient,
    status: JobStatus,
    log: Any,
    offsets_by_stream: dict[str, dict[str, int]],
    *,
    final: bool = False,
) -> None:
    groups = {group.name: group for group in cfg.job.task_groups}
    if not status.allocations:
        return
    version = max(allocation.job_version for allocation in status.allocations)
    for allocation in status.allocations:
        if allocation.job_version != version:
            continue
        group = groups.get(allocation.task_group)
        if group is None:
            continue
        for workload in group.tasks:
            if workload.log_config is not None and workload.log_config.disabled:
                continue
            for stream in ("stdout", "stderr"):
                key = f"{allocation.id}/{workload.name}/{stream}"
                offsets = offsets_by_stream.setdefault(key, {})
                try:
                    for _ in range(16 if final else 1):
                        chunks = client.read_logs(allocation.id, workload.name, stream, offsets=offsets, limit=cfg.log_chunk_size)
                        for chunk in chunks:
                            if chunk.truncated:
                                log.warning("Nomad log truncated: allocation=%s task=%s stream=%s", allocation.id, workload.name, stream)
                            for line in chunk.data.decode("utf-8", errors="replace").splitlines():
                                log.info("[nomad %s/%s %s] %s", allocation.id, workload.name, stream, line)
                            offsets[chunk.file] = chunk.offset + len(chunk.data)
                        if sum(len(chunk.data) for chunk in chunks) < cfg.log_chunk_size:
                            break
                    else:
                        log.warning("Nomad final log drain reached byte limit: allocation=%s task=%s stream=%s", allocation.id, workload.name, stream)
                except (RuntimeError, OSError, ValueError, subprocess.SubprocessError) as error:
                    log.warning("Cannot read Nomad logs: allocation=%s task=%s stream=%s: %s", allocation.id, workload.name, stream, error)


def _log_failure(status: JobStatus, log: Any) -> None:
    for allocation in status.current_allocations:
        if allocation.client_status not in {"failed", "lost"}:
            continue
        log.error("Nomad allocation failed: id=%s group=%s status=%s", allocation.id, allocation.task_group, allocation.client_status)
        for name, state in getattr(allocation, "task_states", {}).items():
            for event in state.events:
                log.error(
                    "Nomad task event: allocation=%s task=%s state=%s type=%s exit_code=%s signal=%s message=%s",
                    allocation.id,
                    name,
                    state.state,
                    event.type,
                    event.exit_code,
                    event.signal,
                    event.display_message,
                )


[docs] class Nomad: """Airflow task group for a Nomad-managed job."""
[docs] def __init__(self, dag: DAG, cfg: NomadAirflowConfiguration | dict[str, Any], **kwargs: Any): if isinstance(cfg, dict): cfg = NomadAirflowConfiguration.model_validate(cfg) self._cfg = cfg self._pool = cfg.pool.pool if isinstance(cfg.pool, Pool) else cfg.pool self._nomad_client = kwargs.pop("nomad_client", NomadClient(cfg)) kwargs.pop("task_id", None) self._operator_kwargs = kwargs self._log_offsets: dict[str, dict[str, int]] = {} self._dag = dag self.setup_dag() self.initialize_tasks() self.configure_nomad >> self.register_job >> self.check_job self.check_job.retrigger_fail >> self.restart_job self.check_job.stop_pass >> self.stop_job >> self.cleanup_nomad self._force_kill = self.get_step_operator("force-kill") PythonOperator = _python_operator() ( PythonOperator( task_id=f"{self._dag.dag_id}-force-kill-dag", python_callable=skip, **self.get_base_operator_kwargs(), ) >> self._force_kill ) any_config_fail = PythonOperator( task_id=f"{self._dag.dag_id}-check-config-failed", python_callable=fail, **{**self.get_base_operator_kwargs(), "trigger_rule": "one_failed"}, ) self.configure_nomad >> any_config_fail self.register_job >> any_config_fail self.stop_job >> any_config_fail self.cleanup_nomad >> any_config_fail
def setup_dag(self) -> None: self._dag.catchup = False self._dag.concurrency = 1 self._dag.max_active_tasks = 1 self._dag.max_active_runs = 1 def initialize_tasks(self) -> None: PythonOperator = _python_operator() self._check_job = self.get_step_operator("check-job") self._configure_nomad = self.get_step_operator("configure-nomad") self._register_job = self.get_step_operator("register-job") if self._cfg.stop_on_exit: self._stop_job = self.get_step_operator("stop-job") if self._cfg.cleanup: self._cleanup_nomad = self.get_step_operator("cleanup-nomad") else: self._cleanup_nomad = PythonOperator( task_id=f"{self._dag.dag_id}-cleanup-nomad", python_callable=skip, **self.get_base_operator_kwargs(), ) else: self._stop_job = PythonOperator( task_id=f"{self._dag.dag_id}-stop-job", python_callable=skip, **self.get_base_operator_kwargs(), ) self._cleanup_nomad = PythonOperator( task_id=f"{self._dag.dag_id}-cleanup-nomad", python_callable=skip, **self.get_base_operator_kwargs(), ) self._restart_job = self.get_step_operator("restart-job") @property def configure_nomad(self) -> Operator: return self._configure_nomad @property def register_job(self) -> Operator: return self._register_job @property def check_job(self) -> HighAvailabilityOperator: return self._check_job @property def stop_job(self) -> Operator: return self._stop_job @property def restart_job(self) -> Operator: return self._restart_job @property def cleanup_nomad(self) -> Operator: return self._cleanup_nomad @property def nomad_client(self) -> NomadClient: return self._nomad_client def get_base_operator_kwargs(self) -> dict[str, Any]: return {**self._operator_kwargs, "dag": self._dag, "pool": self._pool} def _forward_logs(self, status: JobStatus, context: dict[str, Any], *, final: bool = False) -> None: if not self._cfg.forward_logs: return task = context.get("task", self.check_job) log = getattr(task, "log", logging.getLogger("airflow.task")) ti = context.get("ti") or context.get("task_instance") if ti is not None: saved = ti.xcom_pull(task_ids=task.task_id, key="nomad_log_offsets") if saved is None: saved = ti.xcom_pull(task_ids=self.check_job.task_id, key="nomad_log_offsets") self._log_offsets = saved or {} _forward_logs(self._cfg, self.nomad_client, status, log, self._log_offsets, final=final) if ti is not None: ti.xcom_push(key="nomad_log_offsets", value=self._log_offsets) def _command(self, result: CommandResult) -> bool: if result.returncode: from airflow.exceptions import AirflowException raise AirflowException(f"Nomad command failed (exit {result.returncode}): {result.stderr.strip() or result.stdout.strip()}") return True def _flush_logs(self, context: dict[str, Any]) -> None: if not self._cfg.forward_logs: return try: status = self.nomad_client.status() except (RuntimeError, OSError, ValueError, subprocess.SubprocessError) as error: log = getattr(context.get("task", self.check_job), "log", logging.getLogger("airflow.task")) log.warning("Cannot read Nomad status for final logs: %s", error) return self._forward_logs(status, context, final=True) def get_step_kwargs(self, step: Step) -> dict[str, Any]: cfg = self._cfg if step == "configure-nomad": return { "python_callable": lambda **kwargs: self.check_job.check_end_conditions(**kwargs) is None and str(cfg.write()), "do_xcom_push": True, } if step == "register-job": def _register_job(**kwargs: Any) -> bool: if self.check_job.check_end_conditions(**kwargs) is not None: return False registered = self._command(self.nomad_client.register()) restart = cfg.restart_on_retrigger or (cfg.restart_on_initial and self.check_job.is_initial_run(**kwargs)) if registered and restart: self._flush_logs(kwargs) return self._command(self.nomad_client.restart()) return registered return {"python_callable": _register_job, "do_xcom_push": True} if step == "stop-job": def _stop_job(**kwargs: Any) -> bool: self._flush_logs(kwargs) return self._command(self.nomad_client.stop(purge=cfg.purge_on_exit)) return { "python_callable": _stop_job, "do_xcom_push": True, } if step == "check-job": def _check_job(**kwargs: Any) -> CheckResult: from airflow_ha import Action, Result status = self.nomad_client.status() self._forward_logs(status, kwargs, final=status.complete or status.stopped or status.failed) if status.complete or status.stopped: return Result.PASS, Action.STOP if status.failed: _log_failure(status, getattr(kwargs.get("task", self.check_job), "log", logging.getLogger("airflow.task"))) return Result.FAIL, Action.RETRIGGER return Result.PASS, Action.CONTINUE return {"python_callable": _check_job, "do_xcom_push": True} if step == "restart-job": def _restart_job(**kwargs: Any) -> bool: self._flush_logs(kwargs) return self._command(self.nomad_client.restart()) return {"python_callable": _restart_job, "do_xcom_push": True} if step == "cleanup-nomad": return {"python_callable": lambda: _remove_jobspec(cfg.path, cfg.working_dir), "do_xcom_push": True} if step == "force-kill": def _force_kill(**kwargs: Any) -> bool: self._flush_logs(kwargs) return self._command(self.nomad_client.stop(purge=True)) return {"python_callable": _force_kill, "do_xcom_push": True} raise NotImplementedError(f"Unknown step: {step}") def get_step_operator(self, step: Step) -> Operator: from airflow_ha import HighAvailabilityOperator PythonOperator = _python_operator() if step == "check-job": return HighAvailabilityOperator( task_id=f"{self._dag.dag_id}-{step}", poke_interval=self._cfg.check_interval.total_seconds(), timeout=self._cfg.check_timeout.total_seconds(), mode="poke", runtime=self._cfg.runtime, endtime=self._cfg.endtime, maxretrigger=self._cfg.maxretrigger, reference_date=self._cfg.reference_date, **self.get_base_operator_kwargs(), **self.get_step_kwargs(step), ) return PythonOperator( task_id=f"{self._dag.dag_id}-{step}", **self.get_base_operator_kwargs(), **self.get_step_kwargs(step), ) def __lshift__(self, other: Operator) -> Operator: self.configure_nomad << other return self.cleanup_nomad def __rshift__(self, other: Operator) -> Operator: self.cleanup_nomad >> other return other def set_upstream(self, other: Operator) -> None: self.configure_nomad.set_upstream(other) def set_downstream(self, other: Operator) -> None: self.cleanup_nomad.set_downstream(other) def update_relative(self, other: Operator, upstream: bool = True, edge_modifier: Any = None) -> Nomad: if upstream: self.configure_nomad.update_relative(other, upstream=True, edge_modifier=edge_modifier) else: self.cleanup_nomad.update_relative(other, upstream=False, edge_modifier=edge_modifier) return self @property def leaves(self) -> Any: return self.cleanup_nomad.leaves @property def roots(self) -> Any: return self.configure_nomad.roots
[docs] def check_nomad_health( cfg: NomadAirflowConfiguration | dict[str, Any], *, nomad_client: NomadClient | None = None, require_running: bool = True, **context: Any ) -> dict[str, Any]: """Check an existing job once, requiring a running service by default. Use this callable in a scheduled PythonOperator to monitor persistent jobs. Airflow task XComs retain byte cursors across scheduled health checks. """ from airflow.exceptions import AirflowException if isinstance(cfg, dict): cfg = NomadAirflowConfiguration.model_validate(cfg) client = nomad_client or NomadClient(cfg) status = client.status() log = getattr(context.get("task"), "log", logging.getLogger("airflow.task")) if cfg.forward_logs: ti = context.get("ti") or context.get("task_instance") offsets = {} if ti is not None: task_id = getattr(context.get("task"), "task_id", None) or ti.task_id offsets = ti.xcom_pull(task_ids=task_id, key="nomad_log_offsets", include_prior_dates=True) or {} _forward_logs(cfg, client, status, log, offsets, final=True) if ti is not None: ti.xcom_push(key="nomad_log_offsets", value=offsets) _log_failure(status, log) failed_allocations = any( allocation.client_status in {"failed", "lost"} or any(state.failed for state in allocation.task_states.values()) for allocation in status.current_allocations ) if ( status.failed or failed_allocations or status.stopped or not status.allocations or (require_running and not any(allocation.client_status == "running" for allocation in status.current_allocations)) ): details = [ f"{allocation.id}/{name}: {state.state}; " + "; ".join(f"{event.type} exit={event.exit_code} signal={event.signal} {event.display_message}" for event in state.events) for allocation in status.current_allocations for name, state in allocation.task_states.items() ] raise AirflowException( f"Nomad job {status.id} is unhealthy: " + ( "; ".join(details) or f"running={status.running} complete={status.complete} failed={status.failed} stopped={status.stopped} allocations={len(status.allocations)}" ) ) return { "job_id": status.id, "namespace": status.namespace, "running": status.running, "complete": status.complete, "allocations": [allocation.id for allocation in status.current_allocations], }