github-advanced-security[bot] commented on code in PR #73750: URL: https://github.com/apache/airflow/pull/73750#discussion_r4110459638
########## dev/dag_parsing_poc/run_scheduler_parsing.py: ########## @@ -0,0 +1,586 @@ +# 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. +# /// script +# requires-python = ">=3.10" +# dependencies = ["apache-airflow-core", "apache-airflow-providers-celery", "cryptography"] +# /// +"""Run the scheduler-hosting experiment in Breeze with an isolated Celery worker and scheduler.""" + +from __future__ import annotations + +import argparse +import configparser +import json +import multiprocessing +import os +import socket +import sqlite3 +import statistics +import subprocess +import threading +import time +from contextlib import ExitStack +from datetime import datetime, timezone +from pathlib import Path +from uuid import uuid4 + +from cryptography.hazmat.primitives import serialization +from cryptography.hazmat.primitives.asymmetric.ed25519 import Ed25519PrivateKey + +REDIS_IMAGE = "public.ecr.aws/docker/library/redis:7.4-alpine" +ROUTE = "scheduler-parsing" + + +def write_json(path: Path, value) -> None: + path.write_text(json.dumps(value, indent=2) + "\n") + + +def docker(*arguments: str, check: bool = True) -> str: + result = subprocess.run(["docker", *arguments], check=False, capture_output=True, text=True, timeout=45) + if check and result.returncode: + raise RuntimeError(f"Docker {arguments[0]} failed: {result.stderr.strip()}") + return result.stdout.strip() + + +def wait_until(predicate, *, timeout: float = 120, check=None): + deadline = time.monotonic() + timeout + while time.monotonic() < deadline: + if value := predicate(): + return value + if check is not None: + check() + time.sleep(0.1) + raise TimeoutError("Scheduler-hosting checkpoint did not finish") + + +def stop_process(process) -> None: + if process.pid is None: + return + process.terminate() + process.join(5) + if process.is_alive(): + process.kill() + process.join(5) + + +def run_api(database: str, public: Path, listener: socket.socket) -> None: + # These fresh processes import Airflow only after the experiment environment is configured. + from airflow.dag_processing.executor_manager import _serve_api + + _serve_api(database, public, listener) + + +def run_provider(database: str, key: bytes, stop, capacity: int) -> None: + from airflow.api_fastapi.auth.tokens import JWTGenerator + from airflow.api_fastapi.execution_api.parsing import ( + TOKEN_AUDIENCE, + TOKEN_ISSUER, + TOKEN_KEY_ID, + TOKEN_SCOPE, + ) + from airflow.dag_processing.executor_runner import ParsingExecutorRunner + from airflow.dag_processing.parsing_metadata import MetadataOrchestrationStore + from airflow.executors.workloads import WorkloadType + from airflow.providers.celery.executors.celery_executor import CeleryExecutor + + generator = JWTGenerator( + private_key=Ed25519PrivateKey.from_private_bytes(key), + kid=TOKEN_KEY_ID, + issuer=TOKEN_ISSUER, + audience=TOKEN_AUDIENCE, + algorithm="EdDSA", + valid_for=600, + ) + + def issue_token(manifest): + return generator.generate( + { + "sub": manifest["workload_id"], + "scope": TOKEN_SCOPE, + "attempt_ids": [item["attempt_id"] for item in manifest["definitions"]], + } + ) + + executor = CeleryExecutor(parallelism=capacity) + executor.supported_workload_types = frozenset({WorkloadType.PARSE_DAG_DEFINITIONS}) + store = MetadataOrchestrationStore(database) + runner = ParsingExecutorRunner(store, executor, route=ROUTE, token_issuer=issue_token) + runner.start() + try: + while not stop.poll(0.05): + runner.tick() + finally: + runner.close() + store.engine.dispose() + + +class MetricsCapture: + """Capture the scheduler's StatsD measurements independently of its process.""" + + def __init__(self): + self.socket = socket.socket(socket.AF_INET, socket.SOCK_DGRAM) + self.socket.bind(("0.0.0.0", 0)) Review Comment: ## CodeQL / Binding a socket to all network interfaces Binding a socket to all interfaces (using ['0.0.0.0'](1)) is a security risk. [Show more details](https://github.com/apache/airflow/security/code-scanning/662) ########## dev/dag_parsing_poc/run_celery.py: ########## @@ -0,0 +1,306 @@ +# 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. +# /// script +# requires-python = ">=3.10" +# /// +"""Trusted Celery experiment driver; run inside the current worktree's Breeze image.""" + +from __future__ import annotations + +import argparse +import json +import multiprocessing +import os +import shutil +import socket +import time +from collections import Counter +from pathlib import Path + +from cryptography.hazmat.primitives import serialization +from cryptography.hazmat.primitives.asymmetric.ed25519 import Ed25519PrivateKey + +from airflow.api_fastapi.auth.tokens import JWTGenerator +from airflow.api_fastapi.execution_api.parsing import TOKEN_AUDIENCE, TOKEN_ISSUER, TOKEN_KEY_ID +from airflow.dag_processing.parsing_state import ReceiptStore +from airflow.executors.workloads import WorkloadType +from airflow.executors.workloads.parsing import ParseDagDefinitionsState + +from dev.dag_parsing_poc.run import create_archive, create_workloads, serve_api, wait_for_api, write_fixtures + + +def wait_until(predicate, *, description: str, timeout: float = 60) -> None: + deadline = time.monotonic() + timeout + while time.monotonic() < deadline: + if predicate(): + return + time.sleep(0.1) + raise TimeoutError(f"Timed out waiting for {description}") + + +def read_json(path: Path) -> dict: + return json.loads(path.read_text()) if path.exists() else {} + + +def read_task_events(output: Path, task_id: str) -> list[dict]: + path = output / "worker-evidence" / "task-events.jsonl" + if not path.exists(): + return [] + complete_lines = path.read_text().split("\n")[:-1] + return [event for line in complete_lines if (event := json.loads(line))["task_id"] == task_id] + + +def finish_workloads(executor, workloads, *, timeout: float = 120) -> dict: + terminal = {} + started = time.monotonic() + + def poll(): + executor.heartbeat() + for key, (state, info) in executor.get_event_buffer().items(): + if state in {ParseDagDefinitionsState.SUCCESS, ParseDagDefinitionsState.FAILED}: + terminal[str(key)] = {"state": state.value, "info": str(info) if info else None} + return len(terminal) == len(workloads) + + wait_until(poll, description="Celery terminal events", timeout=timeout) + if any(event["state"] != "success" for event in terminal.values()): + raise RuntimeError(f"Celery workload failed: {terminal}") + if executor.slots_available != executor.parallelism: + raise RuntimeError("Terminal workloads did not release executor slots") + return {"events": terminal, "elapsed_seconds": time.monotonic() - started} + + +def run_experiment(args, output: Path, generator: JWTGenerator, store: ReceiptStore) -> dict: + # The Celery provider caches config on import, so configure this fresh driver first. + os.environ["AIRFLOW__CELERY__BROKER_URL"] = args.broker_url + os.environ["AIRFLOW__CELERY__RESULT_BACKEND"] = args.result_backend + os.environ["AIRFLOW__CELERY__SYNC_PARALLELISM"] = "1" + from airflow.configuration import conf + from airflow.providers.celery.executors.celery_executor import CeleryExecutor + + source = output / "submitter-bundle" + worker_root = output / "worker-bundle" + files = write_fixtures(source, args.definitions, True) + for path in files[: args.definitions]: + with path.open("a") as stream: + stream.write( + "\nfrom pathlib import Path\n" + f"with Path('/worker-evidence/{path.stem}.imports').open('a') as marker:\n" + " marker.write('imported\\n')\n" + ) + slow = source / "active_duplicate.py" + slow.write_text( + "import time\nfrom pathlib import Path\nfrom airflow.sdk import DAG\n" + "with Path('/worker-evidence/active_duplicate.imports').open('a') as marker:\n" + " marker.write('imported\\n')\n" + "time.sleep(5)\nwith DAG('celery_active_duplicate', schedule=None):\n pass\n" + ) + archive = create_archive(files, source / "definitions.zip") if args.archive_members else None + shutil.copytree(source, worker_root) + (output / "worker-evidence").mkdir() + (output / "driver-ready.json").write_text(json.dumps({"port": args.port, "queue": args.queue}) + "\n") + wait_until( + lambda: read_json(output / "worker-evidence" / "isolation.json").get("stage") == "worker_ready", + description="isolated worker startup", + timeout=180, + ) + configured_executor = conf.get("core", "executor") + task_executor = CeleryExecutor(parallelism=2) + executor = CeleryExecutor(parallelism=1) + executor.supported_workload_types = frozenset({WorkloadType.PARSE_DAG_DEFINITIONS}) + executor.start() + try: + workloads = [ + workload.model_copy(update={"queue": args.queue}) + for workload in create_workloads( + files, batch_size=args.batch_size, timeout=30, generator=generator, archive_path=archive + ) + ] + for workload in workloads: + store.register_workload(workload) + executor.queue_workload(workload, session=None) + ordinary = finish_workloads(executor, workloads) + results = [result for workload in workloads for result in store.get_results(workload.workload_id)] + expected = {"success": args.definitions, "import_error": 1, "timeout": 1} + outcomes = dict(Counter(result["outcome"] for result in results)) + received_dags = { + result["relative_path"]: [dag.get("dag", {}).get("dag_id") for dag in result["serialized_dags"]] + for result in results + if result["outcome"] == "success" + } + if outcomes != expected or received_dags != { + f"{'definitions.zip/' if archive else ''}success_{index}.py": [f"executor_parsing_poc_{index}"] + for index in range(args.definitions) + }: + raise RuntimeError(f"Unexpected parsing output: {outcomes}, {received_dags}") + + task = executor.celery_app.tasks["execute_workload"] + replay = workloads[0] + replay_id = str(replay.workload_id) + before = store.get_attempts(replay.workload_id) + wait_until( + lambda: any(event["state"] == "SUCCESS" for event in read_task_events(output, replay_id)), + description="first batch completion evidence", + ) + initial_successes = sum(event["state"] == "SUCCESS" for event in read_task_events(output, replay_id)) + task.apply_async( + args=[replay.model_dump_json()], + queue=args.queue, + task_id=replay_id, + argsrepr="(<redacted parsing workload>,)", + ) + wait_until( + lambda: ( + sum(event["state"] == "SUCCESS" for event in read_task_events(output, replay_id)) + > initial_successes + ), + description="accepted batch replay", + ) + if store.get_attempts(replay.workload_id) != before: + raise RuntimeError("Replay changed accepted receipts") + + active = create_workloads([slow], batch_size=1, timeout=30, generator=generator)[0].model_copy( + update={"queue": args.queue} + ) + store.register_workload(active) + executor.queue_workload(active, session=None) + executor.heartbeat() + active_id = str(active.workload_id) + wait_until( + lambda: store.get_attempts(active.workload_id)[0]["status"] == "claimed", + description="original active claim", + ) + original_execution = store.get_attempts(active.workload_id)[0]["execution_id"] + task.apply_async( + args=[active.model_dump_json()], + queue=args.queue, + task_id=active_id, + argsrepr="(<redacted parsing workload>,)", + ) + wait_until( + lambda: any(event["state"] == "IGNORED" for event in read_task_events(output, active_id)), + description="duplicate being ignored", + ) + duplicate_backend_state = executor.celery_app.AsyncResult(active_id).state + if duplicate_backend_state in {"FAILURE", "REVOKED"}: + raise RuntimeError("Duplicate poisoned the original Celery result") + active_run = finish_workloads(executor, [active]) + accepted = store.get_attempts(active.workload_id)[0] + if accepted["execution_id"] != original_execution or accepted["status"] != "accepted": + raise RuntimeError("Duplicate replaced the original execution") + active_results = store.get_results(active.workload_id) + if ( + len(active_results) != 1 + or active_results[0]["outcome"] != "success" + or [dag.get("dag", {}).get("dag_id") for dag in active_results[0]["serialized_dags"]] + != ["celery_active_duplicate"] + ): + raise RuntimeError("Original execution did not publish the expected successful Dag") + import_counts = { + path.stem: len(path.read_text().splitlines()) + for path in (output / "worker-evidence").glob("*.imports") + } + if len(import_counts) != args.definitions + 1 or any(count != 1 for count in import_counts.values()): + raise RuntimeError(f"Definitions were imported more than once: {import_counts}") + if task_executor.slots_available != 2 or conf.get("core", "executor") != configured_executor: + raise RuntimeError("Parsing changed task routing or instance capacity") + if WorkloadType.PARSE_DAG_DEFINITIONS in task_executor.supported_workload_types: + raise RuntimeError("Parsing was enabled by default") + (output / "results.json").write_text(json.dumps(results, indent=2) + "\n") + return { + "mode": "dedicated CeleryExecutor, Redis broker/backend, isolated prefork worker", + "importer": "SDK", + "definition_kind": "zip_member" if archive else "file", + "definitions": len(files), + "dispatches": len(workloads), + "batch_size": args.batch_size, + "parsing_parallelism": 1, + "outcomes": outcomes, + "received_dags": received_dags, + "ordinary": ordinary, + "accepted_replay_preserved_receipts": True, + "active_duplicate": active_run, + "active_duplicate_outcome": active_results[0]["outcome"], + "duplicate_backend_state": duplicate_backend_state, + "import_counts": import_counts, + "independent_task_slots": task_executor.slots_available, + "isolation": read_json(output / "worker-evidence" / "isolation.json"), + "limitations": [ + "Development receipt API; no production metadata ingestion or parse-time reads.", + "SDK importing still uses core validation and serialization; not an SDK-only runtime.", + "No automatic recovery of abandoned claims or orchestrator restart/adoption.", + "Separate worker container on the same Docker host, not a multi-host deployment.", + "No concurrent task benchmark, full manager baseline, Kubernetes or HA proof.", + ], + } + finally: + executor.end() + + +def main() -> None: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--output", type=Path, required=True) + parser.add_argument("--broker-url", required=True) + parser.add_argument("--result-backend", required=True) + parser.add_argument("--queue", default="poc-parsing") + parser.add_argument("--port", type=int, default=8799) + parser.add_argument("--definitions", type=int, default=2) + parser.add_argument("--batch-size", type=int, choices=range(1, 101), default=2) + parser.add_argument("--archive-members", action="store_true") + args = parser.parse_args() + if args.definitions < 1: + parser.error("definitions must be positive") + output = args.output.resolve() + output.mkdir(parents=True, exist_ok=False) + key = Ed25519PrivateKey.generate() + public_key = output / "verification-key.pem" + public_key.write_bytes( + key.public_key().public_bytes( + serialization.Encoding.PEM, serialization.PublicFormat.SubjectPublicKeyInfo + ) + ) + generator = JWTGenerator( + private_key=key, + kid=TOKEN_KEY_ID, + issuer=TOKEN_ISSUER, + audience=TOKEN_AUDIENCE, + algorithm="EdDSA", + valid_for=360, + ) + store_path = output / "receipts.sqlite" + store = ReceiptStore(store_path) + listener = socket.socket() + listener.bind(("0.0.0.0", args.port)) Review Comment: ## CodeQL / Binding a socket to all network interfaces Binding a socket to all interfaces (using ['0.0.0.0'](1)) is a security risk. [Show more details](https://github.com/apache/airflow/security/code-scanning/659) ########## dev/dag_parsing_poc/run_celery_recovery.py: ########## @@ -0,0 +1,544 @@ +# 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. +# /// script +# requires-python = ">=3.10" +# /// +"""Measure current Celery recovery gaps using real broker, worker and HTTP processes.""" + +from __future__ import annotations + +import argparse +import json +import multiprocessing +import os +import re +import shutil +import socket +from datetime import timedelta +from pathlib import Path +from uuid import UUID, uuid4 + +import httpx +from cryptography.hazmat.primitives import serialization +from cryptography.hazmat.primitives.asymmetric.ed25519 import Ed25519PrivateKey + +from airflow._shared.timezones import timezone +from airflow.api_fastapi.auth.tokens import JWTGenerator +from airflow.api_fastapi.execution_api.parsing import TOKEN_AUDIENCE, TOKEN_ISSUER, TOKEN_KEY_ID, TOKEN_SCOPE +from airflow.dag_processing.executor_worker import compute_source_revision +from airflow.dag_processing.parsing_state import ReceiptStore +from airflow.executors.workloads import BundleInfo, WorkloadType +from airflow.executors.workloads.parsing import DagDefinitionAttempt, DagDefinitionResult, ParseDagDefinitions + +from dev.dag_parsing_poc.run import serve_api, wait_for_api +from dev.dag_parsing_poc.run_celery import read_json, read_task_events, wait_until + + +def write_json(path: Path, value: dict) -> None: + temporary = path.with_suffix(".tmp") + temporary.write_text(json.dumps(value, indent=2) + "\n") + temporary.replace(path) + + +def validate_worker_ready(worker: dict, *, run_id: str, queue: str) -> None: + if ( + worker.get("stage") != "worker_ready" + or worker.get("run_id") != run_id + or worker.get("queue") != queue + ): + raise RuntimeError("Worker readiness does not match this recovery run and queue") + try: + UUID(worker["worker_id"]) + except (KeyError, ValueError, TypeError) as error: + raise RuntimeError("Worker readiness has no valid startup identity") from error + + +def matches_worker_event(event: dict, worker: dict, *, task_id: str, state: str) -> bool: + return ( + event.get("task_id") == task_id + and event.get("state") == state + and all(event.get(key) == worker[key] for key in ("run_id", "worker_id", "container_hostname")) + ) + + +def read_import_events(path: Path) -> list[dict]: + return [json.loads(line) for line in path.read_text().split("\n")[:-1]] if path.exists() else [] + + +def validate_replacement( + control: dict, original: dict, replacement: dict, *, run_id: str, queue: str +) -> None: + if control.get("run_id") != run_id: + raise RuntimeError("Docker evidence belongs to a different recovery run") + for role, worker in (("original", original), ("replacement", replacement)): + validate_worker_ready(worker, run_id=run_id, queue=queue) + inspection = control[f"{role}_container"] + container_id = inspection["Id"] + hostname = worker.get("container_hostname") + if ( + not re.fullmatch(r"[0-9a-f]{64}", container_id) + or hostname != container_id[:12] + or hostname != inspection["Config"]["Hostname"] + ): + raise RuntimeError(f"Docker inspection does not identify the {role} worker's container") + if ( + control["original_container"]["Id"] == control["replacement_container"]["Id"] + or original["worker_id"] == replacement["worker_id"] + ): + raise RuntimeError("Replacement must be a different worker and container") + termination = control["original_container"]["State"] + if termination["Running"] or termination["Pid"] != 0 or termination["ExitCode"] != 137: + raise RuntimeError("Whole-worker kill was not confirmed by Docker") + replacement_state = control["replacement_container"]["State"] + if not replacement_state["Running"] or replacement_state["Pid"] <= 0: + raise RuntimeError("Replacement worker container is not running") + + +def delete_running_backend_record(backend, task_id: str) -> int: + if backend.get_task_meta(task_id, cache=False)["status"] != "STARTED": + raise RuntimeError("Missing-record scenario requires an existing STARTED result") + backend_key = backend.get_key_for_task(task_id) + deleted = backend.client.delete(backend_key) + if deleted != 1: + raise RuntimeError("Missing-record scenario must delete exactly one Redis result key") + if backend.client.exists(backend_key): + raise RuntimeError("Redis result key still exists after deletion") + return deleted + + +def create_executor(): + from airflow.providers.celery.executors.celery_executor import CeleryExecutor + + executor = CeleryExecutor(parallelism=1) + executor.supported_workload_types = frozenset({WorkloadType.PARSE_DAG_DEFINITIONS}) + executor.start() + return executor + + +def inspect_fresh_executor(store_path: str, workload_id: str, output_path: str) -> None: + executor = create_executor() + try: + write_json( + Path(output_path), + { + "pid": os.getpid(), + "running": len(executor.running), + "tracked_workloads": len(executor.workloads), + "available_slots": executor.slots_available, + "receipts": ReceiptStore(store_path).get_attempts(workload_id), + "reconstruction_attempted": False, + }, + ) + finally: + executor.end() + + +def create_workload(paths: list[Path], generator: JWTGenerator, *, queue: str, stop_seconds: float): + definitions = tuple( + DagDefinitionAttempt( + attempt_id=uuid4(), + relative_path=path.name, + source_revision=compute_source_revision(path), + timeout_seconds=120 if path.stem == "interrupted" else 10, + ) + for path in paths + ) + workload_id = uuid4() + now = timezone.utcnow() + return ParseDagDefinitions( + workload_id=workload_id, + bundle_info=BundleInfo(name="poc", version="v1"), + definitions=definitions, + start_deadline=now + timedelta(seconds=stop_seconds / 2), + stop_deadline=now + timedelta(seconds=stop_seconds), + queue=queue, + token=generator.generate( + extras={ + "sub": str(workload_id), + "scope": TOKEN_SCOPE, + "attempt_ids": [str(definition.attempt_id) for definition in definitions], + } + ), + ) + + +def write_fixtures(output: Path, *, recoverable: bool = False) -> list[Path]: + source = output / "submitter-bundle" + source.mkdir() + paths = [] + for name in ("accepted", "interrupted", "unrelated"): + path = source / f"{name}.py" + body = ( + "from pathlib import Path\nimport json\nimport os\nimport socket\nimport time\nfrom airflow.sdk import DAG\n" + f"with Path('/worker-evidence/{name}.imports').open('a') as marker:\n" + " marker.write(json.dumps({\n" + " 'run_id': os.environ['AIRFLOW_DAG_PARSING_POC_RUN_ID'],\n" + " 'worker_id': os.environ['AIRFLOW_DAG_PARSING_POC_WORKER_ID'],\n" + " 'container_hostname': socket.gethostname(),\n" + " 'task_id': os.environ['AIRFLOW_DAG_PARSING_POC_TASK_ID'],\n" + " 'state': 'IMPORT_STARTED',\n" + " }) + '\\n')\n" + ) + if name == "interrupted": + if recoverable: + body += "if len(Path('/worker-evidence/interrupted.imports').read_text().splitlines()) == 1:\n time.sleep(120)\n" + else: + body += "time.sleep(120)\n" + body += "Path('/worker-evidence/interrupted.finished').touch()\n" + body += f"with DAG('recovery_{name}', schedule=None):\n pass\n" + path.write_text(body) + paths.append(path) + shutil.copytree(source, output / "worker-bundle") + (output / "worker-evidence").mkdir() + return paths + + +def snapshot_executor(executor, workload) -> dict: + executor.sync() + events = { + str(key): {"state": state.value, "info": str(info) if info else None} + for key, (state, info) in executor.get_event_buffer().items() + } + return { + "backend_state": executor.celery_app.backend.get_task_meta(str(workload.workload_id), cache=False)[ + "status" + ], + "available_slots": executor.slots_available, + "running": len(executor.running), + "tracked_workloads": len(executor.workloads), + "queued": sum(len(queue) for queue in executor.executor_queues.values()), + "events": events, + } + + +def check_late_publication(url: str, workload, store: ReceiptStore, original: list[dict]) -> dict: + accepted_result = store.get_results(workload.workload_id)[0] + definition = workload.definitions[1] + late_result = DagDefinitionResult( + attempt_id=definition.attempt_id, + relative_path=definition.relative_path, + source_revision=definition.source_revision, + outcome="worker_error", + duration_seconds=0, + diagnostics=["Synthetic late publication probe after observed worker termination"], + ).model_dump(mode="json") + bodies = [ + (0, accepted_result, original[0]["execution_id"], 200), + ( + 0, + {**accepted_result, "diagnostics": ["conflicting synthetic retry"]}, + original[0]["execution_id"], + 409, + ), + (1, late_result, original[1]["execution_id"], 410), + (1, late_result, str(uuid4()), 409), + ] + statuses = [] + with httpx.Client( + timeout=5, trust_env=False, headers={"Authorization": f"Bearer {workload.token}"} + ) as client: + for index, result, execution_id, expected in bodies: + attempt = workload.definitions[index] + response = client.post( + f"{url}/execution/poc/parsing/workloads/{workload.workload_id}/attempts/{attempt.attempt_id}/result", + json={"execution_id": execution_id, "result": result}, + ) + if response.status_code != expected: + raise RuntimeError(f"Late publication returned {response.status_code}, expected {expected}") + if expected == 200 and response.json()["digest"] != original[0]["digest"]: + raise RuntimeError("Accepted receipt digest changed during replay") + statuses.append(response.status_code) + if store.get_attempts(workload.workload_id) != original: + raise RuntimeError("Late publication changed durable receipts") + return { + "accepted_identical_replay": statuses[0], + "accepted_conflicting_replay": statuses[1], + "unaccepted_original_execution": statuses[2], + "unaccepted_wrong_execution": statuses[3], + "source": "synthetic trusted-driver HTTP requests with a still-valid workload token", + } + + +def run_experiment(args, output: Path, store: ReceiptStore, generator: JWTGenerator, url: str) -> dict: + os.environ["AIRFLOW__CELERY__BROKER_URL"] = args.broker_url + os.environ["AIRFLOW__CELERY__RESULT_BACKEND"] = args.result_backend + os.environ["AIRFLOW__CELERY__SYNC_PARALLELISM"] = "1" + os.environ["AIRFLOW__CELERY_BROKER_TRANSPORT_OPTIONS__VISIBILITY_TIMEOUT"] = "3600" + files = write_fixtures(output, recoverable=args.recover) + run_id = str(uuid4()) + write_json(output / "driver-ready.json", {"port": args.port, "queue": args.queue, "run_id": run_id}) + wait_until( + lambda: read_json(output / "worker-evidence" / "isolation.json").get("stage") == "worker_ready", + description="original isolated worker", + timeout=180, + ) + original_worker = read_json(output / "worker-evidence" / "isolation.json") + validate_worker_ready(original_worker, run_id=run_id, queue=args.queue) + workload = create_workload(files[:2], generator, queue=args.queue, stop_seconds=args.deadline_seconds) + if args.recover: + from airflow.dag_processing.executor_recovery import PublicationOutcome + + from dev.dag_parsing_poc.recovery_checkpoint import create_coordinator + + coordinator = create_coordinator(store, args.queue, generator) + coordinator.start() + coordinator.admit(workload) + executor = coordinator.executor + else: + store.register_workload(workload) + executor = create_executor() + try: + if args.recover: + if coordinator.dispatch_reserved() != [ + PublicationOutcome(str(workload.workload_id), "published") + ]: + raise RuntimeError("Initial workload publication was not acknowledged") + else: + executor.queue_workload(workload, session=None) + executor.heartbeat() + wait_until( + lambda: ( + [row["status"] for row in store.get_attempts(workload.workload_id)] == ["accepted", "claimed"] + and len(read_import_events(output / "worker-evidence" / "interrupted.imports")) == 1 + ), + description="accepted first definition and active second import", + timeout=30, + ) + original = store.get_attempts(workload.workload_id) + task_id = str(workload.workload_id) + if not any( + matches_worker_event(event, original_worker, task_id=task_id, state="STARTED") + for event in read_task_events(output, task_id) + ): + raise RuntimeError("Original delivery evidence does not identify the ready worker") + for name in ("accepted", "interrupted"): + imports = read_import_events(output / "worker-evidence" / f"{name}.imports") + if len(imports) != 1 or not matches_worker_event( + imports[0], original_worker, task_id=task_id, state="IMPORT_STARTED" + ): + raise RuntimeError("Import evidence does not identify the original delivery and worker") + first = store.get_results(workload.workload_id)[0] + if first["outcome"] != "success" or [ + item.get("dag", {}).get("dag_id") for item in first["serialized_dags"] + ] != ["recovery_accepted"]: + raise RuntimeError("First definition did not publish the expected successful Dag") + write_json( + output / "kill-request.json", + { + "action": "kill the whole original worker container, save Docker termination evidence, start a replacement", + "workload_id": str(workload.workload_id), + "run_id": run_id, + "original_worker": original_worker, + "receipts": original, + "before_kill": snapshot_executor(executor, workload), + "stop_deadline": workload.stop_deadline.isoformat(), + }, + ) + wait_until( + lambda: (output / "replacement-ready.json").exists(), + description="operator-confirmed worker termination and replacement readiness", + timeout=180, + ) + control = read_json(output / "replacement-ready.json") + replacement_worker = read_json(output / "worker-evidence" / "isolation.json") + validate_replacement(control, original_worker, replacement_worker, run_id=run_id, queue=args.queue) + after_kill = snapshot_executor(executor, workload) + if (after_kill["running"], after_kill["tracked_workloads"]) != (1, 1): + raise RuntimeError( + "Worker loss produced terminal provider evidence; missing-record scenario needs a live tracked submission" + ) + if store.get_attempts(workload.workload_id) != original: + raise RuntimeError("Partial receipts changed after worker loss") + + backend = executor.celery_app.backend + deleted = delete_running_backend_record(backend, task_id) + after_delete = snapshot_executor(executor, workload) + if after_delete["backend_state"] != "PENDING" or after_delete["available_slots"] != 0: + raise RuntimeError("Missing backend record did not retain the tracked executor slot") + + if args.recover: + from dev.dag_parsing_poc.recovery_checkpoint import run_checkpoint + + return run_checkpoint( + args, + output, + store, + generator, + workload, + original, + original_worker, + replacement_worker, + control, + after_delete, + deleted, + url, + ) + + unrelated = create_workload([files[2]], generator, queue=args.queue, stop_seconds=300) + store.register_workload(unrelated) + executor.queue_workload(unrelated, session=None) + executor.heartbeat() + if len(executor.executor_queues[WorkloadType.PARSE_DAG_DEFINITIONS]) != 1: + raise RuntimeError("Unrelated work was unexpectedly admitted through an occupied executor") + task = executor.celery_app.tasks["execute_workload"] + task.apply_async( + args=[workload.model_dump_json()], + queue=args.queue, + task_id=task_id, + argsrepr="(<redacted parsing workload>,)", + ) + wait_until( + lambda: any( + matches_worker_event(event, replacement_worker, task_id=task_id, state="IGNORED") + for event in read_task_events(output, task_id) + ), + description="replacement ignoring the unfinished original claim", + timeout=30, + ) + after_replay = snapshot_executor(executor, workload) + if store.get_attempts(workload.workload_id) != original: + raise RuntimeError("Replacement took over an unfinished claim or changed an accepted receipt") + if (after_replay["running"], after_replay["tracked_workloads"], after_replay["queued"]) != (1, 1, 1): + raise RuntimeError("Ignored redelivery changed tracked or queued work") + + fresh_path = output / "fresh-executor.json" + observer = multiprocessing.get_context("spawn").Process( + target=inspect_fresh_executor, + args=(store.path, task_id, str(fresh_path)), + ) + observer.start() + observer.join(timeout=30) + if observer.is_alive(): + observer.kill() + observer.join(timeout=5) + raise TimeoutError("Fresh executor process did not exit") + if observer.exitcode != 0: + raise RuntimeError("Fresh executor process failed") + fresh = read_json(fresh_path) + if fresh["pid"] != observer.pid or fresh["pid"] == os.getpid(): + raise RuntimeError("Executor observation did not come from a fresh spawned process") + if (fresh["running"], fresh["tracked_workloads"], fresh["available_slots"]) != (0, 0, 1) or fresh[ + "receipts" + ] != original: + raise RuntimeError( + "Fresh executor or reopened receipt state differed from the expected prototype gap" + ) + + wait_until( + lambda: timezone.utcnow() > workload.stop_deadline, + description="original stop deadline", + timeout=args.deadline_seconds + 5, + ) + late = check_late_publication(url, workload, store, original) + executor.heartbeat() + final = snapshot_executor(executor, workload) + if (final["running"], final["tracked_workloads"], final["queued"]) != (1, 1, 1): + raise RuntimeError("Expiry changed the unresolved slot or admitted unrelated work") + counts = { + path.stem: len(path.read_text().splitlines()) + for path in (output / "worker-evidence").glob("*.imports") + } + if ( + counts != {"accepted": 1, "interrupted": 1} + or (output / "worker-evidence" / "interrupted.finished").exists() + ): + raise RuntimeError(f"Unexpected reimport or completion after worker loss: {counts}") + write_json(output / "accepted-results.json", {"results": store.get_results(workload.workload_id)}) + return { + "mode": "whole-worker loss, real Redis and HTTP, explicit duplicate injection", + "workload_id": task_id, + "run_id": run_id, + "receipts": original, + "after_kill": after_kill, + "deleted_backend_records": deleted, + "after_backend_deletion": after_delete, + "after_replay": after_replay, + "fresh_executor": fresh, + "after_deadline": final, + "late_publication": late, + "import_counts": counts, + "termination": control, + "original_worker_isolation": original_worker, + "replacement_worker_isolation": replacement_worker, + "limitations": [ + "A fresh executor process was inspected; no scheduler, ownership lease or automatic adoption exists.", + "Explicit replay does not exercise Redis visibility-timeout redelivery or broker restart.", + "Late publication requests were synthesized by the trusted driver, not sent by the killed worker.", + "Docker termination evidence applies only to this killed container, not general Celery cancellation.", + "This measures missing recovery; it does not add claim retirement, replacement attempts or a durable capacity ledger.", + ], + } + finally: + executor.end() + + +def main() -> None: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--output", type=Path, required=True) + parser.add_argument("--broker-url", required=True) + parser.add_argument("--result-backend", required=True) + parser.add_argument("--queue", default="poc-recovery") + parser.add_argument("--port", type=int, default=8799) + parser.add_argument("--deadline-seconds", type=float, default=90) + parser.add_argument("--recover", action="store_true", help="Exercise durable admission and replacement") + args = parser.parse_args() + if not 30 <= args.deadline_seconds <= 300: + parser.error("deadline-seconds must be between 30 and 300") + output = args.output.resolve() + output.mkdir(parents=True, exist_ok=False) + key = Ed25519PrivateKey.generate() + public_key = output / "verification-key.pem" + public_key.write_bytes( + key.public_key().public_bytes( + serialization.Encoding.PEM, + serialization.PublicFormat.SubjectPublicKeyInfo, + ) + ) + generator = JWTGenerator( + private_key=key, + kid=TOKEN_KEY_ID, + issuer=TOKEN_ISSUER, + audience=TOKEN_AUDIENCE, + algorithm="EdDSA", + valid_for=900, + ) + store = ReceiptStore(output / "receipts.sqlite") + listener = socket.socket() + listener.bind(("0.0.0.0", args.port)) Review Comment: ## CodeQL / Binding a socket to all network interfaces Binding a socket to all interfaces (using ['0.0.0.0'](1)) is a security risk. [Show more details](https://github.com/apache/airflow/security/code-scanning/660) ########## dev/dag_parsing_poc/run_scheduler_parsing.py: ########## @@ -0,0 +1,586 @@ +# 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. +# /// script +# requires-python = ">=3.10" +# dependencies = ["apache-airflow-core", "apache-airflow-providers-celery", "cryptography"] +# /// +"""Run the scheduler-hosting experiment in Breeze with an isolated Celery worker and scheduler.""" + +from __future__ import annotations + +import argparse +import configparser +import json +import multiprocessing +import os +import socket +import sqlite3 +import statistics +import subprocess +import threading +import time +from contextlib import ExitStack +from datetime import datetime, timezone +from pathlib import Path +from uuid import uuid4 + +from cryptography.hazmat.primitives import serialization +from cryptography.hazmat.primitives.asymmetric.ed25519 import Ed25519PrivateKey + +REDIS_IMAGE = "public.ecr.aws/docker/library/redis:7.4-alpine" +ROUTE = "scheduler-parsing" + + +def write_json(path: Path, value) -> None: + path.write_text(json.dumps(value, indent=2) + "\n") + + +def docker(*arguments: str, check: bool = True) -> str: + result = subprocess.run(["docker", *arguments], check=False, capture_output=True, text=True, timeout=45) + if check and result.returncode: + raise RuntimeError(f"Docker {arguments[0]} failed: {result.stderr.strip()}") + return result.stdout.strip() + + +def wait_until(predicate, *, timeout: float = 120, check=None): + deadline = time.monotonic() + timeout + while time.monotonic() < deadline: + if value := predicate(): + return value + if check is not None: + check() + time.sleep(0.1) + raise TimeoutError("Scheduler-hosting checkpoint did not finish") + + +def stop_process(process) -> None: + if process.pid is None: + return + process.terminate() + process.join(5) + if process.is_alive(): + process.kill() + process.join(5) + + +def run_api(database: str, public: Path, listener: socket.socket) -> None: + # These fresh processes import Airflow only after the experiment environment is configured. + from airflow.dag_processing.executor_manager import _serve_api + + _serve_api(database, public, listener) + + +def run_provider(database: str, key: bytes, stop, capacity: int) -> None: + from airflow.api_fastapi.auth.tokens import JWTGenerator + from airflow.api_fastapi.execution_api.parsing import ( + TOKEN_AUDIENCE, + TOKEN_ISSUER, + TOKEN_KEY_ID, + TOKEN_SCOPE, + ) + from airflow.dag_processing.executor_runner import ParsingExecutorRunner + from airflow.dag_processing.parsing_metadata import MetadataOrchestrationStore + from airflow.executors.workloads import WorkloadType + from airflow.providers.celery.executors.celery_executor import CeleryExecutor + + generator = JWTGenerator( + private_key=Ed25519PrivateKey.from_private_bytes(key), + kid=TOKEN_KEY_ID, + issuer=TOKEN_ISSUER, + audience=TOKEN_AUDIENCE, + algorithm="EdDSA", + valid_for=600, + ) + + def issue_token(manifest): + return generator.generate( + { + "sub": manifest["workload_id"], + "scope": TOKEN_SCOPE, + "attempt_ids": [item["attempt_id"] for item in manifest["definitions"]], + } + ) + + executor = CeleryExecutor(parallelism=capacity) + executor.supported_workload_types = frozenset({WorkloadType.PARSE_DAG_DEFINITIONS}) + store = MetadataOrchestrationStore(database) + runner = ParsingExecutorRunner(store, executor, route=ROUTE, token_issuer=issue_token) + runner.start() + try: + while not stop.poll(0.05): + runner.tick() + finally: + runner.close() + store.engine.dispose() + + +class MetricsCapture: + """Capture the scheduler's StatsD measurements independently of its process.""" + + def __init__(self): + self.socket = socket.socket(socket.AF_INET, socket.SOCK_DGRAM) + self.socket.bind(("0.0.0.0", 0)) + self.socket.settimeout(0.1) + self.port = self.socket.getsockname()[1] + self.samples: list[dict] = [] + self.stop = threading.Event() + self.thread = threading.Thread(target=self.receive, daemon=True) + self.thread.start() + + def receive(self): + while not self.stop.is_set(): + try: + packet, _ = self.socket.recvfrom(65536) + except TimeoutError: + continue + for metric in packet.decode().splitlines(): + name, value = metric.split(":", 1) + number, kind, *_ = value.split("|") + self.samples.append( + {"at": time.monotonic(), "name": name, "value": float(number), "kind": kind} + ) + + def close(self): + self.stop.set() + self.thread.join(2) + self.socket.close() + + +def summarize_metrics(samples: list[dict]) -> dict: + result = {} + for name in sorted({item["name"] for item in samples}): + if ( + "parsing_step" not in name + and "scheduler_loop_duration" not in name + and "scheduler_heartbeat" not in name + ): + continue + matching = [item for item in samples if item["name"] == name] + values = sorted(item["value"] for item in matching) + result[name] = ( + { + "samples": len(values), + "median_ms": statistics.median(values), + "p95_ms": values[int((len(values) - 1) * 0.95)], + "max_ms": max(values), + } + if matching[0]["kind"] == "ms" + else {"count": sum(values)} + ) + return result + + +def get_state(database: Path) -> dict: + with sqlite3.connect(f"file:{database}?mode=ro", uri=True, timeout=1) as connection: + return { + "heartbeat": connection.execute( + "SELECT max(latest_heartbeat) FROM job WHERE job_type='SchedulerJob'" + ).fetchone()[0], + "completed_tasks": connection.execute( + "SELECT count(*) FROM task_instance WHERE state='success'" + ).fetchone()[0], + "accepted": dict(connection.execute("SELECT path, accepted_count FROM parse_sources")), + } + + +def record_phase(name: str, database: Path, metrics: MetricsCapture, seconds: float, action=None) -> dict: + before = get_state(database) + started = time.monotonic() + if action is None: + time.sleep(seconds) + else: + action(seconds) + after = get_state(database) + phase = { + "name": name, + "before": before, + "after": after, + "metrics": summarize_metrics([item for item in metrics.samples if item["at"] >= started]), + } + if not any(name.endswith("parsing_step_duration") for name in phase["metrics"]): + raise RuntimeError(f"Missing parsing callback measurements during {name}") + if after["heartbeat"] == before["heartbeat"] or after["completed_tasks"] <= before["completed_tasks"]: + raise RuntimeError(f"Scheduler stopped making progress during {name}: {phase}") + print(json.dumps({"event": "phase", **phase}), flush=True) + return phase + + +def contend(database: Path, seconds: float) -> None: + deadline = time.monotonic() + seconds + while time.monotonic() < deadline: + with sqlite3.connect(database, timeout=0) as connection: + try: + connection.execute("BEGIN IMMEDIATE") + except sqlite3.OperationalError: + pass + else: + time.sleep(0.04) + time.sleep(0.08) + + +def run_experiment(output: Path, phase_seconds: float) -> dict: + own = json.loads(docker("inspect", os.environ["HOSTNAME"]))[0] + mounts = {mount["Destination"]: Path(mount["Source"]) for mount in own["Mounts"]} + host_root = mounts["/opt/airflow/airflow-core"].parent + host_output = mounts["/files"] / output.relative_to("/files") + control, source, evidence = (output / name for name in ("control", "worker-bundle", "worker-evidence")) + for path in (control, source, evidence): + path.mkdir(mode=0o777) + path.chmod(0o777) + database = control / "metadata.sqlite" + os.environ.update( + AIRFLOW_HOME=str(control / "api-home"), + AIRFLOW_CONFIG=str(control / "api.cfg"), + AIRFLOW__DATABASE__SQL_ALCHEMY_CONN=f"sqlite:///{database}", + AIRFLOW__CORE__LOAD_EXAMPLES="False", + AIRFLOW__CORE__EXECUTOR="LocalExecutor", + AIRFLOW__CORE__MIN_SERIALIZED_DAG_UPDATE_INTERVAL="0", + AIRFLOW__CORE__DAGS_ARE_PAUSED_AT_CREATION="False", + AIRFLOW__DAG_PROCESSOR__DAG_BUNDLE_CONFIG_LIST="[]", + AIRFLOW__CELERY__BROKER_URL="redis://poc-broker:6379/0", + AIRFLOW__CELERY__RESULT_BACKEND="redis://poc-broker:6379/1", + AIRFLOW__CELERY__SYNC_PARALLELISM="1", + ) + with (output / "migration.log").open("w") as log: + subprocess.run( + ["airflow", "db", "migrate"], check=True, stdout=log, stderr=subprocess.STDOUT, timeout=120 + ) + database.chmod(0o666) + from sqlalchemy.orm import Session + + from airflow.dag_processing.discovery import discover_python_bundle + from airflow.dag_processing.orchestrator import ParseOrchestrator + from airflow.dag_processing.parsing_metadata import MetadataOrchestrationStore + from airflow.executors.workloads import BundleInfo + from airflow.models.dagbundle import DagBundleModel + + (source / "fast.py").write_text( + "from datetime import datetime, timedelta, timezone\nfrom airflow.sdk import DAG\n" + "from airflow.providers.standard.operators.empty import EmptyOperator\n" + "with DAG('scheduler_hosted',schedule=timedelta(seconds=2),catchup=False," + "start_date=datetime(2025,1,1,tzinfo=timezone.utc),is_paused_upon_creation=False):\n" + " EmptyOperator(task_id='tick')\n" + ) + (source / "slow.py").write_text( + "import time\nfrom airflow.sdk import DAG\ntime.sleep(3)\n" + "dag = DAG('scheduler_slow', schedule=None)\n" + ) + config = {"route": ROUTE, "bundle": "poc", "capacity": 2, "batch_size": 1, "parse_interval": 2} + write_json(control / "parsing.json", config) + store = MetadataOrchestrationStore(database) + with Session(store.engine) as session: + session.add(DagBundleModel(name="poc", version="v1")) + session.commit() + ParseOrchestrator(store, **config).update_inventory( + BundleInfo(name="poc", version="v1"), discover_python_bundle(source, bundle_name="poc") + ) + context = multiprocessing.get_context("spawn") + key = Ed25519PrivateKey.generate() + public = control / "public.pem" + public.write_bytes( + key.public_key().public_bytes( + serialization.Encoding.PEM, serialization.PublicFormat.SubjectPublicKeyInfo + ) + ) + listener = socket.socket() + listener.bind(("0.0.0.0", 0)) Review Comment: ## CodeQL / Binding a socket to all network interfaces Binding a socket to all interfaces (using ['0.0.0.0'](1)) is a security risk. [Show more details](https://github.com/apache/airflow/security/code-scanning/663) ########## dev/dag_parsing_poc/run_metadata.py: ########## @@ -0,0 +1,285 @@ +# 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. +# /// script +# requires-python = ">=3.10" +# /// +"""Prove isolated Celery parsing can populate metadata and feed the real scheduler.""" + +from __future__ import annotations + +import argparse +import hashlib +import json +import multiprocessing +import os +import socket +from pathlib import Path +from uuid import uuid4 + + +def serve_metadata_api(store_path: str, public_key_path: str, listener: socket.socket) -> None: + import uvicorn + + from airflow.api_fastapi.execution_api.parsing import create_app + + server = uvicorn.Server( + uvicorn.Config(create_app(store_path, public_key_path, persist_metadata=True), log_level="warning") + ) + server.run(sockets=[listener]) + + +def run_scheduler_checkpoint() -> dict: + from sqlalchemy import select + + from airflow.executors.local_executor import LocalExecutor + from airflow.jobs.job import Job + from airflow.jobs.scheduler_job_runner import SchedulerJobRunner + from airflow.models.dag import DagModel + from airflow.models.dagrun import DagRun + from airflow.models.taskinstance import TaskInstance + from airflow.utils.session import create_session + from airflow.utils.state import DagRunState, TaskInstanceState + + runner = SchedulerJobRunner(job=Job(), executors=[LocalExecutor(parallelism=1)]) + with create_session() as session: + session.add(runner.job) + session.flush() + models = list(session.scalars(select(DagModel).where(DagModel.is_paused == False))) # noqa: E712 + if len(models) != 1 or models[0].next_dagrun is None: + raise RuntimeError("Remote result did not create one schedulable Dag") + runner._create_dag_runs(models, session=session) + session.flush() + runner._start_queued_dagruns(session=session) + session.flush() + dag_run = session.scalar(select(DagRun)) + if dag_run is None or dag_run.state != DagRunState.RUNNING: + raise RuntimeError("Scheduler did not start the remotely parsed Dag") + runner._schedule_dag_run(dag_run, session=session) + session.flush() + task_instance = session.scalar(select(TaskInstance)) + if task_instance is None or task_instance.state != TaskInstanceState.SCHEDULED: + raise RuntimeError("Scheduler did not schedule the remotely parsed task") + return { + "dag_id": dag_run.dag_id, + "run_id": dag_run.run_id, + "run_state": dag_run.state, + "task_id": task_instance.task_id, + "task_state": task_instance.state, + "dag_version_id": str(task_instance.dag_version_id), + } + + +def run_checkpoint(args, output: Path) -> dict: + # Airflow and Celery cache configuration on import. This script starts a fresh interpreter. + os.environ.update( + { + "AIRFLOW__DATABASE__SQL_ALCHEMY_CONN": f"sqlite:///{output / 'metadata.sqlite'}", + "AIRFLOW__CORE__EXECUTOR": "LocalExecutor", + "AIRFLOW__CORE__LOAD_EXAMPLES": "False", + "AIRFLOW__CORE__MIN_SERIALIZED_DAG_UPDATE_INTERVAL": "0", + "AIRFLOW__CELERY__BROKER_URL": args.broker_url, + "AIRFLOW__CELERY__RESULT_BACKEND": args.result_backend, + "AIRFLOW__CELERY__SYNC_PARALLELISM": "1", + } + ) + from cryptography.hazmat.primitives import serialization + from cryptography.hazmat.primitives.asymmetric.ed25519 import Ed25519PrivateKey + from sqlalchemy import select + + from airflow import settings + from airflow.api_fastapi.auth.tokens import JWTGenerator + from airflow.dag_processing.executor_recovery import ParsingRecoveryCoordinator + from airflow.dag_processing.parsing_metadata import MetadataReceiptStore + from airflow.executors.workloads import WorkloadType + from airflow.models import import_all_models + from airflow.models.base import Base + from airflow.models.dagbundle import DagBundleModel + from airflow.models.dagcode import DagCode + from airflow.models.serialized_dag import SerializedDagModel + from airflow.providers.celery.executors.celery_executor import CeleryExecutor + from airflow.utils.db import add_default_pool_if_not_exists, synchronize_log_template + from airflow.utils.session import create_session + + from dev.dag_parsing_poc.run import create_archive, create_workloads, wait_for_api + from dev.dag_parsing_poc.run_celery import finish_workloads, read_json, read_task_events, wait_until + from dev.dag_parsing_poc.run_celery_recovery import validate_worker_ready + + import_all_models() + Base.metadata.create_all(settings.engine) + with create_session() as session: + add_default_pool_if_not_exists(session=session) + synchronize_log_template(session=session) + session.add(DagBundleModel(name="poc", version="v1")) + store = MetadataReceiptStore(output / "metadata.sqlite") + signing_key = Ed25519PrivateKey.generate() + public = output / "verification-key.pem" + public.write_bytes( + signing_key.public_key().public_bytes( + serialization.Encoding.PEM, serialization.PublicFormat.SubjectPublicKeyInfo + ) + ) + generator = JWTGenerator( + private_key=signing_key, + kid="dag-parsing-poc", + issuer="dag-parsing-poc", + audience="dag-parsing-poc", + algorithm="EdDSA", + valid_for=600, + ) + worker_root = output / "worker-bundle" + worker_root.mkdir() + source = worker_root / "scheduled.py" + source.write_text( + "from pathlib import Path\nfrom datetime import datetime, timezone\n" + "from airflow.sdk import DAG, task\n" + "with Path('/worker-evidence/imports.txt').open('a') as marker:\n" + " marker.write('imported\\n')\n" + "with DAG('remote_metadata_checkpoint', schedule='@once', " + "start_date=datetime(2026, 1, 1, tzinfo=timezone.utc), is_paused_upon_creation=False):\n" + " @task\n def sample():\n return 1\n sample()\n" + ) + archive = create_archive([source], worker_root / "definitions.zip") if args.archive_members else None + (output / "worker-evidence").mkdir() + run_id = str(uuid4()) + listener = socket.socket() + listener.bind(("0.0.0.0", args.port)) Review Comment: ## CodeQL / Binding a socket to all network interfaces Binding a socket to all interfaces (using ['0.0.0.0'](1)) is a security risk. [Show more details](https://github.com/apache/airflow/security/code-scanning/661) -- This is an automated message from the Apache Git Service. To respond to the message, please log on to GitHub and use the URL above to go to the specific comment. To unsubscribe, e-mail: [email protected] For queries about this service, please contact Infrastructure at: [email protected]
