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 8043c703de1 Java SDK: Group a native Dag's tasks with task groups
(#74230)
8043c703de1 is described below
commit 8043c703de19af786e2794fc816e9158f802a43e
Author: Jason(Zhe-You) Liu <[email protected]>
AuthorDate: Thu Oct 8 10:32:35 2026 +0800
Java SDK: Group a native Dag's tasks with task groups (#74230)
* Java SDK: Group a native Dag's tasks with task groups
* Java SDK: Reserve the whole group view surface and pin group edge order
The reserved names of a group's wiring view are now read from
Deps.TaskGroup, so
endpoints, before and after are rejected as task method names inside a
group, not
only groupId and nodes. A task named before would otherwise win overload
resolution over the inherited before(Flow...) and silently register a task.
Add tests for that clash, for duplicate group IDs in one scope, and for the
order
group edges resolve in: an empty group prefers an earlier edge over its
parent,
and edges resolve in the order they were drawn.
Correct the Task groups docs and TaskGroupRef KDoc: only what a group holds
is
read once, while edges still resolve in drawing order. Restore the missing
transform, load and audit calls in the group wiring example so it parses.
---
.../language-sdks/java.rst | 63 +++
.../example/nativedag/AnnotationExample.java | 17 +-
.../org/apache/airflow/sdk/BuilderProcessor.kt | 320 ++++++++++---
.../kotlin/org/apache/airflow/sdk/BuilderTest.kt | 510 ++++++++++++++++++++-
java-sdk/sdk/build.gradle.kts | 43 +-
.../main/kotlin/org/apache/airflow/sdk/Bundle.kt | 4 +-
.../main/kotlin/org/apache/airflow/sdk/DagDef.kt | 147 +++++-
.../src/main/kotlin/org/apache/airflow/sdk/Deps.kt | 95 +++-
.../main/kotlin/org/apache/airflow/sdk/Endpoint.kt | 29 ++
.../kotlin/org/apache/airflow/sdk/TaskGroupRef.kt | 111 +++++
.../kotlin/org/apache/airflow/sdk/internal/Ids.kt | 29 ++
.../kotlin/org/apache/airflow/sdk/internal/Refs.kt | 56 ++-
.../org/apache/airflow/sdk/ArgTestSupport.kt | 2 +-
.../kotlin/org/apache/airflow/sdk/TaskGroupTest.kt | 286 ++++++++++++
.../apache/airflow/sdk/internal/ArgValuesTest.kt | 2 +-
.../org/apache/airflow/sdk/internal/RefsTest.kt | 82 +++-
16 files changed, 1680 insertions(+), 116 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 f158f795cca..446557db8ea 100644
--- a/airflow-core/docs/authoring-and-scheduling/language-sdks/java.rst
+++ b/airflow-core/docs/authoring-and-scheduling/language-sdks/java.rst
@@ -694,6 +694,69 @@ class that supplies only task bodies, for a Dag a Python
file declares, carries
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.
+Task groups
+~~~~~~~~~~~
+
+A task group gathers tasks that the Airflow UI shows as one node, as Python's
``TaskGroup`` does.
+Everything declared in a group carries the group's ID as a prefix, so task
``stage`` in group
+``staging`` is the task ``staging.stage``. On the interface surface,
``taskGroup`` declares a group on
+the Dag or inside another group, and the group declares its tasks:
+
+.. code-block:: java
+
+ var staging = dag.taskGroup("staging");
+ var stage = staging.task("stage", Stage.class); //
"staging.stage"
+ staging.taskGroup("checks").task("nulls", Nulls.class).after(stage); //
"staging.checks.nulls"
+ extract.before(staging);
+
+With annotations, a ``@Builder.TaskGroup`` class holds the tasks of one group,
and nesting one in
+another nests the groups:
+
+.. code-block:: java
+
+ @Builder.TaskGroup // the group "Staging", after the
class
+ static class Staging {
+ @Builder.Task
+ public long stage(long rows) { ... } // the task "Staging.stage"
+
+ @Builder.TaskGroup(id = "checks")
+ static class Checks {
+ @Builder.Task
+ public void nulls(long staged) { ... } // "Staging.checks.nulls"
+ }
+ }
+
+ @Builder.Deps
+ static class Wiring implements EtlPipelineDeps {
+ void depends() {
+ var rows = extract();
+ load(transform(rows, lit(0.9)));
+ rows.before(audit());
+ var staged = staging().stage(rows);
+ staging().checks().nulls(staged);
+ extract().before(staging());
+ }
+ }
+
+The generated view nests the same way, so a group is both the namespace of
what it holds and a point
+in the flow: ``staging().checks().nulls(staged)`` reaches a task, and
``extract().before(staging())``
+orders the whole group. Task method names scope to their own group, so two
groups can each declare
+``run()``. A group class is ``static``, non-private, and needs a no-argument
constructor, because the
+generated task bodies instantiate it.
+
+A group stands at either end of ``before``, ``after`` and ``Flow.of``. As an
upstream it stands for
+its leaves, the tasks nothing else in the group runs after; as a downstream,
for its roots, the tasks
+that run after nothing else in the group. A group ID contains only ASCII
letters, digits,
+underscores, or dashes, and no task or other group in the Dag can share it.
+
+.. note::
+
+ What a group holds is read once, when the Dag is serialized, which is what
lets the wiring class
+ above order a whole group before any of its tasks are declared, as
``extract().before(staging())``
+ does. Python instead reads it at each ``>>``. Edges are still resolved
in the order they were
+ drawn, as Python resolves them, so drawing an inner edge before or after
an outer one gives
+ different upstreams.
+
Configuration attributes
~~~~~~~~~~~~~~~~~~~~~~~~
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
index c2a3bc89f75..1c6ce70cc72 100644
---
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
@@ -54,9 +54,14 @@ public class AnnotationExample {
log.log(INFO, "Loaded {0}", transformed);
}
- @Builder.Task(id = "audit")
- public void audit() {
- log.log(INFO, "Audited the run");
+ // A task group: everything it declares is prefixed with its id, so this is
+ // the task "checks.audit".
+ @Builder.TaskGroup(id = "checks")
+ static class Checks {
+ @Builder.Task(id = "audit")
+ public void audit() {
+ log.log(INFO, "Audited the run");
+ }
}
// Implements the generated wiring view, so javac type-checks the graph:
@@ -66,8 +71,10 @@ public class AnnotationExample {
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());
+ // Ordering-only edge: the checks group runs after extract, with no data
+ // flowing.
+ extracted.before(checks());
+ checks().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 93e185a614e..1101dbb9084 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
@@ -32,6 +32,7 @@ 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.GROUP_ID
import org.apache.airflow.sdk.internal.Refs
import org.apache.airflow.sdk.internal.SchemaFields
import org.apache.airflow.sdk.internal.TaskArgs
@@ -57,6 +58,7 @@ import javax.lang.model.element.VariableElement
import javax.lang.model.type.TypeKind
import javax.lang.model.type.TypeMirror
import javax.tools.Diagnostic
+import java.lang.reflect.Modifier as ReflectModifier
import org.apache.airflow.sdk.internal.builderName as generatedBuilderName
/**
@@ -93,6 +95,7 @@ import org.apache.airflow.sdk.internal.builderName as
generatedBuilderName
@SupportedAnnotationTypes(
"org.apache.airflow.sdk.Builder.Dag",
"org.apache.airflow.sdk.Builder.Task",
+ "org.apache.airflow.sdk.Builder.TaskGroup",
"org.apache.airflow.sdk.Builder.TaskHandler",
"org.apache.airflow.sdk.Builder.Deps",
)
@@ -113,6 +116,20 @@ class BuilderProcessor : AbstractProcessor() {
)
}
}
+ roundEnv.getElementsAnnotatedWith(Builder.TaskGroup::class.java).forEach {
el ->
+ val owner = el.enclosingElement
+ val nested =
+ owner is TypeElement &&
+ (owner.getAnnotation(Builder.Dag::class.java) != null ||
owner.getAnnotation(Builder.TaskGroup::class.java) != null)
+ if (!nested) {
+ processingEnv.messager.printMessage(
+ Diagnostic.Kind.ERROR,
+ "@Builder.TaskGroup class '${el.simpleName}' must be nested in a
@Builder.Dag class or in " +
+ "another @Builder.TaskGroup class",
+ el,
+ )
+ }
+ }
roundEnv
.getElementsAnnotatedWith(Builder.TaskHandler::class.java)
.mapNotNull { it.enclosingElement as? TypeElement }
@@ -131,7 +148,8 @@ class BuilderProcessor : AbstractProcessor() {
with(processingEnv) {
runCatching {
val packageName =
elementUtils.getPackageOf(el).qualifiedName.toString()
- val declarations = collectTasks(el)
+ val scope = collectScope(el, emptyList(), emptyList())
+ checkIds(scope)
val builderName =
ClassName.get(
packageName,
@@ -139,12 +157,12 @@ class BuilderProcessor : AbstractProcessor() {
)
val depsName = ClassName.get(packageName, "${el.simpleName}Deps")
val deps = findDeps(el, depsName)
- declarations.forEach { checkViewName(it) }
+ checkViewNames(scope)
JavaFile
- .builder(packageName, buildBuilder(el, declarations, deps,
builderName))
+ .builder(packageName, buildBuilder(el, scope, deps, builderName))
.build()
.writeTo(filer)
- JavaFile.builder(packageName, buildDeps(el, declarations,
builderName, depsName)).build().writeTo(filer)
+ JavaFile.builder(packageName, buildDeps(el, scope, builderName,
depsName)).build().writeTo(filer)
}.onFailure { e ->
messager.printMessage(
Diagnostic.Kind.ERROR,
@@ -193,13 +211,13 @@ class BuilderProcessor : AbstractProcessor() {
require(handler.dag.isNotBlank()) {
"@Builder.TaskHandler on '${inner.simpleName}' must name the Dag the
Python file declares"
}
- val decl = TaskDeclaration(inner, handler.task.ifBlank {
inner.simpleName.toString() }, collectDataParams(inner))
+ val decl = TaskDeclaration(inner, handler.task.ifBlank {
inner.simpleName.toString() }, collectDataParams(inner), el)
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))
+ registrar.addType(buildTask(decl))
registerInto.addStatement(
$$"bundle.register($S, $S, $L.class)",
handler.dag,
@@ -214,10 +232,11 @@ class BuilderProcessor : AbstractProcessor() {
private fun buildBuilder(
el: TypeElement,
- declarations: List<TaskDeclaration>,
+ scope: Scope,
deps: TypeElement,
builderName: ClassName,
): TypeSpec {
+ val declarations = scope.allTasks()
val ann = dagAnnotation(el)
val builderClass =
@@ -234,16 +253,20 @@ class BuilderProcessor : AbstractProcessor() {
explicitConfig(el, DAG_ANNOTATION, DAG_STRUCTURAL_ATTRIBUTES,
SchemaFields.DAG).forEach { (key, value) ->
buildMethod.addStatement($$"dag.config($S, $L)", key, value)
}
+ val taskIds = CodeBlock.join(declarations.map { CodeBlock.of($$"$S",
it.id) }, ", ")
+ val groupIds = CodeBlock.join(scope.allGroups().map { CodeBlock.of($$"$S",
it.fullId) }, ", ")
buildMethod.addStatement(
- $$"return $T.record(dag, $T.of($L), new $T()::depends)",
+ $$"return $T.record(dag, $T.of($L), $T.of($L), new $T()::depends)",
REFS_TYPE,
- ClassName.get(List::class.java),
- CodeBlock.join(declarations.map { CodeBlock.of($$"$S", it.id) }, ", "),
+ LIST_TYPE,
+ taskIds,
+ LIST_TYPE,
+ groupIds,
ClassName.get(deps),
)
builderClass.addMethod(buildMethod.build())
- declarations.forEach { builderClass.addType(buildTask(it, el)) }
+ declarations.forEach { builderClass.addType(buildTask(it)) }
return builderClass.build()
}
@@ -258,7 +281,7 @@ class BuilderProcessor : AbstractProcessor() {
*/
private fun buildDeps(
el: TypeElement,
- declarations: List<TaskDeclaration>,
+ scope: Scope,
builderName: ClassName,
depsName: ClassName,
): TypeSpec {
@@ -271,31 +294,83 @@ class BuilderProcessor : AbstractProcessor() {
"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",
+ "and wires the data edge. {@code before} and {@code after} wire an
ordering-only edge.\n" +
+ "<p>A task group is reached by calling it, and stands at either
end of an edge:\n" +
+ "{@code staging().stage(rows)} and {@code
extract().before(staging())}.\n",
ClassName.get(el),
)
+ addScope(view, scope, builderName, depsName)
+ return view.build()
+ }
- for (decl in declarations) {
- val method =
+ /** Adds one scope's task methods, and a nested interface plus accessor per
group it holds. */
+ private fun addScope(
+ view: TypeSpec.Builder,
+ scope: Scope,
+ builderName: ClassName,
+ viewName: ClassName,
+ inGroup: Boolean = false,
+ ) {
+ scope.tasks.forEach { view.addMethod(viewMethod(it, builderName, inGroup))
}
+ for (group in scope.groups) {
+ val nested = viewName.nestedClass(group.element.simpleName.toString())
+ view.addMethod(
MethodSpec
- .methodBuilder(decl.method.simpleName.toString())
+ .methodBuilder(group.accessor)
.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())
+ .returns(nested)
+ .addJavadoc("The task group {@code \$L}, and everything declared in
it.\n", group.fullId)
+ .addStatement($$"return new $T() {}", nested)
+ .build(),
+ )
+ val groupView =
+ TypeSpec
+ .interfaceBuilder(nested)
+ .addModifiers(Modifier.PUBLIC, Modifier.STATIC)
+ .addSuperinterface(GROUP_TYPE)
+ .addJavadoc("Wiring view of the task group {@code \$L}.\n",
group.fullId)
+ .addMethod(
+ MethodSpec
+ .methodBuilder("groupId")
+ .addAnnotation(Override::class.java)
+ .addModifiers(Modifier.PUBLIC, Modifier.DEFAULT)
+ .returns(String::class.java)
+ .addStatement($$"return $S", group.fullId)
+ .build(),
+ )
+ addScope(groupView, group.scope, builderName, nested, inGroup = true)
+ view.addType(groupView.build())
}
- return view.build()
+ }
+
+ /** One task's method on the wiring view: injected arguments stripped,
inputs lifted to [Arg]. */
+ private fun viewMethod(
+ decl: TaskDeclaration,
+ builderName: ClassName,
+ inGroup: Boolean,
+ ): MethodSpec {
+ 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))
+ // The view knows the group it belongs to, so the recorder is told where
+ // the task goes instead of deriving it from the task's ID.
+ val group = if (inGroup) CodeBlock.of("groupId()") else
CodeBlock.of($$"$S", "")
+ if (decl.dataParams.isEmpty()) {
+ method.addStatement($$"return $T.node($L, $L)", REFS_TYPE, group, def)
+ } else {
+ method.addStatement(
+ $$"return $T.call($L, $L, $L)",
+ REFS_TYPE,
+ group,
+ def,
+ decl.dataParams.joinToString { it.name },
+ )
+ }
+ return method.build()
}
/**
@@ -326,22 +401,64 @@ class BuilderProcessor : AbstractProcessor() {
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>()
+ /** The Dag's tasks and task groups, read from the class tree the author
wrote. */
+ private fun collectScope(
+ el: TypeElement,
+ path: List<String>,
+ classPath: List<String>,
+ ): Scope {
+ val tasks = 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 " +
+ val localId = ann.id.ifBlank { inner.simpleName.toString() }
+ require(tasks.none {
it.method.simpleName.contentEquals(inner.simpleName) }) {
+ "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\")"
+ "task id with @Builder.Task(id = \"$localId\")"
+ }
+ tasks +=
+ TaskDeclaration(inner, (path + localId).joinToString("."),
collectDataParams(inner), el, classPath)
+ }
+
+ val groups = mutableListOf<GroupDeclaration>()
+ for (inner in el.enclosedElements.filterIsInstance<TypeElement>()) {
+ val ann = inner.getAnnotation(Builder.TaskGroup::class.java) ?: continue
+ val localId = checkGroupClass(inner, ann)
+ require(groups.none { it.id == localId }) {
+ "Class ${el.simpleName} declares more than one task group '$localId'"
}
- declarations += TaskDeclaration(inner, id, collectDataParams(inner))
+ val scope = collectScope(inner, path + localId, classPath +
inner.simpleName.toString())
+ groups += GroupDeclaration(inner, localId, (path +
localId).joinToString("."), scope)
}
- return declarations
+ return Scope(tasks, groups)
+ }
+
+ /** Checks that `new <group class>()` compiles and names a valid group, and
returns its local ID. */
+ private fun checkGroupClass(
+ el: TypeElement,
+ ann: Builder.TaskGroup,
+ ): String {
+ val name = el.simpleName
+ require(el.kind == ElementKind.CLASS && Modifier.ABSTRACT !in
el.modifiers) {
+ "@Builder.TaskGroup '$name' must be a concrete class"
+ }
+ require(Modifier.STATIC in el.modifiers && Modifier.PRIVATE !in
el.modifiers) {
+ "@Builder.TaskGroup class '$name' must be static and non-private"
+ }
+ require(
+ el.enclosedElements
+ .filterIsInstance<ExecutableElement>()
+ .any { it.kind == ElementKind.CONSTRUCTOR && it.parameters.isEmpty()
&& Modifier.PRIVATE !in it.modifiers },
+ ) {
+ "@Builder.TaskGroup class '$name' needs a non-private no-argument
constructor"
+ }
+ val id = ann.id.ifBlank { name.toString() }
+ require(GROUP_ID.matches(id)) {
+ "Task group ID '$id' must contain only ASCII letters, digits,
underscores, or dashes"
+ }
+ return id
}
/**
@@ -403,15 +520,67 @@ class BuilderProcessor : AbstractProcessor() {
}
/**
- * 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`.
+ * Rejects two tasks of the Dag sharing an ID, a task and a task group
+ * sharing one, and two task methods whose generated classes would collide.
+ */
+ private fun checkIds(scope: Scope) {
+ val declarations = scope.allTasks()
+ val taskIds = mutableSetOf<String>()
+ declarations.forEach { decl ->
+ require(taskIds.add(decl.id)) { "Tasks in Dag have duplicate ID:
${decl.id}" }
+ }
+ scope.allGroups().forEach { group ->
+ require(group.fullId !in taskIds) {
+ "Dag has both a task and a task group with ID '${group.fullId}';
rename one"
+ }
+ }
+ val byClassName = mutableMapOf<String, TaskDeclaration>()
+ declarations.forEach { decl ->
+ byClassName.put(decl.className, decl)?.let { first ->
+ throw IllegalArgumentException(
+ "Task methods '${first.id}' and '${decl.id}' both generate the task
class " +
+ "'${decl.className}'; rename one of them or an enclosing
@Builder.TaskGroup class",
+ )
+ }
+ }
+ }
+
+ /**
+ * Rejects a task method or task group whose wiring-view twin would clash
+ * with a member the view already has: `depends`, `lit`, a method of
+ * `Object`, or, inside a group, one of `Deps.TaskGroup`'s own. Names scope
to their
+ * own group, so only one scope is compared.
*/
- 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}\")"
+ private fun checkViewNames(
+ scope: Scope,
+ inGroup: Boolean = false,
+ ) {
+ val reserved = if (inGroup) RESERVED_VIEW_NAMES +
RESERVED_GROUP_VIEW_NAMES else RESERVED_VIEW_NAMES
+ scope.tasks.forEach { decl ->
+ val name = decl.method.simpleName.toString()
+ require(name !in reserved) {
+ "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.substringAfterLast('.')}\")"
+ }
+ }
+ val accessors = mutableMapOf<String, GroupDeclaration>()
+ scope.groups.forEach { group ->
+ require(group.accessor !in reserved) {
+ "Task group class '${group.element.simpleName}' clashes with a member
of the wiring view; " +
+ "rename the class and keep the group id with @Builder.TaskGroup(id =
\"${group.id}\")"
+ }
+ require(scope.tasks.none {
it.method.simpleName.contentEquals(group.accessor) }) {
+ "Task group class '${group.element.simpleName}' and task method
'${group.accessor}' would both " +
+ "be '${group.accessor}()' on the wiring view; rename one"
+ }
+ accessors[group.accessor]?.let { first ->
+ throw IllegalArgumentException(
+ "Task group classes '${first.element.simpleName}' and
'${group.element.simpleName}' would both " +
+ "be '${group.accessor}()' on the wiring view; rename one",
+ )
+ }
+ accessors[group.accessor] = group
+ checkViewNames(group.scope, inGroup = true)
}
}
@@ -502,10 +671,7 @@ class BuilderProcessor : AbstractProcessor() {
}
}
- private fun buildTask(
- decl: TaskDeclaration,
- parent: TypeElement,
- ): TypeSpec {
+ private fun buildTask(decl: TaskDeclaration): TypeSpec {
val executeSpec =
MethodSpec
.methodBuilder("execute")
@@ -564,7 +730,7 @@ class BuilderProcessor : AbstractProcessor() {
}.also {
executeSpec.addStatement(
it,
- ClassName.get(parent),
+ ClassName.get(decl.owner),
inner.simpleName,
innerArgs,
)
@@ -682,13 +848,43 @@ class BuilderProcessor : AbstractProcessor() {
}
}
-/** One [Builder.Task]-annotated method with its resolved id and data
parameters. */
+/** The tasks and task groups one class declares. */
+private class Scope(
+ val tasks: List<TaskDeclaration>,
+ val groups: List<GroupDeclaration>,
+) {
+ /** Every task of this scope and the groups beneath it, outermost first. */
+ fun allTasks(): List<TaskDeclaration> = tasks + groups.flatMap {
it.scope.allTasks() }
+
+ /** Every group beneath this scope, parents before the groups nested in
them. */
+ fun allGroups(): List<GroupDeclaration> = groups.flatMap { listOf(it) +
it.scope.allGroups() }
+}
+
+/** One `@Builder.TaskGroup` class, and what it declares. */
+private class GroupDeclaration(
+ val element: TypeElement,
+ val id: String,
+ val fullId: String,
+ val scope: Scope,
+) {
+ /** The view method that reaches this group, named after the class it is
declared as. */
+ val accessor: String =
element.simpleName.toString().replaceFirstChar(Char::lowercase)
+}
+
+/**
+ * One [Builder.Task]-annotated method with its resolved id and data
parameters.
+ * [owner] is the class that declares it, which the generated body
instantiates,
+ * and [classPath] the task-group classes enclosing it.
+ */
private class TaskDeclaration(
val method: ExecutableElement,
val id: String,
val dataParams: List<DataParam>,
+ val owner: TypeElement,
+ val classPath: List<String> = emptyList(),
) {
- val className: String =
method.simpleName.toString().replaceFirstChar(Char::uppercase)
+ val className: String =
+ (classPath +
method.simpleName.toString().replaceFirstChar(Char::uppercase)).joinToString("_")
}
/**
@@ -717,6 +913,8 @@ 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 val GROUP_TYPE = DEPS_TYPE.nestedClass("TaskGroup")
+private val LIST_TYPE = ClassName.get(List::class.java)
private const val DAG_ANNOTATION = "org.apache.airflow.sdk.Builder.Dag"
private const val TASK_ANNOTATION = "org.apache.airflow.sdk.Builder.Task"
@@ -736,6 +934,18 @@ private val RESERVED_VIEW_NAMES =
"wait",
)
+/**
+ * What a group's view inherits from `Deps.TaskGroup`, on top of
+ * [RESERVED_VIEW_NAMES]. Read from the interface so it cannot drift when a
+ * member is added there.
+ */
+private val RESERVED_GROUP_VIEW_NAMES: Set<String> =
+ Deps.TaskGroup::class.java
+ .methods
+ .filterNot { ReflectModifier.isStatic(it.modifiers) }
+ .map { it.name }
+ .toSet()
+
private val DAG_STRUCTURAL_ATTRIBUTES = setOf("id", "to")
private val TASK_STRUCTURAL_ATTRIBUTES = setOf("id")
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 572fde64e9c..d75c33aa45d 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
@@ -99,7 +99,7 @@ class BuilderTest {
public final class TestExampleBuilder {
public static DagDef build() {
var dag = new DagDef("TestExample");
- return Refs.record(dag, List.of("t1", "t2", "t3"), new
TestExample.Wiring()::depends);
+ return Refs.record(dag, List.of("t1", "t2", "t3"), List.of(), new
TestExample.Wiring()::depends);
}
public static final class T1 implements Task {
@@ -151,15 +151,15 @@ class BuilderTest {
*/
public interface TestExampleDeps extends Deps {
default TaskRef<Void> t1() {
- return Refs.node(new TaskDef("t1", TestExampleBuilder.T1.class));
+ return Refs.node("", new TaskDef("t1",
TestExampleBuilder.T1.class));
}
default TaskRef<Integer> t2() {
- return Refs.node(new TaskDef("t2", TestExampleBuilder.T2.class));
+ 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);
+ return Refs.call("", new TaskDef("t3",
TestExampleBuilder.T3.class), value);
}
}
""",
@@ -213,7 +213,7 @@ class BuilderTest {
public final class TestExampleBuilder {
public static DagDef build() {
var dag = new DagDef("TestExample");
- return Refs.record(dag, List.of("t"), new
TestExample.Wiring()::depends);
+ return Refs.record(dag, List.of("t"), List.of(), new
TestExample.Wiring()::depends);
}
public static final class T implements Task {
@@ -281,7 +281,7 @@ class BuilderTest {
public final class TestExampleBuilder {
public static DagDef build() {
var dag = new DagDef("TestExample");
- return Refs.record(dag, List.of("t"), new
TestExample.Wiring()::depends);
+ return Refs.record(dag, List.of("t"), List.of(), new
TestExample.Wiring()::depends);
}
public static final class T implements Task {
@@ -364,24 +364,24 @@ class BuilderTest {
*/
public interface TestExampleDeps extends Deps {
default TaskRef<String> ps() {
- return Refs.node(new TaskDef("ps", TestExampleBuilder.Ps.class));
+ return Refs.node("", new TaskDef("ps",
TestExampleBuilder.Ps.class));
}
default TaskRef<Void> pv() {
- return Refs.node(new TaskDef("pv", TestExampleBuilder.Pv.class));
+ return Refs.node("", new TaskDef("pv",
TestExampleBuilder.Pv.class));
}
default TaskRef<List<String>> pl() {
- return Refs.node(new TaskDef("pl", TestExampleBuilder.Pl.class));
+ return Refs.node("", new TaskDef("pl",
TestExampleBuilder.Pl.class));
}
default TaskRef<Long> pn() {
- return Refs.node(new TaskDef("pn", TestExampleBuilder.Pn.class));
+ 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);
+ return Refs.call("", new TaskDef("t",
TestExampleBuilder.T.class), text, anything, items, boxed);
}
}
""",
@@ -437,7 +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"));
- return Refs.record(dag, List.of("t1"), new
TestExample.Wiring()::depends);
+ return Refs.record(dag, List.of("t1"), List.of(), new
TestExample.Wiring()::depends);
}
public static final class T1 implements Task {
@@ -472,7 +472,7 @@ class BuilderTest {
*/
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));
+ 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));
}
}
""",
@@ -523,7 +523,7 @@ class BuilderTest {
public final class TestExampleBuilder {
public static DagDef build() {
var dag = new DagDef("TestExample");
- return Refs.record(dag, List.of("t"), new
TestExample.Wiring()::depends);
+ return Refs.record(dag, List.of("t"), List.of(), new
TestExample.Wiring()::depends);
}
public static final class T implements Task {
@Override
@@ -593,7 +593,7 @@ class BuilderTest {
public final class TestExampleBuilder {
public static DagDef build() {
var dag = new DagDef("TestExample");
- return Refs.record(dag, List.of("flat", "named"), new
TestExample.Wiring()::depends);
+ return Refs.record(dag, List.of("flat", "named"), List.of(), new
TestExample.Wiring()::depends);
}
public static final class Flat implements Task {
@@ -814,7 +814,7 @@ class BuilderTest {
public final class TestExampleBuilder {
public static DagDef build() {
var dag = new DagDef("foo");
- return Refs.record(dag, List.of(), new
TestExample.Wiring()::depends);
+ return Refs.record(dag, List.of(), List.of(), new
TestExample.Wiring()::depends);
}
}
""",
@@ -850,7 +850,7 @@ class BuilderTest {
public final class Foo {
public static DagDef build() {
var dag = new DagDef("TestExample");
- return Refs.record(dag, List.of(), new
TestExample.Wiring()::depends);
+ return Refs.record(dag, List.of(), List.of(), new
TestExample.Wiring()::depends);
}
}
""",
@@ -898,7 +898,7 @@ class BuilderTest {
public final class TestExampleBuilder {
public static DagDef build() {
var dag = new DagDef("TestExample");
- return Refs.record(dag, List.of("foo"), new
TestExample.Wiring()::depends);
+ return Refs.record(dag, List.of("foo"), List.of(), new
TestExample.Wiring()::depends);
}
public static final class T1 implements Task {
@@ -932,7 +932,7 @@ class BuilderTest {
*/
public interface TestExampleDeps extends Deps {
default TaskRef<Void> t1() {
- return Refs.node(new TaskDef("foo", TestExampleBuilder.T1.class));
+ return Refs.node("", new TaskDef("foo",
TestExampleBuilder.T1.class));
}
}
""",
@@ -1276,6 +1276,474 @@ class BuilderTest {
)
}
+ @Test
+ @DisplayName("nest a wiring view per task group, keyed by the class tree")
+ fun generateBuilderWithTaskGroups() {
+ val compilation =
+ compile(
+ """
+ package org.apache.airflow.example;
+ import org.apache.airflow.sdk.Builder;
+ @Builder.Dag
+ public class TestExample {
+ @Builder.Task public void extract() {}
+
+ @Builder.TaskGroup
+ static class Staging {
+ @Builder.Task public void stage() {}
+
+ @Builder.TaskGroup(id = "checks")
+ static class Checks {
+ @Builder.Task public void nulls() {}
+ }
+ }
+
+ @Builder.Deps
+ static class Wiring implements TestExampleDeps {
+ void depends() {
+ extract().before(staging());
+ staging().stage().before(staging().checks().nulls());
+ }
+ }
+ }
+ """,
+ )
+
+ assertThat(compilation).succeeded()
+ assertThat(compilation)
+ .generatedSourceFile("org.apache.airflow.example.TestExampleBuilder")
+ .contentsAsUtf8String()
+ .contains(
+ "return Refs.record(dag, List.of(\"extract\", \"Staging.stage\",
\"Staging.checks.nulls\"), " +
+ "List.of(\"Staging\", \"Staging.checks\"), new
TestExample.Wiring()::depends);",
+ )
+ val view =
assertThat(compilation).generatedSourceFile("org.apache.airflow.example.TestExampleDeps")
+ view.contentsAsUtf8String().contains("default Staging staging() {")
+ view.contentsAsUtf8String().contains("interface Staging extends
Deps.TaskGroup {")
+ view.contentsAsUtf8String().contains("return \"Staging.checks\";")
+ view.contentsAsUtf8String().contains(
+ "return Refs.node(groupId(), new TaskDef(\"Staging.checks.nulls\", " +
+ "TestExampleBuilder.Staging_Checks_Nulls.class));",
+ )
+ assertThat(compilation)
+ .generatedSourceFile("org.apache.airflow.example.TestExampleBuilder")
+ .contentsAsUtf8String()
+ .contains("public static final class Staging_Checks_Nulls implements
Task {")
+ }
+
+ @Test
+ @DisplayName("scope task method names to their own task group")
+ fun generateBuilderScopesTaskNamesPerGroup() {
+ val compilation =
+ compile(
+ """
+ package org.apache.airflow.example;
+ import org.apache.airflow.sdk.Builder;
+ @Builder.Dag
+ public class TestExample {
+ @Builder.TaskGroup
+ static class First {
+ @Builder.Task public void run() {}
+ }
+
+ @Builder.TaskGroup
+ static class Second {
+ @Builder.Task public void run() {}
+ }
+
+ @Builder.Deps
+ static class Wiring implements TestExampleDeps {
+ void depends() { first().run().before(second().run()); }
+ }
+ }
+ """,
+ )
+
+ assertThat(compilation).succeeded()
+ assertThat(compilation)
+ .generatedSourceFile("org.apache.airflow.example.TestExampleBuilder")
+ .contentsAsUtf8String()
+ .contains("public static final class First_Run implements Task {")
+ }
+
+ @Test
+ @DisplayName("reject a task group ID that is not a plain identifier")
+ fun rejectInvalidTaskGroupId() {
+ val compilation =
+ compile(
+ """
+ package org.apache.airflow.example;
+ import org.apache.airflow.sdk.Builder;
+ @Builder.Dag
+ public class TestExample {
+ @Builder.TaskGroup(id = "staging.checks")
+ static class Staging {
+ @Builder.Task public void t1() {}
+ }
+ }
+ """,
+ )
+ assertThat(compilation).failed()
+ assertThat(compilation).hadErrorContaining(
+ "Task group ID 'staging.checks' must contain only ASCII letters, digits,
underscores, or dashes",
+ )
+ }
+
+ @Test
+ @DisplayName("reject a non-static task group class")
+ fun rejectNonStaticTaskGroupClass() {
+ val compilation =
+ compile(
+ """
+ package org.apache.airflow.example;
+ import org.apache.airflow.sdk.Builder;
+ @Builder.Dag
+ public class TestExample {
+ @Builder.TaskGroup
+ class Staging {
+ @Builder.Task public void t1() {}
+ }
+ }
+ """,
+ )
+ assertThat(compilation).failed()
+ assertThat(compilation).hadErrorContaining(
+ "@Builder.TaskGroup class 'Staging' must be static and non-private",
+ )
+ }
+
+ @Test
+ @DisplayName("reject a task group whose accessor clashes with a task method")
+ fun rejectTaskGroupClashingWithTaskMethod() {
+ val compilation =
+ compile(
+ """
+ package org.apache.airflow.example;
+ import org.apache.airflow.sdk.Builder;
+ @Builder.Dag
+ public class TestExample {
+ @Builder.Task public void staging() {}
+
+ @Builder.TaskGroup
+ static class Staging {
+ @Builder.Task public void t1() {}
+ }
+
+ @Builder.Deps
+ static class Wiring implements TestExampleDeps {
+ void depends() {}
+ }
+ }
+ """,
+ )
+ assertThat(compilation).failed()
+ assertThat(compilation).hadErrorContaining(
+ "Task group class 'Staging' and task method 'staging' would both be
'staging()' on the wiring " +
+ "view; rename one",
+ )
+ }
+
+ @Test
+ @DisplayName("reject an abstract task group class")
+ fun rejectAbstractTaskGroupClass() {
+ val compilation =
+ compile(
+ """
+ package org.apache.airflow.example;
+ import org.apache.airflow.sdk.Builder;
+ @Builder.Dag
+ public class TestExample {
+ @Builder.TaskGroup
+ abstract static class Staging {
+ @Builder.Task public void t1() {}
+ }
+ }
+ """,
+ )
+ assertThat(compilation).failed()
+ assertThat(compilation).hadErrorContaining("@Builder.TaskGroup 'Staging'
must be a concrete class")
+ }
+
+ @Test
+ @DisplayName("reject a task group class with no no-argument constructor")
+ fun rejectTaskGroupClassWithoutNoArgConstructor() {
+ val compilation =
+ compile(
+ """
+ package org.apache.airflow.example;
+ import org.apache.airflow.sdk.Builder;
+ @Builder.Dag
+ public class TestExample {
+ @Builder.TaskGroup
+ static class Staging {
+ Staging(String name) {}
+
+ @Builder.Task public void t1() {}
+ }
+ }
+ """,
+ )
+ assertThat(compilation).failed()
+ assertThat(compilation).hadErrorContaining(
+ "@Builder.TaskGroup class 'Staging' needs a non-private no-argument
constructor",
+ )
+ }
+
+ @Test
+ @DisplayName("reject a task in a group whose name clashes with the group
view")
+ fun rejectTaskClashingWithGroupView() {
+ val compilation =
+ compile(
+ """
+ package org.apache.airflow.example;
+ import org.apache.airflow.sdk.Builder;
+ @Builder.Dag
+ public class TestExample {
+ @Builder.TaskGroup
+ static class Staging {
+ @Builder.Task public void nodes() {}
+ }
+
+ @Builder.Deps
+ static class Wiring implements TestExampleDeps {
+ void depends() {}
+ }
+ }
+ """,
+ )
+ assertThat(compilation).failed()
+ assertThat(compilation).hadErrorContaining(
+ "Task method 'nodes' clashes with a member of the wiring view; rename
the method and keep the " +
+ "task id with @Builder.Task(id = \"nodes\")",
+ )
+ }
+
+ @Test
+ @DisplayName("reject a task in a group named after a member the view
inherits from Flow")
+ fun rejectTaskNamedBeforeInGroup() {
+ val compilation =
+ compile(
+ """
+ package org.apache.airflow.example;
+ import org.apache.airflow.sdk.Builder;
+ @Builder.Dag
+ public class TestExample {
+ @Builder.TaskGroup
+ static class Staging {
+ @Builder.Task public void before() {}
+ }
+
+ @Builder.Deps
+ static class Wiring implements TestExampleDeps {
+ void depends() {}
+ }
+ }
+ """,
+ )
+ assertThat(compilation).failed()
+ assertThat(compilation).hadErrorContaining(
+ "Task method 'before' clashes with a member of the wiring view; rename
the method and keep the " +
+ "task id with @Builder.Task(id = \"before\")",
+ )
+ }
+
+ @Test
+ @DisplayName("reject two task groups in one scope with the same id")
+ fun rejectDuplicateTaskGroupIds() {
+ val compilation =
+ compile(
+ """
+ package org.apache.airflow.example;
+ import org.apache.airflow.sdk.Builder;
+ @Builder.Dag
+ public class TestExample {
+ @Builder.TaskGroup(id = "checks")
+ static class Alpha {
+ @Builder.Task public void t1() {}
+ }
+
+ @Builder.TaskGroup(id = "checks")
+ static class Beta {
+ @Builder.Task public void t2() {}
+ }
+
+ @Builder.Deps
+ static class Wiring implements TestExampleDeps {
+ void depends() {}
+ }
+ }
+ """,
+ )
+ assertThat(compilation).failed()
+ assertThat(compilation).hadErrorContaining("Class TestExample declares
more than one task group 'checks'")
+ }
+
+ @Test
+ @DisplayName("accept a task named after a group view member outside a group")
+ fun acceptTaskNamedAfterGroupViewMemberAtTopLevel() {
+ val compilation =
+ compile(
+ """
+ package org.apache.airflow.example;
+ import org.apache.airflow.sdk.Builder;
+ @Builder.Dag
+ public class TestExample {
+ @Builder.Task public void nodes() {}
+
+ @Builder.Deps
+ static class Wiring implements TestExampleDeps {
+ public void depends() { nodes(); }
+ }
+ }
+ """,
+ )
+ assertThat(compilation).succeeded()
+ }
+
+ @Test
+ @DisplayName("reject two task group classes that share an accessor")
+ fun rejectTaskGroupsSharingAnAccessor() {
+ val compilation =
+ compile(
+ """
+ package org.apache.airflow.example;
+ import org.apache.airflow.sdk.Builder;
+ @Builder.Dag
+ public class TestExample {
+ @Builder.TaskGroup
+ static class Staging {
+ @Builder.Task public void t1() {}
+ }
+
+ @Builder.TaskGroup(id = "lower")
+ static class staging {
+ @Builder.Task public void t2() {}
+ }
+
+ @Builder.Deps
+ static class Wiring implements TestExampleDeps {
+ void depends() {}
+ }
+ }
+ """,
+ )
+ assertThat(compilation).failed()
+ assertThat(compilation).hadErrorContaining(
+ "Task group classes 'Staging' and 'staging' would both be 'staging()' on
the wiring view; rename one",
+ )
+ }
+
+ @Test
+ @DisplayName("accept a task id that carries a dot, as Python's KEY_REGEX
does")
+ fun acceptDottedTaskId() {
+ val compilation =
+ compile(
+ """
+ package org.apache.airflow.example;
+ import org.apache.airflow.sdk.Builder;
+ @Builder.Dag
+ public class TestExample {
+ @Builder.Task(id = "staging.stage") public void stage() {}
+
+ @Builder.Deps
+ static class Wiring implements TestExampleDeps {
+ public void depends() { stage(); }
+ }
+ }
+ """,
+ )
+ assertThat(compilation).succeeded()
+ }
+
+ @Test
+ @DisplayName("reject a task and a task group sharing an id")
+ fun rejectTaskSharingIdWithGroup() {
+ val compilation =
+ compile(
+ """
+ package org.apache.airflow.example;
+ import org.apache.airflow.sdk.Builder;
+ @Builder.Dag
+ public class TestExample {
+ @Builder.Task(id = "Staging") public void staged() {}
+
+ @Builder.TaskGroup
+ static class Staging {
+ @Builder.Task public void t1() {}
+ }
+
+ @Builder.Deps
+ static class Wiring implements TestExampleDeps {
+ void depends() {}
+ }
+ }
+ """,
+ )
+ assertThat(compilation).failed()
+ assertThat(compilation).hadErrorContaining(
+ "Dag has both a task and a task group with ID 'Staging'; rename one",
+ )
+ }
+
+ @Test
+ @DisplayName("reject two task methods whose generated classes would collide")
+ fun rejectCollidingGeneratedTaskClasses() {
+ val compilation =
+ compile(
+ """
+ package org.apache.airflow.example;
+ import org.apache.airflow.sdk.Builder;
+ @Builder.Dag
+ public class TestExample {
+ @Builder.TaskGroup(id = "flat")
+ static class A_B {
+ @Builder.Task public void c() {}
+ }
+
+ @Builder.TaskGroup(id = "outer")
+ static class A {
+ @Builder.TaskGroup(id = "inner")
+ static class B {
+ @Builder.Task public void c() {}
+ }
+ }
+
+ @Builder.Deps
+ static class Wiring implements TestExampleDeps {
+ void depends() {}
+ }
+ }
+ """,
+ )
+ assertThat(compilation).failed()
+ assertThat(compilation).hadErrorContaining(
+ "both generate the task class 'A_B_C'; rename one of them or an
enclosing @Builder.TaskGroup class",
+ )
+ }
+
+ @Test
+ @DisplayName("reject a task group class that is not nested in a dag or
another group")
+ fun rejectMisplacedTaskGroup() {
+ val compilation =
+ compile(
+ """
+ package org.apache.airflow.example;
+ import org.apache.airflow.sdk.Builder;
+ public class TestExample {
+ @Builder.TaskGroup
+ static class Staging {
+ @Builder.Task public void t1() {}
+ }
+ }
+ """,
+ )
+ assertThat(compilation).failed()
+ assertThat(compilation).hadErrorContaining(
+ "@Builder.TaskGroup class 'Staging' must be nested in a @Builder.Dag
class or in another " +
+ "@Builder.TaskGroup class",
+ )
+ }
+
@Test
@DisplayName("generate builder for dag class with varargs task parameter")
fun generateBuilderForDagClassWithVarArgsTaskParameter() {
@@ -1338,7 +1806,7 @@ class BuilderTest {
)
assertThat(compilation).failed()
assertThat(compilation).hadErrorContaining(
- "Dag class TestExample overloads task method 'extract'; a method's name
is the name of its " +
+ "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\")",
)
@@ -1501,7 +1969,7 @@ class BuilderTest {
public final class TestExampleBuilder {
public static DagDef build() {
var dag = new DagDef("TestExample");
- return Refs.record(dag, List.of("score"), new
TestExample.Wiring()::depends);
+ return Refs.record(dag, List.of("score"), List.of(), 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 c4178f5148d..9ef581634bd 100644
--- a/java-sdk/sdk/build.gradle.kts
+++ b/java-sdk/sdk/build.gradle.kts
@@ -494,9 +494,9 @@ abstract class GenerateDagDslTask : DefaultTask() {
(excludedTaskKeys - excludedSeen).takeIf { it.isNotEmpty() }?.let {
throw GradleException("Excluded task keys match no eligible schema
property; remove or fix: $it")
}
- // "id"/"to" name the annotations' structural attributes, so a schema
- // key camel-casing to either would silently shadow them.
- (dagFields + taskFields).firstOrNull { it.attribute == "id" ||
it.attribute == "to" }?.let {
+ // "id" and "to" name the annotations' structural attributes, so a
schema
+ // key camel-casing to one would silently shadow it.
+ (dagFields + taskFields).firstOrNull { it.attribute in setOf("id",
"to") }?.let {
throw GradleException("Schema key '${it.key}' collides with a
structural annotation attribute")
}
@@ -636,6 +636,43 @@ abstract class GenerateDagDslTask : DefaultTask() {
| @Target(AnnotationTarget.CLASS)
| @MustBeDocumented
| annotation class Deps
+ |
+ | /**
+ | * Marks a nested class that groups the tasks declared inside
it, as
+ | * Python's `TaskGroup` does.
+ | *
+ | * Declare it as a `static` nested class of the [Dag] class, or
of
+ | * another [TaskGroup] class to nest one group in another.
Everything it
+ | * declares carries its ID as a prefix, so `stage` in `Staging`
is the
+ | * task `Staging.stage`:
+ | *
+ | * ```java
+ | * @Builder.TaskGroup
+ | * static class Staging {
+ | * @Builder.Task
+ | * public long stage(long rows) { ... }
+ | *
+ | * @Builder.TaskGroup(id = "checks")
+ | * static class Checks {
+ | * @Builder.Task
+ | * public void nulls(long staged) { ... }
+ | * }
+ | * }
+ | * ```
+ | *
+ | * The wiring class reaches them through the generated view,
where the
+ | * group is both a namespace and a point in the flow:
+ | * `staging().checks().nulls(staged)` and
`extract().before(staging())`.
+ | *
+ | * @param id Group ID within its enclosing group. Empty derives
it from
+ | * the annotated class's name. Must contain only ASCII
letters,
+ | * digits, underscores, or dashes.
+ | */
+ | @Target(AnnotationTarget.CLASS)
+ | @MustBeDocumented
+ | annotation class TaskGroup(
+ | val id: String = "",
+ | )
|}
|
""".trimMargin(),
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 605b2b46ddf..bd055ab2c25 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
@@ -196,6 +196,7 @@ class Bundle(
// recursive and could blow up with deep dependency chains. I kept the
recursive implementation
// for readability since the scenario is unlikely; feel free to rewrite if it
blows up for you.
private fun checkNoCycle(dag: DagDef) {
+ val expansion = dag.expandGroupEdges()
val visiting = mutableSetOf<String>()
val done = mutableSetOf<String>()
@@ -204,10 +205,11 @@ private fun checkNoCycle(dag: DagDef) {
require(visiting.add(def.id)) {
"Task dependencies in Dag '${dag.id}' contain a cycle involving task
'${def.id}'"
}
- def.upstreams.forEach(::visit)
+ expansion.upstreamsOf(def).forEach { dag.tasks[it]?.let(::visit) }
visiting -= def.id
done += def.id
}
+
dag.tasks.values.forEach(::visit)
}
diff --git a/java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/DagDef.kt
b/java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/DagDef.kt
index c579ddd517b..d5ba757beab 100644
--- a/java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/DagDef.kt
+++ b/java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/DagDef.kt
@@ -19,6 +19,7 @@
package org.apache.airflow.sdk
+import org.apache.airflow.sdk.internal.GROUP_ID
import org.apache.airflow.sdk.internal.SchemaFields
import org.apache.airflow.sdk.internal.checkConfigValue
import org.apache.airflow.sdk.internal.validateTaskInput
@@ -50,6 +51,12 @@ class DagDef(
internal val tasks = linkedMapOf<String, TaskDef>()
internal val dagConfig = linkedMapOf<String, Any>()
+ /** Task groups keyed by their full ID, parents before the groups nested in
them. */
+ internal val groups = linkedMapOf<String, TaskGroupRef>()
+
+ /** Edges with a task group at either end, in the order drawn. */
+ internal val groupEdges = linkedSetOf<Pair<Endpoint, Endpoint>>()
+
/**
* Sets one Dag-level configuration value.
*
@@ -133,14 +140,152 @@ class DagDef(
task.owner?.let { owner ->
throw IllegalArgumentException("Task '${task.id}' already belongs to Dag
'${owner.id}'")
}
+ require(task.id !in groups) { "Dag '$id' already has a task group with ID:
${task.id}" }
require(tasks.putIfAbsent(task.id, task) == null) {
"Tasks in Dag have duplicate ID: ${task.id}"
}
task.owner = this
return this
}
+
+ /**
+ * Declares a task group of this Dag.
+ *
+ * ```java
+ * var staging = dag.taskGroup("staging");
+ * var stage = staging.task("stage", Stage.class); // task "staging.stage"
+ * extract.before(staging);
+ * ```
+ *
+ * @param id Group ID. Must contain only ASCII letters, digits, underscores,
+ * or dashes, and differ from every task and group ID in this Dag.
+ * @return The group, to declare tasks in and to wire edges with.
+ * @throws IllegalArgumentException if [id] is not a valid group ID, or the
+ * Dag already has a task or task group with that ID.
+ */
+ fun taskGroup(id: String): TaskGroupRef = addGroup(null, id)
+
+ internal fun addGroup(
+ parent: TaskGroupRef?,
+ localId: String,
+ ): TaskGroupRef {
+ require(GROUP_ID.matches(localId)) {
+ "Task group ID '$localId' must contain only ASCII letters, digits,
underscores, or dashes"
+ }
+ val groupId = parent?.qualify(localId) ?: localId
+ require(groupId !in tasks && groupId !in groups) {
+ "Dag '$id' already has a task or task group with ID: $groupId"
+ }
+ return TaskGroupRef(this, groupId, parent).also {
+ groups[groupId] = it
+ parent?.children?.add(it)
+ }
+ }
+
+ /**
+ * What this Dag's task-group edges mean in terms of tasks.
+ *
+ * A group upstream stands for its leaves and a group downstream for its
+ * roots. Edges are read in the order they were drawn, each seeing the ones
+ * before it, which is how Python resolves a group's endpoints at every
+ * `>>`. The result is computed on demand and stored nowhere, so a task
+ * added to a group after the Dag was registered still counts.
+ */
+ internal fun expandGroupEdges(): GroupExpansion {
+ val upstreams = mutableMapOf<String, MutableSet<String>>()
+ val edges = mutableMapOf<String, MutableGroupEdges>()
+
+ fun edgesOf(groupId: String) = edges.getOrPut(groupId) {
MutableGroupEdges() }
+
+ fun upstreamIds(def: TaskDef): Set<String> =
def.upstreams.mapTo(linkedSetOf()) { it.id } + upstreams[def.id].orEmpty()
+
+ fun roots(group: TaskGroupRef): List<TaskDef> {
+ val members = group.nodes()
+ val ids = members.mapTo(mutableSetOf()) { it.id }
+ return members.filter { task -> upstreamIds(task).none { it in ids } }
+ }
+
+ fun leaves(group: TaskGroupRef): List<TaskDef> {
+ val members = group.nodes()
+ val ids = members.mapTo(mutableSetOf()) { it.id }
+ val fedInside = members.flatMapTo(mutableSetOf()) { task ->
upstreamIds(task).filter { it in ids } }
+ return members.filter { it.id !in fedInside }
+ }
+
+ // Python's find_leaves: the group's own leaves, else whatever already runs
+ // before it, else the group it is nested in.
+ fun leavesOf(endpoint: Endpoint): List<TaskDef> =
+ when (endpoint) {
+ is TaskDef -> listOf(endpoint)
+ is TaskGroupRef -> {
+ var group: TaskGroupRef? = endpoint
+ var found: List<TaskDef> = emptyList()
+ while (group != null && found.isEmpty()) {
+ found = leaves(group).ifEmpty {
edgesOf(group.id).upstreamTaskIds.map { tasks.getValue(it) } }
+ group = group.parent
+ }
+ found
+ }
+ }
+
+ fun rootsOf(endpoint: Endpoint): List<TaskDef> =
+ when (endpoint) {
+ is TaskDef -> listOf(endpoint)
+ is TaskGroupRef -> roots(endpoint)
+ }
+
+ for ((upstream, downstream) in groupEdges) {
+ val from = leavesOf(upstream).map { it.id }
+ rootsOf(downstream).forEach { task -> upstreams.getOrPut(task.id) {
linkedSetOf() } += from }
+ if (downstream is TaskGroupRef) {
+ edgesOf(downstream.id).upstreamTaskIds += from
+ if (upstream is TaskGroupRef) edgesOf(downstream.id).upstreamGroupIds
+= upstream.id
+ }
+ // When both ends are groups, the upstream records the downstream group
+ // only, not its tasks, which is how Python leaves it.
+ when {
+ upstream is TaskGroupRef && downstream is TaskGroupRef ->
+ edgesOf(upstream.id).downstreamGroupIds += downstream.id
+ upstream is TaskGroupRef && downstream is TaskDef ->
+ edgesOf(upstream.id).downstreamTaskIds += downstream.id
+ }
+ }
+ return GroupExpansion(upstreams, edges)
+ }
+}
+
+/**
+ * The task edges a Dag's task-group edges stand for, and the edges each group
+ * records for itself, as [DagDef.expandGroupEdges] worked them out.
+ */
+internal class GroupExpansion(
+ private val upstreams: Map<String, Set<String>>,
+ private val edges: Map<String, GroupEdges>,
+) {
+ /** Every task [def] runs after: the edges it carries, plus the ones a group
edge implies. */
+ fun upstreamsOf(def: TaskDef): Set<String> =
def.upstreams.mapTo(linkedSetOf()) { it.id } + upstreams[def.id].orEmpty()
+
+ /** The edges the group with full ID [groupId] records for itself. */
+ fun edgesOf(groupId: String): GroupEdges = edges[groupId] ?:
EMPTY_GROUP_EDGES
+}
+
+/** One task group's own edges, as Python's `TaskGroup` records them. */
+internal open class GroupEdges {
+ open val upstreamGroupIds: Set<String> = emptySet()
+ open val downstreamGroupIds: Set<String> = emptySet()
+ open val upstreamTaskIds: Set<String> = emptySet()
+ open val downstreamTaskIds: Set<String> = emptySet()
+}
+
+private class MutableGroupEdges : GroupEdges() {
+ override val upstreamGroupIds = linkedSetOf<String>()
+ override val downstreamGroupIds = linkedSetOf<String>()
+ override val upstreamTaskIds = linkedSetOf<String>()
+ override val downstreamTaskIds = linkedSetOf<String>()
}
+private val EMPTY_GROUP_EDGES = GroupEdges()
+
/**
* One task definition: its ID, the class that implements it, its upstream
* dependencies, and its task-level configuration.
@@ -164,7 +309,7 @@ class DagDef(
class TaskDef(
val id: String,
val definition: Class<out Task>,
-) {
+) : Endpoint {
init {
validateTaskInput(definition)
}
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 e70df9f8534..0dab9c978f7 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,6 +19,8 @@
package org.apache.airflow.sdk
+import org.apache.airflow.sdk.internal.Refs
+
/**
* Vocabulary for declaring a Dag's task graph in Java, and the base of every
* generated `<Dag>Deps` wiring view.
@@ -28,7 +30,7 @@ package org.apache.airflow.sdk
*/
interface Deps {
/**
- * A point in the task graph: one task, or a set of them.
+ * A point in the task graph: one task, one task group, or a set of them.
*
* [Flow] declares a dependency where nothing flows but the ordering. An
* edge that carries a value is declared by passing the upstream's handle
@@ -38,6 +40,13 @@ interface Deps {
/** The tasks at this point in the flow. */
fun nodes(): List<TaskDef>
+ /**
+ * The ends an edge drawn here attaches to. A task stands for itself, so
+ * the default is [nodes]; a task group stands for the group rather than
+ * for the tasks it holds today.
+ */
+ fun endpoints(): List<Endpoint> = nodes()
+
/**
* Runs the tasks here before each of [next], carrying no value.
*
@@ -53,8 +62,7 @@ interface Deps {
* @return This point in the flow.
*/
fun before(vararg next: Flow): Flow {
- val upstreams = nodes()
- next.flatMap { it.nodes() }.forEach { downstream -> upstreams.forEach {
downstream.dependsOn(it) } }
+ next.forEach { link(this, it) }
return this
}
@@ -69,8 +77,7 @@ interface Deps {
* @return This point in the flow.
*/
fun after(vararg previous: Flow): Flow {
- val downstreams = nodes()
- previous.flatMap { it.nodes() }.forEach { upstream ->
downstreams.forEach { it.dependsOn(upstream) } }
+ previous.forEach { link(it, this) }
return this
}
@@ -84,10 +91,35 @@ interface Deps {
* ```
*/
@JvmStatic
- fun of(vararg flows: Flow): Flow = FlowSet(flows.flatMap { it.nodes() })
+ fun of(vararg flows: Flow): Flow = FlowSet(flows.toList())
}
}
+ /**
+ * One task group of the Dag being wired: a point in the flow, and the
+ * namespace of the tasks and groups declared inside it.
+ *
+ * The generated wiring view nests one of these per [Builder.TaskGroup]
+ * class, so a group is reached by calling it and its contents by calling on
+ * through:
+ *
+ * ```java
+ * staging().stage(rows); // the task "staging.stage"
+ * staging().checks().nulls(id); // the task "staging.checks.nulls"
+ * extract().before(staging()); // the whole group runs after extract
+ * ```
+ */
+ interface TaskGroup : Flow {
+ /** Full ID of this group, as the Dag registered it. */
+ fun groupId(): String
+
+ override fun nodes(): List<TaskDef> = Refs.group(groupId()).nodes()
+
+ // The group itself, not its tasks, so an edge drawn before its tasks
+ // exist still reaches them.
+ override fun endpoints(): List<Endpoint> = listOf(Refs.group(groupId()))
+ }
+
/**
* Wraps an inline constant as a task argument, as in
* `transform(extract(), lit(0.9))`. It is passed to the task as a constant
@@ -98,9 +130,54 @@ interface Deps {
fun <T> lit(value: T?): Arg<T> = Arg.lit(value)
}
-/** Several tasks as one point in the flow, which no single [TaskRef] can
represent. */
+/** Several tasks or groups as one point in the flow, which no single handle
can represent. */
internal class FlowSet(
- private val nodes: List<TaskDef>,
+ internal val flows: List<Deps.Flow>,
) : Deps.Flow {
- override fun nodes(): List<TaskDef> = nodes
+ override fun nodes(): List<TaskDef> = flows.flatMap { it.nodes() }
+
+ override fun endpoints(): List<Endpoint> = flows.flatMap { it.endpoints() }
}
+
+/**
+ * Draws an ordering edge from each endpoint of [upstream] to each of
+ * [downstream]. An edge between two tasks is recorded on the downstream task;
+ * one with a task group at either end is recorded on the group's Dag, and
+ * means whatever tasks the group holds when the Dag is registered.
+ */
+private fun link(
+ upstream: Deps.Flow,
+ downstream: Deps.Flow,
+) {
+ for (up in upstream.endpoints()) {
+ for (down in downstream.endpoints()) {
+ if (up is TaskDef && down is TaskDef) {
+ down.dependsOn(up)
+ } else {
+ val upDag = up.owningDag
+ val downDag = down.owningDag
+ require(upDag == null || downDag == null || upDag === downDag) {
+ "Cannot order ${up.label} of Dag '${upDag?.id}' before ${down.label}
of " +
+ "Dag '${downDag?.id}'; an edge stays inside one Dag"
+ }
+ (upDag ?: downDag)?.let { it.groupEdges += up to down }
+ }
+ }
+ }
+}
+
+/** The Dag an endpoint belongs to, null for a task not registered with one
yet. */
+internal val Endpoint.owningDag: DagDef?
+ get() =
+ when (this) {
+ is TaskDef -> owner
+ is TaskGroupRef -> dag
+ }
+
+/** How an endpoint is named in a diagnostic. */
+internal val Endpoint.label: String
+ get() =
+ when (this) {
+ is TaskDef -> "task '$id'"
+ is TaskGroupRef -> "task group '$id'"
+ }
diff --git a/java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/Endpoint.kt
b/java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/Endpoint.kt
new file mode 100644
index 00000000000..8831e4c091d
--- /dev/null
+++ b/java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/Endpoint.kt
@@ -0,0 +1,29 @@
+/*
+ * 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
+
+/**
+ * One end of an ordering edge: a single task, or a whole task group.
+ *
+ * [Deps.Flow.before] and [Deps.Flow.after] draw edges between endpoints. A
+ * task stands for itself; a task group stands for the group, so an edge drawn
+ * to it reaches whatever it holds when the Dag is registered.
+ */
+sealed interface Endpoint
diff --git
a/java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/TaskGroupRef.kt
b/java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/TaskGroupRef.kt
new file mode 100644
index 00000000000..dc6448a5eb0
--- /dev/null
+++ b/java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/TaskGroupRef.kt
@@ -0,0 +1,111 @@
+/*
+ * 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
+
+/**
+ * A group of tasks in a Dag, shown in the Airflow UI as one node that expands:
+ * Python's `TaskGroup`.
+ *
+ * Everything declared in a group carries the group's ID as a prefix, so task
+ * `stage` in group `staging` is the task `staging.stage`. A group can stand
+ * at either end of an edge, so a whole group can be ordered against a task or
+ * another group:
+ *
+ * ```java
+ * var staging = dag.taskGroup("staging");
+ * staging.task("stage", Stage.class);
+ * extract.before(staging); // every task staging starts with waits for extract
+ * ```
+ *
+ * As an upstream, a group stands for its leaves, the tasks nothing else in the
+ * group runs after; as a downstream, for its roots, the tasks that run after
+ * nothing else in the group. What the group holds is read once, rather than at
+ * each edge as Python reads it, so a group's edges can be drawn before its
+ * tasks are declared. The edges still resolve in the order they were drawn.
+ *
+ * @property id Group ID, including any enclosing group's prefix.
+ */
+class TaskGroupRef internal constructor(
+ internal val dag: DagDef,
+ val id: String,
+ internal val parent: TaskGroupRef? = null,
+) : Deps.Flow,
+ Endpoint {
+ /** IDs of the tasks declared directly in this group, in declaration order.
*/
+ internal val taskIds = mutableListOf<String>()
+
+ /** Groups nested directly in this group, in declaration order. */
+ internal val children = mutableListOf<TaskGroupRef>()
+
+ /**
+ * Creates a task in this group, registers it with the Dag, and hands back
+ * its handle.
+ *
+ * @param id Task ID within this group; the task's ID is `<group ID>.<id>`.
+ * @param definition Class that implements [Task]. Must have a public no-arg
+ * constructor.
+ * @return The handle representing this task.
+ * @throws IllegalArgumentException if the Dag already has a task or task
+ * group with the resulting ID.
+ */
+ fun <T> task(
+ id: String,
+ definition: Class<out Task>,
+ ): TaskRef<T> {
+ val def = TaskDef(qualify(id), definition)
+ adopt(def)
+ return TaskRef(def)
+ }
+
+ /**
+ * Nests a task group inside this one.
+ *
+ * @param id Group ID within this group; the nested group's ID is
+ * `<group ID>.<id>`. Must contain only ASCII letters, digits,
+ * underscores, or dashes.
+ * @return The nested group.
+ * @throws IllegalArgumentException if [id] is not a valid group ID, or the
+ * Dag already has a task or task group with the resulting ID.
+ */
+ fun taskGroup(id: String): TaskGroupRef = dag.addGroup(this, id)
+
+ /** Every task in this group, including those in nested groups. */
+ override fun nodes(): List<TaskDef> = taskIds.map { dag.tasks.getValue(it) }
+ children.flatMap { it.nodes() }
+
+ override fun endpoints(): List<Endpoint> = listOf(this)
+
+ override fun before(vararg next: Deps.Flow): TaskGroupRef {
+ super.before(*next)
+ return this
+ }
+
+ override fun after(vararg previous: Deps.Flow): TaskGroupRef {
+ super.after(*previous)
+ return this
+ }
+
+ internal fun qualify(localId: String): String = "$id.$localId"
+
+ /** Registers [def], whose ID already carries this group's prefix, as a task
of this group. */
+ internal fun adopt(def: TaskDef) {
+ dag.addTask(def)
+ taskIds += def.id
+ }
+}
diff --git
a/java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/internal/Ids.kt
b/java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/internal/Ids.kt
new file mode 100644
index 00000000000..b64551e9b01
--- /dev/null
+++ b/java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/internal/Ids.kt
@@ -0,0 +1,29 @@
+/*
+ * Licensed to the Apache Software Foundation (ASF) under one
+ * or more contributor license agreements. See the NOTICE file
+ * distributed with this work for additional information
+ * regarding copyright ownership. The ASF licenses this file
+ * to you under the Apache License, Version 2.0 (the
+ * "License"); you may not use this file except in compliance
+ * with the License. You may obtain a copy of the License at
+ *
+ * http://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing,
+ * software distributed under the License is distributed on an
+ * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
+ * KIND, either express or implied. See the License for the
+ * specific language governing permissions and limitations
+ * under the License.
+ */
+
+package org.apache.airflow.sdk.internal
+
+/**
+ * @suppress
+ *
+ * What a task group ID may contain, mirroring Python's `GROUP_KEY_REGEX`.
+ * Public so the annotation processor, which is a separate module, can check it
+ * against the same pattern the SDK enforces; not user-facing API.
+ */
+val GROUP_ID: Regex = Regex("[A-Za-z0-9_-]+")
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
index 39021f03b4e..387e6fec365 100644
--- 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
@@ -22,6 +22,7 @@ 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.TaskGroupRef
import org.apache.airflow.sdk.TaskRef
/**
@@ -49,6 +50,10 @@ object Refs {
* Runs one `depends()` call with [dag] in scope, then returns the Dag the
* wiring built.
*
+ * @param groupIds Full ID of every task group, parents before the groups
+ * nested in them. All are created before `depends()` runs, so the wiring
+ * can order a group before calling any of its tasks, and a group holding
+ * no tasks still exists.
* @throws IllegalArgumentException if the wiring left a declared task
* unregistered.
*/
@@ -56,9 +61,11 @@ object Refs {
fun record(
dag: DagDef,
taskIds: List<String>,
+ groupIds: List<String>,
depends: Runnable,
): DagDef {
check(recording.get() == null) { "Dag wiring is already being recorded on
this thread" }
+ groupIds.forEach { createGroup(dag, it) }
recording.set(Recording(dag))
try {
depends.run()
@@ -76,16 +83,23 @@ object Refs {
/**
* Records a task that takes no data arguments.
*
+ * @param groupId Full ID of the task group holding it, empty when it sits in
+ * none. The generated wiring view knows which it is.
* @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)
+ fun <T> node(
+ groupId: String,
+ def: TaskDef,
+ ): TaskRef<T> = call(groupId, def)
/**
* Records a task and the data edge for every [TaskRef] among [args]; a
* literal argument records a baked value and no edge.
*
+ * @param groupId Full ID of the task group holding it, empty when it sits in
+ * none.
* @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`
@@ -94,6 +108,7 @@ object Refs {
@JvmStatic
@Suppress("UNCHECKED_CAST", "SpreadOperator")
fun <T> call(
+ groupId: String,
def: TaskDef,
vararg args: Arg<*>?,
): TaskRef<T> {
@@ -116,7 +131,44 @@ object Refs {
}
inputs.filterIsInstance<TaskRef<*>>().forEach { def.dependsOn(it.def) }
def.inputs += inputs
- active.dag.addTask(def)
+ if (groupId.isEmpty()) {
+ active.dag.addTask(def)
+ } else {
+ active.dag.groups
+ .getValue(groupId)
+ .adopt(def)
+ }
return TaskRef<T>(def).also { active.byTaskId[def.id] = it }
}
+
+ /**
+ * The task group with this full ID in the Dag being recorded. Public so the
+ * generated wiring view can resolve the group it stands for.
+ */
+ @JvmStatic
+ fun group(id: String): TaskGroupRef {
+ val active =
+ checkNotNull(recording.get()) {
+ "Task group '$id' was looked up outside a @Builder.Deps class"
+ }
+ return requireNotNull(active.dag.groups[id]) {
+ "Dag '${active.dag.id}' has no task group '$id'"
+ }
+ }
+
+ /**
+ * Creates the group with full ID [id] in [dag]. Its enclosing group already
+ * exists, because [record] takes parents before the groups nested in them.
+ */
+ private fun createGroup(
+ dag: DagDef,
+ id: String,
+ ) {
+ val parentId = id.substringBeforeLast('.', "")
+ if (parentId.isEmpty()) {
+ dag.taskGroup(id)
+ } else {
+ dag.groups.getValue(parentId).taskGroup(id.substringAfterLast('.'))
+ }
+ }
}
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 ac9e8507c12..50b0a8ef3ae 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
@@ -129,6 +129,6 @@ internal class NoopTask : Task {
internal fun contextWiredWith(inputs: List<Arg<*>>): Context {
val dag = DagDef("d")
val def = TaskDef("t", NoopTask::class.java)
- Refs.record(dag, listOf("t")) { Refs.call<Unit>(def, *inputs.toTypedArray())
}
+ Refs.record(dag, listOf("t"), emptyList()) { 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/TaskGroupTest.kt
b/java-sdk/sdk/src/test/kotlin/org/apache/airflow/sdk/TaskGroupTest.kt
new file mode 100644
index 00000000000..22c5972f32e
--- /dev/null
+++ b/java-sdk/sdk/src/test/kotlin/org/apache/airflow/sdk/TaskGroupTest.kt
@@ -0,0 +1,286 @@
+/*
+ * Licensed to the Apache Software Foundation (ASF) under one
+ * or more contributor license agreements. See the NOTICE file
+ * distributed with this work for additional information
+ * regarding copyright ownership. The ASF licenses this file
+ * to you under the Apache License, Version 2.0 (the
+ * "License"); you may not use this file except in compliance
+ * with the License. You may obtain a copy of the License at
+ *
+ * http://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing,
+ * software distributed under the License is distributed on an
+ * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
+ * KIND, either express or implied. See the License for the
+ * specific language governing permissions and limitations
+ * under the License.
+ */
+
+package org.apache.airflow.sdk
+
+import org.junit.jupiter.api.Assertions.assertEquals
+import org.junit.jupiter.api.Assertions.assertThrows
+import org.junit.jupiter.api.DisplayName
+import org.junit.jupiter.api.Test
+
+internal class TaskGroupTest {
+ private fun upstreamIds(
+ dag: DagDef,
+ taskId: String,
+ ) = dag.expandGroupEdges().upstreamsOf(dag.tasks.getValue(taskId))
+
+ @Test
+ @DisplayName("Should prefix the IDs of tasks and groups declared in a group")
+ fun shouldPrefixIdsDeclaredInGroup() {
+ val dag = DagDef("d")
+ val staging = dag.taskGroup("staging")
+ val stage = staging.task<Unit>("stage", NoopTask::class.java)
+ val checks = staging.taskGroup("checks")
+ checks.task<Unit>("nulls", NoopTask::class.java)
+
+ assertEquals("staging.stage", stage.def.id)
+ assertEquals("staging.checks", checks.id)
+ assertEquals(listOf("staging.stage", "staging.checks.nulls"),
dag.tasks.keys.toList())
+ assertEquals(listOf("staging", "staging.checks"), dag.groups.keys.toList())
+ }
+
+ @Test
+ @DisplayName("Should reject a group ID that is not a plain identifier")
+ fun shouldRejectInvalidGroupId() {
+ val error = assertThrows(IllegalArgumentException::class.java) {
DagDef("d").taskGroup("a.b") }
+
+ assertEquals(
+ "Task group ID 'a.b' must contain only ASCII letters, digits,
underscores, or dashes",
+ error.message,
+ )
+ }
+
+ @Test
+ @DisplayName("Should reject a group ID that a task already uses")
+ fun shouldRejectGroupIdTakenByTask() {
+ val dag = DagDef("d")
+ dag.task<Unit>("staging", NoopTask::class.java)
+
+ val error = assertThrows(IllegalArgumentException::class.java) {
dag.taskGroup("staging") }
+
+ assertEquals("Dag 'd' already has a task or task group with ID: staging",
error.message)
+ }
+
+ @Test
+ @DisplayName("Should reject a task ID that a group already uses")
+ fun shouldRejectTaskIdTakenByGroup() {
+ val dag = DagDef("d")
+ dag.taskGroup("staging")
+
+ val error =
+ assertThrows(IllegalArgumentException::class.java) {
dag.task<Unit>("staging", NoopTask::class.java) }
+
+ assertEquals("Dag 'd' already has a task group with ID: staging",
error.message)
+ }
+
+ @Test
+ @DisplayName("Should wire a group upstream from its leaves and downstream to
its roots on registration")
+ fun shouldExpandGroupEdgesOntoRootsAndLeaves() {
+ val dag = DagDef("d")
+ val extract = dag.task<Unit>("extract", NoopTask::class.java)
+ val staging = dag.taskGroup("staging")
+ val stage = staging.task<Unit>("stage", NoopTask::class.java)
+ staging.taskGroup("checks").task<Unit>("nulls",
NoopTask::class.java).after(stage)
+ val publish = dag.taskGroup("publish")
+ publish.task<Unit>("push", NoopTask::class.java)
+ val load = dag.task<Unit>("load", NoopTask::class.java)
+ extract.before(staging)
+ staging.before(publish)
+ publish.before(load)
+
+ Bundle().register(dag)
+
+ assertEquals(setOf("extract"), upstreamIds(dag, "staging.stage"))
+ assertEquals(setOf("staging.stage"), upstreamIds(dag,
"staging.checks.nulls"))
+ assertEquals(setOf("staging.checks.nulls"), upstreamIds(dag,
"publish.push"))
+ assertEquals(setOf("publish.push"), upstreamIds(dag, "load"))
+ }
+
+ @Test
+ @DisplayName("Should record each group's own edges the way Python's
TaskGroup does")
+ fun shouldRecordGroupEdgesOnGroups() {
+ val dag = DagDef("d")
+ val extract = dag.task<Unit>("extract", NoopTask::class.java)
+ val staging = dag.taskGroup("staging")
+ staging.task<Unit>("stage", NoopTask::class.java)
+ val publish = dag.taskGroup("publish")
+ publish.task<Unit>("push", NoopTask::class.java)
+ val load = dag.task<Unit>("load", NoopTask::class.java)
+ extract.before(staging)
+ staging.before(publish)
+ publish.before(load)
+
+ Bundle().register(dag)
+
+ val edges = dag.expandGroupEdges()
+ assertEquals(setOf("extract"), edges.edgesOf(staging.id).upstreamTaskIds)
+ assertEquals(setOf("publish"),
edges.edgesOf(staging.id).downstreamGroupIds)
+ assertEquals(emptySet<String>(),
edges.edgesOf(staging.id).downstreamTaskIds)
+ assertEquals(setOf("staging"), edges.edgesOf(publish.id).upstreamGroupIds)
+ assertEquals(setOf("staging.stage"),
edges.edgesOf(publish.id).upstreamTaskIds)
+ assertEquals(setOf("load"), edges.edgesOf(publish.id).downstreamTaskIds)
+ }
+
+ @Test
+ @DisplayName("Should step over a group with no tasks to the tasks beyond it")
+ fun shouldStepOverEmptyGroup() {
+ val dag = DagDef("d")
+ val extract = dag.task<Unit>("extract", NoopTask::class.java)
+ val empty = dag.taskGroup("empty")
+ val load = dag.task<Unit>("load", NoopTask::class.java)
+ extract.before(empty)
+ empty.before(load)
+
+ Bundle().register(dag)
+
+ assertEquals(setOf("extract"), upstreamIds(dag, "load"))
+ }
+
+ @Test
+ @DisplayName("Should step over an empty nested group to the tasks of the
group holding it")
+ fun shouldStepOverEmptyNestedGroup() {
+ val dag = DagDef("d")
+ val outer = dag.taskGroup("outer")
+ outer.task<Unit>("t", NoopTask::class.java)
+ val inner = outer.taskGroup("inner")
+ val load = dag.task<Unit>("load", NoopTask::class.java)
+ inner.before(load)
+
+ Bundle().register(dag)
+
+ assertEquals(setOf("outer.t"), upstreamIds(dag, "load"))
+ }
+
+ @Test
+ @DisplayName("Should prefer what already runs before an empty group over the
group holding it")
+ fun shouldPreferEarlierEdgeOverParentForEmptyGroup() {
+ val dag = DagDef("d")
+ val extract = dag.task<Unit>("extract", NoopTask::class.java)
+ val outer = dag.taskGroup("outer")
+ outer.task<Unit>("t", NoopTask::class.java)
+ val inner = outer.taskGroup("inner")
+ val load = dag.task<Unit>("load", NoopTask::class.java)
+ extract.before(inner)
+ inner.before(load)
+
+ Bundle().register(dag)
+
+ assertEquals(setOf("extract"), upstreamIds(dag, "load"))
+ }
+
+ @Test
+ @DisplayName("Should resolve a group's endpoints in the order the edges were
drawn")
+ fun shouldResolveGroupEndpointsInDrawingOrder() {
+ val dag = DagDef("d")
+ val extract = dag.task<Unit>("extract", NoopTask::class.java)
+ val outer = dag.taskGroup("outer")
+ val t = outer.task<Unit>("t", NoopTask::class.java)
+ val inner = outer.taskGroup("inner")
+ inner.task<Unit>("i", NoopTask::class.java)
+ extract.before(outer)
+ inner.before(t)
+
+ Bundle().register(dag)
+
+ // outer had both tasks as roots when the first edge was drawn, so extract
reaches both.
+ assertEquals(setOf("extract", "outer.inner.i"), upstreamIds(dag,
"outer.t"))
+ assertEquals(setOf("extract"), upstreamIds(dag, "outer.inner.i"))
+ }
+
+ @Test
+ @DisplayName("Should resolve a group's endpoints against the edges drawn
before it")
+ fun shouldResolveGroupEndpointsAgainstEarlierEdges() {
+ val dag = DagDef("d")
+ val extract = dag.task<Unit>("extract", NoopTask::class.java)
+ val outer = dag.taskGroup("outer")
+ val t = outer.task<Unit>("t", NoopTask::class.java)
+ val inner = outer.taskGroup("inner")
+ inner.task<Unit>("i", NoopTask::class.java)
+ inner.before(t)
+ extract.before(outer)
+
+ Bundle().register(dag)
+
+ // outer's only root once inner runs before t, so extract reaches nothing
else.
+ assertEquals(setOf("outer.inner.i"), upstreamIds(dag, "outer.t"))
+ assertEquals(setOf("extract"), upstreamIds(dag, "outer.inner.i"))
+ }
+
+ @Test
+ @DisplayName("Should reject a group edge that crosses two Dags")
+ fun shouldRejectCrossDagGroupEdge() {
+ val staging = DagDef("a").taskGroup("staging")
+ val publish = DagDef("b").taskGroup("publish")
+
+ val error = assertThrows(IllegalArgumentException::class.java) {
staging.before(publish) }
+
+ assertEquals(
+ "Cannot order task group 'staging' of Dag 'a' before task group
'publish' of Dag 'b'; " +
+ "an edge stays inside one Dag",
+ error.message,
+ )
+ }
+
+ @Test
+ @DisplayName("Should wire a task added to a group after the Dag was
registered")
+ fun shouldWireTaskAddedAfterRegistration() {
+ val dag = DagDef("d")
+ val extract = dag.task<Unit>("extract", NoopTask::class.java)
+ val staging = dag.taskGroup("staging")
+ staging.task<Unit>("stage", NoopTask::class.java)
+ extract.before(staging)
+ Bundle().register(dag)
+
+ staging.task<Unit>("late", NoopTask::class.java)
+
+ assertEquals(setOf("extract"), upstreamIds(dag, "staging.late"))
+ }
+
+ @Test
+ @DisplayName("Should expand an edge drawn before the group's tasks were
declared")
+ fun shouldExpandEdgeDrawnBeforeGroupFilled() {
+ val dag = DagDef("d")
+ val extract = dag.task<Unit>("extract", NoopTask::class.java)
+ val staging = dag.taskGroup("staging")
+ extract.before(staging)
+ staging.task<Unit>("stage", NoopTask::class.java)
+
+ Bundle().register(dag)
+
+ assertEquals(setOf("extract"), upstreamIds(dag, "staging.stage"))
+ }
+
+ @Test
+ @DisplayName("Should draw edges for every group and task in a combined flow")
+ fun shouldWireGroupsInCombinedFlow() {
+ val dag = DagDef("d")
+ val staging = dag.taskGroup("staging")
+ staging.task<Unit>("stage", NoopTask::class.java)
+ val extract = dag.task<Unit>("extract", NoopTask::class.java)
+ val load = dag.task<Unit>("load", NoopTask::class.java)
+ Deps.Flow.of(staging, extract).before(load)
+
+ Bundle().register(dag)
+
+ assertEquals(setOf("staging.stage", "extract"), upstreamIds(dag, "load"))
+ }
+
+ @Test
+ @DisplayName("Should reject a cycle that a group edge closes")
+ fun shouldRejectCycleThroughGroup() {
+ val dag = DagDef("d")
+ val staging = dag.taskGroup("staging")
+ val stage = staging.task<Unit>("stage", NoopTask::class.java)
+ stage.before(staging)
+
+ val error = assertThrows(IllegalArgumentException::class.java) {
Bundle().register(dag) }
+
+ assertEquals("Task dependencies in Dag 'd' contain a cycle involving task
'staging.stage'", error.message)
+ }
+}
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 9acaa23f6ce..fad223f605f 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
@@ -141,7 +141,7 @@ internal class ArgValuesTest {
.distinct()
.forEach { dag.addTask(it) }
val def = TaskDef("consumer", NoopArgTask::class.java)
- Refs.record(dag, listOf("consumer")) { Refs.call<Unit>(def,
*inputs.toTypedArray()) }
+ Refs.record(dag, listOf("consumer"), emptyList()) { 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
index ffb77fa7b40..975dd8e449c 100644
--- 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
@@ -23,6 +23,7 @@ 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.Deps
import org.apache.airflow.sdk.LiteralArg
import org.apache.airflow.sdk.Task
import org.apache.airflow.sdk.TaskDef
@@ -33,6 +34,12 @@ import org.junit.jupiter.api.Assertions.assertThrows
import org.junit.jupiter.api.DisplayName
import org.junit.jupiter.api.Test
+/** Stands in for the generated wiring view of a task group. */
+private fun groupView(id: String) =
+ object : Deps.TaskGroup {
+ override fun groupId() = id
+ }
+
private class NoopRefTask : Task {
override fun execute(
context: Context,
@@ -45,9 +52,9 @@ internal class RefsTest {
@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))
+ Refs.record(dag, listOf("p", "c"), emptyList()) {
+ 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")
@@ -62,11 +69,11 @@ internal class RefsTest {
@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))
+ Refs.record(dag, listOf("a", "b"), emptyList()) {
+ 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)))
+ first.before(Refs.node<Unit>("", TaskDef("b", NoopRefTask::class.java)))
}
assertEquals(setOf("a", "b"), dag.tasks.keys)
@@ -78,7 +85,7 @@ internal class RefsTest {
fun shouldPassWhenWiringComplete() {
val dag = DagDef("d")
- Refs.record(dag, listOf("t")) { Refs.node<Unit>(TaskDef("t",
NoopRefTask::class.java)) }
+ Refs.record(dag, listOf("t"), emptyList()) { Refs.node<Unit>("",
TaskDef("t", NoopRefTask::class.java)) }
}
@Test
@@ -88,7 +95,7 @@ internal class RefsTest {
val error =
assertThrows(IllegalArgumentException::class.java) {
- Refs.record(dag, listOf("t", "x", "y")) { Refs.node<Unit>(TaskDef("t",
NoopRefTask::class.java)) }
+ Refs.record(dag, listOf("t", "x", "y"), emptyList()) {
Refs.node<Unit>("", TaskDef("t", NoopRefTask::class.java)) }
}
assertEquals(
@@ -103,7 +110,7 @@ internal class RefsTest {
fun shouldRefuseWiringOutsideRecording() {
val error =
assertThrows(IllegalStateException::class.java) {
- Refs.node<Unit>(TaskDef("t", NoopRefTask::class.java))
+ Refs.node<Unit>("", TaskDef("t", NoopRefTask::class.java))
}
assertEquals(
@@ -118,8 +125,8 @@ internal class RefsTest {
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)
+ Refs.record(DagDef("d"), listOf("t"), emptyList()) {
+ Refs.call<Unit>("", TaskDef("t", NoopRefTask::class.java),
Arg.lit(1), null)
}
}
@@ -131,9 +138,9 @@ internal class RefsTest {
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))
+ Refs.record(DagDef("d"), listOf("t"), emptyList()) {
+ Refs.node<Unit>("", TaskDef("t", NoopRefTask::class.java))
+ Refs.call<Unit>("", TaskDef("t", NoopRefTask::class.java),
Arg.lit(1))
}
}
@@ -148,11 +155,52 @@ internal class RefsTest {
fun shouldRefuseNestedRecording() {
val error =
assertThrows(IllegalStateException::class.java) {
- Refs.record(DagDef("outer"), emptyList()) {
- Refs.record(DagDef("inner"), emptyList()) {}
+ Refs.record(DagDef("outer"), emptyList(), emptyList()) {
+ Refs.record(DagDef("inner"), emptyList(), emptyList()) {}
}
}
assertEquals("Dag wiring is already being recorded on this thread",
error.message)
}
+
+ @Test
+ @DisplayName("Should make every group before the wiring runs and register
each task in its own")
+ fun shouldRegisterGroupedTaskInGroup() {
+ val dag = DagDef("d")
+ Refs.record(
+ dag,
+ listOf("extract", "staging.checks.nulls"),
+ listOf("staging", "staging.checks", "staging.empty"),
+ ) {
+ val extract = Refs.node<Unit>("", TaskDef("extract",
NoopRefTask::class.java))
+ extract.before(groupView("staging"))
+ Refs.node<Unit>("staging.checks", TaskDef("staging.checks.nulls",
NoopRefTask::class.java))
+ }
+
+ // staging.empty holds no task, so only the group list can have made it.
+ assertEquals(listOf("staging", "staging.checks", "staging.empty"),
dag.groups.keys.toList())
+ assertEquals(listOf("staging.checks.nulls"),
dag.groups.getValue("staging.checks").taskIds)
+ assertEquals(1, dag.groupEdges.size)
+ }
+
+ @Test
+ @DisplayName("Should resolve the group a wiring-view group stands for")
+ fun shouldResolveGroupOfView() {
+ val dag = DagDef("d")
+ Refs.record(dag, listOf("staging.stage"), listOf("staging")) {
+ Refs.node<Unit>("staging", TaskDef("staging.stage",
NoopRefTask::class.java))
+ assertEquals(listOf("staging.stage"), groupView("staging").nodes().map {
it.id })
+ }
+ }
+
+ @Test
+ @DisplayName("Should fail naming a group the Dag does not have")
+ fun shouldFailOnUnknownGroup() {
+ val error =
+ assertThrows(IllegalArgumentException::class.java) {
+ Refs.record(DagDef("d"), emptyList(), emptyList()) {
Refs.group("staging") }
+ }
+
+ assertEquals("Dag 'd' has no task group 'staging'", error.message)
+ }
}