This is an automated email from the ASF dual-hosted git repository.
jason810496 pushed a commit to branch main
in repository https://gitbox.apache.org/repos/asf/airflow.git
The following commit(s) were added to refs/heads/main by this push:
new ab4da0b5636 Java SDK: Declare a Dag's task graph with @Builder.Deps
(#71189)
ab4da0b5636 is described below
commit ab4da0b5636f75c2bba55ec1205a705a2da903aa
Author: Jason(Zhe-You) Liu <[email protected]>
AuthorDate: Tue Oct 6 10:55:35 2026 +0800
Java SDK: Declare a Dag's task graph with @Builder.Deps (#71189)
* Java SDK: Declare a Dag's task graph with @Builder.Deps
An annotated Dag class could name its tasks but not say how they connect,
so a Dag whose annotations carried everything else still needed a Python
Dag file for its graph. Declaring the graph by calling the tasks lets
javac check that each task is fed something its parameter accepts.
Bundle.register(Class) now also takes the @Builder.Dag class itself and
registers the Dag its generated builder builds, so there is still only one
name to keep in sync.
* Java SDK: Quote wired task ids as Java string literals
* Java SDK: Validate the @Builder.Deps class at compile time
* Java SDK: Reject task method names that clash with the wiring view
* Java SDK: Type a wired numeric input by its parameter's own type
* Java SDK: Register both the Dag and handlers of a class, unwrapping
builder failures
* Java SDK: Reject a raw null wiring argument and test the recorder's guards
* Java SDK: Reword the native-Dag docs without em-dashes
* Java SDK: Drop Bundle tests that repeat existing coverage
* Java SDK: Require a @Builder.Deps class on every @Builder.Dag class
* Java SDK: Back the k8s test's stub tasks with task handlers
* Java SDK: Reject overloaded task and task-handler methods
* Java SDK: Fix the Dag examples a required wiring class invalidates
---
.../language-sdks/java.rst | 77 +++
java-sdk/README.md | 6 +-
java-sdk/adr/0002-native-dag-interface.md | 7 +-
.../airflow/example/ExampleBundleBuilder.java | 3 +
.../example/nativedag/AnnotationExample.java | 73 ++
.../org/apache/airflow/sdk/BuilderProcessor.kt | 335 +++++++--
.../kotlin/org/apache/airflow/sdk/BuilderTest.kt | 761 ++++++++++++++++++++-
java-sdk/sdk/build.gradle.kts | 46 +-
.../src/main/kotlin/org/apache/airflow/sdk/Arg.kt | 22 +-
.../main/kotlin/org/apache/airflow/sdk/Bundle.kt | 68 +-
.../src/main/kotlin/org/apache/airflow/sdk/Deps.kt | 17 +-
.../kotlin/org/apache/airflow/sdk/internal/Refs.kt | 122 ++++
.../org/apache/airflow/sdk/internal/Registrar.kt | 23 +
.../org/apache/airflow/sdk/ArgTestSupport.kt | 5 +-
.../kotlin/org/apache/airflow/sdk/BundleTest.kt | 118 +++-
.../kotlin/org/apache/airflow/sdk/InputTaskTest.kt | 6 +-
.../apache/airflow/sdk/internal/ArgValuesTest.kt | 5 +-
.../org/apache/airflow/sdk/internal/RefsTest.kt | 158 +++++
.../apache/airflow/sdk/internal/RegistrarTest.kt} | 22 +-
.../apache/airflow/k8sexample/CombinedExample.java | 15 +-
.../airflow/k8sexample/K8sBundleBuilder.java | 17 +-
21 files changed, 1733 insertions(+), 173 deletions(-)
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 eafd542df69..f55397367d9 100644
--- a/airflow-core/docs/authoring-and-scheduling/language-sdks/java.rst
+++ b/airflow-core/docs/authoring-and-scheduling/language-sdks/java.rst
@@ -331,6 +331,9 @@ Annotate a plain Java class and let the SDK generate the
boilerplate at compile
- Marks a method as a task of a Java-owned Dag. If ``id`` is omitted the
method name is
used. Further attributes (``retries``, ``queue``, ``retryDelay``, …)
are Airflow's own
task settings; only attributes written explicitly are applied.
+ * - ``@Builder.Deps``
+ - Marks the nested class that declares the task graph in Java,
TaskFlow-style. Required for
+ a Dag that Java owns end to end. See :ref:`java-sdk/native-dags`.
* - ``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
@@ -591,6 +594,80 @@ Edges are checked when the Dag is registered with a
``Bundle``: an upstream that
Dag, or to no Dag, and a cycle anywhere in the graph both fail there rather
than at the first task
run.
+Wiring the graph with ``@Builder.Deps``
+~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~
+
+For a Dag written with annotations, the graph is declared by a nested
``@Builder.Deps`` class. The
+annotation processor generates a ``<ClassName>Deps`` interface, the *wiring
view*, with one method
+per ``@Builder.Task`` method: the injected ``Client`` and ``Context``
parameters are dropped, each
+data parameter becomes an ``Arg<T>``, and the return value becomes a
``TaskRef<T>``. Calling a view
+method registers its task, and passing the handle one returned into another
feeds the upstream's
+output into the downstream's parameter *and* wires the data edge. The call
graph is the task graph,
+and ``javac`` type-checks it.
+
+Declare the wiring class as a ``static`` nested class of the Dag class that
``implements`` the
+generated view, with a no-argument ``depends()`` method:
+
+.. code-block:: java
+
+ @Builder.Dag(
+ id = "java_etl",
+ schedule = "@daily",
+ description = "Pure-Java Dag built with annotations",
+ tags = {"example", "java-sdk"})
+ public class EtlPipeline {
+
+ @Builder.Task(id = "extract", retries = 2)
+ public long extract() {
+ return 42L;
+ }
+
+ @Builder.Task(id = "transform")
+ public long transform(long extracted, double factor) {
+ return (long) (extracted * factor);
+ }
+
+ @Builder.Task(id = "load")
+ public void load(long transformed) {
+ // implement task logic
+ }
+
+ @Builder.Task(id = "audit")
+ public void audit() {
+ // side effect only, no data in or out
+ }
+
+ @Builder.Deps
+ static class Wiring implements EtlPipelineDeps {
+ void depends() {
+ var rows = extract();
+ load(transform(rows, lit(0.9)));
+ rows.before(audit()); // ordering-only edge: audit waits for extract
+ }
+ }
+ }
+
+Every ``@Builder.Task`` method must be called in the wiring class; a task the
wiring missed fails at
+Dag-parse time. ``lit(...)`` wires an inline constant where no upstream feeds
a parameter. A bare
+``double`` cannot be an ``Arg``, so a constant is wrapped. A view method that
takes no arguments
+returns the same handle every time, so it names one node wherever it appears;
one that takes
+arguments is called once, and the wiring fails if it is called again with
arguments, so hold its
+handle in a local and reuse that.
+
+``before``, ``after`` and ``Flow.of`` work here exactly as they do on the
interface surface; inside
+the wiring class ``Flow`` is inherited by simple name, so it needs no import
and never collides with
+``java.util.concurrent.Flow``.
+
+Every ``@Builder.Dag`` class declares a wiring class, because the graph is
what the Dag owns. A
+class that supplies only task bodies, for a Dag a Python file declares, carries
+``@Builder.TaskHandler`` instead and contributes no Dag.
+
+.. note::
+
+ Runtime argument bindings win over Java-declared wiring. When the
supervisor delivers bindings
+ for a run (see :ref:`java-sdk/arg-binding`), the binding at a parameter's
position is what the
+ task receives. Wired inputs are the fallback, which is what a native Java
Dag always uses.
+
Configuration attributes
~~~~~~~~~~~~~~~~~~~~~~~~
diff --git a/java-sdk/README.md b/java-sdk/README.md
index 752253b6ab8..8e78206db39 100644
--- a/java-sdk/README.md
+++ b/java-sdk/README.md
@@ -745,9 +745,9 @@ E2E_TEST_MODE=java_sdk uv run --project airflow-e2e-tests
pytest \
- The annotation processor (`BuilderProcessor.kt`) uses `kapt`. The `Builder`
class holding the `@Builder.Dag` / `@Builder.Task` annotations is generated
from the Dag serialization schema by `:sdk:generateDagDsl` (vendored at
- `sdk/schema/dag-schema.json`). The `Arg`/`TaskRef` and `Deps`/`Flow` graph
- types are hand-written next to the rest of the public surface in
- `sdk/src/main/kotlin/org/apache/airflow/sdk/`. When adding annotation
+ `sdk/schema/dag-schema.json`), `@Builder.Deps` included. The `Arg`/`TaskRef`
+ and `Deps`/`Flow` graph types are hand-written next to the rest of the public
+ surface in `sdk/src/main/kotlin/org/apache/airflow/sdk/`. When adding
annotation
behaviour, handle it in `BuilderProcessor.kt` and add a golden-output test in
`processor/src/test/kotlin/`.
- The Python coordinator subclasses `SubprocessCoordinator`. Do not reach into
diff --git a/java-sdk/adr/0002-native-dag-interface.md
b/java-sdk/adr/0002-native-dag-interface.md
index b3ad667e578..5ab2ffdaff3 100644
--- a/java-sdk/adr/0002-native-dag-interface.md
+++ b/java-sdk/adr/0002-native-dag-interface.md
@@ -65,7 +65,8 @@ public class EtlPipeline { // extends nothing of ours; your
own base class stays
public void audit(Client client) { /* side effect only, no data in or out */
}
@Builder.Task(id = "notify")
- public void notify(Client client) { /* side effect only, no data in or out
*/ }
+ // "notify" is Object.notify, so the method takes another name and keeps the
task id.
+ public void alert(Client client) { /* side effect only, no data in or out */
}
@Builder.Deps
static class Wiring implements EtlPipelineDeps {
@@ -77,7 +78,7 @@ public class EtlPipeline { // extends nothing of ours; your
own base class stays
// non-TaskFlow (ordering-only) edges: sequence with no data flowing
rows.then(audit()); // extract >> audit
- Flow.of(loaded, audit()).then(notify()); // [load, audit] >> notify
+ Flow.of(loaded, audit()).then(alert()); // [load, audit] >> notify
}
}
@@ -101,7 +102,7 @@ interface EtlPipelineDeps extends Deps { // Deps: the
shared base that nests Flo
default TaskRef<Long> transform(Arg<Long> extracted, Arg<Double> threshold)
{ return Flow.call("transform", extracted, threshold); }
default TaskRef<Void> load(Arg<Long> transformed) { return Flow.call("load",
transformed); }
default TaskRef<Void> audit() { return Flow.node("audit"); }
- default TaskRef<Void> notify() { return Flow.node("notify"); }
+ default TaskRef<Void> alert() { return Flow.node("notify"); }
}
```
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 d75f5762ba1..76175db8a6d 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
@@ -21,12 +21,15 @@ package org.apache.airflow.example;
import org.apache.airflow.sdk.*;
+// One bundle serves every surface: Dags built in Java, and the handler classes
+// whose Dags the Python file owns.
public class ExampleBundleBuilder {
public static Bundle build() {
return new Bundle()
.register(InterfaceExampleBuilder.build())
.register(AnnotationExample.class)
.register(XComCastingExample.class)
+ .register(org.apache.airflow.example.nativedag.AnnotationExample.class)
.register(org.apache.airflow.example.nativedag.InterfaceExample.build());
}
diff --git
a/java-sdk/example/src/java/org/apache/airflow/example/nativedag/AnnotationExample.java
b/java-sdk/example/src/java/org/apache/airflow/example/nativedag/AnnotationExample.java
new file mode 100644
index 00000000000..c2a3bc89f75
--- /dev/null
+++
b/java-sdk/example/src/java/org/apache/airflow/example/nativedag/AnnotationExample.java
@@ -0,0 +1,73 @@
+/*
+ * 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.
+ */
+
+// "native" is a Java keyword, so the native-Dag examples live in "nativedag".
+package org.apache.airflow.example.nativedag;
+
+import static java.lang.System.Logger.Level.INFO;
+
+import org.apache.airflow.sdk.*;
+
+// A Dag defined entirely in Java, annotation-style. The @Builder.Dag and
+// @Builder.Task attributes carry the configuration and the @Builder.Deps class
+// declares the graph -- the single place dependencies are defined.
[email protected](
+ id = "java_native_annotation_example",
+ description = "Pure-Java Dag authored with annotations",
+ schedule = "@daily",
+ startDate = "2026-01-01T00:00:00Z",
+ catchup = false,
+ tags = {"example", "java-sdk"})
+public class AnnotationExample {
+ private static final System.Logger log =
System.getLogger(AnnotationExample.class.getName());
+
+ @Builder.Task(id = "extract", retries = 2)
+ public long extract() {
+ log.log(INFO, "Extracting a value");
+ return 42L;
+ }
+
+ @Builder.Task(id = "transform")
+ public long transform(long extracted, double factor) {
+ log.log(INFO, "Transforming {0} by {1}", extracted, factor);
+ return (long) (extracted * factor);
+ }
+
+ @Builder.Task(id = "load")
+ public void load(long transformed) {
+ log.log(INFO, "Loaded {0}", transformed);
+ }
+
+ @Builder.Task(id = "audit")
+ public void audit() {
+ log.log(INFO, "Audited the run");
+ }
+
+ // Implements the generated wiring view, so javac type-checks the graph:
+ // extract() yields a TaskRef<Long>, which is an Arg<Long> for transform.
+ @Builder.Deps
+ static class Wiring implements AnnotationExampleDeps {
+ void depends() {
+ var extracted = extract();
+ load(transform(extracted, lit(1.5)));
+ // Ordering-only edge: audit runs after extract, with no data flowing.
+ extracted.before(audit());
+ }
+ }
+}
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 21b4165728e..93e185a614e 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
@@ -28,9 +28,11 @@ import com.squareup.javapoet.MethodSpec
import com.squareup.javapoet.ParameterizedTypeName
import com.squareup.javapoet.TypeName
import com.squareup.javapoet.TypeSpec
+import com.squareup.javapoet.WildcardTypeName
import org.apache.airflow.sdk.internal.ArgValues
import org.apache.airflow.sdk.internal.Field
import org.apache.airflow.sdk.internal.FieldType
+import org.apache.airflow.sdk.internal.Refs
import org.apache.airflow.sdk.internal.SchemaFields
import org.apache.airflow.sdk.internal.TaskArgs
import org.apache.airflow.sdk.internal.TypeRef
@@ -55,6 +57,7 @@ import javax.lang.model.element.VariableElement
import javax.lang.model.type.TypeKind
import javax.lang.model.type.TypeMirror
import javax.tools.Diagnostic
+import org.apache.airflow.sdk.internal.builderName as generatedBuilderName
/**
* @suppress
@@ -65,24 +68,33 @@ import javax.tools.Diagnostic
* `META-INF/services/javax.annotation.processing.Processor`; not intended to
be
* instantiated or referenced directly.
*
- * For each class annotated with [Builder.Dag], generates a `*Builder` class
- * containing:
+ * For each class annotated with [Builder.Dag], generates:
*
- * - One inner class per [Builder.Task]-annotated method, implementing [Task].
- * - A static `build()` method that constructs the [DagDef], lowers every
- * explicitly-written `@Builder.Dag` attribute into a `DagDef.config` call,
- * and registers those inner classes as [TaskDef]s, each carrying its
- * explicitly-written `@Builder.Task` attributes the same way.
+ * - A `*Builder` class containing one inner class per [Builder.Task]-annotated
+ * method (implementing [Task]), and a static `build()` that constructs the
+ * [DagDef], lowers every explicitly-written `@Builder.Dag` attribute into a
+ * `DagDef.config` call, then runs the class's [Builder.Deps] class and
+ * verifies it registered every task.
+ * - A `*Deps` wiring-view interface whose methods mirror the task methods:
+ * injectable parameters ([Client],
+ * [Context]) are dropped, data parameters become [Arg]-typed inputs, and the
+ * return value becomes a [TaskRef]. Calling one registers the task with its
+ * explicitly-written `@Builder.Task` attributes lowered into
`TaskDef.config`
+ * calls; passing one call's handle to another wires the dependency edge and
+ * feeds the upstream's return-value XCom into the downstream's parameter,
+ * type-checked by javac through the [Arg] / [TaskRef] generics.
*
- * 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]
- * fields through [ArgValues], by argument name. Non-`void` return values are
+ * In the generated `execute` bodies, 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] fields through [ArgValues], by argument name. Non-`void` return
values are
* forwarded to `client.setXCom`.
*/
@SupportedAnnotationTypes(
"org.apache.airflow.sdk.Builder.Dag",
+ "org.apache.airflow.sdk.Builder.Task",
"org.apache.airflow.sdk.Builder.TaskHandler",
+ "org.apache.airflow.sdk.Builder.Deps",
)
@SupportedSourceVersion(SourceVersion.RELEASE_11)
class BuilderProcessor : AbstractProcessor() {
@@ -91,6 +103,16 @@ class BuilderProcessor : AbstractProcessor() {
roundEnv: RoundEnvironment,
): Boolean {
if (annotations.isEmpty()) return false
+ roundEnv.getElementsAnnotatedWith(Builder.Deps::class.java).forEach { el ->
+ val owner = el.enclosingElement
+ if (owner !is TypeElement ||
owner.getAnnotation(Builder.Dag::class.java) == null) {
+ processingEnv.messager.printMessage(
+ Diagnostic.Kind.ERROR,
+ "@Builder.Deps class '${el.simpleName}' must be nested directly in a
@Builder.Dag class",
+ el,
+ )
+ }
+ }
roundEnv
.getElementsAnnotatedWith(Builder.TaskHandler::class.java)
.mapNotNull { it.enclosingElement as? TypeElement }
@@ -108,12 +130,21 @@ class BuilderProcessor : AbstractProcessor() {
roundEnv.getElementsAnnotatedWith(Builder.Dag::class.java).filterIsInstance<TypeElement>().forEach
{ el ->
with(processingEnv) {
runCatching {
+ val packageName =
elementUtils.getPackageOf(el).qualifiedName.toString()
+ val declarations = collectTasks(el)
+ val builderName =
+ ClassName.get(
+ packageName,
+ generatedBuilderName(packageName, el.simpleName.toString(),
dagAnnotation(el).to).substringAfterLast('.'),
+ )
+ val depsName = ClassName.get(packageName, "${el.simpleName}Deps")
+ val deps = findDeps(el, depsName)
+ declarations.forEach { checkViewName(it) }
JavaFile
- .builder(
- elementUtils.getPackageOf(el).qualifiedName.toString(),
- buildDag(el),
- ).build()
+ .builder(packageName, buildBuilder(el, declarations, deps,
builderName))
+ .build()
.writeTo(filer)
+ JavaFile.builder(packageName, buildDeps(el, declarations,
builderName, depsName)).build().writeTo(filer)
}.onFailure { e ->
messager.printMessage(
Diagnostic.Kind.ERROR,
@@ -152,6 +183,7 @@ class BuilderProcessor : AbstractProcessor() {
.addModifiers(Modifier.PUBLIC, Modifier.STATIC)
.addParameter(BUNDLE_TYPE, "bundle")
+ val names = mutableSetOf<String>()
for (inner in el.enclosedElements) {
if (inner !is ExecutableElement) continue
val handler = inner.getAnnotation(Builder.TaskHandler::class.java) ?:
continue
@@ -161,24 +193,36 @@ class BuilderProcessor : AbstractProcessor() {
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))
+ val decl = TaskDeclaration(inner, handler.task.ifBlank {
inner.simpleName.toString() }, collectDataParams(inner))
+ require(names.add(inner.simpleName.toString())) {
+ "Class ${el.simpleName} overloads task-handler method
'${inner.simpleName}'; a method's name is " +
+ "the name of its generated task class, so rename one and keep its
task id with " +
+ "@Builder.TaskHandler(task = \"${decl.id}\")"
+ }
+ registrar.addType(buildTask(decl, el))
registerInto.addStatement(
$$"bundle.register($S, $S, $L.class)",
handler.dag,
- handler.task.ifBlank { inner.simpleName },
- innerName,
+ decl.id,
+ decl.className,
)
}
return registrar.addMethod(registerInto.build()).build()
}
- private fun buildDag(el: TypeElement): TypeSpec {
- val ann = el.getAnnotation(Builder.Dag::class.java)!!
+ private fun dagAnnotation(el: TypeElement): Builder.Dag =
el.getAnnotation(Builder.Dag::class.java)!!
+
+ private fun buildBuilder(
+ el: TypeElement,
+ declarations: List<TaskDeclaration>,
+ deps: TypeElement,
+ builderName: ClassName,
+ ): TypeSpec {
+ val ann = dagAnnotation(el)
val builderClass =
TypeSpec
- .classBuilder(ann.to.ifBlank { "${el.simpleName}Builder" })
+ .classBuilder(builderName)
.addModifiers(Modifier.PUBLIC, Modifier.FINAL)
val buildMethod =
@@ -190,46 +234,206 @@ class BuilderProcessor : AbstractProcessor() {
explicitConfig(el, DAG_ANNOTATION, DAG_STRUCTURAL_ATTRIBUTES,
SchemaFields.DAG).forEach { (key, value) ->
buildMethod.addStatement($$"dag.config($S, $L)", key, value)
}
+ buildMethod.addStatement(
+ $$"return $T.record(dag, $T.of($L), new $T()::depends)",
+ REFS_TYPE,
+ ClassName.get(List::class.java),
+ CodeBlock.join(declarations.map { CodeBlock.of($$"$S", it.id) }, ", "),
+ ClassName.get(deps),
+ )
+ builderClass.addMethod(buildMethod.build())
- for (inner in el.enclosedElements) {
- if (inner !is ExecutableElement) continue
- if (inner.isVarArgs) throw IllegalArgumentException("Cannot create task
from vararg function ${inner.simpleName}")
-
- val taskAnn = inner.getAnnotation(Builder.Task::class.java) ?: continue
- val innerName =
inner.simpleName.toString().replaceFirstChar(Char::uppercase)
+ declarations.forEach { builderClass.addType(buildTask(it, el)) }
+ return builderClass.build()
+ }
- builderClass.addType(buildTask(innerName, inner, el))
+ /**
+ * Generates the Dag's wiring view: one default method per task, with the
+ * injected arguments stripped, each data argument lifted to [Arg] and the
+ * return lifted to [TaskRef].
+ *
+ * It is an interface so the `@Builder.Deps` class can *implement* it and
+ * keep its own `extends` free, and so the Dag class's real task methods --
+ * which differ only in their injected arguments -- do not clash with it.
+ */
+ private fun buildDeps(
+ el: TypeElement,
+ declarations: List<TaskDeclaration>,
+ builderName: ClassName,
+ depsName: ClassName,
+ ): TypeSpec {
+ val view =
+ TypeSpec
+ .interfaceBuilder(depsName)
+ .addModifiers(Modifier.PUBLIC)
+ .addSuperinterface(DEPS_TYPE)
+ .addJavadoc(
+ "Wiring view of {@link \$T}'s task methods, for declaring its task
graph.\n\n" +
+ "<p>Calling one registers its task with the Dag being built;
passing the handle it\n" +
+ "returned into another call feeds the upstream's output into that
task's parameter\n" +
+ "and wires the data edge. {@code before} and {@code after} wire an
ordering-only edge.\n",
+ ClassName.get(el),
+ )
- buildMethod.addStatement(
- $$"dag.addTask($L)",
- taskDefCode(inner, taskAnn.id.ifBlank { inner.simpleName.toString() },
innerName),
- )
+ for (decl in declarations) {
+ val method =
+ MethodSpec
+ .methodBuilder(decl.method.simpleName.toString())
+ .addModifiers(Modifier.PUBLIC, Modifier.DEFAULT)
+ .returns(ParameterizedTypeName.get(TASK_HANDLE_TYPE,
TypeName.get(decl.method.returnType).boxIfPossible()))
+ decl.dataParams.forEach { method.addParameter(inType(it.type), it.name) }
+ val def = taskDefCode(decl, CodeBlock.of($$"$T.$L", builderName,
decl.className))
+ if (decl.dataParams.isEmpty()) {
+ method.addStatement($$"return $T.node($L)", REFS_TYPE, def)
+ } else {
+ method.addStatement(
+ $$"return $T.call($L, $L)",
+ REFS_TYPE,
+ def,
+ decl.dataParams.joinToString { it.name },
+ )
+ }
+ view.addMethod(method.build())
}
-
- buildMethod.addStatement("return dag")
- builderClass.addMethod(buildMethod.build())
- return builderClass.build()
+ return view.build()
}
/**
- * Emits `new TaskDef(id, <className>.class)` with the explicitly-written
+ * Emits `new TaskDef(id, <classRef>.class)` with the explicitly-written
* `@Builder.Task` attributes lowered into chained `.config` calls.
*/
private fun taskDefCode(
- method: ExecutableElement,
- id: String,
- className: String,
+ decl: TaskDeclaration,
+ classRef: CodeBlock,
): CodeBlock {
val taskDef =
CodeBlock
.builder()
- .add($$"new $T($S, $L.class)", TASK_DEF_TYPE, id, className)
- explicitConfig(method, TASK_ANNOTATION, TASK_STRUCTURAL_ATTRIBUTES,
SchemaFields.TASK).forEach { (key, value) ->
+ .add($$"new $T($S, $L.class)", TASK_DEF_TYPE, decl.id, classRef)
+ explicitConfig(decl.method, TASK_ANNOTATION, TASK_STRUCTURAL_ATTRIBUTES,
SchemaFields.TASK).forEach { (key, value) ->
taskDef.add($$".config($S, $L)", key, value)
}
return taskDef.build()
}
+ /**
+ * Maps a data parameter's declared type to its wiring-view input type,
+ * `Arg<? extends T>` of the boxed type. A numeric parameter therefore takes
+ * only its own type, so javac rejects wiring that could lose a value, such
+ * as a `double` upstream into a `long` parameter. An `Object` parameter
+ * takes any upstream, including a `void` task's handle, whose value is null.
+ */
+ private fun inType(paramType: TypeMirror): TypeName =
+ ParameterizedTypeName.get(ARG_TYPE,
WildcardTypeName.subtypeOf(TypeName.get(paramType).boxIfPossible()))
+
+ private fun collectTasks(el: TypeElement): List<TaskDeclaration> {
+ val declarations = mutableListOf<TaskDeclaration>()
+ for (inner in el.enclosedElements) {
+ if (inner !is ExecutableElement) continue
+ val ann = inner.getAnnotation(Builder.Task::class.java) ?: continue
+ if (inner.isVarArgs) throw IllegalArgumentException("Cannot create task
from vararg function ${inner.simpleName}")
+ val id = ann.id.ifBlank { inner.simpleName.toString() }
+ require(declarations.none { it.id == id }) { "Tasks in Dag have
duplicate ID: $id" }
+ require(declarations.none {
it.method.simpleName.contentEquals(inner.simpleName) }) {
+ "Dag class ${el.simpleName} overloads task method
'${inner.simpleName}'; a method's name is the " +
+ "name of its generated task class and of its wiring-view method, so
rename one and keep its " +
+ "task id with @Builder.Task(id = \"$id\")"
+ }
+ declarations += TaskDeclaration(inner, id, collectDataParams(inner))
+ }
+ return declarations
+ }
+
+ /**
+ * Finds and validates the class's `@Builder.Deps` wiring class, which
+ * declares the Dag's task graph and is what makes it a Dag Java owns.
+ *
+ * The generated builder runs `new Wiring()::depends`, so everything that
+ * expression needs is checked here, where the error can name the class.
+ */
+ private fun findDeps(
+ el: TypeElement,
+ view: ClassName,
+ ): TypeElement {
+ val classes =
+ el.enclosedElements
+ .filterIsInstance<TypeElement>()
+ .filter { it.getAnnotation(Builder.Deps::class.java) != null }
+ require(classes.isNotEmpty()) {
+ "Dag class ${el.simpleName} must declare a @Builder.Deps class
implementing ${view.simpleName()} " +
+ "to declare its task graph; a class of task bodies for a Dag the
Python file owns carries " +
+ "@Builder.TaskHandler instead"
+ }
+ val deps =
+ classes.singleOrNull()
+ ?: throw IllegalArgumentException(
+ "Dag class ${el.simpleName} declares more than one @Builder.Deps
class: " +
+ classes.joinToString { it.simpleName.toString() },
+ )
+ val name = deps.simpleName
+ require(deps.kind == ElementKind.CLASS && Modifier.ABSTRACT !in
deps.modifiers) {
+ "@Builder.Deps '$name' must be a concrete class"
+ }
+ require(Modifier.STATIC in deps.modifiers && Modifier.PRIVATE !in
deps.modifiers) {
+ "@Builder.Deps class '$name' must be static and non-private"
+ }
+ require(deps.interfaces.any { it.isView(view) }) {
+ "@Builder.Deps class '$name' must implement ${view.simpleName()}, the
wiring view of ${el.simpleName}"
+ }
+ require(
+ deps.enclosedElements
+ .filterIsInstance<ExecutableElement>()
+ .any { it.kind == ElementKind.CONSTRUCTOR && it.parameters.isEmpty()
&& Modifier.PRIVATE !in it.modifiers },
+ ) {
+ "@Builder.Deps class '$name' needs a non-private no-argument constructor"
+ }
+ val depends =
+ processingEnv.elementUtils
+ .getAllMembers(deps)
+ .filterIsInstance<ExecutableElement>()
+ .firstOrNull { it.isNoArgDepends() }
+ ?: throw IllegalArgumentException(
+ "@Builder.Deps class '$name' must have a non-private, no-argument
depends() method",
+ )
+ val checked = depends.thrownTypes.filterNot { isUnchecked(it) }
+ require(checked.isEmpty()) {
+ "depends() of @Builder.Deps class '$name' must not throw checked
exceptions: ${checked.joinToString()}"
+ }
+ return deps
+ }
+
+ /**
+ * Rejects a task method whose wiring-view twin would clash with a member
+ * the view or the wiring class already has: `depends`, `lit`, or a method
+ * of `Object`.
+ */
+ private fun checkViewName(decl: TaskDeclaration) {
+ val name = decl.method.simpleName.toString()
+ require(name !in RESERVED_VIEW_NAMES) {
+ "Task method '$name' clashes with a member of the wiring view; rename
the method and keep " +
+ "the task id with @Builder.Task(id = \"${decl.id}\")"
+ }
+ }
+
+ /**
+ * Matches the view by the name the class wrote: the view is generated in
+ * this same round, so javac may not have resolved it yet.
+ */
+ private fun TypeMirror.isView(view: ClassName): Boolean = toString().let {
it == view.canonicalName() || it == view.simpleName() }
+
+ private fun isUnchecked(type: TypeMirror): Boolean =
+ with(processingEnv) {
+ listOf(RuntimeException::class.java, Error::class.java).any {
+ typeUtils.isAssignable(type,
elementUtils.getTypeElement(it.canonicalName).asType())
+ }
+ }
+
+ private fun ExecutableElement.isNoArgDepends(): Boolean =
+ simpleName.contentEquals("depends") &&
+ parameters.isEmpty() &&
+ Modifier.PRIVATE !in modifiers &&
+ Modifier.STATIC !in modifiers
+
/**
* Lowers the explicitly-written configuration attributes of [element]'s
* [annotationName] annotation into (schema key, value code) pairs. Only
@@ -299,8 +503,7 @@ class BuilderProcessor : AbstractProcessor() {
}
private fun buildTask(
- name: String,
- inner: ExecutableElement,
+ decl: TaskDeclaration,
parent: TypeElement,
): TypeSpec {
val executeSpec =
@@ -313,8 +516,8 @@ class BuilderProcessor : AbstractProcessor() {
.addParameter(CLIENT_TYPE, "client")
.addException(Exception::class.java)
- val dataParams = collectDataParams(inner)
- val dataByName = dataParams.associateBy { it.name }
+ val inner = decl.method
+ val dataByName = decl.dataParams.associateBy { it.name }
val innerArgs =
with(processingEnv) {
inner.parameters.joinToString { param ->
@@ -327,9 +530,9 @@ class BuilderProcessor : AbstractProcessor() {
}
}
- val taken = dataParams.mapTo(mutableSetOf()) { it.local }
+ val taken = decl.dataParams.mapTo(mutableSetOf()) { it.local }
val argsLocal = generateSequence("args") { "${it}_" }.first { it !in taken
}
- val flatParams = dataParams.filterNot { it.isTaskInput }
+ val flatParams = decl.dataParams.filterNot { it.isTaskInput }
if (flatParams.isNotEmpty()) {
executeSpec.addStatement(
$$"$T $L = $T.of(context, client, $L)",
@@ -339,7 +542,7 @@ class BuilderProcessor : AbstractProcessor() {
flatParams.size,
)
}
- dataParams.forEach { param ->
+ decl.dataParams.forEach { param ->
val paramType = TypeName.get(param.type)
if (param.isTaskInput) {
executeSpec.addStatement(
@@ -368,7 +571,7 @@ class BuilderProcessor : AbstractProcessor() {
}
return TypeSpec
- .classBuilder(name)
+ .classBuilder(decl.className)
.addSuperinterface(Task::class.java)
.addModifiers(Modifier.PUBLIC, Modifier.FINAL, Modifier.STATIC)
.addMethod(executeSpec.build())
@@ -479,6 +682,15 @@ class BuilderProcessor : AbstractProcessor() {
}
}
+/** One [Builder.Task]-annotated method with its resolved id and data
parameters. */
+private class TaskDeclaration(
+ val method: ExecutableElement,
+ val id: String,
+ val dataParams: List<DataParam>,
+) {
+ val className: String =
method.simpleName.toString().replaceFirstChar(Char::uppercase)
+}
+
/**
* One data parameter of a task method, positioned among its peers, read into
* [local] by the generated body. [isTaskInput] marks a [TaskInput] parameter,
@@ -501,13 +713,34 @@ 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 val REFS_TYPE = ClassName.get(Refs::class.java)
+private val ARG_TYPE = ClassName.get(Arg::class.java)
+private val TASK_HANDLE_TYPE = ClassName.get(TaskRef::class.java)
+private val DEPS_TYPE = ClassName.get(Deps::class.java)
private const val DAG_ANNOTATION = "org.apache.airflow.sdk.Builder.Dag"
private const val TASK_ANNOTATION = "org.apache.airflow.sdk.Builder.Task"
+private val RESERVED_VIEW_NAMES =
+ setOf(
+ "depends",
+ "lit",
+ "clone",
+ "equals",
+ "finalize",
+ "getClass",
+ "hashCode",
+ "notify",
+ "notifyAll",
+ "toString",
+ "wait",
+ )
+
private val DAG_STRUCTURAL_ATTRIBUTES = setOf("id", "to")
private val TASK_STRUCTURAL_ATTRIBUTES = setOf("id")
+private fun TypeName.boxIfPossible(): TypeName = if (this == TypeName.VOID ||
isPrimitive) box() else this
+
private fun ProcessingEnvironment.isType(
t: TypeMirror,
c: ClassName,
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 6decad1a906..572fde64e9c 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
@@ -40,8 +40,8 @@ private fun JavaFileObjectSubject.hasSourceEquivalentTo(
class BuilderTest {
@Test
- @DisplayName("generate builder for dag class")
- fun generateBuilderForDagClass() {
+ @DisplayName("generate builder and task-reference twins for dag class")
+ fun generateBuilderAndRefForDagClass() {
val compilation =
compile(
"""
@@ -65,6 +65,14 @@ class BuilderTest {
public void t3(Context ctx, int value) {
System.out.println(String.format("%s %s", ctx.ti, value));
}
+
+ @Builder.Deps
+ static class Wiring implements TestExampleDeps {
+ void depends() {
+ t1();
+ t3(t2());
+ }
+ }
}
""",
)
@@ -80,20 +88,18 @@ class BuilderTest {
import java.lang.Exception;
import java.lang.Integer;
import java.lang.Override;
+ import java.util.List;
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.Refs;
import org.apache.airflow.sdk.internal.TaskArgs;
public final class TestExampleBuilder {
public static DagDef build() {
var dag = new DagDef("TestExample");
- dag.addTask(new TaskDef("t1", T1.class));
- dag.addTask(new TaskDef("t2", T2.class));
- dag.addTask(new TaskDef("t3", T3.class));
- return dag;
+ return Refs.record(dag, List.of("t1", "t2", "t3"), new
TestExample.Wiring()::depends);
}
public static final class T1 implements Task {
@@ -121,6 +127,43 @@ class BuilderTest {
}
""",
)
+ assertThat(compilation)
+ .generatedSourceFile("org.apache.airflow.example.TestExampleDeps")
+ .hasSourceEquivalentTo(
+ "org.apache.airflow.example.TestExampleDeps",
+ """
+ package org.apache.airflow.example;
+
+ import java.lang.Integer;
+ import java.lang.Void;
+ import org.apache.airflow.sdk.Arg;
+ import org.apache.airflow.sdk.Deps;
+ import org.apache.airflow.sdk.TaskDef;
+ import org.apache.airflow.sdk.TaskRef;
+ import org.apache.airflow.sdk.internal.Refs;
+
+ /**
+ * Wiring view of {@link TestExample}'s task methods, for declaring
its task graph.
+ *
+ * <p>Calling one registers its task with the Dag being built;
passing the handle it
+ * returned into another call feeds the upstream's output into that
task's parameter
+ * and wires the data edge. {@code before} and {@code after} wire an
ordering-only edge.
+ */
+ public interface TestExampleDeps extends Deps {
+ default TaskRef<Void> t1() {
+ return Refs.node(new TaskDef("t1", TestExampleBuilder.T1.class));
+ }
+
+ default TaskRef<Integer> t2() {
+ return Refs.node(new TaskDef("t2", TestExampleBuilder.T2.class));
+ }
+
+ default TaskRef<Void> t3(Arg<? extends Integer> value) {
+ return Refs.call(new TaskDef("t3", TestExampleBuilder.T3.class),
value);
+ }
+ }
+ """,
+ )
}
@Test
@@ -137,6 +180,11 @@ class BuilderTest {
public class TestExample {
@Builder.Task
public void t(long first, Client client, String second, Context ctx,
Integer third) {}
+
+ @Builder.Deps
+ static class Wiring implements TestExampleDeps {
+ void depends() { t(lit(1L), lit("second"), lit(3)); }
+ }
}
""",
)
@@ -154,18 +202,18 @@ class BuilderTest {
import java.lang.Long;
import java.lang.Override;
import java.lang.String;
+ import java.util.List;
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.Refs;
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;
+ return Refs.record(dag, List.of("t"), new
TestExample.Wiring()::depends);
}
public static final class T implements Task {
@@ -197,6 +245,11 @@ class BuilderTest {
public class TestExample {
@Builder.Task
public void t(boolean flag, float fraction, Double boxed,
List<String> tags, Map raw) {}
+
+ @Builder.Deps
+ static class Wiring implements TestExampleDeps {
+ void depends() { t(lit(true), lit(1f), lit(2.0),
lit(List.of("a")), lit(Map.of())); }
+ }
}
""",
)
@@ -221,15 +274,14 @@ class BuilderTest {
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.Refs;
import org.apache.airflow.sdk.internal.TaskArgs;
import org.apache.airflow.sdk.internal.TypeRef;
public final class TestExampleBuilder {
public static DagDef build() {
var dag = new DagDef("TestExample");
- dag.addTask(new TaskDef("t", T.class));
- return dag;
+ return Refs.record(dag, List.of("t"), new
TestExample.Wiring()::depends);
}
public static final class T implements Task {
@@ -249,6 +301,93 @@ class BuilderTest {
)
}
+ @Test
+ @DisplayName("type twin inputs by declared parameter type")
+ fun generateRefTypesTwinInputs() {
+ val compilation =
+ compile(
+ """
+ package org.apache.airflow.example;
+ import java.util.List;
+ import org.apache.airflow.sdk.Builder;
+ @Builder.Dag
+ public class TestExample {
+ @Builder.Task
+ public String ps() { return "x"; }
+
+ @Builder.Task
+ public void pv() {}
+
+ @Builder.Task
+ public List<String> pl() { return null; }
+
+ @Builder.Task
+ public long pn() { return 1L; }
+
+ @Builder.Task
+ public void t(String text, Object anything, List<String> items, Long
boxed) {}
+
+ @Builder.Deps
+ static class Wiring implements TestExampleDeps {
+ void depends() {
+ t(ps(), pv(), pl(), pn());
+ }
+ }
+ }
+ """,
+ )
+
+ assertThat(compilation).succeeded()
+ assertThat(compilation)
+ .generatedSourceFile("org.apache.airflow.example.TestExampleDeps")
+ .hasSourceEquivalentTo(
+ "org.apache.airflow.example.TestExampleDeps",
+ """
+ package org.apache.airflow.example;
+
+ import java.lang.Long;
+ import java.lang.String;
+ import java.lang.Void;
+ import java.util.List;
+ import org.apache.airflow.sdk.Arg;
+ import org.apache.airflow.sdk.Deps;
+ import org.apache.airflow.sdk.TaskDef;
+ import org.apache.airflow.sdk.TaskRef;
+ import org.apache.airflow.sdk.internal.Refs;
+
+ /**
+ * Wiring view of {@link TestExample}'s task methods, for declaring
its task graph.
+ *
+ * <p>Calling one registers its task with the Dag being built;
passing the handle it
+ * returned into another call feeds the upstream's output into that
task's parameter
+ * and wires the data edge. {@code before} and {@code after} wire an
ordering-only edge.
+ */
+ public interface TestExampleDeps extends Deps {
+ default TaskRef<String> ps() {
+ return Refs.node(new TaskDef("ps", TestExampleBuilder.Ps.class));
+ }
+
+ default TaskRef<Void> pv() {
+ return Refs.node(new TaskDef("pv", TestExampleBuilder.Pv.class));
+ }
+
+ default TaskRef<List<String>> pl() {
+ return Refs.node(new TaskDef("pl", TestExampleBuilder.Pl.class));
+ }
+
+ default TaskRef<Long> pn() {
+ return Refs.node(new TaskDef("pn", TestExampleBuilder.Pn.class));
+ }
+
+ default TaskRef<Void> t(Arg<? extends String> text, Arg<?> anything,
+ Arg<? extends List<String>> items, Arg<? extends Long> boxed) {
+ return Refs.call(new TaskDef("t", TestExampleBuilder.T.class),
text, anything, items, boxed);
+ }
+ }
+ """,
+ )
+ }
+
@Test
@DisplayName("lower explicit annotation attributes into config calls")
fun generateBuilderLowersConfigAttributes() {
@@ -262,6 +401,13 @@ class BuilderTest {
public class TestExample {
@Builder.Task(retries = 2, queue = "q", retryDelay = "PT5M",
retryExponentialBackoff = 1.5)
public void t1() {}
+
+ @Builder.Deps
+ static class Wiring implements TestExampleDeps {
+ void depends() {
+ t1();
+ }
+ }
}
""",
)
@@ -276,14 +422,13 @@ class BuilderTest {
import java.lang.Exception;
import java.lang.Override;
- import java.time.Duration;
import java.time.OffsetDateTime;
import java.util.List;
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.Refs;
public final class TestExampleBuilder {
public static DagDef build() {
@@ -292,8 +437,7 @@ class BuilderTest {
dag.config("tags", List.of("a", "b"));
dag.config("catchup", true);
dag.config("start_date",
OffsetDateTime.parse("2026-01-01T00:00:00Z"));
- dag.addTask(new TaskDef("t1", T1.class).config("retries",
2).config("queue", "q").config("retry_delay",
Duration.parse("PT5M")).config("retry_exponential_backoff", 1.5));
- return dag;
+ return Refs.record(dag, List.of("t1"), new
TestExample.Wiring()::depends);
}
public static final class T1 implements Task {
@@ -305,6 +449,34 @@ class BuilderTest {
}
""",
)
+ assertThat(compilation)
+ .generatedSourceFile("org.apache.airflow.example.TestExampleDeps")
+ .hasSourceEquivalentTo(
+ "org.apache.airflow.example.TestExampleDeps",
+ """
+ package org.apache.airflow.example;
+
+ import java.lang.Void;
+ import java.time.Duration;
+ import org.apache.airflow.sdk.Deps;
+ import org.apache.airflow.sdk.TaskDef;
+ import org.apache.airflow.sdk.TaskRef;
+ import org.apache.airflow.sdk.internal.Refs;
+
+ /**
+ * Wiring view of {@link TestExample}'s task methods, for declaring
its task graph.
+ *
+ * <p>Calling one registers its task with the Dag being built;
passing the handle it
+ * returned into another call feeds the upstream's output into that
task's parameter
+ * and wires the data edge. {@code before} and {@code after} wire an
ordering-only edge.
+ */
+ public interface TestExampleDeps extends Deps {
+ default TaskRef<Void> t1() {
+ return Refs.node(new TaskDef("t1",
TestExampleBuilder.T1.class).config("retries", 2).config("queue",
"q").config("retry_delay",
Duration.parse("PT5M")).config("retry_exponential_backoff", 1.5));
+ }
+ }
+ """,
+ )
}
@Test
@@ -319,6 +491,11 @@ class BuilderTest {
public class TestExample {
@Builder.Task
public void t(String args, int other) {}
+
+ @Builder.Deps
+ static class Wiring implements TestExampleDeps {
+ void depends() { t(lit("a"), lit(1)); }
+ }
}
""",
)
@@ -335,18 +512,18 @@ class BuilderTest {
import java.lang.Integer;
import java.lang.Override;
import java.lang.String;
+ import java.util.List;
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.Refs;
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;
+ return Refs.record(dag, List.of("t"), new
TestExample.Wiring()::depends);
}
public static final class T implements Task {
@Override
@@ -369,6 +546,7 @@ class BuilderTest {
compile(
"""
package org.apache.airflow.example;
+ import java.util.List;
import org.apache.airflow.sdk.Builder;
import org.apache.airflow.sdk.TaskInput;
@Builder.Dag
@@ -382,6 +560,11 @@ class BuilderTest {
@Builder.Task
public void named(ScoreInput context) {}
+
+ @Builder.Deps
+ static class Wiring implements TestExampleDeps {
+ void depends() { flat(lit("a"), lit(1)); named(lit(null)); }
+ }
}
""",
)
@@ -398,20 +581,19 @@ class BuilderTest {
import java.lang.Integer;
import java.lang.Override;
import java.lang.String;
+ import java.util.List;
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.Refs;
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;
+ return Refs.record(dag, List.of("flat", "named"), new
TestExample.Wiring()::depends);
}
public static final class Flat implements Task {
@@ -611,7 +793,13 @@ class BuilderTest {
"""
package org.apache.airflow.example;
import org.apache.airflow.sdk.Builder;
- @Builder.Dag(id = "foo") public class TestExample {}
+ @Builder.Dag(id = "foo")
+ public class TestExample {
+ @Builder.Deps
+ static class Wiring implements TestExampleDeps {
+ void depends() {}
+ }
+ }
""",
)
assertThat(compilation)
@@ -620,9 +808,14 @@ class BuilderTest {
"org.apache.airflow.example.TestExampleBuilder",
"""
package org.apache.airflow.example;
+ import java.util.List;
import org.apache.airflow.sdk.DagDef;
+ import org.apache.airflow.sdk.internal.Refs;
public final class TestExampleBuilder {
- public static DagDef build() { var dag = new DagDef("foo"); return
dag; }
+ public static DagDef build() {
+ var dag = new DagDef("foo");
+ return Refs.record(dag, List.of(), new
TestExample.Wiring()::depends);
+ }
}
""",
)
@@ -636,7 +829,13 @@ class BuilderTest {
"""
package org.apache.airflow.example;
import org.apache.airflow.sdk.Builder;
- @Builder.Dag(to = "Foo") public class TestExample {}
+ @Builder.Dag(to = "Foo")
+ public class TestExample {
+ @Builder.Deps
+ static class Wiring implements TestExampleDeps {
+ void depends() {}
+ }
+ }
""",
)
assertThat(compilation)
@@ -645,9 +844,14 @@ class BuilderTest {
"org.apache.airflow.example.Foo",
"""
package org.apache.airflow.example;
+ import java.util.List;
import org.apache.airflow.sdk.DagDef;
+ import org.apache.airflow.sdk.internal.Refs;
public final class Foo {
- public static DagDef build() { var dag = new DagDef("TestExample");
return dag; }
+ public static DagDef build() {
+ var dag = new DagDef("TestExample");
+ return Refs.record(dag, List.of(), new
TestExample.Wiring()::depends);
+ }
}
""",
)
@@ -664,6 +868,13 @@ class BuilderTest {
@Builder.Dag
public class TestExample {
@Builder.Task(id = "foo") public void t1() {}
+
+ @Builder.Deps
+ static class Wiring implements TestExampleDeps {
+ void depends() {
+ t1();
+ }
+ }
}
""",
)
@@ -677,17 +888,17 @@ class BuilderTest {
import java.lang.Exception;
import java.lang.Override;
+ import java.util.List;
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.Refs;
public final class TestExampleBuilder {
public static DagDef build() {
var dag = new DagDef("TestExample");
- dag.addTask(new TaskDef("foo", T1.class));
- return dag;
+ return Refs.record(dag, List.of("foo"), new
TestExample.Wiring()::depends);
}
public static final class T1 implements Task {
@@ -699,6 +910,370 @@ class BuilderTest {
}
""",
)
+ assertThat(compilation)
+ .generatedSourceFile("org.apache.airflow.example.TestExampleDeps")
+ .hasSourceEquivalentTo(
+ "org.apache.airflow.example.TestExampleDeps",
+ """
+ package org.apache.airflow.example;
+
+ import java.lang.Void;
+ import org.apache.airflow.sdk.Deps;
+ import org.apache.airflow.sdk.TaskDef;
+ import org.apache.airflow.sdk.TaskRef;
+ import org.apache.airflow.sdk.internal.Refs;
+
+ /**
+ * Wiring view of {@link TestExample}'s task methods, for declaring
its task graph.
+ *
+ * <p>Calling one registers its task with the Dag being built;
passing the handle it
+ * returned into another call feeds the upstream's output into that
task's parameter
+ * and wires the data edge. {@code before} and {@code after} wire an
ordering-only edge.
+ */
+ public interface TestExampleDeps extends Deps {
+ default TaskRef<Void> t1() {
+ return Refs.node(new TaskDef("foo", TestExampleBuilder.T1.class));
+ }
+ }
+ """,
+ )
+ }
+
+ @Test
+ @DisplayName("reject wiring that feeds an incompatible upstream type")
+ fun rejectIncompatibleWiring() {
+ val compilation =
+ compile(
+ """
+ package org.apache.airflow.example;
+ import org.apache.airflow.sdk.Builder;
+ @Builder.Dag
+ public class TestExample {
+ @Builder.Task
+ public String ps() { return "x"; }
+
+ @Builder.Task
+ public void t(int v) {}
+
+ @Builder.Deps
+ static class Wiring implements TestExampleDeps {
+ void depends() { t(ps()); }
+ }
+ }
+ """,
+ )
+ assertThat(compilation).failed()
+ assertThat(compilation).hadErrorContaining("incompatible types")
+ }
+
+ @Test
+ @DisplayName("reject wiring a numeric upstream into a narrower numeric
parameter")
+ fun rejectLossyNumericWiring() {
+ val compilation =
+ compile(
+ """
+ package org.apache.airflow.example;
+ import org.apache.airflow.sdk.Builder;
+ @Builder.Dag
+ public class TestExample {
+ @Builder.Task
+ public double ratio() { return 2.7; }
+
+ @Builder.Task
+ public void load(long rows) {}
+
+ @Builder.Deps
+ static class Wiring implements TestExampleDeps {
+ void depends() { load(ratio()); }
+ }
+ }
+ """,
+ )
+ assertThat(compilation).failed()
+ assertThat(compilation).hadErrorContaining("incompatible types")
+ }
+
+ @Test
+ @DisplayName("reject a dag class with no wiring class")
+ fun rejectDagClassWithoutWiringClass() {
+ val compilation =
+ compile(
+ """
+ package org.apache.airflow.example;
+ import org.apache.airflow.sdk.Builder;
+ @Builder.Dag
+ public class TestExample {
+ @Builder.Task public void t1() {}
+ }
+ """,
+ )
+ assertThat(compilation).failed()
+ assertThat(compilation).hadErrorContaining(
+ "Dag class TestExample must declare a @Builder.Deps class implementing
TestExampleDeps " +
+ "to declare its task graph; a class of task bodies for a Dag the
Python file owns carries " +
+ "@Builder.TaskHandler instead",
+ )
+ }
+
+ @Test
+ @DisplayName("reject more than one wiring class")
+ fun rejectMultipleWiringClasses() {
+ val compilation =
+ compile(
+ """
+ package org.apache.airflow.example;
+ import org.apache.airflow.sdk.Builder;
+ @Builder.Dag
+ public class TestExample {
+ @Builder.Task public void t1() {}
+
+ @Builder.Deps
+ static class One implements TestExampleDeps {
+ void depends() { t1(); }
+ }
+
+ @Builder.Deps
+ static class Two implements TestExampleDeps {
+ void depends() {}
+ }
+ }
+ """,
+ )
+ assertThat(compilation).failed()
+ assertThat(compilation).hadErrorContaining(
+ "Dag class TestExample declares more than one @Builder.Deps class: One,
Two",
+ )
+ }
+
+ @Test
+ @DisplayName("reject a non-static wiring class")
+ fun rejectNonStaticWiringClass() {
+ val compilation =
+ compile(
+ """
+ package org.apache.airflow.example;
+ import org.apache.airflow.sdk.Builder;
+ @Builder.Dag
+ public class TestExample {
+ @Builder.Task public void t1() {}
+
+ @Builder.Deps
+ class Wiring implements TestExampleDeps {
+ void depends() { t1(); }
+ }
+ }
+ """,
+ )
+ assertThat(compilation).failed()
+ assertThat(compilation).hadErrorContaining(
+ "@Builder.Deps class 'Wiring' must be static and non-private",
+ )
+ }
+
+ @Test
+ @DisplayName("reject a wiring class with no depends() method")
+ fun rejectWiringClassWithoutDepends() {
+ val compilation =
+ compile(
+ """
+ package org.apache.airflow.example;
+ import org.apache.airflow.sdk.Builder;
+ @Builder.Dag
+ public class TestExample {
+ @Builder.Task public void t1() {}
+
+ @Builder.Deps
+ static class Wiring implements TestExampleDeps {
+ void wireItUp() { t1(); }
+ }
+ }
+ """,
+ )
+ assertThat(compilation).failed()
+ assertThat(compilation).hadErrorContaining(
+ "@Builder.Deps class 'Wiring' must have a non-private, no-argument
depends() method",
+ )
+ }
+
+ @Test
+ @DisplayName("reject a wiring class that implements another Dag's wiring
view")
+ fun rejectWiringClassOfAnotherDag() {
+ val compilation =
+ compile(
+ """
+ package org.apache.airflow.example;
+ import org.apache.airflow.sdk.Builder;
+ @Builder.Dag
+ public class TestExample {
+ @Builder.Task public void t1() {}
+
+ @Builder.Deps
+ static class Wiring implements OtherDeps {
+ void depends() { t1(); }
+ }
+ }
+
+ @Builder.Dag
+ class Other {
+ @Builder.Task public void t1() {}
+
+ @Builder.Deps
+ static class Wiring implements OtherDeps {
+ void depends() { t1(); }
+ }
+ }
+ """,
+ )
+ assertThat(compilation).failed()
+ assertThat(compilation).hadErrorContaining(
+ "@Builder.Deps class 'Wiring' must implement TestExampleDeps, the wiring
view of TestExample",
+ )
+ }
+
+ @Test
+ @DisplayName("reject an abstract wiring class")
+ fun rejectAbstractWiringClass() {
+ val compilation =
+ compile(
+ """
+ package org.apache.airflow.example;
+ import org.apache.airflow.sdk.Builder;
+ @Builder.Dag
+ public class TestExample {
+ @Builder.Task public void t1() {}
+
+ @Builder.Deps
+ abstract static class Wiring implements TestExampleDeps {
+ void depends() { t1(); }
+ }
+ }
+ """,
+ )
+ assertThat(compilation).failed()
+ assertThat(compilation).hadErrorContaining("@Builder.Deps 'Wiring' must be
a concrete class")
+ }
+
+ @Test
+ @DisplayName("reject a wiring class with no no-argument constructor")
+ fun rejectWiringClassWithoutNoArgConstructor() {
+ val compilation =
+ compile(
+ """
+ package org.apache.airflow.example;
+ import org.apache.airflow.sdk.Builder;
+ @Builder.Dag
+ public class TestExample {
+ @Builder.Task public void t1() {}
+
+ @Builder.Deps
+ static class Wiring implements TestExampleDeps {
+ Wiring(int unused) {}
+ void depends() { t1(); }
+ }
+ }
+ """,
+ )
+ assertThat(compilation).failed()
+ assertThat(compilation).hadErrorContaining(
+ "@Builder.Deps class 'Wiring' needs a non-private no-argument
constructor",
+ )
+ }
+
+ @Test
+ @DisplayName("reject a depends() that throws a checked exception")
+ fun rejectDependsThrowingCheckedException() {
+ val compilation =
+ compile(
+ """
+ package org.apache.airflow.example;
+ import org.apache.airflow.sdk.Builder;
+ @Builder.Dag
+ public class TestExample {
+ @Builder.Task public void t1() {}
+
+ @Builder.Deps
+ static class Wiring implements TestExampleDeps {
+ void depends() throws java.io.IOException { t1(); }
+ }
+ }
+ """,
+ )
+ assertThat(compilation).failed()
+ assertThat(compilation).hadErrorContaining(
+ "depends() of @Builder.Deps class 'Wiring' must not throw checked
exceptions: java.io.IOException",
+ )
+ }
+
+ @Test
+ @DisplayName("accept a depends() the wiring class inherits")
+ fun acceptInheritedDepends() {
+ val compilation =
+ compile(
+ """
+ package org.apache.airflow.example;
+ import org.apache.airflow.sdk.Builder;
+ @Builder.Dag
+ public class TestExample {
+ @Builder.Task public void t1() {}
+
+ static class Base implements TestExampleDeps {
+ void depends() { t1(); }
+ }
+
+ @Builder.Deps
+ static class Wiring extends Base implements TestExampleDeps {}
+ }
+ """,
+ )
+ assertThat(compilation).succeeded()
+ }
+
+ @Test
+ @DisplayName("reject a wiring class outside a Dag class")
+ fun rejectMisplacedWiringClass() {
+ val compilation =
+ compile(
+ """
+ package org.apache.airflow.example;
+ import org.apache.airflow.sdk.Builder;
+ public class TestExample {
+ @Builder.Deps
+ static class Wiring {
+ void depends() {}
+ }
+ }
+ """,
+ )
+ assertThat(compilation).failed()
+ assertThat(compilation).hadErrorContaining(
+ "@Builder.Deps class 'Wiring' must be nested directly in a @Builder.Dag
class",
+ )
+ }
+
+ @Test
+ @DisplayName("reject a task method whose name clashes with a wiring-view
member")
+ fun rejectTaskNameClashingWithViewMember() {
+ val compilation =
+ compile(
+ """
+ package org.apache.airflow.example;
+ import org.apache.airflow.sdk.Builder;
+ @Builder.Dag
+ public class TestExample {
+ @Builder.Task(id = "wire") public void depends() {}
+
+ @Builder.Deps
+ static class Wiring implements TestExampleDeps {
+ void depends() {}
+ }
+ }
+ """,
+ )
+ assertThat(compilation).failed()
+ assertThat(compilation).hadErrorContaining(
+ "Task method 'depends' clashes with a member of the wiring view; rename
the method and keep " +
+ "the task id with @Builder.Task(id = \"wire\")",
+ )
}
@Test
@@ -719,6 +1294,78 @@ class BuilderTest {
)
}
+ @Test
+ @DisplayName("reject duplicate task ids")
+ fun rejectDuplicateTaskIds() {
+ val compilation =
+ compile(
+ """
+ package org.apache.airflow.example;
+ import org.apache.airflow.sdk.Builder;
+ @Builder.Dag
+ public class TestExample {
+ @Builder.Task(id = "x")
+ public void t1() {}
+
+ @Builder.Task(id = "x")
+ public void t2() {}
+ }
+ """,
+ )
+ assertThat(compilation).failed()
+ assertThat(compilation).hadErrorContaining("Tasks in Dag have duplicate
ID: x")
+ }
+
+ @Test
+ @DisplayName("reject overloaded task methods, whatever their parameters")
+ fun rejectOverloadedTaskMethods() {
+ val compilation =
+ compile(
+ """
+ package org.apache.airflow.example;
+ import org.apache.airflow.sdk.Builder;
+ @Builder.Dag
+ public class TestExample {
+ @Builder.Task(id = "a") public void extract() {}
+ @Builder.Task(id = "b") public void extract(String text) {}
+
+ @Builder.Deps
+ static class Wiring implements TestExampleDeps {
+ void depends() {}
+ }
+ }
+ """,
+ )
+ assertThat(compilation).failed()
+ assertThat(compilation).hadErrorContaining(
+ "Dag class TestExample overloads task method 'extract'; a method's name
is the name of its " +
+ "generated task class and of its wiring-view method, so rename one and
keep its task id " +
+ "with @Builder.Task(id = \"b\")",
+ )
+ }
+
+ @Test
+ @DisplayName("reject overloaded task-handler methods")
+ fun rejectOverloadedTaskHandlerMethods() {
+ val compilation =
+ compile(
+ """
+ package org.apache.airflow.example;
+ import org.apache.airflow.sdk.Builder;
+ public class TestExample {
+ @Builder.TaskHandler(dag = "etl", task = "a") public void score() {}
+ @Builder.TaskHandler(dag = "etl", task = "b") public void
score(String text) {}
+ }
+ """,
+ )
+ assertThat(compilation).failed()
+ assertThat(compilation).hadErrorContaining(
+ "Class TestExample overloads task-handler method 'score'; a method's
name is the name of its " +
+ "generated task class, so rename one and keep its task id with " +
+ "@Builder.TaskHandler(task = \"b\")",
+ )
+ }
+
@Test
@DisplayName("reject a duration attribute that is not ISO-8601")
fun rejectInvalidDurationAttribute() {
@@ -730,6 +1377,11 @@ class BuilderTest {
@Builder.Dag
public class TestExample {
@Builder.Task(retryDelay = "5 minutes") public void t() {}
+
+ @Builder.Deps
+ static class Wiring implements TestExampleDeps {
+ void depends() { t(); }
+ }
}
""",
)
@@ -750,6 +1402,11 @@ class BuilderTest {
@Builder.Dag(tags = {"say \"hi\"", "back\\slash"})
public class TestExample {
@Builder.Task public void t() {}
+
+ @Builder.Deps
+ static class Wiring implements TestExampleDeps {
+ void depends() { t(); }
+ }
}
""",
)
@@ -761,6 +1418,37 @@ class BuilderTest {
.contains("""dag.config("tags", List.of("say \"hi\"",
"back\\slash"));""")
}
+ @Test
+ @DisplayName("escape quotes and backslashes in the task ids the wiring must
register")
+ fun generateBuilderEscapesWiredTaskIds() {
+ val compilation =
+ compile(
+ """
+ package org.apache.airflow.example;
+ import org.apache.airflow.sdk.Builder;
+ @Builder.Dag
+ public class TestExample {
+ @Builder.Task(id = "say \"hi\"") public void a() {}
+ @Builder.Task(id = "back\\slash") public void b() {}
+
+ @Builder.Deps
+ static class Wiring implements TestExampleDeps {
+ void depends() {
+ a();
+ b();
+ }
+ }
+ }
+ """,
+ )
+
+ assertThat(compilation).succeeded()
+ assertThat(compilation)
+ .generatedSourceFile("org.apache.airflow.example.TestExampleBuilder")
+ .contentsAsUtf8String()
+ .contains("""List.of("say \"hi\"", "back\\slash")""")
+ }
+
@Test
@DisplayName("bind a TaskInput through the shared populator")
fun generateBuilderBindsTaskInputFields() {
@@ -783,6 +1471,11 @@ class BuilderTest {
@Builder.Task
public double score(Client client, ScoreInput input) { return
input.threshold; }
+
+ @Builder.Deps
+ static class Wiring implements TestExampleDeps {
+ void depends() { score(lit(null)); }
+ }
}
""",
)
@@ -797,18 +1490,18 @@ class BuilderTest {
import java.lang.Exception;
import java.lang.Override;
+ import java.util.List;
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.Refs;
public final class TestExampleBuilder {
public static DagDef build() {
var dag = new DagDef("TestExample");
- dag.addTask(new TaskDef("score", Score.class));
- return dag;
+ return Refs.record(dag, List.of("score"), new
TestExample.Wiring()::depends);
}
public static final class Score implements Task {
diff --git a/java-sdk/sdk/build.gradle.kts b/java-sdk/sdk/build.gradle.kts
index 20b491b3721..c4178f5148d 100644
--- a/java-sdk/sdk/build.gradle.kts
+++ b/java-sdk/sdk/build.gradle.kts
@@ -517,7 +517,8 @@ abstract class GenerateDagDslTask : DefaultTask() {
| * Container for the annotation-based Dag-authoring API.
| *
| * Annotating a class with [Dag] generates a `<Class>Builder`
whose static
- | * `build()` returns the [DagDef] to add to a [Bundle].
+ | * `build()` returns the [DagDef] to add to a [Bundle], and a
`<Class>Deps`
+ | * wiring view for the class's [Deps] class to implement.
| *
| * Example:
| *
@@ -530,13 +531,19 @@ abstract class GenerateDagDslTask : DefaultTask() {
| *
| * @Builder.Task(id = "transform")
| * public long transform(Client client, long extracted) { ...
}
+ | *
+ | * @Builder.Deps
+ | * static class Wiring implements MyPipelineDeps {
+ | * void depends() { transform(extract()); }
+ | * }
| * }
| * ```
| *
- | * A task method's data parameters — everything other than the
injected
- | * [Client] and [Context] — receive, by position, the arguments
the Python
- | * `@task.stub` call site bound. Keyword arguments bind by name
instead
- | * through a single [TaskInput] parameter.
+ | * A task method's data parameters, meaning every parameter other
than the
+ | * injected [Client] and [Context], receive by position the
inputs the
+ | * [Deps] class wired. For a task the Python Dag file declares
with `@task.stub`,
+ | * the arguments bound at that call site take their place. Keyword
+ | * arguments bind by name instead through a single [TaskInput]
parameter.
| */
|class Builder internal constructor() {
| /**
@@ -600,6 +607,35 @@ abstract class GenerateDagDslTask : DefaultTask() {
| val dag: String,
| val task: String = "",
| )
+ |
+ | /**
+ | * Marks the nested class that declares this Dag's task graph.
+ | *
+ | * Declare it as a `static` nested class that implements the
generated
+ | * `<Dag>Deps` wiring view and has a no-argument `depends()`
method.
+ | * Calling a view method registers its task; passing the handle
one
+ | * returned into another call wires a data edge; `before` and
`after`
+ | * wire an ordering-only one:
+ | *
+ | * ```java
+ | * @Builder.Deps
+ | * static class Wiring implements EtlPipelineDeps {
+ | * void depends() {
+ | * var rows = extract();
+ | * var loaded = load(transform(rows, lit(0.9)));
+ | * rows.before(audit());
+ | * report().after(loaded, audit());
+ | * }
+ | * }
+ | * ```
+ | *
+ | * Every [Dag] class declares one, because the graph is what
the Dag
+ | * owns. A class that supplies only task bodies, for a Dag a
Python
+ | * file declares, carries [TaskHandler] instead.
+ | */
+ | @Target(AnnotationTarget.CLASS)
+ | @MustBeDocumented
+ | annotation class Deps
|}
|
""".trimMargin(),
diff --git a/java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/Arg.kt
b/java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/Arg.kt
index 556ed5a7d4a..bff37ad02ef 100644
--- a/java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/Arg.kt
+++ b/java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/Arg.kt
@@ -28,7 +28,21 @@ package org.apache.airflow.sdk
*
* @param T Type of the value.
*/
-sealed class Arg<T>
+sealed class Arg<T> {
+ companion object {
+ /**
+ * Wraps an inline constant as a task argument, passed to the task as a
+ * constant and creating no dependency edge.
+ *
+ * [Deps.lit] is the spelling a [Builder.Deps] class uses; this is the same
+ * call for code with no `Deps` in scope.
+ *
+ * @param value Constant to bind; may be null for a nullable parameter.
+ */
+ @JvmStatic
+ fun <T> lit(value: T?): Arg<T> = LiteralArg(value)
+ }
+}
internal class LiteralArg<T>(
internal val value: T?,
@@ -37,8 +51,10 @@ internal class LiteralArg<T>(
/**
* The output of a registered task, and the task's place in the flow.
*
- * [Deps.Flow.before] and [Deps.Flow.after] wire an ordering-only edge from
- * this task, where nothing flows but the sequence.
+ * Passing this handle into another task's arguments feeds this task's return
+ * value into that parameter and wires the data edge. [Deps.Flow.before] and
+ * [Deps.Flow.after] wire an ordering-only edge instead, where nothing flows
+ * but the sequence.
*
* @param T Return type of the task this handle refers to.
*/
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 e1dff618894..605b2b46ddf 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,7 +19,10 @@
package org.apache.airflow.sdk
+import org.apache.airflow.sdk.internal.builderName
import org.apache.airflow.sdk.internal.registrarName
+import java.lang.reflect.InvocationTargetException
+import java.lang.reflect.Method
/**
* All [DagDef]s that this JVM process can execute.
@@ -89,29 +92,31 @@ class Bundle(
}
/**
- * Registers every task handler a class holds, from the ids each
- * [Builder.TaskHandler] names.
+ * Registers what an annotated class holds, read from the class itself so
+ * there is no second name to keep in sync.
*
- * @param handlerClass A class with [Builder.TaskHandler] methods.
+ * A [Builder.Dag] class contributes the Dag its generated builder builds; a
+ * class of [Builder.TaskHandler] methods contributes each handler, bound to
+ * the Dag the Python file owns. A class can carry both.
+ *
+ * @param annotated A class carrying [Builder.Dag] or [Builder.TaskHandler].
* @return This bundle, for chaining.
- * @throws IllegalArgumentException if the class has no generated
- * registrar, because annotation processing did not run over it.
+ * @throws IllegalArgumentException if the class has no generated code,
+ * because annotation processing did not run over it, or if the Dag's
+ * wiring is invalid, such as a task the `@Builder.Deps` class did not
+ * call.
*/
- fun register(handlerClass: Class<*>): Bundle {
+ fun register(annotated: 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)
+ val dag = annotated.getAnnotation(Builder.Dag::class.java)
+ if (dag != null) {
+ val builder = generated(builderName(annotated.packageName,
annotated.simpleName, dag.to), annotated, "builder")
+ register(invokeGenerated(builder.getMethod("build")) as DagDef)
+ }
+ if (dag == null || annotated.declaredMethods.any {
it.isAnnotationPresent(Builder.TaskHandler::class.java) }) {
+ val registrar = generated(registrarName(annotated.name), annotated,
"registrar")
+ invokeGenerated(registrar.getMethod("registerInto", Bundle::class.java),
this)
+ }
return this
}
@@ -150,6 +155,31 @@ class Bundle(
taskId: String,
): TaskDef? = (dags[dagId] ?: taskHandlers[dagId])?.tasks?.get(taskId)
+ private fun generated(
+ name: String,
+ from: Class<*>,
+ what: String,
+ ): Class<*> =
+ try {
+ Class.forName(name, true, from.classLoader)
+ } catch (e: ClassNotFoundException) {
+ throw IllegalArgumentException(
+ "No generated $what $name for ${from.name}; does it carry @Builder.Dag
or " +
+ "@Builder.TaskHandler, and is airflow-sdk-processor on the
annotationProcessor path?",
+ e,
+ )
+ }
+
+ private fun invokeGenerated(
+ method: Method,
+ vararg args: Any?,
+ ): Any? =
+ try {
+ method.invoke(null, *args)
+ } catch (e: InvocationTargetException) {
+ throw e.cause ?: e
+ }
+
/**
* 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
diff --git a/java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/Deps.kt
b/java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/Deps.kt
index 38a14d00354..e70df9f8534 100644
--- a/java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/Deps.kt
+++ b/java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/Deps.kt
@@ -19,7 +19,13 @@
package org.apache.airflow.sdk
-/** Vocabulary for declaring a Dag's task graph in Java. */
+/**
+ * Vocabulary for declaring a Dag's task graph in Java, and the base of every
+ * generated `<Dag>Deps` wiring view.
+ *
+ * A [Builder.Deps] class inherits [lit] and [Flow] by simple name, so it needs
+ * no import and `Flow` does not collide with `java.util.concurrent.Flow`.
+ */
interface Deps {
/**
* A point in the task graph: one task, or a set of them.
@@ -81,6 +87,15 @@ interface Deps {
fun of(vararg flows: Flow): Flow = FlowSet(flows.flatMap { it.nodes() })
}
}
+
+ /**
+ * Wraps an inline constant as a task argument, as in
+ * `transform(extract(), lit(0.9))`. It is passed to the task as a constant
+ * and creates no dependency edge.
+ *
+ * @param value Constant to bind; may be null for a nullable parameter.
+ */
+ fun <T> lit(value: T?): Arg<T> = Arg.lit(value)
}
/** Several tasks as one point in the flow, which no single [TaskRef] can
represent. */
diff --git
a/java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/internal/Refs.kt
b/java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/internal/Refs.kt
new file mode 100644
index 00000000000..39021f03b4e
--- /dev/null
+++ b/java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/internal/Refs.kt
@@ -0,0 +1,122 @@
+/*
+ * 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.Arg
+import org.apache.airflow.sdk.DagDef
+import org.apache.airflow.sdk.TaskDef
+import org.apache.airflow.sdk.TaskRef
+
+/**
+ * @suppress
+ *
+ * The recorder behind a `@Builder.Deps` wiring class. Public so that
+ * processor-generated wiring views can call it; not user-facing API.
+ *
+ * A wiring view's methods are `default` methods on an interface, so they hold
+ * no Dag of their own. [record] puts the Dag being built in scope for exactly
+ * the duration of one `depends()` call, and the view's methods register into
+ * it. Nothing parses a syntax tree: identity travels with the [TaskRef] a
+ * call returns, so a result held in a local and reused just works.
+ */
+object Refs {
+ private class Recording(
+ val dag: DagDef,
+ ) {
+ val byTaskId = linkedMapOf<String, TaskRef<*>>()
+ }
+
+ private val recording = ThreadLocal<Recording?>()
+
+ /**
+ * Runs one `depends()` call with [dag] in scope, then returns the Dag the
+ * wiring built.
+ *
+ * @throws IllegalArgumentException if the wiring left a declared task
+ * unregistered.
+ */
+ @JvmStatic
+ fun record(
+ dag: DagDef,
+ taskIds: List<String>,
+ depends: Runnable,
+ ): DagDef {
+ check(recording.get() == null) { "Dag wiring is already being recorded on
this thread" }
+ recording.set(Recording(dag))
+ try {
+ depends.run()
+ } finally {
+ recording.remove()
+ }
+ val missing = taskIds.filterNot { it in dag.tasks }
+ require(missing.isEmpty()) {
+ "Wiring for Dag '${dag.id}' did not register task(s)
${missing.joinToString { "'$it'" }}: " +
+ "every @Builder.Task method must be called in the @Builder.Deps class"
+ }
+ return dag
+ }
+
+ /**
+ * Records a task that takes no data arguments.
+ *
+ * @return The handle representing this task, memoized by [TaskDef.id] so
every call
+ * yields the same one.
+ */
+ @JvmStatic
+ fun <T> node(def: TaskDef): TaskRef<T> = call(def)
+
+ /**
+ * Records a task and the data edge for every [TaskRef] among [args]; a
+ * literal argument records a baked value and no edge.
+ *
+ * @return The handle representing this task, memoized by [TaskDef.id] so a
result
+ * held in a local and reused refers to one node.
+ * @throws IllegalArgumentException if an argument is a raw Java `null`
+ * rather than `lit(null)`.
+ */
+ @JvmStatic
+ @Suppress("UNCHECKED_CAST", "SpreadOperator")
+ fun <T> call(
+ def: TaskDef,
+ vararg args: Arg<*>?,
+ ): TaskRef<T> {
+ val inputs =
+ args.mapIndexed { i, arg ->
+ requireNotNull(arg) {
+ "Argument ${i + 1} of task '${def.id}' is null; wrap a null constant
as lit(null)"
+ }
+ }
+ val active =
+ checkNotNull(recording.get()) {
+ "Task '${def.id}' was wired outside a @Builder.Deps class; the wiring
view's methods " +
+ "only record while the generated builder is running depends()"
+ }
+ active.byTaskId[def.id]?.let { existing ->
+ require(inputs.isEmpty()) {
+ "Task '${def.id}' is wired more than once with arguments; call it once
and reuse the handle it returned"
+ }
+ return existing as TaskRef<T>
+ }
+ inputs.filterIsInstance<TaskRef<*>>().forEach { def.dependsOn(it.def) }
+ def.inputs += inputs
+ active.dag.addTask(def)
+ return TaskRef<T>(def).also { active.byTaskId[def.id] = it }
+ }
+}
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
index 3220ae78c02..30edadb3060 100644
--- 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
@@ -34,3 +34,26 @@ package org.apache.airflow.sdk.internal
* returns it.
*/
fun registrarName(binaryName: String): String = "${binaryName.replace('$',
'_')}Handlers"
+
+/**
+ * @suppress
+ *
+ * Names the builder generated for a `@Builder.Dag` class: a top-level class
+ * in the Dag class's package, named by the annotation's `to` or, when that is
+ * blank, `<Class>Builder`.
+ *
+ * Public so the annotation processor emits the name the runtime looks up; not
+ * user-facing API.
+ *
+ * @param packageName Package of the Dag class; empty for the default package.
+ * @param simpleName Simple name of the Dag class.
+ * @param to The annotation's `to` attribute.
+ */
+fun builderName(
+ packageName: String,
+ simpleName: String,
+ to: String,
+): String {
+ val name = to.ifBlank { "${simpleName}Builder" }
+ return if (packageName.isEmpty()) name else "$packageName.$name"
+}
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
index 1cc3ff3dbae..6733621ab0b 100644
--- 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
@@ -24,6 +24,7 @@ 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.internal.Refs
import org.apache.airflow.sdk.execution.comm.TaskInstance as CommTaskInstance
/** Records getXCom calls and serves canned values keyed by task id. */
@@ -104,8 +105,8 @@ internal class NoopTask : Task {
/** A context whose task was wired by its Dag with the given inputs. */
internal fun contextWiredWith(inputs: List<Arg<*>>): Context {
+ val dag = DagDef("d")
val def = TaskDef("t", NoopTask::class.java)
- DagDef("d").addTask(def)
- def.inputs += inputs
+ Refs.record(dag, listOf("t")) { Refs.call<Unit>(def, *inputs.toTypedArray())
}
return taskContext().also { it.taskDef = def }
}
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 29c5fc8a99b..4c645aab89f 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
@@ -34,6 +34,21 @@ internal class BundleTest {
/** A class of handlers that the processor generated
[BundleTest_NestedHandlers] for. */
class Nested
+ /** A Dag class that the processor generated [WiredDagBuilder] for. */
+ @Builder.Dag(id = "wired")
+ class WiredDag
+
+ /** A Dag class whose generated [BrokenDagBuilder] fails, as an invalid
wiring would. */
+ @Builder.Dag
+ class BrokenDag
+
+ /** A Dag class that also holds a task handler, so the processor generated
both. */
+ @Builder.Dag(id = "mixed")
+ class MixedDag {
+ @Builder.TaskHandler(dag = "etl", task = "score")
+ fun score() = Unit
+ }
+
@Test
@DisplayName("Should index dags by dagId")
fun shouldIndexDagsByDagId() {
@@ -134,14 +149,54 @@ internal class BundleTest {
Assertions.assertEquals(emptySet<String>(), bundle.dags.keys)
}
+ @Test
+ @DisplayName("Should register the Dag a @Builder.Dag class's generated
builder builds")
+ fun shouldRegisterDagFromBuilderClass() {
+ val bundle = Bundle().register(WiredDag::class.java)
+
+ Assertions.assertEquals(listOf("wired"), bundle.dags.keys.toList())
+ Assertions.assertEquals(emptySet<String>(), bundle.taskHandlers.keys)
+ }
+
+ @Test
+ @DisplayName("Should rethrow the failure of a generated builder rather than
its reflection wrapper")
+ fun shouldUnwrapBuilderFailure() {
+ val error =
+ Assertions.assertThrows(IllegalArgumentException::class.java) {
+ Bundle().register(BrokenDag::class.java)
+ }
+
+ Assertions.assertEquals("wiring failed", error.message)
+ }
+
+ @Test
+ @DisplayName("Should register both the Dag and the handlers of a class that
carries both")
+ fun shouldRegisterDagAndHandlersOfOneClass() {
+ val bundle = Bundle().register(MixedDag::class.java)
+
+ Assertions.assertEquals(listOf("mixed"), bundle.dags.keys.toList())
+ Assertions.assertEquals(
+ listOf("score"),
+ bundle.taskHandlers
+ .getValue("etl")
+ .tasks.keys
+ .toList(),
+ )
+ }
+
@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())
+ Assertions.assertEquals(
+ listOf("score"),
+ bundle.taskHandlers
+ .getValue("etl")
+ .tasks.keys
+ .toList(),
+ )
}
@Test
@@ -152,10 +207,10 @@ internal class BundleTest {
Bundle().register(NoOp::class.java)
}
- Assertions.assertTrue(
- error.message!!.startsWith(
- "No generated registrar org.apache.airflow.sdk.BundleTest_NoOpHandlers
for ",
- ),
+ Assertions.assertEquals(
+ "No generated registrar org.apache.airflow.sdk.BundleTest_NoOpHandlers
for " +
+ "${NoOp::class.java.name}; does it carry @Builder.Dag or
@Builder.TaskHandler, " +
+ "and is airflow-sdk-processor on the annotationProcessor path?",
error.message,
)
}
@@ -240,6 +295,48 @@ internal class BundleTest {
}
}
+class NoopBundleTask : Task {
+ override fun execute(
+ context: Context,
+ client: Client,
+ ) = Unit
+}
+
+/** Stands in for the builder the annotation processor generates for
[BundleTest.WiredDag]. */
+class WiredDagBuilder {
+ companion object {
+ @JvmStatic
+ fun build() = DagDef("wired")
+ }
+}
+
+/** Stands in for a generated builder whose wiring is invalid. */
+class BrokenDagBuilder {
+ companion object {
+ @JvmStatic
+ fun build(): DagDef = throw IllegalArgumentException("wiring failed")
+ }
+}
+
+/** Stands in for the builder the annotation processor generates for
[BundleTest.MixedDag]. */
+class MixedDagBuilder {
+ companion object {
+ @JvmStatic
+ fun build() = DagDef("mixed")
+ }
+}
+
+/** Stands in for the registrar the annotation processor generates for
[BundleTest.MixedDag]. */
+@Suppress("ktlint:standard:class-naming", "ClassName")
+class BundleTest_MixedDagHandlers {
+ companion object {
+ @JvmStatic
+ fun registerInto(bundle: Bundle) {
+ bundle.register("etl", "score", NoopBundleTask::class.java)
+ }
+ }
+}
+
/**
* Stands in for the registrar the annotation processor generates beside
* [BundleTest.Nested], to pin the name [Bundle.register] looks up.
@@ -249,14 +346,7 @@ class BundleTest_NestedHandlers {
companion object {
@JvmStatic
fun registerInto(bundle: Bundle) {
- bundle.register("etl", "score", NoOpHandler::class.java)
+ bundle.register("etl", "score", NoopBundleTask::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
index ca557b92611..fadbf6fad61 100644
--- 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
@@ -122,9 +122,9 @@ internal class InputTaskTest {
@DisplayName("Should decode a TaskInput wholesale from its wired input when
no bindings arrive")
fun shouldDecodeTaskInputFromWiredInput() {
// A native Dag has no stub call site, so there are no argument names to
- // match fields against: the input the Dag wired to this task decodes into
- // the whole TaskInput at once.
- val context = contextWiredWith(listOf(LiteralArg(mapOf("region" to "emea",
"threshold" to 0.5))))
+ // match fields against: the input the wiring class fed this task decodes
+ // into the whole TaskInput at once.
+ val context = contextWiredWith(listOf(Arg.lit(mapOf("region" to "emea",
"threshold" to 0.5))))
val (client, _) = clientWith(null)
val task = Summarize()
diff --git
a/java-sdk/sdk/src/test/kotlin/org/apache/airflow/sdk/internal/ArgValuesTest.kt
b/java-sdk/sdk/src/test/kotlin/org/apache/airflow/sdk/internal/ArgValuesTest.kt
index 768ad066f5c..7d8439634aa 100644
---
a/java-sdk/sdk/src/test/kotlin/org/apache/airflow/sdk/internal/ArgValuesTest.kt
+++
b/java-sdk/sdk/src/test/kotlin/org/apache/airflow/sdk/internal/ArgValuesTest.kt
@@ -51,7 +51,7 @@ private class NoopArgTask : Task {
) = Unit
}
-/** Resolution of the inputs a Dag declared in Java, without runtime bindings.
*/
+/** Resolution of the inputs a `@Builder.Deps` class recorded, without runtime
bindings. */
internal class ArgValuesTest {
/** Upstream task ids read through the transport, in arrival order. */
private val pulls = java.util.concurrent.CopyOnWriteArrayList<String>()
@@ -119,8 +119,7 @@ internal class ArgValuesTest {
.distinct()
.forEach { dag.addTask(it) }
val def = TaskDef("consumer", NoopArgTask::class.java)
- dag.addTask(def)
- def.inputs += inputs
+ Refs.record(dag, listOf("consumer")) { Refs.call<Unit>(def,
*inputs.toTypedArray()) }
return contextWithoutTaskDef().also { it.taskDef = def }
}
diff --git
a/java-sdk/sdk/src/test/kotlin/org/apache/airflow/sdk/internal/RefsTest.kt
b/java-sdk/sdk/src/test/kotlin/org/apache/airflow/sdk/internal/RefsTest.kt
new file mode 100644
index 00000000000..ffb77fa7b40
--- /dev/null
+++ b/java-sdk/sdk/src/test/kotlin/org/apache/airflow/sdk/internal/RefsTest.kt
@@ -0,0 +1,158 @@
+/*
+ * 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.Arg
+import org.apache.airflow.sdk.Client
+import org.apache.airflow.sdk.Context
+import org.apache.airflow.sdk.DagDef
+import org.apache.airflow.sdk.LiteralArg
+import org.apache.airflow.sdk.Task
+import org.apache.airflow.sdk.TaskDef
+import org.apache.airflow.sdk.TaskRef
+import org.junit.jupiter.api.Assertions.assertEquals
+import org.junit.jupiter.api.Assertions.assertSame
+import org.junit.jupiter.api.Assertions.assertThrows
+import org.junit.jupiter.api.DisplayName
+import org.junit.jupiter.api.Test
+
+private class NoopRefTask : Task {
+ override fun execute(
+ context: Context,
+ client: Client,
+ ) = Unit
+}
+
+internal class RefsTest {
+ @Test
+ @DisplayName("Should register the task, record inputs, and wire handle
edges")
+ fun shouldRegisterTaskWithInputsAndEdges() {
+ val dag = DagDef("d")
+ Refs.record(dag, listOf("p", "c")) {
+ val producer = Refs.node<Long>(TaskDef("p", NoopRefTask::class.java))
+ Refs.call<Unit>(TaskDef("c", NoopRefTask::class.java), producer,
Arg.lit(5))
+ }
+
+ val consumerDef = dag.tasks.getValue("c")
+ assertEquals(setOf("p", "c"), dag.tasks.keys)
+ assertEquals(setOf(dag.tasks.getValue("p")), consumerDef.upstreams)
+ assertEquals(2, consumerDef.inputs.size)
+ assertEquals(dag.tasks.getValue("p"), (consumerDef.inputs[0] as
TaskRef<*>).def)
+ assertEquals(5, (consumerDef.inputs[1] as LiteralArg<*>).value)
+ }
+
+ @Test
+ @DisplayName("Should return the same handle wherever a task is wired")
+ fun shouldMemoizeHandleByTaskId() {
+ val dag = DagDef("d")
+ Refs.record(dag, listOf("a", "b")) {
+ val first = Refs.node<Unit>(TaskDef("a", NoopRefTask::class.java))
+ val again = Refs.node<Unit>(TaskDef("a", NoopRefTask::class.java))
+ assertSame(first, again)
+ first.before(Refs.node<Unit>(TaskDef("b", NoopRefTask::class.java)))
+ }
+
+ assertEquals(setOf("a", "b"), dag.tasks.keys)
+ assertEquals(setOf(dag.tasks.getValue("a")),
dag.tasks.getValue("b").upstreams)
+ }
+
+ @Test
+ @DisplayName("Should pass when the wiring registered every task")
+ fun shouldPassWhenWiringComplete() {
+ val dag = DagDef("d")
+
+ Refs.record(dag, listOf("t")) { Refs.node<Unit>(TaskDef("t",
NoopRefTask::class.java)) }
+ }
+
+ @Test
+ @DisplayName("Should fail naming the tasks the wiring missed")
+ fun shouldFailNamingMissedTasks() {
+ val dag = DagDef("d")
+
+ val error =
+ assertThrows(IllegalArgumentException::class.java) {
+ Refs.record(dag, listOf("t", "x", "y")) { Refs.node<Unit>(TaskDef("t",
NoopRefTask::class.java)) }
+ }
+
+ assertEquals(
+ "Wiring for Dag 'd' did not register task(s) 'x', 'y': " +
+ "every @Builder.Task method must be called in the @Builder.Deps class",
+ error.message,
+ )
+ }
+
+ @Test
+ @DisplayName("Should refuse a wiring call made outside a recording")
+ fun shouldRefuseWiringOutsideRecording() {
+ val error =
+ assertThrows(IllegalStateException::class.java) {
+ Refs.node<Unit>(TaskDef("t", NoopRefTask::class.java))
+ }
+
+ assertEquals(
+ "Task 't' was wired outside a @Builder.Deps class; the wiring view's
methods " +
+ "only record while the generated builder is running depends()",
+ error.message,
+ )
+ }
+
+ @Test
+ @DisplayName("Should reject a raw null argument, pointing to lit(null)")
+ fun shouldRejectRawNullArgument() {
+ val error =
+ assertThrows(IllegalArgumentException::class.java) {
+ Refs.record(DagDef("d"), listOf("t")) {
+ Refs.call<Unit>(TaskDef("t", NoopRefTask::class.java), Arg.lit(1),
null)
+ }
+ }
+
+ assertEquals("Argument 2 of task 't' is null; wrap a null constant as
lit(null)", error.message)
+ }
+
+ @Test
+ @DisplayName("Should reject a task wired a second time with arguments")
+ fun shouldRejectTaskWiredTwiceWithArguments() {
+ val error =
+ assertThrows(IllegalArgumentException::class.java) {
+ Refs.record(DagDef("d"), listOf("t")) {
+ Refs.node<Unit>(TaskDef("t", NoopRefTask::class.java))
+ Refs.call<Unit>(TaskDef("t", NoopRefTask::class.java), Arg.lit(1))
+ }
+ }
+
+ assertEquals(
+ "Task 't' is wired more than once with arguments; call it once and reuse
the handle it returned",
+ error.message,
+ )
+ }
+
+ @Test
+ @DisplayName("Should refuse to record a Dag while another is being recorded")
+ fun shouldRefuseNestedRecording() {
+ val error =
+ assertThrows(IllegalStateException::class.java) {
+ Refs.record(DagDef("outer"), emptyList()) {
+ Refs.record(DagDef("inner"), emptyList()) {}
+ }
+ }
+
+ assertEquals("Dag wiring is already being recorded on this thread",
error.message)
+ }
+}
diff --git
a/java-sdk/example/src/java/org/apache/airflow/example/ExampleBundleBuilder.java
b/java-sdk/sdk/src/test/kotlin/org/apache/airflow/sdk/internal/RegistrarTest.kt
similarity index 63%
copy from
java-sdk/example/src/java/org/apache/airflow/example/ExampleBundleBuilder.java
copy to
java-sdk/sdk/src/test/kotlin/org/apache/airflow/sdk/internal/RegistrarTest.kt
index d75f5762ba1..ad32b92e0b9 100644
---
a/java-sdk/example/src/java/org/apache/airflow/example/ExampleBundleBuilder.java
+++
b/java-sdk/sdk/src/test/kotlin/org/apache/airflow/sdk/internal/RegistrarTest.kt
@@ -17,20 +17,16 @@
* under the License.
*/
-package org.apache.airflow.example;
+package org.apache.airflow.sdk.internal
-import org.apache.airflow.sdk.*;
+import org.junit.jupiter.api.Assertions.assertEquals
+import org.junit.jupiter.api.DisplayName
+import org.junit.jupiter.api.Test
-public class ExampleBundleBuilder {
- public static Bundle build() {
- return new Bundle()
- .register(InterfaceExampleBuilder.build())
- .register(AnnotationExample.class)
- .register(XComCastingExample.class)
-
.register(org.apache.airflow.example.nativedag.InterfaceExample.build());
- }
-
- public static void main(String[] args) {
- Server.create(args).serve(build());
+internal class RegistrarTest {
+ @Test
+ @DisplayName("Should name the builder of a Dag class in the default package
with no leading dot")
+ fun shouldNameBuilderInDefaultPackage() {
+ assertEquals("FooBuilder", builderName("", "Foo", ""))
}
}
diff --git
a/kubernetes-tests/lang_sdk/java_example/src/java/org/apache/airflow/k8sexample/CombinedExample.java
b/kubernetes-tests/lang_sdk/java_example/src/java/org/apache/airflow/k8sexample/CombinedExample.java
index 63aaa5b6f05..09b7b67c8e6 100644
---
a/kubernetes-tests/lang_sdk/java_example/src/java/org/apache/airflow/k8sexample/CombinedExample.java
+++
b/kubernetes-tests/lang_sdk/java_example/src/java/org/apache/airflow/k8sexample/CombinedExample.java
@@ -24,22 +24,21 @@ import org.apache.airflow.sdk.*;
import org.slf4j.Logger;
import org.slf4j.LoggerFactory;
-// Java half of the KubernetesExecutor lang-SDK system test bundle. Registers
the
-// Java tasks of the shared "lang_sdk_combined" Dag; the Go tasks of the same
-// dag_id live in ../go_example and the Python stub Dag in ../dags. The
-// coordinator locates this jar by dag_id, so only the Java tasks are
registered
-// here.
[email protected](id = "lang_sdk_combined")
+// Java half of the KubernetesExecutor lang-SDK system test bundle. Supplies
the
+// bodies of the "lang_sdk_combined" Dag's Java stub tasks; the Dag itself is
+// declared by the Python file in ../dags, and the Go tasks of the same dag_id
+// live in ../go_example. The coordinator locates this jar by dag_id, so only
the
+// Java tasks are registered here.
public class CombinedExample {
private static final Logger logger =
LoggerFactory.getLogger(CombinedExample.class);
- @Builder.Task(id = "java_extract")
+ @Builder.TaskHandler(dag = "lang_sdk_combined", task = "java_extract")
public long extract(Client client) {
logger.info("java_extract running");
return new Date().getTime();
}
- @Builder.Task(id = "java_transform")
+ @Builder.TaskHandler(dag = "lang_sdk_combined", task = "java_transform")
public void transform(Client client) {
var variable = client.getVariable("my_variable");
logger.info("java_transform obtained variable {}", variable);
diff --git
a/kubernetes-tests/lang_sdk/java_example/src/java/org/apache/airflow/k8sexample/K8sBundleBuilder.java
b/kubernetes-tests/lang_sdk/java_example/src/java/org/apache/airflow/k8sexample/K8sBundleBuilder.java
index c3b511b898c..dd5b53b7983 100644
---
a/kubernetes-tests/lang_sdk/java_example/src/java/org/apache/airflow/k8sexample/K8sBundleBuilder.java
+++
b/kubernetes-tests/lang_sdk/java_example/src/java/org/apache/airflow/k8sexample/K8sBundleBuilder.java
@@ -19,21 +19,16 @@
package org.apache.airflow.k8sexample;
-import java.util.List;
import org.apache.airflow.sdk.*;
-import org.jetbrains.annotations.NotNull;
-// CombinedExampleBuilder is generated by the annotation processor from the
-// @Builder.Dag class CombinedExample.
-public class K8sBundleBuilder implements BundleBuilder {
- @NotNull
- @Override
- public Iterable<DagDef> getDags() {
- return List.of(CombinedExampleBuilder.build());
+// The bundle holds task handlers only: the "lang_sdk_combined" Dag is the
Python
+// file's, so nothing here declares a Dag.
+public class K8sBundleBuilder {
+ public static Bundle build() {
+ return new Bundle().register(CombinedExample.class);
}
public static void main(String[] args) {
- var bundle = new K8sBundleBuilder().build();
- Server.create(args).serve(bundle);
+ Server.create(args).serve(build());
}
}