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 f35012832d6e2..1bc631f426d71 100644 --- a/airflow-core/docs/authoring-and-scheduling/language-sdks/go.rst +++ b/airflow-core/docs/authoring-and-scheduling/language-sdks/go.rst @@ -502,8 +502,9 @@ Deploying 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. +launches the matching bundle. If multiple usable bundles declare the same Dag, the first in sorted path order +wins. 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`` 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..876452126abec --- /dev/null +++ b/airflow-core/tests/unit/dag_processing/test_task_handler_processor_go.py @@ -0,0 +1,152 @@ +# 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 and probe it for its task handlers.""" + +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._subprocess import supports_task_handler_parsing +from airflow.sdk.execution_time import supervisor +from airflow.sdk.execution_time.coordinator import get_coordinator_manager, reset_coordinator_manager + +from tests_common.test_utils.config import conf_vars +from tests_common.test_utils.paths import AIRFLOW_ROOT_PATH + +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"}]} + + +@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") + bundle = tmp_path_factory.mktemp("go-task-handlers") / "example_dags" + completed = subprocess.run( + ["go", "tool", "airflow-go-pack", "--output", os.fspath(bundle), "./example/bundle"], + cwd=AIRFLOW_ROOT_PATH / "go-sdk", + env={**os.environ, "CGO_ENABLED": "0"}, + capture_output=True, + text=True, + check=False, + ) + assert completed.returncode == 0, completed.stderr + return bundle + + +@pytest.fixture(autouse=True) +def _go_coordinator(): + 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): + coordinator = get_coordinator_manager().get_coordinator("go") + artifact = coordinator._find_task_handler_artifact(bundle_path=go_bundle.parent, dag_id="simple_dag") + assert supports_task_handler_parsing(artifact.schema_version) + assert artifact.path == go_bundle.resolve() + + result = LangSDKTaskHandlerProcessorProcess.run( + coordinator="go", + path=artifact.path, + bundle_path=go_bundle.parent, + bundle_name="go-task-handlers", + artifact_rel_path=os.fspath(artifact.path.relative_to(go_bundle.parent.resolve())), + 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), + ], + ) diff --git a/contributing-docs/30_new_language_sdk.rst b/contributing-docs/30_new_language_sdk.rst index 86c65efc716ea..5bf40ad999848 100644 --- a/contributing-docs/30_new_language_sdk.rst +++ b/contributing-docs/30_new_language_sdk.rst @@ -126,6 +126,43 @@ 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: answering a task handler parse +~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~ + +When a Python Dag file has stub tasks, the Dag processor checks each stub task +against the handler registered by the artifact that a worker would run for it. +A coordinator opts in by implementing two methods: + +.. code-block:: python + + def _find_task_handler_artifact(self, *, bundle_path: pathlib.Path, dag_id: str) -> ResolvedBundle: ... + + + def _build_parse_task_handler_command(self, *, path: pathlib.Path) -> tuple[list[str], str | None]: ... + +``_find_task_handler_artifact`` returns the artifact a stub task of *dag_id* +runs and its supervisor schema version, found in the Dag bundle at +*bundle_path* the same way ``_build_execute_task_command`` finds it. It scans +*bundle_path* itself, because ``self._get_scan_roots()`` is not set there. It +raises ``FileNotFoundError`` when the task would find no artifact it can run, +or ``ValueError`` when the artifact it found uses a supervisor schema version +this Task SDK does not know (``ExecutableCoordinator`` skips such bundles, so +it raises ``FileNotFoundError`` instead). + +``_build_parse_task_handler_command`` builds the command that starts the +runtime for the artifact at *path*, which ``_find_task_handler_artifact`` +returned. ``self._get_scan_roots()`` returns the root of the Dag bundle holding +*path*, for a command that needs it, such as a classpath. The returned pair +follows the rules of ``_build_execute_task_command``: no ``--comm`` or +``--logs`` flags, and the schema version the runtime understands. A runtime +whose schema version is older than ``TASK_HANDLER_PARSING_SCHEMA_VERSION`` is +not started. The runtime answers as described in +`Answering TaskHandlerParseRequest`_. + +Both defaults raise ``NotImplementedError``, so the stub tasks of a coordinator +that does not implement them are not checked. ``ExecutableCoordinator`` +implements both for executable bundles. + Supervisor Schema ~~~~~~~~~~~~~~~~~ @@ -195,8 +232,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. A ``TaskHandlerParseRequest`` +comes from the Dag processor instead, and asks which task handlers the runtime +registers (see `Answering TaskHandlerParseRequest`_). Wire protocol ~~~~~~~~~~~~~ @@ -273,6 +312,43 @@ 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 response to +it, and exits. No task function runs: the answer lists the handlers +registered through the SDK's ``TaskHandler`` interface, and never the tasks of +a native Dag. ``task_handlers`` MUST declare every registered handler and +depend only on the artifact, never on the request. ``fileloc`` repeats the +request's ``file``. + +* ``task_handlers`` maps each Dag id the artifact registers handlers for to + their declarations, in registration order. A stub task the answer does not + declare makes the Dag fail to import. +* 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 handler does not constrain it. +* An empty ``task_handlers`` mapping is sent empty, never as ``null``. + +``go-sdk/pkg/execution`` is a reference implementation. + Logging ~~~~~~~ diff --git a/go-sdk/README.md b/go-sdk/README.md index 57e6d4378de2e..90fab00486819 100644 --- a/go-sdk/README.md +++ b/go-sdk/README.md @@ -389,6 +389,10 @@ Python supervisor / task runner protocol on the comm socket, with structured JSON-line logs on the logs socket. - The Python runtime is the worker. It proxies every `GetConnection` / `GetVariable` / `GetXCom` / `SetXCom` call through to the Execution API. The Go binary just runs the task function. +- The Dag processor starts the binary a stub task would run in the same way, to check a Python Dag's stub + tasks. It sends a `TaskHandlerParseRequest`, and the binary replies with every task handler registered + with `airflow.TaskHandler`, keyed by Dag id in registration order, and how each handler's parameters + bind. Then it exits. The Go side of the protocol is implemented in `pkg/execution/`. On the Python side it is the `ExecutableCoordinator` in `task-sdk/src/airflow/sdk/coordinators/executable/coordinator.py`. diff --git a/go-sdk/airflow/bundle.go b/go-sdk/airflow/bundle.go index d100231d8320e..9ca70299a10f8 100644 --- a/go-sdk/airflow/bundle.go +++ b/go-sdk/airflow/bundle.go @@ -130,18 +130,15 @@ func (b *BundleRef) Register(items ...Registerable) { } // taskHandlerMap holds the registered task handlers by dag_id and task_id. -// It also keeps registration order. The --airflow-metadata manifest lists the tasks of each Dag -// in that order. +// It also keeps registration order. The --airflow-metadata manifest, and the reply to the Dag +// processor's task handler parse, list the tasks of each Dag in that order. type taskHandlerMap struct { mu sync.RWMutex handlers map[string]map[string]bundle.Task order []bundle.TaskHandlerInfo } -var ( - _ bundle.Bundle = (*taskHandlerMap)(nil) - _ bundle.EnumerableBundle = (*taskHandlerMap)(nil) -) +var _ bundle.Registry = (*taskHandlerMap)(nil) func (m *taskHandlerMap) add(dagID, taskID string, task bundle.Task) { m.mu.Lock() diff --git a/go-sdk/airflow/dag.go b/go-sdk/airflow/dag.go index 8a45c867eb4b8..e4f982a493093 100644 --- a/go-sdk/airflow/dag.go +++ b/go-sdk/airflow/dag.go @@ -58,7 +58,8 @@ type DagRef struct { // panic once the Dag is registered. // // [BundleRef.Serve] does not yet serve the Dags that Dag returns. It leaves them out of the -// --airflow-metadata manifest and cannot run their tasks. +// --airflow-metadata manifest and cannot run their tasks. Their tasks are never task handlers, so +// the reply to the Dag processor's task handler parse leaves them out too. func Dag(dagID string, spec ...DagSpec) *DagRef { if len(spec) > 1 { panic(fmt.Sprintf( diff --git a/go-sdk/airflow/serve.go b/go-sdk/airflow/serve.go index bcfb33ab934d1..f540fa1b82692 100644 --- a/go-sdk/airflow/serve.go +++ b/go-sdk/airflow/serve.go @@ -57,8 +57,8 @@ const ( // The command-line flags of the executable decide what Serve does. // With --airflow-metadata it prints the bundle's manifest and returns, which is how // airflow-go-pack reads the Dag and task ids of the registered task handlers. -// With --comm and --logs, which the Airflow supervisor passes, it runs one task over the -// coordinator protocol. +// With --comm and --logs, which Airflow passes, it speaks the coordinator protocol: it either +// runs one task, or tells the Dag processor every task handler registered with [TaskHandler]. // // main must exit with a non-zero status when Serve returns an error, because the exit status // is how the supervisor learns that the task failed: diff --git a/go-sdk/airflow/serve_test.go b/go-sdk/airflow/serve_test.go index 93de70af0231c..c925750f81230 100644 --- a/go-sdk/airflow/serve_test.go +++ b/go-sdk/airflow/serve_test.go @@ -284,3 +284,65 @@ func TestServeRunsTaskForSupervisor(t *testing.T) { } assert.True(t, ran) } + +// A fake Dag processor sends TaskHandlerParseRequest over the comm socket, as the probe does after +// it starts the bundle with --comm and --logs. The tasks of a Dag from Dag are not task handlers, +// so the reply leaves them out. +func TestServeDeclaresTaskHandlersButNotDags(t *testing.T) { + commLn, err := net.Listen("tcp", "127.0.0.1:0") + require.NoError(t, err) + defer commLn.Close() + logsLn, err := net.Listen("tcp", "127.0.0.1:0") + require.NoError(t, err) + defer logsLn.Close() + deadline := time.Now().Add(10 * time.Second) + require.NoError(t, commLn.(*net.TCPListener).SetDeadline(deadline)) + require.NoError(t, logsLn.(*net.TCPListener).SetDeadline(deadline)) + + native := Dag("native_etl") + native.Task(extract) + b := Bundle() + b.Register(TaskHandler("py_etl", "transform", noop), native) + + done := make(chan error, 1) + go func() { + done <- b.serve( + []string{"--comm", commLn.Addr().String(), "--logs", logsLn.Addr().String()}, + io.Discard, + ) + }() + + commConn, err := commLn.Accept() + require.NoError(t, err) + defer commConn.Close() + logsConn, err := logsLn.Accept() + require.NoError(t, err) + defer logsConn.Close() + require.NoError(t, commConn.SetDeadline(deadline)) + + dagProcessor := execution.NewCoordinatorComm(commConn, commConn, discardLogger()) + require.NoError(t, dagProcessor.SendRequest(0, map[string]any{ + "type": "TaskHandlerParseRequest", + "file": "/bundles/etl", + "bundle_path": "/bundles", + "bundle_name": "go", + })) + + frame, err := dagProcessor.ReadMessage() + require.NoError(t, err) + var body map[string]any + require.NoError(t, msgpack.Unmarshal(frame.Body, &body)) + assert.Equal(t, "TaskHandlerParsingResult", body["type"]) + handlers, ok := body["task_handlers"].(map[string]any) + require.True(t, ok) + assert.Len(t, handlers, 1) + assert.Contains(t, handlers, "py_etl") + + require.NoError(t, dagProcessor.SendRequest(frame.ID, map[string]any{})) + select { + case err := <-done: + require.NoError(t, err) + case <-time.After(2 * time.Second): + t.Fatal("Serve did not return after the result was acknowledged") + } +} diff --git a/go-sdk/internal/bundle/task.go b/go-sdk/internal/bundle/task.go index 9db3626a1ff1b..1037e07b79c4c 100644 --- a/go-sdk/internal/bundle/task.go +++ b/go-sdk/internal/bundle/task.go @@ -26,6 +26,7 @@ import ( "runtime" "github.com/apache/airflow/go-sdk/pkg/binding" + "github.com/apache/airflow/go-sdk/pkg/execution/genmodels" "github.com/apache/airflow/go-sdk/pkg/sdkcontext" "github.com/apache/airflow/go-sdk/sdk" ) @@ -35,6 +36,9 @@ import ( // and airflow.DagRef.If wrap a plain Go function into a Task. type Task interface { Execute(ctx context.Context, logger *slog.Logger, args []binding.Arg) error + // Declare describes the parameters the task binds, for the Dag processor to check the + // Python stub task against. + Declare(taskID string) genmodels.TaskHandlerDeclaration } // Bundle looks up a registered task by dag_id and task_id. The coordinator @@ -57,6 +61,13 @@ type EnumerableBundle interface { ListTaskHandlers() []TaskHandlerInfo } +// Registry is what a bundle binary serves: a task run looks its handler up, and a task +// handler parse lists the handlers. +type Registry interface { + Bundle + EnumerableBundle +} + type taskFunction struct { fn reflect.Value fullName string @@ -129,6 +140,10 @@ func (f *taskFunction) Execute( return f.call(ctx, sdkClient, reflectArgs, logger, branch) } +func (f *taskFunction) Declare(taskID string) genmodels.TaskHandlerDeclaration { + return f.plan.Declare(taskID) +} + func clientFrom(ctx context.Context) (sdk.Client, error) { client, ok := ctx.Value(sdkcontext.SdkClientContextKey).(sdk.Client) if !ok { diff --git a/go-sdk/pkg/binding/binding.go b/go-sdk/pkg/binding/binding.go index 066f479f1bbee..40be16698c83f 100644 --- a/go-sdk/pkg/binding/binding.go +++ b/go-sdk/pkg/binding/binding.go @@ -23,7 +23,8 @@ // Captured defaults may go unclaimed, and a sole untagged struct can decode one // unclaimed argument as a whole value. // -// Analyze validates a function once. Resolve binds each execution. +// Analyze validates a function once. Resolve binds each execution. Declare describes the +// parameters for the Dag processor's check of a stub task against its handler. // AnalyzePositional builds a plan in which a sole struct binds positionally too, as one whole // argument. package binding diff --git a/go-sdk/pkg/binding/declaration.go b/go-sdk/pkg/binding/declaration.go new file mode 100644 index 0000000000000..67540a6c6054c --- /dev/null +++ b/go-sdk/pkg/binding/declaration.go @@ -0,0 +1,215 @@ +// 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. + +package binding + +import ( + "encoding/json" + "math" + "reflect" + "time" + + "github.com/apache/airflow/go-sdk/pkg/execution/genmodels" +) + +// Declare describes the parameters that the task function's TaskFlow arguments fill, so the Dag +// processor can check a stub task against its Go handler before the task runs. +// +// Flat parameters bind by position and carry no name. A lone struct, tagged or not, binds by +// field name, exactly for an `arg:`-tagged field and ignoring case and underscores otherwise. +// A lone struct with no field to bind, such as time.Time, is declared named with nil params, +// because the runtime takes a lone argument as its whole value and no params list can say so. +func (p *Plan) Declare(taskID string) genmodels.TaskHandlerDeclaration { + decl := genmodels.TaskHandlerDeclaration{ + TaskID: taskID, + Binding: genmodels.TaskHandlerDeclarationBindingPositional, + } + // Not nil, except for a lone struct below: a null list tells the Dag processor that the + // params cannot be listed. + params := genmodels.TaskHandlerParams{} + for _, plan := range p.params { + switch plan.kind { + case paramData: + params = append(params, genmodels.TaskHandlerParam{ValueSchema: valueSchema(plan.typ)}) + case paramLoneStruct: + decl.Binding = genmodels.TaskHandlerDeclarationBindingNamed + if len(plan.fields) == 0 { + // Only the handler's presence is checked. An empty list would warn about a + // lone argument that the runtime accepts. + return decl + } + for _, sf := range plan.fields { + params = append(params, genmodels.TaskHandlerParam{ + Name: sf.argName, + ExactName: sf.tagged, + ValueSchema: valueSchema(sf.fieldType), + }) + } + } + } + decl.Params = ¶ms + return decl +} + +// valueSchema returns the JSON Schema of the values a parameter of type t accepts, or nil when +// no schema can state them. +// +// It states only what decoding into t enforces, so a value the schema rejects is one the task +// would fail on. The exception is a null element or map value, which decoding turns into the +// zero value; like the runtime's own argument check, the schema states null only at the top level. +// The vocabulary is the one build_arg_bindings emits for a stub parameter's Python annotation, +// so the common pairs, such as int and int, come out equal. +func valueSchema(t reflect.Type) *genmodels.ArgValueSchema { + fragment := schemaFragment(t, map[reflect.Type]bool{}) + if fragment == nil { + return nil + } + schema := make(genmodels.ArgValueSchema, len(fragment)) + for k, v := range fragment { + schema[k] = v + } + return &schema +} + +var ( + timeType = reflect.TypeFor[time.Time]() + jsonNumberType = reflect.TypeFor[json.Number]() +) + +func schemaFragment(t reflect.Type, visiting map[reflect.Type]bool) map[string]any { + // A type that contains itself is left unconstrained where it recurs. + if visiting[t] { + return nil + } + visiting[t] = true + defer delete(visiting, t) + + switch t.Kind() { + case reflect.Pointer: + return nullable(schemaFragment(t.Elem(), visiting)) + case reflect.Slice, reflect.Map: + // Decoding null into a slice or map leaves it nil. + return nullable(nonNullFragment(t, visiting)) + default: + return nonNullFragment(t, visiting) + } +} + +func nonNullFragment(t reflect.Type, visiting map[reflect.Type]bool) map[string]any { + switch { + case t == timeType: + return map[string]any{"type": "string", "format": "date-time"} + // json.Number takes a number or a numeric string. + case t == jsonNumberType: + return nil + case implementsJSONUnmarshaler(t): + return nil + case implementsTextUnmarshaler(t): + return map[string]any{"type": "string"} + } + + switch t.Kind() { + case reflect.String: + return map[string]any{"type": "string"} + case reflect.Bool: + return map[string]any{"type": "boolean"} + case reflect.Int, reflect.Int8, reflect.Int16, reflect.Int32, reflect.Int64: + return signedIntFragment(t.Bits()) + case reflect.Uint, + reflect.Uint8, + reflect.Uint16, + reflect.Uint32, + reflect.Uint64, + reflect.Uintptr: + return unsignedIntFragment(t.Bits()) + case reflect.Float32: + return map[string]any{"type": "number", "format": "float"} + case reflect.Float64: + return map[string]any{"type": "number", "format": "double"} + case reflect.Slice: + // A byte slice takes a base64 string as well as an array. + if t.Elem().Kind() == reflect.Uint8 { + return nil + } + return arrayFragment(t.Elem(), visiting) + case reflect.Array: + return arrayFragment(t.Elem(), visiting) + case reflect.Map: + var values any = true + if fragment := schemaFragment(t.Elem(), visiting); fragment != nil { + values = fragment + } + return map[string]any{"type": "object", "additionalProperties": values} + case reflect.Struct: + return map[string]any{"type": "object"} + default: + return nil + } +} + +func signedIntFragment(bits int) map[string]any { + switch bits { + case 64: + return map[string]any{"type": "integer", "format": "int64"} + case 32: + return map[string]any{"type": "integer", "format": "int32"} + default: + return map[string]any{ + "type": "integer", + "minimum": -(int64(1) << (bits - 1)), + "maximum": int64(1)<<(bits-1) - 1, + } + } +} + +func unsignedIntFragment(bits int) map[string]any { + return map[string]any{ + "type": "integer", + "minimum": uint64(0), + "maximum": uint64(math.MaxUint64) >> (64 - bits), + } +} + +func arrayFragment(elem reflect.Type, visiting map[reflect.Type]bool) map[string]any { + items := schemaFragment(elem, visiting) + if items == nil { + items = map[string]any{} + } + return map[string]any{"type": "array", "items": items} +} + +func nullable(fragment map[string]any) map[string]any { + if fragment == nil { + return nil + } + if branches, isUnion := fragment["anyOf"].([]any); isUnion { + for _, branch := range branches { + if b, ok := branch.(map[string]any); ok && b["type"] == "null" { + return fragment + } + } + } + return map[string]any{"anyOf": []any{fragment, map[string]any{"type": "null"}}} +} + +func implementsJSONUnmarshaler(t reflect.Type) bool { + return t.Implements(jsonUnmarshalerType) || reflect.PointerTo(t).Implements(jsonUnmarshalerType) +} + +func implementsTextUnmarshaler(t reflect.Type) bool { + return t.Implements(textUnmarshalerType) || reflect.PointerTo(t).Implements(textUnmarshalerType) +} diff --git a/go-sdk/pkg/binding/declaration_test.go b/go-sdk/pkg/binding/declaration_test.go new file mode 100644 index 0000000000000..3ef0df9a7aad3 --- /dev/null +++ b/go-sdk/pkg/binding/declaration_test.go @@ -0,0 +1,269 @@ +// 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. + +package binding + +import ( + "encoding/json" + "math" + "reflect" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/apache/airflow/go-sdk/internal/contexttest" + "github.com/apache/airflow/go-sdk/pkg/execution/genmodels" +) + +type declRegion string + +type declLevel int + +func (l *declLevel) UnmarshalText([]byte) error { return nil } + +type declCustom struct{ V int } + +func (c *declCustom) UnmarshalJSON([]byte) error { return nil } + +type declTree []declTree + +type declEmbedded struct{ Count int } + +type declTaggedInput struct { + Region string `arg:"region_code"` + Threshold float64 + declEmbedded +} + +type declUntaggedInput struct { + RegionCode string + Labels map[string]string +} + +// declUnbindable has an exported field, but a func cannot receive a task argument. +type declUnbindable struct { + Callback func() +} + +var ( + stringFragment = map[string]any{"type": "string"} + int64Fragment = map[string]any{"type": "integer", "format": "int64"} + doubleFragment = map[string]any{"type": "number", "format": "double"} +) + +func nullOr(fragment map[string]any) map[string]any { + return map[string]any{"anyOf": []any{fragment, map[string]any{"type": "null"}}} +} + +func schemaOf(fragment map[string]any) *genmodels.ArgValueSchema { + schema := make(genmodels.ArgValueSchema, len(fragment)) + for k, v := range fragment { + schema[k] = v + } + return &schema +} + +func TestValueSchema(t *testing.T) { + for name, tc := range map[string]struct { + typ reflect.Type + want map[string]any + }{ + "string": {reflect.TypeFor[string](), stringFragment}, + "named string": {reflect.TypeFor[declRegion](), stringFragment}, + "bool": {reflect.TypeFor[bool](), map[string]any{"type": "boolean"}}, + "int": {reflect.TypeFor[int](), int64Fragment}, + "int64": {reflect.TypeFor[int64](), int64Fragment}, + "int32": { + reflect.TypeFor[int32](), + map[string]any{"type": "integer", "format": "int32"}, + }, + "int8": { + reflect.TypeFor[int8](), + map[string]any{"type": "integer", "minimum": int64(-128), "maximum": int64(127)}, + }, + "int16": { + reflect.TypeFor[int16](), + map[string]any{"type": "integer", "minimum": int64(-32768), "maximum": int64(32767)}, + }, + "uint8": { + reflect.TypeFor[uint8](), + map[string]any{"type": "integer", "minimum": uint64(0), "maximum": uint64(255)}, + }, + "uint32": { + reflect.TypeFor[uint32](), + map[string]any{"type": "integer", "minimum": uint64(0), "maximum": uint64(math.MaxUint32)}, + }, + "uint64": { + reflect.TypeFor[uint64](), + map[string]any{"type": "integer", "minimum": uint64(0), "maximum": uint64(math.MaxUint64)}, + }, + "float32": { + reflect.TypeFor[float32](), + map[string]any{"type": "number", "format": "float"}, + }, + "float64": {reflect.TypeFor[float64](), doubleFragment}, + "time.Time": { + reflect.TypeFor[time.Time](), + map[string]any{"type": "string", "format": "date-time"}, + }, + "pointer to time.Time": { + reflect.TypeFor[*time.Time](), + nullOr(map[string]any{"type": "string", "format": "date-time"}), + }, + "text unmarshaler on the pointer": {reflect.TypeFor[declLevel](), stringFragment}, + "json unmarshaler": {reflect.TypeFor[declCustom](), nil}, + "json.Number": {reflect.TypeFor[json.Number](), nil}, + "byte slice": {reflect.TypeFor[[]byte](), nil}, + "any": {reflect.TypeFor[any](), nil}, + "pointer to any": {reflect.TypeFor[*any](), nil}, + "complex": {reflect.TypeFor[complex128](), nil}, + "slice": { + reflect.TypeFor[[]string](), + nullOr(map[string]any{"type": "array", "items": stringFragment}), + }, + "slice of any": { + reflect.TypeFor[[]any](), + nullOr(map[string]any{"type": "array", "items": map[string]any{}}), + }, + "array": { + reflect.TypeFor[[3]int](), + map[string]any{"type": "array", "items": int64Fragment}, + }, + "map": { + reflect.TypeFor[map[string]int](), + nullOr(map[string]any{"type": "object", "additionalProperties": int64Fragment}), + }, + "map of any": { + reflect.TypeFor[map[string]any](), + nullOr(map[string]any{"type": "object", "additionalProperties": true}), + }, + "struct": {reflect.TypeFor[declUntaggedInput](), map[string]any{"type": "object"}}, + "pointer": {reflect.TypeFor[*int](), nullOr(int64Fragment)}, + "pointer to pointer": {reflect.TypeFor[**int](), nullOr(int64Fragment)}, + "pointer to slice": { + reflect.TypeFor[*[]int](), + nullOr(map[string]any{"type": "array", "items": int64Fragment}), + }, + "recursive": { + reflect.TypeFor[declTree](), + nullOr(map[string]any{"type": "array", "items": map[string]any{}}), + }, + } { + t.Run(name, func(t *testing.T) { + got := valueSchema(tc.typ) + if tc.want == nil { + assert.Nil(t, got) + return + } + assert.Equal(t, schemaOf(tc.want), got) + }) + } +} + +func declare(t *testing.T, fn any) genmodels.TaskHandlerDeclaration { + t.Helper() + plan, err := Analyze(reflect.TypeOf(fn), "testFn") + require.NoError(t, err) + return plan.Declare("task") +} + +func TestDeclareFlatParams(t *testing.T) { + decl := declare(t, func( + actx contexttest.Context, region string, count int, note *string, config declUntaggedInput, + ) error { + return nil + }) + + assert.Equal(t, genmodels.TaskHandlerDeclaration{ + TaskID: "task", + Binding: genmodels.TaskHandlerDeclarationBindingPositional, + Params: &genmodels.TaskHandlerParams{ + {ValueSchema: schemaOf(stringFragment)}, + {ValueSchema: schemaOf(int64Fragment)}, + {ValueSchema: schemaOf(nullOr(stringFragment))}, + {ValueSchema: schemaOf(map[string]any{"type": "object"})}, + }, + }, decl) +} + +func TestDeclareTaggedStruct(t *testing.T) { + decl := declare(t, func(actx contexttest.Context, in declTaggedInput) error { return nil }) + + assert.Equal(t, genmodels.TaskHandlerDeclaration{ + TaskID: "task", + Binding: genmodels.TaskHandlerDeclarationBindingNamed, + Params: &genmodels.TaskHandlerParams{ + {Name: "region_code", ExactName: true, ValueSchema: schemaOf(stringFragment)}, + {Name: "Threshold", ValueSchema: schemaOf(doubleFragment)}, + {Name: "Count", ValueSchema: schemaOf(int64Fragment)}, + }, + }, decl) +} + +func TestDeclareUntaggedStruct(t *testing.T) { + want := genmodels.TaskHandlerDeclaration{ + TaskID: "task", + Binding: genmodels.TaskHandlerDeclarationBindingNamed, + Params: &genmodels.TaskHandlerParams{ + {Name: "RegionCode", ValueSchema: schemaOf(stringFragment)}, + { + Name: "Labels", + ValueSchema: schemaOf(nullOr(map[string]any{ + "type": "object", "additionalProperties": stringFragment, + })), + }, + }, + } + for name, fn := range map[string]any{ + "value": func(actx contexttest.Context, in declUntaggedInput) error { return nil }, + "pointer": func(actx contexttest.Context, in *declUntaggedInput) error { return nil }, + } { + t.Run(name, func(t *testing.T) { + assert.Equal(t, want, declare(t, fn)) + }) + } +} + +// A lone struct no field of which can bind an argument still declares as named, but with no +// params to check: the runtime takes a lone argument as its whole value, which a list cannot name. +func TestDeclareLoneStructWithoutBindableFields(t *testing.T) { + for name, fn := range map[string]any{ + "time.Time": func(actx contexttest.Context, at time.Time) error { return nil }, + "pointer to time.Time": func(actx contexttest.Context, at *time.Time) error { return nil }, + "empty struct": func(actx contexttest.Context, in struct{}) error { return nil }, + "unbindable field": func(actx contexttest.Context, in declUnbindable) error { return nil }, + } { + t.Run(name, func(t *testing.T) { + decl := declare(t, fn) + + assert.Equal(t, genmodels.TaskHandlerDeclarationBindingNamed, decl.Binding) + assert.Nil(t, decl.Params) + }) + } +} + +func TestDeclareContextOnly(t *testing.T) { + decl := declare(t, func(actx contexttest.Context) error { return nil }) + + assert.Equal(t, genmodels.TaskHandlerDeclaration{ + TaskID: "task", + Binding: genmodels.TaskHandlerDeclarationBindingPositional, + Params: &genmodels.TaskHandlerParams{}, + }, decl) +} diff --git a/go-sdk/pkg/execution/comms.go b/go-sdk/pkg/execution/comms.go index fa31e9517551e..75ac5e8e7edff 100644 --- a/go-sdk/pkg/execution/comms.go +++ b/go-sdk/pkg/execution/comms.go @@ -39,10 +39,10 @@ import ( // serialising the full send-then-read round trip behind a single mutex, and // guarantees the response a caller receives matches the request it sent. // -// The supervisor's initial StartupDetails frame arrives unsolicited, before -// any client request is in flight, and is read synchronously via -// ReadMessage; ReadMessage must not be called after the dispatcher has been -// started. +// The supervisor's initial frame, StartupDetails or TaskHandlerParseRequest, +// arrives unsolicited, before any client request is in flight, and is read +// synchronously via ReadMessage; ReadMessage must not be called after the +// dispatcher has been started. type CoordinatorComm struct { reader io.Reader writer io.Writer diff --git a/go-sdk/pkg/execution/integration_test.go b/go-sdk/pkg/execution/integration_test.go index 66d6611f08458..6db3611a93f5a 100644 --- a/go-sdk/pkg/execution/integration_test.go +++ b/go-sdk/pkg/execution/integration_test.go @@ -101,7 +101,29 @@ func (b testBundle) LookupTask(dagID, taskID string) (bundle.Task, bool) { return task, ok } -func buildBundle(t *testing.T, register func(testBundle)) bundle.Bundle { +// ListTaskHandlers lists nothing: a map keeps no registration order, and only a task +// handler parse lists handlers. The parse tests use handlerRegistry instead. +func (b testBundle) ListTaskHandlers() []bundle.TaskHandlerInfo { return nil } + +// handlerRegistry lists its task handlers in the order they were added, as airflow.Bundle does. +type handlerRegistry struct { + testBundle + order []bundle.TaskHandlerInfo +} + +func newHandlerRegistry() *handlerRegistry { return &handlerRegistry{testBundle: testBundle{}} } + +func (r *handlerRegistry) add(dagID, taskID string, fn any) { + if _, ok := r.testBundle[dagID]; !ok { + r.AddDag(dagID) + } + r.testBundle[dagID].AddTaskWithName(taskID, fn) + r.order = append(r.order, bundle.TaskHandlerInfo{DagID: dagID, TaskID: taskID}) +} + +func (r *handlerRegistry) ListTaskHandlers() []bundle.TaskHandlerInfo { return r.order } + +func buildBundle(t *testing.T, register func(testBundle)) bundle.Registry { t.Helper() b := testBundle{} register(b) @@ -996,7 +1018,8 @@ func TestServeFailureAfterConnectClosesComm(t *testing.T) { logsConn := <-logsCh defer logsConn.Close() - // Serve expects StartupDetails as the first frame, so it fails to decode a VariableResult. + // Serve expects StartupDetails or TaskHandlerParseRequest as the first frame, so it fails to + // decode a VariableResult. payload, err := encodeRequest( 0, map[string]any{"type": "VariableResult", "key": "k", "value": "v"}, @@ -1017,3 +1040,114 @@ func TestServeFailureAfterConnectClosesComm(t *testing.T) { _, err = readFrame(commConn) require.Error(t, err) } + +type loadInput struct { + Table string `arg:"table_name"` +} + +// startTaskHandlerParse runs Serve for a registry with two handlers, sends it a +// TaskHandlerParseRequest, and returns the comm connection and the reply frame. +func startTaskHandlerParse(t *testing.T) (net.Conn, IncomingFrame, chan error) { + t.Helper() + commAddr, logsAddr, commCh, logsCh, cleanup := startSupervisor(t) + t.Cleanup(cleanup) + + registry := newHandlerRegistry() + registry.add("etl", "extract", func(actx contexttest.Context, region string, day *int) error { + return nil + }) + registry.add("etl", "load", func(actx contexttest.Context, in loadInput) error { return nil }) + registry.add("etl", "notify", simpleTask) + + done := make(chan error, 1) + go func() { done <- Serve(registry, commAddr, logsAddr) }() + + commConn := <-commCh + t.Cleanup(func() { commConn.Close() }) + logsConn := <-logsCh + t.Cleanup(func() { logsConn.Close() }) + require.NoError(t, commConn.SetDeadline(time.Now().Add(10*time.Second))) + + payload, err := encodeRequest(0, map[string]any{ + "type": "TaskHandlerParseRequest", + "file": "/bundles/go/etl", + "bundle_path": "/bundles/go", + "bundle_name": "go", + }) + require.NoError(t, err) + require.NoError(t, writeFrame(commConn, payload)) + + frame, err := readFrame(commConn) + require.NoError(t, err) + require.True(t, isNilRaw(frame.Err)) + return commConn, frame, done +} + +func TestServeTaskHandlerParseRequestEndToEnd(t *testing.T) { + commConn, frame, done := startTaskHandlerParse(t) + + // Decoded as a map, so the wire's keys, nulls and empty lists are what is compared. + assert.Equal(t, map[string]any{ + "type": "TaskHandlerParsingResult", + "fileloc": "/bundles/go/etl", + "task_handlers": map[string]any{ + "etl": []any{ + map[string]any{ + "task_id": "extract", + "binding": "positional", + "params": []any{ + map[string]any{ + "name": nil, + "value_schema": map[string]any{"type": "string"}, + }, + map[string]any{ + "name": nil, + "value_schema": map[string]any{"anyOf": []any{ + map[string]any{"type": "integer", "format": "int64"}, + map[string]any{"type": "null"}, + }}, + }, + }, + }, + map[string]any{ + "task_id": "load", + "binding": "named", + "params": []any{map[string]any{ + "name": "table_name", + "exact_name": true, + "value_schema": map[string]any{"type": "string"}, + }}, + }, + map[string]any{"task_id": "notify", "binding": "positional", "params": []any{}}, + }, + }, + }, rawToMap(t, frame.Body)) + + // Serve returns only once the parent acknowledges the result. + select { + case err := <-done: + t.Fatalf("Serve returned before the result was acknowledged: %v", err) + case <-time.After(100 * time.Millisecond): + } + require.NoError(t, writeFrame(commConn, encodeResponseFrame(t, frame.ID, nil, nil))) + + select { + case err := <-done: + require.NoError(t, err) + case <-time.After(2 * time.Second): + t.Fatal("Serve did not return after the result was acknowledged") + } +} + +func TestServeTaskHandlerParseFailsWithoutAcknowledgement(t *testing.T) { + commConn, _, done := startTaskHandlerParse(t) + + require.NoError(t, commConn.Close()) + + select { + case err := <-done: + require.ErrorContains(t, err, "sending task handler parse result") + case <-time.After(2 * time.Second): + t.Fatal("Serve did not return after the parent closed the socket") + } +} diff --git a/go-sdk/pkg/execution/messages.go b/go-sdk/pkg/execution/messages.go index 1dd596201258a..77df9add847a6 100644 --- a/go-sdk/pkg/execution/messages.go +++ b/go-sdk/pkg/execution/messages.go @@ -75,6 +75,12 @@ func decodeIncomingBody(raw msgpack.RawMessage) (any, error) { return nil, fmt.Errorf("decoding StartupDetails: %w", err) } return &msg, nil + case genmodels.TypeTaskHandlerParseRequest: + var msg genmodels.TaskHandlerParseRequest + if err := msgpack.Unmarshal(raw, &msg); err != nil { + return nil, fmt.Errorf("decoding TaskHandlerParseRequest: %w", err) + } + return &msg, nil case genmodels.TypeDagFileParseRequest: var msg genmodels.DagFileParseRequest if err := msgpack.Unmarshal(raw, &msg); err != nil { diff --git a/go-sdk/pkg/execution/messages_test.go b/go-sdk/pkg/execution/messages_test.go index 32fe80d35083f..40ee9be806a85 100644 --- a/go-sdk/pkg/execution/messages_test.go +++ b/go-sdk/pkg/execution/messages_test.go @@ -455,6 +455,29 @@ func TestDecodeIncomingBodyDispatch(t *testing.T) { assert.True(t, ok) }) + t.Run("TaskHandlerParseRequest", func(t *testing.T) { + raw := marshalBody(t, map[string]any{ + "type": "TaskHandlerParseRequest", + "file": "/bundles/go/etl", + "bundle_path": "/bundles/go", + "bundle_name": "go", + }) + result, err := decodeIncomingBody(raw) + require.NoError(t, err) + assert.Equal(t, &genmodels.TaskHandlerParseRequest{ + Type: "TaskHandlerParseRequest", + File: "/bundles/go/etl", + BundlePath: "/bundles/go", + BundleName: "go", + }, result) + }) + + t.Run("malformed TaskHandlerParseRequest", func(t *testing.T) { + raw := marshalBody(t, map[string]any{"type": "TaskHandlerParseRequest", "file": 42}) + _, err := decodeIncomingBody(raw) + assert.ErrorContains(t, err, "decoding TaskHandlerParseRequest") + }) + t.Run("ErrorResponse", func(t *testing.T) { raw := marshalBody(t, map[string]any{"type": "ErrorResponse", "error": "GENERIC_ERROR"}) result, err := decodeIncomingBody(raw) diff --git a/go-sdk/pkg/execution/server.go b/go-sdk/pkg/execution/server.go index 0f56e005315ca..2596a53ceb823 100644 --- a/go-sdk/pkg/execution/server.go +++ b/go-sdk/pkg/execution/server.go @@ -20,8 +20,10 @@ // the Airflow supervisor (Python ExecutableCoordinator), the Serve method of // airflow.BundleRef dispatches here. // -// The first inbound frame on the comm socket is a StartupDetails message -// that drives multi-round task execution. +// The first inbound frame on the comm socket selects what the runtime does. A +// StartupDetails message drives multi-round task execution. A +// TaskHandlerParseRequest from the Dag processor gets one TaskHandlerParsingResult +// that declares every registered task handler. // // See go-sdk/adr/0003-coordinator-protocol-msgpack-ipc.md. package execution @@ -48,21 +50,23 @@ import ( const dialTimeout = 30 * time.Second // terminalSendTimeout bounds the write of the final TaskState/SucceedTask -// frame. The supervisor normally drains the comm socket promptly, but a -// half-open connection (the supervisor gone without a clean close) could -// otherwise wedge the runtime on a blocked write; the deadline turns that -// into a fast failure -- and thus a non-zero exit -- instead of a hang. +// frame, and the send and acknowledgement of a task handler parse result. The +// supervisor normally drains the comm socket promptly, but a half-open +// connection (the supervisor gone without a clean close) could otherwise wedge +// the runtime on a blocked write or wait; the deadline turns that into a fast +// failure -- and thus a non-zero exit -- instead of a hang. const terminalSendTimeout = 30 * time.Second // Serve runs the bundle binary in coordinator mode. It dials the supervisor's // comm and logs sockets, installs an slog handler that writes JSON-line // records to the logs connection, and dispatches on the first frame. // -// Serve returns nil on a clean shutdown: the task ran and its terminal -// TaskState/SucceedTask frame was delivered, and the caller should exit 0. A -// non-nil error indicates a protocol-level failure (connection loss, -// malformed frames, unknown first message type) that happens before or -// instead of delivering a terminal frame. +// Serve returns nil on a clean shutdown, and the caller should exit 0: either +// the task ran and its terminal TaskState/SucceedTask frame was delivered, or +// the task handler parse result was delivered and acknowledged. A non-nil error +// indicates a protocol-level failure (connection loss, malformed frames, +// unknown first message type) that happens before or instead of delivering a +// terminal frame. // // Failure-signaling contract: the caller (main) must turn a non-nil error // into a non-zero process exit. The supervisor derives the task's final state @@ -73,7 +77,7 @@ const terminalSendTimeout = 30 * time.Second // fails closed without needing to send a frame; the post-connect paths below // log the reason at Error first so it still reaches the supervisor's log // stream over the already-connected logs socket. -func Serve(b bundle.Bundle, commAddr, logsAddr string) error { +func Serve(b bundle.Registry, commAddr, logsAddr string) error { if commAddr == "" { return fmt.Errorf("missing --comm=host:port argument") } @@ -167,6 +171,18 @@ func Serve(b bundle.Bundle, commAddr, logsAddr string) error { } logger.Debug("Task execution complete") + case *genmodels.TaskHandlerParseRequest: + logger.Debug("Task handler parse mode", "file", msg.File) + result := declareTaskHandlers(b, msg) + sendCtx, cancel := context.WithTimeout(ctx, terminalSendTimeout) + defer cancel() + // Waiting for the acknowledgement means the parent holds the result before this + // process exits and closes the socket under it. + if _, err := comm.Communicate(sendCtx, result); err != nil { + return fmt.Errorf("sending task handler parse result: %w", err) + } + logger.Debug("Task handler parse complete") + default: logger.Error("Unexpected initial message type", "type", fmt.Sprintf("%T", body)) return fmt.Errorf("unexpected initial message type: %T", body) diff --git a/go-sdk/pkg/execution/task_handler_parse.go b/go-sdk/pkg/execution/task_handler_parse.go new file mode 100644 index 0000000000000..1e2d5aca2f5f0 --- /dev/null +++ b/go-sdk/pkg/execution/task_handler_parse.go @@ -0,0 +1,42 @@ +// 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. + +package execution + +import ( + "github.com/apache/airflow/go-sdk/internal/bundle" + "github.com/apache/airflow/go-sdk/pkg/execution/genmodels" +) + +// declareTaskHandlers answers a TaskHandlerParseRequest with the declarations of every +// registered task handler, keyed by Dag id, each Dag's in registration order. +func declareTaskHandlers( + b bundle.Registry, + req *genmodels.TaskHandlerParseRequest, +) genmodels.TaskHandlerParsingResult { + // A nil map would go out as null, which the Dag processor rejects. + handlers := genmodels.TaskHandlers{} + for _, info := range b.ListTaskHandlers() { + // Declare only what a task run could find. + task, ok := b.LookupTask(info.DagID, info.TaskID) + if !ok { + continue + } + handlers[info.DagID] = append(handlers[info.DagID], task.Declare(info.TaskID)) + } + return genmodels.TaskHandlerParsingResult{Fileloc: req.File, TaskHandlers: handlers} +} diff --git a/go-sdk/pkg/execution/task_handler_parse_test.go b/go-sdk/pkg/execution/task_handler_parse_test.go new file mode 100644 index 0000000000000..47237b5a8e719 --- /dev/null +++ b/go-sdk/pkg/execution/task_handler_parse_test.go @@ -0,0 +1,77 @@ +// 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. + +package execution + +import ( + "testing" + + "github.com/stretchr/testify/assert" + + "github.com/apache/airflow/go-sdk/internal/bundle" + "github.com/apache/airflow/go-sdk/internal/contexttest" + "github.com/apache/airflow/go-sdk/pkg/execution/genmodels" +) + +func TestDeclareTaskHandlers(t *testing.T) { + registry := newHandlerRegistry() + registry.add("etl", "transform", func(actx contexttest.Context, region string) error { + return nil + }) + registry.add("report", "render", simpleTask) + registry.add("etl", "extract", simpleTask) + // Listed, but a task run could not find it. + registry.order = append(registry.order, bundle.TaskHandlerInfo{DagID: "etl", TaskID: "ghost"}) + + result := declareTaskHandlers( + registry, + &genmodels.TaskHandlerParseRequest{File: "/bundles/go/etl"}, + ) + + assert.Equal(t, genmodels.TaskHandlerParsingResult{ + Fileloc: "/bundles/go/etl", + TaskHandlers: genmodels.TaskHandlers{ + "etl": { + { + TaskID: "transform", + Binding: genmodels.TaskHandlerDeclarationBindingPositional, + Params: &genmodels.TaskHandlerParams{ + {ValueSchema: &genmodels.ArgValueSchema{"type": "string"}}, + }, + }, + { + TaskID: "extract", + Binding: genmodels.TaskHandlerDeclarationBindingPositional, + Params: &genmodels.TaskHandlerParams{}, + }, + }, + "report": { + { + TaskID: "render", + Binding: genmodels.TaskHandlerDeclarationBindingPositional, + Params: &genmodels.TaskHandlerParams{}, + }, + }, + }, + }, result) +} + +func TestDeclareTaskHandlersWithoutHandlers(t *testing.T) { + result := declareTaskHandlers(newHandlerRegistry(), &genmodels.TaskHandlerParseRequest{}) + + assert.Equal(t, genmodels.TaskHandlers{}, result.TaskHandlers) +} diff --git a/task-sdk/docs/executable-bundle-spec.rst b/task-sdk/docs/executable-bundle-spec.rst index 27e48a0d04fde..601efd2aeca2e 100644 --- a/task-sdk/docs/executable-bundle-spec.rst +++ b/task-sdk/docs/executable-bundle-spec.rst @@ -303,15 +303,18 @@ Bundle files are placed **as-is** in the Dag bundle named by the ``task_handler_bundle_name`` kwarg on the :class:`~airflow.sdk.coordinators.executable.ExecutableCoordinator` entry under ``[sdk] coordinators`` (or, when it is unset, in the task's own Dag -bundle). The scanner walks the Dag bundle **recursively** and considers only -regular files whose **executable bit is set** for the invoking user; files -without the executable bit are skipped without reading their trailer, so the -Dag bundle must preserve it. For each candidate it reads the last 64 bytes and -treats files whose magic matches ``"AFBNDL01"`` as bundles. Matched files are -then SHA-256-verified per the reader algorithm; a mismatch demotes the file -back to "ignored, with an error log." Files without the magic are silently -ignored, so non-bundle files (READMEs, dotfiles) MAY share the Dag bundle -without interfering with the scan. +bundle). The scanner walks the Dag bundle **recursively**, in sorted path +order, and considers only regular files whose **executable bit is set** for the +invoking user; files without the executable bit are skipped without reading +their trailer, so the Dag bundle must preserve it. For each candidate it reads +the last 64 bytes and treats files whose magic matches ``"AFBNDL01"`` as +bundles. Matched files are then SHA-256-verified per the reader algorithm; a +mismatch demotes the file back to "ignored, with an error log." The first +usable match in sorted path order is the one that runs: a verified bundle whose +``dags`` lists the task's Dag and whose ``supervisor_schema_version`` the +coordinator knows. Files without the magic are silently ignored, so non-bundle +files (READMEs, dotfiles) MAY share the Dag bundle without interfering with the +scan. :: diff --git a/task-sdk/docs/lang-sdk-spec.rst b/task-sdk/docs/lang-sdk-spec.rst index a260b7a56fc2c..702bcd890afab 100644 --- a/task-sdk/docs/lang-sdk-spec.rst +++ b/task-sdk/docs/lang-sdk-spec.rst @@ -168,3 +168,29 @@ wrote. How ``fn`` reaches ``Context`` and ``Client`` is each SDK's own choice getters that read the scope opened at step 6 — and all this spec asks is that ``fn`` receives the same pair the SDK built at step 3. The terminal state at step 7 is one of ``SucceedTask``, ``RetryTask``, or ``TaskState``, and it is reported exactly once. + +Task handler parse lifecycle +---------------------------- + +To check a Python Dag's ``@task.stub`` tasks, the Dag processor starts the process that a task of the +Dag would run and asks it for its ``TaskHandler`` registrations: + +.. mermaid:: + + sequenceDiagram + participant D as Dag processor + participant R as SDK process + + R->>R: 1. process start + D->>R: 2. send TaskHandlerParseRequest + R->>R: 3. declare every registered TaskHandler,
keyed by dag_id, in registration order + R->>D: 4. send one TaskHandlerParsingResult + D->>R: 5. reply to it + R->>R: 6. exit + +No ``fn`` runs: the answer comes from the registrations alone, and the tasks of a native ``Dag`` are +never declared. It must depend only on the artifact, never on the request. A declaration states whether +the stub task's arguments bind to the handler's parameters by position or by name, and which values each +parameter accepts. An SDK that cannot list a handler's parameters sends ``null`` for them, and then only +the handler's presence is checked. An artifact that registers no ``TaskHandler`` answers with an empty +mapping. diff --git a/task-sdk/src/airflow/sdk/coordinators/_bundle_metadata.py b/task-sdk/src/airflow/sdk/coordinators/_bundle_metadata.py index b97ed80e8d82e..9987f4f051b2d 100644 --- a/task-sdk/src/airflow/sdk/coordinators/_bundle_metadata.py +++ b/task-sdk/src/airflow/sdk/coordinators/_bundle_metadata.py @@ -46,8 +46,7 @@ def walk_files( Roots are visited in order and each directory's entries sorted, so coordinator selection does not depend on filesystem ordering. - ``JavaCoordinator`` and ``ExecutableCoordinator`` still carry equivalent walks and should move - onto this one. + ``JavaCoordinator`` still carries an equivalent walk and should move onto this one. """ yield from _walk_files(roots, match, set()) diff --git a/task-sdk/src/airflow/sdk/coordinators/executable/coordinator.py b/task-sdk/src/airflow/sdk/coordinators/executable/coordinator.py index 237cc8dbbdf90..96ac0bf9c1c7d 100644 --- a/task-sdk/src/airflow/sdk/coordinators/executable/coordinator.py +++ b/task-sdk/src/airflow/sdk/coordinators/executable/coordinator.py @@ -22,7 +22,6 @@ import hashlib import os import pathlib -import stat import struct from collections import OrderedDict from typing import TYPE_CHECKING, Any, BinaryIO @@ -34,11 +33,12 @@ ResolvedBundle, extract_supervisor_schema_version, parse_metadata_mapping, + walk_files, ) from airflow.sdk.coordinators._subprocess import SubprocessCoordinator if TYPE_CHECKING: - from collections.abc import Iterable, Iterator, Sequence + from collections.abc import Sequence from structlog.typing import FilteringBoundLogger from typing_extensions import Self @@ -253,40 +253,8 @@ def _dag_ids(metadata: dict[str, Any]) -> set[str]: return set(dags.keys()) -def _find_executables(items: Iterable[pathlib.Path]) -> Iterator[pathlib.Path]: - """ - Yield executable regular files under *items*, descending into directories. - - A symlink loop or a directory that hardlinks into one of its ancestors - would otherwise recurse until the interpreter stack is exhausted, so - directories are deduplicated by ``(st_dev, st_ino)`` for the duration - of a single scan. - """ - seen_dirs: set[tuple[int, int]] = set() - yield from _walk_executables(items, seen_dirs) - - -def _walk_executables( - items: Iterable[pathlib.Path], seen_dirs: set[tuple[int, int]] -) -> Iterator[pathlib.Path]: - for item in items: - try: - st = item.stat() - except OSError: - continue - if stat.S_ISDIR(st.st_mode): - key = (st.st_dev, st.st_ino) - if key in seen_dirs: - log.debug("Skipping already-visited directory", path=str(item)) - continue - seen_dirs.add(key) - try: - children = list(item.iterdir()) - except OSError: - continue - yield from _walk_executables(children, seen_dirs) - elif stat.S_ISREG(st.st_mode) and os.access(item, os.X_OK): - yield item +def _is_executable(path: pathlib.Path) -> bool: + return os.access(path, os.X_OK) @attrs.define @@ -295,7 +263,7 @@ class _Bundle(ResolvedBundle): def find(cls, roots: Sequence[pathlib.Path], dag_id: str) -> Self: log.debug("Finding executable bundles recursively", roots=roots) rejected: list[tuple[pathlib.Path, str]] = [] - for p in _find_executables(roots): + for p in walk_files(roots, match=_is_executable): if (metadata := _read_bundle_metadata(p)) is None: continue if dag_id not in _dag_ids(metadata): @@ -340,7 +308,10 @@ class ExecutableCoordinator(SubprocessCoordinator): executable bundles a Python stub Dag delegates task execution to. It must be registered in ``[dag_processor] dag_bundle_config_list``. If unset, the task's own Dag bundle is used. Only files with the executable bit set - are considered. + are considered. The Dag bundle is searched recursively, in sorted path + order, and the first executable bundle that declares the task instance's + Dag runs. That bundle also answers a task handler parse request with the + task handlers it registers. :param task_startup_timeout: Maximum time the coordinator waits for a task process to start, in seconds. The default is 10 seconds. """ @@ -349,3 +320,23 @@ def _build_execute_task_command(self, *, what: TaskInstance) -> tuple[list[str], roots = self._get_scan_roots() bundle = _Bundle.find(roots, what.dag_id) return [str(bundle.path)], bundle.schema_version + + def _find_task_handler_artifact(self, *, bundle_path: pathlib.Path, dag_id: str) -> ResolvedBundle: + """ + Return the executable bundle a task of *dag_id* runs, with symlinks in its path resolved. + + Compute its path in the Dag bundle relative to ``bundle_path.resolve()``, since *bundle_path* + may contain symlinks. A bundle with a supervisor schema version this Task SDK does not know + is skipped like any unusable bundle, so it leads to ``FileNotFoundError``, never ``ValueError``. + """ + return _Bundle.find([bundle_path], dag_id) + + def _build_parse_task_handler_command(self, *, path: pathlib.Path) -> tuple[list[str], str | None]: + # The trailer and binary hash check that a task's bundle gets. + if (metadata := _read_bundle_metadata(path)) is None: + raise ValueError( + f"{path} is not a valid executable bundle: it cannot be read, has no AFBNDL01 " + "trailer, or its binary digest or metadata is invalid" + ) + # Absolute, as for a task, so exec never searches PATH for it. + return [os.fspath(path.resolve())], extract_supervisor_schema_version(metadata) diff --git a/task-sdk/tests/task_sdk/coordinators/executable/test_coordinator.py b/task-sdk/tests/task_sdk/coordinators/executable/test_coordinator.py index 7071c9134bfb9..a31d35cf535fb 100644 --- a/task-sdk/tests/task_sdk/coordinators/executable/test_coordinator.py +++ b/task-sdk/tests/task_sdk/coordinators/executable/test_coordinator.py @@ -32,6 +32,7 @@ from uuid6 import uuid7 from airflow.sdk.api.datamodels._generated import TaskInstance +from airflow.sdk.coordinators._subprocess import TASK_HANDLER_PARSING_SCHEMA_VERSION from airflow.sdk.coordinators.executable.coordinator import ( FOOTER_MAGIC, FOOTER_SIZE, @@ -228,6 +229,17 @@ def test_searches_multiple_roots(self, tmp_path): bundle = _Bundle.find([root_a, root_b], "beta_dag") assert bundle.path == target.resolve() + def test_picks_the_first_bundle_in_sorted_path_order(self, tmp_path, monkeypatch): + _build_bundle(tmp_path / "b_bundle", dag_ids=["tutorial_dag"]) + (tmp_path / "a").mkdir() + expected = _build_bundle(tmp_path / "a" / "nested", dag_ids=["tutorial_dag"]) + _build_bundle(tmp_path / "c_bundle", dag_ids=["tutorial_dag"]) + # Directory order depends on the filesystem, so list every directory in reverse sorted order. + iterdir = Path.iterdir + monkeypatch.setattr(Path, "iterdir", lambda self: iter(sorted(iterdir(self), reverse=True))) + + assert _Bundle.find([tmp_path], "tutorial_dag").path == expected.resolve() + def test_skips_non_bundle_files(self, tmp_path): (tmp_path / "README.md").write_text("not a bundle") _make_executable(tmp_path / "stray_executable") @@ -409,6 +421,137 @@ def test_raises_when_dag_id_not_found(self, tmp_path): coordinator._build_execute_task_command(what=ti) +class TestFindTaskHandlerArtifact: + def test_finds_the_bundle_a_task_of_the_dag_runs(self, tmp_path): + _build_bundle(tmp_path / "a_other", dag_ids=["other_dag"]) + _build_bundle(tmp_path / "b_bundle", dag_ids=["tutorial_dag"]) + _build_bundle(tmp_path / "c_bundle", dag_ids=["tutorial_dag"]) + coordinator = ExecutableCoordinator() + with coordinator._set_scan_roots([tmp_path]): + command, schema_version = coordinator._build_execute_task_command(what=_make_ti()) + + artifact = coordinator._find_task_handler_artifact(bundle_path=tmp_path, dag_id="tutorial_dag") + + assert [str(artifact.path)] == command + assert artifact.schema_version == schema_version + + def test_resolves_a_symlinked_bundle_root(self, tmp_path): + nested = tmp_path / "real" / "team-a" + nested.mkdir(parents=True) + target = _build_bundle(nested / "etl", dag_ids=["tutorial_dag"]) + root = tmp_path / "bundle" + try: + root.symlink_to(tmp_path / "real", target_is_directory=True) + except (OSError, NotImplementedError): + pytest.skip("symlinks not supported on this platform") + + artifact = ExecutableCoordinator()._find_task_handler_artifact( + bundle_path=root, dag_id="tutorial_dag" + ) + + assert artifact.path == target.resolve() + assert artifact.path.relative_to(root.resolve()) == Path("team-a", "etl") + + @pytest.mark.parametrize( + ("dag_ids", "schema_version", "message"), + [ + pytest.param( + ["other_dag"], + "2026-06-16", + "cannot find executable bundle containing", + id="no-bundle-lists-the-dag", + ), + pytest.param( + ["tutorial_dag"], + "1999-01-01", + "matching bundles were rejected", + id="unknown-schema-version", + ), + ], + ) + def test_raises_file_not_found_when_no_bundle_can_run_the_dag( + self, tmp_path, dag_ids, schema_version, message + ): + metadata = _make_metadata(dag_ids) + metadata["sdk"]["supervisor_schema_version"] = schema_version + _build_bundle(tmp_path / "etl", metadata=metadata) + + with pytest.raises(FileNotFoundError, match=message): + ExecutableCoordinator()._find_task_handler_artifact(bundle_path=tmp_path, dag_id="tutorial_dag") + + +class TestBuildParseTaskHandlerCommand: + def test_returns_the_bundle_and_its_schema_version(self, tmp_path): + # The Dag processor names the artifact, so the Dag ids in its metadata play no part. + bundle = _build_bundle(tmp_path / "etl", dag_ids=["other_dag"]) + + command, schema_version = ExecutableCoordinator()._build_parse_task_handler_command(path=bundle) + + assert command == [str(bundle.resolve())] + assert schema_version == "2026-06-16" + + def test_returns_an_absolute_path(self, tmp_path, monkeypatch): + _build_bundle(tmp_path / "etl") + monkeypatch.chdir(tmp_path) + + command, _ = ExecutableCoordinator()._build_parse_task_handler_command(path=Path("etl")) + + assert command == [str((tmp_path / "etl").resolve())] + + @pytest.mark.parametrize( + "build", + [ + pytest.param(_make_executable, id="no-trailer"), + pytest.param( + lambda path: _build_bundle(path, binary_sha256=b"\x00" * 32), id="binary-digest-mismatch" + ), + ], + ) + def test_rejects_a_file_that_is_not_a_valid_bundle(self, tmp_path, build): + bundle = build(tmp_path / "etl") + + with pytest.raises(ValueError, match="is not a valid executable bundle"): + ExecutableCoordinator()._build_parse_task_handler_command(path=bundle) + + def test_rejects_a_bundle_without_a_schema_version(self, tmp_path): + metadata = _make_metadata(["etl"]) + del metadata["sdk"]["supervisor_schema_version"] + bundle = _build_bundle(tmp_path / "etl", metadata=metadata) + + with pytest.raises(ValueError, match="supervisor_schema_version"): + ExecutableCoordinator()._build_parse_task_handler_command(path=bundle) + + @patch("airflow.sdk.coordinators._subprocess._set_close_on_exec_above_stderr", autospec=True) + @patch("airflow.sdk.coordinators._subprocess._set_parent_death_signal", autospec=True) + @patch("airflow.sdk.coordinators._subprocess.signal.signal", autospec=True) + @patch( + "airflow.sdk.coordinators._subprocess.os.execvpe", autospec=True, side_effect=OSError("exec failed") + ) + def test_the_probe_execs_what_a_task_of_the_dag_runs( + self, mock_execvpe, mock_signal, mock_death_signal, mock_close_on_exec, tmp_path + ): + metadata = _make_metadata(["tutorial_dag"]) + metadata["sdk"]["supervisor_schema_version"] = TASK_HANDLER_PARSING_SCHEMA_VERSION + _build_bundle(tmp_path / "etl", metadata=metadata) + coordinator = ExecutableCoordinator() + with coordinator._set_scan_roots([tmp_path]): + command, _ = coordinator._build_execute_task_command(what=_make_ti()) + artifact = coordinator._find_task_handler_artifact(bundle_path=tmp_path, dag_id="tutorial_dag") + + with pytest.raises(OSError, match="exec failed"): + coordinator.parse_task_handler( + path=artifact.path, + bundle_path=tmp_path, + comm_address=("127.0.0.1", 1001), + logs_address=("127.0.0.1", 1002), + report_schema_version=lambda schema_version: None, + ) + + mock_execvpe.assert_called_once_with( + command[0], [*command, "--comm=127.0.0.1:1001", "--logs=127.0.0.1:1002"], mock.ANY + ) + + @pytest.fixture def bundles_dir(tmp_path): _build_bundle(tmp_path / "my_bundle", dag_ids=["tutorial_dag"])