Source code for airflow.providers.edge3.worker_api.routes.jobs

# 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 typing import TYPE_CHECKING, Annotated
from uuid import UUID

from fastapi import Body, Depends, HTTPException, status
from sqlalchemy import select

from airflow.api_fastapi.common.db.common import SessionDep  # noqa: TC001
from airflow.api_fastapi.common.router import AirflowRouter
from airflow.api_fastapi.core_api.openapi.exceptions import create_openapi_http_exception_doc
from airflow.executors.workloads import ExecuteTask
from airflow.providers.common.compat.sdk import Stats, timezone
from airflow.providers.edge3.models.edge_job import EdgeJobModel
from airflow.providers.edge3.models.edge_worker import EdgeWorkerModel
from airflow.providers.edge3.version_compat import AIRFLOW_V_3_3_PLUS
from airflow.providers.edge3.worker_api.auth import jwt_token_authorization_rest
from airflow.providers.edge3.worker_api.datamodels import (
    EdgeJobFetched,
    WorkerApiDocs,
    WorkerQueuesBody,
)
from airflow.utils.helpers import prune_dict
from airflow.utils.state import TaskInstanceState

if TYPE_CHECKING:
    from airflow.providers.edge3.models.types import ExecuteTypeBody

[docs] log = logging.getLogger(__name__)
[docs] jobs_router = AirflowRouter(tags=["Jobs"], prefix="/jobs")
[docs] def parse_command(command: str, dag_id: str, run_id: str) -> ExecuteTypeBody: if AIRFLOW_V_3_3_PLUS: from airflow.executors.workloads import ExecuteCallback from airflow.providers.edge3.models.types import EXECUTE_CALLBACK_TAG if dag_id == EXECUTE_CALLBACK_TAG and run_id.startswith(EXECUTE_CALLBACK_TAG): return ExecuteCallback.model_validate_json(command) # type: ignore[return-value] return ExecuteTask.model_validate_json(command)
@jobs_router.post( "/fetch/{worker_name}", dependencies=[Depends(jwt_token_authorization_rest)], responses=create_openapi_http_exception_doc( [ status.HTTP_400_BAD_REQUEST, status.HTTP_403_FORBIDDEN, status.HTTP_404_NOT_FOUND, status.HTTP_409_CONFLICT, ] ), )
[docs] def fetch( worker_name: str, body: Annotated[ WorkerQueuesBody, Body( title="Log data chunks", description="The queues and capacity from which the worker can fetch jobs.", ), ], session: SessionDep, ) -> EdgeJobFetched | None: """Fetch a job to execute on the edge worker.""" worker = session.scalar(select(EdgeWorkerModel).where(EdgeWorkerModel.worker_name == worker_name)) if not worker: raise HTTPException(status.HTTP_404_NOT_FOUND, "Worker not found") query = ( select(EdgeJobModel) .where( EdgeJobModel.state == TaskInstanceState.QUEUED, EdgeJobModel.concurrency_slots <= body.free_concurrency, ) .order_by(EdgeJobModel.queued_dttm) ) if body.queues: query = query.where(EdgeJobModel.queue.in_(body.queues)) if worker.team_name is not None: query = query.where(EdgeJobModel.team_name == worker.team_name) query = query.limit(1) query = query.with_for_update(skip_locked=True) job: EdgeJobModel | None = session.scalar(query) if not job: return None if job.task_instance_id and not body.supports_task_instance_uuid: log.warning("Edge worker %s cannot fetch UUID-keyed jobs; upgrade the worker.", worker_name) raise HTTPException( status.HTTP_409_CONFLICT, "Upgrade this Edge worker to report task-instance UUIDs." ) job.state = TaskInstanceState.RESTARTING # keep this intermediate state until worker sets to running job.edge_worker = worker_name job.last_update = timezone.utcnow() session.commit() # Edge worker does not backport emitted Airflow metrics, so export some metrics tags = prune_dict( {"dag_id": job.dag_id, "task_id": job.task_id, "queue": job.queue, "team_name": job.team_name} ) Stats.incr("edge_worker.ti.start", tags=tags) return EdgeJobFetched( dag_id=job.dag_id, task_id=job.task_id, run_id=job.run_id, map_index=job.map_index, try_number=job.try_number, command=parse_command(job.command, job.dag_id, job.run_id), concurrency_slots=job.concurrency_slots, task_instance_id=UUID(job.task_instance_id) if job.task_instance_id else None, )
@jobs_router.patch( "/state/{dag_id}/{task_id}/{run_id}/{try_number}/{map_index}/{state}", dependencies=[Depends(jwt_token_authorization_rest)], responses=create_openapi_http_exception_doc( [ status.HTTP_400_BAD_REQUEST, status.HTTP_403_FORBIDDEN, ] ), )
[docs] def state( dag_id: Annotated[str, WorkerApiDocs.dag_id], task_id: Annotated[str, WorkerApiDocs.task_id], run_id: Annotated[str, WorkerApiDocs.run_id], try_number: Annotated[int, WorkerApiDocs.try_number], map_index: Annotated[int, WorkerApiDocs.map_index], state: Annotated[TaskInstanceState, WorkerApiDocs.state], session: SessionDep, task_instance_id: Annotated[UUID | None, Body(embed=True)] = None, ) -> None: """Update the state of a job running on the edge worker.""" job = session.scalar( select(EdgeJobModel) .where( EdgeJobModel.dag_id == dag_id, EdgeJobModel.task_id == task_id, EdgeJobModel.run_id == run_id, EdgeJobModel.map_index == map_index, EdgeJobModel.try_number == try_number, EdgeJobModel.task_instance_id == (str(task_instance_id) if task_instance_id else ""), ) .with_for_update() ) if job is None: sibling_identity = session.scalar( select(EdgeJobModel.task_instance_id) .where( EdgeJobModel.dag_id == dag_id, EdgeJobModel.task_id == task_id, EdgeJobModel.run_id == run_id, EdgeJobModel.map_index == map_index, EdgeJobModel.try_number == try_number, ) .limit(1) ) if sibling_identity is not None: log.warning( "Ignoring Edge state report for %s.%s run %s try %s map %s: " "task-instance UUID %s does not match a stored job.", dag_id, task_id, run_id, try_number, map_index, task_instance_id, ) return if job.state == TaskInstanceState.RUNNING and state in ( TaskInstanceState.SUCCESS, TaskInstanceState.FAILED, ): tags = { "dag_id": job.dag_id, "task_id": job.task_id, "queue": job.queue, "state": str(state), "team_name": job.team_name, } Stats.incr("edge_worker.ti.finish", tags=prune_dict(tags)) job.state = state job.last_update = timezone.utcnow()

Was this entry helpful?