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 402d7885194 Java SDK: Resolve a task's arguments from the Dag's own
wiring (#73596)
402d7885194 is described below
commit 402d7885194dc93decc5c5a9c37d46c2de356db7
Author: Jason(Zhe-You) Liu <[email protected]>
AuthorDate: Fri Oct 2 19:14:39 2026 +0800
Java SDK: Resolve a task's arguments from the Dag's own wiring (#73596)
* Java SDK: Resolve a task's arguments from the Dag's own wiring
A Dag authored in Java has no Python call site, so the supervisor sends no
argument bindings for it and a task's data parameters had nothing to
resolve against. Where the Dag wired inputs for a task, those stand in;
a stub-backed task keeps reading its bindings, including when the call
site bound none.
The authoring surface that records those inputs lands in the next commit,
so nothing declares them yet.
* Keep taskDef out of Context's constructor and reuse resolveWiredAll
* Use ADR-0007 wording for the wired-arity warning
---
.../org/apache/airflow/sdk/BuilderProcessor.kt | 2 +-
.../kotlin/org/apache/airflow/sdk/BuilderTest.kt | 4 +-
.../src/main/kotlin/org/apache/airflow/sdk/Arg.kt | 11 +-
.../main/kotlin/org/apache/airflow/sdk/Context.kt | 10 +-
.../main/kotlin/org/apache/airflow/sdk/DagDef.kt | 1 +
.../kotlin/org/apache/airflow/sdk/InputTask.kt | 2 +-
.../org/apache/airflow/sdk/execution/Task.kt | 7 +-
.../org/apache/airflow/sdk/internal/ArgValues.kt | 124 ++++++++++
.../org/apache/airflow/sdk/internal/TaskArgs.kt | 42 +++-
.../org/apache/airflow/sdk/ArgTestSupport.kt | 15 ++
.../kotlin/org/apache/airflow/sdk/ArgValuesTest.kt | 2 +-
.../kotlin/org/apache/airflow/sdk/InputTaskTest.kt | 90 +++++++
.../org/apache/airflow/sdk/execution/TaskTest.kt | 24 ++
.../apache/airflow/sdk/internal/ArgValuesTest.kt | 266 +++++++++++++++++++++
14 files changed, 582 insertions(+), 18 deletions(-)
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 049aba0d060..21b4165728e 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
@@ -343,7 +343,7 @@ class BuilderProcessor : AbstractProcessor() {
val paramType = TypeName.get(param.type)
if (param.isTaskInput) {
executeSpec.addStatement(
- $$"$T $L = $T.bindInput(client, $T.class)",
+ $$"$T $L = $T.bindInput(context, client, $T.class)",
paramType,
param.local,
ARG_VALUES_TYPE,
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 6fd52e01b75..6decad1a906 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
@@ -427,7 +427,7 @@ class BuilderTest {
public static final class Named implements Task {
@Override
public void execute(Context context, Client client) throws
Exception {
- TestExample.ScoreInput context_ = ArgValues.bindInput(client,
TestExample.ScoreInput.class);
+ TestExample.ScoreInput context_ = ArgValues.bindInput(context,
client, TestExample.ScoreInput.class);
new TestExample().named(context_);
}
}
@@ -814,7 +814,7 @@ class BuilderTest {
public static final class Score implements Task {
@Override
public void execute(Context context, Client client) throws
Exception {
- TestExample.ScoreInput input = ArgValues.bindInput(client,
TestExample.ScoreInput.class);
+ TestExample.ScoreInput input = ArgValues.bindInput(context,
client, TestExample.ScoreInput.class);
client.setXCom(new TestExample().score(client, input));
}
}
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 8bedb4f4cb8..556ed5a7d4a 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
@@ -20,13 +20,20 @@
package org.apache.airflow.sdk
/**
- * A value a task can be given, of which [TaskRef] — the output of an upstream
- * task — is the only form so far.
+ * A value a task can be given: the output of an upstream task, carried by the
+ * [TaskRef] that task's registration returned, or an inline constant.
+ *
+ * A constant has to be wrapped because a bare `Double` cannot implement this
+ * type: boxed types only, no primitives.
*
* @param T Type of the value.
*/
sealed class Arg<T>
+internal class LiteralArg<T>(
+ internal val value: T?,
+) : Arg<T>()
+
/**
* The output of a registered task, and the task's place in the flow.
*
diff --git a/java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/Context.kt
b/java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/Context.kt
index ece4d69b4f7..69badbd49f6 100644
--- a/java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/Context.kt
+++ b/java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/Context.kt
@@ -124,8 +124,14 @@ data class Context(
@JvmField val dagRun: DagRun,
@JvmField val ti: TaskInstance,
) {
+ /** Registration of the executing task; resolves wired data inputs. */
+ internal var taskDef: TaskDef? = null
+
internal companion object {
- fun from(request: StartupDetails) =
+ fun from(
+ request: StartupDetails,
+ taskDef: TaskDef? = null,
+ ): Context =
Context(
dagRun =
with(request.tiContext.dagRun) {
@@ -141,6 +147,6 @@ data class Context(
)
},
ti = with(request.ti) { TaskInstance(dagId, runId, taskId, mapIndex,
tryNumber) },
- )
+ ).also { it.taskDef = taskDef }
}
}
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 6ba3ff7f165..c579ddd517b 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
@@ -170,6 +170,7 @@ class TaskDef(
}
internal val configValues = linkedMapOf<String, Any>()
+ internal val inputs = mutableListOf<Arg<*>>()
internal val upstreams = linkedSetOf<TaskDef>()
internal var owner: DagDef? = null
diff --git a/java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/InputTask.kt
b/java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/InputTask.kt
index 52276f36f4e..c8dc446e408 100644
--- a/java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/InputTask.kt
+++ b/java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/InputTask.kt
@@ -57,7 +57,7 @@ interface InputTask<I : TaskInput> : Task {
override fun execute(
context: Context,
client: Client,
- ) = execute(context, client, ArgValues.bindInput(client, inputType()))
+ ) = execute(context, client, ArgValues.bindInput(context, client,
inputType()))
/**
* Executes this task.
diff --git
a/java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/execution/Task.kt
b/java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/execution/Task.kt
index 1c189061ae5..6490e221ee9 100644
--- a/java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/execution/Task.kt
+++ b/java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/execution/Task.kt
@@ -74,9 +74,10 @@ internal object TaskRunner {
request: StartupDetails,
client: Client,
): Any {
- val definition =
- bundle.taskDef(request.ti.dagId, request.ti.taskId)?.definition
+ val taskDef =
+ bundle.taskDef(request.ti.dagId, request.ti.taskId)
?: return TaskResult.of(TaskState.State.REMOVED)
+ val definition = taskDef.definition
val instance =
try {
definition.getDeclaredConstructor().newInstance()
@@ -103,7 +104,7 @@ internal object TaskRunner {
return TaskResult.failure(request.tiContext.shouldRetry)
}
return try {
- instance.execute(Context.from(request), client)
+ instance.execute(Context.from(request, taskDef), client)
TaskResult.success()
} catch (e: CancellationException) {
throw e // Let coroutine cancellation propagate so the task coroutine
unwinds.
diff --git
a/java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/internal/ArgValues.kt
b/java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/internal/ArgValues.kt
index 80e42a259a2..559e975645b 100644
--- a/java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/internal/ArgValues.kt
+++ b/java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/internal/ArgValues.kt
@@ -23,9 +23,17 @@ package org.apache.airflow.sdk.internal
import com.fasterxml.jackson.databind.ObjectMapper
import com.fasterxml.jackson.databind.json.JsonMapper
+import kotlinx.coroutines.Dispatchers
+import kotlinx.coroutines.async
+import kotlinx.coroutines.awaitAll
+import kotlinx.coroutines.runBlocking
+import org.apache.airflow.sdk.Arg
import org.apache.airflow.sdk.Client
+import org.apache.airflow.sdk.Context
+import org.apache.airflow.sdk.LiteralArg
import org.apache.airflow.sdk.MissingXComException
import org.apache.airflow.sdk.TaskInput
+import org.apache.airflow.sdk.TaskRef
import org.apache.airflow.sdk.execution.ArgBinding
import org.apache.airflow.sdk.execution.Logger
import java.lang.reflect.Field
@@ -42,6 +50,15 @@ import java.lang.reflect.Type
* graph the scheduler ordered the run by. Flat data parameters resolve the
* binding at their position (through [TaskArgs]); [TaskInput] fields resolve
* bindings by name.
+ *
+ * A natively authored Dag has no stub call site, so the supervisor sends no
+ * bindings for it and the inputs the Dag itself wired stand in. When the
+ * supervisor sends bindings they are used for every parameter; the Dag's own
+ * inputs are read only when it sends none.
+ *
+ * A count that does not match is fatal for flat parameters and a warning for a
+ * [TaskInput]: a position has no name to fall back on, while a field does, so
+ * the task still runs on what it can bind.
*/
object ArgValues {
private val mapper: ObjectMapper =
JsonMapper.builder().build().findAndRegisterModules()
@@ -68,9 +85,18 @@ object ArgValues {
*/
@JvmStatic
fun <I : TaskInput> bindInput(
+ context: Context,
client: Client,
type: Class<I>,
): I {
+ // Runtime bindings carry argument names to match fields against. A wired
+ // input carries none, so it decodes into the whole input at once -- which
+ // is well defined because a TaskInput is a task's only data parameter.
+ wiredInputs(context, client)?.let { wired ->
+ warnWiredArity(client, type, wired.size)
+ return type.cast(decode(resolveWiredAll(wired.take(1), client).single(),
type))
+ ?: throw missingInput(wired[0], type.simpleName)
+ }
val input = newInput(type)
val arguments = ArgIndex(client.argBindings)
val unfilled = mutableListOf<String>()
@@ -147,6 +173,104 @@ object ArgValues {
type: Type,
): Any? = decode(client.resolveBinding(binding), type)
+ /**
+ * Reports a Dag that wired more inputs than a [TaskInput] can take. A
+ * [TaskInput] is a task's only data parameter, so exactly one input feeds
+ * it; the extras are ignored and the first input is used, leaving the task
+ * to run on what it can bind rather than failing the run outright.
+ */
+ private fun warnWiredArity(
+ client: Client,
+ type: Class<*>,
+ wired: Int,
+ ) {
+ if (wired <= 1) return
+ logger.warning(
+ "Dag's call passed argument(s) the task handler does not declare",
+ mapOf(
+ "task_id" to client.details.ti.taskId,
+ "input" to type.simpleName,
+ "declared" to 1,
+ "wired" to wired,
+ ),
+ )
+ }
+
+ /**
+ * Resolves every wired input at once: each upstream is read once however
+ * many parameters it feeds, and several upstreams are read concurrently.
+ * The supervisor protocol matches responses to requests by id, so a task
+ * wired to several upstreams waits roughly one round trip rather than one
+ * per parameter.
+ */
+ internal fun resolveWiredAll(
+ inputs: List<Arg<*>>,
+ client: Client,
+ ): List<Any?> {
+ val upstreams = inputs.filterIsInstance<TaskRef<*>>().map { it.def.id
}.distinct()
+ val fetched =
+ when (upstreams.size) {
+ 0 -> emptyMap()
+ 1 -> mapOf(upstreams[0] to client.getXCom(taskId = upstreams[0]))
+ else ->
+ runBlocking {
+ upstreams.map { taskId -> async(Dispatchers.IO) { taskId to
client.getXCom(taskId = taskId) } }.awaitAll()
+ }.toMap()
+ }
+ return inputs.map { input ->
+ when (input) {
+ is TaskRef<*> -> fetched[input.def.id]
+ is LiteralArg<*> -> input.value
+ }
+ }
+ }
+
+ /** Decodes an already-resolved wired value into [type], passing null
through. */
+ internal fun decodeWired(
+ value: Any?,
+ type: Type,
+ ): Any? = decode(value, type)
+
+ /**
+ * The inputs the Dag wired for this task, or null when the run's arguments
+ * come from the stub call site. A task with no wired inputs reads the
+ * bindings, so a stub call that bound nothing keeps its own diagnostics.
+ */
+ internal fun wiredInputs(
+ context: Context,
+ client: Client,
+ ): List<Arg<*>>? = if (client.argBindings.isEmpty())
context.taskDef?.inputs?.takeIf { it.isNotEmpty() } else null
+
+ /**
+ * The failure for a wired argument that resolved to nothing where a value is
+ * required, naming [target] — the position of the parameter it feeds.
+ */
+ internal fun missingWired(
+ input: Arg<*>,
+ target: String,
+ ): MissingXComException =
+ when (input) {
+ is TaskRef<*> -> MissingXComException(input.def.id, target)
+ is LiteralArg<*> ->
+ MissingXComException(
+ "Task parameter '$target' is wired to a null literal, but has a
primitive type that cannot " +
+ "be null; declare a boxed type (e.g. Integer instead of int) to
receive null.",
+ )
+ }
+
+ /** The failure for a wired input that resolved to nothing for a
[TaskInput]. */
+ private fun missingInput(
+ input: Arg<*>,
+ target: String,
+ ): MissingXComException =
+ when (input) {
+ is TaskRef<*> ->
+ MissingXComException(
+ "Input '$target' requires an XCom from task '${input.def.id}', but
none was pushed.",
+ )
+ is LiteralArg<*> -> MissingXComException("Input '$target' is wired to a
null literal, so there is nothing to bind.")
+ }
+
/**
* Builds the failure for a binding that resolved to nothing where a value is
* required, naming [target] — the stub argument, or the [TaskInput] field
diff --git
a/java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/internal/TaskArgs.kt
b/java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/internal/TaskArgs.kt
index 31c840662c3..14c6404758b 100644
--- a/java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/internal/TaskArgs.kt
+++ b/java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/internal/TaskArgs.kt
@@ -19,10 +19,12 @@
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.MissingXComException
import org.apache.airflow.sdk.execution.ArgBinding
+import java.lang.reflect.Type
/**
* @suppress
@@ -46,6 +48,8 @@ class TaskArgs private constructor(
private val context: Context,
private val client: Client,
private val arguments: List<ArgBinding>,
+ private val wired: List<Arg<*>>?,
+ private val wiredValues: List<Any?>,
) {
companion object {
/**
@@ -71,6 +75,16 @@ class TaskArgs private constructor(
client: Client,
declared: Int,
): TaskArgs {
+ ArgValues.wiredInputs(context, client)?.let { wired ->
+ // Positional binding is strict in both directions: a parameter has no
+ // name to fall back on, so a count that does not match cannot be
+ // resolved and the run fails here rather than mid-task.
+ check(wired.size == declared) {
+ "Task '${context.ti.taskId}' declares $declared data parameter(s) " +
+ "but the Dag wired ${wired.size} argument(s)"
+ }
+ return TaskArgs(context, client, emptyList(), wired,
ArgValues.resolveWiredAll(wired, client))
+ }
val bound = client.argBindings
val arguments = if (bound.size == declared) bound else bound.filterNot {
it.fromDefault }
check(arguments.size == declared) {
@@ -83,7 +97,7 @@ class TaskArgs private constructor(
"Task '${context.ti.taskId}' declares $declared data parameter(s) " +
"but the stub call bound ${bound.size} argument(s)$defaults"
}
- return TaskArgs(context, client, arguments)
+ return TaskArgs(context, client, arguments, null, emptyList())
}
}
@@ -96,7 +110,7 @@ class TaskArgs private constructor(
fun <T : Any> get(
position: Int,
type: Class<T>,
- ): T? = type.cast(ArgValues.valueAt(client, arguments[position], type))
+ ): T? = type.cast(valueAt(position, type))
/**
* Resolves the argument bound at [position] into the generic [type], passing
@@ -108,7 +122,7 @@ class TaskArgs private constructor(
fun <T : Any> get(
position: Int,
type: TypeRef<T>,
- ): T? = ArgValues.valueAt(client, arguments[position], type.type) as T?
+ ): T? = valueAt(position, type.type) as T?
/**
* Resolves the argument bound at [position] into [type], which must not be
@@ -136,7 +150,23 @@ class TaskArgs private constructor(
type: TypeRef<T>,
): T = get(position, type) ?: throw missingAt(position)
- // The stub signature's own parameter name is the clearest label for a
failure
- // here: it is what the Dag author has to change.
- private fun missingAt(at: Int) = ArgValues.missing(arguments[at],
context.ti.taskId)
+ private fun valueAt(
+ position: Int,
+ type: Type,
+ ): Any? =
+ if (wired != null) {
+ ArgValues.decodeWired(wiredValues[position], type)
+ } else {
+ ArgValues.valueAt(client, arguments[position], type)
+ }
+
+ // A bound argument is labelled with the stub signature's own parameter name,
+ // which is what the Dag author has to change. A wired one has no such name,
+ // so it is labelled by the position it feeds.
+ private fun missingAt(at: Int): MissingXComException =
+ if (wired != null) {
+ ArgValues.missingWired(wired[at], "#$at")
+ } else {
+ ArgValues.missing(arguments[at], context.ti.taskId)
+ }
}
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 1fff219370b..1cc3ff3dbae 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
@@ -94,3 +94,18 @@ internal fun taskContext(): Context =
dagRun = DagRun("d", "r", null, null, null, null, null, emptyMap()),
ti = TaskInstance("d", "r", "t", null, 1),
)
+
+internal class NoopTask : Task {
+ override fun execute(
+ context: Context,
+ client: Client,
+ ) = Unit
+}
+
+/** A context whose task was wired by its Dag with the given inputs. */
+internal fun contextWiredWith(inputs: List<Arg<*>>): Context {
+ val def = TaskDef("t", NoopTask::class.java)
+ DagDef("d").addTask(def)
+ def.inputs += inputs
+ return taskContext().also { it.taskDef = def }
+}
diff --git
a/java-sdk/sdk/src/test/kotlin/org/apache/airflow/sdk/ArgValuesTest.kt
b/java-sdk/sdk/src/test/kotlin/org/apache/airflow/sdk/ArgValuesTest.kt
index 16eaa161fbd..b6b5347ba4e 100644
--- a/java-sdk/sdk/src/test/kotlin/org/apache/airflow/sdk/ArgValuesTest.kt
+++ b/java-sdk/sdk/src/test/kotlin/org/apache/airflow/sdk/ArgValuesTest.kt
@@ -92,7 +92,7 @@ private fun <I : TaskInput> bind(
xcoms: Map<String, Any?> = emptyMap(),
): I {
val (client, _) = clientWith(bindings, xcoms)
- return ArgValues.bindInput(client, type)
+ return ArgValues.bindInput(taskContext(), client, type)
}
private fun literal(
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 b33abcfba0d..ca557b92611 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
@@ -21,8 +21,12 @@
package org.apache.airflow.sdk
+import org.apache.airflow.sdk.execution.Level
+import org.apache.airflow.sdk.execution.LogSender
import org.junit.jupiter.api.Assertions.assertEquals
+import org.junit.jupiter.api.Assertions.assertNull
import org.junit.jupiter.api.Assertions.assertThrows
+import org.junit.jupiter.api.Assertions.assertTrue
import org.junit.jupiter.api.DisplayName
import org.junit.jupiter.api.Test
@@ -114,6 +118,92 @@ internal class InputTaskTest {
assertEquals(0.5, input.threshold)
}
+ @Test
+ @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))))
+ val (client, _) = clientWith(null)
+ val task = Summarize()
+
+ task.execute(context, client)
+
+ val input = requireNotNull(task.received)
+ assertEquals("emea", input.region)
+ assertEquals(0.5, input.threshold)
+ }
+
+ @Test
+ @DisplayName("Should fail when the input wired to a TaskInput resolves to
nothing")
+ fun shouldRejectNullWiredTaskInput() {
+ val context = contextWiredWith(listOf(LiteralArg<Map<String, Any?>>(null)))
+ val (client, _) = clientWith(null)
+
+ val error =
+ assertThrows(MissingXComException::class.java) {
Summarize().execute(context, client) }
+
+ assertEquals(
+ "Input 'SummaryInput' is wired to a null literal, so there is nothing to
bind.",
+ error.message,
+ )
+ }
+
+ @Test
+ @DisplayName("Should warn and bind the first input when the Dag wired more
than a TaskInput takes")
+ fun shouldWarnWhenMoreInputsWiredThanTaskInputTakes() {
+ LogSender.messages.clear()
+ val context =
+ contextWiredWith(
+ listOf(
+ LiteralArg(mapOf("region" to "emea", "threshold" to 0.5)),
+ LiteralArg(mapOf("region" to "apac", "threshold" to 0.1)),
+ ),
+ )
+ val (client, _) = clientWith(null)
+ val task = Summarize()
+
+ task.execute(context, client)
+
+ assertEquals("emea", requireNotNull(task.received).region)
+ val message = LogSender.messages.single { it.level == Level.WARNING }
+ assertEquals("Dag's call passed argument(s) the task handler does not
declare", message.event)
+ assertEquals(1, message.arguments["declared"])
+ assertEquals(2, message.arguments["wired"])
+ assertEquals("SummaryInput", message.arguments["input"])
+ }
+
+ @Test
+ @DisplayName("Should stay quiet when the Dag wired exactly one input for a
TaskInput")
+ fun shouldNotWarnWhenOneInputWiredForTaskInput() {
+ LogSender.messages.clear()
+ val context = contextWiredWith(listOf(LiteralArg(mapOf("region" to "emea",
"threshold" to 0.5))))
+ val (client, _) = clientWith(null)
+
+ Summarize().execute(context, client)
+
+ assertTrue(LogSender.messages.none { it.level == Level.WARNING }) {
+ "unexpected warnings: ${LogSender.messages.map { it.event }}"
+ }
+ }
+
+ @Test
+ @DisplayName("Should still match a TaskInput by name when the stub call
bound no arguments")
+ fun shouldBindTaskInputWhenStubBoundNothing() {
+ // No bindings and no wired inputs: the name-matching path still owns this,
+ // so the fields take their defaults rather than the whole-input decode a
+ // wired TaskInput gets.
+ val (client, _) = clientWith(null)
+ val task = Summarize()
+
+ task.execute(taskContext(), client)
+
+ val input = requireNotNull(task.received)
+ assertNull(input.region)
+ assertEquals(0.0, input.threshold)
+ }
+
@Test
@DisplayName("Should resolve the input type a superclass declared")
fun shouldResolveInheritedInputType() {
diff --git
a/java-sdk/sdk/src/test/kotlin/org/apache/airflow/sdk/execution/TaskTest.kt
b/java-sdk/sdk/src/test/kotlin/org/apache/airflow/sdk/execution/TaskTest.kt
index 5d16eda1ccb..548a099bf81 100644
--- a/java-sdk/sdk/src/test/kotlin/org/apache/airflow/sdk/execution/TaskTest.kt
+++ b/java-sdk/sdk/src/test/kotlin/org/apache/airflow/sdk/execution/TaskTest.kt
@@ -169,6 +169,19 @@ class TaskTest {
}
}
+ @Test
+ @DisplayName("Should thread the task definition into the execution context")
+ fun shouldThreadTaskDefIntoContext() {
+ val result =
+ runTask(
+ bundleWith("asserting", TaskDefAssertingTask::class.java),
+ startupDetails(taskId = "asserting"),
+ noOpClient(),
+ )
+
+ Assertions.assertInstanceOf(SucceedTask::class.java, result)
+ }
+
private fun bundleWith(
taskId: String,
taskClass: Class<out Task>,
@@ -298,4 +311,15 @@ class TaskTest {
client: Client,
): Unit = throw IllegalStateException("should not be reachable")
}
+
+ class TaskDefAssertingTask : Task {
+ override fun execute(
+ context: Context,
+ client: Client,
+ ) {
+ check(context.taskDef?.id == context.ti.taskId) {
+ "expected the runner to thread the task definition into the context"
+ }
+ }
+ }
}
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
new file mode 100644
index 00000000000..768ad066f5c
--- /dev/null
+++
b/java-sdk/sdk/src/test/kotlin/org/apache/airflow/sdk/internal/ArgValuesTest.kt
@@ -0,0 +1,266 @@
+/*
+ * Licensed to the Apache Software Foundation (ASF) under one
+ * or more contributor license agreements. See the NOTICE file
+ * distributed with this work for additional information
+ * regarding copyright ownership. The ASF licenses this file
+ * to you under the Apache License, Version 2.0 (the
+ * "License"); you may not use this file except in compliance
+ * with the License. You may obtain a copy of the License at
+ *
+ * http://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing,
+ * software distributed under the License is distributed on an
+ * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
+ * KIND, either express or implied. See the License for the
+ * specific language governing permissions and limitations
+ * under the License.
+ */
+
+@file:Suppress("PLATFORM_CLASS_MAPPED_TO_KOTLIN")
+
+package org.apache.airflow.sdk.internal
+
+import 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.DagRun
+import org.apache.airflow.sdk.LiteralArg
+import org.apache.airflow.sdk.MissingXComException
+import org.apache.airflow.sdk.Task
+import org.apache.airflow.sdk.TaskDef
+import org.apache.airflow.sdk.TaskInstance
+import org.apache.airflow.sdk.TaskRef
+import org.apache.airflow.sdk.execution.comm.ConnectionResult
+import org.apache.airflow.sdk.execution.comm.StartupDetails
+import org.apache.airflow.sdk.execution.comm.VariableResult
+import org.apache.airflow.sdk.execution.comm.XComResult
+import org.junit.jupiter.api.Assertions.assertEquals
+import org.junit.jupiter.api.Assertions.assertNull
+import org.junit.jupiter.api.Assertions.assertThrows
+import org.junit.jupiter.api.DisplayName
+import org.junit.jupiter.api.Test
+import org.apache.airflow.sdk.execution.Client as Transport
+import org.apache.airflow.sdk.execution.comm.TaskInstance as CommTaskInstance
+
+private class NoopArgTask : Task {
+ override fun execute(
+ context: Context,
+ client: Client,
+ ) = Unit
+}
+
+/** Resolution of the inputs a Dag declared in Java, without runtime bindings.
*/
+internal class ArgValuesTest {
+ /** Upstream task ids read through the transport, in arrival order. */
+ private val pulls = java.util.concurrent.CopyOnWriteArrayList<String>()
+
+ private fun clientWith(xcomsByTask: Map<String, Any?>): Client =
+ Client(
+ StartupDetails().also {
+ it.ti =
+ CommTaskInstance().also { ti ->
+ ti.taskId = "consumer"
+ ti.dagId = "d"
+ ti.runId = "r"
+ ti.tryNumber = 1
+ }
+ },
+ object : Transport {
+ override fun getConnection(id: String): ConnectionResult = throw
NotImplementedError()
+
+ override fun getVariable(key: String): VariableResult = throw
NotImplementedError()
+
+ override fun setVariable(
+ key: String,
+ value: String,
+ description: String?,
+ ): Unit = throw NotImplementedError()
+
+ override fun deleteVariable(key: String): Unit = throw
NotImplementedError()
+
+ override fun getXCom(
+ key: String,
+ dagId: String,
+ taskId: String,
+ runId: String,
+ mapIndex: Int?,
+ includePriorDates: Boolean,
+ ): XComResult {
+ pulls += taskId
+ arrival?.let {
+ it.countDown()
+ check(it.await(5, java.util.concurrent.TimeUnit.SECONDS)) {
+ "wired upstreams were read one at a time"
+ }
+ }
+ return XComResult().also { it.value = xcomsByTask[taskId] }
+ }
+
+ override fun setXCom(
+ key: String,
+ value: Any,
+ dagId: String,
+ taskId: String,
+ runId: String,
+ mapIndex: Int,
+ ): Unit = throw NotImplementedError()
+ },
+ )
+
+ private var arrival: java.util.concurrent.CountDownLatch? = null
+
+ private fun contextFor(inputs: List<Arg<*>>): Context {
+ val dag = DagDef("d")
+ inputs
+ .filterIsInstance<TaskRef<*>>()
+ .map { it.def }
+ .distinct()
+ .forEach { dag.addTask(it) }
+ val def = TaskDef("consumer", NoopArgTask::class.java)
+ dag.addTask(def)
+ def.inputs += inputs
+ return contextWithoutTaskDef().also { it.taskDef = def }
+ }
+
+ private fun contextWithoutTaskDef(): Context =
+ Context(
+ dagRun = DagRun("d", "r", null, null, null, null, null, emptyMap()),
+ ti = TaskInstance("d", "r", "consumer", null, 1),
+ )
+
+ private fun handleFor(taskId: String): TaskRef<Any> =
TaskRef(TaskDef(taskId, NoopArgTask::class.java))
+
+ @Test
+ @DisplayName("Should resolve a handle input from the upstream task's XCom")
+ fun shouldResolveHandleInputFromXCom() {
+ val context = contextFor(listOf(handleFor("producer")))
+
+ val args = TaskArgs.of(context, clientWith(mapOf("producer" to 42L)), 1)
+
+ assertEquals(42L, args.require(0, java.lang.Long::class.java))
+ }
+
+ @Test
+ @DisplayName("Should resolve a literal input without touching the client")
+ fun shouldResolveLiteralInput() {
+ val context = contextFor(listOf(LiteralArg(7)))
+
+ val args = TaskArgs.of(context, clientWith(emptyMap()), 1)
+
+ assertEquals(7L, args.require(0, java.lang.Long::class.java))
+ }
+
+ @Test
+ @DisplayName("Should throw MissingXComException when a required upstream
pushed no value")
+ fun shouldThrowForMissingRequiredValue() {
+ val context = contextFor(listOf(handleFor("producer")))
+ val client = clientWith(mapOf("producer" to null))
+
+ val args = TaskArgs.of(context, client, 1)
+
+ assertThrows(MissingXComException::class.java) { args.require(0,
Integer::class.java) }
+ }
+
+ @Test
+ @DisplayName("Should throw MissingXComException for a required null literal")
+ fun shouldThrowForRequiredNullLiteral() {
+ val context = contextFor(listOf(LiteralArg<Int>(null)))
+ val client = clientWith(emptyMap())
+
+ val args = TaskArgs.of(context, client, 1)
+
+ val error = assertThrows(MissingXComException::class.java) {
args.require(0, Integer::class.java) }
+
+ assertEquals(
+ "Task parameter '#0' is wired to a null literal, but has a primitive
type that cannot " +
+ "be null; declare a boxed type (e.g. Integer instead of int) to
receive null.",
+ error.message,
+ )
+ }
+
+ @Test
+ @DisplayName("Should pass null through for optional inputs")
+ fun shouldPassNullThroughForOptionalInputs() {
+ val context = contextFor(listOf(handleFor("producer")))
+
+ val args = TaskArgs.of(context, clientWith(mapOf("producer" to null)), 1)
+
+ assertNull(args.get(0, Integer::class.java))
+ }
+
+ @Test
+ @DisplayName("Should fail when the Dag wired fewer inputs than the task
declares")
+ fun shouldFailOnUnwiredPosition() {
+ val context = contextFor(listOf(handleFor("producer")))
+
+ val error =
+ assertThrows(IllegalStateException::class.java) {
+ TaskArgs.of(context, clientWith(emptyMap()), 2)
+ }
+
+ assertEquals(
+ "Task 'consumer' declares 2 data parameter(s) but the Dag wired 1
argument(s)",
+ error.message,
+ )
+ }
+
+ @Test
+ @DisplayName("Should read an upstream once when several parameters are wired
to the same handle")
+ fun shouldReadEachUpstreamOnce() {
+ val producer = handleFor("producer")
+ val context = contextFor(listOf(producer, producer))
+
+ val args = TaskArgs.of(context, clientWith(mapOf("producer" to 42L)), 2)
+
+ assertEquals(42L, args.require(0, java.lang.Long::class.java))
+ assertEquals(42L, args.require(1, java.lang.Long::class.java))
+ assertEquals(listOf("producer"), pulls)
+ }
+
+ @Test
+ @DisplayName("Should read wired upstreams concurrently rather than one at a
time")
+ fun shouldReadWiredUpstreamsConcurrently() {
+ val context = contextFor(listOf(handleFor("left"), handleFor("right")))
+ // Each read blocks until both have arrived, so sequential resolution
+ // cannot satisfy it and the check inside the transport fails.
+ arrival = java.util.concurrent.CountDownLatch(2)
+
+ val args = TaskArgs.of(context, clientWith(mapOf("left" to 1L, "right" to
2L)), 2)
+
+ assertEquals(1L, args.require(0, java.lang.Long::class.java))
+ assertEquals(2L, args.require(1, java.lang.Long::class.java))
+ assertEquals(setOf("left", "right"), pulls.toSet())
+ }
+
+ @Test
+ @DisplayName("Should fail when the Dag wired more inputs than the task
declares")
+ fun shouldFailOnSurplusWiredInput() {
+ val context = contextFor(listOf(handleFor("left"), handleFor("right")))
+
+ val error =
+ assertThrows(IllegalStateException::class.java) {
+ TaskArgs.of(context, clientWith(emptyMap()), 1)
+ }
+
+ assertEquals(
+ "Task 'consumer' declares 1 data parameter(s) but the Dag wired 2
argument(s)",
+ error.message,
+ )
+ }
+
+ @Test
+ @DisplayName("Should report the stub call when nothing was wired and nothing
was bound")
+ fun shouldFailWithoutTaskDef() {
+ val error =
+ assertThrows(IllegalStateException::class.java) {
+ TaskArgs.of(contextWithoutTaskDef(), clientWith(emptyMap()), 1)
+ }
+
+ assertEquals(
+ "Task 'consumer' declares 1 data parameter(s) but the stub call bound 0
argument(s)",
+ error.message,
+ )
+ }
+}