diff --git a/.agents/skills/airflow-java-sdk/SKILL.md b/.agents/skills/airflow-java-sdk/SKILL.md index 8edb742679535..9921832eafee5 100644 --- a/.agents/skills/airflow-java-sdk/SKILL.md +++ b/.agents/skills/airflow-java-sdk/SKILL.md @@ -50,8 +50,9 @@ subclasses must only import from `org.apache.airflow.sdk`; any import of ## Bundle composition and coordinator discovery -A **bundle** is a directory of JAR files (typically `build/bundle/`) placed on the coordinator's -`jars_root`. The coordinator scans the directory at task-dispatch time to find: +A **bundle** is a directory of JAR files (typically `build/bundle/`) placed in the Dag bundle named +by the coordinator's `task_handler_bundle_name` (the task's own Dag bundle when unset). The +coordinator scans that Dag bundle at task-dispatch time to find: 1. **`Main-Class`** (standard JAR manifest attribute) — the fully-qualified class name of the entry point that the coordinator invokes with `java -classpath … --comm … --logs …`. @@ -65,7 +66,7 @@ A **bundle** is a directory of JAR files (typically `build/bundle/`) placed on t `runtimeClasspath` and copies it into the shadow JAR manifest. In thin-JAR mode (`fatJar = false`), the value stays in the `airflow-sdk` JAR deployed alongside the bundle JAR. -The Python coordinator (`JavaCoordinator`) scans every JAR under `jars_root` with +The Python coordinator (`JavaCoordinator`) scans every JAR in that Dag bundle with `_JarInfo.find()`, reads `META-INF/MANIFEST.MF` out of each ZIP, and collects `Main-Class` and `Airflow-Supervisor-Schema-Version` from whichever JARs carry them. The resolved schema version is then passed as the `schema_version` return value from `_build_execute_task_command`, which @@ -74,7 +75,7 @@ the base `SubprocessCoordinator` uses to negotiate the supervisor wire protocol. If `main_class` is set explicitly on the `JavaCoordinator` instance (via `[sdk] coordinators` kwargs), the scan uses it as a filter; otherwise the first JAR with a `Main-Class` attribute wins. Either way, `Airflow-Supervisor-Schema-Version` must be present in at least one JAR in -`jars_root` or startup fails. +the Dag bundle or startup fails. Every JAR in the Dag bundle goes on one classpath. --- @@ -117,7 +118,7 @@ E2E_TEST_MODE=java_sdk uv run --project airflow-e2e-tests pytest \ `coordinator.py` extends `SubprocessCoordinator`. The only method subclasses must implement is `_build_execute_task_command`, which returns `(argv, schema_version)`. Look at the existing -implementation for how `jars_root`, `java_executable`, `jvm_args`, and `main_class` are +implementation for how the scanned Dag bundle, `java_executable`, `jvm_args`, and `main_class` are assembled into the command. Do not reach into the JVM process from Python beyond what this method provides. diff --git a/airflow-core/adr/lang-sdk/0011-mixed-language-dag-processing.md b/airflow-core/adr/lang-sdk/0011-mixed-language-dag-processing.md index c2acaadbedf2a..425ca3ce1f3ff 100644 --- a/airflow-core/adr/lang-sdk/0011-mixed-language-dag-processing.md +++ b/airflow-core/adr/lang-sdk/0011-mixed-language-dag-processing.md @@ -83,6 +83,10 @@ The split is per registration, not per file and not per bundle. Comparing a stub against its handler needs two things at once: the `airflow.sdk.DAG` objects the Python file produced, and the `arg_bindings` that only appear once those Dags are serialized. `_parse_file` holds both, between `_serialize_dags` and the `DagFileParsingResult` it returns. That is where the handler query is issued. +Steps 2 to 4 below are superseded by [ADR-0013](0013-persisted-task-handler-bindings.md) Flow 1. A stub task on a queue no coordinator serves is not checked, and the coordinator +lists every candidate in its Dag bundle instead of locating the artifact behind each `dag_id`. Each candidate without a recorded answer is probed until a coordinator that lists it +gets an answer (see ADR-0013 Flow 1), whatever Dag ids the stub tasks need. + ``` DagFileProcessorProcess(etl.py) ← manager spawns, as for any file └── _parse_file_entrypoint → _parse_file @@ -90,7 +94,7 @@ DagFileProcessorProcess(etl.py) ← manager spawn ├── _serialize_dags(bag) → is_stub tasks carry arg_bindings (ADR-0007) │ │ ┌─────────────────────────────────────────────────────────────────────┐ - │ │ Step 1: Collect the dag_ids to ask about — every Dag in this │ + │ │ Step 1: Collect the Dags to validate — every Dag in this │ │ │ file with at least one is_stub task → ["etl"] │ │ └─────────────────────────────────────────────────────────────────────┘ │ @@ -123,27 +127,26 @@ DagFileProcessorProcess(etl.py) ← manager spawn │ │ └── BundleScanner.scanBundles(roots) │ │ │ → "etl" → ResolvedBundle(analytics.jar, mainClass, ...) │ │ │ │ - │ │ Group the dag_ids by (coordinator, artifact) — one group, one │ + │ │ Group the stub tasks by (coordinator, artifact) — one group, one │ │ │ process, one request │ │ └─────────────────────────────────────────────────────────────────────┘ │ │ ┌─────────────────────────────────────────────────────────────────────┐ │ │ Step 4: Query each group — one request, one response │ │ │ │ - │ │ SDKTaskHandlerProcessorProcess.start( │ - │ │ target=_parse_task_handler_entrypoint, │ + │ │ LangSDKTaskHandlerProcessorProcess.start( │ + │ │ target=_start_task_handler_runtime_entrypoint, │ │ │ coordinator=JavaCoordinator("jdk-11"), │ │ │ path=analytics.jar) │ │ │ │ │ │ │ ├── in the child: _build_parse_task_handler_command() │ │ │ │ coordinator.parse_task_handler() — spawn JVM │ │ │ │ │ - │ │ │ ──TaskHandlerParseRequest(file=analytics.jar, │ - │ │ │ dag_ids=["etl"])─────▶ JVM │ + │ │ │ ──TaskHandlerParseRequest(file=analytics.jar)─────▶ JVM │ │ │ │ (ToSDKTaskHandlerProcessor) │ │ │ │ │ - │ │ │ JVM answers from its own TaskHandler registrations │ - │ │ │ whose dagId is one of the requested ids │ + │ │ │ JVM answers with every TaskHandler registration │ + │ │ │ in the artifact, or {} when there is none │ │ │ │ │ │ │ │ ◀─TaskHandlerParsingResult(task_handlers={ │ │ │ │ "etl": [extract, transform, load]})─────── JVM │ @@ -156,15 +159,21 @@ DagFileProcessorProcess(etl.py) ← manager spawn │ └─────────────────────────────────────────────────────────────────────┘ │ │ ┌─────────────────────────────────────────────────────────────────────┐ - │ │ Step 5: Compare per dag_id against the Dag just serialized — │ - │ │ union task_handlers[dag_id] across every process first │ + │ │ Step 5: Check each stub task against the answers of the │ + │ │ coordinator its queue routes to │ │ │ │ │ │ Python Dag "etl" (stub tasks) TaskHandlerDeclaration │ │ │ ────────────────────────────── ────────────────────────────── │ - │ │ task_id ↔ task_id (sets must match) │ - │ │ arg_bindings[*].name ↔ params[*].name (in order) │ + │ │ (dag_id, task_id) ↔ exactly one declaration │ + │ │ none, or two artifacts claiming it → an error │ + │ │ a handler with no stub task → not an error │ + │ │ arg_bindings[*] ↔ params[*] (per binding) │ + │ │ by position, or by folded or exact name │ + │ │ defaulted arguments dropped if that makes the count match │ + │ │ unmatched by name → a warning, not an error │ │ │ arg_bindings[*].value_schema ↔ params[*].value_schema │ - │ │ compared only where neither side is null │ + │ │ top-level JSON types, where both sides have a schema │ + │ │ a mapped stub task, or params=None → the handler's presence only │ │ │ │ │ │ On mismatch → import_errors["etl.py"] │ │ └─────────────────────────────────────────────────────────────────────┘ @@ -175,8 +184,8 @@ DagFileProcessorProcess(etl.py) ← manager spawn ``` The parse owns validation, not an importer. `PythonDagImporter` returns `airflow.sdk.DAG` objects and knows nothing about coordinators or queues, so `@task.stub` keeps working for -any importer that can produce a Dag carrying stub tasks. `_parse_file` is also the only place where the whole file's Dags are visible at once, which is what lets one request cover -every `dag_id` that resolved to the same artifact ([ADR-0012](0012-lang-sdk-parse-protocol.md)). +any importer that can produce a Dag carrying stub tasks. `_parse_file` is also the only place where the whole file's Dags are visible at once, which is what lets one request per +artifact serve every Dag in the file whose stubs resolved to it ([ADR-0012](0012-lang-sdk-parse-protocol.md)). Resolution goes through the coordinator registry, not the filesystem, so the Python Dag and the Lang-SDK artifact **do not need to be in the same DagBundle**. Nothing here needs an `airflow.sdk.DAG` round-trip either — validation compares against the Dag the Python parser already built. Appendix B states exactly what is compared. @@ -186,7 +195,7 @@ Resolution goes through the coordinator registry, not the filesystem, so the Pyt | Caller | Coordinator call | What comes back | Action | |--------------------------------------------|------------------------------------------------------|------------------------------------------|--------------------------------------------------| | `_parse_file` → `PythonDagImporter` | — (the Python file is parsed in process) | its own parsed Dags | PERSIST | -| `_parse_file`, per (coordinator, artifact) | `parse_task_handler`, scoped to that group's dag_ids | `TaskHandlerParsingResult` | VALIDATE only — not a Dag, so nothing to persist | +| `_parse_file`, per (coordinator, artifact) | `parse_task_handler`, for every handler it registers | `TaskHandlerParsingResult` | VALIDATE only — not a Dag, so nothing to persist | | `_parse_file` → `JavaDagImporter` | `parse_dag` | `DagFileParsingResult`, native Dags only | PERSIST | There is no fourth row. A `TaskHandlerRef` has no Dag, so no `DagImporter` — and nothing reading a `DagImporter`'s results — ever sees one. @@ -197,13 +206,15 @@ There is no fourth row. A `TaskHandlerRef` has no Dag, so no `DagImporter` — a - No importer knows about coordinators. `PythonDagImporter` is unchanged by this ADR; the stub-to-handler comparison sits in `_parse_file`, above every importer. - A mixed-language `dag_id` never appears in Dag processing results. No `Dag` registration exists for a `dag_id` a Python file already owns, so everything downstream sees exactly one record per `dag_id`, with no flag to interpret. -- Stub/implementation mismatches — missing handler, extra handler, parameter name or order, incompatible schema — surface as import errors against the Python file at parse time, - alongside the errors the parse already reports. An unannotated stub argument is checked by name and position only. +- Stub/implementation mismatches, such as a missing handler, two handlers for one stub task, a `positional` argument count that does not match or an incompatible schema, surface as import errors against + the Python file at parse time, alongside the errors the parse already reports. A `named` argument or parameter that matches nothing is only logged as a warning, because the + runtime runs the task anyway. An unannotated stub argument is checked only for how it binds, by position or by name. - The Python Dag and Lang-SDK artifact can live in different DagBundles. -- A single Dag can have stubs targeting different queues, some Java, some Go. Each resolves to its own coordinator instance, and validation unions their declarations per `dag_id` - before comparing task ids. -- Validating a file costs one extra process per (coordinator, artifact) pair its stubs resolve to — one for the common case of a file whose stubs all target a single runtime, and - none at all for a file with no stub tasks. +- A single Dag can have stubs targeting different queues, some Java, some Go. Each resolves to its own coordinator instance, and each stub task is checked only against the handlers + of its own coordinator's artifacts. +- Validating a file costs, on a cold start, one extra process per candidate artifact without a recorded answer in the coordinators' Dag bundles, plus one per retry of a failed + probe under the next coordinator that lists it, and none in the steady state or for a file with no stub tasks. This supersedes the cost of one process per (coordinator, artifact) + pair the stubs resolve to; see [ADR-0013](0013-persisted-task-handler-bindings.md) Flow 1. - Mixed-language is Python-primary only. Lang-SDK runtimes cannot define stub operators; a native Dag cannot delegate tasks to Python. - No per-Dag flag, no schema migration, no new `DagModel` column, no REST/UI change. - Terms track Language SDK spec `1.0`. A spec rename of `TaskHandler`, or of the `register` / `serve` verbs, lands here too. @@ -234,12 +245,12 @@ with one `bundle.serve()`. One artifact can hold both, so neither the file nor t ### Appendix B — What is compared -Handler declarations from every process spawned in Step 4 are unioned per `dag_id` before comparison, since one Dag's stubs can target several queues. +Each stub task is compared only against the answers of the coordinator its queue routes to, since one Dag's stubs can target several queues. Each answer covers every Dag its artifact registers handlers for; only the parsed file's stub tasks are looked up in it. -- `task_id` sets must match exactly. A missing or extra handler is an error. -- `arg_bindings[*].name` against `params[*].name`, in order — both sides bind positionally. -- `arg_bindings[*].value_schema` against `params[*].value_schema`, compared only where neither side is null. An unannotated `@task.stub` parameter produces `null` today, so a - strict comparison would make every untyped stub argument a parse error. +- Every stub task has exactly one handler among those answers, the one declaration of its `(dag_id, task_id)`. None, or two artifacts that both register it, is an error. A handler with no stub task is not an error: one artifact serves many Dag files, and a handler can outlive its stub task. +- `arg_bindings` against `params` as each declaration's `binding` says: by position for `positional`, by name for `named`. A `named` parameter takes the argument of its exact name, else, unless `exact_name` is set, the one whose name matches case-insensitively with underscores ignored; a folded name two arguments share matches neither. An argless call passes no arguments. An argument filled from the stub signature's default is type-checked when it binds and is never reported as unmatched. + A `positional` count matches when all arguments, or those left after dropping the defaulted ones, number the parameters; any other count is an error. Under `named`, an argument or parameter that matches nothing is only logged as a warning, and nothing is logged when the argument may be the whole value: no parameter matched, exactly one argument was passed, the handler declares parameters, none of them matches only by its exact name, and the argument may be an object ([ADR-0012](0012-lang-sdk-parse-protocol.md) Appendix B). A mapped stub task, whose arguments are not captured at parse time, and a declaration whose `params` is `None` are checked only for the handler's presence. +- `arg_bindings[*].value_schema` against `params[*].value_schema`, compared only where both sides have a schema. An unannotated `@task.stub` parameter produces `null` today, so a strict comparison would make every untyped stub argument a parse error. Only the top-level JSON types are compared, taken from `type`, the union over `anyOf` or `oneOf`, or the values of `const` or `enum`; a schema of any other shape, such as `$ref` or `allOf`, is not compared. As in the Go runtime, an argument that may be null needs a parameter that accepts null, and at least one of the argument's other types must be one the parameter accepts, with `integer` accepted by `number`. So an argument that may be an integer or a string binds to an integer parameter, and one that may be an integer or null does not. `format`, ranges and nested items are not compared. -Any mismatch is reported against the Python file, which is the definition the author can act on, and travels back on `DagFileParsingResult.import_errors` with everything else the +Every error is reported against the Python file, which is the definition the author can act on, and travels back on `DagFileParsingResult.import_errors` with everything else the parse found. diff --git a/airflow-core/adr/lang-sdk/0012-lang-sdk-parse-protocol.md b/airflow-core/adr/lang-sdk/0012-lang-sdk-parse-protocol.md index e144b87445b58..1a5f18a31a69f 100644 --- a/airflow-core/adr/lang-sdk/0012-lang-sdk-parse-protocol.md +++ b/airflow-core/adr/lang-sdk/0012-lang-sdk-parse-protocol.md @@ -26,8 +26,8 @@ Proposed ## Context The Dag processor asks a Lang-SDK runtime two different questions. "Which Dags does this artifact define?" is answered over the messages [ADR-0004](0004-dag-parsing.md) already -defines. "Which task handlers does this artifact register for a `dag_id` Python already owns?" has no answer in those messages, because a `TaskHandler` registration carries no Dag -([ADR-0011](0011-mixed-language-dag-processing.md)). +defines. "Which task handlers does this artifact register, each for a `dag_id` Python already owns?" has no answer in those messages, because a `TaskHandler` registration carries +no Dag ([ADR-0011](0011-mixed-language-dag-processing.md)). This ADR defines the request that carries the second question, the subprocess classes that carry both, and the two parse-side entry points on the coordinator. @@ -82,7 +82,6 @@ relay is only type-safe because both hops speak the same pair, which is the reas ``` class TaskHandlerParseRequest: file: str # the artifact resolved for this coordinator - dag_ids: list[str] # every Dag in the parsed file with stub tasks that resolved here bundle_path: Path bundle_name: str type: Literal["TaskHandlerParseRequest"] @@ -96,20 +95,24 @@ class TaskHandlerParsingResult: class TaskHandlerDeclaration: task_id: str - params: list[TaskHandlerParam] # ordered — arg_bindings are positional + binding: Literal["positional", "named"] # how stub-task arguments bind to params + params: list[TaskHandlerParam] | None # ordered; the order matters only for "positional". None: the runtime cannot list them class TaskHandlerParam: - name: str + name: str | None # None: the runtime has no name for this positional parameter value_schema: ArgValueSchema | None = None - required: bool # the handler declares no default + exact_name: bool = False # match as spelled, not case-insensitively with underscores ignored ``` -One request carries every `dag_id` that resolved to the same artifact under the same coordinator, so a file whose stubs all target one runtime costs one process. A `dag_id` the -artifact registers nothing for is **omitted** from `task_handlers` rather than returned empty: the key set is not required to match `dag_ids`, because it is the union across -coordinators that has to cover the stubs ([ADR-0011](0011-mixed-language-dag-processing.md)). +The request names no Dags. The runtime answers with every task handler the artifact registers, keyed by `dag_id`, and with `{}` when it registers none. The answer must depend only +on the artifact, never on the request, so the Dag processor can cache it per artifact ([ADR-0013](0013-persisted-task-handler-bindings.md)). One request per candidate artifact +answers for every Dag it registers. The cost of one process per (coordinator, artifact) pair a file's stubs resolve to is superseded by +[ADR-0013](0013-persisted-task-handler-bindings.md) Flow 1: on a cold start the Dag processor probes each candidate without a recorded answer until a coordinator that lists it gets +an answer, and in the steady state it probes none. The key set is not required to match the file's Dags: it can hold other files' Dags, and each stub task needs its handler among +the answers of its own coordinator ([ADR-0011](0011-mixed-language-dag-processing.md)). `value_schema` reuses the `ArgValueSchema` definition `arg_bindings` already carries ([ADR-0007](0007-taskflow-across-language-boundary.md)), so both sides of a comparison are the -same type. Two properties matter to validation: the field is nullable on both sides, and `params` is ordered. Appendix B says what that forces. +same type. Two properties matter to validation: the field is nullable on both sides, and each declaration names its binding mode. Appendix B says what that forces. `task_handlers` is the counterpart to `DagFileParsingResult.serialized_dags`, but fully typed. `serialized_dags` is `list[LazyDeserializedDAG]`, which is an opaque object in the schema snapshot. A handler declaration carries no Dag, so it code-generates and schema-validates in every SDK, and nothing on this path needs a DagSerialization implementation. @@ -129,8 +132,8 @@ WatchedSubprocess │ └── coordinator.parse_dag() — spawn runtime, forward fd 0 ⇄ comm socket │ same request and result types as its base class │ - └── SDKTaskHandlerProcessorProcess (new — ADR-0011) - target = _parse_task_handler_entrypoint + └── LangSDKTaskHandlerProcessorProcess (new — ADR-0011) + target = _start_task_handler_runtime_entrypoint └── coordinator.parse_task_handler() — same forwarding TaskHandlerParseRequest → TaskHandlerParsingResult ``` @@ -208,13 +211,23 @@ bodies that never differ. Reusing `ToManager` leaves one reply union with one ne A boolean on `DagFileParseRequest` was the other alternative. It cannot work: a `DagRef` and a `TaskHandlerRef` are different payloads, not two subsets of one, so the flag would select between shapes the result type cannot both hold. -### Appendix B — What the nullable, ordered parameter list forces +### Appendix B — What the nullable schema and the binding mode force `value_schema` is nullable on both sides. An unannotated `@task.stub` parameter produces `value_schema: null` today, so validation compares schemas only where neither side is null, and falls back to name-and-arity otherwise. A strict comparison would turn every untyped stub argument into a parse error. -`params` is ordered because `LiteralArgBinding` and `XComArgBinding` are each documented as "one positional stub-task argument". Position is part of the contract, not incidental, -and both sides bind positionally. +The stub side always has names and positions: `LiteralArgBinding` and `XComArgBinding` are each documented as "one positional stub-task argument". The handler side binds the way +its runtime does, so each declaration names its `binding` and the check follows it: + +- `positional`: by position. Names are informative only, and absent where the runtime has none (Go flat params). Java's `TaskArgs` binds this way although it has names. An + argument count that matches `params` neither with every argument nor after dropping the defaulted ones, or a value type a parameter does not accept, is an import error. +- `named`: by name in any order, case-insensitively with underscores ignored unless `exact_name` is set (Go `arg:` tags, explicit Java names). A Go struct, tagged or untagged, + and Java's `TaskInput` bind this way. An argument no parameter takes, or a parameter no argument fills, is logged as a warning, and the task still runs: the runtimes allow + both, and an unfilled field keeps its default. When no parameter matches and exactly one argument was passed, it may be the whole value and is not warned about, unless no + parameter is declared, a parameter sets `exact_name` (a Go `arg:` tag), or the argument cannot be an object. A value type a parameter does not accept is an import error. + +`params` is `None` when the runtime cannot list a handler's parameters, and then only the handler's presence is checked. TypeScript declares `named` with `params: None`: its +types are erased, so a handler cannot list what it takes. A declaration carries no class, method, or source location. [ADR-0006](0006-no-lang-sdk-source-display.md) rules out Lang-SDK source display, and putting it on the wire would invite a consumer to render it. It carries no `dag_id` either — the `task_handlers` key supplies it, so a declaration cannot disagree with the bucket it arrived in. diff --git a/airflow-core/adr/lang-sdk/0013-persisted-task-handler-bindings.md b/airflow-core/adr/lang-sdk/0013-persisted-task-handler-bindings.md index cbcace7df5231..4e06cb12d4c39 100644 --- a/airflow-core/adr/lang-sdk/0013-persisted-task-handler-bindings.md +++ b/airflow-core/adr/lang-sdk/0013-persisted-task-handler-bindings.md @@ -21,7 +21,7 @@ ## Status -Proposed +Accepted ## Context @@ -157,31 +157,34 @@ backs many handlers and its fingerprint must have exactly one value. ```sql CREATE TABLE lang_sdk_task_handler_artifact ( - id UUID NOT NULL, - bundle_name VARCHAR(250) NOT NULL, -- the task_handler_bundle_name it was found in - relative_fileloc VARCHAR(2000) NOT NULL, -- path within that bundle - size_bytes BIGINT NOT NULL, -- cheap fingerprint tier - cache_digest VARCHAR(64) NOT NULL, -- content fingerprint tier; see "The fast path" - last_probed_at TIMESTAMP NOT NULL, - PRIMARY KEY (id), - CONSTRAINT lstha_bundle_fileloc_uq UNIQUE (bundle_name, relative_fileloc) + id UUID NOT NULL, + bundle_name VARCHAR(250) NOT NULL, -- the task_handler_bundle_name it was found in + relative_fileloc VARCHAR(2000) NOT NULL, -- path within that bundle + relative_fileloc_hash VARCHAR(32) NOT NULL, -- md5 of relative_fileloc; the path is too long to index + size_bytes BIGINT NOT NULL, -- cheap fingerprint tier + cache_digest VARCHAR(128) NULL, -- content fingerprint tier, NULL if none is stored; see "The fast path" + task_handlers JSON NOT NULL, -- {dag_id: [TaskHandlerDeclaration, ...]}, the probe answer + last_probed_at TIMESTAMP NOT NULL, + CONSTRAINT lang_sdk_task_handler_artifact_pkey PRIMARY KEY (id), + CONSTRAINT lang_sdk_task_handler_artifact_bundle_fileloc_uq UNIQUE (bundle_name, relative_fileloc_hash) ); CREATE TABLE lang_sdk_task_handler ( - dag_id VARCHAR(250) NOT NULL, - task_id VARCHAR(250) NOT NULL, - artifact_id UUID NOT NULL, - dag_bundle_name VARCHAR(250) NOT NULL, -- the *Python* file that owns this row - dag_relative_fileloc VARCHAR(2000) NOT NULL, -- ditto - handler_params JSON NOT NULL, -- list[TaskHandlerParam], ordered - PRIMARY KEY (dag_id, task_id), - CONSTRAINT lsth_dag_fkey FOREIGN KEY (dag_id) + dag_id VARCHAR(250) NOT NULL, + task_id VARCHAR(250) NOT NULL, + artifact_id UUID NOT NULL, + dag_bundle_name VARCHAR(250) NOT NULL, -- the *Python* file that owns this row + dag_relative_fileloc VARCHAR(2000) NOT NULL, -- ditto + dag_relative_fileloc_hash VARCHAR(32) NOT NULL, -- md5 of dag_relative_fileloc + CONSTRAINT lang_sdk_task_handler_pkey PRIMARY KEY (dag_id, task_id), + CONSTRAINT lang_sdk_task_handler_dag_id_fkey FOREIGN KEY (dag_id) REFERENCES dag (dag_id) ON DELETE CASCADE, - CONSTRAINT lsth_artifact_fkey FOREIGN KEY (artifact_id) + CONSTRAINT lang_sdk_task_handler_artifact_id_fkey FOREIGN KEY (artifact_id) REFERENCES lang_sdk_task_handler_artifact (id) ); -CREATE INDEX idx_lsth_dag_file ON lang_sdk_task_handler (dag_bundle_name, dag_relative_fileloc); -CREATE INDEX idx_lsth_artifact_id ON lang_sdk_task_handler (artifact_id); +CREATE INDEX idx_lang_sdk_task_handler_dag_file + ON lang_sdk_task_handler (dag_bundle_name, dag_relative_fileloc_hash); +CREATE INDEX idx_lang_sdk_task_handler_artifact_id ON lang_sdk_task_handler (artifact_id); ``` `PRIMARY KEY (dag_id, task_id)` is the conflict guard: two artifacts claiming the same task cannot @@ -194,11 +197,10 @@ joining through `DagModel`: `dag.relative_fileloc` is not indexed, and the codeb that querying it means "a sequential scan of dag" ([`reassign_dags_with_unconfigured_bundles`](https://github.com/apache/airflow/blob/79991cd4db0c9346a28b23c453377f6df0c6b4ed/airflow-core/src/airflow/dag_processing/bundles/manager.py#L478)). -`handler_params` stores what the runtime declared, so a changed Python file can be re-validated -against a cached declaration with no subprocess. It is deliberately **not** called `arg_bindings`: -that name already denotes the Python side of the comparison (`XComArgBinding` / `LiteralArgBinding`, -carrying wiring and values, [`build_arg_bindings`](https://github.com/apache/airflow/blob/79991cd4db0c9346a28b23c453377f6df0c6b4ed/airflow-core/src/airflow/serialization/stub_arg_bindings.py#L221-L286)), -and reusing it would make the validation read as comparing a thing to itself. +`task_handlers` caches the artifact's whole probe answer, every handler it registers keyed by +`dag_id`, so a changed Python file can be re-validated against it with no subprocess. It is written +only from a probe, together with `size_bytes` and `cache_digest`, so the answer always belongs to the +fingerprint beside it. A binding row is then only a mapping from a stub task to its artifact. `cache_digest` is **opaque and coordinator-defined**, not "SHA-256 of the file". @@ -208,15 +210,16 @@ and reusing it would make the validation read as comparing a thing to itself. already knows about. The parse child processor subprocesses run in the client context without a database connection ([`_parse_file_entrypoint`](https://github.com/apache/airflow/blob/79991cd4db0c9346a28b23c453377f6df0c6b4ed/airflow-core/src/airflow/dag_processing/processor.py#L208-L232)), so the prior cache state must be pushed down from the manager. The alternative is a dedicated -Execution API for the child processor process to retrieve `KnownSDKTaskHandlerArtifact` itself, which +Execution API for the child processor process to retrieve `TaskHandlerArtifact` itself, which was rejected on blast radius. ```python -class KnownSDKTaskHandlerArtifact(BaseModel): +class TaskHandlerArtifact(BaseModel): bundle_name: str relative_fileloc: str size_bytes: int - cache_digest: str + cache_digest: str | None # None: the artifact stores no fingerprint, so it is always probed + task_handlers: dict[str, list[TaskHandlerDeclaration]] # the probe answer: dag_id -> declarations class DagFileParseRequest(BaseModel): @@ -224,64 +227,58 @@ class DagFileParseRequest(BaseModel): bundle_path: Path bundle_name: str callback_requests: list[CallbackRequest] - known_artifacts: list[KnownSDKTaskHandlerArtifact] = [] # new + known_artifacts: list[TaskHandlerArtifact] = [] # new type: Literal["DagFileParseRequest"] ``` -`known_artifacts` is scoped by *artifact* bundle, not by Dag file, so the manager reads it **once per -parsing loop** for every configured `task_handler_bundle_name` and pushes the same list to every -child. Ten Dag files resolving against one twenty-jar bundle therefore probe that bundle once in -total, not once each. +`known_artifacts` is scoped by *artifact* bundle, not by Dag file, so the manager reads it **once per parsing loop** for every bundle that can hold task handlers: the `task_handler_bundle_name` of each coordinator a queue routes to, plus the Dag bundles it parses when one of them sets none and so reads the task's own Dag bundle. A coordinator no queue routes to runs no stub task, so its bundle is not read. Each child gets the rows of the named bundles whose team matches its Dag bundle's team, and of its own Dag bundle when a coordinator falls back to it, never those of another Dag bundle, which no coordinator reads for it. The team rule exists because an answer one file's parse records is trusted by every file that reads it. Teams match only when equal: without `[core] multi_team` every named bundle is in scope, and with it a team-less Dag bundle sees only team-less named bundles. One recorded answer serves every file, so ten Dag files resolving against one bundle probe a changed artifact once in total, apart from the children already started when it changed. **Dag-parsing child → coordinator subprocess.** Introduced here. The child spawns the runtime and forwards bytes in both directions, decoding nothing; the process that spawned the parse decodes the reply. `ToSDKTaskHandlerProcessor` is a new parent-to-child union differing from `ToDagProcessor` in -one member, and `ToManager` gains `SDKTaskHandlerParsingResult`. +one member, and `ToManager` gains `TaskHandlerParsingResult`. ```python -class SDKTaskHandlerParseRequest(BaseModel): # parent -> runtime, on ToSDKTaskHandlerProcessor +class TaskHandlerParseRequest(BaseModel): # parent -> runtime, on ToSDKTaskHandlerProcessor file: str # the candidate artifact being probed - dag_ids: list[str] # every Dag in this file with stub tasks routed here bundle_path: Path bundle_name: str - type: Literal["SDKTaskHandlerParseRequest"] + type: Literal["TaskHandlerParseRequest"] -class SDKTaskHandlerParsingResult(BaseModel): # runtime -> parent, on ToManager +class TaskHandlerParsingResult(BaseModel): # runtime -> parent, on ToManager fileloc: str task_handlers: dict[str, list[TaskHandlerDeclaration]] # dag_id -> declarations import_errors: dict[str, str] | None = None warnings: list | None = None - type: Literal["SDKTaskHandlerParsingResult"] + type: Literal["TaskHandlerParsingResult"] class TaskHandlerDeclaration(BaseModel): task_id: str - params: list[TaskHandlerParam] # ordered; arg bindings are positional + binding: Literal["positional", "named"] # how arguments bind to params + params: list[TaskHandlerParam] | None # ordered; None: the runtime cannot list them class TaskHandlerParam(BaseModel): - name: str + name: str | None # None: the runtime has no name for this positional parameter value_schema: JSONSchema | None = None - required: bool # the handler declares no default + exact_name: bool = False # match as spelled, not case-insensitively with underscores ignored ``` -A `dag_id` the artifact registers nothing for is **omitted** from `task_handlers` rather than returned -empty, so a probe that matches nothing is distinguishable from a probe that matched a Dag with zero -tasks. +`task_handlers` maps every `dag_id` the artifact registers a handler for, and is `{}` when it registers +none. The answer must not depend on the request, because the manager records it once and serves it to +every Dag file that resolves against the artifact. **Dag-parsing child → manager.** [`DagFileParsingResult`](https://github.com/apache/airflow/blob/79991cd4db0c9346a28b23c453377f6df0c6b4ed/airflow-core/src/airflow/dag_processing/processor.py#L133-L145) gains the resolved -bindings. +bindings and the probed artifacts. ```python -class SDKTaskHandlerBinding(BaseModel): +class TaskHandlerBinding(BaseModel): dag_id: str task_id: str artifact_bundle_name: str artifact_rel_path: str - artifact_size_bytes: int - artifact_cache_digest: str - handler_params: list[TaskHandlerParam] class DagFileParsingResult(BaseModel): @@ -289,15 +286,18 @@ class DagFileParsingResult(BaseModel): serialized_dags: list[LazyDeserializedDAG] warnings: list | None = None import_errors: dict[str, str] | None = None - task_handler_bindings: list[SDKTaskHandlerBinding] | None = None # new + task_handler_bindings: list[TaskHandlerBinding] | None = None # new + probed_artifacts: list[TaskHandlerArtifact] = [] # new ``` -`None` and `[]` mean different things, and the difference is load-bearing: +`probed_artifacts` holds every artifact the parse probed, with its answer; it is recorded even when the bindings are `None`, and it is the only way an answer gets written. + +For `task_handler_bindings`, `None` and `[]` mean different things, and the difference is load-bearing: | value | meaning | manager does | |--------|------------------------------------------|-----------------------| | `None` | handlers were not evaluated in this parse | **nothing** — no reconcile | -| `[]` | evaluated, this file has no stub handlers | delete this file's rows | +| `[]` | evaluated, this file has no stub handlers | delete the rows of this result's Dags | | `[…]` | evaluated, these are the bindings | reconcile to this set | `None` covers the stability-check early return ([`_parse_file`](https://github.com/apache/airflow/blob/79991cd4db0c9346a28b23c453377f6df0c6b4ed/airflow-core/src/airflow/dag_processing/processor.py#L245-L251)), callback-only runs @@ -305,24 +305,21 @@ class DagFileParsingResult(BaseModel): `persist_parsing_result` already uses ([`persist_parsing_result`](https://github.com/apache/airflow/blob/79991cd4db0c9346a28b23c453377f6df0c6b4ed/airflow-core/src/airflow/dag_processing/manager.py#L1361-L1364)). Without it, a transient parse failure would silently wipe every binding the file owns. -On a fast-path skip the child **re-emits the bindings it was given**, unchanged. It does not omit -them. Request and result carrying the same information makes that a copy rather than a special -"keep these" signal the reconcile could get wrong. +A candidate the fast path skips is answered from its recorded `task_handlers`, so the child resolves and validates every stub task of the file on every parse, and a list always holds all of the file's bindings. There is no separate "keep these" signal for the reconcile to get wrong. -**Scheduler → worker.** `ExecuteTask` and `StartupDetails` each gain one optional reference. The -artifact bundle is a second, independent bundle, so it needs its own `BundleInfo`. +**Scheduler → worker.** `ExecuteTask` and `StartupDetails` each gain one optional reference to the artifact that implements the stub task. ```python -class SDKTaskHandlerRef(BaseModel): - bundle_info: BundleInfo # the artifact bundle: name, version, version_data - rel_path: str # path within it +class TaskHandlerArtifactRef(BaseModel): + bundle_info: BundleInfo | None = None # the artifact bundle; None: the task's own Dag bundle + rel_path: str # POSIX path within it class ExecuteTask(BaseDagBundleWorkload): ti: TaskInstanceDTO dag_rel_path: os.PathLike[str] # the Python Dag file, unchanged bundle_info: BundleInfo # the Dag bundle, unchanged - task_handler: SDKTaskHandlerRef | None = None # new + task_handler_artifact: TaskHandlerArtifactRef | None = None # new ... @@ -330,10 +327,14 @@ class StartupDetails(BaseModel): ti: TaskInstance dag_rel_path: str bundle_info: BundleInfo - task_handler: SDKTaskHandlerRef | None = None # new + task_handler_artifact: TaskHandlerArtifactRef | None = None # new, as the workload carries it ... ``` +An artifact in a named bundle needs that bundle's own `BundleInfo`, since it is a second, independent bundle. An artifact in the task's own Dag bundle has `bundle_info=None` instead, because a copy of the Dag's `BundleInfo` would send its `version_data`, which can be a whole object manifest, a second time on every workload. + +The task's own bundle is read at the version the run uses: its pinned version, or the version current when the task starts if the run is not pinned. For a pinned run, the artifact then matches the Dag code the run is pinned to. A named bundle carries no version, since artifact rows record none, so it resolves to the version current when the task starts. Either way the resolved version is pinned for the whole task. The scheduler decides "own bundle" by name, so a coordinator whose `task_handler_bundle_name` names the Dag's own bundle also reads it at the version the run uses. + `None` means "this task needs no Lang-SDK artifact" (an ordinary Python task). It never means "unknown": a stub task that failed to resolve is not queued at all (see "Failure handling"). @@ -347,12 +348,17 @@ stub Dag is validated against bindings that have not been written yet. ``` DagProcessorManager [reads DB] │ - │ once per parsing loop, per configured task_handler_bundle_name: - │ SELECT bundle_name, relative_fileloc, size_bytes, cache_digest + │ once per parsing loop, for every bundle that can hold task handlers + │ (each routed coordinator's task_handler_bundle_name, plus the Dag + │ bundles when one sets none and reads the task's own Dag bundle): + │ SELECT bundle_name, relative_fileloc, size_bytes, cache_digest, + │ task_handlers │ FROM lang_sdk_task_handler_artifact - │ WHERE bundle_name IN (:configured bundles) ──▶ known_artifacts + │ WHERE bundle_name IN (:those bundles) ──▶ known_artifacts │ - │ per file about to be dispatched: + │ per file about to be dispatched: the rows of the bundles in its scope + │ (named bundles of its team, plus its own Dag bundle when a routed + │ coordinator sets none) │ (the child re-reads nothing; it has no DB) │ ├── DagFileParseRequest(file=etl.py, bundle_*, known_artifacts=[...]) @@ -372,64 +378,90 @@ DagFileProcessorProcess(etl.py) [no DB — client con │ transform → "java" → jdk-17 → "java-task-handlers" │ ingest → "go" → go-sdk → "go-task-handlers" │ - │ a queue with no coordinator entry is an import error here, - │ not a silent Python fallback at execution time + │ a stub task on a queue with no coordinator entry is left to an + │ external worker: not checked, not bound; without + │ queue_to_coordinator nothing is checked + │ each coordinator reads its task_handler_bundle_name, or the Dag's + │ own bundle when unset; a bundle of another team is an import + │ error, since teams must be equal │ - ├─4─ per coordinator: list candidates in its bundle - │ Go → files carrying the AFBNDL01 trailer magic - │ Java → *.jar with a Main-Class manifest attribute - │ TS → *.min.mjs with a valid //# airflowBundle= layout header + ├─4─ per coordinator: list candidates in its bundle; a coordinator that + │ cannot list its artifacts is an import error for its stub tasks + │ Go → files carrying the AFBNDL01 trailer magic, whatever + │ their executable bit (one without it is rejected) + │ Java → *.jar whose manifest carries Airflow-Cache-Digest + │ (dependency JARs may declare Main-Class too) + │ TS → *.min.mjs whose first line starts with //# airflowBundle= + │ (one with an invalid layout line is rejected) │ walk order is deterministic, so conflicts reproduce │ ├─5─ FAST PATH, per candidate (see "The fast path") - │ size + cache_digest match known_artifacts, and the candidate - │ set is unchanged ──▶ skip the launch, echo the bindings - │ anything differs ──▶ probe + │ size + stored cache_digest equal a known artifact's + │ ──▶ skip the launch, use its task_handlers + │ anything else, or no stored digest ──▶ probe + │ rejected by the listing ──▶ broken, no probe │ - ├─6─ PROBE, per differing candidate — one subprocess - │ SDKTaskHandlerProcessorProcess.start( - │ target=_parse_task_handler_entrypoint, - │ coordinator=JavaCoordinator("jdk-17"), - │ path=) - │ ──SDKTaskHandlerParseRequest(file=…, dag_ids=["etl"])──▶ runtime - │ ◀─SDKTaskHandlerParsingResult(task_handlers={"etl": […]})── runtime + ├─6─ PROBE, per differing candidate: one subprocess each, one at a time, + │ all within 90% of [dag_processor] dag_file_processor_timeout, + │ counted from the creation of the parse child; an answer is shared by + │ every coordinator that lists the artifact, and a failed probe is + │ retried under the next coordinator that lists it; a candidate + │ whose probes all fail, or that is left when time runs out, is broken + │ LangSDKTaskHandlerProcessorProcess.run( + │ coordinator="jdk-17", + │ path=, bundle_path=…, bundle_name=…, + │ artifact_rel_path=…, deadline=…) + │ ──TaskHandlerParseRequest(file=…)──▶ runtime + │ ◀─TaskHandlerParsingResult(task_handlers={"etl": […]})── runtime │ Get* from the runtime is relayed up ToManager unchanged │ - ├─7─ VALIDATE per dag_id, unioned across coordinators - │ task_id sets must match exactly - │ arg_bindings[*].name ↔ handler_params[*].name, in order - │ arg_bindings[*].schema ↔ handler_params[*].value_schema, - │ compared only where neither is null - │ two candidates claiming one (dag_id, task_id) → import error - │ naming both paths + ├─7─ VALIDATE per stub task, against the recorded and the fresh answers + │ of the coordinator its queue routes to + │ exactly one declaration of its (dag_id, task_id); none, or two + │ candidates claiming it → import error naming the candidates + │ a handler with no stub task → not an error + │ arg_bindings[*] ↔ declaration.params[*], per its binding: + │ by position, or by folded or exact name; + │ an argless call passes no arguments; + │ defaulted ones are dropped if that + │ makes a positional count match + │ arg_bindings[*].schema ↔ declaration.params[*].value_schema, + │ top-level JSON types, where both exist + │ mapped stub task, or params=None → the handler's presence only │ └─8─ on success → task_handler_bindings=[…] on mismatch → import_errors[etl.py]=… AND bindings=None + either way → probed_artifacts=[every fresh answer] │ ├── DagFileParsingResult(serialized_dags=[…], import_errors={…}, - ▼ task_handler_bindings=[…] | [] | None) -DagProcessorManager.persist_parsing_result [writes DB] - └── update_dag_parsing_results_in_db — one transaction, in this order: + ▼ task_handler_bindings=[…] | [] | None, + probed_artifacts=[…]) +DagProcessorManager.handle_parsing_result [writes DB] + ├── record the probed artifacts, a transaction of their own, skipped when there are none ← new + │ lang_sdk_task_handler_artifact UPSERT by (bundle_name, relative_fileloc): + │ fingerprint, answer, last_probed_at + │ + └── persist_parsing_result → update_dag_parsing_results_in_db — one transaction, in this order: 1. add_dags / update_dags → DagModel rows exist 2. asset reference tables (existing) - 3. lang_sdk_task_handler_artifact UPSERT by (bundle_name, relative_fileloc) ← new - 4. lang_sdk_task_handler reconcile by dag_id ← new - 5. SerializedDagModel / DagVersion / DagCode - 6. ParseImportError / DagWarning + 3. lang_sdk_task_handler reconcile by dag_id, artifact ids looked up by key ← new + 4. SerializedDagModel / DagVersion / DagCode + 5. ParseImportError / DagWarning ``` -Step 3 must upsert: two parse children can discover the same artifact in the same loop and race on -the unique key. +The probed-artifact write must upsert: two parse children can probe the same artifact in the same loop and race on the unique key. It commits before the persist starts, so its exclusive row locks are gone before step 3 takes shared ones, and an answer is kept even when the persist fails. -Step 4 reconciles **by the `dag_id`s in the result**, not by file path: delete rows for those +Step 3 reconciles **by the `dag_id`s in the result**, not by file path: delete rows for those `dag_id`s whose `task_id` is absent from the returned set, then insert or update the rest. Path-keyed eviction breaks when a Dag moves between files — the old rows stay keyed to a path nothing parses any more, and the primary key then blocks the new insert. With `dag_id` as the key a move simply updates -`dag_relative_fileloc`, and a Dag that disappears entirely is reclaimed by the `ON DELETE CASCADE`. +`dag_relative_fileloc`. A Dag removed from its file is only marked stale, so its rows stay until its `dag` row is deleted, and the `ON DELETE CASCADE` then removes them. Until then they keep their artifacts from the orphan sweep. + +Step 3 writes nothing to the artifact table. A Dag bound to an artifact outside the Dag file's scope, or to one with no row because a concurrent orphan sweep deleted it, keeps its rows as they are, with a warning; in the second case the next parse finds no recorded answer and probes again. Artifact rows are **never** evicted from one file's result. One artifact backs handlers owned by many Python files, so this file seeing fewer candidates says nothing about another file's. They are a -cache; they are reclaimed by orphan sweep or by `db clean`, never by a per-file reconcile. +cache; they are reclaimed by orphan sweep, never by a per-file reconcile. ### Flow 2) Scheduling: the read path @@ -457,10 +489,10 @@ SchedulerJobRunner._enqueue_task_instances_with_queued_state [no further DB re is_stub and no binding? → FAIL the TI with the reason ← new never queue it │ - └── ExecuteTask.make(ti, task_handler=SDKTaskHandlerRef(...)) - dag_rel_path ← ti.dag_model.relative_fileloc (existing) - bundle_info ← Dag bundle, pinned to the run (existing) - task_handler ← artifact bundle + rel_path ← new + └── ExecuteTask.make(ti, task_handler_artifact=TaskHandlerArtifactRef(...)) + dag_rel_path ← ti.dag_model.relative_fileloc (existing) + bundle_info ← Dag bundle, pinned to the run (existing) + task_handler_artifact ← artifact bundle + rel_path ← new │ └── executor.queue_workload(workload) ``` @@ -473,22 +505,22 @@ A stub task with no binding is **failed with its reason**, not skipped. executor worker process └── BaseExecutor.run_workload(workload) └── supervise_task(ti=…, bundle_info=…, dag_rel_path=…, - task_handler=workload.task_handler) ← new + task_handler_artifact=workload.task_handler_artifact) ← new │ ├── coordinator = get_coordinator_manager().for_queue(ti.queue) │ unchanged: execution still routes on queue │ - └── coordinator.execute_task(what=ti, …, task_handler=…) + └── coordinator.execute_task(what=ti, …, task_handler_artifact=…) + ├── bundle = initialize(task_handler_artifact.bundle_info + │ or the task's bundle_info) ← the artifact's bundle + │ (the task's own, or a named one) + │ pinned and held under BundleVersionLock for the whole task + ├── bundle.path / rel_path is not a file in it? → fail the task + │ └── SubprocessCoordinator._build_execute_task_command( - what=ti, task_handler=task_handler) ← signature change + what=ti, task_handler_artifact=…) ← signature change │ - ├── bundle = DagBundlesManager().get_bundle( - │ name=task_handler.bundle_info.name, - │ version=task_handler.bundle_info.version, - │ version_data=task_handler.bundle_info.version_data) - │ bundle.initialize() ← the SECOND bundle - │ - ├── artifact = bundle.path / task_handler.rel_path + ├── artifact = bundle.path / task_handler_artifact.rel_path │ NO directory walk. NO dag_id match. NO metadata["dags"]. │ ├── verify integrity, read supervisor_schema_version @@ -505,7 +537,7 @@ executor worker process │ └── _PopenActivitySubprocess.start(…) ──StartupDetails(ti, dag_rel_path, bundle_info, - task_handler)──▶ runtime + task_handler_artifact)──▶ runtime runtime looks up its own registration by (ti.dag_id, ti.task_id) — unchanged ``` @@ -522,37 +554,35 @@ stored digest, so one cannot substitute for the other. Validation is expensive (one subprocess per candidate) and a file is re-parsed every `[dag_processor] min_file_process_interval` seconds, 30 by default. Re-probing unchanged artifacts -every 30 seconds forever is not acceptable, so the probe is skipped when nothing relevant changed. +every 30 seconds forever is not acceptable, so a candidate is not probed while its recorded answer still holds. -Two tiers, cheapest first: +The probe asks for every handler and the answer depends only on the artifact, so the answer recorded with a fingerprint holds for every Dag file that later sees that fingerprint, whichever file probed it. Each candidate is decided on its own, cheapest check first: ``` for each candidate in the coordinator's bundle: - stat(candidate).st_size ≠ known.size_bytes → PROBE - read stored cache_digest ≠ known.cache_digest → PROBE - otherwise → SKIP, echo the binding + rejected by the listing (unusable artifact) → REPORT, never probe + no known artifact at (bundle, relative path) → PROBE + stat(candidate).st_size ≠ known.size_bytes → PROBE + stored cache_digest missing, or ≠ known.cache_digest → PROBE + otherwise → SKIP, use known.task_handlers ``` -`mtime` is deliberately absent. It is reset by an object-store download and by container rebuilds, so -it produces churn without adding certainty; size plus digest is sufficient, with size acting only as -a free pre-filter. - -Two conditions beyond the per-file comparison: +`mtime` is deliberately absent. It is reset by an object-store download and by container rebuilds, so it produces churn without adding certainty. Size plus digest is sufficient: the stored digest is read, not recomputed, so an artifact edited in place without a repack keeps it, and the free size check catches most such edits. Execution's integrity check stays the safety net. -**The candidate set must be unchanged.** A newly deployed artifact has no known fingerprint, so a set -that differs from `known_artifacts` forces a probe. Without this, an added artifact that collides on -`(dag_id, task_id)` would never be detected, and the conflict rule above would be unenforceable. +**A new artifact is probed by the first file that lists it.** It has no recorded answer yet. Its answer lists every Dag it registers, so a collision on `(dag_id, task_id)` with any file's Dag is found when that file is next validated, even when another file probed the artifact first. -**The Python side is re-validated regardless.** The skip avoids the *subprocess*, not the comparison. -`handler_params` is stored precisely so a changed `.py` (a stub task that gained an argument) is -compared against the cached declaration in process. Skipping the comparison as well would cache a -verdict for a signature that no longer exists, and ship the mismatch to a worker as a runtime -argument error instead of catching it as an import error. +**The Python side is re-validated regardless.** The skip avoids the *subprocess*, not the comparison. The recorded `task_handlers` are kept precisely so a changed `.py` (a stub task that gained an argument) is compared against the cached declaration in process. Skipping the comparison as well would cache a verdict for a signature that no longer exists, and ship the mismatch to a worker as a runtime argument error instead of catching it as an import error. ### Failure handling -**Validation fails.** Record the import error; leave every existing row untouched. `bindings=None` -already expresses "do not reconcile", so this needs no additional mechanism. +**Validation fails.** Record the import error; leave every binding row untouched. `bindings=None` +already expresses "do not reconcile", so this needs no additional mechanism. The answers probed for it are still recorded, so the next parse validates against them without probing. Like any import error, it stops every Dag of the file from being scheduled until it clears. + +**A candidate is broken.** It is rejected by the listing, its probe fails or raises, or the parse runs out of time first. It is ignored: a stub task that finds its handler elsewhere is checked and bound as usual, and one left without a handler fails, its import error naming the broken candidate and why. + +**A coordinator cannot be evaluated.** It cannot be built or cannot list its artifacts, or its bundle is missing, cannot be read or belongs to another team. Its stub tasks are not checked, and it is an import error naming it. An error the check does not expect is an import error of each Dag file with a routed stub task, and the serialized Dags are still sent. + +**A name mismatch under named binding.** A passed argument no param takes, or a param no argument fills, is a warning in the Dag file's parse log; the stub task is still bound. **A stub task has no binding at all.** The scheduler fails it with the reason rather than queueing a workload that would die on the worker at `ValueError("dag_path is required")`, far from the cause. @@ -570,23 +600,39 @@ definition whose author can act) naming both artifact paths, since the fix is in and versioning from `DagBundle`. Deployments that mount artifacts themselves point a `LocalDagBundle` at the mount. - Two new tables and one migration. `lang_sdk_task_handler` is reconciled on every parse of a file that owns rows in it; `lang_sdk_task_handler_artifact` is a cache with no per-file eviction. -- The artifact bundle is a second bundle on the execution path. `ExecuteTask` and `StartupDetails` - each grow one optional `SDKTaskHandlerRef`, and the worker performs a second `initialize()`. Workload payloads grow by roughly one `BundleInfo`. -- `DagFileParseRequest` and `DagFileParsingResult` each gain a field, and `ToSDKTaskHandlerProcessor` +- `ExecuteTask` and `StartupDetails` each grow one optional `TaskHandlerArtifactRef`. On the coordinator + path the worker initializes the artifact's bundle, the task's own or a named one, in place of the Dag bundle, not in addition to it. Workload payloads grow by an artifact path, plus a bundle name for a named bundle. +- `DagFileParseRequest` gains a field and `DagFileParsingResult` two, and `ToSDKTaskHandlerProcessor` becomes a fifth union the supervisor-schema registry introspects. Both messages already appear in the generated schemas of all three SDKs, so the snapshot is regenerated and the two prek hooks guarding it run. -- Every Lang SDK runtime must answer `SDKTaskHandlerParseRequest`. -- A misrouted queue becomes an import error at Dag-parsing stage instead of a runtime failure. Today a stub task on a - queue absent from `queue_to_coordinator` silently falls back to the Python coordinator - ([`for_queue`](https://github.com/apache/airflow/blob/79991cd4db0c9346a28b23c453377f6df0c6b4ed/task-sdk/src/airflow/sdk/execution_time/coordinator.py#L280-L284)) and dies in - `_StubOperator.execute()`. +- Every Lang SDK runtime must answer `TaskHandlerParseRequest`. +- A stub task on a queue absent from `queue_to_coordinator` is left to a worker outside Airflow's coordinators: + the parse neither checks nor binds it. Without `queue_to_coordinator` nothing is checked, so + Python-only deployments are unchanged. +- A coordinator serving stub tasks must list and probe its artifacts, since a stub task it cannot bind cannot + run. +- An artifact whose probe fails, or that the parse runs out of time for, has no recorded answer, so every Dag + file with stub tasks on its coordinator probes it again on each parse until it is fixed or removed. One that + several coordinators list is probed under each of them until one answers, so it can be probed once per + coordinator on every parse. One the listing rejects is never probed: it is logged on each parse and named when + a stub task finds no handler. +- The parse's team check only reports. The parse runs Dag code, so the manager enforces the scope when it + records answers and bindings. +- Probes count towards `[dag_processor] dag_file_processor_timeout`, counted from the creation of the parse + child, and each is also limited by `[core] dagbag_import_timeout`, which `get_dagbag_import_timeout` can set + per artifact path. They run one at a time; the answers a parse got before running out of time are still + recorded, so a cold start with many artifacts converges over parses, unless an artifact whose probe keeps + failing slowly is probed ahead of them and leaves too little time. That artifact has no recorded answer, so + it is probed first again on every parse, and the artifacts after it are asked only once it is fixed or removed, + or the timeouts are raised. Guaranteed convergence would need the manager to record failed probes so that + they are probed last, or a probe order that rotates between parses. + If the manager kills the parse anyway, the kernel kills the runtime with it on Linux; processes the runtime + started itself are not covered. - One artifact bundle is one Java classpath. `_calculate_classpath` joins every JAR under the root ([`_calculate_classpath`](https://github.com/apache/airflow/blob/79991cd4db0c9346a28b23c453377f6df0c6b4ed/task-sdk/src/airflow/sdk/coordinators/java/coordinator.py#L85-L87)), so all handlers in a bundle share one dependency graph. Isolating conflicting dependency versions requires a second bundle, a second coordinator instance, and a second queue. This needs documenting. -- Steady-state parsing costs no subprocesses. A cold start (empty tables, or a bundle whose - candidate set changed) costs one subprocess per changed candidate per coordinator, shared across - all Dag files in that parsing loop through `known_artifacts`. +- Steady-state parsing costs no subprocesses. A new or changed artifact costs one subprocess for each Dag file parsed before its answer is recorded, at most the children the manager starts in one loop iteration, and none after that. A Dag that fails validation does not cause a probe on every parse, since the answers it was checked against stay recorded, and an artifact that stores no cache digest is probed on every parse. - No `DagVersion` coupling. An artifact rebuild does not bump a Dag's version, and a Dag edit does not invalidate an artifact fingerprint. - Mixed-language stays Python-primary. A Lang-SDK runtime cannot declare stub tasks, and a native Dag diff --git a/airflow-core/adr/lang-sdk/0014-bundle-metadata-and-cache-digest.md b/airflow-core/adr/lang-sdk/0014-bundle-metadata-and-cache-digest.md index dfc63beec195a..aecdc4d9bc4f0 100644 --- a/airflow-core/adr/lang-sdk/0014-bundle-metadata-and-cache-digest.md +++ b/airflow-core/adr/lang-sdk/0014-bundle-metadata-and-cache-digest.md @@ -88,8 +88,8 @@ No `dag_id` appears anywhere in it, and no `task_id`. Candidate *detection* is unaffected and still requires executing nothing — it is exactly the "this is an Airflow Lang-SDK artifact" marker doing its job: the `AFBNDL01` trailer magic for Go ([`Magic`](https://github.com/apache/airflow/blob/79991cd4db0c9346a28b23c453377f6df0c6b4ed/go-sdk/internal/bundlefooter/footer.go#L57)), the `.min.mjs` suffix plus a valid layout header for -TypeScript ([`BUNDLE_SUFFIX`](https://github.com/apache/airflow/blob/79991cd4db0c9346a28b23c453377f6df0c6b4ed/task-sdk/src/airflow/sdk/coordinators/node/coordinator.py#L43)), a `Main-Class` attribute -for Java. +TypeScript ([`BUNDLE_SUFFIX`](https://github.com/apache/airflow/blob/79991cd4db0c9346a28b23c453377f6df0c6b4ed/task-sdk/src/airflow/sdk/coordinators/node/coordinator.py#L43)), an `Airflow-Cache-Digest` manifest attribute +for Java (dependency JARs in the same bundle may declare `Main-Class` too). The entrypoint source region stays, for display. diff --git a/airflow-core/docs/authoring-and-scheduling/language-sdks/go.rst b/airflow-core/docs/authoring-and-scheduling/language-sdks/go.rst index bd2a3e4ac761a..a868b41a0238a 100644 --- a/airflow-core/docs/authoring-and-scheduling/language-sdks/go.rst +++ b/airflow-core/docs/authoring-and-scheduling/language-sdks/go.rst @@ -48,7 +48,8 @@ Prerequisites * Go 1.24 or later to build and pack bundles. This is a build-time requirement only; the worker that runs a packed bundle needs no Go toolchain, because the bundle is a self-contained native executable. -* The packed bundle must be accessible from the Airflow worker, under a directory the coordinator scans. +* The packed bundle must be accessible from the Airflow worker and the Dag processor, in the Dag bundle the + coordinator scans, and built for the operating system and CPU architecture of both. * The ``apache-airflow-task-sdk`` package (installed with Airflow) provides the coordinator; no additional Python packages are needed. @@ -178,33 +179,46 @@ to pass on with ``bundle.Register(reports.Handlers()...)``. Coordinator configuration ~~~~~~~~~~~~~~~~~~~~~~~~~~~ -Register the coordinator and route the queue to it under ``[sdk]`` in ``airflow.cfg`` (or the equivalent -``AIRFLOW__SDK__*`` environment variables): +Register a Dag bundle for the packed bundles, register the coordinator, and route the queue to it in +``airflow.cfg`` (or the equivalent ``AIRFLOW__*`` environment variables): .. code-block:: ini + [dag_processor] + dag_bundle_config_list = [ + {"name": "dags-folder", "classpath": "airflow.dag_processing.bundles.local.LocalDagBundle", "kwargs": {}}, + { + "name": "go-task-handlers", + "classpath": "airflow.dag_processing.bundles.local.LocalDagBundle", + "kwargs": {"path": "/opt/airflow/go-task-handlers"} + } + ] + [sdk] coordinators = { "go": { "classpath": "airflow.sdk.coordinators.executable.ExecutableCoordinator", - "kwargs": {"executables_root": ["~/airflow/executable-bundles"]} + "kwargs": {"task_handler_bundle_name": "go-task-handlers"} } } queue_to_coordinator = {"golang": "go"} -``executables_root`` is one or more directories the coordinator scans for bundles; ``queue_to_coordinator`` -routes stub tasks with ``queue="golang"`` to this Go coordinator. See :ref:`go-sdk/coordinator-config` for -the full list of accepted ``kwargs``. +``task_handler_bundle_name`` names the Dag bundle the coordinator scans for packed bundles; +``queue_to_coordinator`` routes stub tasks with ``queue="golang"`` to this Go coordinator. See +:ref:`go-sdk/coordinator-config` for the full list of accepted ``kwargs`` and how bundles are located. There is no separate Go worker to run: the Airflow worker forks the bundle binary once per task instance. .. note:: - The coordinator is part of the Airflow worker, so the ``[sdk]`` config (and the bundle files in - ``executables_root``) only need to be present wherever tasks actually execute. With ``CeleryExecutor``, - setting it on the Celery workers is sufficient. With ``LocalExecutor``, tasks run inside the scheduler - process, so it must be set where the scheduler can read it. The API server and Dag processor do not need - it. + The ``[sdk]`` config and the packed bundle files must be present wherever tasks execute and on the Dag + processor. With ``CeleryExecutor``, tasks execute on the Celery workers; with ``LocalExecutor``, they run + inside the scheduler process. The Dag processor checks the stub tasks of each Python Dag against the task + handlers the packed bundles register, so it runs them too and needs bundles built for its operating + system and CPU architecture. The API server does not need any of it. Register the Dag bundle in + ``[dag_processor] dag_bundle_config_list`` on every component, like your other Dag bundles: the worker + and the Dag processor resolve ``task_handler_bundle_name`` through it, and wherever the ``[sdk]`` config + is read it is rejected if the name is missing there. Writing tasks ------------- @@ -335,7 +349,9 @@ Every parameter after the ``airflow.Context`` is a **data parameter**, filled in arguments of the Python stub Task's TaskFlow call. A literal in the Dag file (``transform("uk", ...)``) decodes straight into the parameter; an upstream task's output (``transform(..., extract())``) is pulled from that task's XCom in the current Dag run. If the argument count does not match, or an argument's -declared type cannot fill the Go type, the task fails before its body runs. +declared type cannot fill the Go type, the Python Dag file fails to import (see +:ref:`language-sdks/dag-processor-checks`). A value that still cannot fill the Go type when the task runs +fails the task before its body runs. .. code-block:: go @@ -369,8 +385,10 @@ how unmatched fields and arguments are treated and when an untagged struct is de argument instead. Stub parameters the Dag author left at their Python defaults are the exception to both shapes: they reach -the wire but need no Go parameter, so adding a defaulted parameter to a stub does not break the Go -functions already bound to it. +the wire but need no Go parameter. Adding a defaulted parameter to a stub does not break a struct handler, +or a function that declares none of the stub's defaulted parameters. When a function's parameter count does +not match, every defaulted argument is dropped and the count is compared again, so a function that declares +only some of the defaulted parameters does not bind. .. _go-sdk/types: @@ -444,24 +462,24 @@ Build and pack in one step; any flags after ``--`` are forwarded verbatim to ``g go tool airflow-go-pack ./example/bundle -- -trimpath -tags=prod -Use ``--output `` to write the packed bundle straight into a directory the coordinator scans -(``executables_root``): +Use ``--output `` to write the packed bundle straight into the directory of the Dag bundle the +coordinator scans (see `Deploying`_): .. code-block:: bash - go tool airflow-go-pack --output ~/airflow/executable-bundles/sample-dag-bundle ./example/bundle + go tool airflow-go-pack --output /opt/airflow/go-task-handlers/sample-dag-bundle ./example/bundle Cross-platform builds ~~~~~~~~~~~~~~~~~~~~~~~ -The worker that runs a bundle often uses a different operating system or CPU architecture than your build -machine (for example, deploying to a Linux host from an Apple-silicon ``darwin/arm64`` laptop). Pass -``--goos`` / ``--goarch`` and the packer cross-builds for you: +The worker and the Dag processor that run a bundle often use a different operating system or CPU +architecture than your build machine (for example, deploying to a Linux host from an Apple-silicon +``darwin/arm64`` laptop). Pass ``--goos`` / ``--goarch`` and the packer cross-builds for you: .. code-block:: bash go tool airflow-go-pack --goos linux --goarch amd64 \ - --output ~/airflow/executable-bundles/sample-dag-bundle \ + --output /opt/airflow/go-task-handlers/sample-dag-bundle \ ./example/bundle Alternatively, pack a pre-built binary with ``--executable`` / ``--source``. The packer normally execs the @@ -485,12 +503,16 @@ with ``--airflow-metadata``: Deploying ~~~~~~~~~ -Copy or mount the packed bundle into a directory listed in the coordinator's ``executables_root``. The -:class:`~airflow.sdk.coordinators.executable.ExecutableCoordinator` scans those directories recursively, +Copy or mount the packed bundle into the Dag bundle named by the coordinator's ``task_handler_bundle_name``. +The :class:`~airflow.sdk.coordinators.executable.ExecutableCoordinator` scans that Dag bundle recursively, matches the incoming ``dag_id`` against each bundle's manifest, verifies the bundle's integrity hash, and launches the matching bundle. Bundles are identified by the trailer magic, not by filename (no extension on Linux/macOS, ``.exe`` on Windows), so the file name on the worker is irrelevant. +Only files with the executable bit set are considered, so the packed bundles need a Dag bundle that keeps +it. A ``LocalDagBundle`` over a directory you manage does; object-store Dag bundles such as ``S3DagBundle`` +do not. + .. _go-sdk/coordinator-config: :class:`~airflow.sdk.coordinators.executable.ExecutableCoordinator` configuration @@ -506,16 +528,13 @@ All ``kwargs`` in the ``coordinators`` config entry are passed to the * - Parameter - Default - Description - * - ``executables_root`` - - *(optional)* - - One or more directories scanned recursively for executable bundles. Accepts a string, - a path, or a list of strings/paths. When omitted, bundles are located through a Dag - bundle instead (see the note below). Explicitly setting this option to ``null`` or - an empty list is invalid. - * - ``dag_bundle_name`` - - *(auto: task's own bundle)* - - Name of a configured Dag bundle to load executable bundles from. Mutually exclusive - with ``executables_root``. + * - ``task_handler_bundle_name`` + - *(task's own Dag bundle)* + - Name of the Dag bundle scanned recursively for executable bundles. It is used only by + mixed-language Dags, to locate the task handlers for the ``@task.stub`` tasks of a Python Dag; + Dags defined natively in a language SDK do not use it. It must be registered in + ``[dag_processor] dag_bundle_config_list``. It is checked when the ``[sdk]`` configuration is + loaded, so a typo fails there rather than on the first task. * - ``task_startup_timeout`` - ``10.0`` - Seconds to wait for the bundle subprocess to connect after launch. Increase this if your @@ -523,15 +542,16 @@ All ``kwargs`` in the ``coordinators`` config entry are passed to the .. note:: - **Locating bundles.** ``executables_root`` and ``dag_bundle_name`` are mutually exclusive, - and both are optional: + **Locating bundles.** The packed bundles for the ``@task.stub`` tasks of a Python Dag are read from a Dag + bundle, so they are delivered, refreshed and versioned by the same machinery as your Dags. - * Set ``executables_root`` to scan explicit filesystem directories you manage yourself. - * Set ``dag_bundle_name`` to load bundles from a configured Dag bundle, so they are delivered - and versioned through the same bundle machinery as your Dags. The task uses the version that - bundle is on when it starts, pinned for the whole task. - * Leave both unset (the default) to load bundles from the **task's own** Dag bundle, pinned - to the version the run was created with. + * The expected layout is a separate Dag bundle for the packed bundles, named by + ``task_handler_bundle_name``, rather than the Dag bundle that holds your ``.py`` files. The task + uses the version that Dag bundle is on when it starts, pinned for the whole task. + * If ``task_handler_bundle_name`` is unset, packed bundles are read from the **task's own** Dag + bundle, at the version the run uses: its pinned version, or the version current when the task + starts if the run is not pinned. + * The Dag bundle must keep the executable bit (see `Deploying`_). .. _go-sdk/limitations: diff --git a/airflow-core/docs/authoring-and-scheduling/language-sdks/index.rst b/airflow-core/docs/authoring-and-scheduling/language-sdks/index.rst index 79a0cb3df1f4d..28c0b815bea5f 100644 --- a/airflow-core/docs/authoring-and-scheduling/language-sdks/index.rst +++ b/airflow-core/docs/authoring-and-scheduling/language-sdks/index.rst @@ -155,9 +155,9 @@ Coordinators are registered in ``airflow.cfg`` (or via environment variables) un } } - The ``classpath`` value must be importable by the worker. The ``kwargs`` are passed directly - to the coordinator's constructor. See the language-specific guide for the accepted kwargs - of each coordinator (e.g. :ref:`java-sdk/coordinator-config` for + The ``classpath`` value must be importable by the worker and the Dag processor. The ``kwargs`` + are passed directly to the coordinator's constructor. See the language-specific guide for the + accepted kwargs of each coordinator (e.g. :ref:`java-sdk/coordinator-config` for :class:`~airflow.sdk.coordinators.java.JavaCoordinator`). ``extra`` is an optional object for any additional information you want to associate with a @@ -187,6 +187,55 @@ Both settings can be supplied as environment variables using the standard Airflo AIRFLOW__SDK__COORDINATORS='{"my-coordinator": {...}}' AIRFLOW__SDK__QUEUE_TO_COORDINATOR='{"jdk17": "my-coordinator"}' +.. _language-sdks/dag-processor-checks: + +What the Dag processor checks +----------------------------- + +When the Dag processor parses a Python Dag file, it checks every stub task whose ``queue`` is in +``[sdk] queue_to_coordinator`` against the task handlers that the artifacts of its coordinator register. Those +are the artifacts in the Dag bundle named by the coordinator's ``task_handler_bundle_name``, or in the Dag's +own bundle when it is unset. + +* A stub task needs exactly one task handler among those artifacts. When no artifact registers it, or two do, + the Dag file fails to import. +* The task handler must take the stub task's arguments. An argument count it cannot bind by position, or an + argument whose declared type its parameter does not accept, makes the Dag file fail to import. When the + count does not match, every argument left at its default in the stub signature is dropped and the count is + compared again, so a handler that declares all or none of those parameters binds, and one that declares + only some does not. +* When the handler binds arguments by name, a passed argument it does not declare, or a parameter no argument + fills, is a warning in the Dag file's parse log. The task still runs. +* A mapped stub task, or one whose task handler does not list its parameters, is checked for the handler only. +* An artifact that cannot be run is a warning in the parse log. When a stub task finds no task handler, the + import error names each such artifact and why it has no answer. +* A stub task routed to a coordinator that cannot list its artifacts fails to import. So does one whose + coordinator reads a Dag bundle of another team than the Dag file's bundle. +* A stub task on a queue that ``queue_to_coordinator`` does not route is not checked, so a worker outside + Airflow's coordinators can run it. Without ``queue_to_coordinator``, nothing is checked. + +Each parse of a Dag file with stub tasks lists the files in the bundle of each coordinator they route to, and +a coordinator that runs executables opens every one of them. Give each coordinator a dedicated, small bundle +through ``task_handler_bundle_name`` rather than letting it search the whole Dag bundle. A Dag processor +started with ``--bundle-name`` refreshes only the bundles it is told to parse, so pass it the task handler +bundles its coordinators name as well, or have them checked out on its host. + +The Dag processor runs a new or changed artifact once to ask for its task handlers, and reuses the answer +until the artifact changes. An artifact that gives no answer, because its run fails or runs out of time, is +run again on each parse until it is fixed or removed. Each run counts towards the +``[dag_processor] dag_file_processor_timeout`` of the Dag file being parsed, and is also limited by +``[core] dagbag_import_timeout``, or by the ``get_dagbag_import_timeout`` policy, which is called with the +artifact's path. An artifact that fails slowly can keep the artifacts after it from being asked within the +timeout, and the stub tasks they serve then fail to import until it is fixed or removed, or the timeouts are +raised. Like any import error, a failed check keeps every Dag of the file from being scheduled until it is +fixed. A failed check reads like this: + +.. code-block:: text + + Stub tasks in etl.py do not match their task handlers: + - Dag 'etl', task 'load': no artifact in Dag bundle 'go-task-handlers' registers it + - Dag 'etl', task 'transform' ('etl' in Dag bundle 'go-task-handlers'): passes 2 arguments, the task handler takes 3 + .. _language-sdks/bundle-spec: Implementing a new compiled language SDK diff --git a/airflow-core/docs/authoring-and-scheduling/language-sdks/java.rst b/airflow-core/docs/authoring-and-scheduling/language-sdks/java.rst index 3f69772e4d92d..92baee8820e76 100644 --- a/airflow-core/docs/authoring-and-scheduling/language-sdks/java.rst +++ b/airflow-core/docs/authoring-and-scheduling/language-sdks/java.rst @@ -41,8 +41,8 @@ Prerequisites * JDK 11 or later is required on the machine that builds the Java project. A local Gradle installation is only needed to generate the Gradle Wrapper for a new project. -* JRE 11 or later must be available on the Airflow worker nodes. -* The compiled task JAR(s) and JVM dependencies must be accessible from the worker. +* JRE 11 or later must be available on the Airflow worker nodes and the Dag processor. +* The compiled task JAR(s) and JVM dependencies must be accessible from the worker and the Dag processor. * The ``apache-airflow-task-sdk`` package (installed with Airflow) provides the coordinator; no additional Python packages are needed. @@ -56,7 +56,8 @@ deployment process are the same. See :ref:`java-sdk/interface-api` for the inter The Python Dag source and the Java Gradle project are independent. They do not need to be in the same repository or have any particular relative filesystem layout. The Dag follows the deployment's normal Dag -delivery process; only the compiled Java bundle is deployed from the Gradle project to ``jars_root``. +delivery process; only the compiled Java bundle is deployed from the Gradle project, into a separate Dag +bundle that holds the JARs. Define the Python Dag ~~~~~~~~~~~~~~~~~~~~~ @@ -245,29 +246,45 @@ Deploy ``sales_pipeline.py`` separately through the deployment's normal Dag deli that process might sync it to ``${AIRFLOW_HOME}/dags/`` or package it in a Dag bundle; neither location is inside or relative to ``sales-pipeline-java/``. -Configure Airflow so the coordinator scans the parent JAR directory recursively and routes the ``java`` queue -to it. Add the following ``[sdk]`` section to the file selected by ``AIRFLOW_CONFIG`` (by default, -``${AIRFLOW_HOME}/airflow.cfg``), or set the equivalent ``AIRFLOW__SDK__*`` environment variables: +Configure Airflow to register the JAR directory as a Dag bundle, point the coordinator at it, and route the +``java`` queue to the coordinator. Add the following sections to the file selected by ``AIRFLOW_CONFIG`` (by +default, ``${AIRFLOW_HOME}/airflow.cfg``), or set the equivalent ``AIRFLOW__*`` environment variables: .. code-block:: ini + [dag_processor] + dag_bundle_config_list = [ + {"name": "dags-folder", "classpath": "airflow.dag_processing.bundles.local.LocalDagBundle", "kwargs": {}}, + { + "name": "java-task-handlers", + "classpath": "airflow.dag_processing.bundles.local.LocalDagBundle", + "kwargs": {"path": "/opt/airflow/jars/sales-pipeline"} + } + ] + [sdk] coordinators = { "java": { "classpath": "airflow.sdk.coordinators.java.JavaCoordinator", - "kwargs": {"jars_root": ["/opt/airflow/jars"]} + "kwargs": {"task_handler_bundle_name": "java-task-handlers"} } } queue_to_coordinator = {"java": "java"} ``java`` is a user-chosen coordinator name, not a reserved value. The value assigned to the queue in -``queue_to_coordinator`` must match a key in ``coordinators``. - -Restart the affected Airflow components after changing this configuration. The coordinator config and JARs -must be available wherever tasks execute. With ``CeleryExecutor``, that means the Celery workers; with -``LocalExecutor``, tasks run in subprocesses on the scheduler's host. The API server and Dag processor do not -need the JARs, while the Dag processor must receive ``sales_pipeline.py`` through the separate Dag delivery -process. +``queue_to_coordinator`` must match a key in ``coordinators``, and ``task_handler_bundle_name`` must match a +Dag bundle name in ``dag_bundle_config_list``. See :ref:`java-sdk/coordinator-config` for how JARs are +located. + +Restart the affected Airflow components after changing this configuration. The coordinator config, the JARs +and a JRE must be available wherever tasks execute and on the Dag processor. With ``CeleryExecutor``, tasks +execute on the Celery workers; with ``LocalExecutor``, they run in subprocesses on the scheduler's host. The +Dag processor checks the stub tasks of ``sales_pipeline.py`` against the task handlers the JARs register, so +it runs them too. The API server does not need any of it. Register the Dag bundle in +``[dag_processor] dag_bundle_config_list`` on every component, like your other Dag bundles: the worker and +the Dag processor resolve ``task_handler_bundle_name`` through it, and wherever the ``[sdk]`` config is read +it is rejected if the name is missing there. The Dag processor still receives ``sales_pipeline.py`` through +the separate Dag delivery process. After Airflow has parsed the Dag, trigger it from the UI or command line: @@ -456,10 +473,13 @@ so renaming one in an IDE never rebinds an input. A primitive parameter cannot hold ``null``, so the task fails with ``MissingXComException`` when its binding resolves to nothing; declare a boxed type (``Long``, ``Double``, …) to receive ``null`` instead. The method must declare as many data parameters as the call site bound: positions carry -the whole meaning of a flat binding, so any other count has already shifted them, and the task fails -rather than running on arguments it has mistaken for others. A parameter the Python call omitted -does not count towards that. Its default still arrives, but a method that does not declare it is -not reading shifted arguments, so the SDK drops it before comparing the two counts. +the whole meaning of a flat binding, so any other count has already shifted them. The Dag processor +reports it as an import error of the Python Dag file (see :ref:`language-sdks/dag-processor-checks`), +and a task that runs anyway fails rather than running on arguments it has mistaken for others. A +parameter the Python call omitted does not count towards that when the counts differ. Its default +still arrives, but a method that declares none of the omitted parameters is not reading shifted +arguments, so they are all dropped and the counts are compared again. A method that declares only +some of them still fails. Generic parameters are decoded element by element. Declare ``List`` and ``values.get(0)`` really is a ``Double``, even though the call site passed whole numbers and the wire carries them as @@ -726,7 +746,7 @@ configuration): "java-jdk17": { "classpath": "airflow.sdk.coordinators.java.JavaCoordinator", "kwargs": { - "jars_root": ["/opt/airflow/jars"], + "task_handler_bundle_name": "java-task-handlers", "jvm_args": ["-Djava.util.logging.config.file=/opt/airflow/logging.properties"] } } @@ -851,9 +871,10 @@ Then run: ./gradlew bundle -The ``build/bundle/`` directory contains all required JAR(s). Copy or mount it into the directory pointed to -by ``jars_root`` in the coordinator configuration. :class:`~airflow.sdk.coordinators.java.JavaCoordinator` -scans ``jars_root`` recursively and builds the classpath automatically. +The ``build/bundle/`` directory contains all required JAR(s). Copy or mount it into the Dag bundle named by +``task_handler_bundle_name`` in the coordinator configuration. +:class:`~airflow.sdk.coordinators.java.JavaCoordinator` scans that Dag bundle recursively and builds the +classpath automatically. .. note:: @@ -865,8 +886,7 @@ scans ``jars_root`` recursively and builds the classpath automatically. The plugin generates a fat JAR with the `Shadow `__ plugin by default. This is generally a good idea since you only deploy one JAR file to avoid dependency issues between projects. If this does not suit you, set ``fatJar = false`` in ``airflowBundle`` to produce thin JARs instead. The rest of the - process stays the same, but you will need to put all dependency JARs somewhere Airflow can find with - ``jars_root``. + process stays the same, but you will need to put all dependency JARs in the same Dag bundle. .. _java-sdk/build/maven: @@ -956,8 +976,8 @@ Then run: mvn package -The fat JAR is written to ``target/-.jar``. Copy it to the directory configured as -``jars_root`` in your coordinator. +The fat JAR is written to ``target/-.jar``. Copy it into the Dag bundle named by +``task_handler_bundle_name`` in your coordinator. **Option 2: thin JAR with separate dependencies** @@ -1017,8 +1037,8 @@ Then run: mvn package -``target/bundle/`` will contain the thin JAR and all runtime dependency JARs. Point ``jars_root`` at -this directory. +``target/bundle/`` will contain the thin JAR and all runtime dependency JARs. Copy or mount this +directory into the Dag bundle named by ``task_handler_bundle_name``. .. note:: @@ -1044,16 +1064,13 @@ All ``kwargs`` in the ``coordinators`` config entry are passed to the * - Parameter - Default - Description - * - ``jars_root`` - - *(optional)* - - One or more directories scanned recursively for ``.jar`` files. Accepts a string, - a path, or a list of strings/paths. When omitted, JARs are located through a Dag - bundle instead (see the note below). Explicitly setting this option to ``null`` or - an empty list is invalid. - * - ``dag_bundle_name`` - - *(auto: task's own bundle)* - - Name of a configured Dag bundle to load JARs from. Mutually exclusive with - ``jars_root``. + * - ``task_handler_bundle_name`` + - *(task's own Dag bundle)* + - Name of the Dag bundle scanned recursively for ``.jar`` files. It is used only by + mixed-language Dags, to locate the task handlers for the ``@task.stub`` tasks of a Python Dag; + Dags defined natively in a language SDK do not use it. It must be registered in + ``[dag_processor] dag_bundle_config_list``. It is checked when the ``[sdk]`` configuration is + loaded, so a typo fails there rather than on the first task. * - ``java_executable`` - ``"java"`` - Path to the ``java`` binary. Defaults to ``java`` on ``$PATH``. @@ -1062,10 +1079,9 @@ All ``kwargs`` in the ``coordinators`` config entry are passed to the - Extra JVM arguments such as ``["-Xmx1g", "-Dsome.property=value"]``. * - ``main_class`` - *(auto-detect)* - - Explicit entry-point class. If omitted, the coordinator scans for a JAR whose - manifest sets ``Main-Class`` — in ``jars_root`` when set, otherwise across the - resolved Dag bundle. If multiple executable JARs match the result is - non-deterministic; set ``main_class`` explicitly in that case. + - Explicit entry-point class. If omitted, the coordinator scans the Dag bundle for a JAR + whose manifest sets ``Main-Class``. If more than one JAR in that Dag bundle sets it, which + one runs is non-deterministic, so set ``main_class`` explicitly in that case. * - ``task_startup_timeout`` - ``10.0`` - Seconds to wait for the JVM subprocess to connect after launch. Increase this if your @@ -1073,22 +1089,26 @@ All ``kwargs`` in the ``coordinators`` config entry are passed to the .. note:: - **Locating JARs.** ``jars_root`` and ``dag_bundle_name`` are mutually exclusive, and both - are optional: + **Locating JARs.** The JARs for the ``@task.stub`` tasks of a Python Dag are read from a Dag bundle, so + they are delivered, refreshed and versioned by the same machinery as your Dags. - * Set ``jars_root`` to scan explicit filesystem directories you manage yourself. - * Set ``dag_bundle_name`` to load JARs from a configured Dag bundle, so they are delivered and - versioned through the same bundle machinery as your Dags. The task uses the version that bundle + * The expected layout is a separate Dag bundle for the JARs, named by ``task_handler_bundle_name``, + rather than the Dag bundle that holds your ``.py`` files. The task uses the version that Dag bundle is on when it starts, pinned for the whole task. - * Leave both unset (the default) to load JARs from the **task's own** Dag bundle, pinned to - the version the run was created with. + * If ``task_handler_bundle_name`` is unset, JARs are read from the **task's own** Dag bundle, at the + version the run uses: its pinned version, or the version current when the task starts if the run + is not pinned. + * Every JAR in the Dag bundle goes on one classpath, so all handlers in it share one set of + dependencies. To isolate conflicting dependency versions, put the handlers in a second Dag bundle + served by a second coordinator on its own queue. .. note:: The ``[sdk]`` configuration is read at startup, so changes to ``coordinators`` or ``queue_to_coordinator`` (for example adding ``jvm_args``) only take effect after you restart the - scheduler (or ``airflow standalone``). A rebuilt bundle JAR, by contrast, is picked up on the next - task launch without a restart, because a fresh JVM is spawned per task instance. + components that read it: the workers (the scheduler with ``LocalExecutor``), the Dag processor, or + ``airflow standalone``. A rebuilt bundle JAR, by contrast, is picked up on the next task launch without a + restart, because a fresh JVM is spawned per task instance. .. _java-sdk/java-executable: @@ -1110,7 +1130,7 @@ point ``java_executable`` at it explicitly: "java-jdk17": { "classpath": "airflow.sdk.coordinators.java.JavaCoordinator", "kwargs": { - "jars_root": ["/opt/airflow/jars"], + "task_handler_bundle_name": "java-task-handlers", "java_executable": "/opt/homebrew/opt/openjdk@17/bin/java" } } diff --git a/airflow-core/docs/authoring-and-scheduling/language-sdks/typescript.rst b/airflow-core/docs/authoring-and-scheduling/language-sdks/typescript.rst index 47070c581da15..45aac83f1e50f 100644 --- a/airflow-core/docs/authoring-and-scheduling/language-sdks/typescript.rst +++ b/airflow-core/docs/authoring-and-scheduling/language-sdks/typescript.rst @@ -48,9 +48,9 @@ The SDK is the ``apache-airflow-ts-sdk`` package (ESM-only). It is currently in Prerequisites ------------- -* Node.js 22 or later must be available on the Airflow worker nodes. +* Node.js 22 or later must be available on the Airflow worker nodes and the Dag processor. * The packed bundle (a single ``bundle.min.mjs`` file, see :ref:`typescript-sdk/build`) must be accessible - from the worker, under a directory the coordinator scans. + from the worker and the Dag processor, in the Dag bundle the coordinator scans. * The ``apache-airflow-task-sdk`` package (installed with Airflow) provides the coordinator; no additional Python packages are needed. * In the TypeScript project, install the ``apache-airflow-ts-sdk`` npm package to author task handlers: @@ -122,7 +122,9 @@ The ``dagId`` a handler binds must match the ``dag_id`` of the Python Dag, and t ``@task.stub`` function in that Dag, including any TaskGroup prefix. ``register`` takes any number of task handlers and Dags, and ``bundle.serve()`` serves exactly what is -registered, so a task left out is not part of the packed bundle and is marked removed at runtime. +registered, so a task left out is not part of the packed bundle. A stub task whose handler is left out +fails the import of its Python Dag file (see :ref:`language-sdks/dag-processor-checks`), and a task that +still reaches a bundle without its handler is marked removed at runtime. A second ``bundle.serve()`` call is rejected. Registering holds no sockets and starts nothing, so a unit test can build a bundle and dispatch a handler through ``bundle.getTaskHandler(dagId, taskId)`` without a coordinator runtime. @@ -209,34 +211,48 @@ check. Coordinator configuration ~~~~~~~~~~~~~~~~~~~~~~~~~ -Register the coordinator and route the queue to it under ``[sdk]`` in ``airflow.cfg`` (or the equivalent -``AIRFLOW__SDK__*`` environment variables): +Register a Dag bundle for the packed bundles, register the coordinator, and route the queue to it in +``airflow.cfg`` (or the equivalent ``AIRFLOW__*`` environment variables): .. code-block:: ini + [dag_processor] + dag_bundle_config_list = [ + {"name": "dags-folder", "classpath": "airflow.dag_processing.bundles.local.LocalDagBundle", "kwargs": {}}, + { + "name": "ts-task-handlers", + "classpath": "airflow.dag_processing.bundles.local.LocalDagBundle", + "kwargs": {"path": "/opt/airflow/ts-bundles"} + } + ] + [sdk] coordinators = { "ts": { "classpath": "airflow.sdk.coordinators.node.NodeCoordinator", - "kwargs": {"bundles_root": ["/opt/airflow/ts-bundles"]} + "kwargs": {"task_handler_bundle_name": "ts-task-handlers"} } } queue_to_coordinator = {"typescript": "ts"} -``bundles_root`` is one or more directories the coordinator scans for bundles; ``queue_to_coordinator`` -routes stub tasks with ``queue="typescript"`` to this coordinator. See -:ref:`typescript-sdk/coordinator-config` for the full list of accepted ``kwargs``. +``task_handler_bundle_name`` names the Dag bundle the coordinator scans for packed bundles; +``queue_to_coordinator`` routes stub tasks with ``queue="typescript"`` to this coordinator. See +:ref:`typescript-sdk/coordinator-config` for the full list of accepted ``kwargs`` and how the bundle is +located. There is no separate Node.js worker to run: the Airflow worker launches the bundle with ``node`` once per task instance. .. note:: - The coordinator runs inside the Airflow worker, so the ``[sdk]`` config (and the packed ``*.min.mjs`` - bundles in ``bundles_root``) only need to be present wherever tasks actually execute. With - ``CeleryExecutor``, setting them on the Celery workers is sufficient. With ``LocalExecutor``, tasks run - inside the scheduler process, so they must be present where the scheduler can read them. The API server - and Dag processor do not need them. + The ``[sdk]`` config, the packed ``*.min.mjs`` bundles and Node.js must be present wherever tasks execute + and on the Dag processor. With ``CeleryExecutor``, tasks execute on the Celery workers; with + ``LocalExecutor``, they run inside the scheduler process. The Dag processor checks the stub tasks of each + Python Dag against the task handlers the packed bundles register, so it runs them too. The API server + does not need any of it. Register the Dag bundle in ``[dag_processor] dag_bundle_config_list`` on every + component, like your other Dag bundles: the worker and the Dag processor resolve + ``task_handler_bundle_name`` through it, and wherever the ``[sdk]`` config is read it is rejected if the + name is missing there. .. _typescript-sdk/native-dag: @@ -499,7 +515,7 @@ that only supply utilities or types are not embedded. npx airflow-ts-pack src/main.ts --outdir dist Use ``--outdir `` to choose the output directory (default ``dist``), ``--outfile `` to name the -artifact exactly, which helps when one ``bundles_root`` holds several bundles, and ``--source `` to set +artifact exactly, which helps when one Dag bundle holds several bundles, and ``--source `` to set the source name displayed in the Airflow UI (default: the entry file's basename). ``--outdir`` and ``--outfile`` are mutually exclusive, and an ``--outfile`` name must end in ``.min.mjs`` so the coordinator can find it. @@ -507,12 +523,11 @@ can find it. Deploying ~~~~~~~~~ -Copy or mount the bundle into a directory listed in the coordinator's ``bundles_root``. -:class:`~airflow.sdk.coordinators.node.NodeCoordinator` searches the configured directories in order, -recursively, and launches the first integrity-verified ``*.min.mjs`` bundle whose metadata declares the task -instance's Dag. The artifact's name does not matter beyond that suffix, so one root can hold several bundles -and a Dag is routed to whichever declares it. If multiple bundles declare the same Dag, the first configured -root wins, and within a root the first in sorted path order. +Copy or mount the bundle into the Dag bundle named by the coordinator's ``task_handler_bundle_name``. +:class:`~airflow.sdk.coordinators.node.NodeCoordinator` searches that Dag bundle recursively and launches the +first integrity-verified ``*.min.mjs`` bundle whose metadata declares the task instance's Dag. The artifact's +name does not matter beyond that suffix, so one Dag bundle can hold several bundles and a Dag is routed to +whichever declares it. If multiple bundles declare the same Dag, the first in sorted path order wins. .. _typescript-sdk/coordinator-config: @@ -529,16 +544,13 @@ All ``kwargs`` in the ``coordinators`` config entry are passed to the * - Parameter - Default - Description - * - ``bundles_root`` - - *(optional)* - - One or more directories searched recursively, in order, for an integrity-verified ``*.min.mjs`` - bundle that declares the requested Dag. Accepts a string, a path, or a list of strings/paths. When - omitted, the bundle is located through a Dag bundle instead (see the note below). Explicitly setting - this option to ``null`` or an empty list is invalid. - * - ``dag_bundle_name`` - - *(auto: task's own bundle)* - - Name of a configured Dag bundle to load the ``*.min.mjs`` bundle from. Mutually exclusive with - ``bundles_root``. + * - ``task_handler_bundle_name`` + - *(task's own Dag bundle)* + - Name of the Dag bundle searched recursively for an integrity-verified ``*.min.mjs`` bundle that + declares the requested Dag. It is used only by mixed-language Dags, to locate the task handlers for + the ``@task.stub`` tasks of a Python Dag; Dags defined natively in a language SDK do not use it. It + must be registered in ``[dag_processor] dag_bundle_config_list``. It is checked when the ``[sdk]`` + configuration is loaded, so a typo fails there rather than on the first task. * - ``node_executable`` - ``"node"`` - Path to the ``node`` binary. Defaults to ``node`` on ``$PATH``. @@ -549,15 +561,15 @@ All ``kwargs`` in the ``coordinators`` config entry are passed to the .. note:: - **Locating the bundle.** ``bundles_root`` and ``dag_bundle_name`` are mutually exclusive, and both - are optional: + **Locating the bundle.** The packed bundles for the ``@task.stub`` tasks of a Python Dag are read from a + Dag bundle, so they are delivered, refreshed and versioned by the same machinery as your Dags. - * Set ``bundles_root`` to scan explicit filesystem directories you manage yourself. - * Set ``dag_bundle_name`` to load the bundle from a configured Dag bundle, so it is delivered - and versioned through the same bundle machinery as your Dags. The task uses the version that - bundle is on when it starts, pinned for the whole task. - * Leave both unset (the default) to load the bundle from the **task's own** Dag bundle, pinned to the - version the run was created with. + * The expected layout is a separate Dag bundle for the packed bundles, named by + ``task_handler_bundle_name``, rather than the Dag bundle that holds your ``.py`` files. The task uses + the version that Dag bundle is on when it starts, pinned for the whole task. + * If ``task_handler_bundle_name`` is unset, the bundle is read from the **task's own** Dag bundle, at + the version the run uses: its pinned version, or the version current when the task starts if the + run is not pinned. Limitations ----------- diff --git a/airflow-core/docs/core-concepts/dags.rst b/airflow-core/docs/core-concepts/dags.rst index 18acf9bbdec3c..3521f52fd1921 100644 --- a/airflow-core/docs/core-concepts/dags.rst +++ b/airflow-core/docs/core-concepts/dags.rst @@ -719,6 +719,8 @@ You can either do this all inside of the Dag bundle, with a standard filesystem package1/__init__.py package1/functions.py +The file name must end in ``.zip``. Other zip archives in the Dag bundle, such as JAR files, are not parsed as Dags. + Note that packaged Dags come with some caveats: * They cannot be used if you have pickling enabled for serialization diff --git a/airflow-core/docs/migrations-ref.rst b/airflow-core/docs/migrations-ref.rst index afd3ed01e8c32..c25f4acf8889f 100644 --- a/airflow-core/docs/migrations-ref.rst +++ b/airflow-core/docs/migrations-ref.rst @@ -39,7 +39,12 @@ Here's the list of all the Database Migrations that are executed via when you ru +-------------------------+------------------+-------------------+--------------------------------------------------------------+ | Revision ID | Revises ID | Airflow Version | Description | +=========================+==================+===================+==============================================================+ -| ``90e4d18ccadf`` (head) | ``e5a91c7f42b3`` | ``3.4.0`` | Add timetable_asset_gated to DagModel. | +| ``37d645374a9c`` (head) | ``f7ed13533d23`` | ``3.4.0`` | Add last_probed_at index to lang_sdk_task_handler_artifact. | ++-------------------------+------------------+-------------------+--------------------------------------------------------------+ +| ``f7ed13533d23`` | ``90e4d18ccadf`` | ``3.4.0`` | Add lang_sdk_task_handler_artifact and lang_sdk_task_handler | +| | | | tables. | ++-------------------------+------------------+-------------------+--------------------------------------------------------------+ +| ``90e4d18ccadf`` | ``e5a91c7f42b3`` | ``3.4.0`` | Add timetable_asset_gated to DagModel. | +-------------------------+------------------+-------------------+--------------------------------------------------------------+ | ``e5a91c7f42b3`` | ``ca8499dc1004`` | ``3.4.0`` | Add language column to dag_code. | +-------------------------+------------------+-------------------+--------------------------------------------------------------+ diff --git a/airflow-core/newsfragments/73970.significant.rst b/airflow-core/newsfragments/73970.significant.rst new file mode 100644 index 0000000000000..886d1d6ca0a8f --- /dev/null +++ b/airflow-core/newsfragments/73970.significant.rst @@ -0,0 +1,5 @@ +Language SDK coordinators load task handlers from a Dag bundle + +The experimental ``jars_root`` (Java) and ``executables_root`` (Go) coordinator kwargs are removed. +Put the artifacts in a Dag bundle registered in ``[dag_processor] dag_bundle_config_list`` and set ``task_handler_bundle_name`` in the coordinator's ``kwargs``, or leave it unset to use the task's own Dag bundle. +A config that still sets a removed kwarg fails with ``Cannot instantiate coordinator ''``. diff --git a/airflow-core/newsfragments/74033.significant.rst b/airflow-core/newsfragments/74033.significant.rst new file mode 100644 index 0000000000000..d02a2c021c974 --- /dev/null +++ b/airflow-core/newsfragments/74033.significant.rst @@ -0,0 +1,3 @@ +Only ``.zip`` archives are discovered as packaged Dags + +Dag discovery now lists a zip archive as a packaged Dag only when its file name ends in ``.zip``. Before, every zip archive in a Dag bundle was parsed whatever its name, including JAR files and archives with no suffix. A packaged Dag stored under another name is no longer found: rename it to end in ``.zip``. ``.py`` files are listed as before. diff --git a/airflow-core/src/airflow/config_templates/config.yml b/airflow-core/src/airflow/config_templates/config.yml index 34849d5eb5e13..bd9be70756a99 100644 --- a/airflow-core/src/airflow/config_templates/config.yml +++ b/airflow-core/src/airflow/config_templates/config.yml @@ -2191,6 +2191,19 @@ sdk: example, two ``JavaCoordinator`` instances pinned to different JDK versions). + The language SDK coordinators read the task handlers for the + ``@task.stub`` tasks of a Python Dag from the Dag bundle named by the + ``task_handler_bundle_name`` kwarg. Dags defined natively in a language + SDK do not use it. It must name a bundle in + ``[dag_processor] dag_bundle_config_list`` and is checked when this + option is loaded. When it is unset, a task reads its handlers from its + own Dag bundle. + + The Dag processor needs this option too, with the language runtime the + coordinator starts and the files of that Dag bundle, because it checks + the stub tasks of each Python Dag against the task handlers those files + register. The API server does not need it. + An entry may also carry an optional ``extra`` object for additional information associated with the coordinator that the coordinator itself does not receive; other components read it as needed. For example, @@ -2205,7 +2218,7 @@ sdk: "jdk-17": { "classpath": "airflow.sdk.coordinators.java.JavaCoordinator", "kwargs": { - "jars_root": ["/opt/airflow/java-bundles"], + "task_handler_bundle_name": "java-task-handlers", "java_executable": "/usr/lib/jvm/java-17-openjdk/bin/java", "jvm_args": ["-Xmx1024m"] }, @@ -2218,7 +2231,7 @@ sdk: "go-sdk": { "classpath": "airflow.sdk.coordinators.executable.ExecutableCoordinator", "kwargs": { - "executables_root": ["/opt/airflow/executable-bundles"] + "task_handler_bundle_name": "go-task-handlers" } } } @@ -2714,7 +2727,8 @@ scheduler: description: | How often (in seconds) to check for stale DAGs (DAGs which are no longer present in the expected files) which should be deactivated, as well as assets that are no longer - referenced and should be marked as orphaned. + referenced and should be marked as orphaned. The same check deletes the Lang-SDK task + handler artifacts that no stub task uses any more. version_added: 2.5.0 type: integer example: ~ @@ -3344,7 +3358,12 @@ dag_processor: default: "50" dag_file_processor_timeout: description: | - How long before timing out a DagFileProcessor, which processes a dag file + How long before timing out a DagFileProcessor, which processes a Dag file + + It also bounds running the language SDK artifacts that the stub tasks of the + Dag file are checked against. Each run is limited by ``[core] dagbag_import_timeout`` too, and an + artifact that gives no answer in time leaves the stub tasks it serves failing to import until a + later parse gets its answer. version_added: ~ type: integer example: ~ diff --git a/airflow-core/src/airflow/dag_processing/bundles/manager.py b/airflow-core/src/airflow/dag_processing/bundles/manager.py index 12b3b50fbca21..9c8e6816431a2 100644 --- a/airflow-core/src/airflow/dag_processing/bundles/manager.py +++ b/airflow-core/src/airflow/dag_processing/bundles/manager.py @@ -693,6 +693,19 @@ def get_bundle( name=name, version=version, version_data=version_data, **cfg_bundle.kwargs ) + def get_bundle_team_name(self, name: str) -> str | None: + """ + Return the team that owns the Dag bundle *name*, or ``None`` when it is not team scoped. + + An empty team name is no team. + + :raises ValueError: when *name* is not a configured Dag bundle. + """ + cfg_bundle = self._bundle_config.get(name) + if not cfg_bundle: + raise ValueError(f"Requested bundle '{name}' is not configured.") + return cfg_bundle.team_name or None + @classmethod def is_bundle_configured(cls, name: str) -> bool: """ diff --git a/airflow-core/src/airflow/dag_processing/collection.py b/airflow-core/src/airflow/dag_processing/collection.py index ec8889db31845..d9bb2c3954e43 100644 --- a/airflow-core/src/airflow/dag_processing/collection.py +++ b/airflow-core/src/airflow/dag_processing/collection.py @@ -28,11 +28,12 @@ from __future__ import annotations import traceback +from collections import defaultdict from typing import TYPE_CHECKING, Any, NamedTuple, TypeVar import structlog from sqlalchemy import delete, false, func, insert, or_, select, tuple_, update -from sqlalchemy.exc import OperationalError +from sqlalchemy.exc import DataError, IntegrityError, OperationalError from sqlalchemy.orm import joinedload, load_only from airflow._shared.timezones.timezone import utcnow @@ -54,6 +55,11 @@ from airflow.models.dagrun import DagRun from airflow.models.dagwarning import DagWarning, DagWarningType from airflow.models.errors import ParseImportError +from airflow.models.lang_sdk_task_handler import ( + LangSDKTaskHandler, + LangSDKTaskHandlerArtifact, + compute_fileloc_hash, +) from airflow.models.serialized_dag import SerializedDagModel from airflow.models.trigger import Trigger from airflow.plugins_manager import get_scheduling_class_teams @@ -68,15 +74,17 @@ from airflow.serialization.serialized_objects import BaseSerialization, LazyDeserializedDAG from airflow.triggers.base import BaseEventTrigger from airflow.utils.retries import MAX_DB_RETRIES, run_with_db_retries -from airflow.utils.sqlalchemy import get_dialect_name, with_row_locks +from airflow.utils.sqlalchemy import build_upsert_stmt, get_dialect_name, with_row_locks from airflow.utils.types import DagRunType if TYPE_CHECKING: from collections.abc import Collection, Iterable, Iterator + from uuid import UUID from sqlalchemy.orm import Session from sqlalchemy.sql import Select + from airflow.dag_processing.processor import TaskHandlerArtifact, TaskHandlerBinding from airflow.models.serialized_dag import DagWriteMetadata from airflow.sdk.importers import DagSourceCode # noqa: SDK001 from airflow.typing_compat import Self, Unpack @@ -581,6 +589,211 @@ def _reject_other_teams_plugin_classes( return accepted +_ArtifactKey = tuple[str, str] +"""``(bundle_name, relative_fileloc_hash)``, the unique key of ``lang_sdk_task_handler_artifact``.""" + + +def record_probed_task_handler_artifacts( + artifacts: Collection[TaskHandlerArtifact], *, bundle_names: Collection[str], session: Session +) -> None: + """Upsert each probed artifact's fingerprint and answer; those outside *bundle_names* are dropped with a warning.""" + if outside := sorted( + { + f"{artifact.bundle_name}/{artifact.relative_fileloc}" + for artifact in artifacts + if artifact.bundle_name not in bundle_names + } + ): + log.warning("Ignoring probed task handler artifacts outside the Dag file's scope", artifacts=outside) + by_key: dict[_ArtifactKey, TaskHandlerArtifact] = {} + for artifact in artifacts: + if artifact.bundle_name in bundle_names: + by_key.setdefault( + (artifact.bundle_name, compute_fileloc_hash(artifact.relative_fileloc)), artifact + ) + if not by_key: + return + dialect = get_dialect_name(session) + now = utcnow() + # In key order, so two transactions recording the same artifacts lock them in the same order. + for key in sorted(by_key): + artifact = by_key[key] + probe = { + "size_bytes": artifact.size_bytes, + "cache_digest": artifact.cache_digest, + "task_handlers": { + dag_id: [declaration.model_dump(mode="json") for declaration in declarations] + for dag_id, declarations in artifact.task_handlers.items() + }, + "last_probed_at": now, + } + session.execute( + build_upsert_stmt( + dialect, + LangSDKTaskHandlerArtifact, + conflict_cols=["bundle_name", "relative_fileloc_hash"], + values={ + "bundle_name": artifact.bundle_name, + "relative_fileloc": artifact.relative_fileloc, + # Upserts skip the model's @validates hook, which keeps the hash in sync on ORM writes. + "relative_fileloc_hash": key[1], + **probe, + }, + update_fields=probe, + ) + ) + + +def _get_artifact_key(binding: TaskHandlerBinding) -> _ArtifactKey: + return binding.artifact_bundle_name, compute_fileloc_hash(binding.artifact_rel_path) + + +def _find_task_handler_artifacts( + keys: Collection[_ArtifactKey], *, session: Session +) -> dict[_ArtifactKey, UUID]: + """Return the ids of the recorded artifacts among *keys*, share-locked until commit.""" + # The orphan sweep skips locked rows, so it cannot delete an artifact between this read and the + # insert of the handler rows that reference it. size_bytes is not in the unique index: reading it + # makes MySQL lock the row itself, which the sweep checks, and not only the index entry. + query = with_row_locks( + select( + LangSDKTaskHandlerArtifact.bundle_name, + LangSDKTaskHandlerArtifact.relative_fileloc_hash, + LangSDKTaskHandlerArtifact.id, + LangSDKTaskHandlerArtifact.size_bytes, + ) + .where( + tuple_( + LangSDKTaskHandlerArtifact.bundle_name, LangSDKTaskHandlerArtifact.relative_fileloc_hash + ).in_(list(keys)) + ) + .order_by(LangSDKTaskHandlerArtifact.bundle_name, LangSDKTaskHandlerArtifact.relative_fileloc_hash), + session, + read=True, + key_share=True, + ) + return { + (bundle_name, fileloc_hash): artifact_id + for bundle_name, fileloc_hash, artifact_id, _ in session.execute(query) + } + + +def _sync_task_handlers( + bindings: dict[tuple[str, str], TaskHandlerBinding], + artifact_ids: dict[_ArtifactKey, UUID], + *, + dag_ids: Collection[str], + dag_bundle_name: str, + dag_relative_fileloc: str, + session: Session, +) -> None: + desired = { + key: { + "dag_id": binding.dag_id, + "task_id": binding.task_id, + "artifact_id": artifact_ids[_get_artifact_key(binding)], + "dag_bundle_name": dag_bundle_name, + "dag_relative_fileloc": dag_relative_fileloc, + # Bulk writes skip the model's @validates hook, which keeps the hash in sync on ORM writes. + "dag_relative_fileloc_hash": compute_fileloc_hash(dag_relative_fileloc), + } + for key, binding in bindings.items() + } + recorded = { + (row.dag_id, row.task_id): row._asdict() + for row in session.execute( + select( + LangSDKTaskHandler.dag_id, + LangSDKTaskHandler.task_id, + LangSDKTaskHandler.artifact_id, + LangSDKTaskHandler.dag_bundle_name, + LangSDKTaskHandler.dag_relative_fileloc, + LangSDKTaskHandler.dag_relative_fileloc_hash, + ).where(LangSDKTaskHandler.dag_id.in_(dag_ids)) + ) + } + if gone := sorted(recorded.keys() - desired.keys()): + session.execute( + delete(LangSDKTaskHandler) + .where(tuple_(LangSDKTaskHandler.dag_id, LangSDKTaskHandler.task_id).in_(gone)) + .execution_options(synchronize_session=False) + ) + if added := [desired[key] for key in sorted(desired.keys() - recorded.keys())]: + session.execute(insert(LangSDKTaskHandler), added) + if changed := [ + desired[key] for key in sorted(desired.keys() & recorded.keys()) if desired[key] != recorded[key] + ]: + session.execute(update(LangSDKTaskHandler), changed) + + +def _reconcile_task_handler_bindings( + bindings: Iterable[TaskHandlerBinding], + *, + dag_ids: Collection[str], + rejected_dag_ids: Collection[str], + artifact_bundle_names: Collection[str], + dag_bundle_name: str, + dag_relative_fileloc: str, + session: Session, +) -> None: + """ + Make the recorded task handler bindings of each Dag in *dag_ids* match *bindings*. + + Rows are keyed by Dag id, not by file. Rows of other Dags are left alone, including those of a Dag + that moved to another file, which that file's parse reconciles. A Dag bound to an artifact outside + *artifact_bundle_names*, or to one with no recorded row, keeps its rows as they are. Artifact rows are + only read: the probed-artifact write records them, and the Dag processor's orphan sweep deletes them. + Bindings of *rejected_dag_ids*, Dags the parse produced but that are not persisted, are dropped + quietly: their import error explains why. + """ + by_task: dict[tuple[str, str], list[TaskHandlerBinding]] = defaultdict(list) + unknown_dag_ids: set[str] = set() + for binding in bindings: + if binding.dag_id in dag_ids: + by_task[(binding.dag_id, binding.task_id)].append(binding) + elif binding.dag_id not in rejected_dag_ids: + unknown_dag_ids.add(binding.dag_id) + if unknown_dag_ids: + log.warning( + "Ignoring task handler bindings of Dags not in the parse result", dag_ids=sorted(unknown_dag_ids) + ) + # The parse reports a task bound twice as an import error; this only keeps the rows as they are. + if conflicting := {dag_id for (dag_id, _), group in by_task.items() if len(group) > 1}: + log.warning( + "Ignoring task handler bindings of Dags that bind a task twice", dag_ids=sorted(conflicting) + ) + resolved = {key: group[0] for key, group in by_task.items() if key[0] not in conflicting} + if out_of_scope := { + binding.dag_id + for binding in resolved.values() + if binding.artifact_bundle_name not in artifact_bundle_names + }: + log.warning( + "Ignoring task handler bindings of Dags bound to an artifact outside their scope", + dag_ids=sorted(out_of_scope), + ) + in_scope = [binding for key, binding in resolved.items() if key[0] not in out_of_scope] + keys = {_get_artifact_key(binding) for binding in in_scope} + artifact_ids = _find_task_handler_artifacts(keys, session=session) if keys else {} + # A concurrent orphan sweep or a failed probed-artifact write leaves no row; the next parse probes again. + if unrecorded := { + binding.dag_id for binding in in_scope if _get_artifact_key(binding) not in artifact_ids + }: + log.warning( + "Ignoring task handler bindings of Dags bound to an unrecorded artifact", + dag_ids=sorted(unrecorded), + ) + unchanged = conflicting | out_of_scope | unrecorded + _sync_task_handlers( + {key: binding for key, binding in resolved.items() if key[0] not in unchanged}, + artifact_ids, + dag_ids=[dag_id for dag_id in dag_ids if dag_id not in unchanged], + dag_bundle_name=dag_bundle_name, + dag_relative_fileloc=dag_relative_fileloc, + session=session, + ) + + def update_dag_parsing_results_in_db( bundle_name: str, bundle_version: str | None, @@ -598,6 +811,9 @@ def update_dag_parsing_results_in_db( ), files_parsed: set[tuple[str, str]] | None = None, dag_source_codes: dict[str, DagSourceCode] | None = None, + relative_fileloc: str | None = None, + task_handler_bindings: list[TaskHandlerBinding] | None = None, + task_handler_artifact_bundles: Collection[str] | None = None, ): """ Update everything to do with DAG parsing in the DB. @@ -609,6 +825,7 @@ def update_dag_parsing_results_in_db( - ParseImportError (including with any errors as a result of serialization, not just parsing) - DagWarning - DAG Permissions + - LangSDKTaskHandler, when ``task_handler_bindings`` is given This function will not remove any rows for dags not passed in. It will remove parse errors and warnings from dags/dag files that are passed in. In order words, if a DAG is passed in with a fileloc of `a.py` @@ -621,8 +838,19 @@ def update_dag_parsing_results_in_db( import errors are cleared for files that were parsed but no longer contain DAGs. :param dag_source_codes: Source code read by the Dag importers, keyed by Dag fileloc. Dags without an entry have their source read from ``fileloc``. + :param relative_fileloc: The parsed file, relative to its bundle. Required with ``task_handler_bindings``. + :param task_handler_bindings: The stub-task bindings of every Dag in ``dags``. ``None`` leaves the + recorded bindings as they are; a list replaces the recorded bindings of each Dag in ``dags``. If + the database rejects that write, the recorded bindings stay and the rest is still written. + :param task_handler_artifact_bundles: The bundles whose artifacts the bindings may reference. Required + with ``task_handler_bindings``; a Dag bound to an artifact outside them keeps its recorded bindings. """ + if task_handler_bindings is not None and relative_fileloc is None: + raise ValueError("relative_fileloc is required with task_handler_bindings") + if task_handler_bindings is not None and task_handler_artifact_bundles is None: + raise ValueError("task_handler_artifact_bundles is required with task_handler_bindings") accepted = _reject_other_teams_plugin_classes(bundle_name, dags, import_errors, session=session) + rejected_ids: set[str] = set() if len(accepted) != len(dags): # A rejected Dag may have no ``dag`` row yet, and dag_warning has a foreign key to it. rejected_ids = {dag.dag_id for dag in dags} - {dag.dag_id for dag in accepted} @@ -653,6 +881,35 @@ def update_dag_parsing_results_in_db( SerializedDAG.bulk_write_to_db( bundle_name, bundle_version, dags, parse_duration, session=session ) + # After the Dag rows, which the handler rows reference, and before the serialized Dags. + if ( + task_handler_bindings is not None + and relative_fileloc is not None + and task_handler_artifact_bundles is not None + and dags + ): + # A savepoint, so bindings the database rejects cannot discard the rest of the result. + savepoint = session.begin_nested() + try: + _reconcile_task_handler_bindings( + task_handler_bindings, + dag_ids={dag.dag_id for dag in dags}, + rejected_dag_ids=rejected_ids, + artifact_bundle_names=task_handler_artifact_bundles, + dag_bundle_name=bundle_name, + dag_relative_fileloc=relative_fileloc, + session=session, + ) + except (IntegrityError, DataError): + savepoint.rollback() + log.warning( + "Failed to record task handler bindings; the rest of the parse result is kept", + bundle_name=bundle_name, + relative_fileloc=relative_fileloc, + exc_info=True, + ) + else: + savepoint.commit() # Bulk prefetch metadata for all DAGs to avoid the standard per-DAG # metadata lookups in write_dag. This replaces the update-interval, # hash, and version queries with 2 bulk queries total; DAGs with diff --git a/airflow-core/src/airflow/dag_processing/manager.py b/airflow-core/src/airflow/dag_processing/manager.py index 8c4fd0588eb87..1d8dfa6417bac 100644 --- a/airflow-core/src/airflow/dag_processing/manager.py +++ b/airflow-core/src/airflow/dag_processing/manager.py @@ -39,7 +39,8 @@ import attrs import structlog -from sqlalchemy import select, update +from pydantic import ValidationError +from sqlalchemy import delete, exists, select, update from sqlalchemy.exc import OperationalError from sqlalchemy.orm import load_only from tabulate import tabulate @@ -54,8 +55,15 @@ unpack_bundle_version, ) from airflow.dag_processing.bundles.manager import DagBundlesManager -from airflow.dag_processing.collection import update_dag_parsing_results_in_db -from airflow.dag_processing.processor import DagFileParsingResult, DagFileProcessorProcess +from airflow.dag_processing.collection import ( + record_probed_task_handler_artifacts, + update_dag_parsing_results_in_db, +) +from airflow.dag_processing.processor import ( + DagFileParsingResult, + DagFileProcessorProcess, + TaskHandlerArtifact, +) from airflow.models.asset import remove_references_to_deleted_dags from airflow.models.dag import DagModel from airflow.models.dagbag import DagPriorityParsingRequest @@ -63,8 +71,10 @@ from airflow.models.dagwarning import DagWarning from airflow.models.db_callback_request import DbCallbackRequest from airflow.models.errors import ParseImportError +from airflow.models.lang_sdk_task_handler import LangSDKTaskHandler, LangSDKTaskHandlerArtifact from airflow.observability.metrics import stats_utils from airflow.sdk import SecretCache +from airflow.sdk.execution_time.coordinator import get_coordinator_manager # noqa: SDK001 from airflow.sdk.log import init_log_file, logging_processors from airflow.typing_compat import assert_never from airflow.utils.file import list_py_file_paths, might_contain_dag @@ -74,7 +84,7 @@ from airflow.utils.process_utils import ( kill_child_processes_by_pids, ) -from airflow.utils.retries import retry_db_transaction +from airflow.utils.retries import retry_db_transaction, run_with_db_retries from airflow.utils.session import NEW_SESSION, create_session, provide_session from airflow.utils.sqlalchemy import ( is_lock_not_available_error, @@ -87,6 +97,7 @@ from collections.abc import Callable, Collection, Iterable, Iterator, Sequence from socket import socket + from sqlalchemy.engine import CursorResult from sqlalchemy.orm import Session from sqlalchemy.sql import Select @@ -151,6 +162,24 @@ def normalized_file_path_for_stats(self) -> str: return normalize_name_for_stats(str(self.rel_path), log_warning=False) +class _TaskHandlerBundles(NamedTuple): + """Where the coordinators that queues route to read task-handler artifacts from.""" + + named: frozenset[str] + """Bundles named by a coordinator's ``task_handler_bundle_name``.""" + + include_dag_bundles: bool + """Whether a coordinator without ``task_handler_bundle_name`` reads the task's own Dag bundle.""" + + +_TASK_HANDLER_ARTIFACT_SWEEP_BATCH_SIZE = 1000 + +# A probed artifact that serves no stub task can have a row too, so that parses do not probe it again, +# and such a row is unreferenced from the start. An unreferenced row is therefore kept for this many +# sweep intervals after its last probe, long enough for the Dag files parsed next to reuse it. +_TASK_HANDLER_ARTIFACT_SWEEP_GRACE_INTERVALS = 5 + + def _config_int_factory(section: str, key: str): return functools.partial(conf.getint, section, key) @@ -241,6 +270,7 @@ class DagFileProcessorManager(LoggingMixin): ) _last_deactivate_stale_dags_time: float = attrs.field(default=0, init=False) + _last_task_handler_artifact_sweep_time: float = attrs.field(default=0, init=False) _last_stale_bundle_cleanup_time: float = attrs.field(default=0, init=False) print_stats_interval: float = attrs.field( factory=_config_int_factory("dag_processor", "print_stats_interval") @@ -431,6 +461,61 @@ def _scan_stale_dags(self): self.deactivate_stale_dags(last_parsed=last_parsed) self._last_deactivate_stale_dags_time = time.monotonic() + def _sweep_task_handler_artifacts(self) -> None: + now = time.monotonic() + if now - self._last_task_handler_artifact_sweep_time < self.parsing_cleanup_interval: + return + try: + self.delete_unreferenced_task_handler_artifacts() + except Exception: + self.log.exception("Error deleting unreferenced task handler artifacts") + finally: + self._last_task_handler_artifact_sweep_time = now + + def delete_unreferenced_task_handler_artifacts(self) -> int: + """ + Delete the task handler artifact rows that no stub task references and no recent parse probed. + + Deletes in batches, each committed on its own, and returns how many rows it deleted. Default + implementation writes to the metadata DB; override to do it through an API. + """ + cutoff = timezone.utcnow() - timedelta( + seconds=self.parsing_cleanup_interval * _TASK_HANDLER_ARTIFACT_SWEEP_GRACE_INTERVALS + ) + sweepable = ( + LangSDKTaskHandlerArtifact.last_probed_at < cutoff, + ~exists().where(LangSDKTaskHandler.artifact_id == LangSDKTaskHandlerArtifact.id), + ) + deleted = 0 + while True: + with create_session() as session: + # SKIP LOCKED passes over the artifacts a concurrent parse holds while it binds them. + ids = session.scalars( + with_row_locks( + select(LangSDKTaskHandlerArtifact.id) + .where(*sweepable) + .order_by(LangSDKTaskHandlerArtifact.last_probed_at) + .limit(_TASK_HANDLER_ARTIFACT_SWEEP_BATCH_SIZE), + session, + of=LangSDKTaskHandlerArtifact, + skip_locked=True, + key_share=False, + ) + ).all() + if ids: + # Checked again for handler rows a parse committed since the select. + result = session.execute( + delete(LangSDKTaskHandlerArtifact) + .where(LangSDKTaskHandlerArtifact.id.in_(ids), *sweepable) + .execution_options(synchronize_session=False) + ) + deleted += cast("CursorResult", result).rowcount + if len(ids) < _TASK_HANDLER_ARTIFACT_SWEEP_BATCH_SIZE: + break + if deleted: + self.log.info("Deleted %i unreferenced task handler artifacts.", deleted) + return deleted + def _cleanup_stale_bundle_versions(self): if self.stale_bundle_cleanup_interval <= 0: return @@ -584,6 +669,7 @@ def _run_parsing_loop(self): for callback in self.fetch_callbacks(): self._add_callback_to_queue(callback) self._scan_stale_dags() + self._sweep_task_handler_artifacts() self._cleanup_stale_bundle_versions() self.purge_inactive_dag_warnings() @@ -1301,6 +1387,19 @@ def handle_parsing_result( ) if proc.parsing_result is not None: + if proc.parsing_result.probed_artifacts: + try: + self.persist_probed_task_handler_artifacts( + bundle_name=file.bundle_name, artifacts=proc.parsing_result.probed_artifacts + ) + except Exception: + # The parse result is still persisted; a binding whose artifact has no row leaves its + # Dag unchanged, and the next parse probes the artifact again. + self.log.exception( + "Failed to record probed task handler artifacts", + bundle_name=file.bundle_name, + relative_fileloc=str(file.rel_path), + ) try: self.persist_parsing_result( bundle_name=file.bundle_name, @@ -1373,8 +1472,31 @@ def persist_parsing_result( session=session, files_parsed=files_parsed, dag_source_codes=parsing_result.dag_source_codes, + relative_fileloc=relative_fileloc, + task_handler_bindings=parsing_result.task_handler_bindings, + task_handler_artifact_bundles=( + None + if parsing_result.task_handler_bindings is None + else self._get_task_handler_artifact_bundle_names(bundle_name) + ), ) + def persist_probed_task_handler_artifacts( + self, *, bundle_name: str, artifacts: Sequence[TaskHandlerArtifact] + ) -> None: + """ + Record the artifacts a Dag file's parse probed, with their answers, in a transaction of their own. + + *bundle_name* is the Dag file's bundle. Default implementation writes to the metadata DB; override + to do it through an API. + """ + artifact_bundle_names = self._get_task_handler_artifact_bundle_names(bundle_name) + for attempt in run_with_db_retries(logger=self.log): + with attempt, create_session() as session: + record_probed_task_handler_artifacts( + artifacts, bundle_names=artifact_bundle_names, session=session + ) + def _collect_results(self): finished = [] for file, proc in self._processors.items(): @@ -1444,7 +1566,97 @@ def client(self) -> Client: client.base_url = "http://in-process.invalid./" return client - def _create_process(self, dag_file: DagFileInfo) -> DagFileProcessorProcess: + @functools.cached_property + def _task_handler_bundles(self) -> _TaskHandlerBundles: + try: + bundle_names = get_coordinator_manager().get_task_handler_bundle_names().values() + except Exception: + # A Lang-SDK configuration error must not stop Python-only parsing. + self.log.exception( + "Cannot read [sdk] coordinators; Dag files are parsed without known task-handler artifacts" + ) + return _TaskHandlerBundles(named=frozenset(), include_dag_bundles=False) + return _TaskHandlerBundles( + named=frozenset(name for name in bundle_names if name is not None), + include_dag_bundles=None in bundle_names, + ) + + def get_known_task_handler_artifacts( + self, bundle_names: Collection[str] + ) -> dict[str, list[TaskHandlerArtifact]]: + """ + Return the recorded task-handler artifacts in *bundle_names*, keyed by bundle name. + + Default implementation reads from the metadata DB; override to source them from an API. + """ + return self._get_known_task_handler_artifacts_from_db(bundle_names) + + @provide_session + def _get_known_task_handler_artifacts_from_db( + self, bundle_names: Collection[str], *, session: Session = NEW_SESSION + ) -> dict[str, list[TaskHandlerArtifact]]: + rows = session.execute( + select( + LangSDKTaskHandlerArtifact.bundle_name, + LangSDKTaskHandlerArtifact.relative_fileloc, + LangSDKTaskHandlerArtifact.size_bytes, + LangSDKTaskHandlerArtifact.cache_digest, + LangSDKTaskHandlerArtifact.task_handlers, + ) + .where(LangSDKTaskHandlerArtifact.bundle_name.in_(bundle_names)) + .order_by(LangSDKTaskHandlerArtifact.bundle_name, LangSDKTaskHandlerArtifact.relative_fileloc) + ) + known: dict[str, list[TaskHandlerArtifact]] = defaultdict(list) + for row in rows: + try: + artifact = TaskHandlerArtifact.model_validate(row._asdict()) + except ValidationError as exc: + # Left out, the artifact is probed again and its new answer replaces the row. + self.log.warning( + "Ignoring a recorded task-handler artifact that fails validation", + bundle_name=row.bundle_name, + relative_fileloc=row.relative_fileloc, + error=str(exc), + ) + continue + known[row.bundle_name].append(artifact) + return dict(known) + + def _query_known_task_handler_artifacts(self) -> dict[str, list[TaskHandlerArtifact]]: + """Return the recorded artifacts of every bundle that can hold task handlers for this manager's Dag files.""" + bundle_names = set(self._task_handler_bundles.named) + if self._task_handler_bundles.include_dag_bundles: + bundle_names.update(bundle.name for bundle in self._dag_bundles) + if not bundle_names: + return {} + try: + return self.get_known_task_handler_artifacts(bundle_names) + except Exception: + # Without known artifacts the children probe their candidates, so parsing keeps running. + self.log.exception( + "Cannot read the recorded task-handler artifacts; Dag files are parsed without them" + ) + return {} + + def _get_task_handler_artifact_bundle_names(self, dag_bundle_name: str) -> frozenset[str]: + """Return the bundles whose task-handler artifacts a Dag file in *dag_bundle_name* may read and record.""" + bundles = self._task_handler_bundles + teams = self._get_team_names(bundles.named | {dag_bundle_name}) + scope = {name for name in bundles.named if teams.get(name) == teams.get(dag_bundle_name)} + if bundles.include_dag_bundles: + scope.add(dag_bundle_name) + return frozenset(scope) + + def _select_known_task_handler_artifacts( + self, known_artifacts: dict[str, list[TaskHandlerArtifact]], dag_bundle_name: str + ) -> list[TaskHandlerArtifact]: + """Return the artifacts a Dag file in *dag_bundle_name* may resolve its stub tasks against.""" + bundle_names = self._get_task_handler_artifact_bundle_names(dag_bundle_name) + return [artifact for name in sorted(bundle_names) for artifact in known_artifacts.get(name, ())] + + def _create_process( + self, dag_file: DagFileInfo, *, known_artifacts: Sequence[TaskHandlerArtifact] = () + ) -> DagFileProcessorProcess: id = uuid7() callback_to_execute_for_file = self._callback_to_execute.pop(dag_file, []) @@ -1457,6 +1669,7 @@ def _create_process(self, dag_file: DagFileInfo) -> DagFileProcessorProcess: bundle_name=dag_file.bundle_name, dag_file_rel_path=str(dag_file.rel_path), callbacks=callback_to_execute_for_file, + known_artifacts=known_artifacts, selector=self.selector, logger=logger, logger_filehandle=logger_filehandle, @@ -1467,6 +1680,9 @@ def _create_process(self, dag_file: DagFileInfo) -> DagFileProcessorProcess: def _start_new_processes(self): """Start more processors if we have enough slots and files to process.""" bundle_to_team = self._get_team_names({file.bundle_name for file in self._file_queue}) + # Read once per loop, and only when a child starts: rows persisted by earlier children are + # then visible to later ones, which a cache held across a whole pass would hide. + known_artifacts: dict[str, list[TaskHandlerArtifact]] | None = None while self._parallelism > len(self._processors) and self._file_queue: file, _ = self._file_queue.popitem(last=False) @@ -1474,7 +1690,12 @@ def _start_new_processes(self): if file in self._processors: continue - processor = self._create_process(file) + if known_artifacts is None: + known_artifacts = self._query_known_task_handler_artifacts() + processor = self._create_process( + file, + known_artifacts=self._select_known_task_handler_artifacts(known_artifacts, file.bundle_name), + ) stats.incr( "dag_processing.processes", tags=prune_dict( diff --git a/airflow-core/src/airflow/dag_processing/processor.py b/airflow-core/src/airflow/dag_processing/processor.py index 80e6d7161dec2..0510a4a60e7e9 100644 --- a/airflow-core/src/airflow/dag_processing/processor.py +++ b/airflow-core/src/airflow/dag_processing/processor.py @@ -20,15 +20,18 @@ import importlib import logging import os +import time import traceback from collections.abc import Callable, Sequence from pathlib import Path -from typing import TYPE_CHECKING, Annotated, BinaryIO, ClassVar, Literal +from typing import TYPE_CHECKING, Annotated, Any, BinaryIO, ClassVar, Literal import attrs +import psutil from pydantic import BaseModel, Field, TypeAdapter from airflow._shared.observability.metrics import stats +from airflow.api_fastapi.execution_api.datamodels.task_arg_binding import ArgValueSchema # noqa: TC001 from airflow.callbacks.callback_requests import ( CallbackRequest, DagCallbackRequest, @@ -38,6 +41,7 @@ from airflow.configuration import conf from airflow.dag_processing.bundles.base import BundleVersionLock from airflow.dag_processing.dagbag import BundleDagBag, DagBag +from airflow.exceptions import AirflowConfigException from airflow.models.dag import DagModel from airflow.sdk.exceptions import TaskNotFound from airflow.sdk.execution_time import supervisor @@ -87,15 +91,77 @@ from structlog.typing import FilteringBoundLogger from airflow.api_fastapi.execution_api.app import InProcessExecutionAPI + from airflow.dag_processing.task_handler_resolution import TaskHandlerResolution from airflow.sdk.api.client import Client from airflow.sdk.bases.operator import BaseOperator from airflow.sdk.definitions.context import Context from airflow.sdk.definitions.dag import DAG from airflow.sdk.definitions.mappedoperator import MappedOperator - from airflow.sdk.execution_time.supervisor import RequestHandler, RequestResult + from airflow.sdk.execution_time.supervisor import RequestHandler, RequestResult, ResponseSent from airflow.typing_compat import Self +TaskHandlerBindingMode = Literal["positional", "named"] + + +class TaskHandlerParam(BaseModel): + """One parameter of a task handler.""" + + name: str | None + """``None`` when the runtime has no name for this positional parameter.""" + + value_schema: ArgValueSchema | None = None + """JSON Schema of the values the parameter accepts; ``None`` when the handler does not constrain it.""" + + exact_name: bool = False + """Whether ``name`` matches only as spelled, not case-insensitively with underscores ignored.""" + + +class TaskHandlerDeclaration(BaseModel): + """A task handler that a Lang-SDK artifact registers for one task.""" + + task_id: str + + binding: TaskHandlerBindingMode + """ + How stub-task arguments bind to ``params``. + + - ``positional``: by position; names are informative only. An argument count that matches ``params`` + neither with every argument nor after dropping the defaulted ones, or a value type the param does not + accept, makes the Dag fail to import. + - ``named``: by name in any order, case-insensitively with underscores ignored unless ``exact_name`` is set. + An argument no param takes, or a param no argument fills, is logged as a warning, and the task still + runs. When no param matches and exactly one argument was passed, it may be the whole value and is not + warned about, unless ``params`` is empty, a param sets ``exact_name``, or the argument cannot be an + object. A value type a param does not accept makes the Dag fail to import. + """ + + # The title keeps Go's generated type for this list apart from the HITL ``Params`` map. + params: Annotated[list[TaskHandlerParam] | None, Field(title="Task Handler Params")] + """ + In declaration order; the order is significant only for ``positional`` binding. + + ``None`` when the runtime cannot list the handler's parameters, so only the handler's presence is checked. + """ + + +class TaskHandlerArtifact(BaseModel): + """A Lang-SDK artifact and every task handler it registers.""" + + bundle_name: str + + relative_fileloc: str = Field(max_length=2000) + """Path of the artifact within its bundle.""" + + size_bytes: int + + cache_digest: str | None = Field(max_length=128) + """Opaque content fingerprint defined by the coordinator; ``None`` when the artifact stores none, so it is always probed.""" + + task_handlers: dict[str, list[TaskHandlerDeclaration]] + """Every Dag id the artifact registers a task handler for, with those handlers.""" + + class DagFileParseRequest(BaseModel): """ Request for DAG File Parsing. @@ -113,9 +179,26 @@ class DagFileParseRequest(BaseModel): """Bundle name for team-specific executor validation.""" callback_requests: list[CallbackRequest] = Field(default_factory=list) + + known_artifacts: list[TaskHandlerArtifact] = Field(default_factory=list) + """The recorded task-handler artifacts, with their answers, that this file's stub tasks may resolve against.""" + type: Literal["DagFileParseRequest"] = "DagFileParseRequest" +class TaskHandlerBinding(BaseModel): + """A stub task resolved to the Lang-SDK artifact that runs it.""" + + dag_id: str + task_id: str + + artifact_bundle_name: str + """The bundle the artifact was found in.""" + + artifact_rel_path: str = Field(max_length=2000) + """Path of the artifact within its bundle.""" + + class DagFileParsingResult(BaseModel): """ Result of DAG File Parsing. @@ -132,11 +215,65 @@ class DagFileParsingResult(BaseModel): """Bundle-relative locations of the Dag definitions imported from ``fileloc``.""" dag_source_codes: dict[str, DagSourceCode] = Field(default_factory=dict) """Source code of the parsed Dags, keyed by Dag fileloc.""" + + task_handler_bindings: list[TaskHandlerBinding] | None = None + """ + The stub-task bindings of every Dag in ``serialized_dags``. + + ``None`` when task handlers were not evaluated, and the recorded bindings are left as they are. A list + replaces the recorded bindings of each Dag in ``serialized_dags``, so a Dag with no entry has none. + """ + + probed_artifacts: list[TaskHandlerArtifact] = Field(default_factory=list) + """ + Every artifact this parse probed successfully, with its answer. + + Recorded even when ``task_handler_bindings`` is ``None``, so a Dag that fails validation is not probed + again on every parse. + """ + type: Literal["DagFileParsingResult"] = "DagFileParsingResult" +class TaskHandlerParseRequest(BaseModel): + """ + Request for Task Handler Parsing. + + Asks a Lang-SDK runtime for every task handler an artifact registers. + """ + + file: str + """The artifact to ask.""" + + bundle_path: Path + + bundle_name: str + + type: Literal["TaskHandlerParseRequest"] = "TaskHandlerParseRequest" + + +class TaskHandlerParsingResult(BaseModel): + """ + Result of Task Handler Parsing. + + Every task handler a Lang-SDK artifact registers, keyed by Dag id. + + The answer depends only on the artifact, never on the request. + """ + + fileloc: str + + task_handlers: dict[str, list[TaskHandlerDeclaration]] + """Every Dag id the artifact registers a task handler for; ``{}`` when it registers none.""" + + import_errors: dict[str, str] | None = None + warnings: list | None = None + type: Literal["TaskHandlerParsingResult"] = "TaskHandlerParsingResult" + + ToManager = Annotated[ DagFileParsingResult + | TaskHandlerParsingResult | GetConnection | GetVariable | GetVariableKeys @@ -155,9 +292,9 @@ class DagFileParsingResult(BaseModel): Field(discriminator="type"), ] -ToDagProcessor = Annotated[ - DagFileParseRequest - | ConnectionResult +# Answers to the child's requests, whichever parse it was started for. +_ParseSideResponses = ( + ConnectionResult | VariableResult | VariableKeysResult | TaskStatesResult @@ -169,8 +306,13 @@ class DagFileParsingResult(BaseModel): | XComCountResponse | XComResult | XComSequenceIndexResult - | XComSequenceSliceResult, - Field(discriminator="type"), + | XComSequenceSliceResult +) + +ToDagProcessor = Annotated[DagFileParseRequest | _ParseSideResponses, Field(discriminator="type")] + +ToSDKTaskHandlerProcessor = Annotated[ + TaskHandlerParseRequest | _ParseSideResponses, Field(discriminator="type") ] @@ -216,14 +358,24 @@ def _parse_file_entrypoint(): task_runner.SUPERVISOR_COMMS = comms_decoder log = structlog.get_logger(logger_name="task") - result = _parse_file(msg, log) + result = _parse_file(msg, log, started=_get_process_start()) if result is not None: comms_decoder.send(result) -def _parse_file(msg: DagFileParseRequest, log: FilteringBoundLogger) -> DagFileParsingResult | None: +def _parse_file( + msg: DagFileParseRequest, log: FilteringBoundLogger, *, started: float | None = None +) -> DagFileParsingResult | None: + """ + Parse the Dag file of *msg* and return the result to send, or ``None`` for a callback request. + + *started* is when the parse began, a :func:`time.monotonic` value that the check of the stub tasks counts + its time budget from. It defaults to the call of this function. + """ # TODO: Set known_pool names on DagBag! + if started is None: + started = time.monotonic() stability_check_result = check_dag_file_stability(os.fspath(msg.file)) @@ -254,6 +406,8 @@ def _parse_file(msg: DagFileParseRequest, log: FilteringBoundLogger) -> DagFileP serialized_dags, serialization_import_errors = _serialize_dags(bag, log) bag.import_errors.update(serialization_import_errors) + task_handlers = _resolve_task_handlers(msg, bag, serialized_dags, log, started=started) + _add_import_errors(bag.import_errors, task_handlers.import_errors) result = DagFileParsingResult( fileloc=msg.file, serialized_dags=serialized_dags, @@ -267,10 +421,93 @@ def _parse_file(msg: DagFileParseRequest, log: FilteringBoundLogger) -> DagFileP ], parsed_definitions=bag.parsed_definitions, dag_source_codes=bag.dag_source_codes, + task_handler_bindings=task_handlers.bindings, + probed_artifacts=task_handlers.probed_artifacts, ) return result +def _get_process_start() -> float: + """ + Return when this process was created, as a :func:`time.monotonic` value. + + The manager counts ``[dag_processor] dag_file_processor_timeout`` from the creation of the Dag-parsing + child, so the start-up of an interpreter that is exec'd counts too. An age that is negative or larger than + the timeout means the clocks disagree, and the process is counted from now, as it is when its creation + time cannot be read. On Linux psutil computes the creation time from the boot time in whole seconds, + so the age can be up to about 1 s too high, which only makes the deadline earlier. + """ + now = time.monotonic() + try: + age = time.time() - psutil.Process().create_time() + except (psutil.Error, OSError): + return now + if not 0 <= age <= conf.getfloat("dag_processor", "dag_file_processor_timeout"): + return now + return now - age + + +def _resolve_task_handlers( + msg: DagFileParseRequest, + bag: DagBag, + serialized_dags: list[LazyDeserializedDAG], + log: FilteringBoundLogger, + *, + started: float, +) -> TaskHandlerResolution: + """ + Check the stub tasks of the serialized Dags whose queues ``[sdk] queue_to_coordinator`` routes. + + Nothing is checked without that option, and a stub task on another queue is left to a worker outside + Airflow's coordinators. The check ends by 90% of ``[dag_processor] dag_file_processor_timeout`` from + *started*, so the result is sent before the manager kills this process. An unexpected error is an + import error of each Dag file with a checked stub task, so the serialized Dags are still sent. + """ + # Imported here: the probe imports this module, and a Dag file without stub tasks needs neither. + from airflow.dag_processing.task_handler_resolution import TaskHandlerResolution, resolve_task_handlers + from airflow.dag_processing.task_handler_validation import collect_stub_tasks + + try: + queue_to_coordinator = conf.getjson("sdk", "queue_to_coordinator", fallback={}) + except AirflowConfigException: + queue_to_coordinator = None + else: + if not queue_to_coordinator: + return TaskHandlerResolution(bindings=None, probed_artifacts=[], import_errors={}) + # An invalid option does not filter here: resolving the stub tasks reports it. + routed_queues = queue_to_coordinator if isinstance(queue_to_coordinator, dict) else None + stub_tasks = [ + stub + for stub in collect_stub_tasks(bag.dags.values(), serialized_dags) + if routed_queues is None or stub.queue in routed_queues + ] + if not stub_tasks: + return TaskHandlerResolution(bindings=[], probed_artifacts=[], import_errors={}) + try: + return resolve_task_handlers( + stub_tasks, + dag_bundle_name=msg.bundle_name, + dag_bundle_path=msg.bundle_path, + known_artifacts=msg.known_artifacts, + deadline=started + 0.9 * conf.getfloat("dag_processor", "dag_file_processor_timeout"), + log=log, + ) + except Exception as e: + log.exception("Failed to check the stub tasks against their task handlers", file=msg.file) + return TaskHandlerResolution.failed( + stub_tasks, f"Unexpected error while checking the stub tasks: {type(e).__name__}: {e}" + ) + + +def _add_import_errors(import_errors: dict[str, str], new_import_errors: dict[str, str]) -> None: + """Add *new_import_errors* to *import_errors*, after a blank line in a file that already has one.""" + for fileloc, message in new_import_errors.items(): + if existing := import_errors.get(fileloc): + import_errors[fileloc] = f"{existing.rstrip()}\n\n{message}" + else: + import_errors[fileloc] = message + + def _serialize_dags( bag: DagBag, log: FilteringBoundLogger, @@ -573,18 +810,17 @@ def in_process_api_server() -> InProcessExecutionAPI: @attrs.define(kw_only=True) -class DagFileProcessorProcess(WatchedSubprocess, LoggingMixin): +class BaseDagFileProcessorProcess(WatchedSubprocess, LoggingMixin): """ - Parses dags with Task SDK API. + Parse one Dag file in a child process for the Dag processor manager. - This class provides a wrapper and management around a subprocess to parse a specific DAG file. - - Since DAGs are written with the Task SDK, we need to parse them in a task SDK process such that - we can use the Task SDK definitions when serializing. This prevents potential conflicts with classes - in core Airflow. + The child's output goes to the file's parse log, and its requests are answered with + :attr:`client`. The parse is done once the child has exited and all its sockets are closed; + :attr:`parsing_result` then holds what it sent. Subclasses start the child and send it the + parse request. """ - logger_filehandle: BinaryIO + logger_filehandle: BinaryIO | None = None parsing_result: DagFileParsingResult | None = None decoder: ClassVar[TypeAdapter[ToManager]] = TypeAdapter[ToManager](ToManager) had_callbacks: bool = False # Track if this process was started with callbacks to prevent stale DAG detection false positives @@ -595,60 +831,6 @@ class DagFileProcessorProcess(WatchedSubprocess, LoggingMixin): bundle_name: str dag_file_rel_path: str - @classmethod - def start( # type: ignore[override] - cls, - *, - path: str | os.PathLike[str], - bundle_path: Path, - bundle_name: str, - dag_file_rel_path: str, - callbacks: list[CallbackRequest], - target: Callable[[], None] = _parse_file_entrypoint, - client: Client, - **kwargs, - ) -> Self: - logger = kwargs["logger"] - - # Parsing DAG files runs user code that can trigger macOS-unsafe ObjC - # initialization (secret backends, connection/variable lookups, HTTP - # clients). Fork+exec a clean interpreter there. Tests override `target` - # with a stub to exercise the base infrastructure; keep bare fork for those. - use_exec = target is _parse_file_entrypoint and supervisor._should_use_exec() - - # Pre-importing only helps the bare-fork child (it inherits the imports via - # copy-on-write). An exec'd child re-imports from scratch, so skip it there - # to avoid leaking user modules into the long-lived processor manager. - if not use_exec: - _pre_import_airflow_modules(os.fspath(path), logger) - - proc: Self = super().start( - target=target, - client=client, - bundle_name=bundle_name, - dag_file_rel_path=dag_file_rel_path, - use_exec=use_exec, - **kwargs, - ) - proc.had_callbacks = bool(callbacks) # Track if this process had callbacks - proc._on_child_started(callbacks, path, bundle_path, bundle_name) - return proc - - def _on_child_started( - self, - callbacks: list[CallbackRequest], - path: str | os.PathLike[str], - bundle_path: Path, - bundle_name: str, - ) -> None: - msg = DagFileParseRequest( - file=os.fspath(path), - bundle_path=bundle_path, - bundle_name=bundle_name, - callback_requests=callbacks, - ) - self.send_msg(msg, request_id=0) - def _get_target_loggers(self) -> tuple[FilteringBoundLogger, ...]: base = super()._get_target_loggers() if not self.subprocess_logs_to_stdout: @@ -674,11 +856,11 @@ def _create_log_forwarder( def _handle_parsing_result( self, msg: DagFileParsingResult, log: FilteringBoundLogger, req_id: int - ) -> RequestResult: + ) -> RequestResult | ResponseSent: self.parsing_result = msg return None, {} - _request_handlers: ClassVar[dict[type[BaseModel], RequestHandler[DagFileProcessorProcess]]] = { + _request_handlers: ClassVar[dict[type[BaseModel], RequestHandler[Any]]] = { **WatchedSubprocess._get_shared_request_handlers( DeleteVariable, GetConnection, @@ -720,6 +902,8 @@ def wait(self) -> int: def close(self): self.cleanup_sockets_after_kill() + if self.logger_filehandle is None: + return try: self.logger_filehandle.close() except OSError: @@ -728,3 +912,76 @@ def close(self): self.dag_file_rel_path, exc_info=True, ) + + +@attrs.define(kw_only=True) +class DagFileProcessorProcess(BaseDagFileProcessorProcess): + """ + Parses dags with Task SDK API. + + This class provides a wrapper and management around a subprocess to parse a specific DAG file. + + Since DAGs are written with the Task SDK, we need to parse them in a task SDK process such that + we can use the Task SDK definitions when serializing. This prevents potential conflicts with classes + in core Airflow. + """ + + logger_filehandle: BinaryIO + + @classmethod + def start( # type: ignore[override] + cls, + *, + path: str | os.PathLike[str], + bundle_path: Path, + bundle_name: str, + dag_file_rel_path: str, + callbacks: list[CallbackRequest], + known_artifacts: Sequence[TaskHandlerArtifact] = (), + target: Callable[[], None] = _parse_file_entrypoint, + client: Client, + **kwargs, + ) -> Self: + logger = kwargs["logger"] + + # Parsing DAG files runs user code that can trigger macOS-unsafe ObjC + # initialization (secret backends, connection/variable lookups, HTTP + # clients). Fork+exec a clean interpreter there. Tests override `target` + # with a stub to exercise the base infrastructure; keep bare fork for those. + use_exec = target is _parse_file_entrypoint and supervisor._should_use_exec() + + # Pre-importing only helps the bare-fork child (it inherits the imports via + # copy-on-write). An exec'd child re-imports from scratch, so skip it there + # to avoid leaking user modules into the long-lived processor manager. + if not use_exec: + _pre_import_airflow_modules(os.fspath(path), logger) + + proc: Self = super().start( + target=target, + client=client, + bundle_name=bundle_name, + dag_file_rel_path=dag_file_rel_path, + use_exec=use_exec, + **kwargs, + ) + proc.had_callbacks = bool(callbacks) # Track if this process had callbacks + proc._on_child_started(callbacks, path, bundle_path, bundle_name, known_artifacts=known_artifacts) + return proc + + def _on_child_started( + self, + callbacks: list[CallbackRequest], + path: str | os.PathLike[str], + bundle_path: Path, + bundle_name: str, + *, + known_artifacts: Sequence[TaskHandlerArtifact] = (), + ) -> None: + msg = DagFileParseRequest( + file=os.fspath(path), + bundle_path=bundle_path, + bundle_name=bundle_name, + callback_requests=callbacks, + known_artifacts=list(known_artifacts), + ) + self.send_msg(msg, request_id=0) diff --git a/airflow-core/src/airflow/dag_processing/task_handler_fast_path.py b/airflow-core/src/airflow/dag_processing/task_handler_fast_path.py new file mode 100644 index 0000000000000..f9484a6ba366a --- /dev/null +++ b/airflow-core/src/airflow/dag_processing/task_handler_fast_path.py @@ -0,0 +1,79 @@ +# 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. +"""Decide which task-handler artifacts a Dag file's parse must probe, and which recorded answers still hold.""" + +from __future__ import annotations + +from typing import TYPE_CHECKING + +import attrs + +if TYPE_CHECKING: + from collections.abc import Iterable, Sequence + + from airflow.dag_processing.processor import TaskHandlerArtifact + from airflow.sdk.execution_time.coordinator import TaskHandlerCandidate # noqa: SDK001 + + +@attrs.frozen(kw_only=True) +class TaskHandlerProbePlan: + """What a parse does with each candidate in one coordinator's artifact bundle.""" + + cached: list[TaskHandlerArtifact] + """Recorded answers that still hold, in candidate order.""" + + probe: list[TaskHandlerCandidate] + """Candidates to probe, in candidate order.""" + + rejected: list[TaskHandlerCandidate] + """Candidates the listing found unusable; they are neither probed nor answered from a record.""" + + +def plan_task_handler_probes( + *, + bundle_name: str, + candidates: Sequence[TaskHandlerCandidate], + known_artifacts: Iterable[TaskHandlerArtifact], +) -> TaskHandlerProbePlan: + """ + Sort the *candidates* listed in *bundle_name* into recorded answers to reuse, probes, and rejects. + + A probe answer depends only on the artifact, so the answer recorded with an artifact still holds while + the candidate has the same size and stored cache digest. The size is compared too, because a stored + digest is read, not recomputed, and survives an in-place edit. A candidate that stores no digest is + always probed. + """ + known = { + artifact.relative_fileloc: artifact + for artifact in known_artifacts + if artifact.bundle_name == bundle_name + } + cached: list[TaskHandlerArtifact] = [] + probe: list[TaskHandlerCandidate] = [] + rejected: list[TaskHandlerCandidate] = [] + for candidate in candidates: + if candidate.error is not None: + rejected.append(candidate) + elif ( + candidate.cache_digest is not None + and (artifact := known.get(candidate.rel_path)) is not None + and (artifact.size_bytes, artifact.cache_digest) == (candidate.size_bytes, candidate.cache_digest) + ): + cached.append(artifact) + else: + probe.append(candidate) + return TaskHandlerProbePlan(cached=cached, probe=probe, rejected=rejected) diff --git a/airflow-core/src/airflow/dag_processing/task_handler_processor.py b/airflow-core/src/airflow/dag_processing/task_handler_processor.py new file mode 100644 index 0000000000000..ae099a2b8366f --- /dev/null +++ b/airflow-core/src/airflow/dag_processing/task_handler_processor.py @@ -0,0 +1,627 @@ +# 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. +"""Ask the Lang-SDK runtime of a coordinator which task handlers an artifact registers.""" + +from __future__ import annotations + +import contextlib +import functools +import os +import selectors +import signal +import time +from pathlib import Path +from socket import MSG_DONTWAIT, socket +from typing import TYPE_CHECKING, Annotated, Any, ClassVar, Literal, cast, get_args + +import attrs +import msgspec +import psutil +from pydantic import BaseModel, Field, TypeAdapter +from uuid6 import uuid7 + +from airflow import settings +from airflow.configuration import conf +from airflow.dag_processing.processor import ( + BaseDagFileProcessorProcess, + DagFileParsingResult, + TaskHandlerParseRequest, + TaskHandlerParsingResult, + ToManager, +) +from airflow.sdk.coordinators._subprocess import _is_connection_from_pid, _start_server +from airflow.sdk.exceptions import AirflowRuntimeError +from airflow.sdk.execution_time import supervisor, task_runner +from airflow.sdk.execution_time.comms import CommsDecoder, ErrorResponse, MaskSecret, _RequestFrame +from airflow.sdk.execution_time.coordinator import get_coordinator_manager +from airflow.sdk.execution_time.supervisor import ( + ResponseSent, + length_prefixed_frame_reader, + make_buffered_socket_reader, + process_log_messages_from_subprocess, + register_request_method, +) + +if TYPE_CHECKING: + from collections.abc import Generator + + from structlog.typing import FilteringBoundLogger + + from airflow.sdk.api.client import Client + from airflow.sdk.execution_time.supervisor import RequestHandler, RequestResult + from airflow.typing_compat import Self + +# How long a runtime may keep running after its parse result, as Node does while a handle stays open. +_EXIT_GRACE_PERIOD = 5.0 + + +# StartTaskHandlerRuntime, LangSDKRuntimeSchemaVersion and LangSDKRuntimeStartFailed pass only between +# the probe and its forked child before the exec, so they are not part of the supervisor schema the +# runtimes speak. + + +class StartTaskHandlerRuntime(BaseModel): + """Ask the parse child to exec the runtime of the coordinator configured under *coordinator*.""" + + file: str + bundle_path: Path + coordinator: str + """The coordinator's key in ``[sdk] coordinators``.""" + comm_address: tuple[str, int] + logs_address: tuple[str, int] + type: Literal["StartTaskHandlerRuntime"] = "StartTaskHandlerRuntime" + + +class LangSDKRuntimeSchemaVersion(BaseModel): + """The schema version and the import timeout of the runtime the parse child is about to exec.""" + + schema_version: str | None + import_timeout: float | None = None + """Seconds from the start of the parse; ``None`` means no timeout.""" + type: Literal["LangSDKRuntimeSchemaVersion"] = "LangSDKRuntimeSchemaVersion" + + +class LangSDKRuntimeStartFailed(BaseModel): + """Why the parse child could not start the runtime.""" + + error: str + type: Literal["LangSDKRuntimeStartFailed"] = "LangSDKRuntimeStartFailed" + + +class TaskHandlerProbeStopped(TaskHandlerParsingResult): + """ + The result of a probe that the caller's deadline stopped before its runtime answered. + + Its import error says so. A runtime that failed on its own, or answered before the deadline, gives a + plain :class:`TaskHandlerParsingResult`. + """ + + +def _get_import_timeout(path: str) -> float | None: + """Return the ``get_dagbag_import_timeout`` policy's timeout for *path*; ``None`` means none.""" + timeout = settings.get_dagbag_import_timeout(path) + if not isinstance(timeout, (int, float)): + raise TypeError(f"Value ({timeout}) from get_dagbag_import_timeout must be int or float") + return timeout if timeout > 0 else None + + +def _start_task_handler_runtime_entrypoint() -> None: + """Exec the runtime that probes the artifact named by the start request, or report why it cannot start.""" + os.environ["_AIRFLOW_PROCESS_CONTEXT"] = "client" + # fd 0 becomes the runtime's stdin, so the request channel moves to a close-on-exec copy. + comms = CommsDecoder[StartTaskHandlerRuntime, LangSDKRuntimeSchemaVersion | LangSDKRuntimeStartFailed]( + socket=socket(fileno=os.dup(0)), + body_decoder=TypeAdapter(StartTaskHandlerRuntime), + ) + devnull = os.open(os.devnull, os.O_RDONLY) + os.dup2(devnull, 0) + os.close(devnull) + + msg = comms._get_response() + if not isinstance(msg, StartTaskHandlerRuntime): + raise RuntimeError(f"Required first message to be a StartTaskHandlerRuntime, it was {msg}") + + def report_schema_version(schema_version: str | None) -> None: + comms.send(LangSDKRuntimeSchemaVersion(schema_version=schema_version, import_timeout=import_timeout)) + + try: + # The policy is user code: it runs in this child, where a failure is only this file's import error. + import_timeout = _get_import_timeout(msg.file) + get_coordinator_manager().get_coordinator(msg.coordinator).parse_task_handler( + path=Path(msg.file), + bundle_path=msg.bundle_path, + comm_address=msg.comm_address, + logs_address=msg.logs_address, + report_schema_version=report_schema_version, + ) + except Exception as e: + comms.send(LangSDKRuntimeStartFailed(error=f"{type(e).__name__}: {e}")) + + +_Channel = Literal["comm", "logs"] + + +class _ReadsWithoutWaiting: + """ + The runtime's comm socket, with reads that return at once instead of waiting for data. + + A runtime that stops in the middle of a frame then cannot block the caller's loop, which keeps + checking the import timeout. Replies are sent on the socket itself, which stays blocking. + """ + + def __init__(self, sock: socket) -> None: + self._sock = sock + + def recv(self, bufsize: int) -> bytes: + return self._sock.recv(bufsize, MSG_DONTWAIT) + + def recv_into(self, buffer: memoryview) -> int: + return self._sock.recv_into(buffer, 0, MSG_DONTWAIT) + + +# The requests relayed to the Dag processor without a client: those it answers, but not the Dag parse's +# result. MaskSecret is not relayed: its handler masks the secret here, and mask_secret sends it on to a +# parent. +_PARENT_REQUESTS = frozenset(BaseDagFileProcessorProcess._request_handlers) - { + DagFileParsingResult, + MaskSecret, +} + + +@attrs.define(kw_only=True) +class LangSDKTaskHandlerProcessorProcess(BaseDagFileProcessorProcess): + """ + Ask a coordinator's Lang-SDK runtime for every task handler an artifact registers. + + The forked parse child finds the coordinator, reports the runtime's schema version and execs the + runtime. The runtime connects back to two listeners this process owns and answers the + ``TaskHandlerParseRequest`` itself, so the request is sent once it has connected. A failed start, + a missing result, an invalid frame or message, or a timeout is an import error on the result, + keyed by the artifact's path in its Dag bundle. Processes the runtime leaves in its process group + are killed when it exits. + + The runtime's requests are answered with :attr:`client`. Without one, as in a Dag-parsing child, + they are relayed up the supervisor channel of the process this runs in, or get an error when there + is none. On Linux the runtime is killed when the thread that started this process exits, so start + it from a thread that outlives the probe. + """ + + client: Client | None = None # type: ignore[assignment] + """Answers the runtime's requests; without one, they are relayed to the parent process if there is one.""" + + parsing_result: TaskHandlerParsingResult | None = None # type: ignore[assignment] + + decoder = TypeAdapter( + Annotated[ + LangSDKRuntimeSchemaVersion | LangSDKRuntimeStartFailed | get_args(ToManager)[0], + Field(discriminator="type"), + ] + ) + + coordinator: str + """The coordinator's key in ``[sdk] coordinators``.""" + + _listeners: dict[_Channel, socket] + _parse_request: TaskHandlerParseRequest + _runtime_schema_version: str | None = attrs.field(default=None, init=False) + _import_timeout: float | None = attrs.field(default=None, init=False) + _schema_version_reported: bool = attrs.field(default=False, init=False) + _parsing_result_monotonic: float | None = attrs.field(default=None, init=False) + _unverified_connections: list[tuple[socket, _Channel]] = attrs.field(factory=list, init=False) + + @classmethod + def start( # type: ignore[override] + cls, + *, + coordinator: str, + path: str | os.PathLike[str], + bundle_path: Path, + bundle_name: str, + artifact_rel_path: str, + **kwargs, + ) -> Self: + """ + Start probing the artifact at *path* for every task handler it registers. + + *bundle_path* and *bundle_name* are those of the Dag bundle holding the artifact, and + *artifact_rel_path* is the artifact's path in it. + """ + listeners: dict[_Channel, socket] = {"comm": _start_server(), "logs": _start_server()} + try: + for listener in listeners.values(): + listener.setblocking(False) + parse_request = TaskHandlerParseRequest( + file=os.fspath(path), bundle_path=bundle_path, bundle_name=bundle_name + ) + proc = super().start( + target=_start_task_handler_runtime_entrypoint, + use_exec=supervisor._should_use_exec(), + new_process_group=True, + coordinator=coordinator, + bundle_name=bundle_name, + dag_file_rel_path=artifact_rel_path, + listeners=listeners, + parse_request=parse_request, + **kwargs, + ) + except BaseException: + for listener in listeners.values(): + listener.close() + raise + try: + for channel, listener in listeners.items(): + proc._open_sockets[listener] = f"{channel}-listener" + proc.selector.register( + listener, + selectors.EVENT_READ, + (functools.partial(proc._accept_connection, channel=channel), proc._on_socket_closed), + ) + proc.send_msg( + StartTaskHandlerRuntime( + file=parse_request.file, + bundle_path=bundle_path, + coordinator=coordinator, + comm_address=listeners["comm"].getsockname()[:2], + logs_address=listeners["logs"].getsockname()[:2], + ), + request_id=0, + ) + except BaseException: + proc._kill_runtime() + proc.close() + raise + return proc + + @classmethod + def run( + cls, + *, + coordinator: str, + path: str | os.PathLike[str], + bundle_path: Path, + bundle_name: str, + artifact_rel_path: str, + logger: FilteringBoundLogger, + deadline: float | None = None, + ) -> TaskHandlerParsingResult: + """ + Probe the artifact at *path* as :meth:`start` does, and wait for the result. + + The artifact's ``get_dagbag_import_timeout`` bounds the probe, and + ``[dag_processor] dag_file_processor_timeout`` until the parse child resolves it. A + *deadline*, a :func:`time.monotonic` value, bounds it too: a probe still running then is + killed, and its result is a :class:`TaskHandlerProbeStopped` import error, unless the runtime + had already answered. + """ + return cls._run_to_completion( + coordinator=coordinator, + path=path, + bundle_path=bundle_path, + bundle_name=bundle_name, + artifact_rel_path=artifact_rel_path, + logger=logger, + deadline=deadline, + ) + + @classmethod + def _run_to_completion( + cls, *, logger: FilteringBoundLogger, deadline: float | None = None, **start_kwargs: Any + ) -> TaskHandlerParsingResult: + """ + Probe outside a caller's selector loop and wait for the result. + + The import timeout the parse child reports bounds the probe, and + ``[dag_processor] dag_file_processor_timeout`` until it is reported. *deadline* bounds it + either way. + """ + processor_timeout = conf.getfloat("dag_processor", "dag_file_processor_timeout") + with selectors.DefaultSelector() as selector: + proc = cls.start(id=uuid7(), selector=selector, logger=logger, **start_kwargs) + try: + while not proc.is_ready: + # is_ready applies the import timeout once the parse child has reported it. + now = time.monotonic() + if deadline is not None and now >= deadline: + proc._stop( + f"The Lang-SDK runtime did not parse {proc._parse_request.file} by its deadline", + at_deadline=True, + ) + break + if not proc._schema_version_reported and now - proc.start_time > processor_timeout: + proc._time_out(processor_timeout) + break + proc._service_subprocess(max_wait_time=0.1) + except BaseException: + proc._kill_runtime() + raise + finally: + proc.close() + return cast("TaskHandlerParsingResult", proc.parsing_result) + + def _accept_connection(self, listener: socket, *, channel: _Channel) -> bool: + try: + conn, _ = listener.accept() + except (BlockingIOError, InterruptedError): + return True + conn.setblocking(True) + self._unverified_connections.append((conn, channel)) + self._verify_connections() + return True + + def _verify_connections(self) -> None: + """ + Use each accepted connection once it is confirmed to come from the runtime. + + A connection that is not visible yet stays pending and is checked again on the next + ``is_ready`` poll, so the caller's loop never waits here. + """ + pending = [] + for conn, channel in self._unverified_connections: + if channel not in self._listeners: + # The runtime already connected this channel. + conn.close() + continue + try: + owned = _is_connection_from_pid(conn, self.pid) + except OSError: + conn.close() + continue + if not owned: + pending.append((conn, channel)) + continue + self._close_listener(channel) + if channel == "comm": + self._register_comm(conn) + else: + self._register_logs(conn) + self._unverified_connections = pending + + def _close_listener(self, channel: _Channel) -> None: + if (listener := self._listeners.pop(channel, None)) is not None: + self._on_socket_closed(listener) + listener.close() + + def _close_listeners(self) -> None: + """Close the listeners of a runtime that did not connect, and connections never verified.""" + for channel in list(self._listeners): + self._close_listener(channel) + for conn, _ in self._unverified_connections: + conn.close() + self._unverified_connections = [] + + def _register_comm(self, conn: socket) -> None: + self.stdin = conn + self._open_sockets[conn] = "requests" + read_frame, on_close = length_prefixed_frame_reader( + self._handle_valid_requests(), on_close=self._on_socket_closed + ) + + def read_valid_frame(sock: socket) -> bool: + try: + return read_frame(cast("socket", _ReadsWithoutWaiting(sock))) + except BlockingIOError: + # The rest of the frame has not arrived; the reader keeps what it has read so far. + return True + except msgspec.DecodeError as e: + # A frame that does not decode would otherwise escape the caller's selector loop. + self._fail_on_invalid_message(f"The Lang-SDK runtime sent an invalid frame: {e}") + return False + + self.selector.register(conn, selectors.EVENT_READ, (read_valid_frame, on_close)) + # The parse child reports the version and waits for the reply before it execs the runtime, + # so the version is known here. It is set only now, so the child's messages are not migrated. + self._subprocess_schema_version = self._runtime_schema_version + self.send_msg(self._parse_request, request_id=0) + + def _handle_valid_requests(self) -> Generator[None, _RequestFrame, None]: + """ + Pass each request on to ``handle_requests``, or kill the runtime at one that does not validate. + + ``handle_requests`` would only log such a request, and the runtime would wait for a reply. The + runtime speaks ``ToManager`` only; the start messages come from the parse child. + """ + requests = self.handle_requests(self.process_log) + next(requests) + while True: + frame = yield + try: + BaseDagFileProcessorProcess.decoder.validate_python(self._deserialize_request(frame.body)) + except ValueError as e: + self._fail_on_invalid_message( + f"The Lang-SDK runtime sent a message that does not validate: {e}" + ) + return + requests.send(frame) + + def _fail_on_invalid_message(self, message: str) -> None: + """Kill the runtime; *message* is the import error unless a parse result was already received.""" + if self.parsing_result is None: + self._set_import_error(message) + else: + self.process_log.warning( + "Ignoring an invalid message from the Lang-SDK runtime after its parse result", error=message + ) + self._kill_runtime() + + def _register_logs(self, conn: socket) -> None: + self._open_sockets[conn] = "logs" + self.selector.register( + conn, + selectors.EVENT_READ, + make_buffered_socket_reader( + process_log_messages_from_subprocess(self._get_target_loggers()), + on_close=self._on_socket_closed, + ), + ) + + def _set_import_error(self, message: str, *, at_deadline: bool = False) -> None: + result_type = TaskHandlerProbeStopped if at_deadline else TaskHandlerParsingResult + self.parsing_result = result_type( + fileloc=self._parse_request.file, + task_handlers={}, + import_errors={self.dag_file_rel_path: message}, + ) + + def _handle_runtime_schema_version( + self, msg: LangSDKRuntimeSchemaVersion, log: FilteringBoundLogger, req_id: int + ) -> RequestResult | ResponseSent: + if self._schema_version_reported: + self._reject_request(msg, log, req_id) + return ResponseSent.ALREADY_SENT + self._runtime_schema_version = msg.schema_version + self._import_timeout = msg.import_timeout + self._schema_version_reported = True + return None, {} + + def _handle_start_failed( + self, msg: LangSDKRuntimeStartFailed, log: FilteringBoundLogger, req_id: int + ) -> RequestResult | ResponseSent: + self._set_import_error(f"Cannot start the Lang-SDK runtime: {msg.error}") + return None, {} + + def _handle_parsing_result( # type: ignore[override] + self, msg: TaskHandlerParsingResult, log: FilteringBoundLogger, req_id: int + ) -> RequestResult | ResponseSent: + if self.parsing_result is not None: + log.warning("Ignoring another parse result from the Lang-SDK runtime", fileloc=msg.fileloc) + self.send_msg( + None, + request_id=req_id, + error=ErrorResponse(detail={"message": "A parse result was already received"}), + ) + return ResponseSent.ALREADY_SENT + self.parsing_result = msg + self._parsing_result_monotonic = time.monotonic() + return None, {} + + _request_handlers: ClassVar[dict[type[BaseModel], RequestHandler[Any]]] = { + **{ + message_type: handler + for message_type, handler in BaseDagFileProcessorProcess._request_handlers.items() + if message_type is not DagFileParsingResult + }, + **dict( + [ + register_request_method(LangSDKRuntimeSchemaVersion, _handle_runtime_schema_version), + register_request_method(LangSDKRuntimeStartFailed, _handle_start_failed), + register_request_method(TaskHandlerParsingResult, _handle_parsing_result), + ] + ), + } + + def _handle_request(self, msg, log: FilteringBoundLogger, req_id: int) -> None: + if self.client is None and type(msg) in _PARENT_REQUESTS: + self._relay_request(msg, req_id) + return + super()._handle_request(msg, log, req_id) + + def _relay_request(self, msg: BaseModel, req_id: int) -> None: + """Answer the runtime's request through the supervisor channel of this process, if it has one.""" + comms = getattr(task_runner, "SUPERVISOR_COMMS", None) + if comms is None: + self.send_msg( + None, + request_id=req_id, + error=ErrorResponse( + detail={"message": f"{type(msg).__name__} is answered only in the Dag processor"} + ), + ) + return + try: + response = comms.send(msg) + except AirflowRuntimeError as e: + self.send_msg(None, request_id=req_id, error=e.error) + return + # Only the fields the parent sent, under their wire names, so the runtime gets the same body. + self.send_msg(response, request_id=req_id, exclude_unset=True, by_alias=True) + + @property + def is_ready(self) -> bool: + self._verify_connections() + if ( + self._parsing_result_monotonic is not None + and self._exit_code is None + and time.monotonic() - self._parsing_result_monotonic > _EXIT_GRACE_PERIOD + ): + self.process_log.warning("The Lang-SDK runtime did not exit after its parse result; killing it") + self._kill_runtime() + if ( + self._import_timeout is not None + and self.parsing_result is None + and self._exit_code is None + and time.monotonic() - self.start_time > self._import_timeout + ): + self._time_out(self._import_timeout) + if self._check_subprocess_exit() is None: + return False + self._close_listeners() + if ( + self._open_sockets + and self._import_timeout is not None + and time.monotonic() - self.start_time > self._import_timeout + ): + # A process the runtime left outside its process group holds these open. + self._time_out(self._import_timeout) + self.cleanup_sockets_after_kill() + if not super().is_ready: + return False + if self.parsing_result is None: + self._set_import_error( + f"The Lang-SDK runtime exited with code {self._exit_code} without a parse result" + ) + return True + + def _time_out(self, timeout: float) -> None: + self._stop(f"The Lang-SDK runtime did not parse {self._parse_request.file} within {timeout}s") + + def _stop(self, error: str, *, at_deadline: bool = False) -> None: + """Kill the runtime; *error* is the import error unless a parse result was already received.""" + if self.parsing_result is None: + self._set_import_error(error, at_deadline=at_deadline) + self._kill_runtime() + + def _check_subprocess_exit( + self, raise_on_timeout: bool = False, expect_signal: None | int = None + ) -> int | None: + if self._exit_code is None and self._is_runtime_exited(): + # Until the exited runtime is reaped below, its pid, and so its process group id, cannot be + # reused, so this reaches only the processes it left behind. + with contextlib.suppress(ProcessLookupError, PermissionError): + os.killpg(self.pid, signal.SIGKILL) + return super()._check_subprocess_exit(raise_on_timeout=raise_on_timeout, expect_signal=expect_signal) + + def _is_runtime_exited(self) -> bool: + """Return whether the runtime has exited and is not reaped yet.""" + try: + return psutil.Process(self.pid).status() == psutil.STATUS_ZOMBIE + except psutil.Error: + return False + + def _kill_runtime(self) -> None: + """Kill the runtime and wait for it, without servicing its sockets, whose handler may have failed.""" + if self._exit_code is not None: + return + try: + self._signal_subprocess(signal.SIGKILL) + self._exit_code = self._process.wait(timeout=None) + except (self._process.ProcessNotFound, ProcessLookupError): + self._exit_code = -1 + + def close(self) -> None: + # A listener has nothing to drain, and cleanup would call its accept handler forever. + self._close_listeners() + super().close() diff --git a/airflow-core/src/airflow/dag_processing/task_handler_resolution.py b/airflow-core/src/airflow/dag_processing/task_handler_resolution.py new file mode 100644 index 0000000000000..c563baff23922 --- /dev/null +++ b/airflow-core/src/airflow/dag_processing/task_handler_resolution.py @@ -0,0 +1,456 @@ +# 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. +"""Resolve a Dag file's stub tasks to the task handler artifacts of the coordinators their queues route to.""" + +from __future__ import annotations + +import os +import time +from collections import defaultdict +from typing import TYPE_CHECKING + +import attrs +from pydantic import ValidationError + +from airflow.dag_processing.bundles.manager import DagBundlesManager +from airflow.dag_processing.processor import TaskHandlerArtifact +from airflow.dag_processing.task_handler_fast_path import plan_task_handler_probes +from airflow.dag_processing.task_handler_validation import ( + TaskHandlerProblem, + format_import_errors, + match_task_handlers, +) + +if TYPE_CHECKING: + from collections.abc import Iterable, Sequence + from pathlib import Path + + from structlog.typing import FilteringBoundLogger + + from airflow.dag_processing.processor import TaskHandlerBinding + from airflow.dag_processing.task_handler_validation import StubTask, StubTaskWarning + from airflow.sdk.execution_time.coordinator import ( # noqa: SDK001 + CoordinatorManager, + TaskHandlerCandidate, + ) + +_OUT_OF_TIME = "the parse ran out of [dag_processor] dag_file_processor_timeout" + + +@attrs.frozen(kw_only=True) +class TaskHandlerResolution: + """What a Dag file's parse reports about the task handlers of its stub tasks.""" + + bindings: list[TaskHandlerBinding] | None + """A binding for each stub task checked; ``None`` when they were not checked, or a check failed.""" + + probed_artifacts: list[TaskHandlerArtifact] + """Every artifact probed with an answer, kept even when a check failed.""" + + import_errors: dict[str, str] + """The problems found, as one import error per Dag file.""" + + @classmethod + def failed(cls, stub_tasks: Iterable[StubTask], message: str) -> TaskHandlerResolution: + """Return a resolution that reports *message* as a problem of each Dag file with a stub task.""" + return cls( + bindings=None, + probed_artifacts=[], + import_errors=format_import_errors(_make_problems(stub_tasks, message)), + ) + + +@attrs.frozen(kw_only=True) +class _CoordinatorArtifacts: + """The candidates one coordinator lists in its Dag bundle, sorted by what the parse does with them.""" + + key: str + bundle_name: str + bundle_path: Path + cached: list[TaskHandlerArtifact] + probe: list[TaskHandlerCandidate] + broken: dict[str, str] + """Candidates the coordinator cannot probe, with why.""" + + +class _CoordinatorProblem(Exception): + """Why a coordinator's stub tasks cannot be checked.""" + + +@attrs.frozen(kw_only=True) +class _Probes: + """What probing the candidates gave.""" + + answers: dict[tuple[str, str], TaskHandlerArtifact | str] + """ + By Dag bundle name and path: the answer, or why the deadline left the candidate unprobed. + + Every coordinator that lists the candidate uses the answer. A skip stands for a coordinator only when + its own probe did not fail first, so a skip left by a later coordinator hides no earlier failure. + """ + + failures: dict[tuple[str, str, str], str] + """Why a probe gave no answer, by coordinator key, Dag bundle name and path.""" + + def get(self, coordinator: str, bundle_name: str, rel_path: str) -> TaskHandlerArtifact | str: + """Return *coordinator*'s answer for the candidate, or why it has none.""" + answer = self.answers.get((bundle_name, rel_path)) + if isinstance(answer, TaskHandlerArtifact): + return answer + if (failure := self.failures.get((coordinator, bundle_name, rel_path))) is not None: + return failure + return self.answers[bundle_name, rel_path] + + +class _DagBundles: + """Where the Dag bundles the coordinators name are, looked up once each in a parse.""" + + def __init__(self, dag_bundle_name: str) -> None: + self._dag_bundle_name = dag_bundle_name + self._manager: DagBundlesManager | None = None + self._paths: dict[str, Path | str] = {} + + def get_path(self, name: str) -> Path: + """Return where the Dag bundle *name* is, if it belongs to the team of the Dag's bundle.""" + if name not in self._paths: + try: + self._paths[name] = self._find_path(name) + except _CoordinatorProblem as e: + self._paths[name] = str(e) + path = self._paths[name] + if isinstance(path, str): + raise _CoordinatorProblem(path) + return path + + def _find_path(self, name: str) -> Path: + try: + if self._manager is None: + self._manager = DagBundlesManager() + team_name = self._manager.get_bundle_team_name(name) + dag_team_name = self._manager.get_bundle_team_name(self._dag_bundle_name) + except Exception as e: + raise _CoordinatorProblem( + f"cannot read the teams of Dag bundles {name!r} and {self._dag_bundle_name!r}: {_describe(e)}" + ) from e + if team_name != dag_team_name: + raise _CoordinatorProblem( + f"Dag bundle {name!r} belongs to {_describe_team(team_name)}, " + f"but Dag bundle {self._dag_bundle_name!r} belongs to {_describe_team(dag_team_name)}" + ) + try: + return self._manager.get_bundle(name).path + except Exception as e: + raise _CoordinatorProblem(f"cannot read Dag bundle {name!r}: {_describe(e)}") from e + + +def resolve_task_handlers( + stub_tasks: Sequence[StubTask], + *, + dag_bundle_name: str, + dag_bundle_path: Path, + known_artifacts: Sequence[TaskHandlerArtifact], + deadline: float, + log: FilteringBoundLogger, +) -> TaskHandlerResolution: + """ + Check *stub_tasks* against the task handlers of the coordinators their queues route to. + + A coordinator lists its artifacts in the Dag bundle named by its ``task_handler_bundle_name``, which must + belong to the team of the Dag bundle *dag_bundle_name*, or in that Dag bundle at *dag_bundle_path* when it + names none. A candidate whose answer in *known_artifacts* still holds is not probed. The others are + probed one at a time until *deadline*, a :func:`time.monotonic` value. An answer is shared by every + coordinator that lists the artifact, and a failed probe is retried under the next coordinator that + lists it. A candidate that cannot be probed, whose probe fails, or that the deadline leaves unprobed + has no answer, and is named only when a stub task finds no task handler. A stub task on a queue routed + to no coordinator is not checked. A coordinator that cannot be evaluated is a problem for each Dag + file with stub tasks on it, and those stub tasks are not checked. + """ + from airflow.sdk.execution_time.coordinator import get_coordinator_manager # noqa: SDK001 + + try: + manager = get_coordinator_manager() + except Exception as e: + return TaskHandlerResolution.failed(stub_tasks, f"Cannot load [sdk] coordinators: {_describe(e)}") + + stubs_by_coordinator: defaultdict[str, list[StubTask]] = defaultdict(list) + for stub in stub_tasks: + if (key := manager.get_coordinator_key(stub.queue)) is not None: + stubs_by_coordinator[key].append(stub) + + problems: list[TaskHandlerProblem] = [] + coordinators: list[_CoordinatorArtifacts] = [] + bundle_names = manager.get_task_handler_bundle_names() + bundles = _DagBundles(dag_bundle_name) + for key in sorted(stubs_by_coordinator): + try: + coordinators.append( + _list_artifacts( + manager, + key, + bundle_name=bundle_names[key], + bundles=bundles, + dag_bundle_name=dag_bundle_name, + dag_bundle_path=dag_bundle_path, + known_artifacts=known_artifacts, + log=log, + ) + ) + except _CoordinatorProblem as e: + problems.extend(_make_problems(stubs_by_coordinator[key], f"Coordinator {key!r}: {e}")) + except Exception as e: + log.warning( + "Cannot list the task handler artifacts of a coordinator", coordinator=key, exc_info=True + ) + problems.extend( + _make_problems( + stubs_by_coordinator[key], + f"Coordinator {key!r}: cannot find its task handler artifacts: {_describe(e)}", + ) + ) + + probes = _probe_candidates(coordinators, deadline=deadline, log=log) + + bindings: list[TaskHandlerBinding] = [] + for coordinator in coordinators: + coordinator_answers = list(coordinator.cached) + broken = dict(coordinator.broken) + for candidate in coordinator.probe: + answer = probes.get(coordinator.key, coordinator.bundle_name, candidate.rel_path) + if isinstance(answer, TaskHandlerArtifact): + coordinator_answers.append(answer) + else: + broken[candidate.rel_path] = answer + match = match_task_handlers( + stubs_by_coordinator[coordinator.key], + coordinator_answers, + bundle_name=coordinator.bundle_name, + broken_candidates=broken, + ) + bindings.extend(match.bindings) + problems.extend(match.problems) + for warning in match.warnings: + _log_name_mismatch(warning, log) + + probed_artifacts = [ + answer for answer in probes.answers.values() if isinstance(answer, TaskHandlerArtifact) + ] + if problems: + return TaskHandlerResolution( + bindings=None, probed_artifacts=probed_artifacts, import_errors=format_import_errors(problems) + ) + return TaskHandlerResolution(bindings=bindings, probed_artifacts=probed_artifacts, import_errors={}) + + +def _list_artifacts( + manager: CoordinatorManager, + key: str, + *, + bundle_name: str | None, + bundles: _DagBundles, + dag_bundle_name: str, + dag_bundle_path: Path, + known_artifacts: Sequence[TaskHandlerArtifact], + log: FilteringBoundLogger, +) -> _CoordinatorArtifacts: + if bundle_name is None: + bundle_name, bundle_path = dag_bundle_name, dag_bundle_path + else: + bundle_path = bundles.get_path(bundle_name) + try: + os.scandir(bundle_path).close() + except (FileNotFoundError, NotADirectoryError): + raise _CoordinatorProblem( + f"Dag bundle {bundle_name!r} resolved to {bundle_path}, which does not exist on this Dag processor" + ) from None + except OSError as e: + raise _CoordinatorProblem( + f"Dag bundle {bundle_name!r} at {bundle_path} cannot be read on this Dag processor: {_describe(e)}" + ) from e + try: + coordinator = manager.get_coordinator(key) + except Exception as e: + raise _CoordinatorProblem(f"cannot be built: {_describe(e)}") from e + try: + candidates = coordinator.list_task_handler_candidates(bundle_path) + except NotImplementedError: + classpath = f"{type(coordinator).__module__}.{type(coordinator).__qualname__}" + raise _CoordinatorProblem( + f"{classpath} cannot list task handler artifacts, so its stub tasks cannot be bound" + ) from None + except Exception as e: + raise _CoordinatorProblem( + f"cannot list the task handler artifacts of Dag bundle {bundle_name!r}: {_describe(e)}" + ) from e + + plan = plan_task_handler_probes( + bundle_name=bundle_name, candidates=candidates, known_artifacts=known_artifacts + ) + broken: dict[str, str] = {} + for candidate in plan.rejected: + log.warning( + "Ignoring a task handler artifact its coordinator cannot probe", + coordinator=key, + bundle_name=bundle_name, + path=candidate.rel_path, + error=candidate.error, + ) + broken[candidate.rel_path] = str(candidate.error) + for artifact in plan.cached: + log.debug( + "Using the recorded answer of a task handler artifact", + coordinator=key, + bundle_name=bundle_name, + path=artifact.relative_fileloc, + ) + return _CoordinatorArtifacts( + key=key, + bundle_name=bundle_name, + bundle_path=bundle_path, + cached=plan.cached, + probe=plan.probe, + broken=broken, + ) + + +def _probe_candidates( + coordinators: Iterable[_CoordinatorArtifacts], *, deadline: float, log: FilteringBoundLogger +) -> _Probes: + """Probe each candidate under each coordinator that lists it, until one has an answer.""" + probes = _Probes(answers={}, failures={}) + for coordinator in coordinators: + for candidate in coordinator.probe: + if (coordinator.bundle_name, candidate.rel_path) in probes.answers: + continue + context = { + "coordinator": coordinator.key, + "bundle_name": coordinator.bundle_name, + "path": candidate.rel_path, + } + if time.monotonic() >= deadline: + reason = f"not probed: {_OUT_OF_TIME}" + log.warning("Not probing a task handler artifact", reason=reason, **context) + probes.answers[coordinator.bundle_name, candidate.rel_path] = reason + continue + outcome = _probe_candidate(coordinator, candidate, deadline=deadline, log=log, context=context) + if isinstance(outcome, TaskHandlerArtifact): + probes.answers[coordinator.bundle_name, candidate.rel_path] = outcome + else: + probes.failures[coordinator.key, coordinator.bundle_name, candidate.rel_path] = outcome + return probes + + +def _probe_candidate( + coordinator: _CoordinatorArtifacts, + candidate: TaskHandlerCandidate, + *, + deadline: float, + log: FilteringBoundLogger, + context: dict[str, str], +) -> TaskHandlerArtifact | str: + """Probe *candidate* under *coordinator*, and return its answer or why it has none.""" + from airflow.dag_processing.task_handler_processor import ( + LangSDKTaskHandlerProcessorProcess, + TaskHandlerProbeStopped, + ) + + log.info("Probing a task handler artifact", **context) + started = time.monotonic() + try: + result = LangSDKTaskHandlerProcessorProcess.run( + coordinator=coordinator.key, + path=coordinator.bundle_path / candidate.rel_path, + bundle_path=coordinator.bundle_path, + bundle_name=coordinator.bundle_name, + artifact_rel_path=candidate.rel_path, + logger=log, + deadline=deadline, + ) + except Exception as e: + log.warning( + "Probing a task handler artifact raised", + seconds=round(time.monotonic() - started, 3), + exc_info=True, + **context, + ) + return f"probe failed: {_describe(e)}" + seconds = round(time.monotonic() - started, 3) + for warning in result.warnings or (): + log.warning("The task handler runtime reported a warning", warning=warning, **context) + if result.import_errors: + error = "; ".join(result.import_errors.values()) + log.warning( + "Probed a task handler artifact without an answer", error=error, seconds=seconds, **context + ) + if isinstance(result, TaskHandlerProbeStopped): + return f"probe stopped: {_OUT_OF_TIME}" + return f"probe failed: {error}" + try: + answer = TaskHandlerArtifact( + bundle_name=coordinator.bundle_name, + relative_fileloc=candidate.rel_path, + size_bytes=candidate.size_bytes, + cache_digest=candidate.cache_digest, + task_handlers=result.task_handlers, + ) + except ValidationError as e: + errors = "; ".join(f"{'.'.join(map(str, error['loc']))}: {error['msg']}" for error in e.errors()) + log.warning("Probed a task handler artifact with an invalid answer", error=errors, **context) + return f"invalid answer: {errors}" + log.info("Probed a task handler artifact", seconds=seconds, **context) + return answer + + +def _log_name_mismatch(warning: StubTaskWarning, log: FilteringBoundLogger) -> None: + context = { + "dag_id": warning.dag_id, + "task_id": warning.task_id, + "artifact_bundle_name": warning.artifact_bundle_name, + "artifact_rel_path": warning.artifact_rel_path, + } + if warning.passed_not_declared: + log.warning( + "Dag's call passed argument(s) the task handler does not declare", + passed_not_declared=warning.passed_not_declared, + **context, + ) + if warning.declared_not_passed: + log.warning( + "Task handler declares argument(s) the Dag's call did not pass", + declared_not_passed=warning.declared_not_passed, + **context, + ) + + +def _make_problems(stub_tasks: Iterable[StubTask], message: str) -> list[TaskHandlerProblem]: + """Return *message* as a problem of each Dag file with one of *stub_tasks*.""" + return [ + TaskHandlerProblem(relative_fileloc=relative_fileloc, message=message) + for relative_fileloc in sorted({stub.relative_fileloc for stub in stub_tasks}) + ] + + +def _describe(error: BaseException) -> str: + """Return *error* with its type, and the error it was raised from, if any.""" + description = f"{type(error).__name__}: {error}" + cause = error.__cause__ or (None if error.__suppress_context__ else error.__context__) + if cause is not None: + description = f"{description} ({type(cause).__name__}: {cause})" + return description + + +def _describe_team(team_name: str | None) -> str: + return "no team" if team_name is None else f"team {team_name!r}" diff --git a/airflow-core/src/airflow/dag_processing/task_handler_validation.py b/airflow-core/src/airflow/dag_processing/task_handler_validation.py new file mode 100644 index 0000000000000..dd8d20d6bf579 --- /dev/null +++ b/airflow-core/src/airflow/dag_processing/task_handler_validation.py @@ -0,0 +1,425 @@ +# 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. +"""Check a Dag file's stub tasks against the task handlers that their coordinator's artifacts register.""" + +from __future__ import annotations + +import json +from collections import defaultdict +from typing import TYPE_CHECKING + +import attrs + +from airflow.api_fastapi.execution_api.datamodels.task_arg_binding import ( + LiteralArgBinding, + get_arg_bindings_adapter, +) +from airflow.dag_processing.processor import TaskHandlerBinding +from airflow.sdk.definitions._internal.abstractoperator import DEFAULT_QUEUE # noqa: SDK001 +from airflow.serialization.enums import Encoding + +if TYPE_CHECKING: + from collections.abc import Iterable, Mapping, Sequence + + from airflow.api_fastapi.execution_api.datamodels.task_arg_binding import ArgValueSchema, TaskArgBinding + from airflow.dag_processing.processor import TaskHandlerArtifact, TaskHandlerDeclaration, TaskHandlerParam + from airflow.sdk.definitions.dag import DAG # noqa: SDK001 + from airflow.serialization.serialized_objects import LazyDeserializedDAG + +# Every top-level JSON type, in the order messages list them. +_JSON_TYPES = ("string", "integer", "number", "boolean", "array", "object", "null") + + +@attrs.frozen(kw_only=True) +class StubTask: + """A ``@task.stub`` task of a Dag that serialized, as the worker that runs it will see it.""" + + dag_id: str + task_id: str + queue: str + relative_fileloc: str + """The Dag's file, which keys its import error.""" + + arg_bindings: list[TaskArgBinding] + """The arguments of the call; ``[]`` for an argless call.""" + + is_mapped: bool + """A mapped stub task has no arguments captured at parse time, so only its handler's presence is checked.""" + + +@attrs.frozen(kw_only=True) +class ArgumentCheck: + """How a stub task's arguments bind to one task handler's parameters.""" + + errors: list[str] + """Each makes the Dag file fail to import.""" + + passed_not_declared: list[str] + """Under ``named`` binding, passed arguments no parameter takes; logged, never an error.""" + + declared_not_passed: list[str] + """Under ``named`` binding, parameters no argument fills; logged, never an error.""" + + +@attrs.frozen(kw_only=True) +class TaskHandlerProblem: + """One line of a Dag file's import error.""" + + relative_fileloc: str + message: str + dag_id: str | None = None + """``None`` for a problem of no single stub task, which is listed first.""" + + task_id: str | None = None + + +@attrs.frozen(kw_only=True) +class StubTaskWarning: + """A name mismatch between a stub task and its ``named`` task handler, for the parse log.""" + + dag_id: str + task_id: str + artifact_bundle_name: str + artifact_rel_path: str + passed_not_declared: list[str] + declared_not_passed: list[str] + + +@attrs.frozen(kw_only=True) +class TaskHandlerMatch: + """The outcome of checking stub tasks against one coordinator's answers.""" + + bindings: list[TaskHandlerBinding] + """One per stub task with exactly one handler, even when its arguments do not match.""" + + problems: list[TaskHandlerProblem] + warnings: list[StubTaskWarning] + + +def collect_stub_tasks(dags: Iterable[DAG], serialized_dags: Iterable[LazyDeserializedDAG]) -> list[StubTask]: + """ + Return the stub tasks of every Dag in *serialized_dags*. + + The arguments are read from the serialized Dag as JSON, so they are the ones the worker gets: a tuple + becomes a list, and a mapping key becomes a string. A stub task with no queue of its own runs on the + default queue. + """ + dags_by_id = {dag.dag_id: dag for dag in dags} + stub_tasks: list[StubTask] = [] + for serialized_dag in serialized_dags: + dag = dags_by_id[serialized_dag.dag_id] + encoded_tasks = { + task[Encoding.VAR]["task_id"]: task[Encoding.VAR] for task in serialized_dag.data["dag"]["tasks"] + } + stub_tasks.extend( + StubTask( + dag_id=dag.dag_id, + task_id=task.task_id, + queue=DEFAULT_QUEUE if task.queue is None else task.queue, + relative_fileloc=dag.relative_fileloc or dag.fileloc, + arg_bindings=get_arg_bindings_adapter().validate_json( + json.dumps(encoded_tasks[task.task_id].get("_arg_bindings") or []) + ), + is_mapped=task.is_mapped, + ) + for task in dag.task_dict.values() + if task.is_stub + ) + return stub_tasks + + +def check_task_handler_arguments( + arg_bindings: Sequence[TaskArgBinding], declaration: TaskHandlerDeclaration +) -> ArgumentCheck: + """ + Check how *arg_bindings* bind to the parameters of *declaration*, as its runtime binds them. + + A defaulted argument, one filled from the stub signature's default, is type-checked when it binds and is + never reported as unmatched. + + - ``positional``: the count matches when all arguments, or those left after dropping the defaulted + ones, number the parameters; any other count is an error. Each argument must have a value type the + parameter at its position accepts. + - ``named``: a parameter takes the argument of its exact name, else, unless ``exact_name`` is set, the one + whose name folds to its own (lower case, underscores removed); a folded name two arguments share + matches neither. Each bound argument must have a value type its parameter accepts. A passed argument + no parameter takes and a parameter no argument fills are reported for the parse log. They are not + when the argument may be the whole value: no parameter matched, exactly one argument was passed, the + handler has parameters and none of them has ``exact_name`` set, and the argument may be an object. + - ``params`` is ``None``: nothing is checked. + """ + params = declaration.params + if params is None: + return ArgumentCheck(errors=[], passed_not_declared=[], declared_not_passed=[]) + passed = [arg for arg in arg_bindings if not _is_defaulted(arg)] + if declaration.binding == "positional": + return _check_positional_arguments(arg_bindings, passed, params) + return _check_named_arguments(arg_bindings, passed, params) + + +def check_value_schema( + stub_schema: ArgValueSchema | None, handler_schema: ArgValueSchema | None +) -> str | None: + """ + Return why a value the stub's schema allows may be one the handler's schema rejects, or ``None``. + + Only the top-level JSON types are compared, as the Go runtime compares them. A stub type of ``null`` + needs a handler type of ``null``, and at least one of the stub's other types must be a handler type, + with ``integer`` accepted by ``number``. The types come from ``type``, the union over ``anyOf`` or + ``oneOf``, or the values of ``const`` or ``enum``. A side with no schema, or a schema of any other shape + such as ``$ref`` or ``allOf``, is not compared. ``format``, ranges and nested items are never compared. + """ + if stub_schema is None or handler_schema is None: + return None + stub_types = _get_json_types(stub_schema) + handler_types = _get_json_types(handler_schema) + if stub_types is None or handler_types is None: + return None + accepted = (handler_types | {"integer"}) if "number" in handler_types else handler_types + non_null_types = stub_types - {"null"} + if ("null" not in stub_types or "null" in handler_types) and ( + not non_null_types or non_null_types & accepted + ): + return None + return f"{_join_types(stub_types)}, the task handler takes {_join_types(handler_types)}" + + +def match_task_handlers( + stub_tasks: Sequence[StubTask], + answers: Sequence[TaskHandlerArtifact], + *, + bundle_name: str, + broken_candidates: Mapping[str, str], +) -> TaskHandlerMatch: + """ + Find the one task handler of each stub task among the *answers* of the coordinator it routes to. + + A stub task that no artifact registers is a problem naming each of *broken_candidates* (artifact path to + why it has no answer), and so is one that two or more artifacts register. A handler with no stub task is + not a problem. A stub task with exactly one handler is bound to its artifact, and its arguments are + checked unless it is mapped or the handler's ``params`` are ``None``. + """ + claims: defaultdict[tuple[str, str], list[tuple[TaskHandlerArtifact, TaskHandlerDeclaration]]] = ( + defaultdict(list) + ) + for artifact in answers: + for dag_id, declarations in artifact.task_handlers.items(): + for declaration in declarations: + claims[dag_id, declaration.task_id].append((artifact, declaration)) + + bindings: list[TaskHandlerBinding] = [] + problems: list[TaskHandlerProblem] = [] + warnings: list[StubTaskWarning] = [] + for stub in stub_tasks: + stub_claims = claims[stub.dag_id, stub.task_id] + prefix = f"Dag {stub.dag_id!r}, task {stub.task_id!r}" + if len(stub_claims) != 1: + if stub_claims: + paths = _join_quoted(sorted(artifact.relative_fileloc for artifact, _ in stub_claims)) + message = f"{prefix}: registered by {paths} in Dag bundle {bundle_name!r}" + else: + message = f"{prefix}: no artifact in Dag bundle {bundle_name!r} registers it" + if broken_candidates: + unanswered = ", ".join( + f"{path!r} ({why})" for path, why in sorted(broken_candidates.items()) + ) + message = f"{message}; no answer from {unanswered}" + problems.append(_make_problem(stub, message)) + continue + + [(artifact, declaration)] = stub_claims + bindings.append( + TaskHandlerBinding( + dag_id=stub.dag_id, + task_id=stub.task_id, + artifact_bundle_name=artifact.bundle_name, + artifact_rel_path=artifact.relative_fileloc, + ) + ) + if stub.is_mapped: + continue + check = check_task_handler_arguments(stub.arg_bindings, declaration) + located = f"{prefix} ({artifact.relative_fileloc!r} in Dag bundle {bundle_name!r})" + problems.extend(_make_problem(stub, f"{located}: {error}") for error in check.errors) + if check.passed_not_declared or check.declared_not_passed: + warnings.append( + StubTaskWarning( + dag_id=stub.dag_id, + task_id=stub.task_id, + artifact_bundle_name=artifact.bundle_name, + artifact_rel_path=artifact.relative_fileloc, + passed_not_declared=check.passed_not_declared, + declared_not_passed=check.declared_not_passed, + ) + ) + return TaskHandlerMatch(bindings=bindings, problems=problems, warnings=warnings) + + +def format_import_errors(problems: Iterable[TaskHandlerProblem]) -> dict[str, str]: + """ + Build one import error per Dag file from its *problems*. + + Problems of no single stub task come first, in the order given, then the others by Dag and task id. + """ + by_file: defaultdict[str, list[TaskHandlerProblem]] = defaultdict(list) + for problem in problems: + by_file[problem.relative_fileloc].append(problem) + import_errors: dict[str, str] = {} + for relative_fileloc in sorted(by_file): + ordered = sorted( + by_file[relative_fileloc], + key=lambda problem: (problem.dag_id is not None, problem.dag_id or "", problem.task_id or ""), + ) + lines = [f"Stub tasks in {relative_fileloc} do not match their task handlers:"] + lines.extend(f"- {problem.message}" for problem in ordered) + import_errors[relative_fileloc] = "\n".join(lines) + return import_errors + + +def _check_positional_arguments( + arg_bindings: Sequence[TaskArgBinding], + passed: Sequence[TaskArgBinding], + params: Sequence[TaskHandlerParam], +) -> ArgumentCheck: + bound = arg_bindings if len(arg_bindings) == len(params) else passed + if len(bound) == len(params): + errors = _check_values(zip(bound, params)) + else: + count = f"{len(arg_bindings)} argument{'' if len(arg_bindings) == 1 else 's'}" + if len(passed) != len(arg_bindings): + count = f"{count} ({len(passed)} without defaults)" + errors = [f"passes {count}, the task handler takes {len(params)}"] + return ArgumentCheck(errors=errors, passed_not_declared=[], declared_not_passed=[]) + + +def _check_named_arguments( + arg_bindings: Sequence[TaskArgBinding], + passed: Sequence[TaskArgBinding], + params: Sequence[TaskHandlerParam], +) -> ArgumentCheck: + by_name = {binding.name: binding for binding in arg_bindings} + by_folded_name: defaultdict[str, list[TaskArgBinding]] = defaultdict(list) + for binding in arg_bindings: + by_folded_name[_fold_name(binding.name)].append(binding) + matched: list[tuple[TaskArgBinding, TaskHandlerParam]] = [] + declared_not_passed: list[str] = [] + for index, param in enumerate(params): + if param.name is None: + declared_not_passed.append(f"#{index}") + continue + arg: TaskArgBinding | None = by_name.get(param.name) + if arg is None and not param.exact_name: + candidates = by_folded_name[_fold_name(param.name)] + arg = candidates[0] if len(candidates) == 1 else None + if arg is None: + declared_not_passed.append(param.name) + else: + matched.append((arg, param)) + if not matched and _may_be_whole_value(passed, params): + return ArgumentCheck(errors=[], passed_not_declared=[], declared_not_passed=[]) + claimed = {arg.name for arg, _ in matched} + return ArgumentCheck( + errors=_check_values(matched), + passed_not_declared=[arg.name for arg in passed if arg.name not in claimed], + declared_not_passed=declared_not_passed, + ) + + +def _may_be_whole_value(passed: Sequence[TaskArgBinding], params: Sequence[TaskHandlerParam]) -> bool: + if len(passed) != 1 or not params or any(param.exact_name for param in params): + return False + schema = passed[0].value_schema + types = None if schema is None else _get_json_types(schema) + return types is None or "object" in types + + +def _is_defaulted(arg: TaskArgBinding) -> bool: + return isinstance(arg, LiteralArgBinding) and arg.from_default + + +def _fold_name(name: str) -> str: + return name.replace("_", "").lower() + + +def _check_values(pairs: Iterable[tuple[TaskArgBinding, TaskHandlerParam]]) -> list[str]: + return [ + f"argument {arg.name!r} is {mismatch}" + for arg, param in pairs + if (mismatch := check_value_schema(arg.value_schema, param.value_schema)) is not None + ] + + +def _get_json_types(schema: Mapping[str, object]) -> frozenset[str] | None: + if "type" in schema: + names = schema["type"] + if isinstance(names, str): + names = [names] + if not isinstance(names, list) or not names or not all(name in _JSON_TYPES for name in names): + return None + return frozenset(names) + for keyword in ("anyOf", "oneOf"): + if keyword in schema: + branches = schema[keyword] + if not isinstance(branches, list) or not branches: + return None + union: set[str] = set() + for branch in branches: + if not isinstance(branch, dict) or (types := _get_json_types(branch)) is None: + return None + union |= types + return frozenset(union) + if "const" in schema: + return frozenset([_get_value_type(schema["const"])]) + if isinstance(values := schema.get("enum"), list) and values: + return frozenset(_get_value_type(value) for value in values) + return None + + +def _get_value_type(value: object) -> str: + if value is None: + return "null" + if isinstance(value, bool): + return "boolean" + if isinstance(value, int): + return "integer" + if isinstance(value, float): + return "number" + if isinstance(value, str): + return "string" + if isinstance(value, list): + return "array" + return "object" + + +def _join_types(types: frozenset[str]) -> str: + return _join([name for name in _JSON_TYPES if name in types], "or") + + +def _join_quoted(values: Sequence[str]) -> str: + return _join([repr(value) for value in values], "and") + + +def _join(values: Sequence[str], conjunction: str) -> str: + if len(values) == 1: + return values[0] + return f"{', '.join(values[:-1])} {conjunction} {values[-1]}" + + +def _make_problem(stub: StubTask, message: str) -> TaskHandlerProblem: + return TaskHandlerProblem( + relative_fileloc=stub.relative_fileloc, message=message, dag_id=stub.dag_id, task_id=stub.task_id + ) diff --git a/airflow-core/src/airflow/executors/base_executor.py b/airflow-core/src/airflow/executors/base_executor.py index b99ad38a4d494..4b1fd77ef884f 100644 --- a/airflow-core/src/airflow/executors/base_executor.py +++ b/airflow-core/src/airflow/executors/base_executor.py @@ -890,7 +890,8 @@ def run_workload( if isinstance(workload, ExecuteTask): from airflow.sdk.execution_time.supervisor import supervise_task - # workload.ti is a TaskInstanceDTO which duck-types as TaskInstance. + # workload.ti is a TaskInstanceDTO which duck-types as TaskInstance, and + # workload.task_handler_artifact duck-types as the Task SDK's TaskHandlerArtifactRef. # TODO: Create a protocol for this. return supervise_task( ti=workload.ti, # type: ignore[arg-type] @@ -902,6 +903,7 @@ def run_workload( log_path=workload.log_path, subprocess_logs_to_stdout=subprocess_logs_to_stdout, sentry_integration=getattr(workload, "sentry_integration", ""), + task_handler_artifact=workload.task_handler_artifact, # type: ignore[arg-type] ) if isinstance(workload, ExecuteCallback): from airflow.sdk.execution_time.callback_supervisor import supervise_callback diff --git a/airflow-core/src/airflow/executors/workloads/__init__.py b/airflow-core/src/airflow/executors/workloads/__init__.py index 57f28d2f94960..4f4d1f3702b1e 100644 --- a/airflow-core/src/airflow/executors/workloads/__init__.py +++ b/airflow-core/src/airflow/executors/workloads/__init__.py @@ -25,7 +25,7 @@ from airflow.executors.workloads.base import WORKLOAD_TYPE_PRIORITY, BaseWorkload, BundleInfo, WorkloadType from airflow.executors.workloads.callback import CallbackFetchMethod, ExecuteCallback from airflow.executors.workloads.connection_test import TestConnection -from airflow.executors.workloads.task import ExecuteTask, TaskInstanceDTO +from airflow.executors.workloads.task import ExecuteTask, TaskHandlerArtifactRef, TaskInstanceDTO from airflow.executors.workloads.trigger import RunTrigger All = Annotated[ @@ -49,6 +49,7 @@ "ExecuteCallback", "ExecuteTask", "ExecutorWorkload", + "TaskHandlerArtifactRef", "TaskInstance", "TaskInstanceDTO", "TestConnection", diff --git a/airflow-core/src/airflow/executors/workloads/task.py b/airflow-core/src/airflow/executors/workloads/task.py index 5b1274d35b7ae..0ddad727ae900 100644 --- a/airflow-core/src/airflow/executors/workloads/task.py +++ b/airflow-core/src/airflow/executors/workloads/task.py @@ -21,7 +21,7 @@ from pathlib import Path from typing import TYPE_CHECKING, Literal -from pydantic import Field +from pydantic import BaseModel, Field from airflow.api_fastapi.execution_api.datamodels.taskinstance import TaskInstance from airflow.executors.workloads.base import BaseDagBundleWorkload, BundleInfo, WorkloadType @@ -61,11 +61,26 @@ def key(self) -> TaskInstanceKey: ) +class TaskHandlerArtifactRef(BaseModel): + """The Lang-SDK artifact that implements a stub task: its Dag bundle and its path in that bundle.""" + + bundle_info: BundleInfo | None = None + """ + The Dag bundle holding the artifact; ``None`` means the task's own Dag bundle, at the version the + run uses: its pinned version, or the version current when the task starts if the run is not pinned. + """ + + rel_path: str + """POSIX path of the artifact within that bundle.""" + + class ExecuteTask(BaseDagBundleWorkload): """Execute the given Task.""" ti: TaskInstanceDTO sentry_integration: str = "" + task_handler_artifact: TaskHandlerArtifactRef | None = None + """The artifact that implements this stub task; ``None`` when the workload names none.""" type: Literal[WorkloadType.EXECUTE_TASK] = Field(init=False, default=WorkloadType.EXECUTE_TASK) diff --git a/airflow-core/src/airflow/migrations/versions/0142_3_4_0_add_lang_sdk_task_handler_tables.py b/airflow-core/src/airflow/migrations/versions/0142_3_4_0_add_lang_sdk_task_handler_tables.py new file mode 100644 index 0000000000000..c9679ede25bec --- /dev/null +++ b/airflow-core/src/airflow/migrations/versions/0142_3_4_0_add_lang_sdk_task_handler_tables.py @@ -0,0 +1,91 @@ +# +# 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 lang_sdk_task_handler_artifact and lang_sdk_task_handler tables. + +Revision ID: f7ed13533d23 +Revises: 90e4d18ccadf +Create Date: 2026-09-30 14:19:35.109799 + +""" + +from __future__ import annotations + +import sqlalchemy as sa +from alembic import op + +from airflow.migrations.db_types import StringID +from airflow.utils.sqlalchemy import UtcDateTime + +# revision identifiers, used by Alembic. +revision = "f7ed13533d23" +down_revision = "90e4d18ccadf" +branch_labels = None +depends_on = None +airflow_version = "3.4.0" + + +def upgrade(): + """Add lang_sdk_task_handler_artifact and lang_sdk_task_handler tables.""" + op.create_table( + "lang_sdk_task_handler_artifact", + sa.Column("id", sa.Uuid(), nullable=False), + sa.Column("bundle_name", StringID(), nullable=False), + sa.Column("relative_fileloc", sa.String(length=2000), nullable=False), + sa.Column("relative_fileloc_hash", sa.String(length=32), nullable=False), + sa.Column("size_bytes", sa.BigInteger(), nullable=False), + sa.Column("cache_digest", sa.String(length=128), nullable=True), + sa.Column("task_handlers", sa.JSON(), nullable=False), + sa.Column("last_probed_at", UtcDateTime(), nullable=False), + sa.PrimaryKeyConstraint("id", name=op.f("lang_sdk_task_handler_artifact_pkey")), + sa.UniqueConstraint( + "bundle_name", + "relative_fileloc_hash", + name=op.f("lang_sdk_task_handler_artifact_bundle_fileloc_uq"), + ), + ) + op.create_table( + "lang_sdk_task_handler", + sa.Column("dag_id", StringID(), nullable=False), + sa.Column("task_id", StringID(), nullable=False), + sa.Column("artifact_id", sa.Uuid(), nullable=False), + sa.Column("dag_bundle_name", StringID(), nullable=False), + sa.Column("dag_relative_fileloc", sa.String(length=2000), nullable=False), + sa.Column("dag_relative_fileloc_hash", sa.String(length=32), nullable=False), + sa.PrimaryKeyConstraint("dag_id", "task_id", name=op.f("lang_sdk_task_handler_pkey")), + sa.ForeignKeyConstraint( + columns=("dag_id",), + refcolumns=["dag.dag_id"], + name="lang_sdk_task_handler_dag_id_fkey", + ondelete="CASCADE", + ), + sa.ForeignKeyConstraint( + columns=("artifact_id",), + refcolumns=["lang_sdk_task_handler_artifact.id"], + name="lang_sdk_task_handler_artifact_id_fkey", + ), + sa.Index("idx_lang_sdk_task_handler_dag_file", "dag_bundle_name", "dag_relative_fileloc_hash"), + sa.Index("idx_lang_sdk_task_handler_artifact_id", "artifact_id"), + ) + + +def downgrade(): + """Drop lang_sdk_task_handler and lang_sdk_task_handler_artifact tables.""" + op.drop_table("lang_sdk_task_handler") + op.drop_table("lang_sdk_task_handler_artifact") diff --git a/airflow-core/src/airflow/migrations/versions/0143_3_4_0_add_last_probed_at_index_to_lang_sdk_task_handler_artifact.py b/airflow-core/src/airflow/migrations/versions/0143_3_4_0_add_last_probed_at_index_to_lang_sdk_task_handler_artifact.py new file mode 100644 index 0000000000000..8eb9c65d2b596 --- /dev/null +++ b/airflow-core/src/airflow/migrations/versions/0143_3_4_0_add_last_probed_at_index_to_lang_sdk_task_handler_artifact.py @@ -0,0 +1,51 @@ +# +# 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 last_probed_at index to lang_sdk_task_handler_artifact. + +Revision ID: 37d645374a9c +Revises: f7ed13533d23 +Create Date: 2026-09-30 19:00:42.253263 + +""" + +from __future__ import annotations + +from alembic import op + +# revision identifiers, used by Alembic. +revision = "37d645374a9c" +down_revision = "f7ed13533d23" +branch_labels = None +depends_on = None +airflow_version = "3.4.0" + + +def upgrade(): + """Add last_probed_at index to lang_sdk_task_handler_artifact.""" + with op.batch_alter_table("lang_sdk_task_handler_artifact", schema=None) as batch_op: + batch_op.create_index( + "idx_lang_sdk_task_handler_artifact_last_probed_at", ["last_probed_at"], unique=False + ) + + +def downgrade(): + """Remove last_probed_at index from lang_sdk_task_handler_artifact.""" + with op.batch_alter_table("lang_sdk_task_handler_artifact", schema=None) as batch_op: + batch_op.drop_index("idx_lang_sdk_task_handler_artifact_last_probed_at") diff --git a/airflow-core/src/airflow/models/__init__.py b/airflow-core/src/airflow/models/__init__.py index f5bffb4e57173..f806537c684df 100644 --- a/airflow-core/src/airflow/models/__init__.py +++ b/airflow-core/src/airflow/models/__init__.py @@ -71,6 +71,7 @@ def import_all_models(): import airflow.models.dagwarning import airflow.models.deadline_alert import airflow.models.errors + import airflow.models.lang_sdk_task_handler import airflow.models.revoked_token import airflow.models.serialized_dag import airflow.models.task_state_store diff --git a/airflow-core/src/airflow/models/lang_sdk_task_handler.py b/airflow-core/src/airflow/models/lang_sdk_task_handler.py new file mode 100644 index 0000000000000..ad8ad1ef59203 --- /dev/null +++ b/airflow-core/src/airflow/models/lang_sdk_task_handler.py @@ -0,0 +1,129 @@ +# +# 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 collections.abc import Callable +from datetime import datetime +from typing import TYPE_CHECKING, Any +from uuid import UUID + +import sqlalchemy as sa +import uuid6 +from sqlalchemy import BigInteger, ForeignKeyConstraint, Index, String, UniqueConstraint, Uuid +from sqlalchemy.orm import Mapped, mapped_column, validates + +from airflow._shared.timezones import timezone +from airflow.models.base import Base, StringID +from airflow.utils.hashlib_wrapper import md5 +from airflow.utils.sqlalchemy import UtcDateTime + +if TYPE_CHECKING: + from sqlalchemy.engine.default import DefaultExecutionContext + + +def compute_fileloc_hash(relative_fileloc: str) -> str: + """ + Return the md5 hex digest that the task handler tables index in place of a relative file location. + + A 2000-character path exceeds MySQL's 3072-byte key limit and Postgres's 2704-byte btree entry limit. + """ + return md5(relative_fileloc.encode()).hexdigest() + + +# The column default fills a hash on INSERT, ORM or Core. ORM updates resync it in each model's +# @validates hook instead of an onupdate default, which would fire on every UPDATE, even one that +# leaves the path alone. A Core UPDATE or upsert that changes a path must set its hash itself. +def _build_fileloc_hash_default(path_key: str) -> Callable[[DefaultExecutionContext], str]: + def compute(context: DefaultExecutionContext) -> str: + return compute_fileloc_hash(context.get_current_parameters()[path_key]) + + return compute + + +class LangSDKTaskHandlerArtifact(Base): + """ + A Language SDK artifact in a task handler Dag bundle, cached with its fingerprint and probe answer. + + ``task_handlers`` is the runtime's whole answer, ``{dag_id: [TaskHandlerDeclaration as JSON, ...]}``. + It is written only from a probe and together with the fingerprint, so the two always match. + """ + + __tablename__ = "lang_sdk_task_handler_artifact" + + id: Mapped[UUID] = mapped_column(Uuid(), primary_key=True, default=uuid6.uuid7) + bundle_name: Mapped[str] = mapped_column(StringID(), nullable=False) + relative_fileloc: Mapped[str] = mapped_column(String(2000), nullable=False) + relative_fileloc_hash: Mapped[str] = mapped_column( + String(32), nullable=False, default=_build_fileloc_hash_default("relative_fileloc") + ) + size_bytes: Mapped[int] = mapped_column(BigInteger, nullable=False) + # NULL means the artifact stores no fingerprint and is probed again on every parse. + cache_digest: Mapped[str | None] = mapped_column(String(128), nullable=True) + task_handlers: Mapped[dict[str, list[dict[str, Any]]]] = mapped_column(sa.JSON(), nullable=False) + last_probed_at: Mapped[datetime] = mapped_column(UtcDateTime, nullable=False, default=timezone.utcnow) + + __table_args__ = ( + UniqueConstraint( + bundle_name, + relative_fileloc_hash, + name="lang_sdk_task_handler_artifact_bundle_fileloc_uq", + ), + Index("idx_lang_sdk_task_handler_artifact_last_probed_at", last_probed_at), + ) + + @validates("relative_fileloc") + def _sync_relative_fileloc_hash(self, key: str, relative_fileloc: str) -> str: + self.relative_fileloc_hash = compute_fileloc_hash(relative_fileloc) + return relative_fileloc + + +class LangSDKTaskHandler(Base): + """The binding of one stub task, by ``(dag_id, task_id)``, to the artifact that registers its handler.""" + + __tablename__ = "lang_sdk_task_handler" + + dag_id: Mapped[str] = mapped_column(StringID(), primary_key=True) + task_id: Mapped[str] = mapped_column(StringID(), primary_key=True) + artifact_id: Mapped[UUID] = mapped_column(Uuid(), nullable=False) + dag_bundle_name: Mapped[str] = mapped_column(StringID(), nullable=False) + dag_relative_fileloc: Mapped[str] = mapped_column(String(2000), nullable=False) + dag_relative_fileloc_hash: Mapped[str] = mapped_column( + String(32), nullable=False, default=_build_fileloc_hash_default("dag_relative_fileloc") + ) + + __table_args__ = ( + ForeignKeyConstraint( + (dag_id,), + ["dag.dag_id"], + name="lang_sdk_task_handler_dag_id_fkey", + ondelete="CASCADE", + ), + # No ON DELETE: deleting an artifact that a handler still references must fail. + ForeignKeyConstraint( + (artifact_id,), + ["lang_sdk_task_handler_artifact.id"], + name="lang_sdk_task_handler_artifact_id_fkey", + ), + Index("idx_lang_sdk_task_handler_dag_file", dag_bundle_name, dag_relative_fileloc_hash), + Index("idx_lang_sdk_task_handler_artifact_id", artifact_id), + ) + + @validates("dag_relative_fileloc") + def _sync_dag_relative_fileloc_hash(self, key: str, dag_relative_fileloc: str) -> str: + self.dag_relative_fileloc_hash = compute_fileloc_hash(dag_relative_fileloc) + return dag_relative_fileloc diff --git a/airflow-core/src/airflow/utils/db.py b/airflow-core/src/airflow/utils/db.py index f160cb0597d84..3a489bfe4bf76 100644 --- a/airflow-core/src/airflow/utils/db.py +++ b/airflow-core/src/airflow/utils/db.py @@ -117,7 +117,7 @@ class MappedClassProtocol(Protocol): "3.1.8": "509b94a1042d", "3.2.0": "1d6611b6ab7c", "3.3.0": "d2f4e1b3c5a7", - "3.4.0": "90e4d18ccadf", + "3.4.0": "37d645374a9c", } # Prefix used to identify tables holding data moved during migration. diff --git a/airflow-core/src/airflow/utils/file.py b/airflow-core/src/airflow/utils/file.py index ff7d7db16623b..12fc294397480 100644 --- a/airflow-core/src/airflow/utils/file.py +++ b/airflow-core/src/airflow/utils/file.py @@ -110,7 +110,9 @@ def find_dag_file_paths(directory: str | os.PathLike[str], safe_mode: bool) -> l for file_path in find_path_from_directory(directory, ".airflowignore", ignore_file_syntax): path = Path(file_path) try: - if path.is_file() and (path.suffix == ".py" or zipfile.is_zipfile(path)): + if path.is_file() and ( + path.suffix == ".py" or (path.suffix.lower() == ".zip" and zipfile.is_zipfile(path)) + ): if might_contain_dag(file_path, safe_mode, conf=conf): file_paths.append(file_path) except Exception: diff --git a/airflow-core/tests/unit/dag_processing/bundles/test_dag_bundle_manager.py b/airflow-core/tests/unit/dag_processing/bundles/test_dag_bundle_manager.py index 6eac0c575fbb0..d607863d9ec9e 100644 --- a/airflow-core/tests/unit/dag_processing/bundles/test_dag_bundle_manager.py +++ b/airflow-core/tests/unit/dag_processing/bundles/test_dag_bundle_manager.py @@ -181,6 +181,24 @@ def test_get_configured_bundle_team_names_without_config(): assert _get_configured_bundle_team_names() == {} +@conf_vars( + { + ("core", "multi_team"): "True", + ("dag_processor", "dag_bundle_config_list"): json.dumps( + [*TEAM_BUNDLE_CONFIG, {**TEAM_BUNDLE_CONFIG[1], "name": "empty-team-bundle", "team_name": ""}] + ), + } +) +def test_get_bundle_team_name(): + bundle_manager = DagBundlesManager() + + assert bundle_manager.get_bundle_team_name("team-bundle") == "team-a" + assert bundle_manager.get_bundle_team_name("unscoped-bundle") is None + assert bundle_manager.get_bundle_team_name("empty-team-bundle") is None + with pytest.raises(ValueError, match="'other-test-bundle' is not configured"): + bundle_manager.get_bundle_team_name("other-test-bundle") + + def test_get_bundle(): """Test that get_bundle builds and returns a bundle.""" with patch.dict( diff --git a/airflow-core/tests/unit/dag_processing/fake_task_handler_runtime.py b/airflow-core/tests/unit/dag_processing/fake_task_handler_runtime.py new file mode 100644 index 0000000000000..dca237139fdbc --- /dev/null +++ b/airflow-core/tests/unit/dag_processing/fake_task_handler_runtime.py @@ -0,0 +1,227 @@ +# +# 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. +"""A coordinator whose artifacts name the command that parses them, and a runtime a test can play.""" + +from __future__ import annotations + +import contextlib +import json +import os +import socket +from pathlib import Path +from typing import TYPE_CHECKING, Any +from unittest import mock + +import attrs +import structlog +from pydantic import TypeAdapter + +from airflow.dag_processing.processor import ( + DagFileParseRequest, + DagFileParsingResult, + TaskHandlerArtifact, + TaskHandlerBinding, + TaskHandlerParseRequest, + TaskHandlerParsingResult, + ToManager, + ToSDKTaskHandlerProcessor, + _parse_file, +) +from airflow.sdk.coordinators._subprocess import SubprocessCoordinator +from airflow.sdk.execution_time import supervisor +from airflow.sdk.execution_time.comms import CommsDecoder +from airflow.sdk.execution_time.coordinator import TaskHandlerCandidate, reset_coordinator_manager + +from tests_common.test_utils.config import conf_vars + +if TYPE_CHECKING: + from collections.abc import Callable, Iterator, Sequence + + +@attrs.define(kw_only=True) +class FakeCoordinator(SubprocessCoordinator): + """ + The JSON of an ``*.artifact`` file names the command that parses it. + + ``argv`` is the command, ``schema_version`` its schema version, and ``command_error`` an error to + raise instead. The listing reads the stored ``cache_digest``, and ``listing_error`` as why the + artifact cannot be probed. ``task_handlers`` is the answer :func:`reply_with_task_handlers` sends. + """ + + def _build_parse_task_handler_command(self, *, path: Path) -> tuple[list[str], str | None]: + spec = json.loads(path.read_text()) + if error := spec.get("command_error"): + raise FileNotFoundError(error) + return spec.get("argv", ["/bin/false"]), spec.get("schema_version") + + def _read_task_handler_candidate(self, path: Path, *, rel_path: str) -> TaskHandlerCandidate | None: + if path.suffix != ".artifact": + return None + content = path.read_bytes() + spec = json.loads(content) + return TaskHandlerCandidate( + rel_path=rel_path, + size_bytes=len(content), + cache_digest=spec.get("cache_digest"), + error=spec.get("listing_error"), + ) + + +@contextlib.contextmanager +def fake_coordinator(**kwargs: Any) -> Iterator[None]: + """ + Configure a ``FakeCoordinator`` as ``fake``, with fresh coordinators inside and after the block. + + The parse child is a bare fork even on macOS, so it sees the test's patches and config. + """ + spec = {"fake": {"classpath": f"{__name__}.FakeCoordinator", "kwargs": kwargs}} + reset_coordinator_manager() + try: + with ( + conf_vars({("sdk", "coordinators"): json.dumps(spec)}), + mock.patch.object(supervisor, "_should_use_exec", return_value=False), + ): + yield + finally: + reset_coordinator_manager() + + +FAKE_COORDINATOR = f"{__name__}.FakeCoordinator" +LOCAL_BUNDLE = "airflow.dag_processing.bundles.local.LocalDagBundle" + + +@contextlib.contextmanager +def task_handler_config( + dag_bundle: Path, + artifacts: Path, + coordinators: dict[str, Any] | None = None, + *, + queue_to_coordinator: dict[str, str] | None = None, + bundles: list[dict[str, Any]] | None = None, + multi_team: bool = False, +) -> Iterator[None]: + """ + Route the queue ``fake-queue`` to a ``FakeCoordinator`` reading the ``task-handlers`` Dag bundle. + + By default *dag_bundle* is the Dag bundle ``dags`` and *artifacts* is ``task-handlers``; the other + arguments replace those parts of the configuration. Coordinators are fresh inside and after the block. + """ + if coordinators is None: + coordinators = { + "fake": {"classpath": FAKE_COORDINATOR, "kwargs": {"task_handler_bundle_name": "task-handlers"}} + } + if bundles is None: + bundles = [ + {"name": "dags", "classpath": LOCAL_BUNDLE, "kwargs": {"path": os.fspath(dag_bundle)}}, + {"name": "task-handlers", "classpath": LOCAL_BUNDLE, "kwargs": {"path": os.fspath(artifacts)}}, + ] + reset_coordinator_manager() + try: + with conf_vars( + { + ("core", "load_examples"): "False", + ("core", "multi_team"): str(multi_team), + ("dag_processor", "dag_bundle_config_list"): json.dumps(bundles), + ("sdk", "coordinators"): json.dumps(coordinators), + ("sdk", "queue_to_coordinator"): json.dumps( + {"fake-queue": "fake"} if queue_to_coordinator is None else queue_to_coordinator + ), + } + ): + yield + finally: + reset_coordinator_manager() + + +def parse_dag_file( + dag_file: Path, *, known_artifacts: Sequence[TaskHandlerArtifact] = () +) -> DagFileParsingResult: + """Parse *dag_file* as the Dag processor's child does, in the Dag bundle ``dags`` at its directory.""" + request = DagFileParseRequest( + file=os.fspath(dag_file), + bundle_path=dag_file.parent, + bundle_name="dags", + known_artifacts=list(known_artifacts), + ) + result = _parse_file(request, log=structlog.get_logger()) + assert result is not None + return result + + +def get_stub_task_ids(result: DagFileParsingResult) -> set[tuple[str, str]]: + """Return the Dag and task id of every stub task in the serialized Dags of *result*.""" + return { + (dag.dag_id, task["__var"]["task_id"]) + for dag in result.serialized_dags + for task in dag.data["dag"]["tasks"] + if task["__var"].get("is_stub") + } + + +def sort_bindings(result: DagFileParsingResult) -> list[TaskHandlerBinding]: + """Return the bindings of *result* by Dag and task id; a second import of a file can reorder its Dags.""" + return sorted(result.task_handler_bindings or [], key=lambda binding: (binding.dag_id, binding.task_id)) + + +def write_artifact(path: Path, **spec: Any) -> Path: + path.write_text(json.dumps(spec)) + return path + + +def reply_with_task_handlers( + request: TaskHandlerParseRequest, comms: CommsDecoder | None +) -> TaskHandlerParsingResult: + """Answer with the ``task_handlers`` of the artifact's JSON, as declarations in their wire form.""" + spec = json.loads(Path(request.file).read_text()) + return TaskHandlerParsingResult.model_validate( + {"fileloc": request.file, "task_handlers": spec.get("task_handlers", {})} + ) + + +def play_runtime( + reply: Callable[[TaskHandlerParseRequest, CommsDecoder], TaskHandlerParsingResult | None], + *, + schema_version: str | None = None, + log_lines: Sequence[dict[str, Any]] = (), +) -> Callable[..., None]: + """ + Return a ``parse_task_handler`` that plays the runtime in the parse child instead of exec'ing one. + + It connects back as a runtime does, writes *log_lines* to its logs channel and answers the parse + request with what *reply* returns; ``None`` sends no result. + """ + + def parse_task_handler( + coordinator, *, comm_address, logs_address, report_schema_version, **kwargs + ) -> None: + report_schema_version(schema_version) + with ( + socket.create_connection(comm_address) as comm, + socket.create_connection(logs_address) as logs, + ): + for line in log_lines: + logs.sendall(json.dumps(line).encode() + b"\n") + comms = CommsDecoder[ToSDKTaskHandlerProcessor, ToManager]( + socket=comm, body_decoder=TypeAdapter(ToSDKTaskHandlerProcessor) + ) + request = comms._get_response() + assert isinstance(request, TaskHandlerParseRequest) + if (result := reply(request, comms)) is not None: + comms.send(result) + + return parse_task_handler diff --git a/airflow-core/tests/unit/dag_processing/test_collection.py b/airflow-core/tests/unit/dag_processing/test_collection.py index d7e5213fad5f1..1be65b46f6dcb 100644 --- a/airflow-core/tests/unit/dag_processing/test_collection.py +++ b/airflow-core/tests/unit/dag_processing/test_collection.py @@ -18,11 +18,13 @@ from __future__ import annotations +import contextlib import importlib import logging import os import sys import textwrap +import uuid import warnings from collections.abc import Generator from datetime import timedelta @@ -32,7 +34,7 @@ import pytest from sqlalchemy import delete, event, func, inspect as sa_inspect, select -from sqlalchemy.exc import OperationalError, SAWarning +from sqlalchemy.exc import DataError, OperationalError, SAWarning import airflow.dag_processing.collection from airflow import plugins_manager @@ -46,8 +48,15 @@ _get_latest_runs_stmt_partitioned, _update_dag_tags, _update_import_errors, + record_probed_task_handler_artifacts, update_dag_parsing_results_in_db, ) +from airflow.dag_processing.processor import ( + TaskHandlerArtifact, + TaskHandlerBinding, + TaskHandlerDeclaration, + TaskHandlerParam, +) from airflow.example_dags.plugins.business_day_window import BusinessDayWindow from airflow.example_dags.plugins.custom_partition_mapper import PrefixStripMapper from airflow.example_dags.plugins.workday import AfterWorkdayTimetable @@ -64,6 +73,11 @@ from airflow.models.dagcode import DagCode from airflow.models.dagwarning import DagWarning, DagWarningType from airflow.models.errors import ParseImportError +from airflow.models.lang_sdk_task_handler import ( + LangSDKTaskHandler, + LangSDKTaskHandlerArtifact, + compute_fileloc_hash, +) from airflow.models.serialized_dag import SerializedDagModel from airflow.models.trigger import Trigger from airflow.partition_mappers.base import RollupMapper @@ -91,6 +105,7 @@ from airflow.timetables.simple import PartitionedAtRuntime from airflow.timetables.trigger import CronTriggerTimetable from airflow.triggers.base import BaseEventTrigger +from airflow.utils.session import create_session from airflow.utils.types import DagRunType from tests_common.test_utils.config import conf_vars @@ -1946,6 +1961,27 @@ def test_team_class_is_only_stored_for_its_team( assert list(errors) == [(bundle_name, "team_dag.py")] assert f"belonging to {plugin_team}" in errors[(bundle_name, "team_dag.py")] + @conf_vars({("core", "multi_team"): "True"}) + def test_task_handler_bindings_of_a_rejected_dag_are_dropped_quietly(self, bundle, session, caplog): + bundle_name = bundle(True) + _record(session, "etl.jar") + with mock_plugin_manager(plugins=[self._plugin("other_team", "timetables", AfterWorkdayTimetable)]): + update_dag_parsing_results_in_db( + bundle_name=bundle_name, + bundle_version=None, + dags=[self._serialized(schedule=AfterWorkdayTimetable()), self._serialized("plain_dag")], + import_errors={}, + parse_duration=None, + warnings=set(), + session=session, + relative_fileloc="team_dag.py", + task_handler_bindings=[_make_binding("team_dag"), _make_binding("plain_dag")], + task_handler_artifact_bundles={ARTIFACT_BUNDLE}, + ) + + assert session.scalars(select(LangSDKTaskHandler.dag_id)).all() == ["plain_dag"] + assert not any("task handler bindings" in entry["event"] for entry in caplog.entries) + @conf_vars({("core", "multi_team"): "True"}) @pytest.mark.parametrize( "weight_rule", @@ -2107,3 +2143,457 @@ class WorkdayPlugin(AirflowPlugin): assert stored == set() assert "belonging to other_team" in errors[(bundle_name, "team_dag.py")] + + +ARTIFACT_BUNDLE = "java-task-handlers" +OTHER_TEAM_BUNDLE = "go-task-handlers" +TASK_HANDLERS = { + "etl": [ + TaskHandlerDeclaration(task_id="extract", binding="positional", params=[TaskHandlerParam(name=None)]) + ] +} + + +def _make_artifact( + rel_path: str = "etl.jar", + *, + bundle_name: str = ARTIFACT_BUNDLE, + size_bytes: int = 1024, + cache_digest: str | None = "a" * 64, + task_handlers: dict[str, list[TaskHandlerDeclaration]] | None = None, +) -> TaskHandlerArtifact: + return TaskHandlerArtifact( + bundle_name=bundle_name, + relative_fileloc=rel_path, + size_bytes=size_bytes, + cache_digest=cache_digest, + task_handlers=TASK_HANDLERS if task_handlers is None else task_handlers, + ) + + +def _record(session, *rel_paths: str, bundle_name: str = ARTIFACT_BUNDLE) -> None: + record_probed_task_handler_artifacts( + [_make_artifact(rel_path, bundle_name=bundle_name) for rel_path in rel_paths], + bundle_names={bundle_name}, + session=session, + ) + + +def _make_binding( + dag_id: str = "etl", + task_id: str = "extract", + *, + rel_path: str = "etl.jar", + bundle_name: str = ARTIFACT_BUNDLE, +) -> TaskHandlerBinding: + return TaskHandlerBinding( + dag_id=dag_id, task_id=task_id, artifact_bundle_name=bundle_name, artifact_rel_path=rel_path + ) + + +def _persist(session, *dag_ids: str, relative_fileloc: str = "etl.py", bindings) -> None: + update_dag_parsing_results_in_db( + bundle_name="testing", + bundle_version=None, + dags=[LazyDeserializedDAG.from_dag(DAG(dag_id=dag_id)) for dag_id in dag_ids], + import_errors={}, + parse_duration=None, + warnings=set(), + session=session, + relative_fileloc=relative_fileloc, + task_handler_bindings=bindings, + task_handler_artifact_bundles={ARTIFACT_BUNDLE}, + ) + + +def _get_recorded_handlers(session) -> set[tuple[str, str, str]]: + """Return ``(dag_id, task_id, artifact path)`` for every recorded handler.""" + rows = session.execute( + select( + LangSDKTaskHandler.dag_id, + LangSDKTaskHandler.task_id, + LangSDKTaskHandlerArtifact.relative_fileloc, + ).join(LangSDKTaskHandlerArtifact, LangSDKTaskHandlerArtifact.id == LangSDKTaskHandler.artifact_id) + ) + return {tuple(row) for row in rows} + + +def _get_recorded_artifacts(session) -> list[tuple]: + return session.execute( + select( + LangSDKTaskHandlerArtifact.relative_fileloc, + LangSDKTaskHandlerArtifact.size_bytes, + LangSDKTaskHandlerArtifact.cache_digest, + LangSDKTaskHandlerArtifact.task_handlers, + LangSDKTaskHandlerArtifact.last_probed_at, + ).order_by(LangSDKTaskHandlerArtifact.relative_fileloc) + ).all() + + +@contextlib.contextmanager +def _capture_task_handler_writes(session) -> Generator[list[str], None, None]: + writes: list[str] = [] + + def _capture(conn, cursor, statement, parameters, context, executemany): + if "lang_sdk_task_handler" in statement and not statement.lstrip().upper().startswith("SELECT"): + writes.append(statement) + + bind = session.get_bind() + event.listen(bind, "before_cursor_execute", _capture) + try: + yield writes + finally: + event.remove(bind, "before_cursor_execute", _capture) + + +def _sweep_after_the_lookup(find): + """Wrap ``_find_task_handler_artifacts`` so ``etl-v2.jar`` gets an id no row has, as if swept since.""" + + def _find(keys, *, session): + artifact_ids = find(keys, session=session) + artifact_ids[(ARTIFACT_BUNDLE, compute_fileloc_hash("etl-v2.jar"))] = uuid.uuid4() + return artifact_ids + + return _find + + +def _fail_after_writing(sync): + """Wrap ``_sync_task_handlers`` so it raises a ``DataError`` after its writes.""" + + def _sync(*args, **kwargs): + sync(*args, **kwargs) + raise DataError("INSERT INTO lang_sdk_task_handler", {}, Exception("value too long")) + + return _sync + + +@pytest.mark.db_test +class TestRecordProbedTaskHandlerArtifacts: + @pytest.fixture(autouse=True) + def _clear_committed_rows(self): + # Autouse, so it tears down after the ``session`` fixture has rolled back; the race test + # commits rows from a second session. + yield + with create_session() as session: + session.execute(delete(LangSDKTaskHandlerArtifact)) + + def test_records_fingerprint_answer_and_probe_time(self, session, time_machine): + probed_at = tz.datetime(2026, 9, 30, 1) + time_machine.move_to(probed_at, tick=False) + + _record(session, "etl.jar") + + assert _get_recorded_artifacts(session) == [ + ("etl.jar", 1024, "a" * 64, {"etl": [TASK_HANDLERS["etl"][0].model_dump(mode="json")]}, probed_at) + ] + assert session.scalar( + select(LangSDKTaskHandlerArtifact.relative_fileloc_hash) + ) == compute_fileloc_hash("etl.jar") + + def test_reprobe_updates_fingerprint_answer_and_probe_time(self, session, time_machine): + time_machine.move_to(tz.datetime(2026, 9, 30, 1), tick=False) + _record(session, "etl.jar") + artifact_id = session.scalar(select(LangSDKTaskHandlerArtifact.id)) + + reprobed_at = tz.datetime(2026, 9, 30, 2) + time_machine.move_to(reprobed_at, tick=False) + record_probed_task_handler_artifacts( + [_make_artifact(size_bytes=2048, cache_digest="b" * 64, task_handlers={})], + bundle_names={ARTIFACT_BUNDLE}, + session=session, + ) + + assert _get_recorded_artifacts(session) == [("etl.jar", 2048, "b" * 64, {}, reprobed_at)] + assert session.scalar(select(LangSDKTaskHandlerArtifact.id)) == artifact_id + + def test_records_an_artifact_without_cache_digest(self, session): + record_probed_task_handler_artifacts( + [_make_artifact(cache_digest=None)], bundle_names={ARTIFACT_BUNDLE}, session=session + ) + + assert session.scalar(select(LangSDKTaskHandlerArtifact.cache_digest)) is None + + def test_artifact_outside_the_scope_is_dropped(self, session, caplog): + record_probed_task_handler_artifacts( + [_make_artifact("etl.jar"), _make_artifact("etl", bundle_name=OTHER_TEAM_BUNDLE)], + bundle_names={ARTIFACT_BUNDLE}, + session=session, + ) + + assert [row[0] for row in _get_recorded_artifacts(session)] == ["etl.jar"] + assert { + "event": "Ignoring probed task handler artifacts outside the Dag file's scope", + "artifacts": [f"{OTHER_TEAM_BUNDLE}/etl"], + } in caplog + + def test_duplicate_artifact_is_written_once(self, session): + with _capture_task_handler_writes(session) as writes: + record_probed_task_handler_artifacts( + [_make_artifact(), _make_artifact()], bundle_names={ARTIFACT_BUNDLE}, session=session + ) + + assert len(writes) == 1 + assert [row[0] for row in _get_recorded_artifacts(session)] == ["etl.jar"] + + def test_artifact_recorded_by_a_concurrent_parse(self, session): + """Two Dag processors probe one new artifact; the second upsert must update, not fail.""" + with create_session(scoped=False) as other: + _record(other, "etl.jar") + other_id = other.scalar(select(LangSDKTaskHandlerArtifact.id)) + + record_probed_task_handler_artifacts( + [_make_artifact(size_bytes=2048)], bundle_names={ARTIFACT_BUNDLE}, session=session + ) + + assert session.execute( + select(LangSDKTaskHandlerArtifact.id, LangSDKTaskHandlerArtifact.size_bytes) + ).all() == [(other_id, 2048)] + + +@pytest.mark.db_test +@pytest.mark.usefixtures("testing_dag_bundle") +class TestTaskHandlerBindingReconcile: + @pytest.fixture + def recorded(self, testing_dag_bundle, session): + _record(session, "etl.jar", "etl-v2.jar", "other.jar") + _persist( + session, + "etl", + bindings=[_make_binding(task_id="extract"), _make_binding(task_id="transform")], + ) + _persist( + session, + "other", + relative_fileloc="other.py", + bindings=[_make_binding("other", "load", rel_path="other.jar")], + ) + + @pytest.mark.parametrize( + ("bindings", "expected_handlers"), + [ + pytest.param( + None, + { + ("etl", "extract", "etl.jar"), + ("etl", "transform", "etl.jar"), + ("other", "load", "other.jar"), + }, + id="none-leaves-rows", + ), + pytest.param( + [], + {("other", "load", "other.jar")}, + id="empty-deletes-this-results-dag-rows", + ), + pytest.param( + [_make_binding(task_id="extract", rel_path="etl-v2.jar"), _make_binding(task_id="publish")], + { + ("etl", "extract", "etl-v2.jar"), + ("etl", "publish", "etl.jar"), + ("other", "load", "other.jar"), + }, + id="list-deletes-updates-and-inserts", + ), + ], + ) + @pytest.mark.usefixtures("recorded") + def test_task_handler_bindings_reconcile_by_value(self, session, bindings, expected_handlers): + _persist(session, "etl", bindings=bindings) + + assert _get_recorded_handlers(session) == expected_handlers + + @pytest.mark.usefixtures("recorded") + def test_task_handler_bindings_of_a_result_without_dags_change_nothing(self, session): + before = _get_recorded_handlers(session) + + _persist(session, bindings=[]) + + assert _get_recorded_handlers(session) == before + + def test_task_handler_rows_follow_a_dag_to_its_new_file(self, session): + _record(session, "etl.jar") + _persist(session, "etl", relative_fileloc="a.py", bindings=[_make_binding()]) + + _persist(session, "etl", relative_fileloc="dags/b.py", bindings=[_make_binding()]) + + assert session.execute( + select(LangSDKTaskHandler.dag_relative_fileloc, LangSDKTaskHandler.dag_relative_fileloc_hash) + ).all() == [("dags/b.py", compute_fileloc_hash("dags/b.py"))] + + def test_task_handler_artifact_is_shared_across_files(self, session): + _record(session, "etl.jar") + _persist(session, "etl", relative_fileloc="etl.py", bindings=[_make_binding("etl")]) + _persist(session, "other", relative_fileloc="other.py", bindings=[_make_binding("other")]) + + artifact_ids = session.scalars(select(LangSDKTaskHandlerArtifact.id)).all() + assert len(artifact_ids) == 1 + assert set(session.scalars(select(LangSDKTaskHandler.artifact_id))) == set(artifact_ids) + + _persist(session, "etl", relative_fileloc="etl.py", bindings=[]) + + assert session.scalars(select(LangSDKTaskHandlerArtifact.id)).all() == artifact_ids + assert _get_recorded_handlers(session) == {("other", "extract", "etl.jar")} + + def test_bindings_never_write_the_artifact_table(self, session): + _record(session, "etl.jar") + + with _capture_task_handler_writes(session) as writes: + _persist(session, "etl", bindings=[_make_binding()]) + assert writes + assert not [statement for statement in writes if "lang_sdk_task_handler_artifact" in statement] + + with _capture_task_handler_writes(session) as writes: + _persist(session, "etl", bindings=[_make_binding()]) + assert writes == [] + + @pytest.mark.usefixtures("recorded") + def test_binding_to_a_missing_artifact_leaves_its_dag_unchanged(self, session, caplog): + before = _get_recorded_handlers(session) + + _persist( + session, + "etl", + "other", + bindings=[ + _make_binding(task_id="extract", rel_path="swept.jar"), + _make_binding("other", "export"), + ], + ) + + assert _get_recorded_handlers(session) == {row for row in before if row[0] == "etl"} | { + ("other", "export", "etl.jar") + } + assert { + "event": "Ignoring task handler bindings of Dags bound to an unrecorded artifact", + "dag_ids": ["etl"], + } in caplog + + @pytest.mark.usefixtures("recorded") + def test_binding_outside_the_scope_leaves_its_dag_unchanged(self, session, caplog): + _record(session, "etl.jar", bundle_name=OTHER_TEAM_BUNDLE) + before = _get_recorded_handlers(session) + + _persist( + session, + "etl", + "other", + bindings=[ + _make_binding(task_id="extract", bundle_name=OTHER_TEAM_BUNDLE), + _make_binding("other", "export"), + ], + ) + + assert _get_recorded_handlers(session) == {row for row in before if row[0] == "etl"} | { + ("other", "export", "etl.jar") + } + assert { + "event": "Ignoring task handler bindings of Dags bound to an artifact outside their scope", + "dag_ids": ["etl"], + } in caplog + + def test_task_handler_bindings_for_dags_outside_the_result_are_dropped(self, session, caplog): + _record(session, "etl.jar") + _persist(session, "etl", bindings=[_make_binding("etl"), _make_binding("ghost")]) + + assert _get_recorded_handlers(session) == {("etl", "extract", "etl.jar")} + assert { + "event": "Ignoring task handler bindings of Dags not in the parse result", + "dag_ids": ["ghost"], + } in caplog + + @pytest.mark.usefixtures("recorded") + def test_task_handler_bindings_with_a_duplicate_task_leave_the_dag_untouched(self, session, caplog): + before = _get_recorded_handlers(session) + + _persist( + session, + "etl", + "other", + bindings=[ + _make_binding(task_id="extract", rel_path="etl.jar"), + _make_binding(task_id="extract", rel_path="etl-v2.jar"), + _make_binding("other", "load", rel_path="etl-v2.jar"), + ], + ) + + assert _get_recorded_handlers(session) == {row for row in before if row[0] == "etl"} | { + ("other", "load", "etl-v2.jar") + } + assert { + "event": "Ignoring task handler bindings of Dags that bind a task twice", + "dag_ids": ["etl"], + } in caplog + + @pytest.mark.parametrize( + ("target", "wrap"), + [ + pytest.param("_find_task_handler_artifacts", _sweep_after_the_lookup, id="integrity-error"), + pytest.param("_sync_task_handlers", _fail_after_writing, id="data-error"), + ], + ) + @pytest.mark.usefixtures("recorded") + def test_bindings_the_database_rejects_leave_the_rest_of_the_result(self, session, caplog, target, wrap): + before = _get_recorded_handlers(session) + + with mock.patch.object( + airflow.dag_processing.collection, + target, + side_effect=wrap(getattr(airflow.dag_processing.collection, target)), + ): + update_dag_parsing_results_in_db( + bundle_name="testing", + bundle_version=None, + dags=[LazyDeserializedDAG.from_dag(DAG(dag_id=dag_id)) for dag_id in ("etl", "fresh")], + import_errors={("testing", "broken.py"): "SyntaxError"}, + parse_duration=None, + warnings=set(), + session=session, + files_parsed={("testing", "etl.py"), ("testing", "broken.py")}, + relative_fileloc="etl.py", + # Deletes the transform row before inserting the publish row. + task_handler_bindings=[ + _make_binding(task_id="extract"), + _make_binding(task_id="publish", rel_path="etl-v2.jar"), + ], + task_handler_artifact_bundles={ARTIFACT_BUNDLE}, + ) + + assert _get_recorded_handlers(session) == before + assert session.scalars( + select(SerializedDagModel.dag_id).where(SerializedDagModel.dag_id == "fresh") + ).all() == ["fresh"] + assert session.scalars(select(ParseImportError.filename)).all() == ["broken.py"] + assert { + "event": "Failed to record task handler bindings; the rest of the parse result is kept", + "bundle_name": "testing", + "relative_fileloc": "etl.py", + "log_level": "warning", + } in caplog + + def test_task_handler_bindings_require_relative_fileloc(self, session): + with pytest.raises(ValueError, match="relative_fileloc"): + update_dag_parsing_results_in_db( + bundle_name="testing", + bundle_version=None, + dags=[LazyDeserializedDAG.from_dag(DAG(dag_id="etl"))], + import_errors={}, + parse_duration=None, + warnings=set(), + session=session, + task_handler_bindings=[], + task_handler_artifact_bundles={ARTIFACT_BUNDLE}, + ) + + def test_task_handler_bindings_require_a_scope(self, session): + with pytest.raises(ValueError, match="task_handler_artifact_bundles"): + update_dag_parsing_results_in_db( + bundle_name="testing", + bundle_version=None, + dags=[LazyDeserializedDAG.from_dag(DAG(dag_id="etl"))], + import_errors={}, + parse_duration=None, + warnings=set(), + session=session, + relative_fileloc="etl.py", + task_handler_bindings=[], + ) diff --git a/airflow-core/tests/unit/dag_processing/test_manager.py b/airflow-core/tests/unit/dag_processing/test_manager.py index a3624e1d4a143..4095cc217db44 100644 --- a/airflow-core/tests/unit/dag_processing/test_manager.py +++ b/airflow-core/tests/unit/dag_processing/test_manager.py @@ -41,15 +41,20 @@ import pytest import structlog import time_machine -from sqlalchemy import event, func, select +from sqlalchemy import delete, event, func, select from sqlalchemy.exc import OperationalError from uuid6 import uuid7 +from airflow import settings from airflow._shared.timezones import timezone from airflow.callbacks.callback_requests import DagCallbackRequest from airflow.dag_processing.bundles.base import BaseDagBundle, BundleVersion from airflow.dag_processing.bundles.manager import DagBundlesManager -from airflow.dag_processing.collection import update_dag_parsing_results_in_db +from airflow.dag_processing.collection import ( + _find_task_handler_artifacts, + record_probed_task_handler_artifacts, + update_dag_parsing_results_in_db, +) from airflow.dag_processing.dagbag import DagBag from airflow.dag_processing.manager import ( BundleState, @@ -61,6 +66,10 @@ DagFileParseRequest, DagFileParsingResult, DagFileProcessorProcess, + TaskHandlerArtifact, + TaskHandlerBinding, + TaskHandlerDeclaration, + TaskHandlerParam, _parse_file, ) from airflow.models import DagModel, DbCallbackRequest @@ -68,10 +77,16 @@ from airflow.models.dag_version import DagVersion from airflow.models.dagbundle import DagBundleModel from airflow.models.dagcode import DagCode +from airflow.models.lang_sdk_task_handler import ( + LangSDKTaskHandler, + LangSDKTaskHandlerArtifact, + compute_fileloc_hash, +) from airflow.models.serialized_dag import SerializedDagModel from airflow.models.team import Team from airflow.providers.standard.operators.empty import EmptyOperator from airflow.sdk import DAG as SdkDAG +from airflow.sdk.execution_time.coordinator import reset_coordinator_manager from airflow.sdk.importers import DagSourceCode from airflow.serialization.serialized_objects import LazyDeserializedDAG from airflow.utils.net import get_hostname @@ -1859,9 +1874,10 @@ def test_collect_results_tolerates_stale_file_handle_on_close(self): @pytest.mark.usefixtures("testing_dag_bundle") @pytest.mark.parametrize( - ("callbacks", "path", "expected_body"), + ("callbacks", "known_artifacts", "path", "expected_body"), [ pytest.param( + [], [], "/opt/airflow/dags/test_dag.py", { @@ -1869,9 +1885,62 @@ def test_collect_results_tolerates_stale_file_handle_on_close(self): "bundle_path": "/opt/airflow/dags", "bundle_name": "testing", "callback_requests": [], + "known_artifacts": [], "type": "DagFileParseRequest", }, ), + pytest.param( + [], + [ + TaskHandlerArtifact( + bundle_name="java-task-handlers", + relative_fileloc="etl.jar", + size_bytes=1024, + cache_digest="ab12", + task_handlers={ + "etl": [ + TaskHandlerDeclaration( + task_id="extract", + binding="positional", + params=[TaskHandlerParam(name=None)], + ) + ] + }, + ) + ], + "/opt/airflow/dags/etl.py", + { + "file": "/opt/airflow/dags/etl.py", + "bundle_path": "/opt/airflow/dags", + "bundle_name": "testing", + "callback_requests": [], + "known_artifacts": [ + { + "bundle_name": "java-task-handlers", + "relative_fileloc": "etl.jar", + "size_bytes": 1024, + "cache_digest": "ab12", + "task_handlers": { + "etl": [ + { + "task_id": "extract", + "binding": "positional", + "params": [ + { + "name": None, + "value_schema": None, + "exact_name": False, + } + ], + } + ] + }, + } + ], + "type": "DagFileParseRequest", + }, + id="known-artifacts", + ), pytest.param( [ DagCallbackRequest( @@ -1884,6 +1953,7 @@ def test_collect_results_tolerates_stale_file_handle_on_close(self): is_failure_callback=False, ) ], + [], "/opt/airflow/dags/dag_callback_dag.py", { "file": "/opt/airflow/dags/dag_callback_dag.py", @@ -1903,17 +1973,22 @@ def test_collect_results_tolerates_stale_file_handle_on_close(self): "type": "DagCallbackRequest", } ], + "known_artifacts": [], "type": "DagFileParseRequest", }, ), ], ) - def test_serialize_callback_requests(self, callbacks, path, expected_body): + def test_serialize_parse_request(self, callbacks, known_artifacts, path, expected_body): from airflow.sdk.execution_time.comms import _ResponseFrame processor, read_socket = self.mock_processor() processor._on_child_started( - callbacks, path, bundle_path=Path("/opt/airflow/dags"), bundle_name="testing" + callbacks, + path, + bundle_path=Path("/opt/airflow/dags"), + bundle_name="testing", + known_artifacts=known_artifacts, ) read_socket.settimeout(0.1) @@ -2623,6 +2698,7 @@ def test_callback_queue(self, mock_get_logger, configure_testing_dag_bundle): bundle_name="testing", dag_file_rel_path=str(dag2_path.rel_path), callbacks=[dag2_req1], + known_artifacts=[], selector=mock.ANY, logger=mock_logger, logger_filehandle=mock_filehandle, @@ -2636,6 +2712,7 @@ def test_callback_queue(self, mock_get_logger, configure_testing_dag_bundle): bundle_name="testing", dag_file_rel_path=str(dag1_path.rel_path), callbacks=[dag1_req1, dag1_req2], + known_artifacts=[], selector=mock.ANY, logger=mock_logger, logger_filehandle=mock_filehandle, @@ -4102,6 +4179,635 @@ def test_a_sweep_pays_the_fixed_cost_once_per_call(self, session, testing_dag_bu ) +JAVA_TASK_HANDLERS = "java-task-handlers" + + +def _coordinator(task_handler_bundle_name: str | None = None) -> dict: + kwargs = ( + {} if task_handler_bundle_name is None else {"task_handler_bundle_name": task_handler_bundle_name} + ) + return {"classpath": "airflow.sdk.coordinators.java.JavaCoordinator", "kwargs": kwargs} + + +def _routed_coordinators(coordinators: dict[str, dict]) -> dict[tuple[str, str], str]: + """Return the ``[sdk]`` options that configure *coordinators* and route a queue to each.""" + return { + ("sdk", "coordinators"): json.dumps(coordinators), + ("sdk", "queue_to_coordinator"): json.dumps({f"queue-{key}": key for key in coordinators}), + } + + +def _known_artifact(bundle_name: str) -> TaskHandlerArtifact: + return TaskHandlerArtifact( + bundle_name=bundle_name, + relative_fileloc=f"{bundle_name}.jar", + size_bytes=1024, + cache_digest="ab12", + task_handlers={ + "etl": [ + TaskHandlerDeclaration( + task_id="extract", + binding="positional", + params=[TaskHandlerParam(name=None)], + ) + ] + }, + ) + + +GO_TASK_HANDLERS = "go-task-handlers" +TS_TASK_HANDLERS = "ts-task-handlers" +TEAM_COORDINATORS = { + "java": _coordinator(JAVA_TASK_HANDLERS), + "go": _coordinator(GO_TASK_HANDLERS), + "ts": _coordinator(TS_TASK_HANDLERS), +} + + +class TestKnownTaskHandlerArtifacts: + @pytest.fixture(autouse=True) + def _reset_coordinator_manager(self): + reset_coordinator_manager() + yield + reset_coordinator_manager() + + @pytest.fixture + def manager(self, configure_dag_bundles, tmp_path): + with ( + conf_vars({("core", "load_examples"): "False"}), + configure_dag_bundles({"dags-a": tmp_path, "dags-b": tmp_path, JAVA_TASK_HANDLERS: tmp_path}), + ): + manager = DagFileProcessorManager(max_runs=1) + manager._dag_bundles = list(DagBundlesManager().get_all_dag_bundles()) + yield manager + + @conf_vars(_routed_coordinators({"java": _coordinator(JAVA_TASK_HANDLERS), "go": _coordinator()})) + @mock.patch.object(DagFileProcessorProcess, "start") + @mock.patch.object(DagFileProcessorManager, "get_known_task_handler_artifacts", autospec=True) + def test_start_new_processes_reads_known_artifacts_once_per_loop(self, get_known, start, manager): + known = {name: [_known_artifact(name)] for name in ("dags-a", "dags-b", JAVA_TASK_HANDLERS)} + get_known.return_value = known + files = [ + DagFileInfo(bundle_name=bundle_name, rel_path=Path(rel_path), bundle_path=TEST_DAGS_FOLDER) + for bundle_name, rel_path in (("dags-a", "one.py"), ("dags-a", "two.py"), ("dags-b", "three.py")) + ] + manager._parallelism = 3 + manager._file_queue = OrderedDict.fromkeys(files) + + manager._start_new_processes() + + get_known.assert_called_once() + assert { + call.kwargs["dag_file_rel_path"]: call.kwargs["known_artifacts"] for call in start.call_args_list + } == { + "one.py": [*known["dags-a"], *known[JAVA_TASK_HANDLERS]], + "two.py": [*known["dags-a"], *known[JAVA_TASK_HANDLERS]], + "three.py": [*known["dags-b"], *known[JAVA_TASK_HANDLERS]], + } + + manager._processors.clear() + manager._file_queue = OrderedDict.fromkeys(files[:1]) + manager._start_new_processes() + + assert get_known.call_count == 2 + + @pytest.mark.parametrize( + ("coordinators", "teams", "expected_bundle_names"), + [ + pytest.param( + TEAM_COORDINATORS, + {}, + {JAVA_TASK_HANDLERS, GO_TASK_HANDLERS, TS_TASK_HANDLERS}, + id="multi-team-off", + ), + pytest.param( + TEAM_COORDINATORS, + {"dags-a": "team-a", JAVA_TASK_HANDLERS: "team-a", GO_TASK_HANDLERS: "team-b"}, + {JAVA_TASK_HANDLERS}, + id="same-team", + ), + pytest.param( + TEAM_COORDINATORS, + {JAVA_TASK_HANDLERS: "team-a", GO_TASK_HANDLERS: "team-b"}, + {TS_TASK_HANDLERS}, + id="no-team", + ), + pytest.param( + {**TEAM_COORDINATORS, "fallback": _coordinator()}, + {"dags-a": "team-a", JAVA_TASK_HANDLERS: "team-a", GO_TASK_HANDLERS: "team-b"}, + {JAVA_TASK_HANDLERS, "dags-a"}, + id="own-bundle-fallback", + ), + ], + ) + def test_known_artifacts_are_scoped_by_team( + self, configure_dag_bundles, tmp_path, coordinators, teams, expected_bundle_names + ): + bundle_names = ("dags-a", JAVA_TASK_HANDLERS, GO_TASK_HANDLERS, TS_TASK_HANDLERS) + with ( + conf_vars({("core", "load_examples"): "False", **_routed_coordinators(coordinators)}), + configure_dag_bundles(dict.fromkeys(bundle_names, tmp_path)), + mock.patch.object(DagFileProcessorManager, "_get_team_names", autospec=True, return_value=teams), + ): + selected = DagFileProcessorManager(max_runs=1)._select_known_task_handler_artifacts( + {name: [_known_artifact(name)] for name in bundle_names}, "dags-a" + ) + + assert {artifact.bundle_name for artifact in selected} == expected_bundle_names + + @conf_vars(_routed_coordinators({"java": _coordinator(JAVA_TASK_HANDLERS)})) + @mock.patch.object(DagFileProcessorProcess, "start") + @mock.patch.object(DagFileProcessorManager, "get_known_task_handler_artifacts", autospec=True) + def test_start_new_processes_keeps_own_bundle_out_without_a_fallback(self, get_known, start, manager): + # An override may return bundles it was not asked for. + get_known.return_value = {name: [_known_artifact(name)] for name in ("dags-a", JAVA_TASK_HANDLERS)} + manager._file_queue = OrderedDict.fromkeys( + [DagFileInfo(bundle_name="dags-a", rel_path=Path("one.py"), bundle_path=TEST_DAGS_FOLDER)] + ) + + manager._start_new_processes() + + assert start.call_args.kwargs["known_artifacts"] == [_known_artifact(JAVA_TASK_HANDLERS)] + + @conf_vars(_routed_coordinators({"java": _coordinator(JAVA_TASK_HANDLERS)})) + @mock.patch.object(DagFileProcessorManager, "get_known_task_handler_artifacts", autospec=True) + def test_start_new_processes_reads_nothing_without_a_child_to_start(self, get_known, manager): + manager._start_new_processes() + + get_known.assert_not_called() + + @pytest.mark.parametrize( + ("coordinators", "expected_bundle_names"), + [ + pytest.param({"java": _coordinator(JAVA_TASK_HANDLERS)}, {JAVA_TASK_HANDLERS}, id="named-only"), + pytest.param( + {"java": _coordinator(JAVA_TASK_HANDLERS), "go": _coordinator()}, + {JAVA_TASK_HANDLERS, "dags-a", "dags-b"}, + id="own-bundle-fallback", + ), + ], + ) + @mock.patch.object(DagFileProcessorManager, "get_known_task_handler_artifacts", autospec=True) + def test_query_asks_only_for_bundles_that_can_hold_artifacts( + self, get_known, manager, coordinators, expected_bundle_names + ): + with conf_vars(_routed_coordinators(coordinators)): + manager._query_known_task_handler_artifacts() + + get_known.assert_called_once_with(manager, expected_bundle_names) + + @pytest.mark.parametrize( + "unrouted", + [ + _coordinator("ghost"), + { + "classpath": "airflow.sdk.coordinators.java.JavaCoordinator", + "kwargs": {"task_handler_bundle_name": ["ghost"]}, + }, + _coordinator(), + ], + ids=["unconfigured", "not-a-string", "own-bundle"], + ) + @mock.patch.object(DagFileProcessorManager, "get_known_task_handler_artifacts", autospec=True) + def test_query_leaves_out_the_bundle_of_an_unrouted_coordinator(self, get_known, manager, unrouted): + with conf_vars( + { + ("sdk", "coordinators"): json.dumps( + {"java": _coordinator(JAVA_TASK_HANDLERS), "unrouted": unrouted} + ), + ("sdk", "queue_to_coordinator"): json.dumps({"queue-java": "java"}), + } + ): + manager._query_known_task_handler_artifacts() + + get_known.assert_called_once_with(manager, {JAVA_TASK_HANDLERS}) + + @mock.patch.object(DagFileProcessorManager, "get_known_task_handler_artifacts", autospec=True) + def test_query_skips_the_read_without_coordinators(self, get_known, manager): + assert manager._query_known_task_handler_artifacts() == {} + + get_known.assert_not_called() + + @conf_vars(_routed_coordinators({"java": _coordinator("ghost")})) + @mock.patch.object(DagFileProcessorManager, "get_known_task_handler_artifacts", autospec=True) + def test_invalid_coordinator_config_is_logged_once_and_sends_nothing(self, get_known, manager, caplog): + assert manager._query_known_task_handler_artifacts() == {} + assert manager._query_known_task_handler_artifacts() == {} + + get_known.assert_not_called() + event = "Cannot read [sdk] coordinators; Dag files are parsed without known task-handler artifacts" + assert [entry["log_level"] for entry in caplog.entries if entry["event"] == event] == ["error"] + + @conf_vars(_routed_coordinators({"java": _coordinator(JAVA_TASK_HANDLERS)})) + @mock.patch.object(DagFileProcessorProcess, "start") + @mock.patch.object(DagFileProcessorManager, "get_known_task_handler_artifacts", autospec=True) + def test_start_new_processes_sends_no_known_artifacts_when_the_read_fails( + self, get_known, start, manager, caplog + ): + get_known.side_effect = OperationalError("SELECT", {}, Exception("connection lost")) + manager._parallelism = 2 + manager._file_queue = OrderedDict.fromkeys( + DagFileInfo(bundle_name="dags-a", rel_path=Path(rel_path), bundle_path=TEST_DAGS_FOLDER) + for rel_path in ("one.py", "two.py") + ) + + manager._start_new_processes() + + get_known.assert_called_once() + assert [call.kwargs["known_artifacts"] for call in start.call_args_list] == [[], []] + assert { + "event": "Cannot read the recorded task-handler artifacts; Dag files are parsed without them", + "log_level": "error", + } in caplog + + def test_get_known_task_handler_artifacts_from_db_skips_a_row_that_fails_validation( + self, manager, session, caplog + ): + valid = _known_artifact(JAVA_TASK_HANDLERS) + session.add_all( + [ + LangSDKTaskHandlerArtifact(**valid.model_dump(mode="json")), + LangSDKTaskHandlerArtifact( + bundle_name=JAVA_TASK_HANDLERS, + relative_fileloc="broken.jar", + size_bytes=1024, + cache_digest="ab12", + task_handlers={"etl": [{"task_id": "extract", "binding": "unknown"}]}, + ), + ] + ) + session.flush() + + known = manager._get_known_task_handler_artifacts_from_db({JAVA_TASK_HANDLERS}, session=session) + + assert known == {JAVA_TASK_HANDLERS: [valid]} + assert { + "event": "Ignoring a recorded task-handler artifact that fails validation", + "bundle_name": JAVA_TASK_HANDLERS, + "relative_fileloc": "broken.jar", + "log_level": "warning", + } in caplog + + def test_get_known_task_handler_artifacts_from_db_orders_by_bundle_and_path(self, manager, session): + session.add_all( + LangSDKTaskHandlerArtifact( + bundle_name=bundle_name, + relative_fileloc=rel_path, + size_bytes=1024, + cache_digest=None, + task_handlers={}, + ) + for bundle_name, rel_path in ( + (JAVA_TASK_HANDLERS, "c.jar"), + ("dags-a", "b.jar"), + (JAVA_TASK_HANDLERS, "a.jar"), + (JAVA_TASK_HANDLERS, "b.jar"), + ) + ) + session.flush() + + known = manager._get_known_task_handler_artifacts_from_db( + {JAVA_TASK_HANDLERS, "dags-a"}, session=session + ) + + assert [ + (bundle_name, [artifact.relative_fileloc for artifact in artifacts]) + for bundle_name, artifacts in known.items() + ] == [("dags-a", ["b.jar"]), (JAVA_TASK_HANDLERS, ["a.jar", "b.jar", "c.jar"])] + + def test_get_known_task_handler_artifacts_from_db_reads_only_the_given_bundles(self, manager, session): + session.add_all( + LangSDKTaskHandlerArtifact( + bundle_name=name, + relative_fileloc=f"{name}.jar", + size_bytes=1024, + cache_digest="ab12", + task_handlers={ + "etl": [ + { + "task_id": "extract", + "binding": "positional", + "params": [{"name": None}], + } + ] + }, + ) + for name in (JAVA_TASK_HANDLERS, "dags-a", "unconfigured") + ) + session.flush() + + known = manager._get_known_task_handler_artifacts_from_db( + {JAVA_TASK_HANDLERS, "dags-a"}, session=session + ) + + assert known == { + JAVA_TASK_HANDLERS: [_known_artifact(JAVA_TASK_HANDLERS)], + "dags-a": [_known_artifact("dags-a")], + } + + +class TestTaskHandlerBindings: + @staticmethod + def _finished_parse(**result) -> tuple[DagFileInfo, MagicMock]: + file = DagFileInfo(bundle_name="dags-a", rel_path=Path("etl.py"), bundle_path=TEST_DAGS_FOLDER) + proc = MagicMock( + had_callbacks=False, + start_time=time.monotonic(), + parsing_result=DagFileParsingResult(fileloc="etl.py", serialized_dags=[], **result), + ) + return file, proc + + @staticmethod + def _manager() -> DagFileProcessorManager: + manager = DagFileProcessorManager(max_runs=1) + manager._bundle_versions["dags-a"] = None + return manager + + @mock.patch.object(DagFileProcessorManager, "persist_parsing_result", autospec=True) + @mock.patch.object(DagFileProcessorManager, "persist_probed_task_handler_artifacts", autospec=True) + def test_handle_parsing_result_records_probed_artifacts_first(self, persist_probed, persist_result): + calls: list[str] = [] + persist_probed.side_effect = lambda *args, **kwargs: calls.append("probed") + persist_result.side_effect = lambda *args, **kwargs: calls.append("result") + file, proc = self._finished_parse(probed_artifacts=[_known_artifact(JAVA_TASK_HANDLERS)]) + manager = self._manager() + + manager.handle_parsing_result(file, proc, session=mock.MagicMock()) + + persist_probed.assert_called_once_with( + manager, bundle_name="dags-a", artifacts=[_known_artifact(JAVA_TASK_HANDLERS)] + ) + assert calls == ["probed", "result"] + + @mock.patch.object(DagFileProcessorManager, "persist_parsing_result", autospec=True) + @mock.patch.object(DagFileProcessorManager, "persist_probed_task_handler_artifacts", autospec=True) + def test_handle_parsing_result_skips_the_artifact_write_without_probes( + self, persist_probed, persist_result + ): + file, proc = self._finished_parse() + + self._manager().handle_parsing_result(file, proc, session=mock.MagicMock()) + + persist_probed.assert_not_called() + persist_result.assert_called_once() + + @mock.patch.object(DagFileProcessorManager, "persist_parsing_result", autospec=True) + @mock.patch.object(DagFileProcessorManager, "persist_probed_task_handler_artifacts", autospec=True) + def test_handle_parsing_result_persists_when_the_artifact_write_fails( + self, persist_probed, persist_result, caplog + ): + persist_probed.side_effect = OperationalError("INSERT", {}, Exception("locked")) + file, proc = self._finished_parse(probed_artifacts=[_known_artifact(JAVA_TASK_HANDLERS)]) + + self._manager().handle_parsing_result(file, proc, session=mock.MagicMock()) + + persist_result.assert_called_once() + assert { + "event": "Failed to record probed task handler artifacts", + "bundle_name": "dags-a", + "relative_fileloc": "etl.py", + "log_level": "error", + } in caplog + + @pytest.mark.parametrize( + ("bindings", "expected_scope"), + [ + pytest.param(None, None, id="not-evaluated"), + pytest.param( + [ + TaskHandlerBinding( + dag_id="etl", + task_id="extract", + artifact_bundle_name=JAVA_TASK_HANDLERS, + artifact_rel_path="etl.jar", + ) + ], + frozenset({JAVA_TASK_HANDLERS}), + id="bindings", + ), + ], + ) + @mock.patch.object( + DagFileProcessorManager, + "_get_task_handler_artifact_bundle_names", + autospec=True, + return_value=frozenset({JAVA_TASK_HANDLERS}), + ) + @mock.patch("airflow.dag_processing.manager.update_dag_parsing_results_in_db", autospec=True) + def test_persist_parsing_result_forwards_task_handler_bindings( + self, update_in_db, get_scope, bindings, expected_scope + ): + DagFileProcessorManager(max_runs=1).persist_parsing_result( + bundle_name="dags-a", + bundle_version=None, + version_data=None, + parsing_result=DagFileParsingResult( + fileloc="/files/dags/etl.py", serialized_dags=[], task_handler_bindings=bindings + ), + run_duration=0.1, + relative_fileloc="etl.py", + session=mock.MagicMock(), + ) + + assert update_in_db.call_args.kwargs["relative_fileloc"] == "etl.py" + assert update_in_db.call_args.kwargs["task_handler_bindings"] == bindings + assert update_in_db.call_args.kwargs["task_handler_artifact_bundles"] == expected_scope + assert get_scope.call_count == (bindings is not None) + + @pytest.mark.db_test + @conf_vars(_routed_coordinators({"java": _coordinator(JAVA_TASK_HANDLERS)})) + def test_persist_probed_task_handler_artifacts_commits_the_scoped_artifacts( + self, configure_dag_bundles, tmp_path + ): + reset_coordinator_manager() + try: + with configure_dag_bundles({"dags-a": tmp_path, JAVA_TASK_HANDLERS: tmp_path}): + DagFileProcessorManager(max_runs=1).persist_probed_task_handler_artifacts( + bundle_name="dags-a", + artifacts=[_known_artifact(JAVA_TASK_HANDLERS), _known_artifact("dags-b")], + ) + + with create_session(scoped=False) as other: + assert other.scalars(select(LangSDKTaskHandlerArtifact.bundle_name)).all() == [ + JAVA_TASK_HANDLERS + ] + finally: + reset_coordinator_manager() + with create_session() as session: + session.execute(delete(LangSDKTaskHandlerArtifact)) + + +@pytest.mark.db_test +@pytest.mark.usefixtures("testing_dag_bundle") +class TestTaskHandlerArtifactSweep: + @pytest.fixture(autouse=True) + def _clear_task_handler_rows(self): + yield + with create_session() as session: + session.execute(delete(LangSDKTaskHandler)) + session.execute(delete(LangSDKTaskHandlerArtifact)) + clear_db_serialized_dags() + clear_db_dags() + + @staticmethod + def _add_artifacts(*paths: str, probed_at: datetime, referenced: tuple[str, ...] = ()) -> None: + """Commit one artifact per path, and a handler row for each path in *referenced*.""" + with create_session() as session: + artifacts = { + path: LangSDKTaskHandlerArtifact( + bundle_name=JAVA_TASK_HANDLERS, + relative_fileloc=path, + size_bytes=1024, + cache_digest="ab12", + task_handlers={}, + last_probed_at=probed_at, + ) + for path in paths + } + session.add_all(artifacts.values()) + if referenced: + session.add(DagModel(dag_id="etl", bundle_name="testing")) + session.flush() + session.add_all( + LangSDKTaskHandler( + dag_id="etl", + task_id=path, + artifact_id=artifacts[path].id, + dag_bundle_name="testing", + dag_relative_fileloc="etl.py", + ) + for path in referenced + ) + + @staticmethod + def _get_artifact_paths() -> set[str]: + with create_session() as session: + return set(session.scalars(select(LangSDKTaskHandlerArtifact.relative_fileloc))) + + @staticmethod + def _sweep() -> int: + manager = DagFileProcessorManager(max_runs=1) + manager.parsing_cleanup_interval = 60 + return manager.delete_unreferenced_task_handler_artifacts() + + @mock.patch.object(DagFileProcessorManager, "delete_unreferenced_task_handler_artifacts", autospec=True) + def test_task_handler_artifact_sweep_runs_once_per_cleanup_interval(self, delete_unreferenced): + manager = DagFileProcessorManager(max_runs=1) + manager.parsing_cleanup_interval = 10 + manager._last_task_handler_artifact_sweep_time = time.monotonic() - 20 + + manager._sweep_task_handler_artifacts() + manager._sweep_task_handler_artifacts() + + delete_unreferenced.assert_called_once_with(manager) + + @mock.patch.object(DagFileProcessorManager, "delete_unreferenced_task_handler_artifacts", autospec=True) + def test_task_handler_artifact_sweep_logs_errors(self, delete_unreferenced, caplog): + delete_unreferenced.side_effect = OperationalError("DELETE", {}, Exception("locked")) + manager = DagFileProcessorManager(max_runs=1) + + manager._sweep_task_handler_artifacts() + + assert {"event": "Error deleting unreferenced task handler artifacts", "log_level": "error"} in caplog + assert manager._last_task_handler_artifact_sweep_time > 0 + + @mock.patch("airflow.dag_processing.manager._TASK_HANDLER_ARTIFACT_SWEEP_BATCH_SIZE", 2) + def test_delete_unreferenced_task_handler_artifacts_commits_each_batch(self): + self._add_artifacts( + "a.jar", + "b.jar", + "c.jar", + "d.jar", + "e.jar", + "used.jar", + probed_at=timezone.utcnow() - timedelta(days=1), + referenced=("used.jar",), + ) + events: list[str] = [] + + def on_execute(conn, cursor, statement, parameters, context, executemany): + if statement.lstrip().upper().startswith("DELETE FROM LANG_SDK_TASK_HANDLER_ARTIFACT"): + events.append("DELETE") + + def on_commit(conn): + events.append("COMMIT") + + bind = settings.engine + event.listen(bind, "before_cursor_execute", on_execute) + event.listen(bind, "commit", on_commit) + try: + deleted = self._sweep() + finally: + event.remove(bind, "before_cursor_execute", on_execute) + event.remove(bind, "commit", on_commit) + + assert deleted == 5 + assert events == ["DELETE", "COMMIT"] * 3 + assert self._get_artifact_paths() == {"used.jar"} + + @pytest.mark.parametrize( + ("age", "referenced", "expected"), + [ + pytest.param(timedelta(seconds=301), (), set(), id="unreferenced-past-grace"), + pytest.param(timedelta(seconds=299), (), {"etl.jar"}, id="unreferenced-within-grace"), + pytest.param(timedelta(days=30), ("etl.jar",), {"etl.jar"}, id="referenced"), + ], + ) + @time_machine.travel(datetime(2026, 9, 30, 12, tzinfo=timezone.utc), tick=False) + def test_task_handler_artifact_sweep_grace_period(self, age, referenced, expected): + self._add_artifacts("etl.jar", probed_at=timezone.utcnow() - age, referenced=referenced) + + self._sweep() + + assert self._get_artifact_paths() == expected + + def test_task_handler_artifact_of_a_deleted_dag_is_swept(self, time_machine): + time_machine.move_to(datetime(2026, 9, 30, 12, tzinfo=timezone.utc), tick=False) + with create_session() as session: + record_probed_task_handler_artifacts( + [_known_artifact(JAVA_TASK_HANDLERS)], bundle_names={JAVA_TASK_HANDLERS}, session=session + ) + update_dag_parsing_results_in_db( + "testing", + None, + [LazyDeserializedDAG.from_dag(SdkDAG(dag_id="etl"))], + {}, + None, + set(), + session, + relative_fileloc="etl.py", + task_handler_bindings=[ + TaskHandlerBinding( + dag_id="etl", + task_id="extract", + artifact_bundle_name=JAVA_TASK_HANDLERS, + artifact_rel_path=f"{JAVA_TASK_HANDLERS}.jar", + ) + ], + task_handler_artifact_bundles={JAVA_TASK_HANDLERS}, + ) + with create_session() as session: + session.execute(delete(DagModel).where(DagModel.dag_id == "etl")) + time_machine.shift(timedelta(hours=1)) + + self._sweep() + + assert self._get_artifact_paths() == set() + + @pytest.mark.backend("postgres", "mysql") + def test_task_handler_artifact_sweep_skips_rows_a_parse_holds(self): + self._add_artifacts("held.jar", "free.jar", probed_at=timezone.utcnow() - timedelta(days=1)) + + with create_session(scoped=False) as parse: + _find_task_handler_artifacts( + [(JAVA_TASK_HANDLERS, compute_fileloc_hash("held.jar"))], session=parse + ) + self._sweep() + parse.rollback() + + assert self._get_artifact_paths() == {"held.jar"} + + class TestMultiTeamMetrics: """Tests for team_name tag on dag processing metrics in multi-team mode.""" diff --git a/airflow-core/tests/unit/dag_processing/test_processor.py b/airflow-core/tests/unit/dag_processing/test_processor.py index 0cb61aa02866a..a135415d43b1a 100644 --- a/airflow-core/tests/unit/dag_processing/test_processor.py +++ b/airflow-core/tests/unit/dag_processing/test_processor.py @@ -18,22 +18,26 @@ from __future__ import annotations import inspect +import json import logging import os import pathlib +import signal import sys import textwrap +import time import typing import uuid import zipfile from collections.abc import Callable, Iterable -from socket import socketpair +from socket import socket, socketpair from typing import TYPE_CHECKING, Any, BinaryIO from unittest.mock import MagicMock, PropertyMock, patch +import psutil import pytest import structlog -from pydantic import TypeAdapter +from pydantic import TypeAdapter, ValidationError from sqlalchemy import select from structlog.typing import FilteringBoundLogger @@ -54,11 +58,19 @@ from airflow.dag_processing.dagbag import DagBag from airflow.dag_processing.manager import DagFileProcessorManager, process_parse_results from airflow.dag_processing.processor import ( + BaseDagFileProcessorProcess, DagFileParseRequest, DagFileParsingResult, DagFileProcessorProcess, + TaskHandlerArtifact, + TaskHandlerBinding, + TaskHandlerDeclaration, + TaskHandlerParam, + TaskHandlerParseRequest, + TaskHandlerParsingResult, ToDagProcessor, ToManager, + ToSDKTaskHandlerProcessor, _execute_callbacks, _execute_dag_callbacks, _execute_email_callbacks, @@ -67,12 +79,14 @@ _parse_file_entrypoint, _pre_import_airflow_modules, ) +from airflow.dag_processing.task_handler_resolution import TaskHandlerResolution from airflow.models import DagRun from airflow.models.dagwarning import DagWarning from airflow.sdk import DAG, BaseOperator from airflow.sdk.api.client import Client from airflow.sdk.api.datamodels._generated import ConnectionResponse, DagRunState, VariableResponse -from airflow.sdk.execution_time import comms, supervisor +from airflow.sdk.coordinators._subprocess import SubprocessCoordinator +from airflow.sdk.execution_time import comms, supervisor, task_runner from airflow.sdk.execution_time.comms import ( GetConnection, GetTaskStates, @@ -89,10 +103,19 @@ ) from airflow.sdk.execution_time.task_runner import RuntimeTaskInstance from airflow.sdk.importers import DagSourceCode +from airflow.serialization.serialized_objects import DagSerialization +from airflow.utils.dag_version_inflation_checker import check_dag_file_stability from airflow.utils.session import create_session from airflow.utils.state import TaskInstanceState from tests_common.test_utils.config import conf_vars, env_vars +from unit.dag_processing.fake_task_handler_runtime import ( + FakeCoordinator, + play_runtime, + reply_with_task_handlers, + task_handler_config, + write_artifact, +) if TYPE_CHECKING: from kgb import SpyAgency @@ -113,6 +136,7 @@ def _force_bare_fork(monkeypatch): DEFAULT_DATE = timezone.datetime(2016, 1, 1) +INTEGER = {"type": "integer"} # Filename to be used for dags that are created in an ad-hoc manner and can be removed/ # created at runtime @@ -615,6 +639,82 @@ def test_parse_file_entrypoint_parses_dag_callbacks(mocker): ] +def _make_known_artifact_body(**overrides) -> dict: + return { + "bundle_name": "java-task-handlers", + "relative_fileloc": "etl.jar", + "size_bytes": 1024, + "cache_digest": "ab12", + "task_handlers": { + "etl": [{"task_id": "extract", "binding": "positional", "params": [{"name": None}]}] + }, + **overrides, + } + + +def test_parse_file_entrypoint_decodes_known_artifacts(): + r, w = socketpair() + frame = comms._ResponseFrame( + id=0, + body={ + "file": "/files/dags/etl.py", + "bundle_path": "/files/dags", + "bundle_name": "testing", + "known_artifacts": [ + _make_known_artifact_body(), + _make_known_artifact_body(relative_fileloc="empty.jar", cache_digest=None, task_handlers={}), + ], + "type": "DagFileParseRequest", + }, + ) + w.sendall(frame.as_bytes()) + + # The same decoder _parse_file_entrypoint builds. + decoder = comms.CommsDecoder[ToDagProcessor, ToManager]( + socket=r, + body_decoder=TypeAdapter[ToDagProcessor](ToDagProcessor), + ) + + msg = decoder._get_response() + assert isinstance(msg, DagFileParseRequest) + assert msg.known_artifacts == [ + TaskHandlerArtifact( + bundle_name="java-task-handlers", + relative_fileloc="etl.jar", + size_bytes=1024, + cache_digest="ab12", + task_handlers={ + "etl": [ + TaskHandlerDeclaration( + task_id="extract", + binding="positional", + params=[TaskHandlerParam(name=None)], + ) + ] + }, + ), + TaskHandlerArtifact( + bundle_name="java-task-handlers", + relative_fileloc="empty.jar", + size_bytes=1024, + cache_digest=None, + task_handlers={}, + ), + ] + + +@pytest.mark.parametrize( + ("field", "value"), + [ + pytest.param("cache_digest", "a" * 129, id="cache-digest"), + pytest.param("relative_fileloc", "a" * 2001, id="rel-path"), + ], +) +def test_known_artifact_rejects_values_wider_than_their_column(field, value): + with pytest.raises(ValidationError, match=field): + TaskHandlerArtifact.model_validate(_make_known_artifact_body(**{field: value})) + + def test_parse_file_with_dag_callbacks(spy_agency): from airflow import DAG @@ -2354,15 +2454,487 @@ def get_type_names(union_type): + "\n\nEither handle these types in ToDagProcessor or update in_task_runner_but_not_in_dag_processing_process list." ) + def test_parse_request_unions_differ_only_in_request(self): + dag_processor_types = set(typing.get_args(typing.get_args(ToDagProcessor)[0])) + task_handler_types = set(typing.get_args(typing.get_args(ToSDKTaskHandlerProcessor)[0])) + + assert dag_processor_types - task_handler_types == {DagFileParseRequest} + assert task_handler_types - dag_processor_types == {TaskHandlerParseRequest} + + +class TestTaskHandlerDeclaration: + @pytest.mark.parametrize( + ("binding", "param", "expected"), + [ + pytest.param("positional", {"name": None}, TaskHandlerParam(name=None), id="positional"), + pytest.param( + "named", + {"name": "day", "exact_name": True}, + TaskHandlerParam(name="day", exact_name=True), + id="named", + ), + ], + ) + def test_decodes_binding(self, binding, param, expected): + declaration = TaskHandlerDeclaration.model_validate( + {"task_id": "extract", "binding": binding, "params": [param]} + ) + + assert declaration.binding == binding + assert declaration.params == [expected] + + def test_decodes_unlisted_params(self): + declaration = TaskHandlerDeclaration.model_validate( + {"task_id": "extract", "binding": "named", "params": None} + ) + + assert declaration.params is None + + @pytest.mark.parametrize( + "declaration", + [ + pytest.param({"task_id": "extract", "params": []}, id="missing"), + pytest.param({"task_id": "extract", "binding": "keyword", "params": []}, id="unknown"), + pytest.param( + {"task_id": "extract", "binding": "named_or_whole", "params": []}, id="named_or_whole" + ), + pytest.param({"task_id": "extract", "binding": "named_open", "params": []}, id="named_open"), + ], + ) + def test_rejects_invalid_binding(self, declaration): + with pytest.raises(ValidationError, match="binding"): + TaskHandlerDeclaration.model_validate(declaration) + + +def _make_binding_body(**overrides) -> dict: + return { + "dag_id": "etl", + "task_id": "extract", + "artifact_bundle_name": "java-task-handlers", + "artifact_rel_path": "etl.jar", + **overrides, + } + + +class TestTaskHandlerBinding: + def test_parsing_result_decodes_bindings(self): + result = TypeAdapter(ToManager).validate_python( + { + "type": "DagFileParsingResult", + "fileloc": "/files/dags/etl.py", + "serialized_dags": [], + "task_handler_bindings": [_make_binding_body()], + } + ) + + assert isinstance(result, DagFileParsingResult) + assert result.task_handler_bindings == [ + TaskHandlerBinding( + dag_id="etl", + task_id="extract", + artifact_bundle_name="java-task-handlers", + artifact_rel_path="etl.jar", + ) + ] + assert result.probed_artifacts == [] + + def test_parsing_result_decodes_probed_artifacts(self): + result = TypeAdapter(ToManager).validate_python( + { + "type": "DagFileParsingResult", + "fileloc": "/files/dags/etl.py", + "serialized_dags": [], + "probed_artifacts": [_make_known_artifact_body()], + } + ) + + assert isinstance(result, DagFileParsingResult) + assert result.task_handler_bindings is None + assert result.probed_artifacts == [ + TaskHandlerArtifact( + bundle_name="java-task-handlers", + relative_fileloc="etl.jar", + size_bytes=1024, + cache_digest="ab12", + task_handlers={ + "etl": [ + TaskHandlerDeclaration( + task_id="extract", + binding="positional", + params=[TaskHandlerParam(name=None)], + ) + ] + }, + ) + ] + + def test_rejects_a_path_wider_than_its_column(self): + with pytest.raises(ValidationError, match="artifact_rel_path"): + TaskHandlerBinding.model_validate(_make_binding_body(artifact_rel_path="a" * 2001)) + + +_STUB_DAG = """ +from airflow.sdk import dag, task + + +@task.stub(queue="fake-queue") +def extract(): ... + + +@task.stub(queue="fake-queue") +def load(count: int): ... + + +@task.stub(queue="elsewhere") +def report(): ... + + +@dag(dag_id="etl", schedule=None) +def etl(): + load(extract()) + report() + + +etl() +""" + +_PYTHON_DAG = """ +from airflow.sdk import DAG + +python_only = DAG("python_only", schedule=None) +""" + +_STUB_TASK_HANDLERS = { + "etl": [ + {"task_id": "extract", "binding": "positional", "params": []}, + {"task_id": "load", "binding": "positional", "params": [{"name": "count", "value_schema": INTEGER}]}, + ] +} + +_TASK_HANDLER_PROBLEMS = "Stub tasks in etl.py do not match their task handlers:" + + +def _is_running(process: psutil.Process) -> bool: + # A killed process stays a zombie until reaped, which an init that does not reap never does. + try: + return process.status() != psutil.STATUS_ZOMBIE + except psutil.NoSuchProcess: + return False + + +def _bind(task_id: str) -> TaskHandlerBinding: + return TaskHandlerBinding( + dag_id="etl", task_id=task_id, artifact_bundle_name="task-handlers", artifact_rel_path="etl.artifact" + ) + + +def _record(artifact: pathlib.Path) -> TaskHandlerArtifact: + content = artifact.read_bytes() + spec = json.loads(content) + return TaskHandlerArtifact.model_validate( + { + "bundle_name": "task-handlers", + "relative_fileloc": artifact.name, + "size_bytes": len(content), + "cache_digest": spec.get("cache_digest"), + "task_handlers": spec["task_handlers"], + } + ) + + +@pytest.mark.usefixtures("disable_load_example") +@patch.object( + FakeCoordinator, "parse_task_handler", autospec=True, side_effect=play_runtime(reply_with_task_handlers) +) +class TestParseFileTaskHandlers: + @pytest.fixture + def dag_bundle(self, tmp_path) -> pathlib.Path: + path = tmp_path / "dags" + path.mkdir() + return path + + @pytest.fixture + def artifacts(self, tmp_path) -> pathlib.Path: + path = tmp_path / "artifacts" + path.mkdir() + return path + + @staticmethod + def _parse(dag_bundle: pathlib.Path, source: str, **kwargs) -> DagFileParsingResult: + dag_file = dag_bundle / "etl.py" + dag_file.write_text(source) + result = _parse_file( + DagFileParseRequest(file=os.fspath(dag_file), bundle_path=dag_bundle, bundle_name="dags"), + log=structlog.get_logger(), + **kwargs, + ) + assert result is not None + return result + + def test_parse_file_binds_stub_tasks_to_their_task_handlers( + self, mock_parse_task_handler, dag_bundle, artifacts + ): + artifact = write_artifact( + artifacts / "etl.artifact", cache_digest="d1", task_handlers=_STUB_TASK_HANDLERS + ) + + with task_handler_config(dag_bundle, artifacts): + result = self._parse(dag_bundle, _STUB_DAG) + + assert result.import_errors == {} + assert result.task_handler_bindings == [_bind("extract"), _bind("load")] + assert result.probed_artifacts == [_record(artifact)] + + def test_parse_file_reports_a_missing_task_handler_as_an_import_error( + self, mock_parse_task_handler, dag_bundle, artifacts + ): + artifact = write_artifact( + artifacts / "etl.artifact", task_handlers={"etl": _STUB_TASK_HANDLERS["etl"][:1]} + ) + + with task_handler_config(dag_bundle, artifacts): + result = self._parse(dag_bundle, _STUB_DAG) + + assert result.import_errors == { + "etl.py": f"{_TASK_HANDLER_PROBLEMS}\n" + "- Dag 'etl', task 'load': no artifact in Dag bundle 'task-handlers' registers it" + } + assert result.task_handler_bindings is None + assert result.probed_artifacts == [_record(artifact)] + assert [dag.dag_id for dag in result.serialized_dags] == ["etl"] + + def test_parse_file_appends_a_task_handler_problem_to_a_dagbag_import_error( + self, mock_parse_task_handler, dag_bundle, artifacts + ): + write_artifact(artifacts / "etl.artifact", task_handlers={"etl": _STUB_TASK_HANDLERS["etl"][:1]}) + to_dict = DagSerialization.to_dict + + def fail_for_the_broken_dag(dag): + if dag.dag_id == "broken": + raise RuntimeError("cannot serialize") + return to_dict(dag) + + with ( + task_handler_config(dag_bundle, artifacts), + patch.object(DagSerialization, "to_dict", autospec=True, side_effect=fail_for_the_broken_dag), + ): + result = self._parse( + dag_bundle, + f"{_STUB_DAG}\nfrom airflow.sdk import DAG\n\nbroken = DAG('broken', schedule=None)\n", + ) + + assert result.import_errors["etl.py"].startswith("Traceback (most recent call last):") + assert result.import_errors["etl.py"].endswith( + f"RuntimeError: cannot serialize\n\n{_TASK_HANDLER_PROBLEMS}\n" + "- Dag 'etl', task 'load': no artifact in Dag bundle 'task-handlers' registers it" + ) + + @patch( + "airflow.dag_processing.task_handler_resolution.resolve_task_handlers", + autospec=True, + side_effect=RuntimeError("boom"), + ) + def test_parse_file_reports_an_unexpected_error_in_the_check_as_an_import_error( + self, mock_resolve, mock_parse_task_handler, dag_bundle, artifacts, cap_structlog + ): + with task_handler_config(dag_bundle, artifacts): + result = self._parse(dag_bundle, _STUB_DAG) + + assert result.import_errors == { + "etl.py": f"{_TASK_HANDLER_PROBLEMS}\n" + "- Unexpected error while checking the stub tasks: RuntimeError: boom" + } + assert result.task_handler_bindings is None + assert result.probed_artifacts == [] + assert [dag.dag_id for dag in result.serialized_dags] == ["etl"] + [entry] = [ + e + for e in cap_structlog.entries + if e["event"] == "Failed to check the stub tasks against their task handlers" + ] + assert entry["exception"][0]["exc_type"] == "RuntimeError" + + @pytest.mark.parametrize( + ("started", "deadline"), + [ + pytest.param(950.0, 1040.0, id="started-before-the-call"), + pytest.param(None, 1090.0, id="started-at-the-call"), + ], + ) + @conf_vars({("dag_processor", "dag_file_processor_timeout"): "100"}) + @patch( + "airflow.dag_processing.task_handler_resolution.resolve_task_handlers", + autospec=True, + return_value=TaskHandlerResolution(bindings=[], probed_artifacts=[], import_errors={}), + ) + @patch("airflow.dag_processing.processor.check_dag_file_stability", autospec=True) + @patch("airflow.dag_processing.processor.time", autospec=True) + def test_the_check_ends_at_90_percent_of_the_timeout_from_the_start_of_the_parse( + self, + mock_time, + mock_check, + mock_resolve, + mock_parse_task_handler, + dag_bundle, + artifacts, + started, + deadline, + ): + mock_time.monotonic.return_value = 1000.0 + + def check_a_while_later(*args, **kwargs): + # A budget counted from after this, or from the DagBag import, would end later. + mock_time.monotonic.return_value = 2000.0 + return check_dag_file_stability(*args, **kwargs) + + mock_check.side_effect = check_a_while_later + + with task_handler_config(dag_bundle, artifacts): + self._parse(dag_bundle, _STUB_DAG, **({} if started is None else {"started": started})) + + assert mock_resolve.call_args.kwargs["deadline"] == pytest.approx(deadline) + + @pytest.mark.parametrize( + ("created", "started"), + [ + pytest.param(4995.0, 995.0, id="created-before-the-parse"), + pytest.param(4900.0, 900.0, id="created-a-timeout-ago"), + pytest.param(4899.0, 1000.0, id="created-over-a-timeout-ago"), + pytest.param(5005.0, 1000.0, id="created-in-the-future"), + pytest.param(None, 1000.0, id="creation-time-unreadable"), + ], + ) + @conf_vars({("dag_processor", "dag_file_processor_timeout"): "100"}) + @patch("airflow.dag_processing.processor._parse_file", autospec=True, return_value=None) + @patch("airflow.dag_processing.processor.psutil", autospec=True) + @patch("airflow.dag_processing.processor.time", autospec=True) + def test_the_entrypoint_counts_the_parse_from_the_creation_of_the_process( + self, mock_time, mock_psutil, mock_parse_file, mock_parse_task_handler, monkeypatch, created, started + ): + mock_time.monotonic.return_value = 1000.0 + mock_time.time.return_value = 5000.0 + if created is None: + mock_psutil.Error = psutil.Error + mock_psutil.Process.side_effect = psutil.NoSuchProcess(1234) + else: + mock_psutil.Process.return_value.create_time.return_value = created + request = DagFileParseRequest( + file="/files/dags/etl.py", bundle_path="/files/dags", bundle_name="dags" + ) + # The entrypoint sets the process context and SUPERVISOR_COMMS, which must not outlast the test. + monkeypatch.setenv("_AIRFLOW_PROCESS_CONTEXT", "server") + monkeypatch.setattr(task_runner, "SUPERVISOR_COMMS", None, raising=False) + + with patch.object(comms, "CommsDecoder") as mock_decoder: + mock_decoder.__getitem__.return_value.return_value._get_response.return_value = request + _parse_file_entrypoint() + + assert mock_parse_file.call_args.kwargs["started"] == pytest.approx(started) + + def test_parse_file_reports_an_invalid_queue_routing( + self, mock_parse_task_handler, dag_bundle, artifacts + ): + with ( + task_handler_config(dag_bundle, artifacts), + conf_vars({("sdk", "queue_to_coordinator"): "{"}), + ): + result = self._parse(dag_bundle, _STUB_DAG) + + [error] = result.import_errors.values() + assert error.startswith( + f"{_TASK_HANDLER_PROBLEMS}\n- Cannot load [sdk] coordinators: " + "AirflowConfigException: Unable to parse [sdk] 'queue_to_coordinator' as valid json" + ) + assert result.task_handler_bindings is None + + @pytest.mark.parametrize( + ("source", "queue_to_coordinator", "expected"), + [ + pytest.param(_PYTHON_DAG, None, [], id="no-stub-tasks"), + pytest.param( + _STUB_DAG, + None, + [_bind("extract"), _bind("load")], + id="unrouted-stub-task", + ), + pytest.param(_STUB_DAG, {}, None, id="no-queue-routing"), + ], + ) + def test_parse_file_task_handler_bindings( + self, mock_parse_task_handler, dag_bundle, artifacts, source, queue_to_coordinator, expected + ): + # report's queue is routed in no case, and no artifact registers it. + write_artifact(artifacts / "etl.artifact", task_handlers=_STUB_TASK_HANDLERS) + + with task_handler_config(dag_bundle, artifacts, queue_to_coordinator=queue_to_coordinator): + result = self._parse(dag_bundle, source) + + assert result.import_errors == {} + assert result.task_handler_bindings == expected + assert len(result.probed_artifacts) == (1 if expected else 0) + + @pytest.mark.skipif(not sys.platform.startswith("linux"), reason="the parent-death signal is Linux only") + @pytest.mark.execution_timeout(60) + def test_a_killed_dag_file_processor_takes_its_probe_runtime_with_it( + self, mock_parse_task_handler, dag_bundle, artifacts + ): + mock_parse_task_handler.side_effect = SubprocessCoordinator.parse_task_handler + # A unique duration tells this runtime apart from any other sleep. + duration = f"600.{uuid.uuid4().int % 10**9}" + write_artifact(artifacts / "etl.artifact", argv=["/bin/sh", "-c", f"exec sleep {duration}"]) + (dag_bundle / "etl.py").write_text(_STUB_DAG) + + def find_runtime() -> psutil.Process | None: + for process in psutil.process_iter(["cmdline"]): + if process.info["cmdline"] == ["sleep", duration]: + return process + return None + + with task_handler_config(dag_bundle, artifacts): + proc = DagFileProcessorProcess.start( + id=uuid.uuid4(), + path=dag_bundle / "etl.py", + bundle_path=dag_bundle, + bundle_name="dags", + dag_file_rel_path="etl.py", + callbacks=[], + logger=structlog.get_logger(), + logger_filehandle=MagicMock(spec=BinaryIO), + client=MagicMock(spec=Client), + ) + runtime = None + try: + deadline = time.monotonic() + 30 + while (runtime := find_runtime()) is None: + assert time.monotonic() < deadline, "the probe runtime did not start" + proc._service_subprocess(max_wait_time=0.1) + + os.kill(proc.pid, signal.SIGKILL) + + deadline = time.monotonic() + 10 + while _is_running(runtime): + assert time.monotonic() < deadline, "the probe runtime outlived the Dag file processor" + time.sleep(0.05) + finally: + if runtime is not None and _is_running(runtime): + runtime.kill() + while not proc.is_ready: + proc._service_subprocess(max_wait_time=0.1) + proc.close() + class TestDagFileProcessorProcess: def test_registered_message_types(self): - expected = set(typing.get_args(typing.get_args(ToManager)[0])) + expected = set(typing.get_args(typing.get_args(ToManager)[0])) - {TaskHandlerParsingResult} assert set(DagFileProcessorProcess._request_handlers) == expected @pytest.mark.parametrize( "message_type", - sorted(set(typing.get_args(typing.get_args(ToManager)[0])) - {DagFileParsingResult}, key=str), + sorted( + set(typing.get_args(typing.get_args(ToManager)[0])) + - {DagFileParsingResult, TaskHandlerParsingResult}, + key=str, + ), ids=lambda message_type: message_type.__name__, ) def test_reuses_shared_request_handlers(self, message_type): @@ -2399,6 +2971,38 @@ def test_dispatch_parsing_result(self, send_msg, proc): assert proc.parsing_result is result send_msg.assert_called_once_with(proc, None, request_id=42, error=None) + @patch.object(DagFileProcessorProcess, "send_msg", autospec=True) + def test_rejects_task_handler_parsing_result(self, send_msg, proc): + msg = proc.decoder.validate_python( + { + "type": "TaskHandlerParsingResult", + "fileloc": "handlers.jar", + "task_handlers": { + "etl": [ + { + "task_id": "extract", + "binding": "positional", + "params": [ + {"name": "day", "value_schema": {"type": "string"}}, + {"name": "limit"}, + ], + } + ] + }, + } + ) + assert isinstance(msg, TaskHandlerParsingResult) + + proc._handle_request(msg, structlog.get_logger(), req_id=42) + + send_msg.assert_called_once_with( + proc, + None, + request_id=42, + error=comms.ErrorResponse(detail={"status_code": 400, "message": "Unhandled request"}), + ) + assert proc.parsing_result is None + @patch.object(DagFileProcessorProcess, "send_msg", autospec=True) def test_previous_successful_run_uses_process_id(self, send_msg, proc): proc.client.task_instances.get_previous_successful_dagrun.return_value = ( @@ -2546,6 +3150,25 @@ def test_handle_request_get_variable_masks_value_with_key(self, proc): } +class TestBaseDagFileProcessorProcess: + @patch.object(BaseDagFileProcessorProcess, "cleanup_sockets_after_kill", autospec=True) + def test_close_without_log_file_handle_cleans_up_sockets(self, cleanup_sockets_after_kill): + proc = BaseDagFileProcessorProcess( + process_log=structlog.get_logger(), + id=uuid.uuid4(), + pid=1234, + process=MagicMock(spec=supervisor.ProcessTracker), + stdin=MagicMock(spec=socket), + client=MagicMock(spec=Client), + bundle_name="mybundle", + dag_file_rel_path="dags/my_dag.py", + ) + + proc.close() + + cleanup_sockets_after_kill.assert_called_once_with(proc) + + class TestMultiTeamCallbackMetrics: """Tests for team_name tag on dag.callback_exceptions in multi-team mode.""" diff --git a/airflow-core/tests/unit/dag_processing/test_task_handler_fast_path.py b/airflow-core/tests/unit/dag_processing/test_task_handler_fast_path.py new file mode 100644 index 0000000000000..7340d0277d458 --- /dev/null +++ b/airflow-core/tests/unit/dag_processing/test_task_handler_fast_path.py @@ -0,0 +1,146 @@ +# Licensed to the Apache Software Foundation (ASF) under one +# or more contributor license agreements. See the NOTICE file +# distributed with this work for additional information +# regarding copyright ownership. The ASF licenses this file +# to you under the Apache License, Version 2.0 (the +# "License"); you may not use this file except in compliance +# with the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, +# software distributed under the License is distributed on an +# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +# KIND, either express or implied. See the License for the +# specific language governing permissions and limitations +# under the License. +from __future__ import annotations + +import pytest + +from airflow.dag_processing.processor import TaskHandlerArtifact, TaskHandlerDeclaration +from airflow.dag_processing.task_handler_fast_path import TaskHandlerProbePlan, plan_task_handler_probes +from airflow.sdk.execution_time.coordinator import TaskHandlerCandidate + +BUNDLE_NAME = "java-task-handlers" +DIGEST = "a" * 64 +OTHER_DIGEST = "b" * 64 + + +def _make_candidate( + rel_path: str = "handlers.jar", + *, + size_bytes: int = 100, + cache_digest: str | None = DIGEST, + error: str | None = None, +) -> TaskHandlerCandidate: + return TaskHandlerCandidate( + rel_path=rel_path, size_bytes=size_bytes, cache_digest=cache_digest, error=error + ) + + +def _make_artifact( + rel_path: str = "handlers.jar", + *, + bundle_name: str = BUNDLE_NAME, + size_bytes: int = 100, + cache_digest: str | None = DIGEST, +) -> TaskHandlerArtifact: + # Two Dags, as if several Dag files resolve against the artifact. + return TaskHandlerArtifact( + bundle_name=bundle_name, + relative_fileloc=rel_path, + size_bytes=size_bytes, + cache_digest=cache_digest, + task_handlers={ + "etl": [TaskHandlerDeclaration(task_id="extract", binding="named", params=[])], + "reporting": [TaskHandlerDeclaration(task_id="publish", binding="positional", params=[])], + }, + ) + + +@pytest.mark.parametrize( + ("candidate", "known", "outcome"), + [ + pytest.param(_make_candidate(), [_make_artifact()], "cached", id="same-size-and-digest"), + pytest.param(_make_candidate(), [], "probe", id="added"), + pytest.param(_make_candidate(size_bytes=101), [_make_artifact()], "probe", id="resized-same-digest"), + pytest.param( + _make_candidate(cache_digest=OTHER_DIGEST), + [_make_artifact()], + "probe", + id="digest-changed-same-size", + ), + pytest.param(_make_candidate(cache_digest=None), [_make_artifact()], "probe", id="stores-no-digest"), + pytest.param(_make_candidate(cache_digest=None), [], "probe", id="stores-no-digest-and-unknown"), + pytest.param( + _make_candidate(), [_make_artifact(cache_digest=None)], "probe", id="recorded-without-a-digest" + ), + pytest.param( + _make_candidate(cache_digest=None), + [_make_artifact(cache_digest=None)], + "probe", + id="neither-has-a-digest", + ), + pytest.param( + _make_candidate(), + [_make_artifact(bundle_name="other-bundle")], + "probe", + id="same-path-other-bundle", + ), + pytest.param( + _make_candidate( + error="handlers.jar has an Airflow-Cache-Digest manifest attribute but no Main-Class" + ), + [_make_artifact()], + "rejected", + id="listing-error-despite-a-matching-record", + ), + ], +) +def test_decides_each_candidate(candidate, known, outcome): + plan = plan_task_handler_probes(bundle_name=BUNDLE_NAME, candidates=[candidate], known_artifacts=known) + + expected = { + "cached": TaskHandlerProbePlan(cached=known, probe=[], rejected=[]), + "probe": TaskHandlerProbePlan(cached=[], probe=[candidate], rejected=[]), + "rejected": TaskHandlerProbePlan(cached=[], probe=[], rejected=[candidate]), + } + assert plan == expected[outcome] + + +def test_reuses_the_whole_recorded_answer(): + artifact = _make_artifact() + + plan = plan_task_handler_probes( + bundle_name=BUNDLE_NAME, candidates=[_make_candidate()], known_artifacts=[artifact] + ) + + assert plan.cached[0] is artifact + assert set(plan.cached[0].task_handlers) == {"etl", "reporting"} + + +def test_decides_candidates_independently_in_listing_order(): + unchanged_a = _make_candidate("a.jar") + rebuilt = _make_candidate("b.jar", cache_digest=OTHER_DIGEST) + unusable = _make_candidate( + "c.jar", error="c.jar has an Airflow-Cache-Digest manifest attribute but no Main-Class" + ) + unchanged_d = _make_candidate("d.jar") + added = _make_candidate("e.jar") + known = [ + _make_artifact("d.jar"), + _make_artifact("gone.jar"), + _make_artifact("b.jar"), + _make_artifact("a.jar"), + ] + + plan = plan_task_handler_probes( + bundle_name=BUNDLE_NAME, + candidates=[unchanged_a, rebuilt, unusable, unchanged_d, added], + known_artifacts=known, + ) + + assert plan == TaskHandlerProbePlan( + cached=[known[3], known[0]], probe=[rebuilt, added], rejected=[unusable] + ) diff --git a/airflow-core/tests/unit/dag_processing/test_task_handler_processor.py b/airflow-core/tests/unit/dag_processing/test_task_handler_processor.py new file mode 100644 index 0000000000000..ff115346bae0e --- /dev/null +++ b/airflow-core/tests/unit/dag_processing/test_task_handler_processor.py @@ -0,0 +1,904 @@ +# +# Licensed to the Apache Software Foundation (ASF) under one +# or more contributor license agreements. See the NOTICE file +# distributed with this work for additional information +# regarding copyright ownership. The ASF licenses this file +# to you under the Apache License, Version 2.0 (the +# "License"); you may not use this file except in compliance +# with the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, +# software distributed under the License is distributed on an +# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +# KIND, either express or implied. See the License for the +# specific language governing permissions and limitations +# under the License. +from __future__ import annotations + +import contextlib +import io +import os +import selectors +import signal +import socket +import sys +import threading +import time +import uuid +from pathlib import Path +from unittest.mock import ANY, MagicMock, patch + +import psutil +import pytest +import structlog +from pydantic import TypeAdapter +from structlog.typing import FilteringBoundLogger + +from airflow.configuration import conf +from airflow.dag_processing.processor import ( + BaseDagFileProcessorProcess, + DagFileParseRequest, + DagFileParsingResult, + DagFileProcessorProcess, + TaskHandlerDeclaration, + TaskHandlerParam, + TaskHandlerParseRequest, + TaskHandlerParsingResult, + ToDagProcessor, + ToManager, +) +from airflow.dag_processing.task_handler_processor import ( + LangSDKRuntimeSchemaVersion, + LangSDKTaskHandlerProcessorProcess, + TaskHandlerProbeStopped, + _get_import_timeout, +) +from airflow.sdk.api.client import Client, VariableOperations +from airflow.sdk.api.datamodels._generated import VariableResponse +from airflow.sdk.coordinators._subprocess import SubprocessCoordinator +from airflow.sdk.exceptions import AirflowRuntimeError, ErrorType +from airflow.sdk.execution_time import supervisor, task_runner +from airflow.sdk.execution_time.comms import ( + CommsDecoder, + ErrorResponse, + GetVariable, + MaskSecret, + VariableResult, + _RequestFrame, +) + +from tests_common.test_utils.config import conf_vars +from unit.dag_processing.fake_task_handler_runtime import ( + FakeCoordinator, + fake_coordinator, + play_runtime, + write_artifact, +) + +# The oldest supervisor schema version, so the parse request is downgraded. +OLDEST_SCHEMA_VERSION = "2026-06-16" + + +def _reply_with(*task_ids: str, **result): + """Reply with a handler for each of *task_ids* under the Dag ``etl``.""" + + def reply(request: TaskHandlerParseRequest, comms) -> TaskHandlerParsingResult: + declarations = [ + TaskHandlerDeclaration(task_id=task_id, binding="positional", params=[]) for task_id in task_ids + ] + return TaskHandlerParsingResult(fileloc=request.file, task_handlers={"etl": declarations}, **result) + + return reply + + +def _get_task_ids(result: TaskHandlerParsingResult) -> list[str]: + return [declaration.task_id for declaration in result.task_handlers["etl"]] + + +def _get_open_fds() -> set[int]: + # Without /proc, as on macOS, this is empty, so the fd leak checks pass trivially. + return {int(fd) for fd in os.listdir("/proc/self/fd")} if os.path.isdir("/proc/self/fd") else set() + + +# A runtime that exits once a child has left its process group; the child keeps the runtime's +# stdout and stderr open and writes its pid to argv[1]. +_LEAVE_A_CHILD_OUTSIDE_THE_GROUP = """ +import os, sys, time +read_fd, write_fd = os.pipe() +if os.fork(): + os.read(read_fd, 1) + os._exit(0) +os.setsid() +with open(sys.argv[1], "w") as f: + f.write(str(os.getpid())) +os.write(write_fd, b"x") +time.sleep(30) +""" + + +def _is_running(pid: int) -> bool: + try: + return psutil.Process(pid).status() != psutil.STATUS_ZOMBIE + except psutil.NoSuchProcess: + return False + + +@pytest.fixture(autouse=True) +def _coordinator(): + with fake_coordinator(): + yield + + +@pytest.fixture +def supervisor_comms(monkeypatch): + """Give this process a supervisor channel, as a Dag-parsing child has.""" + comms = MagicMock(spec=CommsDecoder) + monkeypatch.setattr(task_runner, "SUPERVISOR_COMMS", comms, raising=False) + return comms + + +def _start(tmp_path, selector, *, client: Client | None = None, **spec) -> LangSDKTaskHandlerProcessorProcess: + return LangSDKTaskHandlerProcessorProcess.start( + id=uuid.uuid4(), + coordinator="fake", + path=write_artifact(tmp_path / "etl.artifact", **spec), + bundle_path=tmp_path, + bundle_name="task-handlers", + artifact_rel_path="etl.artifact", + selector=selector, + logger=structlog.get_logger(), + client=client, + ) + + +@pytest.fixture +def parse(tmp_path): + """Probe ``etl.artifact`` under a caller's selector loop, and check that nothing is left open.""" + + def _parse(**kwargs) -> LangSDKTaskHandlerProcessorProcess: + fds_before = _get_open_fds() + with selectors.DefaultSelector() as selector: + proc = _start(tmp_path, selector, **kwargs) + deadline = time.monotonic() + 30 + while not proc.is_ready: + assert time.monotonic() < deadline, "the Lang-SDK parse did not finish" + proc._service_subprocess(max_wait_time=0.1) + assert selector.get_map() == {} + proc.close() + assert _get_open_fds() <= fds_before + return proc + + return _parse + + +def _block_until_killed(comms) -> None: + """Hang as a stuck runtime does, until it is killed or the parent closes the comm socket.""" + comms.socket.recv(1) + + +def _send_an_invalid_frame(request, comms) -> None: + comms.socket.sendall(bytes.fromhex("00000003c1c1c1")) + _block_until_killed(comms) + + +def _send_a_frame(comms, body: dict) -> None: + comms.socket.sendall(_RequestFrame(id=1, body=body).as_bytes()) + _block_until_killed(comms) + + +def _send_an_invalid_result(request, comms) -> None: + _send_a_frame( + comms, {"type": "TaskHandlerParsingResult", "fileloc": request.file, "task_handlers": "none"} + ) + + +def _send_a_start_message(request, comms) -> None: + _send_a_frame(comms, LangSDKRuntimeSchemaVersion(schema_version=None).model_dump()) + + +def _run( + tmp_path, *, coordinator: str = "fake", deadline: float | None = None, **spec +) -> TaskHandlerParsingResult: + return LangSDKTaskHandlerProcessorProcess.run( + coordinator=coordinator, + path=write_artifact(tmp_path / "etl.artifact", **spec), + bundle_path=tmp_path, + bundle_name="task-handlers", + artifact_rel_path="etl.artifact", + logger=structlog.get_logger(), + deadline=deadline, + ) + + +class TestLangSDKTaskHandlerProcessorProcess: + @patch.object(FakeCoordinator, "parse_task_handler", autospec=True) + def test_collects_the_result_the_runtime_returns( + self, mock_parse_task_handler, parse, tmp_path, cap_structlog + ): + mock_parse_task_handler.side_effect = play_runtime( + _reply_with("extract"), + schema_version=OLDEST_SCHEMA_VERSION, + log_lines=[{"event": "Registering handlers", "level": "info"}], + ) + + proc = parse() + + assert proc.parsing_result.fileloc == os.fspath(tmp_path / "etl.artifact") + assert proc.parsing_result.import_errors is None + assert _get_task_ids(proc.parsing_result) == ["extract"] + assert proc._subprocess_schema_version == OLDEST_SCHEMA_VERSION + assert "Registering handlers" in cap_structlog + + @patch("airflow.dag_processing.task_handler_processor._is_connection_from_pid", autospec=True) + @patch.object(FakeCoordinator, "parse_task_handler", autospec=True) + def test_a_connection_is_used_once_it_is_verified(self, mock_parse_task_handler, mock_owned, tmp_path): + mock_parse_task_handler.side_effect = play_runtime(_reply_with("extract")) + mock_owned.return_value = False + with selectors.DefaultSelector() as selector: + proc = _start(tmp_path, selector) + child_stdin = proc.stdin + deadline = time.monotonic() + 30 + while len(proc._unverified_connections) < 2: + assert proc.stdin is child_stdin, "an unverified connection was used" + assert time.monotonic() < deadline, "the runtime did not connect" + proc._service_subprocess(max_wait_time=0.1) + assert proc.stdin is child_stdin + + mock_owned.return_value = True + while not proc.is_ready: + assert time.monotonic() < deadline, "the Lang-SDK parse did not finish" + proc._service_subprocess(max_wait_time=0.1) + proc.close() + + assert _get_task_ids(proc.parsing_result) == ["extract"] + + @patch.object( + LangSDKTaskHandlerProcessorProcess, + "send_msg", + autospec=True, + side_effect=OSError("cannot send the start request"), + ) + def test_a_start_that_fails_after_the_fork_leaves_nothing_behind(self, mock_send_msg, tmp_path): + fds_before = _get_open_fds() + children_before = {child.pid for child in psutil.Process().children()} + + with selectors.DefaultSelector() as selector: + with pytest.raises(OSError, match="cannot send the start request"): + _start(tmp_path, selector) + assert selector.get_map() == {} + + assert {child.pid for child in psutil.Process().children()} == children_before + assert _get_open_fds() <= fds_before + + @patch.object(BaseDagFileProcessorProcess, "start", autospec=True, side_effect=OSError("fork failed")) + def test_a_start_that_fails_before_the_fork_closes_its_listeners(self, mock_start, tmp_path): + fds_before = _get_open_fds() + + with selectors.DefaultSelector() as selector: + with pytest.raises(OSError, match="fork failed"): + _start(tmp_path, selector) + + listeners = mock_start.call_args.kwargs["listeners"] + assert [listener.fileno() for listener in listeners.values()] == [-1, -1] + assert _get_open_fds() <= fds_before + + @patch.object(FakeCoordinator, "parse_task_handler", autospec=True) + def test_requests_are_answered_by_the_client(self, mock_parse_task_handler, parse): + def reply(request, comms): + variable = comms.send(GetVariable(key="probe_var")) + return _reply_with(variable.value)(request, comms) + + mock_parse_task_handler.side_effect = play_runtime(reply) + client = MagicMock(spec=Client) + client.variables = MagicMock(spec=VariableOperations) + client.variables.get.return_value = VariableResponse(key="probe_var", value="from-the-client") + + proc = parse(client=client) + + assert _get_task_ids(proc.parsing_result) == ["from-the-client"] + + @pytest.mark.parametrize( + ("spec", "reply", "error"), + [ + pytest.param( + {"command_error": "no runtime"}, + None, + "Cannot start the Lang-SDK runtime: FileNotFoundError: no runtime", + id="command-not-resolved", + ), + pytest.param( + {"argv": ["/no/such/runtime"]}, + None, + "Cannot start the Lang-SDK runtime: FileNotFoundError: " + "[Errno 2] No such file or directory: '/no/such/runtime'", + id="exec-failed", + ), + pytest.param( + {"argv": ["/bin/sh", "-c", "exit 3"]}, + None, + "The Lang-SDK runtime exited with code 3 without a parse result", + id="exits-before-connecting", + ), + pytest.param( + {}, + lambda request, comms: None, + "The Lang-SDK runtime exited with code 0 without a parse result", + id="exits-without-a-result", + ), + pytest.param( + {}, + _send_an_invalid_frame, + "The Lang-SDK runtime sent an invalid frame: MessagePack data is malformed: " + "invalid opcode '\\xc1' (byte 0)", + id="invalid-frame", + ), + ], + ) + @patch.object(FakeCoordinator, "parse_task_handler", autospec=True) + def test_a_failed_parse_is_an_import_error( + self, mock_parse_task_handler, parse, tmp_path, spec, reply, error + ): + mock_parse_task_handler.side_effect = ( + play_runtime(reply) if reply else SubprocessCoordinator.parse_task_handler + ) + + proc = parse(**spec) + + assert proc.parsing_result == TaskHandlerParsingResult( + fileloc=os.fspath(tmp_path / "etl.artifact"), + task_handlers={}, + import_errors={"etl.artifact": error}, + ) + + def test_a_coordinator_that_is_not_configured_is_an_import_error(self, tmp_path): + result = _run(tmp_path, coordinator="missing") + + assert result == TaskHandlerParsingResult( + fileloc=os.fspath(tmp_path / "etl.artifact"), + task_handlers={}, + import_errors={ + "etl.artifact": "Cannot start the Lang-SDK runtime: " + "InvalidCoordinatorError: No coordinator 'missing' in [sdk] coordinators" + }, + ) + + @patch.object( + FakeCoordinator, + "_build_parse_task_handler_command", + SubprocessCoordinator._build_parse_task_handler_command, + ) + def test_a_coordinator_that_does_not_parse_task_handlers_is_an_import_error(self, parse): + proc = parse() + + assert proc.parsing_result.import_errors == { + "etl.artifact": "Cannot start the Lang-SDK runtime: " + "NotImplementedError: FakeCoordinator does not parse task handlers" + } + + @pytest.mark.parametrize( + ("reply", "error"), + [ + pytest.param( + _send_an_invalid_result, + "TaskHandlerParsingResult.task_handlers\n Input should be a valid dictionary", + id="result", + ), + pytest.param( + _send_a_start_message, "does not match any of the expected tags", id="start-message" + ), + ], + ) + @patch.object(FakeCoordinator, "parse_task_handler", autospec=True) + def test_a_message_that_does_not_validate_is_an_import_error( + self, mock_parse_task_handler, parse, reply, error + ): + mock_parse_task_handler.side_effect = play_runtime(reply) + + proc = parse() + + [message] = proc.parsing_result.import_errors.values() + assert message.startswith("The Lang-SDK runtime sent a message that does not validate: ") + assert error in message + assert proc._exit_code == -signal.SIGKILL + + @patch.object(FakeCoordinator, "parse_task_handler", autospec=True) + def test_a_dag_file_parsing_result_is_not_a_task_handler_result(self, mock_parse_task_handler, tmp_path): + def reply(request, comms): + with pytest.raises(AirflowRuntimeError, match="Unhandled request"): + comms.send(DagFileParsingResult(fileloc=request.file, serialized_dags=[])) + + mock_parse_task_handler.side_effect = play_runtime(reply) + + result = _run(tmp_path) + + assert result.import_errors == { + "etl.artifact": "The Lang-SDK runtime exited with code 0 without a parse result" + } + + @patch.object( + FakeCoordinator, "parse_task_handler", autospec=True, side_effect=play_runtime(_send_an_invalid_frame) + ) + def test_killing_the_runtime_is_not_reported_as_out_of_memory( + self, mock_parse_task_handler, parse, cap_structlog + ): + proc = parse() + + assert proc._exit_code == -signal.SIGKILL + assert not any("Likely out of memory" in str(entry.get("event")) for entry in cap_structlog.entries) + + @patch("airflow.dag_processing.task_handler_processor._EXIT_GRACE_PERIOD", 0.5) + @patch.object(FakeCoordinator, "parse_task_handler", autospec=True) + def test_a_runtime_that_runs_on_after_its_result_is_killed( + self, mock_parse_task_handler, parse, cap_structlog + ): + def reply(request, comms): + comms.send(_reply_with("extract")(request, comms)) + _block_until_killed(comms) + + mock_parse_task_handler.side_effect = play_runtime(reply) + + proc = parse() + + assert _get_task_ids(proc.parsing_result) == ["extract"] + assert proc._exit_code == -signal.SIGKILL + assert "The Lang-SDK runtime did not exit after its parse result; killing it" in cap_structlog + + @patch("airflow.settings.get_dagbag_import_timeout", autospec=True, return_value=5) + def test_what_the_runtime_leaves_in_its_process_group_is_killed_when_it_exits( + self, mock_timeout, parse, tmp_path + ): + pid_file = tmp_path / "leftover.pid" + + proc = parse(argv=["/bin/sh", "-c", f"sleep 30 & echo $! > {pid_file}; exit 0"]) + + leftover = int(pid_file.read_text()) + try: + assert proc.parsing_result.import_errors == { + "etl.artifact": "The Lang-SDK runtime exited with code 0 without a parse result" + } + # A poll: the leftover is not a child to wait for, and time_machine cannot reach another process. + deadline = time.monotonic() + 10 + while _is_running(leftover): + assert time.monotonic() < deadline, "the runtime's leftover process is still running" + time.sleep(0.05) + finally: + with contextlib.suppress(ProcessLookupError): + os.kill(leftover, signal.SIGKILL) + + @patch("airflow.settings.get_dagbag_import_timeout", autospec=True, return_value=1) + def test_the_import_timeout_holds_after_the_runtime_exits(self, mock_timeout, parse, tmp_path): + pid_file = tmp_path / "leftover.pid" + + try: + proc = parse(argv=[sys.executable, "-c", _LEAVE_A_CHILD_OUTSIDE_THE_GROUP, os.fspath(pid_file)]) + finally: + if pid_file.exists(): + os.kill(int(pid_file.read_text()), signal.SIGKILL) + + assert proc.parsing_result.import_errors == { + "etl.artifact": f"The Lang-SDK runtime did not parse {tmp_path / 'etl.artifact'} within 1.0s" + } + assert proc._exit_code == 0 + + @pytest.mark.parametrize( + ("policy", "error"), + [ + pytest.param(RuntimeError("policy bug"), "RuntimeError: policy bug", id="raises"), + pytest.param( + lambda path: "30", + "TypeError: Value (30) from get_dagbag_import_timeout must be int or float", + id="not-a-number", + ), + ], + ) + @patch("airflow.settings.get_dagbag_import_timeout", autospec=True) + def test_a_failing_import_timeout_policy_is_an_import_error(self, mock_timeout, parse, policy, error): + mock_timeout.side_effect = policy + + proc = parse() + + assert proc.parsing_result.import_errors == { + "etl.artifact": f"Cannot start the Lang-SDK runtime: {error}" + } + + @pytest.mark.skipif(not Path("/proc/self/fd").is_dir(), reason="reads /proc") + @pytest.mark.parametrize("use_exec", [False, True], ids=["fork", "spawn"]) + def test_the_runtime_inherits_only_its_standard_streams(self, monkeypatch, tmp_path, use_exec): + if use_exec: + # The spawned interpreter finds the coordinator again from its environment. + monkeypatch.setattr(supervisor, "_should_use_exec", lambda: True) + monkeypatch.setenv("PYTHONPATH", os.pathsep.join(sys.path)) + monkeypatch.setenv("AIRFLOW__SDK__COORDINATORS", conf.get("sdk", "coordinators")) + with selectors.DefaultSelector() as selector: + proc = _start(tmp_path, selector, argv=["/bin/sh", "-c", "exec sleep 30"]) + deadline = time.monotonic() + 30 + while psutil.Process(proc.pid).name() != "sleep": + assert time.monotonic() < deadline, "the runtime did not start" + proc._service_subprocess(max_wait_time=0.1) + fd_dir = Path(f"/proc/{proc.pid}/fd") + fds = {fd.name: os.readlink(fd) for fd in fd_dir.iterdir()} + proc.kill(signal.SIGKILL) + proc.close() + + assert sorted(fds) == ["0", "1", "2"] + assert fds["0"] == "/dev/null" + + +def _probe_from_a_dag_parsing_child() -> None: + """Stand in for ``_parse_file_entrypoint``: probe the file it is asked to parse, and return the result.""" + comms_decoder = CommsDecoder[ToDagProcessor, ToManager](body_decoder=TypeAdapter(ToDagProcessor)) + request = comms_decoder._get_response() + assert isinstance(request, DagFileParseRequest) + task_runner.SUPERVISOR_COMMS = comms_decoder # type: ignore[assignment] + + result = LangSDKTaskHandlerProcessorProcess.run( + coordinator="fake", + path=request.file, + bundle_path=request.bundle_path, + bundle_name=request.bundle_name, + artifact_rel_path="etl.artifact", + logger=structlog.get_logger(logger_name="task"), + ) + comms_decoder.send( + DagFileParsingResult( + fileloc=request.file, serialized_dags=[], warnings=[result.model_dump(mode="json")] + ) + ) + + +class TestRequestsWithoutAClient: + @patch.object(FakeCoordinator, "parse_task_handler", autospec=True) + def test_a_request_is_relayed_to_the_parent(self, mock_parse_task_handler, parse, supervisor_comms): + def reply(request, comms): + variable = comms.send(GetVariable(key="probe_var")) + return _reply_with(variable.value)(request, comms) + + mock_parse_task_handler.side_effect = play_runtime(reply, schema_version=OLDEST_SCHEMA_VERSION) + # As decoded from the parent's frame, which always carries its type. + supervisor_comms.send.return_value = VariableResult.model_validate( + {"key": "probe_var", "value": "from-the-parent", "type": "VariableResult"} + ) + + proc = parse() + + assert _get_task_ids(proc.parsing_result) == ["from-the-parent"] + supervisor_comms.send.assert_called_once_with(GetVariable(key="probe_var")) + + @patch.object(FakeCoordinator, "parse_task_handler", autospec=True) + def test_the_parent_s_error_reaches_the_runtime(self, mock_parse_task_handler, parse, supervisor_comms): + def reply(request, comms): + with pytest.raises(AirflowRuntimeError) as ctx: + comms.send(GetVariable(key="probe_var")) + error = ctx.value.error + return _reply_with(f"{error.error.value}:{error.detail['key']}")(request, comms) + + mock_parse_task_handler.side_effect = play_runtime(reply) + supervisor_comms.send.side_effect = AirflowRuntimeError( + ErrorResponse(error=ErrorType.VARIABLE_NOT_FOUND, detail={"key": "probe_var"}) + ) + + proc = parse() + + assert _get_task_ids(proc.parsing_result) == [f"{ErrorType.VARIABLE_NOT_FOUND.value}:probe_var"] + + @patch.object(FakeCoordinator, "parse_task_handler", autospec=True) + def test_a_request_gets_an_error_without_a_parent(self, mock_parse_task_handler, monkeypatch, tmp_path): + monkeypatch.delattr(task_runner, "SUPERVISOR_COMMS", raising=False) + + def reply(request, comms): + with pytest.raises(AirflowRuntimeError) as ctx: + comms.send(GetVariable(key="probe_var")) + return _reply_with(ctx.value.error.detail["message"])(request, comms) + + mock_parse_task_handler.side_effect = play_runtime(reply) + + result = _run(tmp_path) + + assert _get_task_ids(result) == ["GetVariable is answered only in the Dag processor"] + + @patch("airflow.sdk.log._secrets_masker", autospec=True) + @patch.object(FakeCoordinator, "parse_task_handler", autospec=True) + def test_a_secret_is_masked_here_and_by_the_parent( + self, mock_parse_task_handler, mock_secrets_masker, parse, supervisor_comms + ): + def reply(request, comms): + comms.send(MaskSecret(value="probe-secret", name="probe_conn")) + return _reply_with("extract")(request, comms) + + mock_parse_task_handler.side_effect = play_runtime(reply) + + proc = parse() + + assert _get_task_ids(proc.parsing_result) == ["extract"] + mock_secrets_masker.return_value.add_mask.assert_called_once_with("probe-secret", "probe_conn") + supervisor_comms.send.assert_called_once_with(MaskSecret(value="probe-secret", name="probe_conn")) + + @patch.object(FakeCoordinator, "parse_task_handler", autospec=True) + def test_a_runtime_request_is_answered_by_the_dag_processor(self, mock_parse_task_handler, tmp_path): + def reply(request, comms): + variable = comms.send(GetVariable(key="probe_var")) + return _reply_with(variable.value)(request, comms) + + mock_parse_task_handler.side_effect = play_runtime(reply) + client = MagicMock(spec=Client) + client.variables = MagicMock(spec=VariableOperations) + client.variables.get.return_value = VariableResponse(key="probe_var", value="from-the-dag-processor") + artifact = write_artifact(tmp_path / "etl.artifact") + + proc = DagFileProcessorProcess.start( + id=1, + path=artifact, + bundle_path=tmp_path, + bundle_name="task-handlers", + dag_file_rel_path="etl.artifact", + callbacks=[], + target=_probe_from_a_dag_parsing_child, + logger=structlog.get_logger(), + logger_filehandle=io.BytesIO(), + client=client, + ) + deadline = time.monotonic() + 30 + while not proc.is_ready: + assert time.monotonic() < deadline, "the Dag-parsing child did not finish" + proc._service_subprocess(max_wait_time=0.1) + proc.close() + + client.variables.get.assert_called_once_with("probe_var") + [probe_result] = proc.parsing_result.warnings + assert TaskHandlerParsingResult.model_validate(probe_result) == TaskHandlerParsingResult( + fileloc=os.fspath(artifact), + task_handlers={ + "etl": [ + TaskHandlerDeclaration(task_id="from-the-dag-processor", binding="positional", params=[]) + ] + }, + ) + + +class TestRun: + @patch.object(FakeCoordinator, "parse_task_handler", autospec=True) + def test_returns_every_handler_the_runtime_registers(self, mock_parse_task_handler, tmp_path): + def reply(request, comms): + # Echo what the request carries, so the test can check it. + param = TaskHandlerParam(name=request.bundle_name, value_schema={"type": "string"}) + return TaskHandlerParsingResult( + fileloc=request.file, + task_handlers={ + dag_id: [ + TaskHandlerDeclaration( + task_id=os.fspath(request.bundle_path), binding="positional", params=[param] + ) + ] + for dag_id in ("etl", "report") + }, + ) + + mock_parse_task_handler.side_effect = play_runtime(reply) + + result = _run(tmp_path) + + param = TaskHandlerParam(name="task-handlers", value_schema={"type": "string"}) + declaration = TaskHandlerDeclaration( + task_id=os.fspath(tmp_path), binding="positional", params=[param] + ) + assert result == TaskHandlerParsingResult( + fileloc=os.fspath(tmp_path / "etl.artifact"), + task_handlers={"etl": [declaration], "report": [declaration]}, + ) + + @pytest.mark.execution_timeout(30) + @pytest.mark.parametrize("connected", [False, True], ids=["before-connecting", "after-connecting"]) + @patch("airflow.settings.get_dagbag_import_timeout", autospec=True, return_value=1) + @patch.object( + LangSDKTaskHandlerProcessorProcess, + "close", + autospec=True, + side_effect=LangSDKTaskHandlerProcessorProcess.close, + ) + @patch.object(FakeCoordinator, "parse_task_handler", autospec=True) + def test_a_parse_past_the_import_timeout_is_killed( + self, mock_parse_task_handler, mock_close, mock_timeout, tmp_path, connected + ): + mock_parse_task_handler.side_effect = ( + play_runtime(lambda request, comms: _block_until_killed(comms)) + if connected + else SubprocessCoordinator.parse_task_handler + ) + # A runtime that never connects leaves both listeners open when it is killed. + fds_before = _get_open_fds() + + result = _run(tmp_path, argv=["/bin/sh", "-c", "exec sleep 60"]) + + assert result.import_errors == { + "etl.artifact": f"The Lang-SDK runtime did not parse {tmp_path / 'etl.artifact'} within 1.0s" + } + assert not isinstance(result, TaskHandlerProbeStopped) + [proc] = [c.args[0] for c in mock_close.call_args_list] + assert proc._exit_code == -9 + assert not proc._open_sockets + assert _get_open_fds() <= fds_before + + @pytest.mark.execution_timeout(30) + @patch("airflow.settings.get_dagbag_import_timeout", autospec=True, return_value=1) + @patch.object(FakeCoordinator, "parse_task_handler", autospec=True) + def test_a_runtime_that_stops_mid_frame_is_timed_out( + self, mock_parse_task_handler, mock_timeout, tmp_path + ): + def reply(request, comms): + comms.socket.sendall((100).to_bytes(4, byteorder="big")) + _block_until_killed(comms) + + mock_parse_task_handler.side_effect = play_runtime(reply) + + result = _run(tmp_path) + + assert result.import_errors == { + "etl.artifact": f"The Lang-SDK runtime did not parse {tmp_path / 'etl.artifact'} within 1.0s" + } + + @pytest.mark.execution_timeout(30) + @conf_vars({("dag_processor", "dag_file_processor_timeout"): "1"}) + @patch.object( + FakeCoordinator, + "_build_parse_task_handler_command", + autospec=True, + side_effect=lambda self, *, path: threading.Event().wait(), + ) + def test_the_dag_file_processor_timeout_applies_until_the_import_timeout_is_reported( + self, mock_build_command, tmp_path + ): + result = _run(tmp_path) + + assert result.import_errors == { + "etl.artifact": f"The Lang-SDK runtime did not parse {tmp_path / 'etl.artifact'} within 1.0s" + } + + @pytest.mark.execution_timeout(30) + @pytest.mark.parametrize("reported", [False, True], ids=["before-the-import-timeout", "after-it"]) + @patch.object( + LangSDKTaskHandlerProcessorProcess, + "close", + autospec=True, + side_effect=LangSDKTaskHandlerProcessorProcess.close, + ) + @patch.object(FakeCoordinator, "parse_task_handler", autospec=True) + def test_a_probe_past_its_deadline_is_killed( + self, mock_parse_task_handler, mock_close, tmp_path, reported + ): + if reported: + mock_parse_task_handler.side_effect = play_runtime( + lambda request, comms: _block_until_killed(comms) + ) + else: + mock_parse_task_handler.side_effect = lambda self, **kwargs: threading.Event().wait() + + result = _run(tmp_path, deadline=time.monotonic() + 1) + + assert isinstance(result, TaskHandlerProbeStopped) + assert result.import_errors == { + "etl.artifact": f"The Lang-SDK runtime did not parse {tmp_path / 'etl.artifact'} by its deadline" + } + [proc] = [c.args[0] for c in mock_close.call_args_list] + assert proc._exit_code == -signal.SIGKILL + + @pytest.mark.execution_timeout(30) + @patch.object(FakeCoordinator, "parse_task_handler", autospec=True) + def test_an_error_the_runtime_reported_is_kept_when_the_deadline_stops_it( + self, mock_parse_task_handler, tmp_path + ): + def reply(request, comms): + comms.send(_reply_with(import_errors={"etl.artifact": "handler registry failed"})(request, comms)) + _block_until_killed(comms) + + mock_parse_task_handler.side_effect = play_runtime(reply) + + result = _run(tmp_path, deadline=time.monotonic() + 1) + + assert not isinstance(result, TaskHandlerProbeStopped) + assert result.import_errors == {"etl.artifact": "handler registry failed"} + + +@pytest.mark.parametrize(("configured", "expected"), [(30, 30), (0.5, 0.5), (0, None), (-1, None)]) +@patch("airflow.settings.get_dagbag_import_timeout", autospec=True) +def test_only_a_positive_import_timeout_applies(mock_timeout, configured, expected): + mock_timeout.return_value = configured + + assert _get_import_timeout("/b/etl.artifact") == expected + mock_timeout.assert_called_once_with("/b/etl.artifact") + + +def _make_process(**kwargs) -> LangSDKTaskHandlerProcessorProcess: + return LangSDKTaskHandlerProcessorProcess( + id=uuid.uuid4(), + pid=1, + stdin=MagicMock(spec=socket.socket), + process=MagicMock(spec=supervisor.ProcessTracker), + process_log=MagicMock(spec=FilteringBoundLogger), + selector=MagicMock(spec=selectors.BaseSelector), + bundle_name="task-handlers", + dag_file_rel_path="etl.artifact", + coordinator="fake", + listeners={}, + parse_request=TaskHandlerParseRequest( + file="/b/etl.artifact", bundle_path=Path("/b"), bundle_name="task-handlers" + ), + **kwargs, + ) + + +def _build_result(*task_ids: str) -> TaskHandlerParsingResult: + return TaskHandlerParsingResult( + fileloc="/b/etl.artifact", + task_handlers={ + "etl": [ + TaskHandlerDeclaration(task_id=task_id, binding="positional", params=[]) + for task_id in task_ids + ] + }, + ) + + +@patch.object(LangSDKTaskHandlerProcessorProcess, "send_msg", autospec=True) +def test_the_first_parse_result_wins(mock_send_msg): + proc = _make_process() + + proc._handle_request(_build_result("first"), structlog.get_logger(), 1) + proc._handle_request(_build_result("second"), structlog.get_logger(), 2) + + assert _get_task_ids(proc.parsing_result) == ["first"] + assert mock_send_msg.call_args.kwargs["error"].detail == { + "message": "A parse result was already received" + } + + +@pytest.mark.parametrize( + "invalid_frame", + [ + pytest.param(bytes.fromhex("00000003c1c1c1"), id="does-not-decode"), + pytest.param( + _RequestFrame( + id=2, + body={ + "type": "TaskHandlerParsingResult", + "fileloc": "/b/etl.artifact", + "task_handlers": "none", + }, + ).as_bytes(), + id="does-not-validate", + ), + ], +) +@patch.object(LangSDKTaskHandlerProcessorProcess, "_kill_runtime", autospec=True) +@patch.object(LangSDKTaskHandlerProcessorProcess, "send_msg", autospec=True) +def test_an_invalid_message_after_the_parse_result_keeps_it(mock_send_msg, mock_kill_runtime, invalid_frame): + proc = _make_process() + runtime, conn = socket.socketpair() + with runtime, conn: + proc._register_comm(conn) + read_frame, _ = proc.selector.register.call_args.args[2] + runtime.sendall(_RequestFrame(id=1, body=_build_result("extract").model_dump(mode="json")).as_bytes()) + assert read_frame(conn) + runtime.sendall(invalid_frame) + assert not read_frame(conn) + + assert _get_task_ids(proc.parsing_result) == ["extract"] + assert proc.parsing_result.import_errors is None + proc.process_log.warning.assert_called_once_with( + "Ignoring an invalid message from the Lang-SDK runtime after its parse result", error=ANY + ) + mock_kill_runtime.assert_called_once_with(proc) + + +@patch.object(LangSDKTaskHandlerProcessorProcess, "send_msg", autospec=True) +def test_the_schema_version_is_reported_once(mock_send_msg): + proc = _make_process() + + proc._handle_request( + LangSDKRuntimeSchemaVersion(schema_version=OLDEST_SCHEMA_VERSION), structlog.get_logger(), 1 + ) + proc._handle_request(LangSDKRuntimeSchemaVersion(schema_version=None), structlog.get_logger(), 2) + + assert proc._runtime_schema_version == OLDEST_SCHEMA_VERSION + assert mock_send_msg.call_args.kwargs["error"].detail["message"] == "Unhandled request" diff --git a/airflow-core/tests/unit/dag_processing/test_task_handler_processor_go.py b/airflow-core/tests/unit/dag_processing/test_task_handler_processor_go.py new file mode 100644 index 0000000000000..8cf2f485fb71f --- /dev/null +++ b/airflow-core/tests/unit/dag_processing/test_task_handler_processor_go.py @@ -0,0 +1,225 @@ +# 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. +"""Pack the Go SDK example bundle, probe it for its task handlers, check the example Dags against it, and read its cache digest.""" + +from __future__ import annotations + +import json +import os +import shutil +import subprocess +from typing import TYPE_CHECKING +from unittest import mock + +import pytest +import structlog + +from airflow.dag_processing.processor import TaskHandlerDeclaration, TaskHandlerParam +from airflow.dag_processing.task_handler_processor import LangSDKTaskHandlerProcessorProcess +from airflow.sdk.coordinators.executable.coordinator import read_cache_digest +from airflow.sdk.execution_time import supervisor +from airflow.sdk.execution_time.coordinator import reset_coordinator_manager + +from tests_common.test_utils.config import conf_vars +from tests_common.test_utils.paths import AIRFLOW_ROOT_PATH +from unit.dag_processing.fake_task_handler_runtime import ( + LOCAL_BUNDLE, + get_stub_task_ids, + parse_dag_file, + sort_bindings, + task_handler_config, +) + +if TYPE_CHECKING: + from pathlib import Path + +pytestmark = pytest.mark.skipif( + os.environ.get("AIRFLOW_LANG_SDK_REAL_PROBE_TESTS") != "1", + reason="set AIRFLOW_LANG_SDK_REAL_PROBE_TESTS=1 to build and probe a real Go bundle", +) + +STRING = {"type": "string"} +INT64 = {"type": "integer", "format": "int64"} +DOUBLE = {"type": "number", "format": "double"} + + +def _make_nullable(schema: dict) -> dict: + return {"anyOf": [schema, {"type": "null"}]} + + +GO_SDK_PATH = AIRFLOW_ROOT_PATH / "go-sdk" + + +def _pack(bundle: Path, *flags: str) -> Path: + completed = subprocess.run( + ["go", "tool", "airflow-go-pack", "--output", os.fspath(bundle), *flags, "./example/bundle"], + cwd=GO_SDK_PATH, + env={**os.environ, "CGO_ENABLED": "0"}, + capture_output=True, + text=True, + check=False, + ) + assert completed.returncode == 0, completed.stderr + return bundle + + +@pytest.fixture(scope="module") +def go_bundle(tmp_path_factory) -> Path: + if shutil.which("go") is None: + pytest.skip("needs a Go toolchain on PATH") + return _pack(tmp_path_factory.mktemp("go-task-handlers") / "example_dags") + + +@pytest.fixture(autouse=True) +def _go_coordinator(monkeypatch, tmp_path): + # Without cgo the Go SDK reads the current user from USER and HOME, and fails without them. + for name, value in (("USER", "airflow"), ("HOME", os.fspath(tmp_path))): + if not os.environ.get(name): + monkeypatch.setenv(name, value) + spec = {"go": {"classpath": "airflow.sdk.coordinators.executable.ExecutableCoordinator"}} + reset_coordinator_manager() + try: + # A bare fork, so the parse child sees this config. + with ( + conf_vars({("sdk", "coordinators"): json.dumps(spec)}), + mock.patch.object(supervisor, "_should_use_exec", autospec=True, return_value=False), + ): + yield + finally: + reset_coordinator_manager() + + +def test_probes_the_task_handlers_of_a_packed_go_bundle(go_bundle): + result = LangSDKTaskHandlerProcessorProcess.run( + coordinator="go", + path=go_bundle, + bundle_path=go_bundle.parent, + bundle_name="go-task-handlers", + artifact_rel_path=go_bundle.name, + logger=structlog.get_logger(), + ) + + assert result.import_errors is None + assert set(result.task_handlers) == { + "simple_dag", + "concurrent_xcom_dag", + "taskflow_binding_dag", + "variable_write_dag", + } + assert [d.task_id for d in result.task_handlers["simple_dag"]] == ["extract", "transform", "load"] + + declarations = {d.task_id: d for d in result.task_handlers["taskflow_binding_dag"]} + assert declarations["via_flat_args"] == TaskHandlerDeclaration( + task_id="via_flat_args", + binding="positional", + params=[ + TaskHandlerParam(name=None, value_schema=schema) + for schema in ( + STRING, + INT64, + DOUBLE, + {"type": "boolean"}, + _make_nullable({"type": "array", "items": STRING}), + {"type": "object"}, + _make_nullable({"type": "array", "items": INT64}), + _make_nullable(STRING), + ) + ], + ) + assert declarations["via_struct_arg_tag"] == TaskHandlerDeclaration( + task_id="via_struct_arg_tag", + binding="named", + params=[ + TaskHandlerParam(name="region_code", exact_name=True, value_schema=STRING), + TaskHandlerParam(name="threshold", exact_name=True, value_schema=DOUBLE), + ], + ) + # The Dag passes this handler an argument it does not take, which only warns under named binding. + assert declarations["via_struct_more_args"] == TaskHandlerDeclaration( + task_id="via_struct_more_args", + binding="named", + params=[TaskHandlerParam(name="region_code", exact_name=True, value_schema=STRING)], + ) + assert declarations["via_flat_map"] == TaskHandlerDeclaration( + task_id="via_flat_map", + binding="named", + params=[ + TaskHandlerParam(name="Region", value_schema=STRING), + TaskHandlerParam(name="Count", value_schema=INT64), + ], + ) + + +def test_the_example_dags_match_the_packed_bundle(go_bundle, cap_structlog): + dag_file = GO_SDK_PATH / "dags" / "go_examples.py" + coordinators = { + "go-sdk": { + "classpath": "airflow.sdk.coordinators.executable.ExecutableCoordinator", + "kwargs": {"task_handler_bundle_name": "go-task-handlers"}, + } + } + bundles = [ + {"name": "dags", "classpath": LOCAL_BUNDLE, "kwargs": {"path": os.fspath(dag_file.parent)}}, + { + "name": "go-task-handlers", + "classpath": LOCAL_BUNDLE, + "kwargs": {"path": os.fspath(go_bundle.parent)}, + }, + ] + + with task_handler_config( + dag_file.parent, + go_bundle.parent, + coordinators, + queue_to_coordinator={"golang": "go-sdk"}, + bundles=bundles, + ): + first = parse_dag_file(dag_file) + second = parse_dag_file(dag_file, known_artifacts=first.probed_artifacts) + + assert first.import_errors == {} + assert {(b.dag_id, b.task_id) for b in first.task_handler_bindings} == get_stub_task_ids(first) + assert [a.relative_fileloc for a in first.probed_artifacts] == [go_bundle.name] + # The Dag passes via_struct_more_args an argument its struct does not declare, and + # via_struct_fewer_args's struct declares one the Dag does not pass: warnings, not errors. + assert { + "event": "Dag's call passed argument(s) the task handler does not declare", + "dag_id": "taskflow_binding_dag", + "task_id": "via_struct_more_args", + "passed_not_declared": ["unused_label"], + } in cap_structlog + assert { + "event": "Task handler declares argument(s) the Dag's call did not pass", + "dag_id": "taskflow_binding_dag", + "task_id": "via_struct_fewer_args", + "declared_not_passed": ["not_in_dag"], + } in cap_structlog + + assert second.import_errors == {} + assert second.probed_artifacts == [] + assert sort_bindings(second) == sort_bindings(first) + + +def test_a_repack_keeps_the_cache_digest_until_a_source_byte_changes(go_bundle, tmp_path): + digest = read_cache_digest(go_bundle) + assert digest is not None + + assert read_cache_digest(_pack(tmp_path / "unchanged")) == digest + + source = tmp_path / "main.go" + source.write_bytes((GO_SDK_PATH / "example" / "bundle" / "main.go").read_bytes() + b"\n") + assert read_cache_digest(_pack(tmp_path / "changed", "--source", os.fspath(source))) != digest diff --git a/airflow-core/tests/unit/dag_processing/test_task_handler_processor_java.py b/airflow-core/tests/unit/dag_processing/test_task_handler_processor_java.py new file mode 100644 index 0000000000000..8aec77d022fcf --- /dev/null +++ b/airflow-core/tests/unit/dag_processing/test_task_handler_processor_java.py @@ -0,0 +1,253 @@ +# +# 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. +"""Probe a bundle built with the Java SDK's Gradle plugin through a real JVM, and check the example Dags against it.""" + +from __future__ import annotations + +import json +import os +import re +import shutil +import subprocess +import zipfile +from typing import TYPE_CHECKING +from unittest import mock + +import pytest +import structlog + +from airflow.dag_processing.processor import ( + TaskHandlerDeclaration, + TaskHandlerParam, + TaskHandlerParsingResult, +) +from airflow.dag_processing.task_handler_processor import LangSDKTaskHandlerProcessorProcess +from airflow.sdk.coordinators.java.coordinator import _parse_manifest +from airflow.sdk.execution_time import supervisor +from airflow.sdk.execution_time.coordinator import reset_coordinator_manager + +from tests_common.pytest_plugin import AIRFLOW_ROOT_PATH +from tests_common.test_utils.config import conf_vars +from unit.dag_processing.fake_task_handler_runtime import ( + LOCAL_BUNDLE, + get_stub_task_ids, + parse_dag_file, + sort_bindings, + task_handler_config, +) + +if TYPE_CHECKING: + from collections.abc import Iterator + from pathlib import Path + +JAVA_SDK_PATH = AIRFLOW_ROOT_PATH / "java-sdk" + +# The example's sources, built as a thin bundle against the SDK as an included build, so nothing is +# published. The plugin's fat mode looks the SDK up as a published airflow-sdk artifact. +_SETTINGS = """ +pluginManagement { + includeBuild("../java-sdk") +} +includeBuild("../java-sdk") +rootProject.name = "probe-example" +""" + +_BUILD = """ +plugins { + id("org.apache.airflow.sdk") +} +repositories { + mavenCentral() +} +dependencies { + annotationProcessor("org.apache.airflow:processor") + implementation("org.apache.airflow:sdk") + implementation("org.apache.airflow:jpl") +} +sourceSets { + main { + java.srcDir("../java-sdk/example/src/java") + } +} +airflowBundle { + mainClass = "org.apache.airflow.example.ExampleBundleBuilder" + fatJar = false +} +""" + +pytestmark = pytest.mark.skipif( + os.environ.get("AIRFLOW_LANG_SDK_REAL_PROBE_TESTS") != "1", + reason="set AIRFLOW_LANG_SDK_REAL_PROBE_TESTS=1 to probe real Lang-SDK artifacts", +) + + +@pytest.fixture(scope="module") +def example_bundle(tmp_path_factory) -> Path: + """Build ``java-sdk/example`` from a copy of the SDK sources, so the checkout gets no build output.""" + if not JAVA_SDK_PATH.is_dir(): + pytest.skip("the Java SDK sources are absent") + if shutil.which("java") is None or shutil.which("javac") is None: + pytest.skip("needs a JDK") + root = tmp_path_factory.mktemp("java-sdk") + sdk = root / "java-sdk" + shutil.copytree(JAVA_SDK_PATH, sdk, ignore=shutil.ignore_patterns("build", ".gradle", ".kotlin")) + project = root / "probe-example" + project.mkdir() + (project / "settings.gradle").write_text(_SETTINGS) + (project / "build.gradle").write_text(_BUILD) + subprocess.run([sdk / "gradlew", "--no-daemon", "--quiet", "bundle"], cwd=project, check=True) + return project / "build" / "bundle" + + +@pytest.fixture +def java_coordinator() -> Iterator[None]: + spec = {"java": {"classpath": "airflow.sdk.coordinators.java.JavaCoordinator"}} + reset_coordinator_manager() + try: + # The parse child is a bare fork even on macOS, so it sees this config. + with ( + conf_vars({("sdk", "coordinators"): json.dumps(spec)}), + mock.patch.object(supervisor, "_should_use_exec", autospec=True, return_value=False), + ): + yield + finally: + reset_coordinator_manager() + + +def _make_nullable(schema: dict) -> dict: + return {"anyOf": [schema, {"type": "null"}]} + + +INT32 = {"type": "integer", "format": "int32"} +INT64 = {"type": "integer", "format": "int64"} +FLOAT = {"type": "number", "format": "float"} +DOUBLE = {"type": "number", "format": "double"} +STRING = {"type": "string"} + + +def _declare_positional(task_id: str, *params: tuple[str, dict]) -> TaskHandlerDeclaration: + return TaskHandlerDeclaration( + task_id=task_id, + binding="positional", + params=[TaskHandlerParam(name=name, value_schema=schema) for name, schema in params], + ) + + +def _declare_named(task_id: str, *params: tuple[str, dict, bool]) -> TaskHandlerDeclaration: + return TaskHandlerDeclaration( + task_id=task_id, + binding="named", + params=[ + TaskHandlerParam(name=name, value_schema=schema, exact_name=exact) + for name, schema, exact in params + ], + ) + + +@pytest.mark.usefixtures("java_coordinator") +def test_a_built_bundle_declares_its_task_handlers(example_bundle): + jar = example_bundle / "probe-example.jar" + + result = LangSDKTaskHandlerProcessorProcess.run( + coordinator="java", + path=jar, + bundle_path=example_bundle, + bundle_name="java-task-handlers", + artifact_rel_path=jar.name, + logger=structlog.get_logger(), + ) + + assert result == TaskHandlerParsingResult( + fileloc=os.fspath(jar), + task_handlers={ + "java_xcom_casting_example": [ + _declare_named("produce_number"), + _declare_positional("widen_to_long", ("value", INT64)), + _declare_positional("widen_to_double", ("value", DOUBLE)), + _declare_named("produce_nothing"), + _declare_positional("consume_nullable", ("value", _make_nullable(INT32))), + _declare_named("produce_fraction"), + _declare_positional("consume_float", ("value", FLOAT)), + _declare_positional( + "consume_double_list", + ("values", _make_nullable({"type": "array", "items": _make_nullable(DOUBLE)})), + ), + ], + "java_interface_example": [ + _declare_named("extract"), + _declare_named("transform", ("extracted", INT64, False)), + _declare_named( + "summarize", ("region_code", _make_nullable(STRING), True), ("transformed", INT64, False) + ), + ], + "java_annotation_example": [ + _declare_named("extract"), + _declare_positional("transform", ("extracted", INT64)), + _declare_positional("load", ("transformed", INT64)), + _declare_named( + "report", ("runLabel", _make_nullable(STRING), False), ("transformed", INT64, False) + ), + _declare_named("concurrent"), + ], + }, + ) + assert list(result.task_handlers) == [ + "java_interface_example", + "java_annotation_example", + "java_xcom_casting_example", + ] + with zipfile.ZipFile(jar) as zf: + manifest = _parse_manifest(zf.read("META-INF/MANIFEST.MF")) + assert re.fullmatch(r"[0-9a-f]{64}", manifest["airflow-cache-digest"]) + + +@pytest.mark.usefixtures("java_coordinator") +def test_the_example_dags_match_the_built_bundle(example_bundle): + dag_file = JAVA_SDK_PATH / "example" / "src" / "resources" / "dags" / "java_examples.py" + coordinators = { + "java-jdk": { + "classpath": "airflow.sdk.coordinators.java.JavaCoordinator", + "kwargs": {"task_handler_bundle_name": "java-task-handlers"}, + } + } + bundles = [ + {"name": "dags", "classpath": LOCAL_BUNDLE, "kwargs": {"path": os.fspath(dag_file.parent)}}, + { + "name": "java-task-handlers", + "classpath": LOCAL_BUNDLE, + "kwargs": {"path": os.fspath(example_bundle)}, + }, + ] + + with task_handler_config( + dag_file.parent, + example_bundle, + coordinators, + queue_to_coordinator={"java": "java-jdk"}, + bundles=bundles, + ): + first = parse_dag_file(dag_file) + second = parse_dag_file(dag_file, known_artifacts=first.probed_artifacts) + + assert first.import_errors == {} + assert {(b.dag_id, b.task_id) for b in first.task_handler_bindings} == get_stub_task_ids(first) + assert [a.relative_fileloc for a in first.probed_artifacts] == ["probe-example.jar"] + + assert second.import_errors == {} + assert second.probed_artifacts == [] + assert sort_bindings(second) == sort_bindings(first) diff --git a/airflow-core/tests/unit/dag_processing/test_task_handler_processor_ts.py b/airflow-core/tests/unit/dag_processing/test_task_handler_processor_ts.py new file mode 100644 index 0000000000000..54211b64992e9 --- /dev/null +++ b/airflow-core/tests/unit/dag_processing/test_task_handler_processor_ts.py @@ -0,0 +1,186 @@ +# +# 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. +"""Probe a bundle packed by the TypeScript SDK through the real Node runtime, and check the example Dags against it.""" + +from __future__ import annotations + +import json +import os +import shutil +import subprocess +from typing import TYPE_CHECKING +from unittest import mock + +import pytest +import structlog + +from airflow.dag_processing.processor import TaskHandlerDeclaration, TaskHandlerParsingResult +from airflow.dag_processing.task_handler_processor import LangSDKTaskHandlerProcessorProcess +from airflow.sdk.execution_time import supervisor +from airflow.sdk.execution_time.coordinator import reset_coordinator_manager + +from tests_common.pytest_plugin import AIRFLOW_ROOT_PATH +from tests_common.test_utils.config import conf_vars +from unit.dag_processing.fake_task_handler_runtime import ( + LOCAL_BUNDLE, + get_stub_task_ids, + parse_dag_file, + sort_bindings, + task_handler_config, +) + +if TYPE_CHECKING: + from collections.abc import Iterator + from pathlib import Path + +TS_SDK_PATH = AIRFLOW_ROOT_PATH / "ts-sdk" +# The SDK's `engines` requirement. +MIN_NODE_MAJOR = 22 +COORDINATORS = { + "ts": { + "classpath": "airflow.sdk.coordinators.node.NodeCoordinator", + "kwargs": {"node_executable": shutil.which("node")}, + } +} + + +def _get_toolchain_problem() -> str | None: + """Return why this machine cannot pack the example bundle, or ``None`` when it can.""" + if not TS_SDK_PATH.is_dir(): + return "the TypeScript SDK sources are absent" + node = shutil.which("node") + if node is None or shutil.which("pnpm") is None: + return "needs node and pnpm" + try: + version = subprocess.run( + [node, "--version"], capture_output=True, text=True, check=True + ).stdout.strip() + major = int(version.lstrip("v").split(".")[0]) + except (OSError, subprocess.CalledProcessError, ValueError) as e: + return f"cannot read the Node.js version: {e}" + if major < MIN_NODE_MAJOR: + return f"needs Node.js {MIN_NODE_MAJOR} or later, found {version}" + return None + + +pytestmark = pytest.mark.skipif( + os.environ.get("AIRFLOW_LANG_SDK_REAL_PROBE_TESTS") != "1", + reason="set AIRFLOW_LANG_SDK_REAL_PROBE_TESTS=1 to probe real Lang-SDK artifacts", +) + + +@pytest.fixture(scope="module") +def example_bundle(tmp_path_factory) -> Path: + """Pack ``ts-sdk/example`` from a copy of the SDK sources, so the checkout gets no build output.""" + if (problem := _get_toolchain_problem()) is not None: + pytest.skip(problem) + sdk = tmp_path_factory.mktemp("ts-sdk") / "ts-sdk" + shutil.copytree(TS_SDK_PATH, sdk, ignore=shutil.ignore_patterns("node_modules", "dist", ".pnpm-store")) + env = {**os.environ, "CI": "true", "COREPACK_ENABLE_DOWNLOAD_PROMPT": "0"} + # The example links the SDK, so it is installed again once the SDK is built. + for cwd, args in [ + (sdk, ["install", "--frozen-lockfile"]), + (sdk, ["run", "build"]), + (sdk / "example", ["install"]), + (sdk / "example", ["run", "build"]), + ]: + subprocess.run(["pnpm", *args], cwd=cwd, env=env, check=True) + return sdk / "example" / "dist" / "bundle.min.mjs" + + +@pytest.fixture +def fresh_coordinator_manager() -> Iterator[None]: + reset_coordinator_manager() + yield + reset_coordinator_manager() + + +def _declare(task_id: str) -> TaskHandlerDeclaration: + return TaskHandlerDeclaration(task_id=task_id, binding="named", params=None) + + +@pytest.mark.usefixtures("fresh_coordinator_manager") +@conf_vars({("sdk", "coordinators"): json.dumps(COORDINATORS)}) +@mock.patch.object(supervisor, "_should_use_exec", autospec=True, return_value=False) +def test_a_packed_bundle_declares_its_task_handlers(mock_should_use_exec, example_bundle): + result = LangSDKTaskHandlerProcessorProcess.run( + coordinator="ts", + path=example_bundle, + bundle_path=example_bundle.parent, + bundle_name="ts-task-handlers", + artifact_rel_path=example_bundle.name, + logger=structlog.get_logger(), + ) + + assert result == TaskHandlerParsingResult( + fileloc=os.fspath(example_bundle), + task_handlers={ + "typescript_example": [ + _declare("build_message"), + _declare("read_connection"), + _declare("write_and_delete_variable"), + ], + "typescript_taskflow_example": [ + _declare("summarize"), + _declare("report"), + _declare("build_message"), + ], + }, + ) + assert list(result.task_handlers) == ["typescript_example", "typescript_taskflow_example"] + + +@pytest.mark.parametrize("dag_file_name", ["typescript_example.py", "typescript_taskflow_example.py"]) +@mock.patch.object(supervisor, "_should_use_exec", autospec=True, return_value=False) +def test_the_example_dags_match_the_packed_bundle(mock_should_use_exec, example_bundle, dag_file_name): + dag_file = TS_SDK_PATH / "example" / "dags" / dag_file_name + coordinators = { + "ts": { + "classpath": "airflow.sdk.coordinators.node.NodeCoordinator", + "kwargs": { + "task_handler_bundle_name": "ts-task-handlers", + "node_executable": shutil.which("node"), + }, + } + } + bundles = [ + {"name": "dags", "classpath": LOCAL_BUNDLE, "kwargs": {"path": os.fspath(dag_file.parent)}}, + { + "name": "ts-task-handlers", + "classpath": LOCAL_BUNDLE, + "kwargs": {"path": os.fspath(example_bundle.parent)}, + }, + ] + + with task_handler_config( + dag_file.parent, + example_bundle.parent, + coordinators, + queue_to_coordinator={"typescript": "ts"}, + bundles=bundles, + ): + first = parse_dag_file(dag_file) + second = parse_dag_file(dag_file, known_artifacts=first.probed_artifacts) + + assert first.import_errors == {} + assert {(b.dag_id, b.task_id) for b in first.task_handler_bindings} == get_stub_task_ids(first) + assert [a.relative_fileloc for a in first.probed_artifacts] == [example_bundle.name] + + assert second.import_errors == {} + assert second.probed_artifacts == [] + assert sort_bindings(second) == sort_bindings(first) diff --git a/airflow-core/tests/unit/dag_processing/test_task_handler_resolution.py b/airflow-core/tests/unit/dag_processing/test_task_handler_resolution.py new file mode 100644 index 0000000000000..d354bd9383d9e --- /dev/null +++ b/airflow-core/tests/unit/dag_processing/test_task_handler_resolution.py @@ -0,0 +1,734 @@ +# Licensed to the Apache Software Foundation (ASF) under one +# or more contributor license agreements. See the NOTICE file +# distributed with this work for additional information +# regarding copyright ownership. The ASF licenses this file +# to you under the Apache License, Version 2.0 (the +# "License"); you may not use this file except in compliance +# with the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, +# software distributed under the License is distributed on an +# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +# KIND, either express or implied. See the License for the +# specific language governing permissions and limitations +# under the License. +from __future__ import annotations + +import contextlib +import json +import os +import time +from typing import TYPE_CHECKING +from unittest.mock import patch + +import pytest +import structlog + +from airflow.api_fastapi.execution_api.datamodels.task_arg_binding import LiteralArgBinding +from airflow.dag_processing.bundles.manager import DagBundlesManager +from airflow.dag_processing.processor import ( + TaskHandlerArtifact, + TaskHandlerBinding, + TaskHandlerParseRequest, + TaskHandlerParsingResult, +) +from airflow.dag_processing.task_handler_processor import ( + LangSDKTaskHandlerProcessorProcess, + TaskHandlerProbeStopped, +) +from airflow.dag_processing.task_handler_resolution import TaskHandlerResolution, resolve_task_handlers +from airflow.dag_processing.task_handler_validation import StubTask +from airflow.sdk.coordinators._subprocess import SubprocessCoordinator + +from unit.dag_processing.fake_task_handler_runtime import ( + FAKE_COORDINATOR, + LOCAL_BUNDLE, + FakeCoordinator, + reply_with_task_handlers, + task_handler_config, + write_artifact, +) + +if TYPE_CHECKING: + from pathlib import Path + +PROBED = "Stub tasks in etl.py do not match their task handlers:" + + +def _declare(*task_ids: str) -> dict: + return {"etl": [{"task_id": task_id, "binding": "positional", "params": []} for task_id in task_ids]} + + +def _stub(task_id: str, *, queue: str = "fake-queue", relative_fileloc: str = "etl.py", arg_bindings=()): + return StubTask( + dag_id="etl", + task_id=task_id, + queue=queue, + relative_fileloc=relative_fileloc, + arg_bindings=list(arg_bindings), + is_mapped=False, + ) + + +def _binding(task_id: str, *, rel_path: str = "etl.artifact", bundle_name: str = "task-handlers"): + return TaskHandlerBinding( + dag_id="etl", task_id=task_id, artifact_bundle_name=bundle_name, artifact_rel_path=rel_path + ) + + +def _answer(path: Path, *, bundle_name: str = "task-handlers") -> TaskHandlerArtifact: + content = path.read_bytes() + spec = json.loads(content) + return TaskHandlerArtifact.model_validate( + { + "bundle_name": bundle_name, + "relative_fileloc": path.name, + "size_bytes": len(content), + "cache_digest": spec.get("cache_digest"), + "task_handlers": spec["task_handlers"], + } + ) + + +def _probe(*, path, bundle_path, bundle_name, **kwargs) -> TaskHandlerParsingResult: + """Answer as the runtime does, from the ``task_handlers`` of the artifact's JSON.""" + request = TaskHandlerParseRequest(file=os.fspath(path), bundle_path=bundle_path, bundle_name=bundle_name) + return reply_with_task_handlers(request, None) + + +def _fail_probe( + *, + path, + artifact_rel_path, + error="The Lang-SDK runtime exited with code 1 without a parse result", + result_type=TaskHandlerParsingResult, + **kwargs, +) -> TaskHandlerParsingResult: + return result_type(fileloc=os.fspath(path), task_handlers={}, import_errors={artifact_rel_path: error}) + + +def _stop_probe(*, path, artifact_rel_path, **kwargs) -> TaskHandlerParsingResult: + """Answer as a probe the deadline stopped before its runtime answered.""" + return _fail_probe( + path=path, + artifact_rel_path=artifact_rel_path, + error=f"The Lang-SDK runtime did not parse {path} by its deadline", + result_type=TaskHandlerProbeStopped, + ) + + +@pytest.fixture +def dag_bundle(tmp_path) -> Path: + path = tmp_path / "dags" + path.mkdir() + return path + + +@pytest.fixture +def artifacts(tmp_path) -> Path: + path = tmp_path / "artifacts" + path.mkdir() + return path + + +@pytest.fixture +def configure(dag_bundle, artifacts): + """Configure coordinators, their queues and the Dag bundles, as ``task_handler_config`` does.""" + with contextlib.ExitStack() as stack: + yield lambda *args, **kwargs: stack.enter_context( + task_handler_config(dag_bundle, artifacts, *args, **kwargs) + ) + + +def _resolve(stub_tasks, dag_bundle, *, known_artifacts=(), deadline=None) -> TaskHandlerResolution: + return resolve_task_handlers( + stub_tasks, + dag_bundle_name="dags", + dag_bundle_path=dag_bundle, + known_artifacts=list(known_artifacts), + deadline=time.monotonic() + 60 if deadline is None else deadline, + log=structlog.get_logger(), + ) + + +def _configure_two_coordinators(configure) -> None: + """Route ``fake-queue`` to ``fake`` and ``other-queue`` to ``other``, both reading ``task-handlers``.""" + kwargs = {"task_handler_bundle_name": "task-handlers"} + configure( + { + "fake": {"classpath": FAKE_COORDINATOR, "kwargs": kwargs}, + "other": {"classpath": FAKE_COORDINATOR, "kwargs": kwargs}, + }, + queue_to_coordinator={"fake-queue": "fake", "other-queue": "other"}, + ) + + +@patch.object(LangSDKTaskHandlerProcessorProcess, "run", autospec=True, side_effect=_probe) +class TestResolveTaskHandlers: + def test_bindings_name_every_routed_stub_task( + self, mock_run, configure, dag_bundle, artifacts, cap_structlog + ): + configure() + artifact = write_artifact( + artifacts / "etl.artifact", cache_digest="d1", task_handlers=_declare("extract", "load") + ) + deadline = time.monotonic() + 60 + + resolution = _resolve( + [_stub("extract"), _stub("load"), _stub("report", queue="default")], dag_bundle, deadline=deadline + ) + + assert resolution == TaskHandlerResolution( + bindings=[_binding("extract"), _binding("load")], + probed_artifacts=[_answer(artifact)], + import_errors={}, + ) + mock_run.assert_called_once_with( + coordinator="fake", + path=artifact, + bundle_path=artifacts, + bundle_name="task-handlers", + artifact_rel_path="etl.artifact", + logger=mock_run.call_args.kwargs["logger"], + deadline=deadline, + ) + assert {"event": "Probing a task handler artifact", "path": "etl.artifact"} in cap_structlog + assert {"event": "Probed a task handler artifact", "path": "etl.artifact"} in cap_structlog + + def test_a_recorded_answer_is_used_without_probing(self, mock_run, configure, dag_bundle, artifacts): + configure() + artifact = write_artifact( + artifacts / "etl.artifact", cache_digest="d1", task_handlers=_declare("extract") + ) + + resolution = _resolve([_stub("extract")], dag_bundle, known_artifacts=[_answer(artifact)]) + + assert resolution == TaskHandlerResolution( + bindings=[_binding("extract")], probed_artifacts=[], import_errors={} + ) + mock_run.assert_not_called() + + def test_an_artifact_two_coordinators_list_is_probed_once( + self, mock_run, configure, dag_bundle, artifacts + ): + _configure_two_coordinators(configure) + artifact = write_artifact(artifacts / "etl.artifact", task_handlers=_declare("extract", "load")) + + resolution = _resolve([_stub("extract"), _stub("load", queue="other-queue")], dag_bundle) + + assert resolution == TaskHandlerResolution( + bindings=[_binding("extract"), _binding("load")], + probed_artifacts=[_answer(artifact)], + import_errors={}, + ) + assert mock_run.call_count == 1 + assert mock_run.call_args.kwargs["coordinator"] == "fake" + + def test_a_failed_probe_is_retried_under_the_next_coordinator_that_lists_it( + self, mock_run, configure, dag_bundle, artifacts + ): + _configure_two_coordinators(configure) + artifact = write_artifact(artifacts / "etl.artifact", task_handlers=_declare("extract", "load")) + mock_run.side_effect = lambda **kwargs: ( + _fail_probe(**kwargs) if kwargs["coordinator"] == "fake" else _probe(**kwargs) + ) + + resolution = _resolve([_stub("extract"), _stub("load", queue="other-queue")], dag_bundle) + + assert resolution == TaskHandlerResolution( + bindings=[_binding("extract"), _binding("load")], + probed_artifacts=[_answer(artifact)], + import_errors={}, + ) + assert [c.kwargs["coordinator"] for c in mock_run.call_args_list] == ["fake", "other"] + + def test_a_failed_probe_is_named_for_the_coordinator_it_failed_under( + self, mock_run, configure, dag_bundle, artifacts + ): + _configure_two_coordinators(configure) + write_artifact(artifacts / "etl.artifact", task_handlers=_declare("extract", "load")) + mock_run.side_effect = lambda **kwargs: _fail_probe( + error=f"{kwargs['coordinator']} cannot run it", **kwargs + ) + + resolution = _resolve( + [_stub("extract"), _stub("load", queue="other-queue", relative_fileloc="other.py")], dag_bundle + ) + + no_handler = "no artifact in Dag bundle 'task-handlers' registers it; no answer from 'etl.artifact'" + assert resolution == TaskHandlerResolution( + bindings=None, + probed_artifacts=[], + import_errors={ + "etl.py": f"{PROBED}\n- Dag 'etl', task 'extract': {no_handler} (probe failed: fake cannot run it)", + "other.py": "Stub tasks in other.py do not match their task handlers:\n" + f"- Dag 'etl', task 'load': {no_handler} (probe failed: other cannot run it)", + }, + ) + assert [c.kwargs["coordinator"] for c in mock_run.call_args_list] == ["fake", "other"] + + def test_probed_answers_are_returned_when_validation_fails( + self, mock_run, configure, dag_bundle, artifacts + ): + configure() + artifact = write_artifact(artifacts / "etl.artifact", task_handlers=_declare("extract")) + + resolution = _resolve([_stub("extract"), _stub("load")], dag_bundle) + + assert resolution == TaskHandlerResolution( + bindings=None, + probed_artifacts=[_answer(artifact)], + import_errors={ + "etl.py": f"{PROBED}\n- Dag 'etl', task 'load': no artifact in Dag bundle 'task-handlers' registers it" + }, + ) + + @pytest.mark.parametrize( + ("task_id", "import_errors"), + [ + pytest.param("extract", {}, id="handler-found"), + pytest.param( + "load", + { + "etl.py": f"{PROBED}\n- Dag 'etl', task 'load': no artifact in Dag bundle 'task-handlers' " + "registers it; no answer from 'bad.artifact' (bad.artifact is not executable)" + }, + id="handler-missing", + ), + ], + ) + def test_a_rejected_candidate_is_ignored_and_named_when_a_handler_is_missing( + self, mock_run, configure, dag_bundle, artifacts, cap_structlog, task_id, import_errors + ): + configure() + write_artifact(artifacts / "bad.artifact", listing_error="bad.artifact is not executable") + good = write_artifact(artifacts / "good.artifact", task_handlers=_declare("extract")) + + resolution = _resolve([_stub(task_id)], dag_bundle) + + assert resolution.import_errors == import_errors + assert resolution.probed_artifacts == [_answer(good)] + assert [c.kwargs["artifact_rel_path"] for c in mock_run.call_args_list] == ["good.artifact"] + assert { + "event": "Ignoring a task handler artifact its coordinator cannot probe", + "path": "bad.artifact", + "error": "bad.artifact is not executable", + } in cap_structlog + + @pytest.mark.parametrize( + ("task_id", "import_errors"), + [ + pytest.param("extract", {}, id="handler-found"), + pytest.param( + "load", + { + "etl.py": f"{PROBED}\n- Dag 'etl', task 'load': no artifact in Dag bundle 'task-handlers' " + "registers it; no answer from 'bad.artifact' " + "(probe failed: The Lang-SDK runtime exited with code 1 without a parse result)" + }, + id="handler-missing", + ), + ], + ) + def test_a_failed_probe_is_ignored_and_named_when_a_handler_is_missing( + self, mock_run, configure, dag_bundle, artifacts, task_id, import_errors + ): + configure() + write_artifact(artifacts / "bad.artifact") + good = write_artifact(artifacts / "good.artifact", task_handlers=_declare("extract")) + mock_run.side_effect = lambda **kwargs: ( + _fail_probe(**kwargs) if kwargs["artifact_rel_path"] == "bad.artifact" else _probe(**kwargs) + ) + + resolution = _resolve([_stub(task_id)], dag_bundle) + + assert resolution.import_errors == import_errors + assert resolution.probed_artifacts == [_answer(good)] + assert resolution.bindings == ( + None if import_errors else [_binding("extract", rel_path="good.artifact")] + ) + + @pytest.mark.parametrize( + ("task_id", "import_errors"), + [ + pytest.param("extract", {}, id="handler-found"), + pytest.param( + "load", + { + "etl.py": f"{PROBED}\n- Dag 'etl', task 'load': no artifact in Dag bundle 'task-handlers' " + "registers it; no answer from 'bad.artifact' " + "(probe failed: OSError: Cannot allocate memory)" + }, + id="handler-missing", + ), + ], + ) + def test_a_probe_that_raises_is_ignored_and_named_when_a_handler_is_missing( + self, mock_run, configure, dag_bundle, artifacts, cap_structlog, task_id, import_errors + ): + configure() + write_artifact(artifacts / "bad.artifact") + good = write_artifact(artifacts / "good.artifact", task_handlers=_declare("extract")) + + def probe(**kwargs): + if kwargs["artifact_rel_path"] == "bad.artifact": + raise OSError("Cannot allocate memory") + return _probe(**kwargs) + + mock_run.side_effect = probe + + resolution = _resolve([_stub(task_id)], dag_bundle) + + assert resolution.import_errors == import_errors + assert resolution.probed_artifacts == [_answer(good)] + [entry] = [e for e in cap_structlog.entries if e["event"] == "Probing a task handler artifact raised"] + assert entry["path"] == "bad.artifact" + assert entry["exception"][0]["exc_type"] == "OSError" + + def test_a_candidate_past_the_deadline_is_not_probed(self, mock_run, configure, dag_bundle, artifacts): + configure() + write_artifact(artifacts / "a.artifact", task_handlers=_declare("extract")) + write_artifact(artifacts / "b.artifact", task_handlers=_declare("extract")) + deadline = time.monotonic() + 0.5 + + def stop_at_the_deadline(**kwargs): + # As a probe still running at its deadline is stopped. + time.sleep(max(deadline - time.monotonic(), 0) + 0.05) + return _stop_probe(**kwargs) + + mock_run.side_effect = stop_at_the_deadline + + resolution = _resolve([_stub("load")], dag_bundle, deadline=deadline) + + assert resolution == TaskHandlerResolution( + bindings=None, + probed_artifacts=[], + import_errors={ + "etl.py": f"{PROBED}\n- Dag 'etl', task 'load': no artifact in Dag bundle 'task-handlers' " + "registers it; no answer from " + "'a.artifact' (probe stopped: the parse ran out of [dag_processor] dag_file_processor_timeout), " + "'b.artifact' (not probed: the parse ran out of [dag_processor] dag_file_processor_timeout)" + }, + ) + assert [c.kwargs["artifact_rel_path"] for c in mock_run.call_args_list] == ["a.artifact"] + + def test_a_later_coordinators_skip_does_not_hide_an_earlier_coordinators_failure( + self, mock_run, configure, dag_bundle, artifacts + ): + _configure_two_coordinators(configure) + write_artifact(artifacts / "x.artifact") + write_artifact(artifacts / "y.artifact") + deadline = time.monotonic() + 0.5 + + def fail_x_and_run_y_to_the_deadline(**kwargs): + if kwargs["artifact_rel_path"] == "x.artifact": + return _fail_probe(**kwargs) + time.sleep(max(deadline - time.monotonic(), 0) + 0.05) + return _stop_probe(**kwargs) + + mock_run.side_effect = fail_x_and_run_y_to_the_deadline + + resolution = _resolve( + [_stub("extract"), _stub("extract", queue="other-queue", relative_fileloc="other.py")], + dag_bundle, + deadline=deadline, + ) + + no_handler = "no artifact in Dag bundle 'task-handlers' registers it; no answer from" + out_of_time = "the parse ran out of [dag_processor] dag_file_processor_timeout" + assert resolution == TaskHandlerResolution( + bindings=None, + probed_artifacts=[], + import_errors={ + "etl.py": f"{PROBED}\n- Dag 'etl', task 'extract': {no_handler} " + "'x.artifact' (probe failed: The Lang-SDK runtime exited with code 1 without a parse result), " + f"'y.artifact' (probe stopped: {out_of_time})", + "other.py": "Stub tasks in other.py do not match their task handlers:\n" + f"- Dag 'etl', task 'extract': {no_handler} " + f"'x.artifact' (not probed: {out_of_time}), 'y.artifact' (not probed: {out_of_time})", + }, + ) + probed = [(c.kwargs["coordinator"], c.kwargs["artifact_rel_path"]) for c in mock_run.call_args_list] + assert probed == [("fake", "x.artifact"), ("fake", "y.artifact")] + + def test_an_error_the_runtime_reported_is_not_blamed_on_the_deadline( + self, mock_run, configure, dag_bundle, artifacts + ): + configure() + write_artifact(artifacts / "etl.artifact") + deadline = time.monotonic() + 0.5 + + def answer_with_an_error_and_run_past_the_deadline(**kwargs): + time.sleep(max(deadline - time.monotonic(), 0) + 0.05) + return _fail_probe(error="handler registry failed", **kwargs) + + mock_run.side_effect = answer_with_an_error_and_run_past_the_deadline + + resolution = _resolve([_stub("extract")], dag_bundle, deadline=deadline) + + assert resolution.import_errors == { + "etl.py": f"{PROBED}\n- Dag 'etl', task 'extract': no artifact in Dag bundle 'task-handlers' " + "registers it; no answer from 'etl.artifact' (probe failed: handler registry failed)" + } + + def test_a_bundle_of_another_team_is_an_import_error(self, mock_run, configure, dag_bundle, artifacts): + configure( + bundles=[ + { + "name": "dags", + "classpath": LOCAL_BUNDLE, + "kwargs": {"path": os.fspath(dag_bundle)}, + "team_name": "a", + }, + { + "name": "task-handlers", + "classpath": LOCAL_BUNDLE, + "kwargs": {"path": os.fspath(artifacts)}, + "team_name": "b", + }, + ], + multi_team=True, + ) + write_artifact(artifacts / "etl.artifact", task_handlers=_declare("extract")) + + resolution = _resolve([_stub("extract")], dag_bundle) + + assert resolution == TaskHandlerResolution( + bindings=None, + probed_artifacts=[], + import_errors={ + "etl.py": f"{PROBED}\n- Coordinator 'fake': Dag bundle 'task-handlers' belongs to team 'b', " + "but Dag bundle 'dags' belongs to team 'a'" + }, + ) + mock_run.assert_not_called() + + def test_a_coordinator_without_a_bundle_name_reads_the_dag_bundle( + self, mock_run, configure, dag_bundle, artifacts + ): + # The Dag bundle is not looked up in the configuration: the parse request carries its path. + configure( + {"fake": {"classpath": FAKE_COORDINATOR}}, + bundles=[ + {"name": "elsewhere", "classpath": LOCAL_BUNDLE, "kwargs": {"path": os.fspath(artifacts)}} + ], + ) + (dag_bundle / "etl.py").write_text("") + artifact = write_artifact(dag_bundle / "etl.artifact", task_handlers=_declare("extract")) + + resolution = _resolve([_stub("extract")], dag_bundle) + + assert resolution == TaskHandlerResolution( + bindings=[_binding("extract", bundle_name="dags")], + probed_artifacts=[_answer(artifact, bundle_name="dags")], + import_errors={}, + ) + assert mock_run.call_args.kwargs["bundle_path"] == dag_bundle + + @patch.object(DagBundlesManager, "get_bundle", autospec=True, side_effect=DagBundlesManager.get_bundle) + @patch.object(DagBundlesManager, "__init__", autospec=True, side_effect=DagBundlesManager.__init__) + def test_a_bundle_two_coordinators_name_is_looked_up_once( + self, mock_init, mock_get_bundle, mock_run, configure, dag_bundle, artifacts + ): + _configure_two_coordinators(configure) + write_artifact(artifacts / "etl.artifact", task_handlers=_declare("extract", "load")) + + resolution = _resolve([_stub("extract"), _stub("load", queue="other-queue")], dag_bundle) + + assert resolution.import_errors == {} + assert mock_init.call_count == 1 + assert [c.args[1] for c in mock_get_bundle.call_args_list] == ["task-handlers"] + + def test_a_missing_bundle_path_is_an_import_error(self, mock_run, configure, dag_bundle, artifacts): + configure() + artifacts.rmdir() + + resolution = _resolve([_stub("extract")], dag_bundle) + + assert resolution.import_errors == { + "etl.py": f"{PROBED}\n- Coordinator 'fake': Dag bundle 'task-handlers' resolved to {artifacts}, " + "which does not exist on this Dag processor" + } + assert resolution.bindings is None + + @pytest.mark.skipif(os.geteuid() == 0, reason="root reads every directory") + @pytest.mark.parametrize("locked", ["parent", "root"]) + def test_an_unreadable_bundle_path_is_an_import_error( + self, mock_run, configure, dag_bundle, tmp_path, locked + ): + parent = tmp_path / "locked" + root = parent / "artifacts" + root.mkdir(parents=True) + configure( + bundles=[ + {"name": "dags", "classpath": LOCAL_BUNDLE, "kwargs": {"path": os.fspath(dag_bundle)}}, + {"name": "task-handlers", "classpath": LOCAL_BUNDLE, "kwargs": {"path": os.fspath(root)}}, + ] + ) + write_artifact(root / "etl.artifact", task_handlers=_declare("extract")) + (parent if locked == "parent" else root).chmod(0) + try: + resolution = _resolve([_stub("extract")], dag_bundle) + finally: + parent.chmod(0o700) + root.chmod(0o700) + + assert resolution.bindings is None + assert resolution.import_errors == { + "etl.py": f"{PROBED}\n- Coordinator 'fake': Dag bundle 'task-handlers' at {root} cannot be read " + f"on this Dag processor: PermissionError: [Errno 13] Permission denied: '{root}'" + } + mock_run.assert_not_called() + + @pytest.mark.parametrize( + ("listing", "error"), + [ + pytest.param( + SubprocessCoordinator.list_task_handler_candidates, + f"{FAKE_COORDINATOR} cannot list task handler artifacts, so its stub tasks cannot be bound", + id="not-implemented", + ), + pytest.param( + OSError("permission denied"), + "cannot list the task handler artifacts of Dag bundle 'task-handlers': " + "OSError: permission denied", + id="fails", + ), + ], + ) + def test_a_coordinator_that_cannot_list_is_an_import_error( + self, mock_run, configure, dag_bundle, listing, error + ): + configure() + if isinstance(listing, Exception): + patcher = patch.object( + FakeCoordinator, "list_task_handler_candidates", autospec=True, side_effect=listing + ) + else: + patcher = patch.object( + FakeCoordinator, + "_read_task_handler_candidate", + SubprocessCoordinator._read_task_handler_candidate, + ) + + with patcher: + resolution = _resolve( + [ + _stub("extract"), + _stub("load", relative_fileloc="other.py"), + _stub("report", queue="default"), + ], + dag_bundle, + ) + + assert resolution.import_errors == { + "etl.py": f"{PROBED}\n- Coordinator 'fake': {error}", + "other.py": f"Stub tasks in other.py do not match their task handlers:\n- Coordinator 'fake': {error}", + } + assert resolution.bindings is None + + def test_a_coordinator_that_cannot_be_built_is_an_import_error(self, mock_run, configure, dag_bundle): + configure( + { + "fake": { + "classpath": "no_such_module.Coordinator", + "kwargs": {"task_handler_bundle_name": "task-handlers"}, + } + } + ) + + resolution = _resolve([_stub("extract")], dag_bundle) + + assert resolution.import_errors == { + "etl.py": f"{PROBED}\n- Coordinator 'fake': cannot be built: InvalidCoordinatorError: " + "Cannot import coordinator 'fake' (ModuleNotFoundError: No module named 'no_such_module')" + } + + @patch( + "airflow.dag_processing.task_handler_resolution.plan_task_handler_probes", + autospec=True, + side_effect=ValueError("unexpected listing"), + ) + def test_an_unexpected_error_in_a_coordinator_is_an_import_error( + self, mock_plan, mock_run, configure, dag_bundle, artifacts, cap_structlog + ): + configure() + write_artifact(artifacts / "etl.artifact", task_handlers=_declare("extract")) + + resolution = _resolve([_stub("extract"), _stub("load", relative_fileloc="other.py")], dag_bundle) + + error = "Coordinator 'fake': cannot find its task handler artifacts: ValueError: unexpected listing" + assert resolution == TaskHandlerResolution( + bindings=None, + probed_artifacts=[], + import_errors={ + "etl.py": f"{PROBED}\n- {error}", + "other.py": f"Stub tasks in other.py do not match their task handlers:\n- {error}", + }, + ) + [entry] = [ + e + for e in cap_structlog.entries + if e["event"] == "Cannot list the task handler artifacts of a coordinator" + ] + assert entry["coordinator"] == "fake" + assert entry["exception"][0]["exc_type"] == "ValueError" + + def test_an_invalid_sdk_config_is_an_import_error(self, mock_run, configure, dag_bundle): + configure(queue_to_coordinator={"fake-queue": "missing"}) + + resolution = _resolve([_stub("extract"), _stub("load", relative_fileloc="other.py")], dag_bundle) + + error = ( + "Cannot load [sdk] coordinators: " + "ValueError: [sdk] queue_to_coordinator references invalid coordinator key: 'missing'" + ) + assert resolution == TaskHandlerResolution( + bindings=None, + probed_artifacts=[], + import_errors={ + "etl.py": f"{PROBED}\n- {error}", + "other.py": f"Stub tasks in other.py do not match their task handlers:\n- {error}", + }, + ) + + def test_a_name_mismatch_is_logged_in_the_parse_log( + self, mock_run, configure, dag_bundle, artifacts, cap_structlog + ): + configure() + params = [ + {"name": "region_code", "exact_name": True, "value_schema": {"type": "string"}}, + {"name": "not_in_dag", "exact_name": True}, + ] + write_artifact( + artifacts / "etl.artifact", + task_handlers={"etl": [{"task_id": "extract", "binding": "named", "params": params}]}, + ) + arg_bindings = [ + LiteralArgBinding(kind="literal", name=name, value="eu") + for name in ("region_code", "unused_label") + ] + + resolution = _resolve([_stub("extract", arg_bindings=arg_bindings)], dag_bundle) + + assert resolution.bindings == [_binding("extract")] + assert resolution.import_errors == {} + context = { + "dag_id": "etl", + "task_id": "extract", + "artifact_bundle_name": "task-handlers", + "artifact_rel_path": "etl.artifact", + "log_level": "warning", + } + assert { + "event": "Dag's call passed argument(s) the task handler does not declare", + "passed_not_declared": ["unused_label"], + **context, + } in cap_structlog + assert { + "event": "Task handler declares argument(s) the Dag's call did not pass", + "declared_not_passed": ["not_in_dag"], + **context, + } in cap_structlog diff --git a/airflow-core/tests/unit/dag_processing/test_task_handler_validation.py b/airflow-core/tests/unit/dag_processing/test_task_handler_validation.py new file mode 100644 index 0000000000000..4baa9f2dfb69e --- /dev/null +++ b/airflow-core/tests/unit/dag_processing/test_task_handler_validation.py @@ -0,0 +1,650 @@ +# Licensed to the Apache Software Foundation (ASF) under one +# or more contributor license agreements. See the NOTICE file +# distributed with this work for additional information +# regarding copyright ownership. The ASF licenses this file +# to you under the Apache License, Version 2.0 (the +# "License"); you may not use this file except in compliance +# with the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, +# software distributed under the License is distributed on an +# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +# KIND, either express or implied. See the License for the +# specific language governing permissions and limitations +# under the License. +from __future__ import annotations + +import pytest + +from airflow.api_fastapi.execution_api.datamodels.task_arg_binding import LiteralArgBinding, XComArgBinding +from airflow.dag_processing.processor import ( + TaskHandlerArtifact, + TaskHandlerBinding, + TaskHandlerDeclaration, + TaskHandlerParam, +) +from airflow.dag_processing.task_handler_validation import ( + ArgumentCheck, + StubTask, + StubTaskWarning, + TaskHandlerMatch, + TaskHandlerProblem, + check_task_handler_arguments, + check_value_schema, + collect_stub_tasks, + format_import_errors, + match_task_handlers, +) +from airflow.sdk import DAG, task +from airflow.serialization.serialized_objects import LazyDeserializedDAG + +BUNDLE_NAME = "go-task-handlers" +DAG_FILE = "dags/etl.py" + +INTEGER = {"type": "integer"} +NUMBER = {"type": "number"} +STRING = {"type": "string"} +NULL = {"type": "null"} +NULLABLE_INTEGER = {"anyOf": [INTEGER, NULL]} +OBJECT = {"type": "object", "additionalProperties": True} + + +def _literal(name, value=None, *, schema=None, from_default=False): + return LiteralArgBinding( + kind="literal", name=name, value=value, value_schema=schema, from_default=from_default + ) + + +def _default(name, value=None, *, schema=None): + return _literal(name, value, schema=schema, from_default=True) + + +def _xcom(name, *, schema=None): + return XComArgBinding(kind="xcom", name=name, task_id="extract", value_schema=schema) + + +def _param(name=None, *, schema=None, exact_name=False): + return TaskHandlerParam(name=name, value_schema=schema, exact_name=exact_name) + + +def _positional(*params, task_id="transform"): + return TaskHandlerDeclaration(task_id=task_id, binding="positional", params=list(params)) + + +def _named(*params, task_id="transform"): + return TaskHandlerDeclaration(task_id=task_id, binding="named", params=list(params)) + + +def _check(*, errors=(), passed_not_declared=(), declared_not_passed=()): + return ArgumentCheck( + errors=list(errors), + passed_not_declared=list(passed_not_declared), + declared_not_passed=list(declared_not_passed), + ) + + +def _stub(task_id="transform", *, arg_bindings=(), is_mapped=False, dag_id="etl"): + return StubTask( + dag_id=dag_id, + task_id=task_id, + queue="go", + relative_fileloc=DAG_FILE, + arg_bindings=list(arg_bindings), + is_mapped=is_mapped, + ) + + +def _artifact(rel_path="bin/etl", **task_handlers): + return TaskHandlerArtifact( + bundle_name=BUNDLE_NAME, + relative_fileloc=rel_path, + size_bytes=100, + cache_digest="a" * 64, + task_handlers=task_handlers, + ) + + +def _binding(task_id="transform", *, rel_path="bin/etl", dag_id="etl"): + return TaskHandlerBinding( + dag_id=dag_id, task_id=task_id, artifact_bundle_name=BUNDLE_NAME, artifact_rel_path=rel_path + ) + + +def _match(stub_tasks, answers, broken_candidates=None): + return match_task_handlers( + stub_tasks, answers, bundle_name=BUNDLE_NAME, broken_candidates=broken_candidates or {} + ) + + +@pytest.mark.parametrize( + ("arg_bindings", "declaration", "expected"), + [ + pytest.param( + [_literal("a"), _xcom("b")], + _positional(_param(), _param()), + _check(), + id="positional-count-matches", + ), + pytest.param( + [_literal("a"), _literal("b"), _literal("c")], + _positional(_param(), _param()), + _check(errors=["passes 3 arguments, the task handler takes 2"]), + id="positional-too-many", + ), + pytest.param( + [_literal("a")], + _positional(_param(), _param()), + _check(errors=["passes 1 argument, the task handler takes 2"]), + id="positional-too-few", + ), + pytest.param( + [_literal("a"), _default("b")], _positional(_param()), _check(), id="defaulted-dropped-to-match" + ), + pytest.param( + [_literal("a"), _default("b")], + _positional(_param(), _param()), + _check(), + id="defaulted-kept-when-the-full-count-matches", + ), + pytest.param( + [_literal("a"), _default("b"), _literal("c")], + _positional(_param()), + _check(errors=["passes 3 arguments (2 without defaults), the task handler takes 1"]), + id="dropping-defaulted-still-mismatches", + ), + pytest.param( + [_literal("x", schema=INTEGER), _literal("y", schema=STRING)], + _positional(_param("y", schema=INTEGER), _param("x", schema=STRING)), + _check(), + id="positional-names-ignored", + ), + pytest.param( + [_literal("a", schema=INTEGER), _xcom("b", schema=STRING)], + _positional(_param(schema=INTEGER), _param(schema=INTEGER)), + _check(errors=["argument 'b' is string, the task handler takes integer"]), + id="positional-type-mismatch", + ), + pytest.param([], _positional(), _check(), id="argless-against-no-params"), + pytest.param( + [], + _positional(_param()), + _check(errors=["passes 0 arguments, the task handler takes 1"]), + id="argless-against-positional-params", + ), + pytest.param( + [_literal("user_id", schema=INTEGER)], + _named(_param("user_id", schema=INTEGER)), + _check(), + id="named-exact", + ), + pytest.param( + [_literal("user_id"), _literal("region")], + _named(_param("UserId"), _param("region")), + _check(), + id="named-folded", + ), + pytest.param( + [_literal("user_id"), _literal("region")], + _named(_param("UserId", exact_name=True), _param("region")), + _check(passed_not_declared=["user_id"], declared_not_passed=["UserId"]), + id="exact-name-blocks-folding", + ), + pytest.param( + [_literal("region"), _literal("extra")], + _named(_param("region")), + _check(passed_not_declared=["extra"]), + id="passed-argument-no-param-takes", + ), + pytest.param( + [_literal("region"), _default("extra")], + _named(_param("region")), + _check(), + id="defaulted-argument-no-param-takes", + ), + pytest.param( + [_literal("region")], + _named(_param("region"), _param("limit")), + _check(declared_not_passed=["limit"]), + id="param-no-argument-fills", + ), + pytest.param( + [], + _named(_param("region")), + _check(declared_not_passed=["region"]), + id="argless-against-named-params", + ), + pytest.param( + [_literal("config"), _default("extra")], + _named(_param("region"), _param("limit")), + _check(), + id="lone-argument-without-a-schema-may-be-the-whole-value", + ), + pytest.param( + [_literal("config", schema=OBJECT)], + _named(_param("region"), _param("limit")), + _check(), + id="lone-object-argument-may-be-the-whole-value", + ), + pytest.param( + [_literal("config", schema={"$ref": "#/$defs/Config"})], + _named(_param("region"), _param("limit")), + _check(), + id="lone-argument-of-unknown-type-may-be-the-whole-value", + ), + pytest.param( + [_literal("config", schema=STRING)], + _named(_param("region"), _param("limit")), + _check(passed_not_declared=["config"], declared_not_passed=["region", "limit"]), + id="lone-string-argument-is-not-the-whole-value", + ), + pytest.param( + [_literal("config")], + _named(_param("region", exact_name=True), _param("limit")), + _check(passed_not_declared=["config"], declared_not_passed=["region", "limit"]), + id="lone-argument-with-an-exact-name-param", + ), + pytest.param( + [_literal("config")], + _named(), + _check(passed_not_declared=["config"]), + id="lone-argument-against-no-params", + ), + pytest.param( + [_literal("a"), _literal("b")], + _named(_param("region")), + _check(passed_not_declared=["a", "b"], declared_not_passed=["region"]), + id="two-unmatched-arguments", + ), + pytest.param( + [_literal("user_id"), _literal("userId"), _literal("region")], + _named(_param("UserID"), _param("region")), + _check(passed_not_declared=["user_id", "userId"], declared_not_passed=["UserID"]), + id="folded-name-two-arguments-share", + ), + pytest.param( + [_literal("user_id"), _literal("userId")], + _named(_param("user_id")), + _check(passed_not_declared=["userId"]), + id="exact-name-despite-a-shared-fold", + ), + pytest.param( + [_literal("region"), _literal("limit")], + _named(_param("region"), _param()), + _check(passed_not_declared=["limit"], declared_not_passed=["#1"]), + id="nameless-named-param", + ), + pytest.param( + [_literal("region", schema=INTEGER), _literal("limit", schema=INTEGER)], + _named(_param("limit", schema=INTEGER), _param("region", schema=STRING)), + _check(errors=["argument 'region' is integer, the task handler takes string"]), + id="named-type-mismatch", + ), + pytest.param( + [_literal("a", schema=INTEGER), _default("b", schema=STRING)], + _positional(_param(schema=INTEGER), _param(schema=INTEGER)), + _check(errors=["argument 'b' is string, the task handler takes integer"]), + id="positional-defaulted-argument-type-mismatch", + ), + pytest.param( + [_literal("region", schema=STRING), _default("limit", schema=STRING)], + _named(_param("region", schema=STRING), _param("limit", schema=INTEGER)), + _check(errors=["argument 'limit' is string, the task handler takes integer"]), + id="named-defaulted-argument-type-mismatch", + ), + pytest.param( + [_literal("a", schema=STRING)], + TaskHandlerDeclaration(task_id="transform", binding="named", params=None), + _check(), + id="unlisted-params", + ), + ], +) +def test_check_task_handler_arguments(arg_bindings, declaration, expected): + assert check_task_handler_arguments(arg_bindings, declaration) == expected + + +@pytest.mark.parametrize( + ("stub_schema", "handler_schema", "expected"), + [ + pytest.param(INTEGER, INTEGER, None, id="equal"), + pytest.param(INTEGER, NUMBER, None, id="integer-into-number"), + pytest.param(NUMBER, INTEGER, "number, the task handler takes integer", id="number-into-integer"), + pytest.param(STRING, INTEGER, "string, the task handler takes integer", id="string-into-integer"), + pytest.param( + NULLABLE_INTEGER, + INTEGER, + "integer or null, the task handler takes integer", + id="nullable-into-non-nullable", + ), + pytest.param(INTEGER, NULLABLE_INTEGER, None, id="non-nullable-into-nullable"), + pytest.param(NULLABLE_INTEGER, NULLABLE_INTEGER, None, id="nullable-into-nullable"), + pytest.param( + {"const": None}, STRING, "null, the task handler takes string", id="null-into-non-nullable" + ), + pytest.param({"anyOf": [INTEGER, STRING]}, INTEGER, None, id="union-with-one-accepted-type"), + pytest.param( + {"anyOf": [STRING, {"type": "boolean"}]}, + INTEGER, + "string or boolean, the task handler takes integer", + id="union-with-no-accepted-type", + ), + pytest.param( + {"anyOf": [INTEGER, STRING]}, {"anyOf": [STRING, INTEGER, NULL]}, None, id="any-of-subset" + ), + pytest.param( + {"anyOf": [INTEGER, STRING, NULL]}, + {"oneOf": [STRING, INTEGER]}, + "string, integer or null, the task handler takes string or integer", + id="any-of-superset", + ), + pytest.param( + {"type": ["integer", "null"]}, + INTEGER, + "integer or null, the task handler takes integer", + id="type-list", + ), + pytest.param( + {"type": ["string", "null"]}, {"anyOf": [STRING, NULL]}, None, id="type-list-into-any-of" + ), + pytest.param({"const": "eu"}, STRING, None, id="const"), + pytest.param({"const": 1}, STRING, "integer, the task handler takes string", id="const-mismatch"), + pytest.param({"enum": [1, 2.5]}, NUMBER, None, id="enum"), + pytest.param( + {"enum": ["eu", None]}, + {"enum": ["eu", "us"]}, + "string or null, the task handler takes string", + id="enum-mismatch", + ), + pytest.param( + {"enum": [True]}, INTEGER, "boolean, the task handler takes integer", id="boolean-into-integer" + ), + pytest.param({"$ref": "#/$defs/Config"}, INTEGER, None, id="ref-skipped"), + pytest.param({"allOf": [INTEGER]}, STRING, None, id="all-of-skipped"), + pytest.param( + {"anyOf": [{"$ref": "#/$defs/Config"}, NULL]}, INTEGER, None, id="any-of-with-a-ref-skipped" + ), + pytest.param({}, INTEGER, None, id="empty-stub-schema-skipped"), + pytest.param(STRING, {}, None, id="empty-handler-schema-skipped"), + pytest.param(None, INTEGER, None, id="no-stub-schema"), + pytest.param(STRING, None, None, id="no-handler-schema"), + pytest.param( + {"type": "integer", "format": "int64"}, + {"type": "integer", "format": "int32", "minimum": 0}, + None, + id="format-and-range-ignored", + ), + ], +) +def test_check_value_schema(stub_schema, handler_schema, expected): + assert check_value_schema(stub_schema, handler_schema) == expected + + +def test_a_stub_task_with_one_handler_is_bound(): + stub_tasks = [_stub("extract"), _stub("transform", arg_bindings=[_literal("a")])] + answers = [ + _artifact("bin/extract", etl=[_positional(task_id="extract")]), + _artifact("bin/etl", etl=[_positional(_param(), task_id="transform")]), + ] + + assert _match(stub_tasks, answers) == TaskHandlerMatch( + bindings=[_binding("extract", rel_path="bin/extract"), _binding("transform")], + problems=[], + warnings=[], + ) + + +def test_a_stub_task_without_a_handler_is_a_problem(): + answers = [_artifact(other=[_positional()], etl=[_positional(task_id="extract")])] + + assert _match([_stub()], answers) == TaskHandlerMatch( + bindings=[], + problems=[ + TaskHandlerProblem( + relative_fileloc=DAG_FILE, + message="Dag 'etl', task 'transform': no artifact in Dag bundle 'go-task-handlers' registers it", + dag_id="etl", + task_id="transform", + ) + ], + warnings=[], + ) + + +def test_a_missing_handler_names_the_broken_candidates(): + broken_candidates = { + "bin/old": "bin/old is not executable", + "bin/new": "probe failed: the runtime exited with code 1", + } + + match = _match([_stub()], [_artifact(etl=[_positional(task_id="extract")])], broken_candidates) + + assert [problem.message for problem in match.problems] == [ + "Dag 'etl', task 'transform': no artifact in Dag bundle 'go-task-handlers' registers it; " + "no answer from 'bin/new' (probe failed: the runtime exited with code 1), " + "'bin/old' (bin/old is not executable)" + ] + + +def test_a_stub_task_claimed_by_two_artifacts_names_both(): + answers = [ + _artifact("bin/b", etl=[_positional()]), + _artifact("bin/a", etl=[_positional()]), + _artifact("bin/c", etl=[_positional(task_id="extract")]), + ] + + match = _match([_stub()], answers) + + assert match.bindings == [] + assert [problem.message for problem in match.problems] == [ + "Dag 'etl', task 'transform': registered by 'bin/a' and 'bin/b' in Dag bundle 'go-task-handlers'" + ] + + +def test_a_handler_without_a_stub_task_is_not_a_problem(): + answers = [ + _artifact(etl=[_positional(), _positional(task_id="load")], reporting=[_named(task_id="publish")]) + ] + + assert _match([_stub()], answers) == TaskHandlerMatch(bindings=[_binding()], problems=[], warnings=[]) + + +def test_a_mapped_stub_task_is_checked_for_its_handler_only(): + answers = [_artifact(etl=[_positional(_param(), _param())])] + + assert _match([_stub(is_mapped=True)], answers) == TaskHandlerMatch( + bindings=[_binding()], problems=[], warnings=[] + ) + + +def test_unlisted_params_are_checked_for_the_handler_only(): + stub = _stub(arg_bindings=[_literal("a", schema=STRING)]) + answers = [_artifact(etl=[TaskHandlerDeclaration(task_id="transform", binding="named", params=None)])] + + assert _match([stub], answers) == TaskHandlerMatch(bindings=[_binding()], problems=[], warnings=[]) + + +def test_a_name_mismatch_is_a_warning_and_still_binds(): + stub = _stub(arg_bindings=[_literal("region"), _literal("unused_label")]) + answers = [_artifact(etl=[_named(_param("region"), _param("limit"))])] + + assert _match([stub], answers) == TaskHandlerMatch( + bindings=[_binding()], + problems=[], + warnings=[ + StubTaskWarning( + dag_id="etl", + task_id="transform", + artifact_bundle_name=BUNDLE_NAME, + artifact_rel_path="bin/etl", + passed_not_declared=["unused_label"], + declared_not_passed=["limit"], + ) + ], + ) + + +@pytest.mark.parametrize( + ("params", "error"), + [ + pytest.param( + [_param(schema=INTEGER), _param(), _param()], + "passes 2 arguments, the task handler takes 3", + id="count", + ), + pytest.param( + [_param(schema=INTEGER), _param()], + "argument 'count' is integer or null, the task handler takes integer", + id="value-type", + ), + ], +) +def test_an_argument_problem_names_the_artifact(params, error): + stub = _stub(arg_bindings=[_literal("count", schema=NULLABLE_INTEGER), _literal("label")]) + + match = _match([stub], [_artifact(etl=[_positional(*params)])]) + + assert match.bindings == [_binding()] + assert match.problems == [ + TaskHandlerProblem( + relative_fileloc=DAG_FILE, + message=f"Dag 'etl', task 'transform' ('bin/etl' in Dag bundle 'go-task-handlers'): {error}", + dag_id="etl", + task_id="transform", + ) + ] + + +def test_collect_stub_tasks(): + with DAG(dag_id="etl", schedule=None) as dag: + + @task.stub(queue="go") + def extract(): ... + + @task.stub(queue="go") + def transform(data: str, region: str = "eu"): ... + + @task.stub(queue="go") + def fan_out(item: int): ... + + @task.stub + def unrouted(): ... + + @task + def report(): ... + + transform(extract()) + fan_out.expand(item=[1, 2]) + unrouted() + report() + + with DAG(dag_id="unserialized", schedule=None) as unserialized_dag: + + @task.stub(queue="go") + def load(): ... + + load() + + dag.relative_fileloc = DAG_FILE + + stub_tasks = collect_stub_tasks([dag, unserialized_dag], [LazyDeserializedDAG.from_dag(dag)]) + + assert stub_tasks == [ + StubTask( + dag_id="etl", + task_id="extract", + queue="go", + relative_fileloc=DAG_FILE, + arg_bindings=[], + is_mapped=False, + ), + StubTask( + dag_id="etl", + task_id="transform", + queue="go", + relative_fileloc=DAG_FILE, + arg_bindings=[ + XComArgBinding(kind="xcom", name="data", task_id="extract", value_schema=STRING), + _default("region", "eu", schema=STRING), + ], + is_mapped=False, + ), + StubTask( + dag_id="etl", + task_id="fan_out", + queue="go", + relative_fileloc=DAG_FILE, + arg_bindings=[], + is_mapped=True, + ), + StubTask( + dag_id="etl", + task_id="unrouted", + queue="default", + relative_fileloc=DAG_FILE, + arg_bindings=[], + is_mapped=False, + ), + ] + + +def test_collect_stub_tasks_reads_the_arguments_the_worker_gets(): + with DAG(dag_id="etl", schedule=None) as dag: + + @task.stub(queue="go") + def load(table: str, columns: tuple = ("id", "name")): ... + + @task.stub(queue="go") + def merge(mapping: dict): ... + + load("users") + merge({1: "a"}) + + stub_tasks = collect_stub_tasks([dag], [LazyDeserializedDAG.from_dag(dag)]) + + assert [stub.arg_bindings for stub in stub_tasks] == [ + [ + _literal("table", "users", schema=STRING), + _default("columns", ["id", "name"], schema={"type": "array", "items": {}}), + ], + [_literal("mapping", {"1": "a"}, schema=OBJECT)], + ] + + +def test_format_import_errors_groups_by_file_in_a_stable_order(): + problems = [ + TaskHandlerProblem(relative_fileloc="dags/z.py", message="z: load", dag_id="z", task_id="load"), + TaskHandlerProblem( + relative_fileloc=DAG_FILE, message="etl: transform", dag_id="etl", task_id="transform" + ), + TaskHandlerProblem(relative_fileloc=DAG_FILE, message="Coordinator 'java'"), + TaskHandlerProblem( + relative_fileloc=DAG_FILE, message="etl: load, count", dag_id="etl", task_id="load" + ), + TaskHandlerProblem( + relative_fileloc=DAG_FILE, message="audit: check", dag_id="audit", task_id="check" + ), + TaskHandlerProblem(relative_fileloc=DAG_FILE, message="Coordinator 'go-sdk'"), + TaskHandlerProblem( + relative_fileloc=DAG_FILE, message="etl: load, type", dag_id="etl", task_id="load" + ), + ] + + import_errors = format_import_errors(problems) + + assert list(import_errors) == [DAG_FILE, "dags/z.py"] + assert import_errors == { + "dags/z.py": "Stub tasks in dags/z.py do not match their task handlers:\n- z: load", + DAG_FILE: "\n".join( + [ + "Stub tasks in dags/etl.py do not match their task handlers:", + "- Coordinator 'java'", + "- Coordinator 'go-sdk'", + "- audit: check", + "- etl: load, count", + "- etl: load, type", + "- etl: transform", + ] + ), + } diff --git a/airflow-core/tests/unit/executors/test_base_executor.py b/airflow-core/tests/unit/executors/test_base_executor.py index 437e21a498e59..5e622f9c6d703 100644 --- a/airflow-core/tests/unit/executors/test_base_executor.py +++ b/airflow-core/tests/unit/executors/test_base_executor.py @@ -930,6 +930,37 @@ def test_run_workload_passes_team_name_to_connection_test_supervisor(mock_superv ) +@mock.patch("airflow.sdk.execution_time.supervisor.supervise_task", autospec=True) +def test_run_workload_passes_task_handler_artifact_to_task_supervisor(mock_supervise): + mock_supervise.return_value = 0 + reference = workloads.TaskHandlerArtifactRef( + bundle_info=BundleInfo(name="java-task-handlers"), rel_path="libs/etl.jar" + ) + wl = workloads.ExecuteTask( + ti=workloads.TaskInstanceDTO( + id=uuid4(), + dag_version_id=uuid4(), + task_id="extract", + dag_id="etl", + run_id="r", + try_number=1, + map_index=-1, + pool_slots=1, + queue="jdk-17", + priority_weight=1, + ), + dag_rel_path=Path("etl.py"), + token="test-token", + bundle_info=BundleInfo(name="dags-folder", version="v1"), + log_path="etl.log", + task_handler_artifact=reference, + ) + + BaseExecutor.run_workload(wl, server="http://localhost:8080/execution/") + + assert mock_supervise.call_args.kwargs["task_handler_artifact"] is reference + + @mock.patch.dict("os.environ", {}, clear=True) class TestExecutorConf: """Test ExecutorConf shim class that provides team-specific configuration access.""" diff --git a/airflow-core/tests/unit/executors/test_workloads.py b/airflow-core/tests/unit/executors/test_workloads.py index 2a0a8829d1f68..f5a5d9c9583ab 100644 --- a/airflow-core/tests/unit/executors/test_workloads.py +++ b/airflow-core/tests/unit/executors/test_workloads.py @@ -24,6 +24,7 @@ import jwt import pytest +from pydantic import TypeAdapter from airflow.api_fastapi.auth.tokens import JWTGenerator from airflow.executors import workloads @@ -252,6 +253,59 @@ def test_workload_ti_round_trips_through_sdk_generated_model(): assert not hasattr(received, "pool_slots") +class TestExecuteTaskTaskHandlerArtifact: + @staticmethod + def _workload(task_handler_artifact: workloads.TaskHandlerArtifactRef | None) -> ExecuteTask: + return ExecuteTask( + ti=TaskInstanceDTO( + id=uuid4(), + dag_version_id=uuid4(), + task_id="extract", + dag_id="etl", + run_id="r", + try_number=1, + map_index=-1, + pool_slots=1, + queue="jdk-17", + priority_weight=1, + ), + dag_rel_path=PurePosixPath("etl.py"), + token="token", + bundle_info=BundleInfo(name="dags-folder", version="v1"), + log_path="etl.log", + task_handler_artifact=task_handler_artifact, + ) + + @pytest.mark.parametrize( + "task_handler_artifact", + [ + pytest.param(workloads.TaskHandlerArtifactRef(rel_path="etl.jar"), id="own-bundle"), + pytest.param( + workloads.TaskHandlerArtifactRef( + bundle_info=BundleInfo(name="java-task-handlers"), rel_path="libs/etl.jar" + ), + id="named-bundle", + ), + pytest.param(None, id="none"), + ], + ) + def test_round_trips_through_the_workload_union(self, task_handler_artifact): + workload = self._workload(task_handler_artifact) + + received = TypeAdapter(workloads.All).validate_json(workload.model_dump_json()) + + assert isinstance(received, ExecuteTask) + assert received.task_handler_artifact == task_handler_artifact + + def test_workload_without_the_field_names_no_artifact(self): + payload = self._workload(None).model_dump(mode="json") + del payload["task_handler_artifact"] + + received = TypeAdapter(workloads.All).validate_python(payload) + + assert received.task_handler_artifact is None + + class TestExecuteTaskMakeVersionData: """Tests for ExecuteTask.make() threading version_data through BundleInfo.""" diff --git a/airflow-core/tests/unit/models/test_lang_sdk_task_handler.py b/airflow-core/tests/unit/models/test_lang_sdk_task_handler.py new file mode 100644 index 0000000000000..510e2117dd299 --- /dev/null +++ b/airflow-core/tests/unit/models/test_lang_sdk_task_handler.py @@ -0,0 +1,233 @@ +# +# 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 typing import Any + +import pytest +import uuid6 +from sqlalchemy import delete, insert, select +from sqlalchemy.exc import IntegrityError + +from airflow.models.dag import DagModel +from airflow.models.lang_sdk_task_handler import ( + LangSDKTaskHandler, + LangSDKTaskHandlerArtifact, + compute_fileloc_hash, +) + +pytestmark = pytest.mark.db_test + +ARTIFACT_BUNDLE = "java-task-handlers" +# 6000 bytes in UTF-8, over both the MySQL key limit and the Postgres btree entry limit. +LONG_NON_ASCII_PATH = "任" * 2000 +TASK_HANDLERS = { + "etl": [ + { + "task_id": "extract", + "binding": "positional", + "params": [ + {"name": "path", "value_schema": {"type": "string"}, "required": True, "exact_name": True}, + {"name": None, "value_schema": None, "required": False, "exact_name": False}, + ], + } + ], +} + + +def _make_artifact( + *, + bundle_name: str = ARTIFACT_BUNDLE, + relative_fileloc: str = "etl.jar", + cache_digest: str | None = "0" * 64, + task_handlers: dict[str, list[dict[str, Any]]] | None = None, +) -> LangSDKTaskHandlerArtifact: + return LangSDKTaskHandlerArtifact( + bundle_name=bundle_name, + relative_fileloc=relative_fileloc, + size_bytes=1024, + cache_digest=cache_digest, + task_handlers=TASK_HANDLERS if task_handlers is None else task_handlers, + ) + + +def _make_handler(*, dag_id: str, artifact: LangSDKTaskHandlerArtifact) -> LangSDKTaskHandler: + return LangSDKTaskHandler( + dag_id=dag_id, + task_id="extract", + artifact_id=artifact.id, + dag_bundle_name="testing", + dag_relative_fileloc=f"{dag_id}.py", + ) + + +def _add_dags_and_artifact(*dag_ids: str, session) -> LangSDKTaskHandlerArtifact: + session.add_all(DagModel(dag_id=dag_id, bundle_name="testing") for dag_id in dag_ids) + artifact = _make_artifact() + session.add(artifact) + session.flush() + return artifact + + +def test_deleting_dag_deletes_its_handlers(testing_dag_bundle, session): + artifact = _add_dags_and_artifact("dag_a", "dag_b", session=session) + session.add_all(_make_handler(dag_id=dag_id, artifact=artifact) for dag_id in ("dag_a", "dag_b")) + session.flush() + + session.execute(delete(DagModel).where(DagModel.dag_id == "dag_a")) + + assert session.scalars(select(LangSDKTaskHandler.dag_id)).all() == ["dag_b"] + assert session.scalars(select(LangSDKTaskHandlerArtifact.id)).all() == [artifact.id] + + +@pytest.mark.parametrize( + "task_handlers", + [ + pytest.param(TASK_HANDLERS, id="handlers"), + pytest.param({}, id="none"), + ], +) +def test_artifact_stores_its_task_handlers(session, task_handlers): + session.add(_make_artifact(task_handlers=task_handlers)) + session.flush() + session.expire_all() + + assert session.scalar(select(LangSDKTaskHandlerArtifact.task_handlers)) == task_handlers + + +def test_artifact_can_store_no_cache_digest(session): + session.add(_make_artifact(cache_digest=None)) + session.flush() + session.expire_all() + + assert session.scalar(select(LangSDKTaskHandlerArtifact.cache_digest)) is None + + +def test_deleting_referenced_artifact_fails(testing_dag_bundle, session): + artifact = _add_dags_and_artifact("dag_a", session=session) + session.add(_make_handler(dag_id="dag_a", artifact=artifact)) + session.flush() + + with pytest.raises(IntegrityError): + session.execute( + delete(LangSDKTaskHandlerArtifact).where(LangSDKTaskHandlerArtifact.id == artifact.id) + ) + + +@pytest.mark.parametrize( + "relative_fileloc", + [ + pytest.param("etl.jar", id="short"), + pytest.param(LONG_NON_ASCII_PATH, id="long-non-ascii"), + ], +) +def test_artifact_path_is_unique_per_bundle(session, relative_fileloc): + session.add_all( + [ + _make_artifact(relative_fileloc=relative_fileloc), + _make_artifact(bundle_name="go-task-handlers", relative_fileloc=relative_fileloc), + ] + ) + session.flush() + + session.add(_make_artifact(relative_fileloc=relative_fileloc)) + with pytest.raises(IntegrityError): + session.flush() + + +def test_core_insert_fills_fileloc_hashes(testing_dag_bundle, session): + session.add(DagModel(dag_id="dag_a", bundle_name="testing")) + session.flush() + artifact_id = uuid6.uuid7() + + session.execute( + insert(LangSDKTaskHandlerArtifact).values( + id=artifact_id, + bundle_name=ARTIFACT_BUNDLE, + relative_fileloc="etl.jar", + size_bytes=1024, + cache_digest="0" * 64, + task_handlers={}, + ) + ) + session.execute( + insert(LangSDKTaskHandler).values( + dag_id="dag_a", + task_id="extract", + artifact_id=artifact_id, + dag_bundle_name="testing", + dag_relative_fileloc="dags/etl.py", + ) + ) + + assert session.scalar(select(LangSDKTaskHandlerArtifact.relative_fileloc_hash)) == compute_fileloc_hash( + "etl.jar" + ) + assert session.scalar(select(LangSDKTaskHandler.dag_relative_fileloc_hash)) == compute_fileloc_hash( + "dags/etl.py" + ) + + +def _add_artifact(*, session) -> LangSDKTaskHandlerArtifact: + artifact = _make_artifact() + session.add(artifact) + session.flush() + return artifact + + +def _add_handler(*, session) -> LangSDKTaskHandler: + handler = _make_handler(dag_id="dag_a", artifact=_add_dags_and_artifact("dag_a", session=session)) + session.add(handler) + session.flush() + return handler + + +@pytest.mark.parametrize( + ("add_row", "path_attr", "hash_column", "other_attr", "other_value"), + [ + pytest.param( + _add_artifact, + "relative_fileloc", + LangSDKTaskHandlerArtifact.relative_fileloc_hash, + "size_bytes", + 2048, + id="artifact", + ), + pytest.param( + _add_handler, + "dag_relative_fileloc", + LangSDKTaskHandler.dag_relative_fileloc_hash, + "dag_bundle_name", + "other-bundle", + id="handler", + ), + ], +) +def test_fileloc_hash_follows_orm_updates( + testing_dag_bundle, session, add_row, path_attr, hash_column, other_attr, other_value +): + row = add_row(session=session) + original_path = getattr(row, path_attr) + + setattr(row, other_attr, other_value) + session.flush() + assert session.scalar(select(hash_column)) == compute_fileloc_hash(original_path) + + setattr(row, path_attr, "moved/file") + session.flush() + assert session.scalar(select(hash_column)) == compute_fileloc_hash("moved/file") diff --git a/airflow-core/tests/unit/utils/test_db_cleanup.py b/airflow-core/tests/unit/utils/test_db_cleanup.py index 375461b6848ac..41ec3ff7de9a7 100644 --- a/airflow-core/tests/unit/utils/test_db_cleanup.py +++ b/airflow-core/tests/unit/utils/test_db_cleanup.py @@ -1087,6 +1087,7 @@ def test_no_models_missing(self): # leave alone - per-asset key/value state, upserted in place (PK is asset_id+key), # so it is bounded and current, not accumulating history; removed with its asset "asset_state_store", + "lang_sdk_task_handler_artifact", # parse-time cache of task handler artifacts, not run data # Purged indirectly: each of these hangs off a cleaned table by an # ON DELETE CASCADE foreign key, so the rows go when the parent does. # cascade from dag_run once the partition run has fired; while it is still @@ -1098,6 +1099,7 @@ def test_no_models_missing(self): "hitl_detail", # cascade from task_instance "hitl_detail_history", # cascade from task_instance_history "job_team", # cascade from job + "lang_sdk_task_handler", # cascade from dag "task_inlet_asset_reference", # cascade from dag } diff --git a/airflow-core/tests/unit/utils/test_file.py b/airflow-core/tests/unit/utils/test_file.py index 115027771284e..a5062d40765c7 100644 --- a/airflow-core/tests/unit/utils/test_file.py +++ b/airflow-core/tests/unit/utils/test_file.py @@ -213,6 +213,18 @@ def test_list_py_file_paths(self, test_zip_path): f"Detected files mismatched expected files:\ndetected_files: {pformat(detected_files)}\nexpected_files: {pformat(expected_files)}" ) + @pytest.mark.parametrize("archive_name", ["task_handler.jar", "no_suffix"]) + def test_list_py_file_paths_skips_zip_archive_without_zip_suffix(self, tmp_path, archive_name): + dag_source = "from airflow.sdk import DAG\n" + (tmp_path / "dag.py").write_text(dag_source) + for name in ("lower.zip", "upper.ZIP", archive_name): + with zipfile.ZipFile(tmp_path / name, "w") as zf: + zf.writestr("dag.py", dag_source) + + assert set(list_py_file_paths(tmp_path)) == { + str(tmp_path / name) for name in ("dag.py", "lower.zip", "upper.ZIP") + } + @pytest.mark.parametrize( ("edge_filename", "expected_modification"), diff --git a/airflow-e2e-tests/docker/go.yml b/airflow-e2e-tests/docker/go.yml index c142696376586..780a8485c6fde 100644 --- a/airflow-e2e-tests/docker/go.yml +++ b/airflow-e2e-tests/docker/go.yml @@ -18,12 +18,22 @@ # Docker Compose override for go_sdk E2E test mode. # # The Go bundle compiled by airflow-go-pack (conftest._setup_go_sdk_integration) -# is a self-contained, statically linked executable, so the stock Airflow worker -# image can exec it directly -- no extra runtime needs to be installed. We only -# bind-mount the bundle into the directory the ExecutableCoordinator scans and -# point the worker at the "golang" queue where @task.stub tasks are routed. +# is a self-contained, statically linked executable, so the stock Airflow image +# can exec it directly -- no extra runtime needs to be installed. We only +# bind-mount the bundle into the directory the ExecutableCoordinator scans, on +# the worker and the Dag processor, and point the worker at the "golang" queue +# where @task.stub tasks are routed. The Dag processor runs the bundle to check +# the stub tasks of go_examples.py against the task handlers it registers. --- services: + airflow-dag-processor: + # As on the worker, below. + environment: + USER: airflow + HOME: /home/airflow + volumes: + - ./go-bundles:/opt/airflow/go-bundles:ro + airflow-worker: # The bundle is built with CGO_ENABLED=0, so the SDK's user.Current() call at # init falls back to Go's pure-Go resolver, which reads $USER / $HOME and diff --git a/airflow-e2e-tests/docker/java.yml b/airflow-e2e-tests/docker/java.yml index db4897f44c818..ed6dd51949a62 100644 --- a/airflow-e2e-tests/docker/java.yml +++ b/airflow-e2e-tests/docker/java.yml @@ -17,15 +17,24 @@ # Docker Compose override for java_sdk E2E test mode. # -# Replaces the stock airflow-worker image with one that has a JRE installed -# (built by conftest._setup_java_sdk_integration via Dockerfile.java), mounts -# the pre-built bundle JARs (the Java example under /opt/airflow/java-jars, the -# Scala Spark example under /opt/airflow/scala-jars, and the runner-behaviour -# test fixtures under /opt/airflow/java-test-jars), and configures the worker to -# consume the "java", "scala", and "java-test" Celery queues where @task.stub -# tasks are routed. +# Replaces the stock airflow-worker and airflow-dag-processor image with one +# that has a JRE installed (built by conftest._setup_java_sdk_integration via +# Dockerfile.java), mounts the pre-built bundle JARs on both (the Java example +# under /opt/airflow/java-jars, the Scala Spark example under +# /opt/airflow/scala-jars, and the runner-behaviour test fixtures under +# /opt/airflow/java-test-jars), and configures the worker to consume the "java", +# "scala", and "java-test" Celery queues where @task.stub tasks are routed. The +# Dag processor runs the JARs to check the stub tasks of the Python Dags against +# the task handlers they register. --- services: + airflow-dag-processor: + image: airflow-java-worker + volumes: + - ./java-jars:/opt/airflow/java-jars:ro + - ./scala-jars:/opt/airflow/scala-jars:ro + - ./java-test-jars:/opt/airflow/java-test-jars:ro + airflow-worker: image: airflow-java-worker volumes: diff --git a/airflow-e2e-tests/docker/ts.yml b/airflow-e2e-tests/docker/ts.yml index 056106ff9d54a..474e675e78784 100644 --- a/airflow-e2e-tests/docker/ts.yml +++ b/airflow-e2e-tests/docker/ts.yml @@ -17,10 +17,12 @@ # Docker Compose override for ts_sdk E2E test mode. # -# The stock worker image ships no Node.js runtime, so node-provider copies the +# The stock Airflow image ships no Node.js runtime, so node-provider copies the # node binary from the same image conftest builds the bundle with into a -# shared volume. The bundle is bind-mounted where NodeCoordinator scans, and -# the worker consumes the "typescript" queue where @task.stub tasks are routed. +# shared volume. The bundle and the runtime are mounted on the worker and the +# Dag processor, which runs the bundle to check the stub tasks of the example +# Dags against the task handlers it registers. The worker consumes the +# "typescript" queue where @task.stub tasks are routed. --- services: node-provider: @@ -29,6 +31,14 @@ services: volumes: - nodejs-bin:/opt/nodejs + airflow-dag-processor: + volumes: + - ./ts-bundles:/opt/airflow/ts-bundles:ro + - nodejs-bin:/opt/nodejs:ro + depends_on: + node-provider: + condition: service_completed_successfully + airflow-worker: volumes: - ./ts-bundles:/opt/airflow/ts-bundles:ro diff --git a/airflow-e2e-tests/java-test-bundle/src/java/org/apache/airflow/e2e/TestBundleBuilder.java b/airflow-e2e-tests/java-test-bundle/src/java/org/apache/airflow/e2e/TestBundleBuilder.java index ad33a15708a46..1c15b425932df 100644 --- a/airflow-e2e-tests/java-test-bundle/src/java/org/apache/airflow/e2e/TestBundleBuilder.java +++ b/airflow-e2e-tests/java-test-bundle/src/java/org/apache/airflow/e2e/TestBundleBuilder.java @@ -19,7 +19,6 @@ package org.apache.airflow.e2e; -import java.util.List; import org.apache.airflow.sdk.*; import org.jetbrains.annotations.NotNull; @@ -27,7 +26,7 @@ * Bundle for the runner-behaviour E2E tests: deliberately broken task classes that exercise * instantiation failures, and a task that round-trips Airflow Variables through the supervisor. */ -public class TestBundleBuilder implements BundleBuilder { +public class TestBundleBuilder { public static class MissingNoArgConstructor implements Task { public MissingNoArgConstructor(String unused) {} @@ -60,19 +59,15 @@ public void execute(@NotNull Context context, Client client) { } } - @NotNull - @Override - public Iterable getDags() { - var uninstantiable = new DagDef("java_uninstantiable"); - uninstantiable.addTask("missing_no_arg_constructor", MissingNoArgConstructor.class); - uninstantiable.addTask("non_static_inner", NonStaticInner.class); - var variableWrite = new DagDef("java_variable_write"); - variableWrite.addTask("write_and_delete", WriteAndDeleteVariable.class); - return List.of(uninstantiable, variableWrite); + public static Bundle build() { + return new Bundle() + .register( + "java_uninstantiable", "missing_no_arg_constructor", MissingNoArgConstructor.class) + .register("java_uninstantiable", "non_static_inner", NonStaticInner.class) + .register("java_variable_write", "write_and_delete", WriteAndDeleteVariable.class); } public static void main(String[] args) { - var bundle = new TestBundleBuilder().build(); - Server.create(args).serve(bundle); + Server.create(args).serve(build()); } } diff --git a/airflow-e2e-tests/tests/airflow_e2e_tests/conftest.py b/airflow-e2e-tests/tests/airflow_e2e_tests/conftest.py index 39ff3fd0a64b8..08ad58f4eea14 100644 --- a/airflow-e2e-tests/tests/airflow_e2e_tests/conftest.py +++ b/airflow-e2e-tests/tests/airflow_e2e_tests/conftest.py @@ -293,6 +293,22 @@ def _setup_xcom_object_storage_integration(dot_env_file, tmp_dir): ] +def _build_dag_bundle_config(artifact_bundles: dict[str, str]) -> str: + """Return a ``dag_bundle_config_list`` of the Dags folder plus a ``LocalDagBundle`` per artifact path. + + Registration is what makes a coordinator's ``task_handler_bundle_name`` resolvable on the + worker and the Dag processor. The artifact directories are mounted on both; the Dag processor + runs the artifacts in them to check the stub tasks, and finds no Dag files there. + """ + local_bundle = "airflow.dag_processing.bundles.local.LocalDagBundle" + bundles = [{"name": "dags-folder", "classpath": local_bundle, "kwargs": {}}] + bundles.extend( + {"name": name, "classpath": local_bundle, "kwargs": {"path": path}} + for name, path in artifact_bundles.items() + ) + return json.dumps(bundles) + + def _run_java_sdk_gradle(workdir, *gradle_argv, capture_output=False, native=False): """Run the Java SDK Gradle wrapper natively or inside the pinned JDK container. @@ -416,7 +432,8 @@ def _setup_java_sdk_integration(dot_env_file, tmp_dir): copyfile(JAVA_DOCKERFILE_PATH, tmp_dir / "Dockerfile.java") # Copy each bundle's JARs into its own directory; the compose bind-mounts - # expose them to the worker, and each JavaCoordinator globs its own dir. + # expose them to the worker and the Dag processor, where each is registered + # as its own Dag bundle. copytree(JAVA_SDK_EXAMPLE_LIBS_PATH, tmp_dir / "java-jars") copytree(SCALA_SPARK_EXAMPLE_LIBS_PATH, tmp_dir / "scala-jars") copytree(JAVA_TEST_BUNDLE_LIBS_PATH, tmp_dir / "java-test-jars") @@ -432,7 +449,8 @@ def _setup_java_sdk_integration(dot_env_file, tmp_dir): # Keep the bundle JARs out of the build context: Dockerfile.java only adds a # JRE and copies nothing from the context, so without this docker build would # tar and stream the bundles (hundreds of MB of Spark JARs) to the daemon for - # nothing. The JARs reach the worker via the compose bind-mounts, not the image. + # nothing. The JARs reach the worker and the Dag processor via the compose + # bind-mounts, not the image. (tmp_dir / ".dockerignore").write_text("java-jars/\nscala-jars/\njava-test-jars/\n") # Build a local Docker image that extends DOCKER_IMAGE with a JRE. @@ -456,20 +474,27 @@ def _setup_java_sdk_integration(dot_env_file, tmp_dir): ) # One JavaCoordinator per queue on the same worker image, each serving its - # own bundle. The scala-jdk entry pins main_class (Spark's large classpath - # makes Main-Class discovery ambiguous) and carries Spark's Java 17 module - # openings, a small driver heap, and a longer startup timeout for its large - # dependency classpath. + # own artifact bundle (one bundle is one classpath). The scala-jdk entry pins + # main_class (Spark's large classpath makes Main-Class discovery ambiguous) + # and carries Spark's Java 17 module openings, a small driver heap, and a + # longer startup timeout for its large dependency classpath. + dag_bundle_config = _build_dag_bundle_config( + { + "java-task-handlers": "/opt/airflow/java-jars", + "scala-task-handlers": "/opt/airflow/scala-jars", + "java-test-task-handlers": "/opt/airflow/java-test-jars", + } + ) coordinator_config = json.dumps( { "java-jdk": { "classpath": "airflow.sdk.coordinators.java.JavaCoordinator", - "kwargs": {"jars_root": ["/opt/airflow/java-jars"]}, + "kwargs": {"task_handler_bundle_name": "java-task-handlers"}, }, "scala-jdk": { "classpath": "airflow.sdk.coordinators.java.JavaCoordinator", "kwargs": { - "jars_root": ["/opt/airflow/scala-jars"], + "task_handler_bundle_name": "scala-task-handlers", "main_class": "org.apache.airflow.example.ScalaSparkBundleBuilder", "jvm_args": ["-Xmx512m", *_SPARK_JAVA_MODULE_OPTIONS], "task_startup_timeout": 60.0, @@ -477,7 +502,7 @@ def _setup_java_sdk_integration(dot_env_file, tmp_dir): }, "java-test-jdk": { "classpath": "airflow.sdk.coordinators.java.JavaCoordinator", - "kwargs": {"jars_root": ["/opt/airflow/java-test-jars"]}, + "kwargs": {"task_handler_bundle_name": "java-test-task-handlers"}, }, } ) @@ -502,6 +527,7 @@ def _setup_java_sdk_integration(dot_env_file, tmp_dir): dot_env_file.write_text( f"AIRFLOW_UID={os.getuid()}\n" # Single-quote the JSON values so Docker Compose reads them literally. + f"AIRFLOW__DAG_PROCESSOR__DAG_BUNDLE_CONFIG_LIST='{dag_bundle_config}'\n" f"AIRFLOW__SDK__COORDINATORS='{coordinator_config}'\n" f"AIRFLOW__SDK__QUEUE_TO_COORDINATOR='{queue_to_coordinator}'\n" f"AIRFLOW_CONN_TEST_HTTP='{test_http_conn}'\n" @@ -532,8 +558,8 @@ def _run_go_sdk_pack(output_path, *, capture_output=False, native=False): subsequent runs skip straight to compilation). * USER/HOME must be set because the SDK calls user.Current() at init; with cgo disabled Go's pure-Go resolver reads those env vars instead of libc, - and panics if either is empty (the same vars are set on the worker in - go.yml so the packed binary runs the same way at execution time). + and panics if either is empty (the same vars are set on the worker and + the Dag processor in go.yml so the packed binary runs the same way there). """ if native: cwd = GO_SDK_ROOT_PATH @@ -598,13 +624,14 @@ def _setup_go_sdk_integration(dot_env_file, tmp_dir): """Set up the go_sdk E2E test mode. Compiles the Go SDK example bundle into a self-contained executable bundle - via the ``airflow-go-pack`` tooling, drops it into the directory the - ``ExecutableCoordinator`` scans, copies the Python stub Dag, and writes the - coordinator configuration. + via the ``airflow-go-pack`` tooling, drops it into the directory registered + as the ``go-task-handlers`` Dag bundle, copies the Python stub Dag, and + writes the coordinator configuration. The packed bundle is a statically linked native executable (built with - ``CGO_ENABLED=0``), so the stock Airflow worker image can exec it directly - without a Go toolchain or any extra runtime installed -- see ``go.yml``. + ``CGO_ENABLED=0``), so the stock Airflow image can exec it directly on the + worker and the Dag processor without a Go toolchain or any extra runtime + installed -- see ``go.yml``. """ _pack_go_sdk_example_bundle(native=LANG_SDK_NATIVE_TOOLCHAIN) @@ -612,8 +639,8 @@ def _setup_go_sdk_integration(dot_env_file, tmp_dir): copyfile(GO_COMPOSE_PATH, tmp_dir / "go.yml") # Place the packed bundle where the compose bind-mount (./go-bundles) exposes - # it to the worker at /opt/airflow/go-bundles. The bundle scanner requires - # the file to be executable, so preserve the exec bit. + # it to the worker and the Dag processor at /opt/airflow/go-bundles. The + # coordinator runs only an executable file, so preserve the exec bit. go_bundles_dir = tmp_dir / "go-bundles" go_bundles_dir.mkdir() packed_bundle = go_bundles_dir / GO_SDK_BUNDLE_NAME @@ -624,13 +651,14 @@ def _setup_go_sdk_integration(dot_env_file, tmp_dir): copyfile(GO_SDK_DAGS_PATH / "go_examples.py", tmp_dir / "dags" / "go_examples.py") # Coordinator registry: maps the logical name "go-sdk" to ExecutableCoordinator, - # which scans executables_root for the packed bundle by dag_id. + # which scans the go-task-handlers Dag bundle for the packed bundle by dag_id. # Queue mapping: routes tasks on the "golang" queue to "go-sdk". + dag_bundle_config = _build_dag_bundle_config({"go-task-handlers": "/opt/airflow/go-bundles"}) coordinator_config = json.dumps( { "go-sdk": { "classpath": "airflow.sdk.coordinators.executable.ExecutableCoordinator", - "kwargs": {"executables_root": ["/opt/airflow/go-bundles"]}, + "kwargs": {"task_handler_bundle_name": "go-task-handlers"}, } } ) @@ -639,6 +667,7 @@ def _setup_go_sdk_integration(dot_env_file, tmp_dir): dot_env_file.write_text( f"AIRFLOW_UID={os.getuid()}\n" # Single-quote the JSON values so Docker Compose reads them literally. + f"AIRFLOW__DAG_PROCESSOR__DAG_BUNDLE_CONFIG_LIST='{dag_bundle_config}'\n" f"AIRFLOW__SDK__COORDINATORS='{coordinator_config}'\n" f"AIRFLOW__SDK__QUEUE_TO_COORDINATOR='{queue_to_coordinator}'\n" # Connection and variable read by the Go example bundle tasks. @@ -730,12 +759,13 @@ def _setup_ts_sdk_integration(dot_env_file, tmp_dir): for dag_file in ("typescript_example.py", "typescript_taskflow_example.py"): copyfile(TS_SDK_EXAMPLE_PATH / "dags" / dag_file, tmp_dir / "dags" / dag_file) + dag_bundle_config = _build_dag_bundle_config({"ts-task-handlers": "/opt/airflow/ts-bundles"}) coordinator_config = json.dumps( { "ts": { "classpath": "airflow.sdk.coordinators.node.NodeCoordinator", "kwargs": { - "bundles_root": ["/opt/airflow/ts-bundles"], + "task_handler_bundle_name": "ts-task-handlers", "node_executable": "/opt/nodejs/node", }, } @@ -747,6 +777,7 @@ def _setup_ts_sdk_integration(dot_env_file, tmp_dir): f"AIRFLOW_UID={os.getuid()}\n" f"NODE_IMAGE={NODE_IMAGE}\n" # single-quoted so Docker Compose reads the JSON literally + f"AIRFLOW__DAG_PROCESSOR__DAG_BUNDLE_CONFIG_LIST='{dag_bundle_config}'\n" f"AIRFLOW__SDK__COORDINATORS='{coordinator_config}'\n" f"AIRFLOW__SDK__QUEUE_TO_COORDINATOR='{queue_to_coordinator}'\n" "AIRFLOW_CONN_TYPESCRIPT_EXAMPLE_HTTP=http://user:pass@example.com/\n" diff --git a/contributing-docs/30_new_language_sdk.rst b/contributing-docs/30_new_language_sdk.rst index 34c4cdbb4a27a..91ceb7f6df6df 100644 --- a/contributing-docs/30_new_language_sdk.rst +++ b/contributing-docs/30_new_language_sdk.rst @@ -121,10 +121,67 @@ The method returns a ``(command, subprocess_schema_version)`` pair: across SDK versions. See `Supervisor Schema`_ below. Call ``self._get_scan_roots()`` to retrieve the artifact directories the base -class has already resolved from the coordinator's configured source — an -explicit filesystem root, a named Dag bundle (``dag_bundle_name``), or the -task's own bundle. Subclasses should scan those roots rather than reading the -configured root directly. +class has already resolved: the Dag bundle named by ``task_handler_bundle_name``, +or the task's own Dag bundle when it is unset, pinned for the whole task. +Subclasses should scan those directories rather than locating artifacts +themselves. + +SubprocessCoordinator: implementing ``_build_parse_task_handler_command`` +~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~ + +When a Python Dag file has stub tasks, the Dag processor asks the artifact that +implements them which task handlers it registers, and checks each stub task +against its handler. A coordinator opts in by building the command that starts +its runtime for one artifact: + +.. code-block:: python + + def _build_parse_task_handler_command(self, *, path: pathlib.Path) -> tuple[list[str], str | None]: ... + +*path* is the artifact the Dag processor picked among the candidates of the next +section, so the method does not search for one. The returned pair follows the +rules of ``_build_execute_task_command``: no ``--comm`` or ``--logs`` flags, and +the schema version the runtime understands. The runtime then answers as +described in `Answering TaskHandlerParseRequest`_. + +The default raises ``NotImplementedError``, so a coordinator that does not +implement it cannot be probed. Its artifacts then have no answer, and a stub +task that no other artifact registers fails to import. ``ExecutableCoordinator`` +implements it for executable bundles. + +SubprocessCoordinator: implementing ``_read_task_handler_candidate`` +~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~ + +The Dag processor finds the artifacts to ask by walking the coordinator's Dag +bundle in a stable order. It calls this method once for every regular file, +except one it already listed under another path through a symlink, and keeps +the candidates it returns: + +.. code-block:: python + + def _read_task_handler_candidate( + self, path: pathlib.Path, *, rel_path: str + ) -> TaskHandlerCandidate | None: ... + +Return ``None`` for a file that is not one of your artifacts, such as a +dependency your runtime loads. For an artifact, return a +:class:`~airflow.sdk.execution_time.coordinator.TaskHandlerCandidate` with +*rel_path*, the size of the file you read, and the cache digest the artifact +stores (``None`` when it stores none). The method runs for every file on every +parse, so read the stored digest instead of hashing the artifact. The Dag +processor asks an artifact again when its size or digest changes, so the digest +must change whenever the task handlers the artifact registers can change. A +digest longer than 128 characters counts as none, so the artifact is then asked +on every parse. + +Set ``error`` on an artifact of yours that cannot be asked, for example a file +that is not executable. The Dag processor logs it and does not run it, and names +it in the import error of a stub task that finds no task handler. + +The default raises ``NotImplementedError``, so a coordinator that does not +implement it cannot be probed. The Dag processor then reports the stub tasks +routed to the coordinator as an import error of their Dag file, since they +cannot be bound to an artifact. Supervisor Schema ~~~~~~~~~~~~~~~~~ @@ -195,8 +252,10 @@ as soon as possible. The supervisor verifies that the connecting peer belongs to the launched process tree, so the SDK MUST connect from the same process or one of its descendants. -Once both connections are accepted, the supervisor sends a ``StartupDetails`` -message on the comm socket to initiate execution. +Once both connections are accepted, the supervisor sends the first message on +the comm socket. ``StartupDetails`` starts a task. ``TaskHandlerParseRequest`` +comes from the Dag processor instead, and asks which task handlers the runtime +registers (see `Answering TaskHandlerParseRequest`_). Wire protocol ~~~~~~~~~~~~~ @@ -273,6 +332,40 @@ calls themselves) can already produce log records before the ``--logs`` socket in stage 3 exists to carry them. See `Logging`_ below for how to handle that gap. +Answering ``TaskHandlerParseRequest`` +~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~ + +The SDK sends back one ``TaskHandlerParsingResult``, waits for the supervisor's +response to it, and exits. No user code runs; the answer comes from the SDK's +own task handler registrations. It MUST declare every registered handler and +depend only on the artifact, never on the request. + +* ``task_handlers`` maps each Dag id the artifact registers handlers for to + their declarations, in registration order. +* Each ``TaskHandlerDeclaration`` states in ``binding`` how stub-task arguments + bind to its ``params``: + + * ``positional``: by position. An argument count or value type the handler + cannot take makes the Dag fail to import. + * ``named``: by name in any order, ignoring case and underscores unless a + param sets ``exact_name``. An argument or param that matches nothing is + logged as a warning and the task still runs, so the runtime must accept + both. When no param matches and exactly one argument was passed, it may be + the whole value and is not warned about, unless ``params`` is empty, a + param sets ``exact_name``, or the argument cannot be an object. A value + type a param does not accept makes the Dag fail to import. + +* ``params`` lists the handler's parameters in order, ``[]`` when it has none, + or is ``null`` when the SDK cannot list them. Then only the handler's + presence is checked. +* Each ``TaskHandlerParam`` has a ``name`` (``null`` when the SDK has no name + for a positional parameter) and a ``value_schema``: the JSON Schema of the + values it accepts, in the vocabulary ``@task.stub`` uses for Python + annotations, or ``null`` when the SDK cannot state one. +* An empty ``task_handlers`` mapping is sent empty, never as ``null``. + +``go-sdk/pkg/execution`` is a reference implementation. + Logging ~~~~~~~ diff --git a/dev/breeze/doc/images/output_k8s_setup-lang-sdk-test.svg b/dev/breeze/doc/images/output_k8s_setup-lang-sdk-test.svg index 55b1d78afae62..01db1bbce13e0 100644 --- a/dev/breeze/doc/images/output_k8s_setup-lang-sdk-test.svg +++ b/dev/breeze/doc/images/output_k8s_setup-lang-sdk-test.svg @@ -1,4 +1,4 @@ - +