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,
+    )
+  }
+}

Reply via email to