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 27ad40ab675f3..c3d385cc5e00e 100644 --- a/airflow-core/docs/authoring-and-scheduling/language-sdks/java.rst +++ b/airflow-core/docs/authoring-and-scheduling/language-sdks/java.rst @@ -179,24 +179,26 @@ Save the task implementations as import org.apache.airflow.sdk.Builder; - @Builder.Dag(id = "sales_pipeline") public class SalesPipeline { - @Builder.Task(id = "extract") + @Builder.TaskHandler(dag = "sales_pipeline", task = "extract") public long extract() { return 3; } - @Builder.Task(id = "transform") - public long transform(@Builder.XCom(task = "extract") long recordCount) { + @Builder.TaskHandler(dag = "sales_pipeline", task = "transform") + public long transform(long recordCount) { return recordCount * 2; } } .. note:: - See how both ``transform`` in Python and Java need to have an argument to accept upstream XCom. The - Python one is needed to declare dependency, and the Java one is needed to actually retrieve the value. + The graph is declared once, in the Python Dag file: ``transform(extract())`` feeds the upstream's + return value into the downstream's parameter by calling tasks like functions. The supervisor sends + the resulting *argument bindings* to the Java runtime, and each Java data parameter receives + whatever the Python call site bound at its position — an upstream task's XCom or an inline + literal. See :ref:`java-sdk/arg-binding`. Add the Java entry point ~~~~~~~~~~~~~~~~~~~~~~~~ @@ -207,24 +209,17 @@ Save the entry point as ``src/main/java/com/mycompany/airflow/sales/Main.java``: package com.mycompany.airflow.sales; - import java.util.List; - - import org.apache.airflow.sdk.BundleBuilder; - import org.apache.airflow.sdk.DagDef; + import org.apache.airflow.sdk.Bundle; import org.apache.airflow.sdk.Server; - public class Main implements BundleBuilder { - @Override - public Iterable getDags() { - return List.of(SalesPipelineBuilder.build()); - } - + public class Main { public static void main(String[] args) { - Server.create(args).serve(new Main().build()); + Server.create(args).serve(new Bundle().register(SalesPipeline.class)); } } -``SalesPipelineBuilder`` is generated by the annotation processor during compilation. +``register`` takes the class the annotations are on, so there is no second name to keep in sync: the +annotation processor generates the registrar it reads during compilation. Build and deploy ~~~~~~~~~~~~~~~~ @@ -306,35 +301,38 @@ Annotate a plain Java class and let the SDK generate the boilerplate at compile * - Annotation - Purpose - * - ``@Builder.Dag(id = "...")`` - - Marks the class as a task container. The ``id`` must match the ``dag_id`` in the Python Dag. - * - ``@Builder.Task(id = "...")`` - - Marks a method as a task implementation. The ``id`` must match the ``@task.stub`` function - name in the Python Dag. If ``id`` is omitted the method name is used. - * - ``@Builder.XCom(task = "...")`` - - Injects the ``return_value`` XCom from the named upstream task as a method parameter. - The parameter type must be compatible with the stored value (see :ref:`java-sdk/types`). + * - ``@Builder.TaskHandler(dag = "...", task = "...")`` + - Marks a method as the Java body of a task the Python Dag file declares with ``@task.stub``. + ``dag`` must match the ``dag_id`` and ``task`` the stub function name; omitting ``task`` + derives it from the method name. There is no class-level annotation on this surface — the + Dag is the Python file's, so the handler names the pair it binds to. + * - ``TaskInput`` / ``@ArgName("...")`` + - Marks a class as a task's input, so keyword arguments bind by name instead of by position: + each public field receives the argument whose name matches it, ignoring case and + underscores. ``@ArgName`` pins a name the match cannot reach, or renames the argument + deliberately. See :ref:`java-sdk/arg-binding`. + +Besides the annotations, a task method may declare a ``Client`` and a ``Context`` parameter in any +position; the SDK injects both. Every other parameter is a *data parameter* and receives an +argument bound by the Python ``@task.stub`` call site. The annotation processor generates a ``Builder`` class that wires up the task -registry and handles XCom injection automatically. +registry and resolves data parameters and XCom pushes automatically. .. code-block:: java - @Builder.Dag(id = "my_dag") + public class MyDag { - @Builder.Task(id = "fetch") + @Builder.TaskHandler(dag = "my_dag", task = "fetch") public String fetch(Client client) throws Exception { var conn = client.getConnection("my_api"); // implement task logic return result; } - @Builder.Task(id = "process") - public long process( - Client client, - @Builder.XCom(task = "fetch") String fetched - ) { + @Builder.TaskHandler(dag = "my_dag", task = "process") + public long process(Client client, String fetched) { var threshold = (String) client.getVariable("process_threshold"); // implement task logic return count; @@ -376,41 +374,164 @@ task log. } } -Register tasks manually in a ``BundleBuilder``. A task class can be top-level like ``FetchTask``, or -nested ``static`` class like ``ProcessTask``: +Implement ``InputTask`` instead when the Python Dag calls the stub with TaskFlow arguments: the SDK +resolves them from the call site and passes them in. The type argument is a ``TaskInput`` whose public +fields declare what the task expects. See :ref:`java-sdk/arg-binding`. + +Register each task against the Dag and task the Python file declares. ``register`` is one +overloaded verb: these ids, or the class an annotated handler lives on. .. code-block:: java - public class MyBundle implements BundleBuilder { - public static class ProcessTask implements Task { + public class MyBundle { + public static class ProcessInput implements TaskInput { + public String fetched; + } + + public static class ProcessTask implements InputTask { @Override - public void execute(Context context, Client client) throws Exception { - var fetched = (String) client.getXCom("fetch"); + public void execute(Context context, Client client, ProcessInput input) throws Exception { // implement task logic - client.setXCom(fetched); + client.setXCom(input.fetched); } } - @Override - public Iterable getDags() { - var dag = new DagDef("my_dag") - .addTask("fetch", FetchTask.class) - .addTask("process", ProcessTask.class); - return List.of(dag); - } - public static void main(String[] args) { - Server.create(args).serve(new MyBundle().build()); + var bundle = new Bundle() + .register("my_dag", "fetch", FetchTask.class) + .register("my_dag", "process", ProcessTask.class); + Server.create(args).serve(bundle); } } -Place the task classes and ``BundleBuilder`` under the standard ``src/main/java//`` source tree. -The ``BundleBuilder`` can provide the ``main`` method itself, as above, or a separate entry-point class can -call it. Set ``airflowBundle.mainClass`` to the class that provides ``main``. From that point onward, both APIs +A task class can be top-level like ``FetchTask``, or a nested ``static`` class like ``ProcessTask``. +Place them under the standard ``src/main/java//`` source tree, and set +``airflowBundle.mainClass`` to the class that provides ``main``. From that point onward, both APIs use the same ``./gradlew bundle`` command and deploy the resulting ``build/bundle/`` directory in the same way. See the `Java SDK API Reference `__ for more details. +.. _java-sdk/arg-binding: + +Binding stub arguments +~~~~~~~~~~~~~~~~~~~~~~ + +Calling a ``@task.stub`` TaskFlow-style in the Python Dag is what declares the graph, and the +supervisor delivers the resulting argument bindings to the Java runtime with every task run. A +binding carries either an upstream task's ``return_value`` XCom or an inline literal written at the +call site. + +Positional binding +^^^^^^^^^^^^^^^^^^ + +A task method's data parameters bind **by position**, in declaration order — the injected ``Client`` +and ``Context`` parameters do not take up a position. Java parameter names are not part of the API, +so renaming one in an IDE never rebinds an input. + +.. code-block:: python + + @task.stub(queue="java") + def score(rows, threshold): ... + + + score(load_rows(), 0.75) + +.. code-block:: java + + @Builder.TaskHandler(dag = "etl", task = "score") + public long score(Client client, long rows, double threshold) { + // rows <- the load_rows XCom (position 0) + // threshold <- the literal 0.75 (position 1) + } + +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. + +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 +integers. Without that the list would hold ``Long`` values while claiming to hold ``Double``, and +the ``ClassCastException`` would land on whichever line first read an element rather than on the +binding that got it wrong. + +Named binding with a ``TaskInput`` +^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^ + +To bind keyword arguments by name, declare a class implementing ``TaskInput``. Each of its public +non-final fields receives the argument whose name matches it, **ignoring case and underscores**, so +the stub's ``snake_case`` arguments reach ``camelCase`` Java fields with nothing declared. That is +the same fold the Go and TypeScript SDKs apply, so one Python signature binds identically in every +SDK. + +The class needs a public no-argument constructor. A name mismatch in either direction is logged +rather than failed, because a field binds by name: a field nothing supplies keeps its Java default, +and an argument no field claims changes nothing the task reads. Once a field has claimed its +argument, an argument that resolves to nothing is a value and not a mistake, so a boxed or reference +field takes ``null`` and a primitive field fails. + +No two fields may claim argument names that differ only in case or underscores; that fails when the +bundle is built, because the fold cannot tell them apart. Two *arguments* that collide that way +reach only a field naming one of them exactly, and a field that would match both is left unfilled +and reported rather than handed the wrong value. + +.. code-block:: python + + @task.stub(queue="java") + def score(region_code, threshold): ... + + + score(region_code="emea", threshold=load_threshold()) + +.. code-block:: java + + public static class ScoreInput implements TaskInput { + // Pinned so the field can be called region. Or drop the annotation and + // write: public String regionCode; + @ArgName("region_code") + public String region; + + public double threshold; // binds threshold + } + + @Builder.TaskHandler(dag = "etl", task = "score") + public long score(Client client, ScoreInput input) { ... } + +Reach for ``@ArgName`` when the argument name is not a legal or usable Java identifier — a Python +keyword such as ``class``, say — or when the field should read differently from the argument. A +pinned name is matched as written, with no folding. + +A task method declares flat data parameters **or** one ``TaskInput``, never both, so field names and +flat positions cannot shift each other. Mixing them, or declaring two, fails the build. + +Binding in the interface-based API +^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^ + +A task written against the interface has no parameter list to bind, so it declares its input as the +type argument of ``InputTask`` instead — the same ``TaskInput`` an annotated task would declare, +binding the same way: + +.. code-block:: java + + public class ScoreTask implements InputTask { + @Override + public void execute(Context context, Client client, ScoreInput input) throws Exception { + // input.region, input.threshold + } + } + +A ``TaskInput`` is the only way an interface task receives bound values; there is no positional form. +A positional read is safe in the code the annotation processor writes, because it type-checks each +position against the method signature it serves, and is the wrong thing to ask of code a person +writes and later reads. A call site with a single argument is worth the one-field class. + +An ``InputTask`` whose type argument is not a concrete ``TaskInput`` fails when the bundle is built, +rather than mid-run. Plain ``Task`` remains the right interface for a task the Dag file +calls with no arguments. + .. _java-sdk/logging: Logging @@ -428,7 +549,7 @@ is the conventional pattern regardless of which logging framework you choose: private static final System.Logger log = System.getLogger(SalesPipeline.class.getName()); - @Builder.Task(id = "extract") + @Builder.TaskHandler(dag = "etl", task = "extract") public long extract(Client client) { log.log(System.Logger.Level.INFO, "Starting extraction"); return recordCount; @@ -605,15 +726,16 @@ represented as Java objects when read back via ``getXCom``. .. note:: - ``char`` and ``Character`` are not supported. JSON has no single-character type, so a - character value is stored as a JSON string (or a number) and is read back as one of the - Java types in the table above. Declaring ``char`` or ``Character`` as an - ``@Builder.XCom`` parameter compiles, but reading a pushed value fails at runtime with a - ``ClassCastException``. Use ``String`` instead. + Avoid ``char`` and ``Character``. JSON has no single-character type, so the value arrives + as a string or a number and the SDK narrows it: a one-character string or an integer code + point binds, an empty string binds as ``null``, and anything longer fails at runtime with + an ``IllegalArgumentException``. Whether a value binds therefore depends on its length + rather than on the stub signature, which no compile-time check can catch. Use ``String`` + instead. .. note:: - An ``@Builder.XCom`` parameter that reads a value which was never pushed resolves to + A data parameter whose binding resolves to a value that was never pushed receives ``null``. A boxed parameter (``Integer``, ``Long``, ``Boolean``, …) receives ``null`` safely, but a primitive parameter (``int``, ``long``, ``boolean``, …) cannot represent ``null`` and the task fails with ``MissingXComException``. Declare the parameter with a @@ -758,7 +880,7 @@ simplest deployment: one file, no dependency management at runtime. - + com.example.Main diff --git a/airflow-e2e-tests/tests/airflow_e2e_tests/java_sdk_tests/test_java_sdk_dag.py b/airflow-e2e-tests/tests/airflow_e2e_tests/java_sdk_tests/test_java_sdk_dag.py index 07180d2c8a0ff..60ddf7a8d284d 100644 --- a/airflow-e2e-tests/tests/airflow_e2e_tests/java_sdk_tests/test_java_sdk_dag.py +++ b/airflow-e2e-tests/tests/airflow_e2e_tests/java_sdk_tests/test_java_sdk_dag.py @@ -54,6 +54,15 @@ 4. A Java task that throws with retries left returns ``RetryTask`` rather than a terminal ``FAILED``, so the supervisor marks it UP_FOR_RETRY and re-runs it; ``load`` therefore ends ``success`` on its second attempt (try_number 2). + +5. Keyword arguments at the stub call site bind to a ``TaskInput``'s fields by + name, matched ignoring case and underscores; ``report`` asserts the value it + received, so it only succeeds if ``run_label`` reached its ``runLabel`` + field. + +The ``java_interface_example`` and ``scala_spark_example`` Dags are covered too: +the former for tasks written against ``InputTask``, whose arguments the SDK +resolves from the stub call site into a ``TaskInput`` and injects. """ from __future__ import annotations @@ -81,6 +90,7 @@ _LOG_FETCH_TIMEOUT = 60 _ANNOTATION_DAG_ID = "java_annotation_example" +_INTERFACE_DAG_ID = "java_interface_example" _XCOM_CASTING_DAG_ID = "java_xcom_casting_example" _SCALA_SPARK_DAG_ID = "scala_spark_example" _VARIABLE_WRITE_DAG_ID = "java_variable_write" @@ -171,6 +181,12 @@ def annotation_example_run() -> _CompletedRun: return _trigger_and_wait_for_dag(_ANNOTATION_DAG_ID, _JAVA_TASK_TIMEOUT) +@pytest.fixture(scope="module") +def interface_example_run() -> _CompletedRun: + """Trigger the interface example once for all of its assertions.""" + return _trigger_and_wait_for_dag(_INTERFACE_DAG_ID, _JAVA_TASK_TIMEOUT) + + @pytest.fixture(scope="module") def xcom_casting_example_run() -> _CompletedRun: """Trigger the XCom casting example once for all of its assertions.""" @@ -213,6 +229,19 @@ def test_transform_xcom_is_numeric_timestamp(self, annotation_example_run: _Comp f"Expected 'transform' XCom to be a positive integer (millisecond timestamp), got {value!r}" ) + def test_report_binds_keyword_arguments_by_folded_name(self, annotation_example_run: _CompletedRun): + """``report`` declares a ``TaskInput`` with no ``@ArgName``, so reaching ``success`` + proves ``run_label`` folded onto its ``runLabel`` field: the task throws unless that + field holds the ``run_label="nightly"`` literal from the keyword call site.""" + report_ti = annotation_example_run.get_task_instance("report") + + assert report_ti.get("state") == "success", ( + f"Java 'report' task did not succeed.\n" + f" task state : {report_ti.get('state')!r}\n" + f" dag state : {annotation_example_run.state!r}\n" + f" all tasks : {annotation_example_run.ti_states}" + ) + def test_concurrent_client_calls_succeed(self, annotation_example_run: _CompletedRun): """A Java task calling the client from many threads must succeed.""" concurrent_ti = annotation_example_run.get_task_instance("concurrent") @@ -377,13 +406,32 @@ def test_scratch_variable_deleted_by_java_task_is_gone(self, variable_write_run: ) +class TestJavaSDKInterfaceExample: + """Verify tasks written against ``InputTask`` receive their bound arguments.""" + + def test_input_tasks_receive_their_bound_arguments(self, interface_example_run: _CompletedRun): + """``transform`` and ``summarize`` implement ``InputTask``, so reaching ``success`` + proves the supervisor's bindings arrived: ``transform`` binds the ``extract`` XCom + onto its bundle's ``extracted`` field, and ``summarize`` throws unless its bundle + holds the ``region_code="emea"`` literal from the keyword call site.""" + for task_id in ("transform", "summarize"): + task_instance = interface_example_run.get_task_instance(task_id) + assert task_instance.get("state") == "success", ( + f"Java {task_id!r} task did not succeed.\n" + f" task state : {task_instance.get('state')!r}\n" + f" dag state : {interface_example_run.state!r}\n" + f" all tasks : {interface_example_run.ti_states}" + ) + + class TestJavaSDKXComCastingExample: """Verify numeric XCom values are cast across declared Java types.""" def test_numeric_xcom_casting(self, xcom_casting_example_run: _CompletedRun): """Numeric XComs are read across declared types (int -> long -> double, and a wire - double back as a float), and a boxed param stays null when its XCom is absent.""" - for task_id in ("widen_to_double", "consume_nullable", "consume_float"): + double back as a float), a boxed param stays null when its XCom is absent, and a + ``List`` param keeps its element type over a literal list of integers.""" + for task_id in ("widen_to_double", "consume_nullable", "consume_float", "consume_double_list"): task_instance = xcom_casting_example_run.get_task_instance(task_id) assert task_instance.get("state") == "success", ( f"Java {task_id!r} task did not succeed.\n" diff --git a/java-sdk/README.md b/java-sdk/README.md index 027babb6eedf3..97b81903fb1d3 100644 --- a/java-sdk/README.md +++ b/java-sdk/README.md @@ -590,7 +590,7 @@ prek hook regenerate it. -*Min. Airflow version: 3.3 · supervisor schema: 2026-06-16* +*Min. Airflow version: 3.3 · supervisor schema: 2026-10-30* | Dimension | Tier | Supported | Since | Notes | |---|---|---|---|---| diff --git a/java-sdk/adr/0001-mixed-lang-dag-interface.md b/java-sdk/adr/0001-mixed-lang-dag-interface.md index 305f9c7df6ac5..9c500de206648 100644 --- a/java-sdk/adr/0001-mixed-lang-dag-interface.md +++ b/java-sdk/adr/0001-mixed-lang-dag-interface.md @@ -61,7 +61,7 @@ The processor compiles this into a `Task`: public static final class Score implements Task { @Override public void execute(Context context, Client client) throws Exception { - TaskArgs args = TaskArgs.of(context); + TaskArgs args = TaskArgs.of(context, client, 3); long rows = args.require(0, Long.class); double threshold = args.require(1, Double.class); List regions = args.require(2, new TypeReference>() {}); @@ -103,7 +103,7 @@ public static class SummarizeInput implements TaskInput { } // summarize(region_code=..., transformed=...) is called with keyword - // arguments, which bind to the bundle's fields by name. + // arguments, which bind to the `TaskInput`'s fields by name. public static class Summarize implements InputTask { public void execute(@NotNull Context context, Client client, SummarizeInput input) { log.log( @@ -155,11 +155,17 @@ One verb, overloaded, and everything registers as a class: nothing is constructe - Annotated task data parameters bind by position; injected `Client` and `Context` do not consume positions. Positional binding also needs no parameter names in the class file, so it does not depend on compiling with `-parameters`. -- Keyword arguments bind through one `TaskInput` bundle, with `@ArgName` only when Python and Java +- Keyword arguments bind through one `TaskInput`, with `@ArgName` only when Python and Java names differ. - Interface tasks use `InputTask` for named fields. There is no public positional form. - `TaskArgs` is internal. It is what generated code reads, and what every syntax above resolves to; - nothing in the public API returns or accepts it. It exposes `require(index, type)`, which fails + nothing in the public API returns or accepts it. It is opened with + `of(context, client, declared)` — the client because resolving a binding may have to pull an + upstream's XCom, and the declared count because positions carry the whole meaning of a flat + binding, so a call site that bound a different number of arguments than the method takes has + already shifted them. A parameter the call site omitted is exempt: its default still + arrives, but a method that does not declare it is not reading shifted arguments, so those are + dropped before the counts are compared. It exposes `require(index, type)`, which fails when the position resolves to nothing, and `get(index, type)`, which yields `null` — each with a `Class` and a `TypeReference` form. - Annotation processing generates code, it never rewrites it, so the method a user writes stays @@ -169,7 +175,24 @@ One verb, overloaded, and everything registers as a class: nothing is constructe the generated body calls `new Parent().method(...)`. - A parameter whose type has type arguments binds through the `TypeReference` overload, so the element type survives and the generated code never contains an unchecked cast. -- Missing reference or boxed values become `null`; primitive inputs fail clearly. +- **Neither direction of a `TaskInput` name mismatch fails the task.** A field binds by name, so a + field nothing supplies keeps its Java default and an argument no field claims changes nothing the + task reads. Both are logged, since either one means the Java signature and the Python stub + disagree. Captured defaults are not logged, since the Dag author never passed them. Flat + positional binding still fails either way, because a dropped or added argument shifts every later + one. +- **Two fields whose argument names fold alike fail when the bundle is built**, because the fold + cannot tell them apart and there is no value either could safely take. Two *arguments* that fold + alike are the other way round: they reach a field naming one of them exactly, and a field that + would match both is left unfilled and reported. +- **An argument that resolves to nothing is a value, not a mistake.** Once a parameter has claimed + its argument, a null literal or an upstream that pushed no XCom gives `null` to a reference or + boxed type, and fails for a primitive, which cannot hold it. Declaring a boxed type is how a + task says the value is optional. +- This matches the Go SDK, which logs both directions and zero-values an unmatched struct field + ([go-sdk ADR-0006](../../go-sdk/adr/0006-cross-language-argument-binding.md)), and it follows the + cross-language rule in + [lang-SDK ADR-0007](../../airflow-core/adr/lang-sdk/0007-taskflow-across-language-boundary.md). - `@Builder.XCom` is removed, keeping the Python call site as the single source of data-flow wiring. - **The generated `Task` is one more branch of an existing classifier.** `BuilderProcessor` already diff --git a/java-sdk/adr/0002-native-dag-interface.md b/java-sdk/adr/0002-native-dag-interface.md index 80c2d7101cdbd..b3ad667e57823 100644 --- a/java-sdk/adr/0002-native-dag-interface.md +++ b/java-sdk/adr/0002-native-dag-interface.md @@ -139,7 +139,7 @@ public final class EtlPipeline_Dag { public static final class Transform implements Task { @Override public void execute(Context context, Client client) throws Exception { - TaskArgs args = TaskArgs.of(context); + TaskArgs args = TaskArgs.of(context, client, 1); long extracted = args.require(0, Long.class); double threshold = 0.9; // baked from lit(0.9) at Dag-build time client.setXCom(new EtlPipeline().transform(client, context, extracted, threshold)); diff --git a/java-sdk/capabilities.yaml b/java-sdk/capabilities.yaml index 45724e433bda1..252f52f0faf98 100644 --- a/java-sdk/capabilities.yaml +++ b/java-sdk/capabilities.yaml @@ -21,7 +21,7 @@ min_airflow_version: "3.3" # Keep in sync with airflowSupervisorSchemaVersion in gradle.properties, which is what stamps the # JAR manifest. The render hook fails if the two disagree. -supervisor_schema_version: "2026-06-16" +supervisor_schema_version: "2026-10-30" # The runtime terminates a task with SucceedTask, RetryTask, or TaskState (failed/removed); it does # not yet emit skipped, DeferTask, RescheduleTask, or AwaitInputTask. diff --git a/java-sdk/example/src/java/org/apache/airflow/example/AnnotationExample.java b/java-sdk/example/src/java/org/apache/airflow/example/AnnotationExample.java index bb715a73cb502..c7f08be99a615 100644 --- a/java-sdk/example/src/java/org/apache/airflow/example/AnnotationExample.java +++ b/java-sdk/example/src/java/org/apache/airflow/example/AnnotationExample.java @@ -27,12 +27,14 @@ import java.util.concurrent.Executors; import org.apache.airflow.sdk.*; +// The Python Dag file (src/resources/dags/java_examples.py) owns the graph: each +// data parameter below receives whatever the @task.stub call site bound at its +// position, so this class registers task implementations only. @SuppressWarnings("DuplicatedCode") -@Builder.Dag(id = "java_annotation_example") public class AnnotationExample { private static final System.Logger log = System.getLogger(AnnotationExample.class.getName()); - @Builder.Task(id = "extract") + @Builder.TaskHandler(dag = "java_annotation_example", task = "extract") public long extractValue(Client client) throws InterruptedException { log.log(INFO, "Hello from task"); @@ -51,8 +53,8 @@ public long extractValue(Client client) throws InterruptedException { return new Date().getTime(); } - @Builder.Task(id = "transform") - public long transformValue(Client client, @Builder.XCom(task = "extract") long extracted) { + @Builder.TaskHandler(dag = "java_annotation_example", task = "transform") + public long transformValue(Client client, long extracted) { log.log(INFO, "Got XCom from extract: {0}", extracted); var variable = client.getVariable("my_variable"); @@ -67,8 +69,8 @@ public long transformValue(Client client, @Builder.XCom(task = "extract") long e // the task UP_FOR_RETRY -- which only works because the Java SDK now returns // RetryTask (instead of a terminal FAILED) when ti_context.should_retry is // set. The retry then runs this task again and it returns normally. - @Builder.Task - public void load(Context context, @Builder.XCom(task = "transform") long transformed) { + @Builder.TaskHandler(dag = "java_annotation_example") + public void load(Context context, long transformed) { log.log(INFO, "Got XCom from transform: {0}", transformed); if (context.ti.tryNumber == 1) { throw new RuntimeException("I failed"); @@ -76,8 +78,24 @@ public void load(Context context, @Builder.XCom(task = "transform") long transfo log.log(INFO, "Recovered on retry, try number {0}", context.ti.tryNumber); } + // Keyword arguments bind by name instead of by position, through a TaskInput. + // Field names match ignoring case and underscores, so the stub's snake_case + // arguments reach camelCase Java fields on their own. + public static class ReportInput implements TaskInput { + public String runLabel; // binds run_label through the fold + public long transformed; + } + + @Builder.TaskHandler(dag = "java_annotation_example", task = "report") + public void report(ReportInput input) { + log.log(INFO, "Report {0} for transformed value {1}", input.runLabel, input.transformed); + if (!"nightly".equals(input.runLabel)) { + throw new RuntimeException("expected run label 'nightly' but got " + input.runLabel); + } + } + // Verify one supervisor channel can handle client calls across threads. - @Builder.Task(id = "concurrent") + @Builder.TaskHandler(dag = "java_annotation_example", task = "concurrent") public void concurrentClientCalls(Client client) throws Exception { var pool = Executors.newFixedThreadPool(8); try { diff --git a/java-sdk/example/src/java/org/apache/airflow/example/ExampleBundleBuilder.java b/java-sdk/example/src/java/org/apache/airflow/example/ExampleBundleBuilder.java index fa1a860755863..c8e0ac75f4e0d 100644 --- a/java-sdk/example/src/java/org/apache/airflow/example/ExampleBundleBuilder.java +++ b/java-sdk/example/src/java/org/apache/airflow/example/ExampleBundleBuilder.java @@ -19,22 +19,17 @@ package org.apache.airflow.example; -import java.util.List; import org.apache.airflow.sdk.*; -import org.jetbrains.annotations.NotNull; -public class ExampleBundleBuilder implements BundleBuilder { - @NotNull - @Override - public Iterable getDags() { - return List.of( - InterfaceExampleBuilder.build(), - AnnotationExampleBuilder.build(), - XComCastingExampleBuilder.build()); +public class ExampleBundleBuilder { + public static Bundle build() { + return new Bundle() + .register(InterfaceExampleBuilder.build()) + .register(AnnotationExample.class) + .register(XComCastingExample.class); } public static void main(String[] args) { - var bundle = new ExampleBundleBuilder().build(); - Server.create(args).serve(bundle); + Server.create(args).serve(build()); } } diff --git a/java-sdk/example/src/java/org/apache/airflow/example/InterfaceExampleBuilder.java b/java-sdk/example/src/java/org/apache/airflow/example/InterfaceExampleBuilder.java index 78e78eed26998..4d0a494d6dcaa 100644 --- a/java-sdk/example/src/java/org/apache/airflow/example/InterfaceExampleBuilder.java +++ b/java-sdk/example/src/java/org/apache/airflow/example/InterfaceExampleBuilder.java @@ -50,10 +50,15 @@ public void execute(@NotNull Context context, Client client) throws Exception { } } - public static class Transform implements Task { - public void execute(@NotNull Context context, Client client) { - var extracted = client.getXCom("extract"); - log.log(INFO, "Got XCom from extract: {0}", extracted); + public static class TransformInput implements TaskInput { + public long extracted; + } + + // The Python Dag file calls transform(extracted), so the field of that name + // receives the extract task's XCom. + public static class Transform implements InputTask { + public void execute(@NotNull Context context, Client client, TransformInput input) { + log.log(INFO, "Got extracted value from the bound argument: {0}", input.extracted); var variable = client.getVariable("my_variable"); log.log(INFO, "Got variable: {0}", variable); @@ -63,11 +68,24 @@ public void execute(@NotNull Context context, Client client) { } } - public static class Load implements Task { - public void execute(@NotNull Context context, Client client) { - var transformed = client.getXCom("transform"); - log.log(INFO, "Got XCom from transform: {0}", transformed); - throw new RuntimeException("I failed"); + public static class SummarizeInput implements TaskInput { + // Pinned so the field can be called region rather than regionCode. Or drop + // it: public String regionCode; binds region_code with nothing declared. + @ArgName("region_code") + public String region; + + public long transformed; + } + + // summarize(region_code=..., transformed=...) is called with keyword + // arguments, which bind to the fields by name. + public static class Summarize implements InputTask { + public void execute(@NotNull Context context, Client client, SummarizeInput input) { + log.log( + INFO, "Summarize region {0} for transformed value {1}", input.region, input.transformed); + if (!"emea".equals(input.region)) { + throw new RuntimeException("expected region 'emea' but got " + input.region); + } } } @@ -75,6 +93,6 @@ public static DagDef build() { return new DagDef("java_interface_example") .addTask("extract", Extract.class) .addTask("transform", Transform.class) - .addTask("load", Load.class); + .addTask("summarize", Summarize.class); } } diff --git a/java-sdk/example/src/java/org/apache/airflow/example/XComCastingExample.java b/java-sdk/example/src/java/org/apache/airflow/example/XComCastingExample.java index c4945af865df1..5e15dd778d656 100644 --- a/java-sdk/example/src/java/org/apache/airflow/example/XComCastingExample.java +++ b/java-sdk/example/src/java/org/apache/airflow/example/XComCastingExample.java @@ -21,57 +21,72 @@ import static java.lang.System.Logger.Level.INFO; +import java.util.List; import org.apache.airflow.sdk.*; -@Builder.Dag(id = "java_xcom_casting_example") +// Stub-backed tasks wired by the Python Dag file: each parameter receives the +// value the stub call bound at its position, widening or narrowing to the +// declared type at run time. public class XComCastingExample { private static final System.Logger log = System.getLogger(XComCastingExample.class.getName()); - @Builder.Task(id = "produce_number") + @Builder.TaskHandler(dag = "java_xcom_casting_example", task = "produce_number") public int produceNumber() { log.log(INFO, "Producing int 7"); return 7; } // Any primitive numeric type (byte, short, int, long, float, double) and its boxed form works the same way. - @Builder.Task(id = "widen_to_long") - public long widenToLong(@Builder.XCom(task = "produce_number") long value) { + @Builder.TaskHandler(dag = "java_xcom_casting_example", task = "widen_to_long") + public long widenToLong(long value) { log.log(INFO, "Got long {0}", value); return value + 1; } - @Builder.Task(id = "widen_to_double") - public void widenToDouble(@Builder.XCom(task = "widen_to_long") double value) { + @Builder.TaskHandler(dag = "java_xcom_casting_example", task = "widen_to_double") + public void widenToDouble(double value) { log.log(INFO, "Got double {0}", value); if (value != 8.0) { throw new RuntimeException("expected 8.0 but got " + value); } } - @Builder.Task(id = "produce_nothing") + @Builder.TaskHandler(dag = "java_xcom_casting_example", task = "produce_nothing") public void produceNothing() { // Pushes no return_value XCom. } - @Builder.Task(id = "consume_nullable") - public void consumeNullable(@Builder.XCom(task = "produce_nothing") Integer value) { + @Builder.TaskHandler(dag = "java_xcom_casting_example", task = "consume_nullable") + public void consumeNullable(Integer value) { log.log(INFO, "Got nullable int {0}", value); if (value != null) { throw new RuntimeException("expected null but got " + value); } } - @Builder.Task(id = "produce_fraction") + @Builder.TaskHandler(dag = "java_xcom_casting_example", task = "produce_fraction") public double produceFraction() { log.log(INFO, "Producing double 1.5"); return 1.5; } - @Builder.Task(id = "consume_float") - public void consumeFloat(@Builder.XCom(task = "produce_fraction") float value) { + @Builder.TaskHandler(dag = "java_xcom_casting_example", task = "consume_float") + public void consumeFloat(float value) { log.log(INFO, "Got float {0}", value); if (value != 1.5f) { throw new RuntimeException("expected 1.5 but got " + value); } } + + // A parameter with type arguments keeps its element type. Reading an element + // is what would fail if it did not: the wire integers decode to Long unless + // the declared List survives the binding. + @Builder.TaskHandler(dag = "java_xcom_casting_example", task = "consume_double_list") + public void consumeDoubleList(List values) { + log.log(INFO, "Got list {0}", values); + double total = values.get(0) + values.get(1); + if (total != 3.0) { + throw new RuntimeException("expected 3.0 but got " + total); + } + } } diff --git a/java-sdk/example/src/resources/dags/java_examples.py b/java-sdk/example/src/resources/dags/java_examples.py index 5426911b0faee..c948d2a565237 100644 --- a/java-sdk/example/src/resources/dags/java_examples.py +++ b/java-sdk/example/src/resources/dags/java_examples.py @@ -34,27 +34,36 @@ def extract(): ... @task.stub(queue="java") -def transform(): ... +def transform(extracted): ... @task.stub(queue="java", retries=1, retry_delay=timedelta(seconds=5)) -def load(): ... +def load(transformed): ... @task.stub(queue="java") def concurrent(): ... +# Keyword arguments bind to the public fields of the Java task's TaskInput. +@task.stub(queue="java") +def report(run_label, transformed): ... + + +@task.stub(queue="java") +def summarize(region_code, transformed): ... + + @task.stub(queue="java") def produce_number(): ... @task.stub(queue="java") -def widen_to_long(): ... +def widen_to_long(value): ... @task.stub(queue="java") -def widen_to_double(): ... +def widen_to_double(value): ... @task.stub(queue="java") @@ -62,7 +71,7 @@ def produce_nothing(): ... @task.stub(queue="java") -def consume_nullable(): ... +def consume_nullable(value): ... @task.stub(queue="java") @@ -70,7 +79,11 @@ def produce_fraction(): ... @task.stub(queue="java") -def consume_float(): ... +def consume_float(value): ... + + +@task.stub(queue="java") +def consume_double_list(values): ... @task() @@ -82,25 +95,30 @@ def python_task_2(transformed): @dag(dag_id="java_interface_example") def java_interface_example(): - transformed = transform() - python_task_1() >> extract() >> transformed + extracted = extract() + python_task_1() >> extracted + transformed = transform(extracted) python_task_2(transformed) + summarize(region_code="emea", transformed=transformed) @dag(dag_id="java_annotation_example") def java_annotation_example(): - transformed = transform() - python_task_1() >> extract() >> transformed + extracted = extract() + python_task_1() >> extracted + transformed = transform(extracted) python_task_2(transformed) - transformed >> load() + load(transformed) + report(run_label="nightly", transformed=transformed) concurrent() @dag(dag_id="java_xcom_casting_example") def java_xcom_casting_example(): - produce_number() >> widen_to_long() >> widen_to_double() - produce_nothing() >> consume_nullable() - produce_fraction() >> consume_float() + widen_to_double(widen_to_long(produce_number())) + consume_nullable(produce_nothing()) + consume_float(produce_fraction()) + consume_double_list([1, 2]) java_interface_example() diff --git a/java-sdk/gradle.properties b/java-sdk/gradle.properties index 9438ba6435532..477b31da386c4 100644 --- a/java-sdk/gradle.properties +++ b/java-sdk/gradle.properties @@ -17,7 +17,7 @@ org.gradle.configuration-cache=true -airflowSupervisorSchemaVersion=2026-06-16 +airflowSupervisorSchemaVersion=2026-10-30 projectVersion=1.0.0-SNAPSHOT diff --git a/java-sdk/processor/src/main/kotlin/org/apache/airflow/sdk/BuilderProcessor.kt b/java-sdk/processor/src/main/kotlin/org/apache/airflow/sdk/BuilderProcessor.kt index 56cbf1e76ad9f..dfee528a30d3f 100644 --- a/java-sdk/processor/src/main/kotlin/org/apache/airflow/sdk/BuilderProcessor.kt +++ b/java-sdk/processor/src/main/kotlin/org/apache/airflow/sdk/BuilderProcessor.kt @@ -25,18 +25,25 @@ import com.squareup.javapoet.ClassName import com.squareup.javapoet.CodeBlock import com.squareup.javapoet.JavaFile import com.squareup.javapoet.MethodSpec +import com.squareup.javapoet.ParameterizedTypeName import com.squareup.javapoet.TypeName import com.squareup.javapoet.TypeSpec -import java.util.Optional +import org.apache.airflow.sdk.internal.ArgValues +import org.apache.airflow.sdk.internal.TaskArgs +import org.apache.airflow.sdk.internal.TypeRef +import org.apache.airflow.sdk.internal.foldArgName +import org.apache.airflow.sdk.internal.registrarName import javax.annotation.processing.AbstractProcessor import javax.annotation.processing.ProcessingEnvironment import javax.annotation.processing.RoundEnvironment import javax.annotation.processing.SupportedAnnotationTypes import javax.annotation.processing.SupportedSourceVersion import javax.lang.model.SourceVersion +import javax.lang.model.element.ElementKind import javax.lang.model.element.ExecutableElement import javax.lang.model.element.Modifier import javax.lang.model.element.TypeElement +import javax.lang.model.element.VariableElement import javax.lang.model.type.TypeKind import javax.lang.model.type.TypeMirror import javax.tools.Diagnostic @@ -57,11 +64,16 @@ import javax.tools.Diagnostic * - A static `build()` method that constructs the [DagDef] and registers those * inner classes as [TaskDef]s. * - * [Builder.XCom]-annotated parameters are resolved via `client.getXCom` in the - * generated `execute` body, with the result cast to the parameter's declared - * type. Non-`void` return values are forwarded to `client.setXCom`. + * In the generated `execute` body, a task's data parameters resolve against the + * arg bindings the supervisor delivered for the run: flat parameters through + * [TaskArgs], by their position among the data parameters, and [TaskInput] + * [TaskInput] fields through [ArgValues], by argument name. Non-`void` return values are + * forwarded to `client.setXCom`. */ -@SupportedAnnotationTypes("org.apache.airflow.sdk.Builder.Dag") +@SupportedAnnotationTypes( + "org.apache.airflow.sdk.Builder.Dag", + "org.apache.airflow.sdk.Builder.TaskHandler", +) @SupportedSourceVersion(SourceVersion.RELEASE_11) class BuilderProcessor : AbstractProcessor() { override fun process( @@ -69,6 +81,20 @@ class BuilderProcessor : AbstractProcessor() { roundEnv: RoundEnvironment, ): Boolean { if (annotations.isEmpty()) return false + roundEnv + .getElementsAnnotatedWith(Builder.TaskHandler::class.java) + .mapNotNull { it.enclosingElement as? TypeElement } + .distinct() + .forEach { el -> + with(processingEnv) { + runCatching { + JavaFile + .builder(elementUtils.getPackageOf(el).qualifiedName.toString(), buildHandlers(el)) + .build() + .writeTo(filer) + }.onFailure { e -> messager.printMessage(Diagnostic.Kind.ERROR, e.message ?: "Unknown error", el) } + } + } roundEnv.getElementsAnnotatedWith(Builder.Dag::class.java).filterIsInstance().forEach { el -> with(processingEnv) { runCatching { @@ -90,6 +116,53 @@ class BuilderProcessor : AbstractProcessor() { return true } + /** + * Generates the registrar for a class of [Builder.TaskHandler] methods: one + * [Task] implementation per handler, and a `registerInto` that binds each to + * the Dag and task the annotation names. + * + * There is no Dag to build here — the Python Dag file owns it — so this is a + * registrar rather than a builder. + */ + private fun buildHandlers(el: TypeElement): TypeSpec { + require(el.enclosingElement !is TypeElement || Modifier.STATIC in el.modifiers) { + "Nested class '${el.simpleName}' holding @Builder.TaskHandler methods must be static" + } + val registrar = + TypeSpec + .classBuilder(registrarName(ClassName.get(el).reflectionName()).substringAfterLast('.')) + .addModifiers(Modifier.PUBLIC, Modifier.FINAL) + .addJavadoc( + "Registers {@link \$T}'s task handlers against the Dags the Python file owns.\n", + ClassName.get(el), + ) + val registerInto = + MethodSpec + .methodBuilder("registerInto") + .addModifiers(Modifier.PUBLIC, Modifier.STATIC) + .addParameter(BUNDLE_TYPE, "bundle") + + for (inner in el.enclosedElements) { + if (inner !is ExecutableElement) continue + val handler = inner.getAnnotation(Builder.TaskHandler::class.java) ?: continue + if (inner.isVarArgs) { + throw IllegalArgumentException("Cannot create task from vararg function ${inner.simpleName}") + } + require(handler.dag.isNotBlank()) { + "@Builder.TaskHandler on '${inner.simpleName}' must name the Dag the Python file declares" + } + val innerName = inner.simpleName.toString().replaceFirstChar(Char::uppercase) + registrar.addType(buildTask(innerName, inner, el)) + registerInto.addStatement( + $$"bundle.register($S, $S, $L.class)", + handler.dag, + handler.task.ifBlank { inner.simpleName }, + innerName, + ) + } + return registrar.addMethod(registerInto.build()).build() + } + private fun buildDag(el: TypeElement): TypeSpec { val ann = el.getAnnotation(Builder.Dag::class.java)!! @@ -102,23 +175,22 @@ class BuilderProcessor : AbstractProcessor() { MethodSpec .methodBuilder("build") .addModifiers(Modifier.PUBLIC, Modifier.STATIC) - .returns(ClassName.get(DagDef::class.java)) - .addStatement($$"var dag = new $T($S)", ClassName.get(DagDef::class.java), ann.id.ifBlank { el.simpleName }) + .returns(DAG_DEF_TYPE) + .addStatement($$"var dag = new $T($S)", DAG_DEF_TYPE, ann.id.ifBlank { el.simpleName }) for (inner in el.enclosedElements) { if (inner !is ExecutableElement) continue if (inner.isVarArgs) throw IllegalArgumentException("Cannot create task from vararg function ${inner.simpleName}") - val ann = inner.getAnnotation(Builder.Task::class.java) ?: continue + val taskAnn = inner.getAnnotation(Builder.Task::class.java) ?: continue val innerName = inner.simpleName.toString().replaceFirstChar(Char::uppercase) - val task = buildTask(innerName, inner, el) - builderClass.addType(task.spec) + builderClass.addType(buildTask(innerName, inner, el)) buildMethod.addStatement( $$"dag.addTask(new $T($S, $L.class))", - ClassName.get(TaskDef::class.java), - ann.id.ifBlank { inner.simpleName }, + TASK_DEF_TYPE, + taskAnn.id.ifBlank { inner.simpleName }, innerName, ) } @@ -132,40 +204,58 @@ class BuilderProcessor : AbstractProcessor() { name: String, inner: ExecutableElement, parent: TypeElement, - ): BuildTaskResult { - val clientType = ClassName.get(Client::class.java) - val contextType = ClassName.get(Context::class.java) - + ): TypeSpec { val executeSpec = MethodSpec .methodBuilder("execute") .addAnnotation(Override::class.java) .addModifiers(Modifier.PUBLIC) .returns(TypeName.VOID) - .addParameter(contextType, "context") - .addParameter(clientType, "client") + .addParameter(CONTEXT_TYPE, "context") + .addParameter(CLIENT_TYPE, "client") .addException(Exception::class.java) - val required = mutableListOf() + val dataParams = collectDataParams(inner) + val dataByName = dataParams.associateBy { it.name } val innerArgs = with(processingEnv) { inner.parameters.joinToString { param -> - val anno = param.getAnnotation(Builder.XCom::class.java) val type = param.asType() when { - anno != null -> - param.simpleName.toString().also { - required += RequiredXCom(type, it, anno.task.ifBlank { it }) - } - isType(type, clientType) -> "client" - isType(type, contextType) -> "context" - else -> throw IllegalArgumentException("Unsupported task parameter '${param.simpleName}' with type: $type") + isType(type, CLIENT_TYPE) -> "client" + isType(type, CONTEXT_TYPE) -> "context" + else -> dataByName.getValue(param.simpleName.toString()).local } } } - required.forEach { - executeSpec.addStatement($$"var $L = $L", it.paramName, xcomAccess(it)) + + val taken = dataParams.mapTo(mutableSetOf()) { it.local } + val argsLocal = generateSequence("args") { "${it}_" }.first { it !in taken } + val flatParams = dataParams.filterNot { it.isTaskInput } + if (flatParams.isNotEmpty()) { + executeSpec.addStatement( + $$"$T $L = $T.of(context, client, $L)", + TASK_ARGS_TYPE, + argsLocal, + TASK_ARGS_TYPE, + flatParams.size, + ) } + dataParams.forEach { param -> + val paramType = TypeName.get(param.type) + if (param.isTaskInput) { + executeSpec.addStatement( + $$"$T $L = $T.bindInput(client, $T.class)", + paramType, + param.local, + ARG_VALUES_TYPE, + paramType, + ) + } else { + executeSpec.addStatement($$"$T $L = $L", paramType, param.local, positionalAccess(argsLocal, param)) + } + } + if (inner.returnType.kind == TypeKind.VOID) { $$"new $T().$L($L)" } else { @@ -179,80 +269,168 @@ class BuilderProcessor : AbstractProcessor() { ) } - val spec = - TypeSpec - .classBuilder(name) - .addSuperinterface(Task::class.java) - .addModifiers(Modifier.PUBLIC, Modifier.FINAL, Modifier.STATIC) - .addMethod(executeSpec.build()) - .build() - return BuildTaskResult(spec) + return TypeSpec + .classBuilder(name) + .addSuperinterface(Task::class.java) + .addModifiers(Modifier.PUBLIC, Modifier.FINAL, Modifier.STATIC) + .addMethod(executeSpec.build()) + .build() } -} -private fun ProcessingEnvironment.isType( - t: TypeMirror, - c: ClassName, -): Boolean = typeUtils.isSameType(t, elementUtils.getTypeElement(c.canonicalName()).asType()) + /** + * Collects the task method's data parameters — every parameter the SDK does + * not inject — in declaration order. A parameter's index in the returned + * list is the position it binds at: Java parameter names are not API, so + * renaming one must never rebind an input. + * + * Each gets the local the generated body reads it into, which is its own + * name unless that is one the body already uses: `execute`'s injected + * `context` and `client` are in scope for the whole method, so a data + * parameter sharing a name with one binds through a suffixed local instead. + */ + private fun collectDataParams(method: ExecutableElement): List { + val params = mutableListOf() + val taken = mutableSetOf("context", "client") + with(processingEnv) { + for (param in method.parameters) { + val type = param.asType() + if (isType(type, CLIENT_TYPE) || isType(type, CONTEXT_TYPE)) continue + val declaresTaskInput = isTaskInput(type) + if (declaresTaskInput) validateTaskInput(method, param) + val name = param.simpleName.toString() + val local = generateSequence(name) { "${it}_" }.first { it !in taken } + taken += local + params += DataParam(type, name, local, params.size, declaresTaskInput) + } + } + val inputs = params.filter { it.isTaskInput } + require(inputs.size <= 1) { + "Task method '${method.simpleName}' declares more than one TaskInput parameter: " + + inputs.joinToString { "'${it.name}'" } + } + inputs.singleOrNull()?.let { input -> + require(params.size == 1) { + "Task method '${method.simpleName}' declares TaskInput parameter '${input.name}' and other data " + + "parameters; a TaskInput owns the whole named-argument surface, so it must be the only one" + } + } + return params + } -private data class RequiredXCom( - val paramType: TypeMirror, - val paramName: String, - val taskId: String, -) + private fun ProcessingEnvironment.isTaskInput(type: TypeMirror): Boolean { + val marker = elementUtils.getTypeElement(TASK_INPUT_TYPE.canonicalName()) ?: return false + return !type.kind.isPrimitive && typeUtils.isAssignable(type, marker.asType()) + } -private val NUMBER_ACCESSORS: Map = - buildMap { - mapOf( - TypeName.BYTE to "byteValue", - TypeName.SHORT to "shortValue", - TypeName.INT to "intValue", - TypeName.LONG to "longValue", - TypeName.FLOAT to "floatValue", - TypeName.DOUBLE to "doubleValue", - ).forEach { (primitive, accessor) -> - put(primitive, accessor) - put(primitive.box(), accessor) + /** + * Checks at compile time that a [TaskInput] class can be populated at + * runtime: [ArgValues.bindInput] assigns each public non-static non-final + * field the argument it claims, by its [ArgName] value or by its own name + * folded. Two fields whose names fold alike are rejected here rather than + * at run time, since neither could be reached. + */ + private fun ProcessingEnvironment.validateTaskInput( + method: ExecutableElement, + param: VariableElement, + ) { + val inputType = + typeUtils.asElement(param.asType()) as? TypeElement + ?: throw IllegalArgumentException( + "TaskInput parameter '${param.simpleName}' of task method '${method.simpleName}' has no class type", + ) + val hasNoArgConstructor = + inputType.enclosedElements + .filterIsInstance() + .any { it.kind == ElementKind.CONSTRUCTOR && it.parameters.isEmpty() && Modifier.PUBLIC in it.modifiers } + require(hasNoArgConstructor) { + "TaskInput class ${inputType.simpleName} needs a public no-argument constructor" + } + val claimed = mutableMapOf() + instanceFields(inputType).forEach { field -> + require(Modifier.PUBLIC in field.modifiers && Modifier.FINAL !in field.modifiers) { + "TaskInput field ${inputType.simpleName}.${field.simpleName} must be public and non-final " + + "so the SDK can assign its binding" + } + val argName = field.getAnnotation(ArgName::class.java)?.value ?: field.simpleName.toString() + val previous = claimed.put(foldArgName(argName), field.simpleName.toString()) + require(previous == null) { + "TaskInput fields ${inputType.simpleName}.$previous and ${inputType.simpleName}.${field.simpleName} " + + "claim argument names that differ only in case or underscores, which the fold cannot tell " + + "apart; rename one of them" + } } } -private fun xcomAccess(xcom: RequiredXCom): CodeBlock { - val type = TypeName.get(xcom.paramType) - val accessor = NUMBER_ACCESSORS[type] - val number = ClassName.get(Number::class.java) - val optional = ClassName.get(Optional::class.java) - // A primitive parameter cannot hold null, so fail with a clear error instead of an - // opaque NullPointerException while unboxing when the XCom is absent. - val value = - if (type.isPrimitive) { - CodeBlock.of( - $$"$T.ofNullable(client.getXCom($S)).orElseThrow(() -> new $T($S, $S))", - optional, - xcom.taskId, - ClassName.get(MissingXComException::class.java), - xcom.taskId, - xcom.paramName, - ) - } else { - CodeBlock.of($$"client.getXCom($S)", xcom.taskId) + /** + * Every instance field [ArgValues.bindInput] will reach, subclass first — + * the same walk up the superclass chain the runtime makes. Declared members + * alone would miss an inherited field, and a private one would then surface + * mid-run as the very failure the build-time check exists to prevent. + */ + private fun ProcessingEnvironment.instanceFields(inputType: TypeElement): List { + val fields = mutableListOf() + var current: TypeElement? = inputType + while (current != null && !current.qualifiedName.contentEquals("java.lang.Object")) { + fields += + current.enclosedElements + .filterIsInstance() + .filter { it.kind == ElementKind.FIELD && Modifier.STATIC !in it.modifiers } + current = typeUtils.asElement(current.superclass) as? TypeElement } - // Wire integers decode to Long and floats to Double, so a direct (Integer)/(Float) - // cast throws ClassCastException; widen via Number instead. - return when { - accessor == null -> CodeBlock.of($$"($T) $L", if (type.isPrimitive) type.box() else type, value) - type.isPrimitive -> CodeBlock.of($$"(($T) $L).$L()", number, value, accessor) - else -> - CodeBlock.of( - $$"$T.ofNullable(($T) $L).map($T::$L).orElse(null)", - optional, - number, - value, - number, - accessor, - ) + return fields } } -private data class BuildTaskResult( - val spec: TypeSpec, +/** + * One data parameter of a task method, positioned among its peers, read into + * [local] by the generated body. [isTaskInput] marks a [TaskInput] parameter, + * which binds by field name instead. + */ +private class DataParam( + val type: TypeMirror, + val name: String, + val local: String, + val position: Int, + val isTaskInput: Boolean, ) + +private val DAG_DEF_TYPE = ClassName.get(DagDef::class.java) +private val TASK_DEF_TYPE = ClassName.get(TaskDef::class.java) +private val BUNDLE_TYPE = ClassName.get(Bundle::class.java) +private val CLIENT_TYPE = ClassName.get(Client::class.java) +private val CONTEXT_TYPE = ClassName.get(Context::class.java) +private val TASK_INPUT_TYPE = ClassName.get(TaskInput::class.java) +private val TASK_ARGS_TYPE = ClassName.get(TaskArgs::class.java) +private val TYPE_REF_TYPE = ClassName.get(TypeRef::class.java) +private val ARG_VALUES_TYPE = ClassName.get(ArgValues::class.java) + +private fun ProcessingEnvironment.isType( + t: TypeMirror, + c: ClassName, +): Boolean = typeUtils.isSameType(t, elementUtils.getTypeElement(c.canonicalName()).asType()) + +/** + * Emits the read for one flat data parameter, bound at its position. A + * primitive parameter cannot hold null, so `require` fails with a clear + * [MissingXComException] when the binding resolves to nothing; boxed and + * reference parameters take `get` and receive null instead. + * + * A parameter whose declared type has type arguments reads through [TypeRef] + * so its element type survives the decode — a cast cannot express + * `List`, and erasure would let one succeed over a list of numbers + * and fail much later at the first element read. + */ +private fun positionalAccess( + argsLocal: String, + param: DataParam, +): CodeBlock { + val type = TypeName.get(param.type) + val reader = if (type.isPrimitive) "require" else "get" + val target = + if (type is ParameterizedTypeName) { + CodeBlock.of($$"new $T<$T>() {}", TYPE_REF_TYPE, type) + } else { + CodeBlock.of($$"$T.class", type.box()) + } + return CodeBlock.of($$"$L.$L($L, $L)", argsLocal, reader, param.position, target) +} diff --git a/java-sdk/processor/src/test/kotlin/org/apache/airflow/sdk/BuilderTest.kt b/java-sdk/processor/src/test/kotlin/org/apache/airflow/sdk/BuilderTest.kt index 6a08979c56163..18e4d1501a178 100644 --- a/java-sdk/processor/src/test/kotlin/org/apache/airflow/sdk/BuilderTest.kt +++ b/java-sdk/processor/src/test/kotlin/org/apache/airflow/sdk/BuilderTest.kt @@ -62,7 +62,7 @@ class BuilderTest { } @Builder.Task - public void t3(Context ctx, @Builder.XCom(task = "t2") int value) { + public void t3(Context ctx, int value) { System.out.println(String.format("%s %s", ctx.ti, value)); } } @@ -78,15 +78,14 @@ class BuilderTest { package org.apache.airflow.example; import java.lang.Exception; - import java.lang.Number; + import java.lang.Integer; import java.lang.Override; - import java.util.Optional; import org.apache.airflow.sdk.Client; import org.apache.airflow.sdk.Context; import org.apache.airflow.sdk.DagDef; - import org.apache.airflow.sdk.MissingXComException; import org.apache.airflow.sdk.Task; import org.apache.airflow.sdk.TaskDef; + import org.apache.airflow.sdk.internal.TaskArgs; public final class TestExampleBuilder { public static DagDef build() { @@ -111,7 +110,8 @@ class BuilderTest { public static final class T3 implements Task { @Override public void execute(Context context, Client client) throws Exception { - var value = ((Number) Optional.ofNullable(client.getXCom("t2")).orElseThrow(() -> new MissingXComException("t2", "value"))).intValue(); + TaskArgs args = TaskArgs.of(context, client, 1); + int value = args.require(0, Integer.class); new TestExample().t3(context, value); } } @@ -121,25 +121,19 @@ class BuilderTest { } @Test - @DisplayName("widen primitive numerics directly and boxed numerics null-safely") - fun generateBuilderWidensNumericXCom() { + @DisplayName("bind data parameters by position, skipping the injected Client and Context") + fun generateBuilderBindsDataParametersByPosition() { val compilation = compile( """ package org.apache.airflow.example; import org.apache.airflow.sdk.Builder; + import org.apache.airflow.sdk.Client; + import org.apache.airflow.sdk.Context; @Builder.Dag public class TestExample { @Builder.Task - public void t( - @Builder.XCom(task = "a") int i, - @Builder.XCom(task = "b") long l, - @Builder.XCom(task = "c") double d, - @Builder.XCom(task = "f") float fl, - @Builder.XCom(task = "e") Integer boxedInteger, - @Builder.XCom(task = "g") Long boxedLong, - @Builder.XCom(task = "h") Double boxedDouble, - @Builder.XCom(task = "j") Float boxedFloat) {} + public void t(long first, Client client, String second, Context ctx, Integer third) {} } """, ) @@ -153,15 +147,16 @@ class BuilderTest { package org.apache.airflow.example; import java.lang.Exception; - import java.lang.Number; + import java.lang.Integer; + import java.lang.Long; import java.lang.Override; - import java.util.Optional; + import java.lang.String; import org.apache.airflow.sdk.Client; import org.apache.airflow.sdk.Context; import org.apache.airflow.sdk.DagDef; - import org.apache.airflow.sdk.MissingXComException; import org.apache.airflow.sdk.Task; import org.apache.airflow.sdk.TaskDef; + import org.apache.airflow.sdk.internal.TaskArgs; public final class TestExampleBuilder { public static DagDef build() { @@ -172,15 +167,11 @@ class BuilderTest { public static final class T implements Task { @Override public void execute(Context context, Client client) throws Exception { - var i = ((Number) Optional.ofNullable(client.getXCom("a")).orElseThrow(() -> new MissingXComException("a", "i"))).intValue(); - var l = ((Number) Optional.ofNullable(client.getXCom("b")).orElseThrow(() -> new MissingXComException("b", "l"))).longValue(); - var d = ((Number) Optional.ofNullable(client.getXCom("c")).orElseThrow(() -> new MissingXComException("c", "d"))).doubleValue(); - var fl = ((Number) Optional.ofNullable(client.getXCom("f")).orElseThrow(() -> new MissingXComException("f", "fl"))).floatValue(); - var boxedInteger = Optional.ofNullable((Number) client.getXCom("e")).map(Number::intValue).orElse(null); - var boxedLong = Optional.ofNullable((Number) client.getXCom("g")).map(Number::longValue).orElse(null); - var boxedDouble = Optional.ofNullable((Number) client.getXCom("h")).map(Number::doubleValue).orElse(null); - var boxedFloat = Optional.ofNullable((Number) client.getXCom("j")).map(Number::floatValue).orElse(null); - new TestExample().t(i, l, d, fl, boxedInteger, boxedLong, boxedDouble, boxedFloat); + TaskArgs args = TaskArgs.of(context, client, 3); + long first = args.require(0, Long.class); + String second = args.get(1, String.class); + Integer third = args.get(2, Integer.class); + new TestExample().t(first, client, second, context, third); } } } @@ -189,20 +180,19 @@ class BuilderTest { } @Test - @DisplayName("guard non-numeric primitives, leave objects and boxed types nullable") - fun generateBuilderGuardsNonNumericPrimitiveXCom() { + @DisplayName("require primitive parameters, leave boxed and parameterized types nullable") + fun generateBuilderRequiresPrimitivesOnly() { val compilation = compile( """ package org.apache.airflow.example; + import java.util.List; + import java.util.Map; import org.apache.airflow.sdk.Builder; @Builder.Dag public class TestExample { @Builder.Task - public void t( - @Builder.XCom(task = "a") boolean flag, - @Builder.XCom(task = "b") String text, - @Builder.XCom(task = "c") Boolean boxed) {} + public void t(boolean flag, float fraction, Double boxed, List tags, Map raw) {} } """, ) @@ -216,16 +206,20 @@ class BuilderTest { package org.apache.airflow.example; import java.lang.Boolean; + import java.lang.Double; import java.lang.Exception; + import java.lang.Float; import java.lang.Override; import java.lang.String; - import java.util.Optional; + import java.util.List; + import java.util.Map; import org.apache.airflow.sdk.Client; import org.apache.airflow.sdk.Context; import org.apache.airflow.sdk.DagDef; - import org.apache.airflow.sdk.MissingXComException; import org.apache.airflow.sdk.Task; import org.apache.airflow.sdk.TaskDef; + import org.apache.airflow.sdk.internal.TaskArgs; + import org.apache.airflow.sdk.internal.TypeRef; public final class TestExampleBuilder { public static DagDef build() { @@ -236,10 +230,13 @@ class BuilderTest { public static final class T implements Task { @Override public void execute(Context context, Client client) throws Exception { - var flag = (Boolean) Optional.ofNullable(client.getXCom("a")).orElseThrow(() -> new MissingXComException("a", "flag")); - var text = (String) client.getXCom("b"); - var boxed = (Boolean) client.getXCom("c"); - new TestExample().t(flag, text, boxed); + TaskArgs args = TaskArgs.of(context, client, 5); + boolean flag = args.require(0, Boolean.class); + float fraction = args.require(1, Float.class); + Double boxed = args.get(2, Double.class); + List tags = args.get(3, new TypeRef>() {}); + Map raw = args.get(4, Map.class); + new TestExample().t(flag, fraction, boxed, tags, raw); } } } @@ -247,6 +244,361 @@ class BuilderTest { ) } + @Test + @DisplayName("bind a TaskInput through the shared populator") + fun generateBuilderBindsTaskInputFields() { + val compilation = + compile( + """ + package org.apache.airflow.example; + import java.util.List; + import org.apache.airflow.sdk.ArgName; + import org.apache.airflow.sdk.Builder; + import org.apache.airflow.sdk.Client; + import org.apache.airflow.sdk.TaskInput; + @Builder.Dag + public class TestExample { + public static class ScoreInput implements TaskInput { + @ArgName("region_code") public String region; + public double threshold; + public List tags; + } + + @Builder.Task + public double score(Client client, ScoreInput input) { return input.threshold; } + } + """, + ) + + assertThat(compilation).succeeded() + assertThat(compilation) + .generatedSourceFile("org.apache.airflow.example.TestExampleBuilder") + .hasSourceEquivalentTo( + "org.apache.airflow.example.TestExampleBuilder", + """ + package org.apache.airflow.example; + + import java.lang.Exception; + import java.lang.Override; + import org.apache.airflow.sdk.Client; + import org.apache.airflow.sdk.Context; + import org.apache.airflow.sdk.DagDef; + import org.apache.airflow.sdk.Task; + import org.apache.airflow.sdk.TaskDef; + import org.apache.airflow.sdk.internal.ArgValues; + + public final class TestExampleBuilder { + public static DagDef build() { + var dag = new DagDef("TestExample"); + dag.addTask(new TaskDef("score", Score.class)); + return dag; + } + public static final class Score implements Task { + @Override + public void execute(Context context, Client client) throws Exception { + TestExample.ScoreInput input = ArgValues.bindInput(client, TestExample.ScoreInput.class); + client.setXCom(new TestExample().score(client, input)); + } + } + } + """, + ) + } + + @Test + @DisplayName("keep the positional handle from clashing with a parameter named args") + fun generateBuilderAvoidsClashWithParamNamedArgs() { + val compilation = + compile( + """ + package org.apache.airflow.example; + import org.apache.airflow.sdk.Builder; + @Builder.Dag + public class TestExample { + @Builder.Task + public void t(String args, int other) {} + } + """, + ) + + assertThat(compilation).succeeded() + assertThat(compilation) + .generatedSourceFile("org.apache.airflow.example.TestExampleBuilder") + .hasSourceEquivalentTo( + "org.apache.airflow.example.TestExampleBuilder", + """ + package org.apache.airflow.example; + + import java.lang.Exception; + import java.lang.Integer; + import java.lang.Override; + import java.lang.String; + import org.apache.airflow.sdk.Client; + import org.apache.airflow.sdk.Context; + import org.apache.airflow.sdk.DagDef; + import org.apache.airflow.sdk.Task; + import org.apache.airflow.sdk.TaskDef; + import org.apache.airflow.sdk.internal.TaskArgs; + + public final class TestExampleBuilder { + public static DagDef build() { + var dag = new DagDef("TestExample"); + dag.addTask(new TaskDef("t", T.class)); + return dag; + } + public static final class T implements Task { + @Override + public void execute(Context context, Client client) throws Exception { + TaskArgs args_ = TaskArgs.of(context, client, 2); + String args = args_.get(0, String.class); + int other = args_.require(1, Integer.class); + new TestExample().t(args, other); + } + } + } + """, + ) + } + + @Test + @DisplayName("keep generated locals from clashing with the injected client and context") + fun generateBuilderAvoidsClashWithInjectedParamNames() { + val compilation = + compile( + """ + package org.apache.airflow.example; + import org.apache.airflow.sdk.Builder; + import org.apache.airflow.sdk.TaskInput; + @Builder.Dag + public class TestExample { + public static class ScoreInput implements TaskInput { + public double threshold; + } + + @Builder.Task + public void flat(String client, int context) {} + + @Builder.Task + public void named(ScoreInput context) {} + } + """, + ) + + assertThat(compilation).succeeded() + assertThat(compilation) + .generatedSourceFile("org.apache.airflow.example.TestExampleBuilder") + .hasSourceEquivalentTo( + "org.apache.airflow.example.TestExampleBuilder", + """ + package org.apache.airflow.example; + + import java.lang.Exception; + import java.lang.Integer; + import java.lang.Override; + import java.lang.String; + import org.apache.airflow.sdk.Client; + import org.apache.airflow.sdk.Context; + import org.apache.airflow.sdk.DagDef; + import org.apache.airflow.sdk.Task; + import org.apache.airflow.sdk.TaskDef; + import org.apache.airflow.sdk.internal.ArgValues; + import org.apache.airflow.sdk.internal.TaskArgs; + + public final class TestExampleBuilder { + public static DagDef build() { + var dag = new DagDef("TestExample"); + dag.addTask(new TaskDef("flat", Flat.class)); + dag.addTask(new TaskDef("named", Named.class)); + return dag; + } + public static final class Flat implements Task { + @Override + public void execute(Context context, Client client) throws Exception { + TaskArgs args = TaskArgs.of(context, client, 2); + String client_ = args.get(0, String.class); + int context_ = args.require(1, Integer.class); + new TestExample().flat(client_, context_); + } + } + public static final class Named implements Task { + @Override + public void execute(Context context, Client client) throws Exception { + TestExample.ScoreInput context_ = ArgValues.bindInput(client, TestExample.ScoreInput.class); + new TestExample().named(context_); + } + } + } + """, + ) + } + + @Test + @DisplayName("reject a TaskInput whose inherited field cannot be assigned") + fun rejectTaskInputWithNonPublicInheritedField() { + val compilation = + compile( + """ + package org.apache.airflow.example; + import org.apache.airflow.sdk.Builder; + import org.apache.airflow.sdk.TaskInput; + @Builder.Dag + public class TestExample { + public static class BaseInput { + private String secret; + } + + public static class ScoreInput extends BaseInput implements TaskInput { + public double threshold; + } + + @Builder.Task + public void t(ScoreInput input) {} + } + """, + ) + assertThat(compilation).failed() + assertThat(compilation).hadErrorContaining( + "TaskInput field ScoreInput.secret must be public and non-final", + ) + } + + @Test + @DisplayName("reject a TaskInput whose inherited field folds onto a declared one") + fun rejectTaskInputWithCollidingInheritedField() { + val compilation = + compile( + """ + package org.apache.airflow.example; + import org.apache.airflow.sdk.Builder; + import org.apache.airflow.sdk.TaskInput; + @Builder.Dag + public class TestExample { + public static class BaseInput { + public String regionCode; + } + + public static class ScoreInput extends BaseInput implements TaskInput { + public String region_code; + } + + @Builder.Task + public void t(ScoreInput input) {} + } + """, + ) + assertThat(compilation).failed() + assertThat(compilation).hadErrorContaining( + "TaskInput fields ScoreInput.region_code and ScoreInput.regionCode claim argument names that " + + "differ only in case or underscores", + ) + } + + @Test + @DisplayName("reject a TaskInput mixed with flat data parameters") + fun rejectTaskInputMixedWithFlatParams() { + val compilation = + compile( + """ + package org.apache.airflow.example; + import org.apache.airflow.sdk.Builder; + import org.apache.airflow.sdk.TaskInput; + @Builder.Dag + public class TestExample { + public static class ScoreInput implements TaskInput { + public double threshold; + } + + @Builder.Task + public void t(ScoreInput input, int extra) {} + } + """, + ) + assertThat(compilation).failed() + assertThat(compilation).hadErrorContaining( + "Task method 't' declares TaskInput parameter 'input' and other data parameters", + ) + } + + @Test + @DisplayName("reject a task declaring more than one TaskInput") + fun rejectMultipleTaskInputs() { + val compilation = + compile( + """ + package org.apache.airflow.example; + import org.apache.airflow.sdk.Builder; + import org.apache.airflow.sdk.TaskInput; + @Builder.Dag + public class TestExample { + public static class ScoreInput implements TaskInput { + public double threshold; + } + + @Builder.Task + public void t(ScoreInput first, ScoreInput second) {} + } + """, + ) + assertThat(compilation).failed() + assertThat(compilation).hadErrorContaining( + "Task method 't' declares more than one TaskInput parameter: 'first', 'second'", + ) + } + + @Test + @DisplayName("reject a TaskInput with a non-public field") + fun rejectTaskInputWithNonPublicField() { + val compilation = + compile( + """ + package org.apache.airflow.example; + import org.apache.airflow.sdk.Builder; + import org.apache.airflow.sdk.TaskInput; + @Builder.Dag + public class TestExample { + public static class ScoreInput implements TaskInput { + double threshold; + } + + @Builder.Task + public void t(ScoreInput input) {} + } + """, + ) + assertThat(compilation).failed() + assertThat(compilation).hadErrorContaining( + "TaskInput field ScoreInput.threshold must be public and non-final", + ) + } + + @Test + @DisplayName("reject a TaskInput without a public no-argument constructor") + fun rejectTaskInputWithoutNoArgConstructor() { + val compilation = + compile( + """ + package org.apache.airflow.example; + import org.apache.airflow.sdk.Builder; + import org.apache.airflow.sdk.TaskInput; + @Builder.Dag + public class TestExample { + public static class ScoreInput implements TaskInput { + public double threshold; + + public ScoreInput(double threshold) { this.threshold = threshold; } + } + + @Builder.Task + public void t(ScoreInput input) {} + } + """, + ) + assertThat(compilation).failed() + assertThat(compilation).hadErrorContaining( + "TaskInput class ScoreInput needs a public no-argument constructor", + ) + } + @Test @DisplayName("generate builder for dag class with custom dag id") fun generateBuilderWithCustomDagId() { @@ -334,38 +686,183 @@ class BuilderTest { } @Test - @DisplayName("generate builder for dag class with invalid task parameter") - fun generateBuilderForDagClassWithInvalidTaskParameter() { + @DisplayName("generate builder for dag class with varargs task parameter") + fun generateBuilderForDagClassWithVarArgsTaskParameter() { val compilation = compile( """ package org.apache.airflow.example; import org.apache.airflow.sdk.Builder; @Builder.Dag - public class TestExample { @Builder.Task(id = "foo") public void t1(String client) {} } + public class TestExample { @Builder.Task(id = "foo") public void t1(String... client) {} } """, ) assertThat(compilation).failed() assertThat(compilation).hadErrorContaining( - "Unsupported task parameter 'client' with type: java.lang.String", + "Cannot create task from vararg function t1", ) } @Test - @DisplayName("generate builder for dag class with varargs task parameter") - fun generateBuilderForDagClassWithVarArgsTaskParameter() { + @DisplayName("generate a registrar binding each handler to the ids its annotation names") + fun generateHandlerRegistrar() { val compilation = compile( """ package org.apache.airflow.example; import org.apache.airflow.sdk.Builder; - @Builder.Dag - public class TestExample { @Builder.Task(id = "foo") public void t1(String... client) {} } + import org.apache.airflow.sdk.Client; + public class TestExample { + @Builder.TaskHandler(dag = "etl", task = "score") + public long score(Client client, long rows) { return rows; } + + @Builder.TaskHandler(dag = "etl") + public void audit() {} + } + """, + ) + + assertThat(compilation).succeeded() + assertThat(compilation) + .generatedSourceFile("org.apache.airflow.example.TestExampleHandlers") + .hasSourceEquivalentTo( + "org.apache.airflow.example.TestExampleHandlers", + """ + package org.apache.airflow.example; + + import java.lang.Exception; + import java.lang.Long; + import java.lang.Override; + import org.apache.airflow.sdk.Bundle; + import org.apache.airflow.sdk.Client; + import org.apache.airflow.sdk.Context; + import org.apache.airflow.sdk.Task; + import org.apache.airflow.sdk.internal.TaskArgs; + + /** + * Registers {@link TestExample}'s task handlers against the Dags the Python file owns. + */ + public final class TestExampleHandlers { + public static void registerInto(Bundle bundle) { + bundle.register("etl", "score", Score.class); + bundle.register("etl", "audit", Audit.class); + } + + public static final class Score implements Task { + @Override + public void execute(Context context, Client client) throws Exception { + TaskArgs args = TaskArgs.of(context, client, 1); + long rows = args.require(0, Long.class); + client.setXCom(new TestExample().score(client, rows)); + } + } + + public static final class Audit implements Task { + @Override + public void execute(Context context, Client client) throws Exception { + new TestExample().audit(); + } + } + } + """, + ) + } + + @Test + @DisplayName("reject a handler that names no Dag") + fun rejectHandlerWithoutDag() { + val compilation = + compile( + """ + package org.apache.airflow.example; + import org.apache.airflow.sdk.Builder; + public class TestExample { + @Builder.TaskHandler(dag = "") + public void t() {} + } """, ) assertThat(compilation).failed() assertThat(compilation).hadErrorContaining( - "Cannot create task from vararg function t1", + "@Builder.TaskHandler on 't' must name the Dag the Python file declares", + ) + } + + @Test + @DisplayName("name the registrar of a nested handler class after the classes enclosing it") + fun generateRegistrarForNestedHandlerClass() { + val compilation = + compile( + """ + package org.apache.airflow.example; + import org.apache.airflow.sdk.Builder; + public class TestExample { + public static class Inner { + @Builder.TaskHandler(dag = "etl", task = "score") + public long score(long rows) { return rows; } + } + } + """, + ) + + assertThat(compilation).succeeded() + assertThat(compilation) + .generatedSourceFile("org.apache.airflow.example.TestExample_InnerHandlers") + .hasSourceEquivalentTo( + "org.apache.airflow.example.TestExample_InnerHandlers", + """ + package org.apache.airflow.example; + + import java.lang.Exception; + import java.lang.Long; + import java.lang.Override; + import org.apache.airflow.sdk.Bundle; + import org.apache.airflow.sdk.Client; + import org.apache.airflow.sdk.Context; + import org.apache.airflow.sdk.Task; + import org.apache.airflow.sdk.internal.TaskArgs; + + /** + * Registers {@link TestExample.Inner}'s task handlers against the Dags the Python file owns. + */ + public final class TestExample_InnerHandlers { + public static void registerInto(Bundle bundle) { + bundle.register("etl", "score", Score.class); + } + + public static final class Score implements Task { + @Override + public void execute(Context context, Client client) throws Exception { + TaskArgs args = TaskArgs.of(context, client, 1); + long rows = args.require(0, Long.class); + client.setXCom(new TestExample.Inner().score(rows)); + } + } + } + """, + ) + } + + @Test + @DisplayName("reject handlers on a nested class that is not static") + fun rejectHandlersOnInnerClass() { + val compilation = + compile( + """ + package org.apache.airflow.example; + import org.apache.airflow.sdk.Builder; + public class TestExample { + public class Inner { + @Builder.TaskHandler(dag = "etl") + public void t() {} + } + } + """, + ) + + assertThat(compilation).failed() + assertThat(compilation).hadErrorContaining( + "Nested class 'Inner' holding @Builder.TaskHandler methods must be static", ) } } diff --git a/java-sdk/sdk/build.gradle.kts b/java-sdk/sdk/build.gradle.kts index 98a30c653a8df..fb74fcf9c9d01 100644 --- a/java-sdk/sdk/build.gradle.kts +++ b/java-sdk/sdk/build.gradle.kts @@ -276,11 +276,17 @@ dokka { // Dokka rejects the file unless "# Module sdk" is its very first line, so module.md carries // the ASF license header just below the heading instead of above it. includes.from("module.md") - // Suppress everything in 'execution' since it's implementation detail. + // Suppress everything in 'execution' and 'internal' since they're + // implementation detail: generated task classes call into them, so they + // are public on the JVM, but no Dag author writes against them. perPackageOption { matchingRegex = """org\.apache\.airflow\.sdk\.execution.*""" suppress.set(true) } + perPackageOption { + matchingRegex = """org\.apache\.airflow\.sdk\.internal.*""" + suppress.set(true) + } } } diff --git a/java-sdk/sdk/module.md b/java-sdk/sdk/module.md index 9e1c86d895c7a..c17da64966e56 100644 --- a/java-sdk/sdk/module.md +++ b/java-sdk/sdk/module.md @@ -29,7 +29,7 @@ meaning of each dimension is defined in the -*Min. Airflow version: 3.3 · supervisor schema: 2026-06-16* +*Min. Airflow version: 3.3 · supervisor schema: 2026-10-30* | Dimension | Tier | Supported | Since | Notes | |---|---|---|---|---| diff --git a/java-sdk/sdk/schema/schema.json b/java-sdk/sdk/schema/schema.json index e6ce8aa3d066e..b671959c50a00 100644 --- a/java-sdk/sdk/schema/schema.json +++ b/java-sdk/sdk/schema/schema.json @@ -1,6 +1,6 @@ { "$schema": "https://json-schema.org/draft/2020-12/schema", - "api_version": "2026-06-16", + "api_version": "2026-10-30", "description": "Apache Airflow SDK Supervisor Schema", "$defs": { "AssetAliasReferenceAssetEventDagRun": { @@ -753,6 +753,19 @@ ], "title": "Bundle Version" }, + "version_data": { + "anyOf": [ + { + "additionalProperties": true, + "type": "object" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Version Data" + }, "msg": { "anyOf": [ { @@ -1651,6 +1664,19 @@ ], "title": "Bundle Version" }, + "version_data": { + "anyOf": [ + { + "additionalProperties": true, + "type": "object" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Version Data" + }, "msg": { "anyOf": [ { @@ -1846,6 +1872,45 @@ "title": "Ascending", "type": "boolean" }, + "partition_key": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Partition Key" + }, + "partition_key_regexp_pattern": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Partition Key Regexp Pattern" + }, + "extra": { + "anyOf": [ + { + "additionalProperties": { + "type": "string" + }, + "type": "object" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Extra" + }, "type": { "const": "GetAssetEventByAsset", "default": "GetAssetEventByAsset", @@ -1909,6 +1974,45 @@ "title": "Ascending", "type": "boolean" }, + "partition_key": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Partition Key" + }, + "partition_key_regexp_pattern": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Partition Key Regexp Pattern" + }, + "extra": { + "anyOf": [ + { + "additionalProperties": { + "type": "string" + }, + "type": "object" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Extra" + }, "type": { "const": "GetAssetEventByAssetAlias", "default": "GetAssetEventByAssetAlias", @@ -3859,6 +3963,19 @@ ], "title": "Bundle Version" }, + "version_data": { + "anyOf": [ + { + "additionalProperties": true, + "type": "object" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Version Data" + }, "msg": { "anyOf": [ { @@ -4446,6 +4563,131 @@ "title": "ConnectionResponse", "type": "object" }, + "ArgValueSchema": { + "additionalProperties": { + "$ref": "#/$defs/JsonValue" + }, + "title": "ArgValueSchema", + "type": "object" + }, + "LiteralArgBinding": { + "description": "One positional stub-task argument carrying an inline literal from the Dag file.", + "properties": { + "kind": { + "const": "literal", + "title": "Kind", + "type": "string" + }, + "name": { + "title": "Name", + "type": "string" + }, + "value_schema": { + "anyOf": [ + { + "$ref": "#/$defs/ArgValueSchema" + }, + { + "type": "null" + } + ], + "default": null + }, + "value": { + "anyOf": [ + { + "$ref": "#/$defs/JsonValue" + }, + { + "type": "null" + } + ], + "default": null + }, + "from_default": { + "default": false, + "title": "From Default", + "type": "boolean" + } + }, + "required": [ + "kind", + "name" + ], + "title": "LiteralArgBinding", + "type": "object" + }, + "TaskArgBinding": { + "discriminator": { + "mapping": { + "literal": "#/$defs/LiteralArgBinding", + "xcom": "#/$defs/XComArgBinding" + }, + "propertyName": "kind" + }, + "oneOf": [ + { + "$ref": "#/$defs/XComArgBinding" + }, + { + "$ref": "#/$defs/LiteralArgBinding" + } + ], + "title": "TaskArgBinding" + }, + "XComArgBinding": { + "description": "One positional stub-task argument pulled from an upstream task's XCom.", + "properties": { + "kind": { + "const": "xcom", + "title": "Kind", + "type": "string" + }, + "name": { + "title": "Name", + "type": "string" + }, + "value_schema": { + "anyOf": [ + { + "$ref": "#/$defs/ArgValueSchema" + }, + { + "type": "null" + } + ], + "default": null + }, + "task_id": { + "title": "Task Id", + "type": "string" + }, + "map_index": { + "default": -1, + "title": "Map Index", + "type": "integer" + }, + "element_index": { + "anyOf": [ + { + "type": "integer" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Element Index" + } + }, + "required": [ + "kind", + "name", + "task_id" + ], + "title": "XComArgBinding", + "type": "object" + }, "AssetEventDagRunReference": { "additionalProperties": false, "description": "Schema for AssetEvent model used in DagRun.", @@ -4809,6 +5051,21 @@ ], "default": null, "title": "Start Date" + }, + "arg_bindings": { + "anyOf": [ + { + "items": { + "$ref": "#/$defs/TaskArgBinding" + }, + "type": "array" + }, + { + "type": "null" + } + ], + "default": null, + "title": "Arg Bindings" } }, "required": [ diff --git a/java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/ArgName.kt b/java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/ArgName.kt new file mode 100644 index 0000000000000..47334d23475a7 --- /dev/null +++ b/java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/ArgName.kt @@ -0,0 +1,44 @@ +/* + * 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 org.apache.airflow.sdk + +/** + * Pins the stub argument a [TaskInput] field binds, for a Python name that no + * Java identifier reaches on its own. + * + * A field without the annotation already matches its argument ignoring case + * and underscores, so an ordinary `snake_case` parameter needs no annotation: + * + * ```java + * public String regionCode; // binds region_code + * @ArgName("class") public String klass; // binds a Python keyword + * ``` + * + * Reach for it when the argument name is not a legal or usable Java + * identifier. A pinned name is matched as written, with no folding, so the + * annotation says exactly which argument the field takes. + * + * @param value Argument name as declared in the stub task's signature. + */ +@Target(AnnotationTarget.FIELD) +@MustBeDocumented +annotation class ArgName( + val value: String, +) diff --git a/java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/Builder.kt b/java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/Builder.kt index 9dbfbfbefc390..8d00b7ca7a9f6 100644 --- a/java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/Builder.kt +++ b/java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/Builder.kt @@ -36,10 +36,15 @@ package org.apache.airflow.sdk * public long extract(Client client) { ... } * * @Builder.Task(id = "transform") - * public long transform(Client client, @Builder.XCom(task = "extract") long extracted) { ... } + * public long transform(Client client, long extracted) { ... } * } * ``` * + * A task method's data parameters — everything other than the injected + * [Client] and [Context] — receive the arguments the Python `@task.stub` call + * site bound, by position. Keyword arguments bind by name instead through a + * single [TaskInput] parameter. + * * The processor generates `MyPipelineBuilder.build()`, which returns a * fully wired-up [DagDef] ready to add to a [Bundle]. */ @@ -75,14 +80,30 @@ class Builder internal constructor() { ) /** - * Annotation to mark a task definition's method parameter as an XCom input. + * Marks a method as the Java body of a task the Python Dag file declares + * with `@task.stub`. + * + * This is not [Task] under another name. Python declares the task and Java + * supplies only its body, so the handler names the pair it binds to rather + * than an id it owns — Python owns both, and the processor generates the + * registration from them: + * + * ```java + * @Builder.TaskHandler(dag = "etl", task = "score") + * public long score(Client client, long rows, double threshold) { ... } + * ``` * - * @param task The task ID to pull. If empty or not given, the annotated - * parameter's name is used by default. + * Register every handler a class holds with [Bundle.register]; there is no + * [Dag] annotation on this surface, because the Dag is the Python file's. + * + * @param dag Dag ID as declared in the Python Dag file. + * @param task Task ID as declared by the `@task.stub` function. Empty + * derives it from the annotated method's name. */ - @Target(AnnotationTarget.VALUE_PARAMETER) + @Target(AnnotationTarget.FUNCTION) @MustBeDocumented - annotation class XCom( + annotation class TaskHandler( + val dag: String, val task: String = "", ) } diff --git a/java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/Bundle.kt b/java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/Bundle.kt index 6cd549270a8a4..447148b21524a 100644 --- a/java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/Bundle.kt +++ b/java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/Bundle.kt @@ -19,29 +19,129 @@ package org.apache.airflow.sdk +import org.apache.airflow.sdk.internal.registrarName + /** - * An immutable snapshot of all [DagDef]s that this JVM process can execute. + * All [DagDef]s that this JVM process can execute. * - * Build a [Bundle] by implementing [BundleBuilder], then pass it to - * [Server.serve] to start accepting task-execution requests. + * Register everything before passing the bundle to [Server.serve]: serving + * ends registration, so a `register` left below it fails rather than racing + * the running task. * - * @property dags All registered Dags keyed by [DagDef.id]. + * @property dags Dags declared in Java, keyed by [DagDef.id]. * @throws IllegalArgumentException if any two Dags share the same ID. */ class Bundle( dags: Iterable, ) { - internal val dags: Map = dags.associateByDagId() -} + /** Dags declared in Java, which own their own tasks. */ + internal val dags = linkedMapOf() + + /** Dags the Python file owns, holding the task handlers registered for them. */ + internal val taskHandlers = linkedMapOf() + + @Volatile + private var served = false -private fun Iterable.associateByDagId(): Map { - val dagMap = linkedMapOf() - for (dag in this) { - require(dagMap.putIfAbsent(dag.id, dag) == null) { + /** Creates an empty bundle to [register] into. */ + constructor() : this(emptyList()) + + init { + dags.forEach { register(it) } + } + + /** + * Registers a Dag. + * + * @return This bundle, for chaining. + * @throws IllegalArgumentException if another Dag shares its ID, or task + * handlers are already registered against it. + * @throws IllegalStateException if [Server.serve] has already been called. + */ + fun register(dag: DagDef): Bundle { + checkOpen() + require(dag.id !in taskHandlers) { + "Dag '${dag.id}' already has registered task handlers; a Dag declared in Java owns its " + + "own tasks, so one Dag ID cannot have both" + } + require(dags.putIfAbsent(dag.id, dag) == null) { "Dags in bundle have duplicate ID: ${dag.id}" } + return this } - return dagMap + + /** + * Registers every task handler a class holds, from the ids each + * [Builder.TaskHandler] names. + * + * @param handlerClass A class with [Builder.TaskHandler] methods. + * @return This bundle, for chaining. + * @throws IllegalArgumentException if the class has no generated + * registrar, because annotation processing did not run over it. + */ + fun register(handlerClass: Class<*>): Bundle { + checkOpen() + val name = registrarName(handlerClass.name) + val registrar = + try { + Class.forName(name, true, handlerClass.classLoader) + } catch (e: ClassNotFoundException) { + throw IllegalArgumentException( + "No generated registrar $name for ${handlerClass.name}; does it declare " + + "@Builder.TaskHandler methods, and is airflow-sdk-processor on the " + + "annotationProcessor path?", + e, + ) + } + registrar.getMethod("registerInto", Bundle::class.java).invoke(null, this) + return this + } + + /** + * Registers one task implementation against a Dag the Python file owns, for + * a task with no annotation to read the ids from. + * + * The Dag is created on first use: a stub-backed Dag exists only so the + * runtime can find the task, and its graph lives in the Python Dag file. + * + * @param dagId Dag ID as declared in the Python Dag file. + * @param taskId Task ID as declared by the `@task.stub` function. + * @param definition Class that implements [Task]. + * @return This bundle, for chaining. + * @throws IllegalArgumentException if a Dag declared in Java already holds + * that ID. + * @throws IllegalStateException if [Server.serve] has already been called. + */ + fun register( + dagId: String, + taskId: String, + definition: Class, + ): Bundle { + checkOpen() + require(dagId !in dags) { + "Dag '$dagId' is declared in Java; attach its tasks with addTask(...) rather than " + + "registering task handlers for them" + } + taskHandlers.getOrPut(dagId) { DagDef(dagId) }.addTask(taskId, definition) + return this + } + + /** The task to run for a request, from whichever side registered its Dag. */ + internal fun taskDef( + dagId: String, + taskId: String, + ): TaskDef? = (dags[dagId] ?: taskHandlers[dagId])?.tasks?.get(taskId) + + /** + * Ends registration, so a `register` left below `serve` is reported as the + * mistake it is rather than racing the runtime. [Server] calls it when it + * starts serving, whatever the run turns out to do. + */ + internal fun finalizeRegistration() { + served = true + } + + private fun checkOpen() = check(!served) { "Server.serve has already been called; register everything before serve" } } /** diff --git a/java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/Client.kt b/java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/Client.kt index be7e582eef764..b01e604ad12f7 100644 --- a/java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/Client.kt +++ b/java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/Client.kt @@ -19,8 +19,10 @@ package org.apache.airflow.sdk +import org.apache.airflow.sdk.execution.ArgBinding import org.apache.airflow.sdk.execution.Client import org.apache.airflow.sdk.execution.comm.StartupDetails +import org.apache.airflow.sdk.execution.decodeArgBindings /** * A connection registered in Airflow's connection store. @@ -177,10 +179,45 @@ class Client internal constructor( runId = details.ti.runId, mapIndex = details.ti.mapIndex ?: -1, ) + + internal val argBindings: List by lazy { + decodeArgBindings(details.tiContext?.argBindings) + } + + // A literal binding carries the inline value from the Dag file; an XCom + // binding pulls the bound upstream task's return-value XCom, honouring the + // bound map index and element index. + internal fun resolveBinding(binding: ArgBinding): Any? = + when (binding) { + is ArgBinding.Literal -> binding.value + is ArgBinding.XCom -> { + val value = getXCom(taskId = binding.taskId, mapIndex = binding.mapIndex.takeIf { it >= 0 }) + binding.elementIndex?.let { elementOf(value, it, binding) } ?: value + } + } + + /** + * Reads the element a binding indexes out of an upstream's list XCom. An + * upstream that pushed nothing resolves to null like any other unpushed + * binding, so whether a parameter can be null stays the parameter's own + * question rather than the call site's. + */ + private fun elementOf( + value: Any?, + index: Int, + binding: ArgBinding.XCom, + ): Any? { + if (value == null) return null + val bound = "Argument '${binding.name}' binds element $index of task '${binding.taskId}'" + check(value is List<*>) { "$bound, but its XCom is not a list" } + check(index in value.indices) { "$bound, but its XCom holds only ${value.size} element(s)" } + return value[index] + } } /** - * Thrown when a task parameter with a primitive type reads an XCom that was never pushed. + * Thrown when a task's input resolves to nothing where a value is required — + * a data parameter or a [TaskInput] field with a primitive type. */ class MissingXComException( message: String, diff --git a/java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/DagDef.kt b/java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/DagDef.kt index e3ba700a0fa1c..135d651f5a912 100644 --- a/java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/DagDef.kt +++ b/java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/DagDef.kt @@ -19,6 +19,7 @@ package org.apache.airflow.sdk +import org.apache.airflow.sdk.internal.validateTaskInput import kotlin.Throws /** @@ -95,6 +96,9 @@ class DagDef( * @param id Task identifier, unique within a [DagDef]. * @param definition Class that implements [Task]. Must have a public no-arg * constructor. + * @throws IllegalArgumentException if [definition] is an [InputTask] whose + * declared input cannot be bound, so that a mis-declared input fails while + * the [Bundle] is built rather than mid-run. * * @see Builder.Task */ @@ -102,6 +106,10 @@ class TaskDef( val id: String, val definition: Class, ) { + init { + validateTaskInput(definition) + } + internal var owner: DagDef? = null } @@ -116,8 +124,12 @@ class TaskDef( * via its no-argument constructor, then calls [execute] once per task-instance * run. * + * Implement [InputTask] instead for a task the Python Dag file calls with + * TaskFlow arguments; the SDK then resolves those arguments and injects them. + * * @see Builder.Dag * @see Builder.Task + * @see InputTask */ interface Task { /** diff --git a/java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/InputTask.kt b/java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/InputTask.kt new file mode 100644 index 0000000000000..52276f36f4ebf --- /dev/null +++ b/java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/InputTask.kt @@ -0,0 +1,88 @@ +/* + * 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 org.apache.airflow.sdk + +import org.apache.airflow.sdk.internal.ArgValues +import org.apache.airflow.sdk.internal.resolveInputType +import kotlin.Throws + +/** + * A [Task] whose input the SDK resolves from the `@task.stub` TaskFlow call + * site in the Python Dag file and injects. + * + * The type argument is the [TaskInput] the task expects, and it is the only + * place that declaration lives. Each of its public fields binds **by name**, + * ignoring case and underscores, so the stub's `snake_case` arguments reach + * `camelCase` Java fields with nothing declared. + * + * ```java + * public class Summarize implements InputTask { + * @Override + * public void execute(Context context, Client client, SummarizeInput input) { + * // input.region, input.transformed + * } + * } + * ``` + * + * Plain [Task] stays the right choice for a task the Dag file calls with no + * arguments. + * + * @param I This task's input. + * + * @see TaskInput + */ +interface InputTask : Task { + /** + * Resolves this task's declared input, then runs [execute]. Implementations + * override the three-argument [execute] instead of this method. + */ + @Throws(Exception::class) + override fun execute( + context: Context, + client: Client, + ) = execute(context, client, ArgValues.bindInput(client, inputType())) + + /** + * Executes this task. + * + * Any exception thrown marks the task instance as failed. Use [client] to + * read connections, variables, pull XComs, or to push an XCom for downstream + * tasks. + * + * @param context Runtime context for the current execution workload. + * @param client Client for Airflow API calls scoped to this execution. + * @param input This task's arguments, as bound at the stub call site. + * @throws Exception on failure; the task instance is marked failed. + */ + @Throws(Exception::class) + fun execute( + context: Context, + client: Client, + input: I, + ) +} + +/** + * Recovers the [TaskInput] type this task bound to [InputTask]'s type + * parameter. [TaskDef] resolves it up front, so reaching a task run means it + * is resolvable. + */ +@Suppress("UNCHECKED_CAST") +private fun InputTask.inputType(): Class = resolveInputType(javaClass) as Class diff --git a/java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/Server.kt b/java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/Server.kt index 361fd0fe319fc..4b42eb68d4c39 100644 --- a/java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/Server.kt +++ b/java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/Server.kt @@ -143,6 +143,7 @@ class Server( */ suspend fun serveAsync(bundle: Bundle) = coroutineScope { + bundle.finalizeRegistration() val deferral = CompletableDeferred() launch { diff --git a/java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/TaskInput.kt b/java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/TaskInput.kt new file mode 100644 index 0000000000000..8fd78fa632ce6 --- /dev/null +++ b/java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/TaskInput.kt @@ -0,0 +1,60 @@ +/* + * 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 org.apache.airflow.sdk + +/** + * Marks a class as a task's input: when the Python Dag file calls the stub + * task with keyword arguments, each public field receives the argument whose + * name matches it, ignoring case and underscores. That is the same fold the + * Go and TypeScript SDKs apply, so one Python signature binds identically in + * every SDK with nothing declared; [ArgName] pins a name the fold cannot + * reach. + * + * ```java + * public static class ScoreInput implements TaskInput { + * // Pinned to region_code. Or drop it: public String regionCode; + * @ArgName("region_code") + * public String region; + * + * public double threshold; // binds threshold + * } + * + * @Builder.TaskHandler(dag = "etl", task = "score") + * public Result score(Client client, ScoreInput input) { ... } + * ``` + * + * A `TaskInput` binds the same way whichever authoring API declares it: as a + * `@Builder.TaskHandler` parameter, as above, or as the input type of an + * [InputTask]. A task method may declare at most one `TaskInput` parameter + * and, if it does, no other data parameters — the `TaskInput` owns the whole named-argument + * surface, so field names and flat positions cannot shift each other. + * + * The class needs a public no-argument constructor and public non-final + * fields. No two of those fields may claim argument names that differ only in + * case or underscores, since the fold cannot tell them apart. + * + * A name mismatch in either direction is logged rather than failed: a field + * nothing supplies keeps its Java default, and an argument no field claims + * changes nothing the task reads. Once a field has claimed its argument, + * declare a boxed type when that argument may resolve to nothing. + * + * @see InputTask + */ +interface TaskInput diff --git a/java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/execution/ArgBinding.kt b/java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/execution/ArgBinding.kt new file mode 100644 index 0000000000000..33730ea045a67 --- /dev/null +++ b/java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/execution/ArgBinding.kt @@ -0,0 +1,90 @@ +/* + * 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 org.apache.airflow.sdk.execution + +/** + * One stub-task argument bound at the `@task.stub` TaskFlow call site in the + * Python Dag file, delivered via `TIRunContext.arg_bindings`. + * + * The supervisor schema models this as a `kind`-discriminated union + * (`XComArgBinding` / `LiteralArgBinding`), which jsonSchema2Pojo cannot + * express as a typed field — the generated `TIRunContext.argBindings` is a + * plain `Object` holding the msgpack-decoded list of maps — so this hand- + * written decoder materializes the typed view. + */ +internal sealed class ArgBinding { + abstract val name: String + + /** + * Whether the call site omitted this argument, so the value carried here is + * the stub signature's own default rather than one the Dag author wrote. + */ + open val fromDefault: Boolean + get() = false + + internal data class XCom( + override val name: String, + val taskId: String, + val mapIndex: Int, + val elementIndex: Int?, + ) : ArgBinding() + + internal data class Literal( + override val name: String, + val value: Any?, + override val fromDefault: Boolean = false, + ) : ArgBinding() +} + +/** + * Decodes the raw `TIRunContext.argBindings` payload into a list of bindings + * preserving the stub signature's parameter order — flat task parameters + * bind by that position, input-bundle fields by [ArgBinding.name]. + * + * @throws IllegalStateException on a malformed payload, an unsupported + * binding kind, or a duplicate argument name; the task cannot bind its + * arguments correctly, so it must fail rather than run with wrong inputs. + */ +internal fun decodeArgBindings(raw: Any?): List { + if (raw == null) return emptyList() + check(raw is List<*>) { "arg_bindings payload is not a list: ${raw.javaClass.name}" } + val seen = mutableSetOf() + return raw.map { entry -> + check(entry is Map<*, *>) { "arg_bindings entry is not a map: $entry" } + val name = checkNotNull(entry["name"] as? String) { "arg_bindings entry has no name: $entry" } + check(seen.add(name)) { "arg_bindings entries have duplicate name: '$name'" } + when (val kind = entry["kind"]) { + "literal" -> + ArgBinding.Literal( + name = name, + value = entry["value"], + fromDefault = entry["from_default"] == true, + ) + "xcom" -> + ArgBinding.XCom( + name = name, + taskId = checkNotNull(entry["task_id"] as? String) { "xcom arg binding '$name' has no task_id" }, + mapIndex = (entry["map_index"] as? Number)?.toInt() ?: -1, + elementIndex = (entry["element_index"] as? Number)?.toInt(), + ) + else -> error("Unsupported arg binding kind '$kind' for argument '$name'") + } + } +} diff --git a/java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/execution/Task.kt b/java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/execution/Task.kt index 9f1a753704665..1c189061ae503 100644 --- a/java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/execution/Task.kt +++ b/java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/execution/Task.kt @@ -75,7 +75,7 @@ internal object TaskRunner { client: Client, ): Any { val definition = - bundle.dags[request.ti.dagId]?.tasks[request.ti.taskId]?.definition + bundle.taskDef(request.ti.dagId, request.ti.taskId)?.definition ?: return TaskResult.of(TaskState.State.REMOVED) val instance = try { diff --git a/java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/internal/ArgValues.kt b/java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/internal/ArgValues.kt new file mode 100644 index 0000000000000..80e42a259a2a8 --- /dev/null +++ b/java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/internal/ArgValues.kt @@ -0,0 +1,294 @@ +/* + * 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. + */ + +@file:Suppress("PLATFORM_CLASS_MAPPED_TO_KOTLIN") + +package org.apache.airflow.sdk.internal + +import com.fasterxml.jackson.databind.ObjectMapper +import com.fasterxml.jackson.databind.json.JsonMapper +import org.apache.airflow.sdk.Client +import org.apache.airflow.sdk.MissingXComException +import org.apache.airflow.sdk.TaskInput +import org.apache.airflow.sdk.execution.ArgBinding +import org.apache.airflow.sdk.execution.Logger +import java.lang.reflect.Field +import java.lang.reflect.Type + +/** + * @suppress + * + * Resolves a task's data parameters from the arg bindings the supervisor + * delivered, and decodes their raw wire values into the declared types. Public + * so that processor-generated task classes can call it; not user-facing API. + * + * The bindings come from the Python `@task.stub` call site, which is also the + * graph the scheduler ordered the run by. Flat data parameters resolve the + * binding at their position (through [TaskArgs]); [TaskInput] fields resolve + * bindings by name. + */ +object ArgValues { + private val mapper: ObjectMapper = JsonMapper.builder().build().findAndRegisterModules() + private val logger = Logger(ArgValues::class) + + /** + * Materializes a [TaskInput] with every field bound by the argument name it + * claims. + * + * The single populator behind both authoring APIs — the annotation processor + * emits a call to it for a `@Builder.Task` [TaskInput] parameter, and + * [org.apache.airflow.sdk.InputTask] calls it before handing the input to a + * task written against the interface. + * + * Every field has to find an argument: a field nothing binds means the input + * and the stub signature disagree. Neither direction of that disagreement + * fails the task, because a field binds by name: a field nothing supplies + * keeps its Java default, and an argument no field claims changes nothing + * the task reads. Both are logged so the mismatch stays visible. + * + * @throws IllegalArgumentException if the input cannot be populated. + * @throws MissingXComException if a field's argument resolves to nothing and + * the field is primitive. + */ + @JvmStatic + fun bindInput( + client: Client, + type: Class, + ): I { + val input = newInput(type) + val arguments = ArgIndex(client.argBindings) + val unfilled = mutableListOf() + bindableFields(type).forEach { field -> + val argName = argNameOf(field) + val pinned = isPinned(field) + val binding = arguments.claim(argName, pinned) + if (binding == null) { + unfilled += unfilledField(arguments, field, argName, pinned) + } else { + field.set(input, resolveClaimed(client, binding, field)) + } + } + warnUnfilled(client, type, unfilled, arguments) + warnUnclaimed(client, type, arguments) + return input + } + + /** + * Reports the fields the call site supplied nothing for. Each keeps its Java + * default, so the task runs on a value nobody passed. + */ + private fun warnUnfilled( + client: Client, + type: Class<*>, + unfilled: List, + arguments: ArgIndex, + ) { + if (unfilled.isEmpty()) return + logger.warning( + "Task handler declares argument(s) the Dag's call did not pass", + mapOf( + "task_id" to client.details.ti.taskId, + "input" to type.simpleName, + "declared_not_passed" to unfilled, + "passed" to arguments.passed(), + ), + ) + } + + /** + * Reports the arguments the call site passed that no field took. A captured + * default is not one of them: the Dag author did not write it, so a field + * has no reason to exist for it. + */ + private fun warnUnclaimed( + client: Client, + type: Class<*>, + arguments: ArgIndex, + ) { + val unclaimed = arguments.unclaimed() + if (unclaimed.isEmpty()) return + logger.warning( + "Dag's call passed argument(s) the task handler does not declare", + mapOf( + "task_id" to client.details.ti.taskId, + "input" to type.simpleName, + "passed_not_declared" to unclaimed, + "declared" to bindableFields(type).map(::argNameOf), + ), + ) + } + + /** + * Resolves one data parameter into [type], passing null through. Backs + * [TaskArgs]; a parameter that cannot be null goes through + * [TaskArgs.require], which turns null into [missing]. + * + * @param binding The argument [TaskArgs] holds for that position. + */ + internal fun valueAt( + client: Client, + binding: ArgBinding, + type: Type, + ): Any? = decode(client.resolveBinding(binding), type) + + /** + * Builds the failure for a binding that resolved to nothing where a value is + * required, naming [target] — the stub argument, or the [TaskInput] field + * that claimed it. + */ + internal fun missing( + binding: ArgBinding, + taskId: String, + target: String = binding.name, + ): MissingXComException = + when (binding) { + is ArgBinding.XCom -> MissingXComException(binding.taskId, target) + is ArgBinding.Literal -> + MissingXComException( + "Task parameter '$target' of task '$taskId' is bound to a null literal, but has a primitive " + + "type that cannot be null; declare a boxed type (e.g. Integer instead of int) to receive null.", + ) + } + + /** + * Resolves one [TaskInput] field from the argument it claimed. + * + * An argument that resolves to nothing is a value, not a mistake: a boxed or + * reference field takes null, and a primitive field cannot, so it fails with + * a clear [MissingXComException]. + */ + private fun resolveClaimed( + client: Client, + binding: ArgBinding, + field: Field, + ): Any? { + if (!field.type.isPrimitive) return decode(client.resolveBinding(binding), field.genericType) + // The msgpack decoder yields boxed values, so a primitive field decodes + // into its wrapper and unboxes on assignment. + return decode(client.resolveBinding(binding), field.type.kotlin.javaObjectType) + ?: throw missing(binding, client.details.ti.taskId, field.name) + } + + /** + * Names a field the call site supplied nothing for, saying which of the two + * reasons it was: no argument of that name, or two that the fold cannot tell + * apart, where `@ArgName` is the way to say which one is meant. + */ + private fun unfilledField( + arguments: ArgIndex, + field: Field, + argName: String, + pinned: Boolean, + ): String = + if (arguments.foldIsShared(argName, pinned)) { + "${field.name} (argument '$argName' matches more than one passed argument differing only " + + "in case or underscores; add @ArgName)" + } else { + "${field.name} (argument '$argName')" + } + + /** + * Decodes a raw wire value into [type], which carries the full generic type + * where the declared type has one, so an element type survives the decode. + */ + internal fun decode( + value: Any?, + type: Type, + ): Any? { + if (value == null) return null + if (type is Class<*>) { + if (type.isInstance(value)) return value + // The msgpack decoder yields Long for wire integers and Double for wire + // floats, so widen numerics via Number instead of casting. + if (value is Number) numberConverter(type)?.let { return it(value) } + } + // Structured wire values (maps, lists) convert into the declared POJO or + // collection type; unknown fields fail the task, mirroring the Go SDK's + // strict decode of task inputs. + return mapper.convertValue(value, mapper.constructType(type)) + } + + /** + * The run's bindings, addressable by the exact argument name and by the + * cross-language fold, and tracking which of them a field has taken. A fold + * two arguments share matches neither: either could be the one meant, and + * handing a field the wrong value is worse than failing. + */ + private class ArgIndex( + private val bindings: List, + ) { + private val byName = bindings.associateBy { it.name } + private val byFold = mutableMapOf() + private val sharedFolds = mutableSetOf() + private val claimed = mutableSetOf() + + init { + bindings.forEach { binding -> + val fold = foldArgName(binding.name) + if (byFold.put(fold, binding) != null) sharedFolds += fold + } + } + + /** + * Takes the argument a field claims, marking it claimed. An + * `@ArgName`-pinned name is taken literally, which is what makes the + * annotation an escape hatch for a Python name no Java identifier folds + * to. + */ + fun claim( + name: String, + pinned: Boolean, + ): ArgBinding? { + val match = byName[name] ?: foldMatch(name, pinned) + return match?.also { claimed += it.name } + } + + private fun foldMatch( + name: String, + pinned: Boolean, + ): ArgBinding? { + if (pinned) return null + val fold = foldArgName(name) + return if (fold in sharedFolds) null else byFold[fold] + } + + /** Whether [name] reaches two arguments the fold cannot tell apart. */ + fun foldIsShared( + name: String, + pinned: Boolean, + ): Boolean = !pinned && name !in byName && foldArgName(name) in sharedFolds + + /** Explicitly passed argument names no field took. */ + fun unclaimed(): List = bindings.filterNot { it.fromDefault || it.name in claimed }.map { it.name } + + /** Every argument name the call site passed, captured defaults included. */ + fun passed(): List = bindings.map { it.name } + } + + private fun numberConverter(type: Class<*>): ((Number) -> Any)? = + when (type) { + java.lang.Byte::class.java -> { n -> n.toByte() } + java.lang.Short::class.java -> { n -> n.toShort() } + java.lang.Integer::class.java -> { n -> n.toInt() } + java.lang.Long::class.java -> { n -> n.toLong() } + java.lang.Float::class.java -> { n -> n.toFloat() } + java.lang.Double::class.java -> { n -> n.toDouble() } + else -> null + } +} diff --git a/java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/internal/InputTypes.kt b/java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/internal/InputTypes.kt new file mode 100644 index 0000000000000..fad312ffcf354 --- /dev/null +++ b/java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/internal/InputTypes.kt @@ -0,0 +1,176 @@ +/* + * 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 org.apache.airflow.sdk.internal + +import org.apache.airflow.sdk.ArgName +import org.apache.airflow.sdk.InputTask +import org.apache.airflow.sdk.Task +import org.apache.airflow.sdk.TaskInput +import java.lang.reflect.Constructor +import java.lang.reflect.Field +import java.lang.reflect.Modifier +import java.lang.reflect.ParameterizedType +import java.lang.reflect.Type +import java.util.concurrent.ConcurrentHashMap + +/** + * Resolves the [TaskInput] type argument that [taskClass] bound to + * [InputTask]'s type parameter. + * + * @throws IllegalArgumentException if the type argument is not a concrete + * [TaskInput] class. + */ +internal fun resolveInputType(taskClass: Class<*>): Class = + findInputType(taskClass) + ?: throw IllegalArgumentException( + "Task class ${taskClass.name} implements InputTask with an input type that cannot be resolved; " + + "declare a concrete type argument, e.g. 'implements InputTask'", + ) + +/** + * Checks that an [InputTask]'s declared input can be bound, before the task + * ever runs — a mis-declared input then fails while the [org.apache.airflow.sdk.Bundle] + * is being built rather than mid-run. A plain [Task] declares no input and passes. + * + * The annotation processor enforces the same rules at compile time for the + * [TaskInput] a `@Builder.Task` method declares. + * + * @throws IllegalArgumentException if the declared input type is unresolvable + * or cannot be populated. + */ +internal fun validateTaskInput(definition: Class) { + if (!InputTask::class.java.isAssignableFrom(definition)) return + val inputType = resolveInputType(definition) + requirePublicNoArgConstructor(inputType) + bindableFields(inputType) +} + +/** + * Every field of a [TaskInput] that binds an argument: each public non-final + * instance field, its own and inherited. + * + * Memoized, since the answer is fixed by the class: a worker binds the same + * input once per task instance it runs, and reflecting over the whole + * hierarchy each time buys nothing. + * + * @throws IllegalArgumentException if any instance field cannot be assigned, + * which would leave an argument silently unbound, or if two fields claim + * argument names that [foldArgName] cannot tell apart. + */ +internal fun bindableFields(inputType: Class<*>): List = + bindableFieldsByType.computeIfAbsent(inputType) { collectBindableFields(it) } + +private val bindableFieldsByType = ConcurrentHashMap, List>() + +private fun collectBindableFields(inputType: Class<*>): List { + val fields = mutableListOf() + var current: Class<*>? = inputType + while (current != null && current != Any::class.java) { + current.declaredFields + .filterNot { Modifier.isStatic(it.modifiers) || it.isSynthetic } + .forEach { field -> + require(Modifier.isPublic(field.modifiers) && !Modifier.isFinal(field.modifiers)) { + "TaskInput field ${inputType.simpleName}.${field.name} must be public and non-final " + + "so the SDK can assign its binding" + } + // The declaring class may be package-private even though the field is + // public, which reflection from the SDK needs opening. + field.isAccessible = true + fields += field + } + current = current.superclass + } + requireDistinctArgNames(inputType, fields) + return fields +} + +/** The argument name a field claims: its [ArgName] value, or its own name. */ +internal fun argNameOf(field: Field): String = field.getAnnotation(ArgName::class.java)?.value ?: field.name + +/** Whether [ArgName] pinned this field's argument name, which forbids the fold. */ +internal fun isPinned(field: Field): Boolean = field.isAnnotationPresent(ArgName::class.java) + +/** + * @suppress + * + * Reduces an argument name to the token that matches across languages — + * lowercased with underscores removed, the same rule the Go and TypeScript + * SDKs fold by, so one Python signature binds identically in all three. + * + * Public so the annotation processor can reject a clash at compile time with + * the same rule the runtime binds by; not user-facing API. + */ +fun foldArgName(name: String): String = name.replace("_", "").lowercase() + +private fun requireDistinctArgNames( + inputType: Class<*>, + fields: List, +) { + val claimed = mutableMapOf() + fields.forEach { field -> + val previous = claimed.put(foldArgName(argNameOf(field)), field.name) + require(previous == null) { + "TaskInput fields ${inputType.simpleName}.$previous and ${inputType.simpleName}.${field.name} " + + "claim argument names that differ only in case or underscores, which the fold cannot tell " + + "apart; rename one of them" + } + } +} + +/** Instantiates a [TaskInput] for the SDK to populate. */ +@Suppress("UNCHECKED_CAST") +internal fun newInput(inputType: Class): I = requirePublicNoArgConstructor(inputType).newInstance() as I + +/** Memoized alongside [bindableFields], and for the same reason. */ +private val constructorsByType = ConcurrentHashMap, Constructor<*>>() + +private fun requirePublicNoArgConstructor(inputType: Class<*>): Constructor<*> = + constructorsByType.computeIfAbsent(inputType) { + val constructor = it.declaredConstructors.firstOrNull { c -> c.parameterCount == 0 } + require(constructor != null && Modifier.isPublic(constructor.modifiers)) { + "TaskInput class ${it.simpleName} needs a public no-argument constructor" + } + // The class may be package-private even though its constructor is public, + // which reflection from the SDK needs opening. + constructor.also { c -> c.isAccessible = true } + } + +/** + * Walks [type]'s supertypes for the [InputTask] type argument. A type variable + * yields null: only the class that fixes it to a concrete [TaskInput] can say + * what to bind. + */ +private fun findInputType(type: Type?): Class? = + when (type) { + is ParameterizedType -> + if (type.rawType == InputTask::class.java) { + type.actualTypeArguments.firstOrNull().asTaskInputClass() + } else { + findInputType(type.rawType) + } + is Class<*> -> + (type.genericInterfaces.asSequence() + sequenceOf(type.genericSuperclass)) + .firstNotNullOfOrNull { findInputType(it) } + else -> null + } + +@Suppress("UNCHECKED_CAST") +private fun Type?.asTaskInputClass(): Class? = + (this as? Class<*>)?.takeIf { TaskInput::class.java.isAssignableFrom(it) } as Class? diff --git a/java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/internal/Registrar.kt b/java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/internal/Registrar.kt new file mode 100644 index 0000000000000..3220ae78c02aa --- /dev/null +++ b/java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/internal/Registrar.kt @@ -0,0 +1,36 @@ +/* + * 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 org.apache.airflow.sdk.internal + +/** + * @suppress + * + * Names the registrar generated for a class of `@Builder.TaskHandler` + * methods: a top-level class in the handler class's package, named after that + * class and any class enclosing it, so `Outer.Inner` is served by + * `Outer_InnerHandlers`. + * + * Public so the annotation processor emits the name the runtime looks up; not + * user-facing API. + * + * @param binaryName Binary name of the handler class, as `Class.getName` + * returns it. + */ +fun registrarName(binaryName: String): String = "${binaryName.replace('$', '_')}Handlers" diff --git a/java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/internal/TaskArgs.kt b/java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/internal/TaskArgs.kt new file mode 100644 index 0000000000000..31c840662c386 --- /dev/null +++ b/java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/internal/TaskArgs.kt @@ -0,0 +1,142 @@ +/* + * 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 org.apache.airflow.sdk.internal + +import org.apache.airflow.sdk.Client +import org.apache.airflow.sdk.Context +import org.apache.airflow.sdk.MissingXComException +import org.apache.airflow.sdk.execution.ArgBinding + +/** + * @suppress + * + * A task's arguments addressed **by position**, in the stub signature's + * declaration order. + * + * This is what the annotation processor emits for a task method's flat data + * parameters, and it is not user-facing API: positional access is safe in code + * the processor writes and type-checks against the method signature, and wrong + * to ask a Dag author to track by hand. A task written against the interface + * declares an [org.apache.airflow.sdk.TaskInput] instead. + * + * ```java + * TaskArgs args = TaskArgs.of(context, client, 2); + * long rows = args.require(0, Long.class); + * List regions = args.get(1, new TypeRef>() {}); + * ``` + */ +class TaskArgs private constructor( + private val context: Context, + private val client: Client, + private val arguments: List, +) { + companion object { + /** + * Opens a positional view over the arguments bound for this run, for a + * task method declaring [declared] data parameters. + * + * The counts must match. Positions carry the whole meaning of a flat + * binding, so a call site that bound a different number of arguments than + * the method takes has already shifted them: too few leaves a parameter + * unbound, and too many means the extras — or the ones ahead of them — + * are not the arguments the method believes it is reading. + * + * A parameter the call site omitted still arrives, carrying the stub + * signature's default. A method that does not declare it is not reading + * shifted arguments, so those are dropped before the counts are compared. + * + * @throws IllegalStateException if the call site bound a different number + * of arguments than the task declares. + */ + @JvmStatic + fun of( + context: Context, + client: Client, + declared: Int, + ): TaskArgs { + val bound = client.argBindings + val arguments = if (bound.size == declared) bound else bound.filterNot { it.fromDefault } + check(arguments.size == declared) { + val defaults = + if (arguments.size == bound.size) { + "" + } else { + "; ${arguments.size} remain after dropping captured defaults" + } + "Task '${context.ti.taskId}' declares $declared data parameter(s) " + + "but the stub call bound ${bound.size} argument(s)$defaults" + } + return TaskArgs(context, client, arguments) + } + } + + /** + * Resolves the argument bound at [position] into [type], passing null + * through. + * + * @throws org.apache.airflow.sdk.ApiError if the underlying XCom read fails. + */ + fun get( + position: Int, + type: Class, + ): T? = type.cast(ArgValues.valueAt(client, arguments[position], type)) + + /** + * Resolves the argument bound at [position] into the generic [type], passing + * null through. + * + * @throws org.apache.airflow.sdk.ApiError if the underlying XCom read fails. + */ + @Suppress("UNCHECKED_CAST") + fun get( + position: Int, + type: TypeRef, + ): T? = ArgValues.valueAt(client, arguments[position], type.type) as T? + + /** + * Resolves the argument bound at [position] into [type], which must not be + * null. + * + * @throws MissingXComException if the binding resolves to nothing — a null + * literal, or an upstream that pushed no XCom. + * @throws org.apache.airflow.sdk.ApiError if the underlying XCom read fails. + */ + fun require( + position: Int, + type: Class, + ): T = get(position, type) ?: throw missingAt(position) + + /** + * Resolves the argument bound at [position] into the generic [type], which + * must not be null. + * + * @throws MissingXComException if the binding resolves to nothing — a null + * literal, or an upstream that pushed no XCom. + * @throws org.apache.airflow.sdk.ApiError if the underlying XCom read fails. + */ + fun require( + position: Int, + type: TypeRef, + ): T = get(position, type) ?: throw missingAt(position) + + // The stub signature's own parameter name is the clearest label for a failure + // here: it is what the Dag author has to change. + private fun missingAt(at: Int) = ArgValues.missing(arguments[at], context.ti.taskId) +} diff --git a/java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/internal/TypeRef.kt b/java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/internal/TypeRef.kt new file mode 100644 index 0000000000000..55f1a5d01dfa7 --- /dev/null +++ b/java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/internal/TypeRef.kt @@ -0,0 +1,47 @@ +/* + * 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 org.apache.airflow.sdk.internal + +import java.lang.reflect.ParameterizedType +import java.lang.reflect.Type + +/** + * @suppress + * + * Carries a full generic type into [TaskArgs], which a `Class` literal cannot + * express. Subclass it anonymously so the type argument survives erasure on + * the subclass's signature: + * + * ```java + * args.require(0, new TypeRef>() {}); + * ``` + * + * Emitted by the annotation processor for a data parameter whose declared type + * has type arguments; not user-facing API. The SDK owns this rather than + * reusing Jackson's `TypeReference` so that Jackson stays off a consumer's + * compile classpath. + */ +abstract class TypeRef protected constructor() { + internal val type: Type = + (javaClass.genericSuperclass as? ParameterizedType)?.actualTypeArguments?.firstOrNull() + ?: throw IllegalArgumentException( + "TypeRef needs a concrete type argument on an anonymous subclass, e.g. new TypeRef>() {}", + ) +} diff --git a/java-sdk/sdk/src/test/kotlin/org/apache/airflow/sdk/ArgTestSupport.kt b/java-sdk/sdk/src/test/kotlin/org/apache/airflow/sdk/ArgTestSupport.kt new file mode 100644 index 0000000000000..1fff219370b98 --- /dev/null +++ b/java-sdk/sdk/src/test/kotlin/org/apache/airflow/sdk/ArgTestSupport.kt @@ -0,0 +1,96 @@ +/* + * 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 org.apache.airflow.sdk + +import org.apache.airflow.sdk.execution.comm.ConnectionResult +import org.apache.airflow.sdk.execution.comm.StartupDetails +import org.apache.airflow.sdk.execution.comm.TIRunContext +import org.apache.airflow.sdk.execution.comm.VariableResult +import org.apache.airflow.sdk.execution.comm.XComResult +import org.apache.airflow.sdk.execution.comm.TaskInstance as CommTaskInstance + +/** Records getXCom calls and serves canned values keyed by task id. */ +internal class FakeXComTransport( + val xcoms: Map = emptyMap(), +) : org.apache.airflow.sdk.execution.Client { + val pulls = mutableListOf>() + + override fun getConnection(id: String): ConnectionResult = throw NotImplementedError() + + override fun getVariable(key: String): VariableResult = throw NotImplementedError() + + override fun setVariable( + key: String, + value: String, + description: String?, + ) = throw NotImplementedError() + + override fun deleteVariable(key: String) = throw NotImplementedError() + + override fun getXCom( + key: String, + dagId: String, + taskId: String, + runId: String, + mapIndex: Int?, + includePriorDates: Boolean, + ): XComResult { + pulls += taskId to mapIndex + return XComResult().also { + it.key = key + it.value = xcoms[taskId] + } + } + + override fun setXCom( + key: String, + value: Any, + dagId: String, + taskId: String, + runId: String, + mapIndex: Int, + ) = throw NotImplementedError() +} + +internal fun startupDetails(argBindings: List>?): StartupDetails = + StartupDetails().also { details -> + details.ti = + CommTaskInstance().also { + it.dagId = "d" + it.runId = "r" + it.taskId = "t" + it.tryNumber = 1 + } + details.tiContext = TIRunContext().also { it.argBindings = argBindings } + } + +internal fun clientWith( + argBindings: List>?, + xcoms: Map = emptyMap(), +): Pair { + val transport = FakeXComTransport(xcoms) + return Client(startupDetails(argBindings), transport) to transport +} + +internal fun taskContext(): Context = + Context( + dagRun = DagRun("d", "r", null, null, null, null, null, emptyMap()), + ti = TaskInstance("d", "r", "t", null, 1), + ) diff --git a/java-sdk/sdk/src/test/kotlin/org/apache/airflow/sdk/ArgValuesTest.kt b/java-sdk/sdk/src/test/kotlin/org/apache/airflow/sdk/ArgValuesTest.kt new file mode 100644 index 0000000000000..16eaa161fbdd0 --- /dev/null +++ b/java-sdk/sdk/src/test/kotlin/org/apache/airflow/sdk/ArgValuesTest.kt @@ -0,0 +1,303 @@ +/* + * 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. + */ + +@file:Suppress("PLATFORM_CLASS_MAPPED_TO_KOTLIN") + +package org.apache.airflow.sdk + +import org.apache.airflow.sdk.execution.Level +import org.apache.airflow.sdk.execution.LogSender +import org.apache.airflow.sdk.internal.ArgValues +import org.junit.jupiter.api.Assertions.assertEquals +import org.junit.jupiter.api.Assertions.assertNull +import org.junit.jupiter.api.Assertions.assertThrows +import org.junit.jupiter.api.Assertions.assertTrue +import org.junit.jupiter.api.DisplayName +import org.junit.jupiter.api.Test + +class ScoreInput : TaskInput { + /** Pinned, so it binds `region_code` and never folds. */ + @JvmField + @ArgName("region_code") + var region: String? = null + + /** Unpinned camelCase: binds `run_label` through the fold. */ + @JvmField + var runLabel: String? = null + + @JvmField + var threshold: Double = 0.0 +} + +class FoldedInput : TaskInput { + @JvmField + var regionCode: String? = null +} + +/** One primitive field, so resolving to nothing is the only failure possible. */ +class ThresholdInput : TaskInput { + @JvmField + var threshold: Double = 0.0 +} + +/** The boxed twin of [ThresholdInput], which can take the null. */ +class BoxedThresholdInput : TaskInput { + @JvmField + var threshold: Double? = null +} + +/** A generic field: its element type has to survive the decode. */ +class TagsInput : TaskInput { + @JvmField + var tags: List? = null +} + +class CollidingInput : TaskInput { + @JvmField + @ArgName("region_code") + var first: String? = null + + @JvmField + @ArgName("regionCode") + var second: String? = null +} + +private class CollidingInputTask : InputTask { + override fun execute( + context: Context, + client: Client, + input: CollidingInput, + ) = Unit +} + +private fun bind( + type: Class, + bindings: List>, + xcoms: Map = emptyMap(), +): I { + val (client, _) = clientWith(bindings, xcoms) + return ArgValues.bindInput(client, type) +} + +private fun literal( + name: String, + value: Any?, + fromDefault: Boolean = false, +): Map = mapOf("kind" to "literal", "name" to name, "value" to value, "from_default" to fromDefault) + +internal class ArgValuesTest { + @Test + @DisplayName("Should bind fields by their argument name, whatever the call-site order") + fun shouldBindFieldsByArgName() { + val input = + bind( + ScoreInput::class.java, + listOf( + mapOf("kind" to "xcom", "name" to "threshold", "task_id" to "upstream"), + literal("run_label", "nightly"), + literal("region_code", "emea"), + ), + xcoms = mapOf("upstream" to 0.5), + ) + + assertEquals("emea", input.region) + assertEquals("nightly", input.runLabel) + assertEquals(0.5, input.threshold) + } + + @Test + @DisplayName("Should match a camelCase field to a snake_case argument through the fold") + fun shouldFoldSnakeCaseArgument() { + val input = bind(FoldedInput::class.java, listOf(literal("region_code", "emea"))) + + assertEquals("emea", input.regionCode) + } + + @Test + @DisplayName("Should keep the element type of a generic field") + fun shouldKeepGenericFieldElementType() { + val input = bind(TagsInput::class.java, listOf(literal("tags", listOf("a", "b")))) + + assertEquals(listOf("a", "b"), input.tags) + } + + @Test + @DisplayName("Should still bind an exact name that another argument folds onto") + fun shouldPreferExactNameOverAmbiguousFold() { + val input = + bind( + FoldedInput::class.java, + listOf(literal("regionCode", "emea"), literal("region_code", "apac")), + ) + + assertEquals("emea", input.regionCode) + } + + @Test + @DisplayName("Should reject a TaskInput whose fields claim names that fold alike") + fun shouldRejectFieldsWithCollidingFolds() { + val error = + assertThrows(IllegalArgumentException::class.java) { + TaskDef("t", CollidingInputTask::class.java) + } + + assertEquals( + "TaskInput fields CollidingInput.first and CollidingInput.second claim argument names that " + + "differ only in case or underscores, which the fold cannot tell apart; rename one of them", + error.message, + ) + } + + @Test + @DisplayName("Should warn in each direction at once and default the unfilled field") + fun shouldWarnOnFieldMatchingNoArgument() { + LogSender.messages.clear() + + // The field and the argument miss each other, so both directions report. + val input = bind(FoldedInput::class.java, listOf(literal("threshold", 0.5))) + + assertNull(input.regionCode) + val unfilled = LogSender.messages.single { it.event == "Task handler declares argument(s) the Dag's call did not pass" } + assertEquals(listOf("regionCode (argument 'regionCode')"), unfilled.arguments["declared_not_passed"]) + assertEquals(listOf("threshold"), unfilled.arguments["passed"]) + val unclaimed = LogSender.messages.single { it.event == "Dag's call passed argument(s) the task handler does not declare" } + assertEquals(listOf("threshold"), unclaimed.arguments["passed_not_declared"]) + assertEquals(listOf("regionCode"), unclaimed.arguments["declared"]) + } + + @Test + @DisplayName("Should warn when an @ArgName-pinned name is not among the arguments") + fun shouldWarnWhenPinnedNameDoesNotFold() { + LogSender.messages.clear() + + // 'region' is pinned to region_code, which the camelCase argument cannot reach. + val input = + bind( + ScoreInput::class.java, + listOf( + literal("regionCode", "emea"), + literal("run_label", "nightly"), + literal("threshold", 0.5), + ), + ) + + assertNull(input.region) + assertEquals("nightly", input.runLabel) + val unfilled = LogSender.messages.single { it.event == "Task handler declares argument(s) the Dag's call did not pass" } + assertEquals(listOf("region (argument 'region_code')"), unfilled.arguments["declared_not_passed"]) + } + + @Test + @DisplayName("Should warn rather than guess when two arguments fold to the field's name") + fun shouldWarnOnAmbiguousFold() { + LogSender.messages.clear() + + val input = + bind( + FoldedInput::class.java, + listOf(literal("region_code", "emea"), literal("regioncode", "apac")), + ) + + assertNull(input.regionCode) + val message = LogSender.messages.single { it.event == "Task handler declares argument(s) the Dag's call did not pass" } + assertEquals( + listOf( + "regionCode (argument 'regionCode' matches more than one passed argument differing only " + + "in case or underscores; add @ArgName)", + ), + message.arguments["declared_not_passed"], + ) + } + + @Test + @DisplayName("Should bind and warn when the call site passes an argument no field claims") + fun shouldWarnOnUnclaimedArgument() { + LogSender.messages.clear() + + val input = + bind( + FoldedInput::class.java, + listOf(literal("region_code", "emea"), literal("extra", 1L)), + ) + + assertEquals("emea", input.regionCode) + val message = LogSender.messages.single { it.level == Level.WARNING } + assertEquals("Dag's call passed argument(s) the task handler does not declare", message.event) + assertEquals(listOf("extra"), message.arguments["passed_not_declared"]) + assertEquals(listOf("regionCode"), message.arguments["declared"]) + assertEquals("FoldedInput", message.arguments["input"]) + } + + @Test + @DisplayName("Should stay quiet about an unclaimed argument the call site never passed") + fun shouldNotWarnOnUnclaimedCapturedDefault() { + LogSender.messages.clear() + + bind( + FoldedInput::class.java, + listOf(literal("region_code", "emea"), literal("extra", 1L, fromDefault = true)), + ) + + assertTrue(LogSender.messages.none { it.level == Level.WARNING }) { + "unexpected warnings: ${LogSender.messages.map { it.event }}" + } + } + + @Test + @DisplayName("Should give a boxed field null when its argument resolves to nothing") + fun shouldPassNullToBoxedField() { + val input = bind(BoxedThresholdInput::class.java, listOf(literal("threshold", null))) + + assertNull(input.threshold) + } + + @Test + @DisplayName("Should name the field when a primitive one is bound to a null literal") + fun shouldNameFieldBoundToNullLiteral() { + val error = + assertThrows(MissingXComException::class.java) { + bind(ThresholdInput::class.java, listOf(literal("threshold", null))) + } + + assertEquals( + "Task parameter 'threshold' of task 't' is bound to a null literal, but has a primitive type " + + "that cannot be null; declare a boxed type (e.g. Integer instead of int) to receive null.", + error.message, + ) + } + + @Test + @DisplayName("Should name the upstream when a primitive field's XCom was never pushed") + fun shouldNameUpstreamForPrimitiveField() { + val error = + assertThrows(MissingXComException::class.java) { + bind( + ThresholdInput::class.java, + listOf(mapOf("kind" to "xcom", "name" to "threshold", "task_id" to "upstream")), + ) + } + + assertEquals( + "Task parameter 'threshold' requires an XCom from task 'upstream', but none was pushed. " + + "This parameter has a primitive type that cannot be null; declare it with a boxed type " + + "(e.g. Integer instead of int) to receive null.", + error.message, + ) + } +} diff --git a/java-sdk/sdk/src/test/kotlin/org/apache/airflow/sdk/BundleTest.kt b/java-sdk/sdk/src/test/kotlin/org/apache/airflow/sdk/BundleTest.kt index 57050754e887e..207fe8b244248 100644 --- a/java-sdk/sdk/src/test/kotlin/org/apache/airflow/sdk/BundleTest.kt +++ b/java-sdk/sdk/src/test/kotlin/org/apache/airflow/sdk/BundleTest.kt @@ -24,6 +24,16 @@ import org.junit.jupiter.api.DisplayName import org.junit.jupiter.api.Test internal class BundleTest { + private class NoOp : Task { + override fun execute( + context: Context, + client: Client, + ) = Unit + } + + /** A class of handlers that the processor generated [BundleTest_NestedHandlers] for. */ + class Nested + @Test @DisplayName("Should index dags by dagId") fun shouldIndexDagsByDagId() { @@ -44,4 +54,130 @@ internal class BundleTest { Assertions.assertEquals("Dags in bundle have duplicate ID: dag", error.message) } + + @Test + @DisplayName("Should find the registrar generated for a nested handler class") + fun shouldFindRegistrarOfNestedHandlerClass() { + val bundle = Bundle().register(Nested::class.java) + + val etl = bundle.taskHandlers.getValue("etl") + Assertions.assertEquals(listOf("etl"), bundle.taskHandlers.keys.toList()) + Assertions.assertEquals(listOf("score"), etl.tasks.keys.toList()) + } + + @Test + @DisplayName("Should name the registrar it looked for when there is none") + fun shouldNameTheRegistrarItLookedFor() { + val error = + Assertions.assertThrows(IllegalArgumentException::class.java) { + Bundle().register(NoOp::class.java) + } + + Assertions.assertTrue( + error.message!!.startsWith( + "No generated registrar org.apache.airflow.sdk.BundleTest_NoOpHandlers for ", + ), + error.message, + ) + } + + @Test + @DisplayName("Should reject a Dag whose ID task handlers already hold") + fun shouldRejectDagWhoseIdTaskHandlersHold() { + val bundle = Bundle().register("etl", "score", NoOp::class.java) + + val error = + Assertions.assertThrows(IllegalArgumentException::class.java) { + bundle.register(DagDef("etl")) + } + + Assertions.assertEquals( + "Dag 'etl' already has registered task handlers; a Dag declared in Java owns its own " + + "tasks, so one Dag ID cannot have both", + error.message, + ) + } + + @Test + @DisplayName("Should reject task handlers for a Dag ID declared in Java") + fun shouldRejectTaskHandlersForJavaDeclaredDag() { + val bundle = Bundle().register(DagDef("etl")) + + val error = + Assertions.assertThrows(IllegalArgumentException::class.java) { + bundle.register("etl", "score", NoOp::class.java) + } + + Assertions.assertEquals( + "Dag 'etl' is declared in Java; attach its tasks with addTask(...) rather than " + + "registering task handlers for them", + error.message, + ) + } + + @Test + @DisplayName("Should keep accepting more handlers for a Dag the Python file owns") + fun shouldAcceptMoreHandlersForSameDag() { + val bundle = + Bundle() + .register("etl", "score", NoOp::class.java) + .register("etl", "report", NoOp::class.java) + + Assertions.assertEquals( + setOf("score", "report"), + bundle.taskHandlers + .getValue("etl") + .tasks.keys, + ) + } + + @Test + @DisplayName("Should find a task whichever side registered its Dag") + fun shouldFindTaskFromEitherSide() { + val declared = DagDef("java_etl").addTask("extract", NoOp::class.java) + val bundle = Bundle().register(declared).register("py_etl", "score", NoOp::class.java) + + Assertions.assertEquals("extract", bundle.taskDef("java_etl", "extract")?.id) + Assertions.assertEquals("score", bundle.taskDef("py_etl", "score")?.id) + Assertions.assertNull(bundle.taskDef("java_etl", "score")) + Assertions.assertNull(bundle.taskDef("absent", "extract")) + } + + @Test + @DisplayName("Should reject every register once serving has started") + fun shouldRejectRegisterAfterServing() { + val bundle = Bundle() + bundle.finalizeRegistration() + + val message = "Server.serve has already been called; register everything before serve" + listOf<() -> Unit>( + { bundle.register(DagDef("etl")) }, + { bundle.register(Nested::class.java) }, + { bundle.register("etl", "score", NoOp::class.java) }, + ).forEach { register -> + val error = Assertions.assertThrows(IllegalStateException::class.java, register) + Assertions.assertEquals(message, error.message) + } + } +} + +/** + * Stands in for the registrar the annotation processor generates beside + * [BundleTest.Nested], to pin the name [Bundle.register] looks up. + */ +@Suppress("ktlint:standard:class-naming", "ClassName") +class BundleTest_NestedHandlers { + companion object { + @JvmStatic + fun registerInto(bundle: Bundle) { + bundle.register("etl", "score", NoOpHandler::class.java) + } + } + + class NoOpHandler : Task { + override fun execute( + context: Context, + client: Client, + ) = Unit + } } diff --git a/java-sdk/sdk/src/test/kotlin/org/apache/airflow/sdk/InputTaskTest.kt b/java-sdk/sdk/src/test/kotlin/org/apache/airflow/sdk/InputTaskTest.kt new file mode 100644 index 0000000000000..b33abcfba0df1 --- /dev/null +++ b/java-sdk/sdk/src/test/kotlin/org/apache/airflow/sdk/InputTaskTest.kt @@ -0,0 +1,167 @@ +/* + * 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. + */ + +@file:Suppress("PLATFORM_CLASS_MAPPED_TO_KOTLIN") + +package org.apache.airflow.sdk + +import org.junit.jupiter.api.Assertions.assertEquals +import org.junit.jupiter.api.Assertions.assertThrows +import org.junit.jupiter.api.DisplayName +import org.junit.jupiter.api.Test + +/** The bindings a keyword call site delivers for [SummaryInput]. */ +private val NAMED_BINDINGS = + listOf( + mapOf("kind" to "literal", "name" to "region_code", "value" to "emea"), + mapOf("kind" to "xcom", "name" to "threshold", "task_id" to "upstream"), + ) + +class SummaryInput : TaskInput { + @JvmField + @ArgName("region_code") + var region: String? = null + + @JvmField + var threshold: Double = 0.0 +} + +private class Summarize : InputTask { + var received: SummaryInput? = null + + override fun execute( + context: Context, + client: Client, + input: SummaryInput, + ) { + received = input + } +} + +private abstract class SummarizeBase : InputTask { + var received: SummaryInput? = null +} + +private class InheritedSummarize : SummarizeBase() { + override fun execute( + context: Context, + client: Client, + input: SummaryInput, + ) { + received = input + } +} + +private class Unresolved : InputTask { + override fun execute( + context: Context, + client: Client, + input: I, + ) = Unit +} + +class HiddenFieldInput : TaskInput { + var hidden: String? = null +} + +private class HiddenFieldTask : InputTask { + override fun execute( + context: Context, + client: Client, + input: HiddenFieldInput, + ) = Unit +} + +class ConstructedInput( + @JvmField val value: String, +) : TaskInput + +private class ConstructedInputTask : InputTask { + override fun execute( + context: Context, + client: Client, + input: ConstructedInput, + ) = Unit +} + +internal class InputTaskTest { + @Test + @DisplayName("Should inject a TaskInput whose fields are bound by argument name") + fun shouldInjectNamedTaskInput() { + val (client, _) = clientWith(NAMED_BINDINGS, xcoms = mapOf("upstream" to 0.5)) + val task = Summarize() + + task.execute(taskContext(), client) + + val input = requireNotNull(task.received) + assertEquals("emea", input.region) + assertEquals(0.5, input.threshold) + } + + @Test + @DisplayName("Should resolve the input type a superclass declared") + fun shouldResolveInheritedInputType() { + val (client, _) = clientWith(NAMED_BINDINGS, xcoms = mapOf("upstream" to 0.5)) + val task = InheritedSummarize() + + task.execute(taskContext(), client) + + assertEquals("emea", requireNotNull(task.received).region) + } + + @Test + @DisplayName("Should reject registering a task whose input type cannot be resolved") + fun shouldRejectUnresolvableInputType() { + val error = + assertThrows(IllegalArgumentException::class.java) { + TaskDef("t", Unresolved::class.java) + } + + assertEquals( + "Task class ${Unresolved::class.java.name} implements InputTask with an input type that cannot be " + + "resolved; declare a concrete type argument, e.g. 'implements InputTask'", + error.message, + ) + } + + @Test + @DisplayName("Should reject registering a task whose TaskInput field cannot be assigned") + fun shouldRejectUnassignableField() { + val error = + assertThrows(IllegalArgumentException::class.java) { + TaskDef("t", HiddenFieldTask::class.java) + } + + assertEquals( + "TaskInput field HiddenFieldInput.hidden must be public and non-final so the SDK can assign its binding", + error.message, + ) + } + + @Test + @DisplayName("Should reject registering a task whose TaskInput has no public no-argument constructor") + fun shouldRejectTaskInputWithoutNoArgConstructor() { + val error = + assertThrows(IllegalArgumentException::class.java) { + TaskDef("t", ConstructedInputTask::class.java) + } + + assertEquals("TaskInput class ConstructedInput needs a public no-argument constructor", error.message) + } +} diff --git a/java-sdk/sdk/src/test/kotlin/org/apache/airflow/sdk/TaskArgsTest.kt b/java-sdk/sdk/src/test/kotlin/org/apache/airflow/sdk/TaskArgsTest.kt new file mode 100644 index 0000000000000..ea7d9387a84c5 --- /dev/null +++ b/java-sdk/sdk/src/test/kotlin/org/apache/airflow/sdk/TaskArgsTest.kt @@ -0,0 +1,336 @@ +/* + * 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. + */ + +@file:Suppress("PLATFORM_CLASS_MAPPED_TO_KOTLIN") + +package org.apache.airflow.sdk + +import org.apache.airflow.sdk.internal.TaskArgs +import org.apache.airflow.sdk.internal.TypeRef +import org.junit.jupiter.api.Assertions.assertDoesNotThrow +import org.junit.jupiter.api.Assertions.assertEquals +import org.junit.jupiter.api.Assertions.assertNull +import org.junit.jupiter.api.Assertions.assertThrows +import org.junit.jupiter.api.DisplayName +import org.junit.jupiter.api.Test + +/** Fixes the type argument one level up, so the anonymous subclass carries none. */ +private abstract class NamesRef : TypeRef>() + +private fun argsWith( + bindings: List>?, + xcoms: Map = emptyMap(), + declared: Int = bindings?.size ?: 0, +): Pair { + val (client, transport) = clientWith(bindings, xcoms) + return TaskArgs.of(taskContext(), client, declared) to transport +} + +internal class TaskArgsTest { + @Test + @DisplayName("Should resolve a literal binding to its inline value without reading an XCom") + fun shouldResolveLiteralBinding() { + val (args, transport) = argsWith(listOf(mapOf("kind" to "literal", "name" to "x", "value" to 42L))) + + assertEquals(42L, args.get(0, java.lang.Long::class.java)) + assertEquals(emptyList>(), transport.pulls) + } + + @Test + @DisplayName("Should resolve an xcom binding by pulling the bound task's return value") + fun shouldResolveXComBinding() { + val (args, transport) = + argsWith( + listOf(mapOf("kind" to "xcom", "name" to "x", "task_id" to "upstream", "map_index" to -1L)), + xcoms = mapOf("upstream" to 7L), + ) + + assertEquals(7L, args.get(0, java.lang.Long::class.java)) + assertEquals(listOf("upstream" to null), transport.pulls) + } + + @Test + @DisplayName("Should keep bindings in stub-signature order") + fun shouldKeepBindingOrder() { + val (args, _) = + argsWith( + listOf( + mapOf("kind" to "literal", "name" to "b", "value" to 2L), + mapOf("kind" to "literal", "name" to "a", "value" to 1L), + ), + ) + + assertEquals(2L, args.get(0, java.lang.Long::class.java)) + assertEquals(1L, args.get(1, java.lang.Long::class.java)) + } + + @Test + @DisplayName("Should open a task declaring no data parameters when the supervisor sent no bindings") + fun shouldAcceptNoArguments() { + assertDoesNotThrow { argsWith(null) } + } + + @Test + @DisplayName("Should pass a non-negative bound map index to the XCom read") + fun shouldPassBoundMapIndex() { + val (args, transport) = + argsWith( + listOf(mapOf("kind" to "xcom", "name" to "x", "task_id" to "upstream", "map_index" to 2L)), + xcoms = mapOf("upstream" to 7L), + ) + + args.get(0, java.lang.Long::class.java) + + assertEquals(listOf("upstream" to 2), transport.pulls) + } + + @Test + @DisplayName("Should index into a list XCom when the binding has an element index") + fun shouldResolveElementIndex() { + val (args, _) = + argsWith( + listOf(mapOf("kind" to "xcom", "name" to "x", "task_id" to "upstream", "element_index" to 1L)), + xcoms = mapOf("upstream" to listOf("a", "b", "c")), + ) + + assertEquals("b", args.get(0, String::class.java)) + } + + @Test + @DisplayName("Should fail when an element index points into a non-list XCom") + fun shouldRejectElementIndexOnNonList() { + val (args, _) = + argsWith( + listOf(mapOf("kind" to "xcom", "name" to "x", "task_id" to "upstream", "element_index" to 1L)), + xcoms = mapOf("upstream" to "scalar"), + ) + + assertThrows(IllegalStateException::class.java) { args.get(0, String::class.java) } + } + + @Test + @DisplayName("Should fail when an element index points past the end of a list XCom") + fun shouldRejectElementIndexOutOfBounds() { + val (args, _) = + argsWith( + listOf(mapOf("kind" to "xcom", "name" to "x", "task_id" to "upstream", "element_index" to 3L)), + xcoms = mapOf("upstream" to listOf("a", "b")), + ) + + val error = assertThrows(IllegalStateException::class.java) { args.get(0, String::class.java) } + + assertEquals( + "Argument 'x' binds element 3 of task 'upstream', but its XCom holds only 2 element(s)", + error.message, + ) + } + + @Test + @DisplayName("Should pass null through an element index when the upstream pushed nothing") + fun shouldPassNullThroughElementIndex() { + val (args, _) = + argsWith( + listOf(mapOf("kind" to "xcom", "name" to "x", "task_id" to "upstream", "element_index" to 1L)), + ) + + assertNull(args.get(0, String::class.java)) + } + + @Test + @DisplayName("Should fail on an unsupported binding kind") + fun shouldRejectUnknownBindingKind() { + assertThrows(IllegalStateException::class.java) { + argsWith(listOf(mapOf("kind" to "mystery", "name" to "x"))) + } + } + + @Test + @DisplayName("Should fail on duplicate binding names") + fun shouldRejectDuplicateBindingNames() { + assertThrows(IllegalStateException::class.java) { + argsWith( + listOf( + mapOf("kind" to "literal", "name" to "x", "value" to 1L), + mapOf("kind" to "literal", "name" to "x", "value" to 2L), + ), + ) + } + } + + @Test + @DisplayName("Should widen a wire integer into the declared numeric type") + fun shouldWidenNumericBinding() { + val (args, _) = argsWith(listOf(mapOf("kind" to "literal", "name" to "x", "value" to 5L))) + + assertEquals(5, args.require(0, Integer::class.java).toInt()) + } + + @Test + @DisplayName("Should keep the element type of a generic parameter read through TypeRef") + fun shouldKeepGenericElementType() { + val (args, _) = + argsWith(listOf(mapOf("kind" to "literal", "name" to "values", "value" to listOf(1L, 2L)))) + + val values = args.require(0, object : TypeRef>() {}) + + assertEquals(listOf(1.0, 2.0), values) + } + + @Test + @DisplayName("Should pass null through get for both the plain and the generic read") + fun shouldPassNullThrough() { + val (args, _) = + argsWith( + listOf( + mapOf("kind" to "literal", "name" to "scalar", "value" to null), + mapOf("kind" to "literal", "name" to "values", "value" to null), + ), + ) + + assertNull(args.get(0, String::class.java)) + assertNull(args.get(1, object : TypeRef>() {})) + } + + @Test + @DisplayName("Should fail fast when the stub call bound fewer arguments than the task declares") + fun shouldFailWhenFewerArgumentsBoundThanDeclared() { + val error = + assertThrows(IllegalStateException::class.java) { + argsWith(listOf(mapOf("kind" to "literal", "name" to "only", "value" to 1L)), declared = 2) + } + + assertEquals( + "Task 't' declares 2 data parameter(s) but the stub call bound 1 argument(s)", + error.message, + ) + } + + @Test + @DisplayName("Should drop a captured default the method does not declare") + fun shouldDropCapturedDefaultTheMethodOmits() { + val (args, _) = + argsWith( + listOf( + mapOf("kind" to "literal", "name" to "rows", "value" to 1L), + mapOf("kind" to "literal", "name" to "note", "value" to "unset", "from_default" to true), + ), + declared = 1, + ) + + assertEquals(1L, args.get(0, java.lang.Long::class.java)) + } + + @Test + @DisplayName("Should keep a captured default the method does declare") + fun shouldKeepCapturedDefaultTheMethodDeclares() { + val (args, _) = + argsWith( + listOf( + mapOf("kind" to "literal", "name" to "rows", "value" to 1L), + mapOf("kind" to "literal", "name" to "note", "value" to "unset", "from_default" to true), + ), + declared = 2, + ) + + assertEquals("unset", args.get(1, String::class.java)) + } + + @Test + @DisplayName("Should report both counts when dropping captured defaults still leaves a mismatch") + fun shouldFailWhenDroppingDefaultsStillMismatches() { + val error = + assertThrows(IllegalStateException::class.java) { + argsWith( + listOf( + mapOf("kind" to "literal", "name" to "rows", "value" to 1L), + mapOf("kind" to "literal", "name" to "ratio", "value" to 2.5), + mapOf("kind" to "literal", "name" to "note", "value" to "unset", "from_default" to true), + ), + declared = 1, + ) + } + + assertEquals( + "Task 't' declares 1 data parameter(s) but the stub call bound 3 argument(s); " + + "2 remain after dropping captured defaults", + error.message, + ) + } + + @Test + @DisplayName("Should fail fast when the stub call bound more arguments than the task declares") + fun shouldFailWhenMoreArgumentsBoundThanDeclared() { + val error = + assertThrows(IllegalStateException::class.java) { + argsWith( + listOf( + mapOf("kind" to "literal", "name" to "kept", "value" to 1L), + mapOf("kind" to "literal", "name" to "dropped", "value" to 2L), + ), + declared = 1, + ) + } + + assertEquals( + "Task 't' declares 1 data parameter(s) but the stub call bound 2 argument(s)", + error.message, + ) + } + + @Test + @DisplayName("Should name the stub argument when a required position is bound to a null literal") + fun shouldNameArgumentBoundToNullLiteral() { + val (args, _) = argsWith(listOf(mapOf("kind" to "literal", "name" to "region_code", "value" to null))) + + val error = assertThrows(MissingXComException::class.java) { args.require(0, String::class.java) } + + assertEquals( + "Task parameter 'region_code' of task 't' is bound to a null literal, but has a primitive type " + + "that cannot be null; declare a boxed type (e.g. Integer instead of int) to receive null.", + error.message, + ) + } + + @Test + @DisplayName("Should name the upstream when a required position's XCom was never pushed") + fun shouldNameUpstreamWhenXComMissing() { + val (args, _) = + argsWith(listOf(mapOf("kind" to "xcom", "name" to "threshold", "task_id" to "upstream"))) + + val error = assertThrows(MissingXComException::class.java) { args.require(0, Integer::class.java) } + + assertEquals( + "Task parameter 'threshold' requires an XCom from task 'upstream', but none was pushed. " + + "This parameter has a primitive type that cannot be null; declare it with a boxed type " + + "(e.g. Integer instead of int) to receive null.", + error.message, + ) + } + + @Test + @DisplayName("Should reject a TypeRef whose type argument is fixed by an intermediate class") + fun shouldRejectIndirectTypeRef() { + val error = assertThrows(IllegalArgumentException::class.java) { object : NamesRef() {} } + + assertEquals( + "TypeRef needs a concrete type argument on an anonymous subclass, e.g. new TypeRef>() {}", + error.message, + ) + } +}