diff --git a/providers/edge3/docs/migrations-ref.rst b/providers/edge3/docs/migrations-ref.rst index 7d0846c0029a6..45600d053a2be 100644 --- a/providers/edge3/docs/migrations-ref.rst +++ b/providers/edge3/docs/migrations-ref.rst @@ -34,7 +34,9 @@ Here's the list of all the Database Migrations that are executed via when you ru +-------------------------+------------------+-----------------+----------------------------------------------------------+ | Revision ID | Revises ID | Edge3 Version | Description | +=========================+==================+=================+==========================================================+ -| ``c6b3c3d093fd`` (head) | ``a09c3ee8e1d3`` | ``3.5.0`` | Replace individual counters with extended JSON based | +| ``f2a4b6c8d0e1`` (head) | ``c6b3c3d093fd`` | ``5.0.0`` | Add task instance identity to Edge jobs. | ++-------------------------+------------------+-----------------+----------------------------------------------------------+ +| ``c6b3c3d093fd`` | ``a09c3ee8e1d3`` | ``3.5.0`` | Replace individual counters with extended JSON based | | | | | sysinfo. | +-------------------------+------------------+-----------------+----------------------------------------------------------+ | ``a09c3ee8e1d3`` | ``8c275b6fbaa8`` | ``3.4.0`` | Add team_name column to edge_job and edge_worker tables. | diff --git a/providers/edge3/src/airflow/providers/edge3/cli/api_client.py b/providers/edge3/src/airflow/providers/edge3/cli/api_client.py index b7ce2316bf61d..3f81fa79ca946 100644 --- a/providers/edge3/src/airflow/providers/edge3/cli/api_client.py +++ b/providers/edge3/src/airflow/providers/edge3/cli/api_client.py @@ -16,6 +16,7 @@ # under the License. from __future__ import annotations +import json import logging import os from datetime import datetime @@ -24,6 +25,7 @@ from pathlib import Path from typing import TYPE_CHECKING, Any from urllib.parse import quote, urljoin +from uuid import UUID from aiohttp import ClientConnectionError, ClientResponseError, ServerTimeoutError, request from retryhttp import retry, wait_retry_after @@ -174,13 +176,18 @@ async def jobs_fetch( queues: list[str] | None, free_concurrency: int, team_name: str | None = None, + *, + supports_task_instance_uuid: bool = False, ) -> EdgeJobFetched | None: """Fetch a job to execute on the edge worker.""" result = await _make_generic_request( "POST", f"jobs/fetch/{quote(hostname)}", WorkerQueuesBody( - queues=queues, free_concurrency=free_concurrency, team_name=team_name + queues=queues, + free_concurrency=free_concurrency, + team_name=team_name, + supports_task_instance_uuid=supports_task_instance_uuid, ).model_dump_json(exclude_unset=True), ) if result: @@ -188,11 +195,14 @@ async def jobs_fetch( return None -async def jobs_set_state(key: TaskInstanceKey, state: TaskInstanceState) -> None: +async def jobs_set_state( + key: TaskInstanceKey, state: TaskInstanceState, *, task_instance_id: UUID | None = None +) -> None: """Set the state of a job.""" await _make_generic_request( "PATCH", f"jobs/state/{key.dag_id}/{key.task_id}/{key.run_id}/{key.try_number}/{key.map_index}/{state}", + json.dumps({"task_instance_id": str(task_instance_id)}) if task_instance_id else None, ) diff --git a/providers/edge3/src/airflow/providers/edge3/cli/worker.py b/providers/edge3/src/airflow/providers/edge3/cli/worker.py index a661f634e6e7a..6f4f6be791845 100644 --- a/providers/edge3/src/airflow/providers/edge3/cli/worker.py +++ b/providers/edge3/src/airflow/providers/edge3/cli/worker.py @@ -672,7 +672,13 @@ async def loop(self): async def fetch_and_run_job(self) -> None: """Fetch, start and monitor a new job.""" logger.debug("Attempting to fetch a new job...") - edge_job = await jobs_fetch(self.hostname, self.queues, self.free_concurrency, self.team_name) + edge_job = await jobs_fetch( + self.hostname, + self.queues, + self.free_concurrency, + self.team_name, + supports_task_instance_uuid=True, + ) if not edge_job: logger.debug( "No new job to process%s", @@ -689,7 +695,9 @@ async def fetch_and_run_job(self) -> None: job = self._launch_job(edge_job, workload, logfile) self.jobs.append(job) try: - await jobs_set_state(edge_job.key, TaskInstanceState.RUNNING) + await jobs_set_state( + edge_job.key, TaskInstanceState.RUNNING, task_instance_id=edge_job.task_instance_id + ) # As we got one job, directly fetch another one if possible if self.free_concurrency > 0: @@ -707,7 +715,11 @@ async def fetch_and_run_job(self) -> None: if job.is_success: logger.info("Job completed: %s", job.edge_job.identifier) - await jobs_set_state(job.edge_job.key, TaskInstanceState.SUCCESS) + await jobs_set_state( + job.edge_job.key, + TaskInstanceState.SUCCESS, + task_instance_id=job.edge_job.task_instance_id, + ) else: ex_txt = job.failure_details() logger.error("Job failed: %s with:\n%s", job.edge_job.identifier, ex_txt) @@ -718,7 +730,9 @@ async def fetch_and_run_job(self) -> None: log_chunk_time=timezone.utcnow(), log_chunk_data=f"Error executing job:\n{ex_txt}", ) - await jobs_set_state(job.edge_job.key, TaskInstanceState.FAILED) + await jobs_set_state( + job.edge_job.key, TaskInstanceState.FAILED, task_instance_id=job.edge_job.task_instance_id + ) finally: self.jobs.remove(job) # Cleanup temp files used for the job 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 cbc5e5e981141..3f3a3a423c421 100644 --- a/providers/edge3/src/airflow/providers/edge3/executors/edge_executor.py +++ b/providers/edge3/src/airflow/providers/edge3/executors/edge_executor.py @@ -17,17 +17,21 @@ from __future__ import annotations +import json import logging from collections.abc import Sequence +from contextlib import suppress from copy import deepcopy from datetime import datetime, timedelta from typing import TYPE_CHECKING, Any +from uuid import UUID -from sqlalchemy import delete, select +from sqlalchemy import case, delete, literal, select from airflow.executors import workloads from airflow.executors.base_executor import BaseExecutor from airflow.models.taskinstance import TaskInstance +from airflow.models.taskinstancekey import TaskInstanceKey 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 @@ -49,12 +53,14 @@ if AIRFLOW_V_3_4_PLUS: from airflow.executors.workloads.base import WorkloadType +with suppress(ImportError): + from airflow.executors.workloads.types import TaskInstanceUuid + 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 CommandType = Sequence[str] @@ -77,10 +83,13 @@ class EdgeExecutor(BaseExecutor): """Implementation of the EdgeExecutor to distribute work to Edge Workers via HTTP.""" supports_multi_team: bool = True + supports_task_instance_uuid = hasattr(BaseExecutor, "get_task_key") def __init__(self, *args, **kwargs): super().__init__(*args, **kwargs) - self.last_reported_state: dict[TaskInstanceKey | CallbackKey, TaskInstanceState | str] = {} + self.last_reported_state: dict[ + TaskInstanceUuid | 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. @@ -125,7 +134,7 @@ def queue_workload( session: Session, ) -> None: """Put new workload to queue. Airflow 3 entry point to execute a task.""" - key: TaskInstanceKey | CallbackKey + key: TaskInstanceUuid | TaskInstanceKey | CallbackKey if is_callback_execute(workload): existing_job = session.scalars( select(EdgeJobModel).where( @@ -156,18 +165,22 @@ def queue_workload( key = workload.key elif isinstance(workload, workloads.ExecuteTask): task_instance = workload.ti - key = task_instance.key + coordinates = task_instance.key + key = self.get_task_key(task_instance) if self.supports_task_instance_uuid else coordinates + task_instance_id = str(task_instance.id) if self.supports_task_instance_uuid else "" # Check if job already exists with same dag_id, task_id, run_id, map_index, try_number - existing_job = session.scalars( + matching_jobs = 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, + EdgeJobModel.dag_id == coordinates.dag_id, + EdgeJobModel.task_id == coordinates.task_id, + EdgeJobModel.run_id == coordinates.run_id, + EdgeJobModel.map_index == coordinates.map_index, + EdgeJobModel.try_number == coordinates.try_number, + EdgeJobModel.task_instance_id.in_((task_instance_id, "")), ) - ).first() + ) + existing_job = next((job for job in matching_jobs if self._job_key(job) == key), None) if existing_job: existing_job.state = TaskInstanceState.QUEUED @@ -178,11 +191,12 @@ def queue_workload( 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, + dag_id=coordinates.dag_id, + task_id=coordinates.task_id, + run_id=coordinates.run_id, + map_index=coordinates.map_index, + try_number=coordinates.try_number, + task_instance_id=task_instance_id, state=TaskInstanceState.QUEUED, queue=task_instance.queue, concurrency_slots=task_instance.pool_slots, @@ -261,13 +275,20 @@ def _update_orphaned_jobs(self, session: Session) -> bool: ).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, - ) + key = self._job_key(job) + if self.supports_task_instance_uuid and isinstance(key, TaskInstanceUuid): + ti = session.scalar(select(TaskInstance).where(TaskInstance.id == key.id)) + if ti is None: + self.running.discard(key) + self.last_reported_state.pop(key, None) + else: + 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: @@ -284,23 +305,44 @@ def _update_orphaned_jobs(self, session: Session) -> bool: return bool(lifeless_jobs) + def _job_key(self, job: EdgeJobModel) -> TaskInstanceUuid | TaskInstanceKey | CallbackKey: + key = job.key + if self.supports_task_instance_uuid and isinstance(key, TaskInstanceKey): + return TaskInstanceUuid(UUID(job.task_instance_id or json.loads(job.command)["ti"]["id"])) + return key + def _get_tracked_job_keys( self, session: Session, states: Sequence[TaskInstanceState] - ) -> set[TaskInstanceKey | CallbackKey]: + ) -> set[TaskInstanceUuid | 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. """ + command = ( + case((EdgeJobModel.task_instance_id == "", EdgeJobModel.command), else_=None) + if self.supports_task_instance_uuid + else literal(None) + ) query = select( EdgeJobModel.dag_id, EdgeJobModel.task_id, EdgeJobModel.run_id, EdgeJobModel.try_number, EdgeJobModel.map_index, + EdgeJobModel.task_instance_id, + command.label("command"), ).where(EdgeJobModel.team_name == self.team_name, EdgeJobModel.state.in_(states)) - return {build_job_key(*row) for row in session.execute(query)} + keys: set[TaskInstanceUuid | TaskInstanceKey | CallbackKey] = set() + for job in session.execute(query): + key: TaskInstanceUuid | TaskInstanceKey | CallbackKey = build_job_key( + job.dag_id, job.task_id, job.run_id, job.try_number, job.map_index + ) + if self.supports_task_instance_uuid and isinstance(key, TaskInstanceKey): + key = TaskInstanceUuid(UUID(job.task_instance_id or json.loads(job.command)["ti"]["id"])) + keys.add(key) + return keys def _purge_jobs(self, session: Session) -> bool: """Clean finished jobs.""" @@ -324,26 +366,24 @@ def _purge_jobs(self, session: Session) -> bool: ) for job in jobs: - if job.key in self.running: + key = self._job_key(job) + if 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 + if key not in self.last_reported_state or self.last_reported_state[key] != job.state: + self.running_state(key) + self.last_reported_state[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) + if key in self.last_reported_state: + del self.last_reported_state[key] + self.success(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) + if key in self.last_reported_state: + del self.last_reported_state[key] + self.fail(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) + self.last_reported_state[key] = TaskInstanceState(job.state) if ( job.state == TaskInstanceState.SUCCESS and job.last_update_t < (datetime.now() - timedelta(minutes=job_success_purge)).timestamp() @@ -357,22 +397,39 @@ def _purge_jobs(self, session: Session) -> bool: ) 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] + if key in self.last_reported_state: + del self.last_reported_state[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, - ) - ) + self._delete_job(job, session) return purged_marker + @staticmethod + def _delete_job(job: EdgeJobModel, session: Session) -> None: + session.delete(job) + session.flush() + 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, + ~select(EdgeJobModel.dag_id) + .where( + EdgeJobModel.dag_id == job.dag_id, + EdgeJobModel.run_id == job.run_id, + EdgeJobModel.task_id == job.task_id, + EdgeJobModel.map_index == job.map_index, + EdgeJobModel.try_number == job.try_number, + EdgeJobModel.task_instance_id != job.task_instance_id, + ) + .exists(), + ) + .execution_options(synchronize_session=False) + ) + @provide_session def sync(self, *, session: Session = NEW_SESSION) -> None: """Sync will get called periodically by the heartbeat method.""" @@ -401,26 +458,30 @@ def revoke_task(self, *, ti: TaskInstance, session: Session = NEW_SESSION): :param ti: Task instance to revoke :param session: Database session """ - # Remove from executor's internal state - self.running.discard(ti.key) + key = self.get_task_key(ti) if self.supports_task_instance_uuid else ti.key + self.running.discard(key) if AIRFLOW_V_3_4_PLUS: - self.executor_queues[WorkloadType.EXECUTE_TASK].pop(ti.key, None) + self.executor_queues[WorkloadType.EXECUTE_TASK].pop(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] + self.queued_tasks.pop(key, None) + self.last_reported_state.pop(key, None) - # Delete the job from the database to prevent edge workers from picking it up - session.execute( - delete(EdgeJobModel).where( + jobs = session.scalars( + select(EdgeJobModel) + .with_for_update() + .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, + EdgeJobModel.team_name == self.team_name, ) ) - self.log.info("Revoked task instance %s from EdgeExecutor", ti.key) + for job in jobs: + if self._job_key(job) == key: + self._delete_job(job, session) + self.log.info("Revoked task instance %s from EdgeExecutor", key) def try_adopt_task_instances(self, tis: Sequence[TaskInstance]) -> Sequence[TaskInstance]: """ @@ -440,8 +501,14 @@ def try_adopt_task_instances(self, tis: Sequence[TaskInstance]) -> Sequence[Task 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] + rejected = [] + for ti in tis: + key = self.get_task_key(ti) if self.supports_task_instance_uuid else ti.key + if key in tracked_keys: + self.running.add(key) + else: + rejected.append(ti) + return rejected @staticmethod def get_cli_commands() -> list[GroupCommand]: diff --git a/providers/edge3/src/airflow/providers/edge3/migrations/versions/0006_5_0_0_add_task_instance_id_to_edge_job.py b/providers/edge3/src/airflow/providers/edge3/migrations/versions/0006_5_0_0_add_task_instance_id_to_edge_job.py new file mode 100644 index 0000000000000..fc67adfa78203 --- /dev/null +++ b/providers/edge3/src/airflow/providers/edge3/migrations/versions/0006_5_0_0_add_task_instance_id_to_edge_job.py @@ -0,0 +1,71 @@ +# 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. +""" +Add task instance identity to Edge jobs. + +Revision ID: f2a4b6c8d0e1 +Revises: c6b3c3d093fd +Create Date: 2026-09-29 00:00:00.000000 +""" + +from __future__ import annotations + +import sqlalchemy as sa +from alembic import op +from sqlalchemy.dialects import mysql + +revision = "f2a4b6c8d0e1" +down_revision = "c6b3c3d093fd" +branch_labels = None +depends_on = None +edge3_version = "5.0.0" + + +_COORDINATES = ["dag_id", "task_id", "run_id", "map_index", "try_number"] +_PK_NAMING = {"pk": "%(table_name)s_pkey"} + + +def upgrade() -> None: + identity_type = sa.String(36).with_variant( + mysql.VARCHAR(36, charset="ascii", collation="ascii_bin"), "mysql" + ) + with op.batch_alter_table("edge_job", naming_convention=_PK_NAMING) as batch_op: + batch_op.add_column(sa.Column("task_instance_id", identity_type, nullable=False, server_default="")) + batch_op.drop_constraint("edge_job_pkey", type_="primary") + batch_op.create_primary_key("edge_job_pkey", [*_COORDINATES, "task_instance_id"]) + + +def downgrade() -> None: + if op.get_context().as_sql: + raise RuntimeError( + "Edge job identity downgrade requires an online check for duplicate task coordinates" + ) + jobs = sa.table("edge_job", *(sa.column(name) for name in _COORDINATES)) + duplicate = ( + op.get_bind() + .execute(sa.select(*jobs.c).group_by(*jobs.c).having(sa.func.count() > 1).limit(1)) + .first() + ) + if duplicate is not None: + raise RuntimeError( + "Cannot downgrade Edge jobs: multiple task instances share the same coordinates. " + "Remove duplicate attempt jobs before downgrading." + ) + with op.batch_alter_table("edge_job", naming_convention=_PK_NAMING) as batch_op: + batch_op.drop_constraint("edge_job_pkey", type_="primary") + batch_op.create_primary_key("edge_job_pkey", _COORDINATES) + batch_op.drop_column("task_instance_id") diff --git a/providers/edge3/src/airflow/providers/edge3/models/db.py b/providers/edge3/src/airflow/providers/edge3/models/db.py index c93fb454ecef6..f997afe49b0ad 100644 --- a/providers/edge3/src/airflow/providers/edge3/models/db.py +++ b/providers/edge3/src/airflow/providers/edge3/models/db.py @@ -47,6 +47,7 @@ def _callable_accepts_use_migration_files(callable_: Any) -> bool: "3.2.0": "8c275b6fbaa8", "3.4.0": "a09c3ee8e1d3", "3.5.0": "c6b3c3d093fd", + "5.0.0": "f2a4b6c8d0e1", } 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 0309a57d6d434..4519218100337 100644 --- a/providers/edge3/src/airflow/providers/edge3/models/edge_job.py +++ b/providers/edge3/src/airflow/providers/edge3/models/edge_job.py @@ -25,6 +25,7 @@ String, text, ) +from sqlalchemy.dialects import mysql from sqlalchemy.orm import Mapped from airflow.models.base import StringID @@ -73,6 +74,13 @@ class EdgeJobModel(Base, LoggingMixin): Integer, primary_key=True, nullable=False, server_default=text("-1") ) try_number: Mapped[int] = mapped_column(Integer, primary_key=True, default=0) + task_instance_id: Mapped[str] = mapped_column( + String(36).with_variant(mysql.VARCHAR(36, charset="ascii", collation="ascii_bin"), "mysql"), + primary_key=True, + nullable=False, + default="", + server_default="", + ) state: Mapped[str] = mapped_column(String(20)) queue: Mapped[str] = mapped_column(String(256)) concurrency_slots: Mapped[int] = mapped_column(Integer) @@ -97,6 +105,7 @@ def __init__( edge_worker: str | None = None, last_update: datetime | None = None, team_name: str | None = None, + task_instance_id: str = "", ): self.dag_id = dag_id self.task_id = task_id @@ -111,6 +120,7 @@ def __init__( self.edge_worker = edge_worker self.last_update = last_update self.team_name = team_name + self.task_instance_id = task_instance_id super().__init__() __table_args__ = (Index("rj_order", state, queued_dttm, queue),) diff --git a/providers/edge3/src/airflow/providers/edge3/plugins/www/openapi-gen/requests/schemas.gen.ts b/providers/edge3/src/airflow/providers/edge3/plugins/www/openapi-gen/requests/schemas.gen.ts index 51756edfb4452..ea682def7f983 100644 --- a/providers/edge3/src/airflow/providers/edge3/plugins/www/openapi-gen/requests/schemas.gen.ts +++ b/providers/edge3/src/airflow/providers/edge3/plugins/www/openapi-gen/requests/schemas.gen.ts @@ -183,6 +183,12 @@ export const $Job = { title: 'Try Number', description: 'The number of attempt to execute this task.' }, + task_instance_id: { + type: 'string', + title: 'Task Instance Id', + description: 'Task-instance UUID, or empty for a legacy job.', + default: '' + }, state: { '$ref': '#/components/schemas/TaskInstanceState', description: 'State of the job from the view of the executor.' @@ -601,6 +607,12 @@ export const $WorkerQueuesBody = { type: 'integer', title: 'Free Concurrency', description: 'Number of free concurrency slots on the worker.' + }, + supports_task_instance_uuid: { + type: 'boolean', + title: 'Supports Task Instance Uuid', + description: 'Whether the worker reports task-instance UUIDs when updating jobs.', + default: false } }, type: 'object', diff --git a/providers/edge3/src/airflow/providers/edge3/plugins/www/openapi-gen/requests/types.gen.ts b/providers/edge3/src/airflow/providers/edge3/plugins/www/openapi-gen/requests/types.gen.ts index 4ac26c50b1a40..900cb62b36728 100644 --- a/providers/edge3/src/airflow/providers/edge3/plugins/www/openapi-gen/requests/types.gen.ts +++ b/providers/edge3/src/airflow/providers/edge3/plugins/www/openapi-gen/requests/types.gen.ts @@ -97,6 +97,10 @@ export type Job = { * The number of attempt to execute this task. */ try_number: number; + /** + * Task-instance UUID, or empty for a legacy job. + */ + task_instance_id?: string; /** * State of the job from the view of the executor. */ @@ -260,6 +264,12 @@ export type WorkerQueuesBody = { * Number of free concurrency slots on the worker. */ free_concurrency: number; + /** + * Supports Task Instance Uuid + * + * Whether the worker reports task-instance UUIDs when updating jobs. + */ + supports_task_instance_uuid?: boolean; }; /** diff --git a/providers/edge3/src/airflow/providers/edge3/plugins/www/src/pages/JobsPage.test.tsx b/providers/edge3/src/airflow/providers/edge3/plugins/www/src/pages/JobsPage.test.tsx new file mode 100644 index 0000000000000..9bf3d32500113 --- /dev/null +++ b/providers/edge3/src/airflow/providers/edge3/plugins/www/src/pages/JobsPage.test.tsx @@ -0,0 +1,105 @@ +/*! + * 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. + */ +import { render, screen } from "@testing-library/react"; +import type { Job, JobCollectionResponse } from "openapi/requests/types.gen"; +import { MemoryRouter } from "react-router-dom"; +import { expect, it, vi } from "vitest"; + +import { JobsPage } from "src/pages/JobsPage"; + +const query = vi.hoisted((): { data: JobCollectionResponse } => ({ + data: { jobs: [], total_entries: 0 }, +})); + +vi.mock("openapi/queries", () => ({ useUiServiceJobs: () => query })); +vi.mock("@chakra-ui/react", () => ({ + Box: "div", + HStack: "div", + Text: "span", + Table: { Root: "table", Header: "thead", Row: "tr", ColumnHeader: "th", Body: "tbody", Cell: "td" }, +})); +vi.mock("src/components/ui", () => ({ + Select: { Root: "div", Trigger: "div", ValueText: () => null, Content: "div", Item: "div" }, +})); +vi.mock("src/components/SearchBar", () => ({ SearchBar: () => null })); +vi.mock("src/components/ErrorAlert", () => ({ ErrorAlert: () => null })); +vi.mock("src/components/StateBadge", () => ({ StateBadge: "span" })); +vi.mock("src/constants", () => ({ jobStateOptions: { items: [] } })); +vi.mock("src/utils", () => ({ autoRefreshInterval: 1000 })); + +it("keeps same-coordinate attempts distinct when jobs reorder or disappear", () => { + const error = vi.spyOn(console, "error").mockImplementation(() => undefined); + const coordinates: Omit = { + dag_id: "dag", + task_id: "task", + run_id: "run", + map_index: -1, + try_number: 1, + state: "running", + queue: "default", + }; + const first: Job = { + ...coordinates, + task_instance_id: "00000000-0000-0000-0000-000000000001", + edge_worker: "first", + }; + const second: Job = { + ...coordinates, + task_instance_id: "00000000-0000-0000-0000-000000000002", + edge_worker: "second", + }; + const legacy: Job = { ...coordinates, task_instance_id: "", edge_worker: "legacy" }; + query.data = { jobs: [first, second, legacy], total_entries: 3 }; + const page = ( + + + + ); + const { rerender } = render(page); + const rowText = () => + screen + .getAllByRole("row") + .slice(1) + .map((row) => row.textContent); + + expect(rowText()).toEqual([ + expect.stringContaining("first"), + expect.stringContaining("second"), + expect.stringContaining("legacy"), + ]); + query.data = { jobs: [second, first, legacy], total_entries: 3 }; + rerender( + + + , + ); + expect(rowText()).toEqual([ + expect.stringContaining("second"), + expect.stringContaining("first"), + expect.stringContaining("legacy"), + ]); + query.data = { jobs: [first, legacy], total_entries: 2 }; + rerender( + + + , + ); + expect(rowText()).toEqual([expect.stringContaining("first"), expect.stringContaining("legacy")]); + expect(error.mock.calls.filter(([message]) => String(message).includes("same key"))).toHaveLength(0); +}); diff --git a/providers/edge3/src/airflow/providers/edge3/plugins/www/src/pages/JobsPage.tsx b/providers/edge3/src/airflow/providers/edge3/plugins/www/src/pages/JobsPage.tsx index cf86d89dbdb9d..a9f106e351a5c 100644 --- a/providers/edge3/src/airflow/providers/edge3/plugins/www/src/pages/JobsPage.tsx +++ b/providers/edge3/src/airflow/providers/edge3/plugins/www/src/pages/JobsPage.tsx @@ -17,8 +17,9 @@ * under the License. */ import { Box, HStack, Table, Text, type SelectValueChangeDetails } from "@chakra-ui/react"; -import { useState, useCallback, useEffect } from "react"; import { useUiServiceJobs } from "openapi/queries"; +import type { TaskInstanceState } from "openapi/requests/types.gen"; +import { useState, useCallback, useEffect } from "react"; import { Link, useSearchParams } from "react-router-dom"; import TimeAgo from "react-timeago"; @@ -28,7 +29,6 @@ import { StateBadge } from "src/components/StateBadge"; import { Select } from "src/components/ui"; import { jobStateOptions } from "src/constants"; import { autoRefreshInterval } from "src/utils"; -import type { TaskInstanceState } from "openapi/requests/types.gen"; export const JobsPage = () => { const [searchParams] = useSearchParams(); @@ -207,7 +207,14 @@ export const JobsPage = () => { {data.jobs.map((job) => ( {job.dag_id} @@ -240,7 +247,9 @@ export const JobsPage = () => { {job.queued_dttm ? : undefined} - {job.edge_worker} + + {job.edge_worker} + {job.last_update ? : undefined} diff --git a/providers/edge3/src/airflow/providers/edge3/plugins/www/vite.config.ts b/providers/edge3/src/airflow/providers/edge3/plugins/www/vite.config.ts index 441adeaf31d77..2c843ebc6bc53 100644 --- a/providers/edge3/src/airflow/providers/edge3/plugins/www/vite.config.ts +++ b/providers/edge3/src/airflow/providers/edge3/plugins/www/vite.config.ts @@ -23,7 +23,7 @@ import dts from "vite-plugin-dts"; import { defineConfig } from "vitest/config"; // https://vitejs.dev/config/ -export default defineConfig(({ command }) => { +export default defineConfig(({ command, mode }) => { const isLibraryBuild = command === "build"; return { @@ -57,7 +57,7 @@ export default defineConfig(({ command }) => { global: "globalThis", "process.env": "{}", // Define process.env for browser compatibility - "process.env.NODE_ENV": JSON.stringify("production"), + "process.env.NODE_ENV": JSON.stringify(mode === "test" ? "test" : "production"), }, plugins: [ react(), diff --git a/providers/edge3/src/airflow/providers/edge3/worker_api/datamodels.py b/providers/edge3/src/airflow/providers/edge3/worker_api/datamodels.py index fde62612df856..59c9e5c1710c7 100644 --- a/providers/edge3/src/airflow/providers/edge3/worker_api/datamodels.py +++ b/providers/edge3/src/airflow/providers/edge3/worker_api/datamodels.py @@ -18,6 +18,7 @@ from datetime import datetime from typing import Annotated +from uuid import UUID from fastapi import Path from pydantic import BaseModel, Field @@ -76,6 +77,9 @@ class EdgeJobFetched(EdgeJobBase): ), ] concurrency_slots: Annotated[int, Field(description="Number of concurrency slots the job requires.")] + task_instance_id: UUID | None = Field( + default=None, description="Attempt UUID required when reporting this job's state." + ) @property def identifier(self) -> str: @@ -119,6 +123,9 @@ class WorkerQueuesBody(WorkerQueuesBase): """Queues that a worker supports to run jobs on.""" free_concurrency: Annotated[int, Field(description="Number of free concurrency slots on the worker.")] + supports_task_instance_uuid: bool = Field( + default=False, description="Whether the worker reports task-instance UUIDs when updating jobs." + ) class WorkerStateBody(WorkerQueuesBase): diff --git a/providers/edge3/src/airflow/providers/edge3/worker_api/datamodels_ui.py b/providers/edge3/src/airflow/providers/edge3/worker_api/datamodels_ui.py index e38671928a60a..bdbde0aec49fa 100644 --- a/providers/edge3/src/airflow/providers/edge3/worker_api/datamodels_ui.py +++ b/providers/edge3/src/airflow/providers/edge3/worker_api/datamodels_ui.py @@ -48,6 +48,8 @@ class WorkerCollectionResponse(BaseModel): class Job(EdgeJobBase): """Details of the job sent to the scheduler.""" + task_instance_id: str = Field(default="", description="Task-instance UUID, or empty for a legacy job.") + state: Annotated[TaskInstanceState, Field(description="State of the job from the view of the executor.")] queue: Annotated[ str, diff --git a/providers/edge3/src/airflow/providers/edge3/worker_api/routes/jobs.py b/providers/edge3/src/airflow/providers/edge3/worker_api/routes/jobs.py index cf992b7dc8deb..5d1b791b7637c 100644 --- a/providers/edge3/src/airflow/providers/edge3/worker_api/routes/jobs.py +++ b/providers/edge3/src/airflow/providers/edge3/worker_api/routes/jobs.py @@ -17,10 +17,12 @@ 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, update +from sqlalchemy import select from airflow.api_fastapi.common.db.common import SessionDep # noqa: TC001 from airflow.api_fastapi.common.router import AirflowRouter @@ -42,6 +44,8 @@ if TYPE_CHECKING: from airflow.providers.edge3.models.types import ExecuteTypeBody +log = logging.getLogger(__name__) + jobs_router = AirflowRouter(tags=["Jobs"], prefix="/jobs") @@ -64,6 +68,7 @@ def parse_command(command: str, dag_id: str, run_id: str) -> ExecuteTypeBody: status.HTTP_400_BAD_REQUEST, status.HTTP_403_FORBIDDEN, status.HTTP_404_NOT_FOUND, + status.HTTP_409_CONFLICT, ] ), ) @@ -100,6 +105,11 @@ def fetch( 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() @@ -117,6 +127,7 @@ def fetch( 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, ) @@ -138,41 +149,56 @@ def state( 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.""" - # execute query to catch the queue and check if state toggles to success or failed - # otherwise possible that Executor resets orphaned jobs and stats are exported 2 times - if state in [TaskInstanceState.SUCCESS, state == TaskInstanceState.FAILED]: - query = 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.state == TaskInstanceState.RUNNING, - ) - job = session.scalar(query) - - if job: - # Edge worker does not backport emitted Airflow metrics, so export some metrics - 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)) - - query2 = ( - update(EdgeJobModel) + 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 ""), ) - .values(state=state, last_update=timezone.utcnow()) + .with_for_update() ) - session.execute(query2) + 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() diff --git a/providers/edge3/src/airflow/providers/edge3/worker_api/routes/ui.py b/providers/edge3/src/airflow/providers/edge3/worker_api/routes/ui.py index 59a39f9b4cfc0..6f2d3192dbbbc 100644 --- a/providers/edge3/src/airflow/providers/edge3/worker_api/routes/ui.py +++ b/providers/edge3/src/airflow/providers/edge3/worker_api/routes/ui.py @@ -138,6 +138,7 @@ def jobs( run_id=j.run_id, map_index=j.map_index, try_number=j.try_number, + task_instance_id=j.task_instance_id, state=TaskInstanceState(j.state), queue=j.queue, queued_dttm=j.queued_dttm, diff --git a/providers/edge3/src/airflow/providers/edge3/worker_api/v2-edge-generated.yaml b/providers/edge3/src/airflow/providers/edge3/worker_api/v2-edge-generated.yaml index 13cacba97be36..9ec177fa3194d 100644 --- a/providers/edge3/src/airflow/providers/edge3/worker_api/v2-edge-generated.yaml +++ b/providers/edge3/src/airflow/providers/edge3/worker_api/v2-edge-generated.yaml @@ -71,6 +71,12 @@ paths: schema: $ref: '#/components/schemas/HTTPExceptionResponse' description: Not Found + '409': + content: + application/json: + schema: + $ref: '#/components/schemas/HTTPExceptionResponse' + description: Conflict '422': description: Validation Error content: @@ -143,6 +149,11 @@ paths: description: JWT Authorization Token title: Authorization description: JWT Authorization Token + requestBody: + content: + application/json: + schema: + $ref: '#/components/schemas/Body_state' responses: '200': description: Successful Response @@ -955,6 +966,16 @@ paths: $ref: '#/components/schemas/HTTPValidationError' components: schemas: + Body_state: + properties: + task_instance_id: + anyOf: + - type: string + format: uuid + - type: 'null' + title: Task Instance Id + type: object + title: Body_state BundleInfo: properties: name: @@ -1054,6 +1075,13 @@ components: type: integer title: Concurrency Slots description: Number of concurrency slots the job requires. + task_instance_id: + anyOf: + - type: string + format: uuid + - type: 'null' + title: Task Instance Id + description: Attempt UUID required when reporting this job's state. type: object required: - dag_id @@ -1195,6 +1223,11 @@ components: type: integer title: Try Number description: The number of attempt to execute this task. + task_instance_id: + type: string + title: Task Instance Id + description: Task-instance UUID, or empty for a legacy job. + default: '' state: $ref: '#/components/schemas/TaskInstanceState' description: State of the job from the view of the executor. @@ -1571,6 +1604,12 @@ components: type: integer title: Free Concurrency description: Number of free concurrency slots on the worker. + supports_task_instance_uuid: + type: boolean + title: Supports Task Instance Uuid + description: Whether the worker reports task-instance UUIDs when updating + jobs. + default: false type: object required: - free_concurrency diff --git a/providers/edge3/tests/unit/edge3/cli/test_api_client.py b/providers/edge3/tests/unit/edge3/cli/test_api_client.py index 8554d4be79629..3e2b71e3a2ede 100644 --- a/providers/edge3/tests/unit/edge3/cli/test_api_client.py +++ b/providers/edge3/tests/unit/edge3/cli/test_api_client.py @@ -16,14 +16,25 @@ # under the License. from __future__ import annotations +import json from dataclasses import dataclass from http import HTTPStatus from unittest.mock import patch +from uuid import uuid4 import pytest from aiohttp import ClientResponseError, ConnectionTimeoutError - -from airflow.providers.edge3.cli.api_client import _make_generic_request +from fastapi import Request + +from airflow.providers.common.compat.sdk import TaskInstanceKey +from airflow.providers.edge3.cli.api_client import ( + _make_generic_request, + jobs_fetch, + jobs_set_state, + jwt_generator, +) +from airflow.providers.edge3.worker_api import auth +from airflow.utils.state import TaskInstanceState from tests_common.test_utils.aiohttp import MockAiohttpClientResponse from tests_common.test_utils.config import conf_vars @@ -78,6 +89,25 @@ def _request(method: str, *, url: str, data: str | None = None, headers=None): return _request +@pytest.mark.parametrize("uuid_capability", [None, False, True]) +async def test_jobs_fetch_sends_worker_capability(mocker, uuid_capability): + request = mocker.patch( + "airflow.providers.edge3.cli.api_client._make_generic_request", autospec=True, return_value=None + ) + kwargs = {} if uuid_capability is None else {"supports_task_instance_uuid": uuid_capability} + + assert await jobs_fetch("worker", ["queue"], 2, **kwargs) is None + + request.assert_awaited_once() + assert request.call_args.args[:2] == ("POST", "jobs/fetch/worker") + assert json.loads(request.call_args.args[2]) == { + "queues": ["queue"], + "free_concurrency": 2, + "team_name": None, + "supports_task_instance_uuid": bool(uuid_capability), + } + + class TestApiClient: @patch.dict("os.environ", {"AIRFLOW__EDGE__API_RETRIES": "10"}, clear=False) async def test_make_generic_request_success(self): @@ -145,3 +175,45 @@ async def test_make_generic_request_unrecoverable_error(self, mock_sleep): assert err.value.status == HTTPStatus.INTERNAL_SERVER_ERROR assert len(calls) == 10 assert all(call[0] == "POST" and call[1] == unreliable_service for call in calls) + + +@pytest.mark.parametrize("uuid_job", [False, True]) +async def test_job_state_identity_preserves_signed_request_path(uuid_job): + task_id = uuid4() if uuid_job else None + captured = {} + key = TaskInstanceKey("dag", "task", "run", 1, -1) + + def send_request(method, *, url, data, headers): + captured.update(url=url, data=data, headers=headers) + return _MockRequestContext( + response=MockAiohttpClientResponse( + status=HTTPStatus.NO_CONTENT, method=method, url=url, reason="No Content" + ) + ) + + with ( + conf_vars( + { + ("edge", "api_url"): "https://worker/edge_worker/v1/", + ("api_auth", "jwt_secret"): "uuid-test-secret", + } + ), + patch("airflow.providers.edge3.cli.api_client.request", autospec=True, side_effect=send_request), + ): + jwt_generator.cache_clear() + auth.jwt_validator.cache_clear() + try: + await jobs_set_state(key, TaskInstanceState.SUCCESS, task_instance_id=task_id) + request = Request( + { + "type": "http", + "path": "/edge_worker/v1/jobs/state/dag/task/run/1/-1/success", + "headers": [], + } + ) + await auth.jwt_token_authorization_rest(request, captured["headers"]["Authorization"]) + finally: + jwt_generator.cache_clear() + auth.jwt_validator.cache_clear() + assert captured["url"] == "https://worker/edge_worker/v1/jobs/state/dag/task/run/1/-1/success" + assert captured["data"] == (json.dumps({"task_instance_id": str(task_id)}) if task_id else None) diff --git a/providers/edge3/tests/unit/edge3/cli/test_worker.py b/providers/edge3/tests/unit/edge3/cli/test_worker.py index fb5cd1eb4480c..def0cbb46eca5 100644 --- a/providers/edge3/tests/unit/edge3/cli/test_worker.py +++ b/providers/edge3/tests/unit/edge3/cli/test_worker.py @@ -31,6 +31,7 @@ from typing import TYPE_CHECKING from unittest import mock from unittest.mock import call, patch +from uuid import uuid4 import anyio import pytest @@ -489,6 +490,31 @@ def test_fork_child_exits_nonzero_when_supervisor_raises( assert error_file_path.exists() assert "supervisor crashed" in error_file_path.read_text() + @pytest.mark.asyncio + async def test_fetch_capability_is_not_overridden_by_sysinfo_hook( + self, worker_with_job_and_sysinfo: EdgeWorker, mocker + ): + mocker.patch.object( + worker_with_job_and_sysinfo, + "extended_sysinfo", + autospec=True, + return_value={"supports_task_instance_uuid": False}, + ) + fetch = mocker.patch( + "airflow.providers.edge3.cli.worker.jobs_fetch", autospec=True, return_value=None + ) + + assert (await worker_with_job_and_sysinfo._get_sysinfo())["supports_task_instance_uuid"] is False + await worker_with_job_and_sysinfo.fetch_and_run_job() + + fetch.assert_awaited_once_with( + worker_with_job_and_sysinfo.hostname, + worker_with_job_and_sysinfo.queues, + worker_with_job_and_sysinfo.free_concurrency, + worker_with_job_and_sysinfo.team_name, + supports_task_instance_uuid=True, + ) + @patch("airflow.providers.edge3.cli.worker.jobs_fetch") @patch("airflow.providers.edge3.cli.worker.EdgeWorker._launch_job") @pytest.mark.asyncio @@ -508,12 +534,13 @@ async def test_fetch_and_run_job_no_job( @patch("airflow.providers.edge3.cli.worker.jobs_fetch") @patch("airflow.providers.edge3.cli.worker.EdgeWorker._launch_job") - @patch("airflow.providers.edge3.cli.worker.jobs_set_state") + @patch("airflow.providers.edge3.cli.worker.jobs_set_state", autospec=True) @patch("airflow.providers.edge3.cli.worker.EdgeWorker._push_logs_in_chunks") @patch("airflow.providers.edge3.cli.worker.logs_push") @patch.object(Job, "is_running", property(lambda _: False)) @patch.object(Job, "is_success", property(lambda _: True)) @pytest.mark.asyncio + @pytest.mark.parametrize("uuid_job", [False, True]) async def test_fetch_and_run_job_one_job( self, mock_logs_push, @@ -523,8 +550,11 @@ async def test_fetch_and_run_job_one_job( mock_jobs_fetch, tmp_path: Path, worker_with_job: EdgeWorker, + uuid_job, ): + task_instance_id = uuid4() if uuid_job else None edge_job = EdgeJobFetched( + task_instance_id=task_instance_id, dag_id="test", task_id="test", run_id="test", @@ -546,13 +576,16 @@ async def test_fetch_and_run_job_one_job( mock_launch_job.assert_called_once_with( edge_job, edge_job.command, Path(worker_with_job.base_log_folder, "mock.log") ) - assert mock_jobs_set_state.call_count == 2 + assert mock_jobs_set_state.call_args_list == [ + call(edge_job.key, TaskInstanceState.RUNNING, task_instance_id=task_instance_id), + call(edge_job.key, TaskInstanceState.SUCCESS, task_instance_id=task_instance_id), + ] mock_push_log_chunks.assert_called_once() assert len(worker_with_job.jobs) == 1 # no new job added (was removed at the end...) mock_logs_push.assert_not_called() @patch("airflow.providers.edge3.cli.worker.jobs_fetch") - @patch("airflow.providers.edge3.cli.worker.jobs_set_state") + @patch("airflow.providers.edge3.cli.worker.jobs_set_state", autospec=True) @patch("airflow.providers.edge3.cli.worker.EdgeWorker._push_logs_in_chunks") @patch("airflow.providers.edge3.cli.worker.logs_push") @pytest.mark.asyncio @@ -565,7 +598,9 @@ async def test_fetch_and_run_job_fork_failure_pushes_error_to_logs( tmp_path: Path, worker_with_job: EdgeWorker, ): + task_instance_id = uuid4() edge_job = EdgeJobFetched( + task_instance_id=task_instance_id, dag_id="test", task_id="test", run_id="test", @@ -588,7 +623,9 @@ async def test_fetch_and_run_job_fork_failure_pushes_error_to_logs( mock_jobs_fetch.assert_called_once() mock_push_log_chunks.assert_called_once() - assert mock_jobs_set_state.call_args_list[-1].args[1] == TaskInstanceState.FAILED + mock_jobs_set_state.assert_called_with( + edge_job.key, TaskInstanceState.FAILED, task_instance_id=task_instance_id + ) log_chunk_data = mock_logs_push.call_args.kwargs["log_chunk_data"] assert "Task fork exited with code 1" in log_chunk_data assert "RuntimeError: supervisor crashed" in log_chunk_data @@ -1056,6 +1093,7 @@ async def test_get_sysinfo(self, worker_with_job: EdgeWorker): assert sysinfo["status"] == logging.INFO assert "status_text" not in sysinfo # is only defined if extended sysinfo provides this field assert sysinfo["concurrency"] == concurrency + assert "supports_task_instance_uuid" not in sysinfo @pytest.mark.asyncio async def test_get_sysinfo_version_mismatch(self, worker_with_job: EdgeWorker): 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 2f5ab17c0c2e2..ef9c7124d6de4 100644 --- a/providers/edge3/tests/unit/edge3/executors/test_edge_executor.py +++ b/providers/edge3/tests/unit/edge3/executors/test_edge_executor.py @@ -22,11 +22,11 @@ from pathlib import Path from unittest import mock from unittest.mock import MagicMock, patch -from uuid import uuid4 +from uuid import UUID, uuid4 import pytest import time_machine -from sqlalchemy import delete, select +from sqlalchemy import delete, event, select from airflow.executors.workloads import BundleInfo, ExecuteTask from airflow.jobs.job import Job @@ -35,11 +35,14 @@ 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 +from airflow.providers.edge3.models.edge_logs import EdgeLogsModel from airflow.providers.edge3.models.edge_worker import EdgeWorkerModel, EdgeWorkerState from airflow.providers.edge3.models.types import EXECUTE_CALLBACK_TAG +from airflow.providers.edge3.worker_api.routes.jobs import state as set_job_state from airflow.utils.session import create_session from airflow.utils.state import TaskInstanceState +from tests_common.test_utils.asserts import assert_queries_count from tests_common.test_utils.compat import EmptyOperator from tests_common.test_utils.config import conf_vars from tests_common.test_utils.version_compat import AIRFLOW_V_3_2_PLUS, AIRFLOW_V_3_3_PLUS @@ -62,9 +65,104 @@ ) +@pytest.mark.parametrize( + "uuid_executor", + [ + False, + pytest.param( + True, + marks=pytest.mark.skipif( + not hasattr(EdgeExecutor, "get_task_key"), + reason="UUID executor contract requires a supporting Airflow version", + ), + ), + ], +) +def test_tracking_projects_keys_and_only_needed_commands(session, monkeypatch, uuid_executor): + monkeypatch.setattr(EdgeExecutor, "supports_task_instance_uuid", uuid_executor) + executor = EdgeExecutor() + executor.team_name = "projection" + legacy_command = '{"ti":{"id":"00000000-0000-0000-0000-000000000002"}}' + callback_id = str(UUID(int=3)) + rows = [ + EdgeJobModel( + dag_id=dag_id, + task_id=task_id, + run_id=run_id, + try_number=try_number, + map_index=-1, + task_instance_id=task_instance_id, + command=command, + state=state, + team_name=team, + queue="default", + concurrency_slots=1, + ) + for dag_id, task_id, run_id, try_number, task_instance_id, command, state, team in [ + ("dag", "uuid", "run", 1, str(UUID(int=1)), "invalid", TaskInstanceState.QUEUED, "projection"), + ("dag", "legacy", "run", 1, "", legacy_command, TaskInstanceState.QUEUED, "projection"), + ( + EXECUTE_CALLBACK_TAG, + callback_id, + f"{EXECUTE_CALLBACK_TAG}-{callback_id}", + 0, + "", + "invalid", + TaskInstanceState.QUEUED, + "projection", + ), + ( + EXECUTE_CALLBACK_TAG, + "task", + "run", + 1, + str(UUID(int=4)), + "invalid", + TaskInstanceState.QUEUED, + "projection", + ), + ("dag", "other_team", "run", 1, str(UUID(int=5)), "invalid", TaskInstanceState.QUEUED, "other"), + ( + "dag", + "finished", + "run", + 1, + str(UUID(int=6)), + "invalid", + TaskInstanceState.SUCCESS, + "projection", + ), + ] + ] + session.add_all(rows) + session.flush() + expected = {executor._job_key(job) for job in rows[:4]} + statements = [] + + def capture(execution): + statements.append(execution.statement) + + event.listen(session, "do_orm_execute", capture) + try: + with assert_queries_count(1, session=session): + assert executor._get_tracked_job_keys(session, (TaskInstanceState.QUEUED,)) == expected + finally: + event.remove(session, "do_orm_execute", capture) + + assert len(statements) == 1 + fetched = session.connection().execute(statements[0]).all() + assert {job.task_id: job.command for job in fetched} == { + "uuid": None, + "legacy": legacy_command if uuid_executor else None, + callback_id: "invalid" if uuid_executor else None, + "task": None, + } + + class TestEdgeExecutor: @pytest.fixture(autouse=True) - def setup_test_cases(self): + def setup_test_cases(self, monkeypatch): + monkeypatch.setattr(EdgeExecutor, "supports_task_instance_uuid", False) with create_session() as session: session.execute(delete(EdgeJobModel)) @@ -493,7 +591,8 @@ class TestEdgeExecutorMultiTeam: """Tests for multi-team (AIP-67) support in EdgeExecutor.""" @pytest.fixture(autouse=True) - def setup_test_cases(self): + def setup_test_cases(self, monkeypatch): + monkeypatch.setattr(EdgeExecutor, "supports_task_instance_uuid", False) with create_session() as session: session.execute(delete(EdgeJobModel)) session.execute(delete(EdgeWorkerModel)) @@ -699,9 +798,46 @@ def test_no_team_executor_processes_all_jobs(self): remaining_jobs = session.scalars(select(EdgeJobModel)).all() assert len(remaining_jobs) == 2 + def test_purging_all_coordinate_siblings_removes_logs(self): + executor = EdgeExecutor() + with create_session() as session: + session.execute(delete(EdgeLogsModel)) + for identity in (UUID(int=1), UUID(int=2)): + session.add( + EdgeJobModel( + dag_id="dag", + task_id="task", + run_id="run", + map_index=-1, + try_number=1, + task_instance_id=str(identity), + state=TaskInstanceState.SUCCESS, + queue="default", + concurrency_slots=1, + command="mock", + last_update=timezone.utcnow() - timedelta(days=1), + ) + ) + session.add( + EdgeLogsModel( + dag_id="dag", + task_id="task", + run_id="run", + map_index=-1, + try_number=1, + log_chunk_time=timezone.utcnow(), + log_chunk_data="remove", + ) + ) + session.flush() + assert executor._purge_jobs(session) + session.flush() + assert session.scalars(select(EdgeJobModel)).all() == [] + assert session.scalars(select(EdgeLogsModel)).all() == [] + @pytest.mark.skipif(not AIRFLOW_V_3_3_PLUS, reason="ExecuteTypeBody union requires Airflow 3.3+") -class TestQueueWorkload: +class _WorkloadFactory: @pytest.fixture(autouse=True) def setup(self): with create_session() as session: @@ -743,6 +879,12 @@ def _make_execute_callback(self) -> ExecuteCallback: log_path="test.log", ) + +class TestQueueWorkload(_WorkloadFactory): + @pytest.fixture(autouse=True) + def legacy_executor(self, monkeypatch): + monkeypatch.setattr(EdgeExecutor, "supports_task_instance_uuid", False) + def test_queue_workload_execute_task(self): executor = EdgeExecutor() workload = self._make_execute_task() @@ -967,3 +1109,226 @@ def test_queue_workload_unknown_type_raises(self): with create_session() as session: with pytest.raises(TypeError, match="Don't know how to queue workload"): executor.queue_workload(MagicMock(spec=[]), session=session) + + +@pytest.mark.skipif( + not hasattr(EdgeExecutor, "get_task_key"), + reason="UUID executor contract requires a supporting Airflow version", +) +class TestUUIDTaskIdentity(_WorkloadFactory): + @pytest.fixture(autouse=True) + def uuid_executor(self, monkeypatch): + monkeypatch.setattr(EdgeExecutor, "supports_task_instance_uuid", True) + + @pytest.mark.parametrize("terminal_state", [TaskInstanceState.SUCCESS, TaskInstanceState.FAILED]) + @pytest.mark.parametrize("selected", [0, 1]) + def test_same_coordinate_jobs_report_their_own_uuid(self, terminal_state, selected): + first = self._make_execute_task() + first.ti.id = UUID(int=1) + second = first.model_copy(update={"ti": first.ti.model_copy(update={"id": UUID(int=2)})}) + target, sibling = (first, second) if selected == 0 else (second, first) + executor = EdgeExecutor() + target_key, sibling_key = executor.get_task_key(target.ti), executor.get_task_key(sibling.ti) + with create_session() as session: + executor.queue_workload(first, session=session) + executor.queue_workload(second, session=session) + session.flush() + assert len(session.scalars(select(EdgeJobModel)).all()) == 2 + assert executor.running == {target_key, sibling_key} + set_job_state( + dag_id=target.ti.dag_id, + task_id=target.ti.task_id, + run_id=target.ti.run_id, + try_number=1, + map_index=-1, + state=terminal_state, + session=session, + task_instance_id=target.ti.id, + ) + session.flush() + persisted = {job.task_instance_id: job.state for job in session.scalars(select(EdgeJobModel))} + assert persisted == { + str(target.ti.id): terminal_state, + str(sibling.ti.id): TaskInstanceState.QUEUED, + } + executor._purge_jobs(session) + assert executor.get_event_buffer(dag_ids={target.ti.dag_id}) == {target_key: (terminal_state, None)} + assert executor.running == {sibling_key} + + def test_adoption_and_revocation_match_uuid(self): + workload = self._make_execute_task() + replacement = workload.ti.model_copy(update={"id": uuid4()}) + executor = EdgeExecutor() + with create_session() as session: + executor.queue_workload(workload, session=session) + session.commit() + adopter = EdgeExecutor() + assert adopter.try_adopt_task_instances([workload.ti, replacement]) == [replacement] + adopter.revoke_task(ti=replacement, session=session) + assert session.scalar(select(EdgeJobModel)) is not None + adopter.revoke_task(ti=workload.ti, session=session) + session.flush() + assert session.scalar(select(EdgeJobModel)) is None + assert not adopter.running + + def test_legacy_worker_report_cannot_complete_new_uuid_job(self, monkeypatch): + first = self._make_execute_task() + second = first.model_copy(update={"ti": first.ti.model_copy(update={"id": uuid4()})}) + executor = EdgeExecutor() + with create_session() as session: + monkeypatch.setattr(EdgeExecutor, "supports_task_instance_uuid", False) + executor.queue_workload(first, session=session) + session.commit() + monkeypatch.setattr(EdgeExecutor, "supports_task_instance_uuid", True) + adopter = EdgeExecutor() + assert adopter.try_adopt_task_instances([first.ti]) == [] + adopter.queue_workload(second, session=session) + session.flush() + set_job_state( + dag_id=first.ti.dag_id, + task_id=first.ti.task_id, + run_id=first.ti.run_id, + try_number=1, + map_index=-1, + state=TaskInstanceState.SUCCESS, + session=session, + ) + session.flush() + adopter._purge_jobs(session) + assert adopter.get_event_buffer() == { + adopter.get_task_key(first.ti): (TaskInstanceState.SUCCESS, None) + } + assert adopter.running == {adopter.get_task_key(second.ti)} + + def test_resuming_legacy_attempt_reuses_its_job_row(self, monkeypatch): + workload = self._make_execute_task() + with create_session() as session: + monkeypatch.setattr(EdgeExecutor, "supports_task_instance_uuid", False) + EdgeExecutor().queue_workload(workload, session=session) + session.flush() + job = session.scalar(select(EdgeJobModel)) + job.state = TaskInstanceState.SUCCESS + session.flush() + monkeypatch.setattr(EdgeExecutor, "supports_task_instance_uuid", True) + executor = EdgeExecutor() + executor.queue_workload(workload, session=session) + session.flush() + assert len(session.scalars(select(EdgeJobModel)).all()) == 1 + assert job.task_instance_id == "" + assert job.state == TaskInstanceState.QUEUED + executor._purge_jobs(session) + assert executor.running == {executor.get_task_key(workload.ti)} + assert executor.get_event_buffer() == {} + + @pytest.mark.parametrize("revoke_last", [False, True]) + def test_purging_retired_job_keeps_logs_while_coordinate_sibling_exists(self, revoke_last): + first = self._make_execute_task() + second = first.model_copy(update={"ti": first.ti.model_copy(update={"id": uuid4()})}) + executor = EdgeExecutor() + with create_session() as session: + session.execute(delete(EdgeLogsModel)) + executor.queue_workload(first, session=session) + executor.queue_workload(second, session=session) + session.flush() + retired = session.scalar( + select(EdgeJobModel).where(EdgeJobModel.task_instance_id == str(first.ti.id)) + ) + retired.state = TaskInstanceState.SUCCESS + retired.last_update = timezone.utcnow() - timedelta(days=1) + session.add( + EdgeLogsModel( + dag_id=first.ti.dag_id, + task_id=first.ti.task_id, + run_id=first.ti.run_id, + try_number=1, + map_index=-1, + log_chunk_time=timezone.utcnow(), + log_chunk_data="keep", + ) + ) + session.flush() + executor._purge_jobs(session) + session.flush() + assert session.scalar(select(EdgeLogsModel)).log_chunk_data == "keep" + assert session.scalars(select(EdgeJobModel)).one().task_instance_id == str(second.ti.id) + if revoke_last: + executor.revoke_task(ti=second.ti, session=session) + assert session.scalar(select(EdgeJobModel)) is None + assert session.scalar(select(EdgeLogsModel)) is None + assert not executor.running + else: + session.execute(delete(EdgeLogsModel)) + + @pytest.mark.parametrize("legacy_job", [False, True]) + @pytest.mark.parametrize("state", [TaskInstanceState.RUNNING, TaskInstanceState.SUCCESS]) + def test_orphaned_job_uses_attempt_uuid( + self, create_task_instance, session, monkeypatch, legacy_job, state + ): + ti = create_task_instance(state=state) + ti.try_number = 1 + session.flush() + workload = ExecuteTask.make(ti, bundle_info=BundleInfo(name="testing")) + executor = EdgeExecutor() + if legacy_job: + monkeypatch.setattr(EdgeExecutor, "supports_task_instance_uuid", False) + executor.queue_workload(workload, session=session) + session.commit() + monkeypatch.setattr(EdgeExecutor, "supports_task_instance_uuid", True) + executor = EdgeExecutor() + assert executor.try_adopt_task_instances([ti]) == [] + key = executor.get_task_key(ti) + job = session.scalar(select(EdgeJobModel)) + job.state = TaskInstanceState.RUNNING + job.last_update = timezone.utcnow() - timedelta( + seconds=conf.getint("scheduler", "task_instance_heartbeat_timeout") + 1 + ) + session.flush() + assert executor._update_orphaned_jobs(session) + assert job.state == state + executor._purge_jobs(session) + assert executor.get_event_buffer(dag_ids={ti.dag_id}) == {key: (state, None)} + assert executor.running == ({key} if state == TaskInstanceState.RUNNING else set()) + + @pytest.mark.parametrize("legacy_job", [False, True]) + def test_retired_orphan_releases_its_slot_without_changing_retry( + self, create_task_instance, session, monkeypatch, legacy_job + ): + ti = create_task_instance(state=TaskInstanceState.RUNNING) + ti.task.retries = 1 + ti.try_number = 1 + ti.max_tries = 1 + session.commit() + workload = ExecuteTask.make(ti, bundle_info=BundleInfo(name="testing")) + executor = EdgeExecutor() + old_key = executor.get_task_key(ti) + if legacy_job: + monkeypatch.setattr(EdgeExecutor, "supports_task_instance_uuid", False) + executor.queue_workload(workload, session=session) + session.commit() + monkeypatch.setattr(EdgeExecutor, "supports_task_instance_uuid", True) + executor = EdgeExecutor() + assert executor.try_adopt_task_instances([ti]) == [] + retired = session.scalar(select(EdgeJobModel)) + retired.state = TaskInstanceState.RUNNING + retired.last_update = timezone.utcnow() - timedelta( + seconds=conf.getint("scheduler", "task_instance_heartbeat_timeout") + 1 + ) + session.flush() + ti.handle_failure("worker lost", session=session) + session.refresh(ti) + assert ti.id != workload.ti.id + assert ti.state == TaskInstanceState.UP_FOR_RETRY + ti.state = TaskInstanceState.RUNNING + replacement = ExecuteTask.make(ti, bundle_info=BundleInfo(name="testing")) + executor.queue_workload(replacement, session=session) + session.flush() + new_key = executor.get_task_key(ti) + assert executor.running == {old_key, new_key} + assert executor._update_orphaned_jobs(session) + assert retired.state == TaskInstanceState.REMOVED + assert executor.running == {new_key} + assert old_key not in executor.last_reported_state + executor._purge_jobs(session) + assert executor.running == {new_key} + assert executor.get_event_buffer() == {} + assert ti.state == TaskInstanceState.RUNNING diff --git a/providers/edge3/tests/unit/edge3/migrations/test_task_instance_identity.py b/providers/edge3/tests/unit/edge3/migrations/test_task_instance_identity.py new file mode 100644 index 0000000000000..9e97a60e8c96c --- /dev/null +++ b/providers/edge3/tests/unit/edge3/migrations/test_task_instance_identity.py @@ -0,0 +1,170 @@ +# 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 + +from importlib import import_module +from io import StringIO +from pathlib import Path +from uuid import uuid4 + +import pytest +import sqlalchemy as sa +from alembic.config import Config +from alembic.environment import EnvironmentContext +from alembic.migration import MigrationContext +from alembic.operations import Operations +from alembic.script import ScriptDirectory + +from airflow.migrations import db_types + +migration = import_module( + "airflow.providers.edge3.migrations.versions.0006_5_0_0_add_task_instance_id_to_edge_job" +) +COORDINATES = ["dag_id", "task_id", "run_id", "map_index", "try_number"] + + +def reflected_columns(connection): + return [ + column | {"type": str(column["type"])} for column in sa.inspect(connection).get_columns("edge_job") + ] + + +@pytest.fixture +def legacy_jobs(monkeypatch): + for name in ("TIMESTAMP", "StringID"): + monkeypatch.delitem(vars(db_types), name, raising=False) + for name in ("TIMESTAMP", "StringID"): + monkeypatch.setitem(vars(db_types), name, getattr(db_types, name)) + engine = sa.create_engine("sqlite://") + with engine.begin() as connection: + metadata = sa.MetaData(naming_convention={"pk": "%(table_name)s_pkey"}) + scripts = ScriptDirectory(str(Path(migration.__file__).parents[1])) + with EnvironmentContext(Config(), scripts) as environment: + environment.configure(connection=connection, target_metadata=metadata) + with Operations.context(environment.get_context()): + for revision in reversed( + list(scripts.walk_revisions(base="base", head=migration.down_revision)) + ): + revision.module.upgrade() + jobs = sa.Table("edge_job", sa.MetaData(), autoload_with=connection) + row = dict( + dag_id="dag", + task_id="task", + run_id="run", + map_index=-1, + try_number=1, + state="queued", + queue="default", + concurrency_slots=1, + command='{"ti":{"id":"00000000-0000-0000-0000-000000000001"}}', + ) + connection.execute(jobs.insert().values(**row)) + yield connection, row + engine.dispose() + + +def test_upgrade_preserves_existing_job_and_default_legacy_identity(legacy_jobs): + connection, original = legacy_jobs + migration.upgrade() + jobs = sa.Table("edge_job", sa.MetaData(), autoload_with=connection) + row = connection.execute(sa.select(jobs)).mappings().one() + assert row["task_instance_id"] == "" + assert all(row[key] == value for key, value in original.items()) + connection.execute(jobs.insert().values(**(original | {"task_id": "another"}))) + assert connection.scalar(sa.select(jobs.c.task_instance_id).where(jobs.c.task_id == "another")) == "" + inspector = sa.inspect(connection) + assert inspector.get_pk_constraint("edge_job")["constrained_columns"] == [ + *COORDINATES, + "task_instance_id", + ] + assert "rj_order" in {index["name"] for index in inspector.get_indexes("edge_job")} + + +def test_upgrade_allows_distinct_attempts_at_identical_coordinates(legacy_jobs): + connection, original = legacy_jobs + migration.upgrade() + jobs = sa.Table("edge_job", sa.MetaData(), autoload_with=connection) + identities = {"", str(uuid4()), str(uuid4())} + for identity in identities - {""}: + connection.execute(jobs.insert().values(**original, task_instance_id=identity)) + assert set(connection.scalars(sa.select(jobs.c.task_instance_id))) == identities + with pytest.raises(sa.exc.IntegrityError): + connection.execute(jobs.insert().values(**original, task_instance_id="")) + + +def test_downgrade_rejects_duplicate_legacy_keys_before_changing_schema(legacy_jobs): + connection, original = legacy_jobs + migration.upgrade() + jobs = sa.Table("edge_job", sa.MetaData(), autoload_with=connection) + identity = str(uuid4()) + connection.execute(jobs.insert().values(**original, task_instance_id=identity)) + before_columns = reflected_columns(connection) + before_pk = sa.inspect(connection).get_pk_constraint("edge_job") + with pytest.raises(RuntimeError, match="multiple task instances.*same coordinates"): + migration.downgrade() + assert reflected_columns(connection) == before_columns + assert sa.inspect(connection).get_pk_constraint("edge_job") == before_pk + assert set(connection.scalars(sa.select(jobs.c.task_instance_id))) == {"", identity} + + +@pytest.mark.parametrize("native", [False, True]) +def test_safe_downgrade_restores_original_schema_and_preserves_jobs(legacy_jobs, native): + connection, original = legacy_jobs + before_columns = reflected_columns(connection) + before_pk = sa.inspect(connection).get_pk_constraint("edge_job") + before_indexes = sa.inspect(connection).get_indexes("edge_job") + migration.upgrade() + jobs = sa.Table("edge_job", sa.MetaData(), autoload_with=connection) + if native: + connection.execute(jobs.update().values(task_instance_id=str(uuid4()))) + migration.downgrade() + inspector = sa.inspect(connection) + assert reflected_columns(connection) == before_columns + assert inspector.get_pk_constraint("edge_job") == before_pk + assert inspector.get_indexes("edge_job") == before_indexes + restored = sa.Table("edge_job", sa.MetaData(), autoload_with=connection) + row = connection.execute(sa.select(restored)).mappings().one() + assert all(row[key] == value for key, value in original.items()) + migration.upgrade() + assert "task_instance_id" in {column["name"] for column in sa.inspect(connection).get_columns("edge_job")} + + +@pytest.mark.parametrize("dialect", ["postgresql", "mysql"]) +def test_upgrade_compiles_primary_key_ddl_for_server_databases(dialect): + output = StringIO() + context = MigrationContext.configure(dialect_name=dialect, opts={"as_sql": True, "output_buffer": output}) + with Operations.context(context): + migration.upgrade() + ddl = output.getvalue() + assert "PRIMARY KEY (dag_id, task_id, run_id, map_index, try_number, task_instance_id)" in ddl + assert "task_instance_id VARCHAR(36)" in ddl + assert "NOT NULL" in ddl + assert "DEFAULT ''" in ddl + if dialect == "mysql": + assert "CHARACTER SET ascii COLLATE ascii_bin" in ddl + assert "DROP PRIMARY KEY" in ddl + else: + assert "DROP CONSTRAINT edge_job_pkey" in ddl + + +@pytest.mark.parametrize("dialect", ["postgresql", "mysql"]) +def test_offline_downgrade_refuses_before_emitting_schema_changes(dialect): + output = StringIO() + context = MigrationContext.configure(dialect_name=dialect, opts={"as_sql": True, "output_buffer": output}) + with Operations.context(context), pytest.raises(RuntimeError, match="requires an online check"): + migration.downgrade() + assert output.getvalue() == "" diff --git a/providers/edge3/tests/unit/edge3/models/test_db.py b/providers/edge3/tests/unit/edge3/models/test_db.py index 812526c9e564d..2405f39fdb445 100644 --- a/providers/edge3/tests/unit/edge3/models/test_db.py +++ b/providers/edge3/tests/unit/edge3/models/test_db.py @@ -17,11 +17,16 @@ from __future__ import annotations import warnings +from importlib import import_module from unittest import mock import pytest import sqlalchemy as sa +from alembic.migration import MigrationContext +from alembic.operations import Operations +from airflow import settings +from airflow.providers.edge3.models.edge_base import edge_metadata from airflow.utils.db_manager import RunDBManager from tests_common.test_utils.config import conf_vars @@ -29,6 +34,19 @@ pytestmark = [pytest.mark.db_test] +@pytest.fixture +def legacy_edge_job_table(): + migration = import_module("airflow.providers.edge3.migrations.versions.0001_3_0_0_create_edge_tables") + with settings.engine.begin() as connection: + edge_metadata.drop_all(connection) + with Operations.context(MigrationContext.configure(connection)): + migration.upgrade() + yield + with settings.engine.begin() as connection: + edge_metadata.drop_all(connection) + edge_metadata.create_all(connection) + + class TestEdgeDBManager: """Test EdgeDBManager functionality.""" @@ -220,7 +238,9 @@ def test_revision_heads_map_populated(self): assert "3.4.0" in _REVISION_HEADS_MAP assert _REVISION_HEADS_MAP["3.4.0"] == "a09c3ee8e1d3" - def test_initdb_stamps_and_upgrades_when_tables_exist_without_version(self, session): + def test_initdb_stamps_and_upgrades_when_tables_exist_without_version( + self, session, legacy_edge_job_table + ): """Test that initdb runs incremental migrations when tables exist but alembic version table does not.""" from sqlalchemy import inspect, text @@ -261,11 +281,13 @@ def test_initdb_stamps_and_upgrades_when_tables_exist_without_version(self, sess version = conn.execute(text("SELECT version_num FROM alembic_version_edge3")).scalar() columns = {col["name"] for col in inspect(conn).get_columns("edge_worker")} - assert version == "c6b3c3d093fd" + assert version == "f2a4b6c8d0e1" assert "concurrency" in columns assert "team_name" in columns - def test_upgradedb_stamps_and_upgrades_when_tables_exist_without_version(self, session): + def test_upgradedb_stamps_and_upgrades_when_tables_exist_without_version( + self, session, legacy_edge_job_table + ): """Test upgradedb runs incremental migrations when tables exist but alembic version table does not.""" from sqlalchemy import inspect, text @@ -310,7 +332,7 @@ def test_upgradedb_stamps_and_upgrades_when_tables_exist_without_version(self, s assert "concurrency" in columns assert "team_name" in columns - def test_migration_adds_concurrency_column(self, session): + def test_migration_adds_concurrency_column(self, session, legacy_edge_job_table): """Test that upgrading from 3.0.0 actually adds the concurrency column.""" from alembic import command from alembic.migration import MigrationContext diff --git a/providers/edge3/tests/unit/edge3/worker_api/routes/test_jobs.py b/providers/edge3/tests/unit/edge3/worker_api/routes/test_jobs.py index 0b09e47b633f8..4a53dad41581c 100644 --- a/providers/edge3/tests/unit/edge3/worker_api/routes/test_jobs.py +++ b/providers/edge3/tests/unit/edge3/worker_api/routes/test_jobs.py @@ -20,7 +20,7 @@ from pathlib import Path from typing import TYPE_CHECKING from unittest.mock import patch -from uuid import uuid4 +from uuid import UUID, uuid4 import pytest from fastapi import HTTPException, status @@ -32,6 +32,7 @@ from airflow.providers.edge3.models.edge_worker import EdgeWorkerModel, EdgeWorkerState from airflow.providers.edge3.models.types import EXECUTE_CALLBACK_TAG from airflow.providers.edge3.worker_api.datamodels import WorkerQueuesBody +from airflow.providers.edge3.worker_api.routes import jobs from airflow.providers.edge3.worker_api.routes.jobs import fetch, parse_command, state from airflow.utils.session import create_session from airflow.utils.state import TaskInstanceState @@ -84,9 +85,57 @@ def setup_test_cases(self, dag_maker, session: Session): session.execute(delete(EdgeJobModel)) session.execute(delete(EdgeWorkerModel)) session.commit() + yield + session.execute(delete(EdgeJobModel)) + session.execute(delete(EdgeWorkerModel)) + session.commit() - @patch(f"{Stats.__module__}.Stats.incr") - def test_state(self, mock_stats_incr, session: Session): + @pytest.mark.parametrize("uuid_worker", [None, False, True]) + def test_fetch_uuid_job_requires_worker_capability(self, session, uuid_worker, mocker): + warning = mocker.patch.object(jobs.log, "warning", autospec=True) + task_id = uuid4() + worker = EdgeWorkerModel( + worker_name="uuid_worker", + state=EdgeWorkerState.IDLE, + queues=[QUEUE], + ) + worker.sysinfo = {"supports_task_instance_uuid": not uuid_worker} + job = EdgeJobModel( + dag_id=DAG_ID, + task_id=TASK_ID, + run_id=RUN_ID, + try_number=1, + map_index=-1, + task_instance_id=str(task_id), + state=TaskInstanceState.QUEUED, + queue=QUEUE, + concurrency_slots=1, + command=MOCK_COMMAND_STR, + ) + session.add_all([worker, job]) + session.flush() + body = WorkerQueuesBody( + free_concurrency=1, + queues=[QUEUE], + **({"supports_task_instance_uuid": uuid_worker} if uuid_worker is not None else {}), + ) + if uuid_worker: + result = fetch("uuid_worker", body, session) + assert result.task_instance_id == task_id + assert job.state == TaskInstanceState.RESTARTING + warning.assert_not_called() + else: + with pytest.raises(HTTPException) as error: + fetch("uuid_worker", body, session) + assert error.value.status_code == 409 + assert job.state == TaskInstanceState.QUEUED + warning.assert_called_once_with( + "Edge worker %s cannot fetch UUID-keyed jobs; upgrade the worker.", "uuid_worker" + ) + + @pytest.mark.parametrize("terminal_state", [TaskInstanceState.SUCCESS, TaskInstanceState.FAILED]) + @patch(f"{Stats.__module__}.Stats.incr", autospec=True) + def test_state(self, mock_stats_incr, session: Session, terminal_state): with create_session() as session: job = EdgeJobModel( dag_id=DAG_ID, @@ -120,7 +169,7 @@ def test_state(self, mock_stats_incr, session: Session): run_id=RUN_ID, try_number=1, map_index=-1, - state=TaskInstanceState.SUCCESS, + state=terminal_state, session=session, ) @@ -129,7 +178,7 @@ def test_state(self, mock_stats_incr, session: Session): tags={ "dag_id": DAG_ID, "queue": QUEUE, - "state": TaskInstanceState.SUCCESS, + "state": terminal_state, "task_id": TASK_ID, "team_name": "team_a", }, @@ -138,7 +187,90 @@ def test_state(self, mock_stats_incr, session: Session): db_job: EdgeJobModel | None = session.scalar(select(EdgeJobModel)) assert db_job is not None - assert db_job.state == TaskInstanceState.SUCCESS + assert db_job.state == terminal_state + + @pytest.mark.parametrize( + "target", ["", "00000000-0000-0000-0000-000000000001", "00000000-0000-0000-0000-000000000002"] + ) + def test_state_updates_only_matching_attempt(self, session, target, mocker): + warning = mocker.patch.object(jobs.log, "warning", autospec=True) + identities = ["", "00000000-0000-0000-0000-000000000001", "00000000-0000-0000-0000-000000000002"] + for identity in identities: + session.add( + EdgeJobModel( + dag_id=DAG_ID, + task_id=TASK_ID, + run_id=RUN_ID, + try_number=1, + map_index=-1, + task_instance_id=identity, + state=TaskInstanceState.RUNNING, + queue=QUEUE, + concurrency_slots=1, + command="execute", + ) + ) + session.flush() + state( + DAG_ID, + TASK_ID, + RUN_ID, + 1, + -1, + TaskInstanceState.SUCCESS, + session, + task_instance_id=UUID(target) if target else None, + ) + session.flush() + session.expire_all() + actual = {job.task_instance_id: job.state for job in session.scalars(select(EdgeJobModel))} + assert actual == { + identity: TaskInstanceState.SUCCESS if identity == target else TaskInstanceState.RUNNING + for identity in identities + } + warning.assert_not_called() + + @pytest.mark.parametrize( + ("stored", "reported", "known_coordinates"), + [ + ("00000000-0000-0000-0000-000000000001", None, True), + ("00000000-0000-0000-0000-000000000001", "00000000-0000-0000-0000-000000000002", True), + ("", "00000000-0000-0000-0000-000000000002", True), + ("", None, False), + ], + ) + def test_state_warns_only_for_attempt_identity_mismatch( + self, session, mocker, stored, reported, known_coordinates + ): + warning = mocker.patch.object(jobs.log, "warning", autospec=True) + job = EdgeJobModel( + dag_id=DAG_ID, + task_id=TASK_ID, + run_id=RUN_ID, + try_number=1, + map_index=-1, + task_instance_id=stored, + state=TaskInstanceState.RUNNING, + queue=QUEUE, + concurrency_slots=1, + command="execute", + ) + session.add(job) + session.flush() + state( + DAG_ID if known_coordinates else "unknown", + TASK_ID, + RUN_ID, + 1, + -1, + TaskInstanceState.SUCCESS, + session, + task_instance_id=UUID(reported) if reported else None, + ) + session.flush() + session.expire_all() + assert session.scalar(select(EdgeJobModel)).state == TaskInstanceState.RUNNING + assert warning.call_count == int(known_coordinates) @patch(f"{Stats.__module__}.Stats.incr") def test_state_finish_metric_omits_team_name_for_global_job(self, mock_stats_incr, session: Session): diff --git a/providers/edge3/tests/unit/edge3/worker_api/routes/test_ui.py b/providers/edge3/tests/unit/edge3/worker_api/routes/test_ui.py index 256cdd977d396..c92ce4bb11ab3 100644 --- a/providers/edge3/tests/unit/edge3/worker_api/routes/test_ui.py +++ b/providers/edge3/tests/unit/edge3/worker_api/routes/test_ui.py @@ -17,14 +17,20 @@ from __future__ import annotations from typing import TYPE_CHECKING +from uuid import UUID import pytest from sqlalchemy import delete +from airflow.providers.edge3.models.edge_job import EdgeJobModel from airflow.providers.edge3.models.edge_worker import EdgeWorkerModel, EdgeWorkerState +from airflow.utils.state import TaskInstanceState from tests_common.test_utils.version_compat import AIRFLOW_V_3_1_PLUS +if AIRFLOW_V_3_1_PLUS: + from airflow.providers.edge3.worker_api.routes.ui import jobs + if TYPE_CHECKING: from sqlalchemy.orm import Session @@ -48,6 +54,31 @@ def test_worker(self, session: Session): assert len(worker_response.workers) == 1 assert worker_response.workers[0].worker_name == "worker1" + def test_jobs_preserves_same_coordinate_attempts_and_legacy_identity(self, session: Session): + attempts = {str(UUID(int=1)): "first", str(UUID(int=2)): "second", "": "legacy"} + session.add_all( + EdgeJobModel( + dag_id="ui_identity", + task_id="task", + run_id="run", + map_index=-1, + try_number=1, + task_instance_id=task_instance_id, + state=TaskInstanceState.RUNNING, + queue="default", + concurrency_slots=1, + command="unused", + edge_worker=worker, + ) + for task_instance_id, worker in attempts.items() + ) + session.flush() + + response = jobs(session=session, dag_id_pattern="ui_identity").model_dump(mode="json") + + assert response["total_entries"] == 3 + assert {job["task_instance_id"]: job["edge_worker"] for job in response["jobs"]} == attempts + def test_set_worker_concurrency_limit(self, session: Session): from airflow.providers.edge3.worker_api.datamodels_ui import ConcurrencyRequest from airflow.providers.edge3.worker_api.routes.ui import set_worker_concurrency_limit