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 75d9ebf60cb Add task state store API to the Java SDK (#73464)
75d9ebf60cb is described below

commit 75d9ebf60cbb6923943e390e8d517bd33da83543
Author: Andrew Chang <[email protected]>
AuthorDate: Wed Oct 7 02:07:03 2026 +0800

    Add task state store API to the Java SDK (#73464)
    
    * Add task state store API to the Java SDK
    
    Airflow 3.3 added a task state store (AIP-103) and the supervisor already
    handles the Get/Set/Delete/ClearTaskStateStore messages, but the Java SDK
    had no task-facing API for it and capabilities.yaml listed it as
    unsupported. Java tasks could not keep state such as an external job ID
    across retries.
    
    * Expose taskStateStore as a getter and fix the retry scope wording
    
    Review found that a @JvmField on Client cannot be stubbed by Mockito, so a
    Java unit test of a task that reads the store hit a NullPointerException.
    The docs also said entries survive later runs, but the store is scoped to
    one task instance, so they only survive retries within the same Dag run.
    
    * Apply the default retention to task state keys stored from Java
    
    A key stored without a retention never expired, while Python applies
    [state_store] default_retention_days. The coordinator passes that setting
    to the JVM as AIRFLOW__STATE_STORE__DEFAULT_RETENTION_DAYS, so read it
    the same way the Go SDK does: absent means 30 days, 0 means never expire,
    and a malformed value fails instead of retaining for a different period.
    TaskStateStore.NEVER_EXPIRE replaces the old "no retention" meaning, and a
    zero or negative retention is rejected.
    
    * Fail instead of guessing when the default retention is not passed
    
    Review on the Go and TS task state store PRs asked to drop the hardcoded
    30-day fallback: the coordinator always passes
    AIRFLOW__STATE_STORE__DEFAULT_RETENTION_DAYS, so a missing variable means
    the JVM was not launched by the coordinator, and guessing a period there
    would drift from the deployment's config. Also leave a TODO for the
    max_value_storage_bytes warning the Python accessor emits.
---
 .../language-sdks/java.rst                         |  32 ++++
 java-sdk/README.md                                 |   2 +-
 java-sdk/capabilities.yaml                         |   5 +-
 java-sdk/sdk/module.md                             |   2 +-
 .../main/kotlin/org/apache/airflow/sdk/Client.kt   | 128 ++++++++++++++
 .../org/apache/airflow/sdk/execution/Client.kt     |  74 +++++++++
 .../org/apache/airflow/sdk/execution/Comm.kt       |  16 ++
 .../org/apache/airflow/sdk/execution/Frame.kt      |   3 +-
 .../org/apache/airflow/sdk/ArgTestSupport.kt       |  22 +++
 .../kotlin/org/apache/airflow/sdk/ClientTest.kt    | 184 ++++++++++++++++++++-
 .../org/apache/airflow/sdk/execution/CommTest.kt   | 124 +++++++++++++-
 .../org/apache/airflow/sdk/execution/TaskTest.kt   |  19 +++
 .../apache/airflow/sdk/internal/ArgValuesTest.kt   |  22 +++
 13 files changed, 626 insertions(+), 7 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 0f8a591864f..f158f795cca 100644
--- a/airflow-core/docs/authoring-and-scheduling/language-sdks/java.rst
+++ b/airflow-core/docs/authoring-and-scheduling/language-sdks/java.rst
@@ -709,6 +709,38 @@ Durations and date-times are ISO-8601 strings in 
annotations (``retryDelay = "PT
 ``java.time.OffsetDateTime`` values in ``config`` calls.  An unknown key or a 
mismatched value type
 fails the build for an annotation, and the ``config`` call itself for an 
object.
 
+.. _java-sdk/task-state-store:
+
+Task state store
+~~~~~~~~~~~~~~~~
+
+``client.getTaskStateStore()`` gives a task key-value state that is scoped to 
the task instance and
+survives retry attempts within the same Dag run (see 
:doc:`/core-concepts/task-state-store`). Use it to
+remember things like an external job ID so a retried task can resume instead 
of starting over:
+
+.. code-block:: java
+
+    @Builder.Task(id = "submit")
+    public void submit(Client client) throws Exception {
+      var store = client.getTaskStateStore();
+      var jobId = (String) store.get("job_id");
+      if (jobId == null) {
+        jobId = submitJob();
+        store.set("job_id", jobId, Duration.ofHours(6));
+      }
+      waitForJob(jobId);
+      store.delete("job_id");
+    }
+
+``get`` returns ``null`` when the key is not set. ``set`` stores any 
JSON-serializable value. Pass a positive
+``java.time.Duration`` to expire the key after that long, 
``TaskStateStore.NEVER_EXPIRE`` for a key that
+garbage collection skips, or omit the retention to use the deployment's 
``[state_store] default_retention_days``
+(0 means never expire). The coordinator passes that setting to the JVM as
+``AIRFLOW__STATE_STORE__DEFAULT_RETENTION_DAYS``. A zero or negative retention 
is rejected. ``delete`` removes
+one key and ``clear`` removes every key for the task instance. The Java SDK 
does not use a
+``[workers] state_store_backend``: values always go to the metadata database 
as-is, so keys written by Python
+tasks through a custom backend are returned to Java as the raw reference 
marker rather than the stored value.
+
 .. _java-sdk/native-dag-parsing:
 
 Parsing native Java Dags
diff --git a/java-sdk/README.md b/java-sdk/README.md
index 8e78206db39..33a3c560f7b 100644
--- a/java-sdk/README.md
+++ b/java-sdk/README.md
@@ -630,7 +630,7 @@ prek hook regenerate it.
 | capability: `variable-read-write` | MUST | ✓ | 3.3 |  |
 | capability: `self-contained-bundle` | MUST | ✓ | 3.3 | Airflow metadata 
embedded in the jar artifact |
 | capability: `retry-policy` | MAY | ✗ | – | no task-facing retry-policy API 
yet |
-| capability: `task-state-store` | MAY | ✗ | – | no task-facing state-store 
API yet |
+| capability: `task-state-store` | MAY | ✓ | 3.3 | Client.getTaskStateStore() 
get/set/delete/clear |
 | capability: `asset-state-store` | MAY | ✗ | – | no task-facing state-store 
API yet |
 | capability: `asset-event-emit` | MAY | ✗ | – | runtime does not emit asset 
events yet |
 | capability: `asset-event-read` | MAY | ✗ | – | no task-facing asset-event 
API yet |
diff --git a/java-sdk/capabilities.yaml b/java-sdk/capabilities.yaml
index 252f52f0faf..6f61f999240 100644
--- a/java-sdk/capabilities.yaml
+++ b/java-sdk/capabilities.yaml
@@ -83,8 +83,9 @@ capabilities:
     supported: false
     note: "no task-facing retry-policy API yet"
   task-state-store:
-    supported: false
-    note: "no task-facing state-store API yet"
+    supported: true
+    since: "3.3"
+    note: "Client.getTaskStateStore() get/set/delete/clear"
   asset-state-store:
     supported: false
     note: "no task-facing state-store API yet"
diff --git a/java-sdk/sdk/module.md b/java-sdk/sdk/module.md
index c17da64966e..07e272d1431 100644
--- a/java-sdk/sdk/module.md
+++ b/java-sdk/sdk/module.md
@@ -51,7 +51,7 @@ meaning of each dimension is defined in the
 | capability: `variable-read-write` | MUST | ✓ | 3.3 |  |
 | capability: `self-contained-bundle` | MUST | ✓ | 3.3 | Airflow metadata 
embedded in the jar artifact |
 | capability: `retry-policy` | MAY | ✗ | – | no task-facing retry-policy API 
yet |
-| capability: `task-state-store` | MAY | ✗ | – | no task-facing state-store 
API yet |
+| capability: `task-state-store` | MAY | ✓ | 3.3 | Client.getTaskStateStore() 
get/set/delete/clear |
 | capability: `asset-state-store` | MAY | ✗ | – | no task-facing state-store 
API yet |
 | capability: `asset-event-emit` | MAY | ✗ | – | runtime does not emit asset 
events yet |
 | capability: `asset-event-read` | MAY | ✗ | – | no task-facing asset-event 
API yet |
diff --git a/java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/Client.kt 
b/java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/Client.kt
index b01e604ad12..c394e270436 100644
--- a/java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/Client.kt
+++ b/java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/Client.kt
@@ -23,6 +23,11 @@ import org.apache.airflow.sdk.execution.ArgBinding
 import org.apache.airflow.sdk.execution.Client
 import org.apache.airflow.sdk.execution.comm.StartupDetails
 import org.apache.airflow.sdk.execution.decodeArgBindings
+import java.time.Duration
+import java.time.OffsetDateTime
+import java.time.ZoneOffset
+import java.time.temporal.ChronoUnit
+import kotlin.math.floor
 
 /**
  * A connection registered in Airflow's connection store.
@@ -57,6 +62,7 @@ data class Connection(
 class Client internal constructor(
   internal val details: StartupDetails,
   internal val impl: Client,
+  env: (String) -> String? = System::getenv,
 ) {
   internal companion object {
     /**
@@ -65,6 +71,14 @@ class Client internal constructor(
     const val XCOM_RETURN_KEY = "return_value"
   }
 
+  /**
+   * Key-value state scoped to the current task instance.
+   *
+   * Entries survive retries of the task instance within the same Dag run, so
+   * they can carry things like an external job ID across attempts.
+   */
+  val taskStateStore: TaskStateStore = TaskStateStore(details, impl, env)
+
   /**
    * Retrieves a connection from the Airflow connection store.
    *
@@ -215,6 +229,120 @@ class Client internal constructor(
   }
 }
 
+/**
+ * Key-value state scoped to one task instance, shared across its retries
+ * within the same Dag run.
+ *
+ * Values must be JSON-serializable. Every key has an expiry: by default the
+ * deployment's `[state_store] default_retention_days`, which the coordinator
+ * passes to the JVM as `AIRFLOW__STATE_STORE__DEFAULT_RETENTION_DAYS`; [set]
+ * also takes an explicit retention, or [NEVER_EXPIRE] for a key that garbage
+ * collection skips.
+ *
+ * Values are stored in the metadata database as-is; the `[workers]
+ * state_store_backend` used by Python tasks is not applied here.
+ */
+class TaskStateStore internal constructor(
+  private val details: StartupDetails,
+  private val impl: Client,
+  private val env: (String) -> String?,
+) {
+  companion object {
+    /**
+     * Pass as the retention of [set] to store a key that never expires and is
+     * skipped by Airflow's periodic garbage collection.
+     */
+    @JvmField val NEVER_EXPIRE: Duration = ChronoUnit.FOREVER.duration
+
+    internal const val DEFAULT_RETENTION_DAYS_ENV = 
"AIRFLOW__STATE_STORE__DEFAULT_RETENTION_DAYS"
+  }
+
+  /**
+   * Reads the value stored under [key].
+   *
+   * @return The stored value, or `null` if the key is not set.
+   * @throws ApiError if the API call fails.
+   */
+  fun get(key: String): Any? = impl.getTaskStateStore(details.ti.id, 
key)?.value
+
+  /**
+   * Stores [value] under [key], replacing any existing value.
+   *
+   * @param key State key.
+   * @param value Value to store. Must be JSON-serializable.
+   * @param retention How long to keep the key. Must be positive, or
+   *   [NEVER_EXPIRE]; `null` uses `[state_store] default_retention_days`.
+   * @throws IllegalArgumentException if [retention] is zero or negative, or 
the
+   *   default retention from the environment is not a non-negative integer.
+   * @throws IllegalStateException if [retention] is `null` and the coordinator
+   *   did not pass `[state_store] default_retention_days` to the JVM.
+   * @throws ApiError if the API call fails.
+   */
+  @JvmOverloads fun set(
+    key: String,
+    value: Any,
+    retention: Duration? = null,
+  ) {
+    val now = OffsetDateTime.now(ZoneOffset.UTC)
+    val expiresAt =
+      when {
+        retention == null -> resolveDefaultExpiry(now)
+        retention == NEVER_EXPIRE -> null
+        retention.isNegative || retention.isZero ->
+          throw IllegalArgumentException(
+            "Task state retention must be positive or 
TaskStateStore.NEVER_EXPIRE, got $retention for key '$key'",
+          )
+        else -> now.plus(retention)
+      }
+    // TODO: warn when the serialized value exceeds [state_store] 
max_value_storage_bytes once the
+    //   coordinator passes it to the JVM, as the Python accessor does.
+    impl.setTaskStateStore(tiId = details.ti.id, key = key, value = value, 
expiresAt = expiresAt)
+  }
+
+  /**
+   * Deletes the value stored under [key]. Does nothing if the key is not set.
+   *
+   * @throws ApiError if the API call fails.
+   */
+  fun delete(key: String) = impl.deleteTaskStateStore(details.ti.id, key)
+
+  /**
+   * Deletes every key stored for this task instance.
+   *
+   * @throws ApiError if the API call fails.
+   */
+  fun clear() = impl.clearTaskStateStore(details.ti.id)
+
+  private fun resolveDefaultExpiry(now: OffsetDateTime): OffsetDateTime? {
+    val raw =
+      env(DEFAULT_RETENTION_DAYS_ENV)
+        ?: throw IllegalStateException(
+          "$DEFAULT_RETENTION_DAYS_ENV is not set, so the default retention is 
unknown. The coordinator passes " +
+            "[state_store] default_retention_days to the JVM; pass a retention 
or TaskStateStore.NEVER_EXPIRE " +
+            "to set the expiry explicitly.",
+        )
+    val days = parseRetentionDays(raw)
+    return if (days == 0) null else now.plusDays(days.toLong())
+  }
+
+  // Accepts "7.0" because Python's conf.getint does.
+  private fun parseRetentionDays(raw: String): Int {
+    val days =
+      raw.trim().toIntOrNull()
+        ?: raw
+          .trim()
+          .toDoubleOrNull()
+          ?.takeIf { it.isFinite() && it == floor(it) }
+          ?.toInt()
+        ?: throw IllegalArgumentException(
+          "Failed to convert value to int. Please check 
'default_retention_days' key in 'state_store' section. " +
+            "Current value: '$raw'",
+        )
+    require(days >= 0) { "[state_store] default_retention_days must be >= 0, 
got $days. Set to 0 to disable expiry." }
+    return days
+  }
+}
+
 /**
  * Thrown when a task's input resolves to nothing where a value is required —
  * a data parameter or a [TaskInput] field with a primitive type.
diff --git 
a/java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/execution/Client.kt 
b/java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/execution/Client.kt
index d01b151dd36..a4bb15c44cf 100644
--- a/java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/execution/Client.kt
+++ b/java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/execution/Client.kt
@@ -20,16 +20,24 @@
 package org.apache.airflow.sdk.execution
 
 import kotlinx.coroutines.runBlocking
+import org.apache.airflow.sdk.execution.comm.ClearTaskStateStore
 import org.apache.airflow.sdk.execution.comm.ConnectionResult
+import org.apache.airflow.sdk.execution.comm.DeleteTaskStateStore
 import org.apache.airflow.sdk.execution.comm.DeleteVariable
+import org.apache.airflow.sdk.execution.comm.ErrorResponse
 import org.apache.airflow.sdk.execution.comm.GetConnection
+import org.apache.airflow.sdk.execution.comm.GetTaskStateStore
 import org.apache.airflow.sdk.execution.comm.GetVariable
 import org.apache.airflow.sdk.execution.comm.GetXCom
 import org.apache.airflow.sdk.execution.comm.OKResponse
 import org.apache.airflow.sdk.execution.comm.PutVariable
+import org.apache.airflow.sdk.execution.comm.SetTaskStateStore
 import org.apache.airflow.sdk.execution.comm.SetXCom
+import org.apache.airflow.sdk.execution.comm.TaskStateStoreResult
 import org.apache.airflow.sdk.execution.comm.VariableResult
 import org.apache.airflow.sdk.execution.comm.XComResult
+import java.time.OffsetDateTime
+import java.util.UUID
 
 /**
  * @suppress
@@ -74,6 +82,26 @@ interface Client {
     runId: String,
     mapIndex: Int,
   )
+
+  /** Returns `null` when the key is not stored for the task instance. */
+  fun getTaskStateStore(
+    tiId: UUID,
+    key: String,
+  ): TaskStateStoreResult?
+
+  fun setTaskStateStore(
+    tiId: UUID,
+    key: String,
+    value: Any,
+    expiresAt: OffsetDateTime?,
+  )
+
+  fun deleteTaskStateStore(
+    tiId: UUID,
+    key: String,
+  )
+
+  fun clearTaskStateStore(tiId: UUID)
 }
 
 /**
@@ -157,4 +185,50 @@ class CoordinatorClient(
       }
     return runBlocking { exec.communicate<XComResult>(message) }
   }
+
+  override fun getTaskStateStore(
+    tiId: UUID,
+    key: String,
+  ): TaskStateStoreResult? {
+    val message =
+      GetTaskStateStore().also {
+        it.tiId = tiId
+        it.key = key
+      }
+    return runBlocking {
+      exec.communicateOrNullIf<TaskStateStoreResult>(message, 
ErrorResponse.ErrorType.TASK_STORE_NOT_FOUND)
+    }
+  }
+
+  override fun setTaskStateStore(
+    tiId: UUID,
+    key: String,
+    value: Any,
+    expiresAt: OffsetDateTime?,
+  ) {
+    val message =
+      SetTaskStateStore().also {
+        it.tiId = tiId
+        it.key = key
+        it.value = value
+        it.expiresAt = expiresAt
+      }
+    runBlocking { exec.communicate<OKResponse>(message) }
+  }
+
+  override fun deleteTaskStateStore(
+    tiId: UUID,
+    key: String,
+  ) {
+    val message =
+      DeleteTaskStateStore().also {
+        it.tiId = tiId
+        it.key = key
+      }
+    runBlocking { exec.communicate<OKResponse>(message) }
+  }
+
+  override fun clearTaskStateStore(tiId: UUID) {
+    runBlocking { exec.communicate<OKResponse>(ClearTaskStateStore().also { 
it.tiId = tiId }) }
+  }
 }
diff --git 
a/java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/execution/Comm.kt 
b/java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/execution/Comm.kt
index 19bf7922462..caed68d8dbe 100644
--- a/java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/execution/Comm.kt
+++ b/java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/execution/Comm.kt
@@ -209,6 +209,22 @@ class CoordinatorComm(
     }
   }
 
+  /**
+   * Like [communicate], but maps an [ErrorResponse] of type [absent] to `null`
+   * so callers can treat "not found" as a value instead of an exception.
+   */
+  @Throws(ApiError::class)
+  suspend inline fun <reified T> communicateOrNullIf(
+    request: Any,
+    absent: ErrorResponse.ErrorType,
+  ): T? {
+    when (val response = communicateImpl(request)) {
+      is ErrorResponse -> if (response.error == absent) return null else throw 
ApiError("[${response.error}] ${response.detail}")
+      is T -> return response
+      else -> throw ApiError("Unexpected response type 
${response::class.java}")
+    }
+  }
+
   /**
    * Stop the dispatcher and fail anything still awaiting a response.
    */
diff --git 
a/java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/execution/Frame.kt 
b/java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/execution/Frame.kt
index ca61ddd28c2..9cd432692e2 100644
--- a/java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/execution/Frame.kt
+++ b/java-sdk/sdk/src/main/kotlin/org/apache/airflow/sdk/execution/Frame.kt
@@ -27,6 +27,7 @@ import com.fasterxml.jackson.databind.util.StdDateFormat
 import com.fasterxml.jackson.datatype.jsr310.JavaTimeModule
 import org.apache.airflow.sdk.execution.comm.Discriminator
 import org.apache.airflow.sdk.execution.comm.PutVariable
+import org.apache.airflow.sdk.execution.comm.SetTaskStateStore
 import org.msgpack.core.MessagePack
 import org.msgpack.core.MessageUnpacker
 import org.msgpack.core.buffer.ArrayBufferInput
@@ -45,7 +46,7 @@ data class RawFrame(
  * fields to be present. When one is missing it fails validation, logs the
  * frame and never replies, so the caller blocks forever.
  */
-private val REQUIRED_NULLABLE_REQUESTS = setOf(PutVariable::class.java)
+private val REQUIRED_NULLABLE_REQUESTS = setOf(PutVariable::class.java, 
SetTaskStateStore::class.java)
 
 @JsonInclude(JsonInclude.Include.ALWAYS)
 private abstract class KeepNullFields
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 6733621ab0b..ac9e8507c12 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
@@ -22,9 +22,12 @@ package org.apache.airflow.sdk
 import org.apache.airflow.sdk.execution.comm.ConnectionResult
 import org.apache.airflow.sdk.execution.comm.StartupDetails
 import org.apache.airflow.sdk.execution.comm.TIRunContext
+import org.apache.airflow.sdk.execution.comm.TaskStateStoreResult
 import org.apache.airflow.sdk.execution.comm.VariableResult
 import org.apache.airflow.sdk.execution.comm.XComResult
 import org.apache.airflow.sdk.internal.Refs
+import java.time.OffsetDateTime
+import java.util.UUID
 import org.apache.airflow.sdk.execution.comm.TaskInstance as CommTaskInstance
 
 /** Records getXCom calls and serves canned values keyed by task id. */
@@ -68,6 +71,25 @@ internal class FakeXComTransport(
     runId: String,
     mapIndex: Int,
   ) = throw NotImplementedError()
+
+  override fun getTaskStateStore(
+    tiId: UUID,
+    key: String,
+  ): TaskStateStoreResult? = throw NotImplementedError()
+
+  override fun setTaskStateStore(
+    tiId: UUID,
+    key: String,
+    value: Any,
+    expiresAt: OffsetDateTime?,
+  ) = throw NotImplementedError()
+
+  override fun deleteTaskStateStore(
+    tiId: UUID,
+    key: String,
+  ) = throw NotImplementedError()
+
+  override fun clearTaskStateStore(tiId: UUID) = throw NotImplementedError()
 }
 
 internal fun startupDetails(argBindings: List<Map<String, Any?>>?): 
StartupDetails =
diff --git a/java-sdk/sdk/src/test/kotlin/org/apache/airflow/sdk/ClientTest.kt 
b/java-sdk/sdk/src/test/kotlin/org/apache/airflow/sdk/ClientTest.kt
index bd39991abb2..dce06ca0977 100644
--- a/java-sdk/sdk/src/test/kotlin/org/apache/airflow/sdk/ClientTest.kt
+++ b/java-sdk/sdk/src/test/kotlin/org/apache/airflow/sdk/ClientTest.kt
@@ -21,15 +21,32 @@ package org.apache.airflow.sdk
 
 import org.apache.airflow.sdk.execution.comm.ConnectionResult
 import org.apache.airflow.sdk.execution.comm.StartupDetails
+import org.apache.airflow.sdk.execution.comm.TaskInstance
+import org.apache.airflow.sdk.execution.comm.TaskStateStoreResult
 import org.apache.airflow.sdk.execution.comm.VariableResult
 import org.apache.airflow.sdk.execution.comm.XComResult
 import org.junit.jupiter.api.Assertions
 import org.junit.jupiter.api.DisplayName
 import org.junit.jupiter.api.Test
+import java.time.Duration
+import java.time.OffsetDateTime
+import java.time.ZoneOffset
+import java.util.UUID
+
+private data class StateStoreCall(
+  val method: String,
+  val tiId: UUID,
+  val key: String? = null,
+  val value: Any? = null,
+  val expiresAt: OffsetDateTime? = null,
+)
 
 private class FakeTransport(
-  val connection: ConnectionResult,
+  val connection: ConnectionResult = ConnectionResult(),
+  val stored: TaskStateStoreResult? = null,
 ) : org.apache.airflow.sdk.execution.Client {
+  val calls = mutableListOf<StateStoreCall>()
+
   override fun getConnection(id: String): ConnectionResult = connection
 
   override fun getVariable(key: String): VariableResult = throw 
NotImplementedError()
@@ -59,11 +76,63 @@ private class FakeTransport(
     runId: String,
     mapIndex: Int,
   ) = throw NotImplementedError()
+
+  override fun getTaskStateStore(
+    tiId: UUID,
+    key: String,
+  ): TaskStateStoreResult? {
+    calls.add(StateStoreCall("get", tiId, key))
+    return stored
+  }
+
+  override fun setTaskStateStore(
+    tiId: UUID,
+    key: String,
+    value: Any,
+    expiresAt: OffsetDateTime?,
+  ) {
+    calls.add(StateStoreCall("set", tiId, key, value, expiresAt))
+  }
+
+  override fun deleteTaskStateStore(
+    tiId: UUID,
+    key: String,
+  ) {
+    calls.add(StateStoreCall("delete", tiId, key))
+  }
+
+  override fun clearTaskStateStore(tiId: UUID) {
+    calls.add(StateStoreCall("clear", tiId))
+  }
 }
 
 class ClientTest {
   private fun clientWith(connection: ConnectionResult) = 
Client(StartupDetails(), FakeTransport(connection))
 
+  private val tiId: UUID = UUID.randomUUID()
+
+  private fun startupDetails() = StartupDetails().also { it.ti = 
TaskInstance().also { ti -> ti.id = tiId } }
+
+  private fun stateStoreClient(
+    stored: TaskStateStoreResult? = null,
+    env: Map<String, String> = emptyMap(),
+  ): Pair<Client, FakeTransport> {
+    val transport = FakeTransport(stored = stored)
+    return Client(startupDetails(), transport) { env[it] } to transport
+  }
+
+  private fun assertExpiresAbout(
+    expected: Duration,
+    before: OffsetDateTime,
+    after: OffsetDateTime,
+    expiresAt: OffsetDateTime?,
+  ) {
+    Assertions.assertNotNull(expiresAt)
+    Assertions.assertEquals(ZoneOffset.UTC, expiresAt!!.offset)
+    Assertions.assertFalse(expiresAt.isBefore(before.plus(expected)), 
"expiresAt $expiresAt before $before + $expected")
+    Assertions.assertFalse(expiresAt.isAfter(after.plus(expected)), "expiresAt 
$expiresAt after $after + $expected")
+  }
+
   @Test
   @DisplayName("Should convert a Long port from the wire to Int")
   fun shouldConvertLongPort() {
@@ -98,6 +167,119 @@ class ClientTest {
     Assertions.assertNull(connection.port)
   }
 
+  @Test
+  @DisplayName("taskStateStore is exposed to Java as a getter so mocking 
frameworks can stub it")
+  fun taskStateStoreIsExposedAsGetter() {
+    val getter = Client::class.java.getMethod("getTaskStateStore")
+
+    Assertions.assertEquals(TaskStateStore::class.java, getter.returnType)
+    Assertions.assertTrue(Client::class.java.fields.none { it.name == 
"taskStateStore" })
+  }
+
+  @Test
+  @DisplayName("taskStateStore.get returns the stored value scoped to the 
current task instance")
+  fun taskStateStoreGetReturnsStoredValue() {
+    val (client, transport) = stateStoreClient(TaskStateStoreResult().apply { 
value = "job-42" })
+
+    Assertions.assertEquals("job-42", client.taskStateStore.get("job_id"))
+    Assertions.assertEquals(listOf(StateStoreCall("get", tiId, "job_id")), 
transport.calls)
+  }
+
+  @Test
+  @DisplayName("taskStateStore.get returns null when the key is not stored")
+  fun taskStateStoreGetReturnsNullWhenMissing() {
+    val (client, _) = stateStoreClient(stored = null)
+
+    Assertions.assertNull(client.taskStateStore.get("job_id"))
+  }
+
+  @Test
+  @DisplayName("taskStateStore.set without retention uses 
AIRFLOW__STATE_STORE__DEFAULT_RETENTION_DAYS")
+  fun taskStateStoreSetWithoutRetentionUsesDefaultRetentionDays() {
+    listOf("7" to 7L, "7.0" to 7L, " 2 " to 2L).forEach { (raw, days) ->
+      val (client, transport) = stateStoreClient(env = 
mapOf(TaskStateStore.DEFAULT_RETENTION_DAYS_ENV to raw))
+
+      val before = OffsetDateTime.now(ZoneOffset.UTC)
+      client.taskStateStore.set("job_id", 42)
+      val after = OffsetDateTime.now(ZoneOffset.UTC)
+
+      assertExpiresAbout(Duration.ofDays(days), before, after, 
transport.calls.single().expiresAt)
+    }
+  }
+
+  @Test
+  @DisplayName("taskStateStore.set without retention fails when the 
coordinator did not pass the default")
+  fun taskStateStoreSetWithoutRetentionFailsWhenVariableAbsent() {
+    val (client, transport) = stateStoreClient()
+
+    val error = Assertions.assertThrows(IllegalStateException::class.java) { 
client.taskStateStore.set("job_id", 42) }
+
+    
Assertions.assertTrue(error.message!!.startsWith(TaskStateStore.DEFAULT_RETENTION_DAYS_ENV),
 error.message)
+    Assertions.assertTrue(transport.calls.isEmpty(), "no call expected: 
${transport.calls}")
+  }
+
+  @Test
+  @DisplayName("taskStateStore.set never expires for a default of 0 days or an 
explicit NEVER_EXPIRE")
+  fun taskStateStoreSetNeverExpires() {
+    val (byConfig, configTransport) = stateStoreClient(env = 
mapOf(TaskStateStore.DEFAULT_RETENTION_DAYS_ENV to "0"))
+    val (byArgument, argumentTransport) = stateStoreClient()
+
+    byConfig.taskStateStore.set("job_id", 42)
+    byArgument.taskStateStore.set("job_id", 42, TaskStateStore.NEVER_EXPIRE)
+
+    Assertions.assertEquals(listOf(StateStoreCall("set", tiId, "job_id", 42, 
expiresAt = null)), configTransport.calls)
+    Assertions.assertEquals(listOf(StateStoreCall("set", tiId, "job_id", 42, 
expiresAt = null)), argumentTransport.calls)
+  }
+
+  @Test
+  @DisplayName("taskStateStore.set with retention expires that long after now, 
in UTC")
+  fun taskStateStoreSetWithRetentionComputesExpiry() {
+    val (client, transport) = stateStoreClient()
+    val retention = Duration.ofHours(6)
+
+    val before = OffsetDateTime.now(ZoneOffset.UTC)
+    client.taskStateStore.set("job_id", 42, retention)
+    val after = OffsetDateTime.now(ZoneOffset.UTC)
+
+    assertExpiresAbout(retention, before, after, 
transport.calls.single().expiresAt)
+  }
+
+  @Test
+  @DisplayName("taskStateStore.set rejects a zero or negative retention 
without calling the supervisor")
+  fun taskStateStoreSetRejectsNonPositiveRetention() {
+    listOf(Duration.ZERO, Duration.ofSeconds(-1)).forEach { retention ->
+      val (client, transport) = stateStoreClient()
+
+      Assertions.assertThrows(IllegalArgumentException::class.java) { 
client.taskStateStore.set("job_id", 42, retention) }
+      Assertions.assertTrue(transport.calls.isEmpty(), "no call expected for 
$retention: ${transport.calls}")
+    }
+  }
+
+  @Test
+  @DisplayName("taskStateStore.set rejects a default retention that is not a 
non-negative integer")
+  fun taskStateStoreSetRejectsBadDefaultRetention() {
+    listOf("abc", "-1", "1.5", "").forEach { raw ->
+      val (client, transport) = stateStoreClient(env = 
mapOf(TaskStateStore.DEFAULT_RETENTION_DAYS_ENV to raw))
+
+      Assertions.assertThrows(IllegalArgumentException::class.java) { 
client.taskStateStore.set("job_id", 42) }
+      Assertions.assertTrue(transport.calls.isEmpty(), "no call expected for 
'$raw': ${transport.calls}")
+    }
+  }
+
+  @Test
+  @DisplayName("taskStateStore.delete and clear address the current task 
instance")
+  fun taskStateStoreDeleteAndClearUseCurrentTaskInstance() {
+    val (client, transport) = stateStoreClient()
+
+    client.taskStateStore.delete("job_id")
+    client.taskStateStore.clear()
+
+    Assertions.assertEquals(
+      listOf(StateStoreCall("delete", tiId, "job_id"), StateStoreCall("clear", 
tiId)),
+      transport.calls,
+    )
+  }
+
   @Test
   @DisplayName("MissingXComException builds the full message naming the task 
and parameter")
   fun missingXComExceptionBuildsFullMessage() {
diff --git 
a/java-sdk/sdk/src/test/kotlin/org/apache/airflow/sdk/execution/CommTest.kt 
b/java-sdk/sdk/src/test/kotlin/org/apache/airflow/sdk/execution/CommTest.kt
index 0c2ef844f6c..7f161ed9b5c 100644
--- a/java-sdk/sdk/src/test/kotlin/org/apache/airflow/sdk/execution/CommTest.kt
+++ b/java-sdk/sdk/src/test/kotlin/org/apache/airflow/sdk/execution/CommTest.kt
@@ -26,6 +26,7 @@ import kotlinx.coroutines.async
 import kotlinx.coroutines.runBlocking
 import kotlinx.coroutines.supervisorScope
 import org.apache.airflow.sdk.ApiError
+import org.apache.airflow.sdk.TaskStateStore
 import org.apache.airflow.sdk.execution.comm.GetVariable
 import org.apache.airflow.sdk.execution.comm.StartupDetails
 import org.apache.airflow.sdk.execution.comm.TaskInstance
@@ -37,8 +38,10 @@ import org.junit.jupiter.api.Timeout
 import org.msgpack.core.MessagePack
 import org.msgpack.core.buffer.ArrayBufferInput
 import java.io.ByteArrayOutputStream
+import java.time.Duration
 import java.time.OffsetDateTime
 import java.time.ZoneOffset
+import java.util.UUID
 import java.util.concurrent.ConcurrentLinkedQueue
 import java.util.concurrent.TimeUnit
 import org.apache.airflow.sdk.Client as PublicClient
@@ -179,6 +182,40 @@ class CommsTest {
     return out.toByteArray()
   }
 
+  private fun notFoundResponseFrame(id: Int): ByteArray {
+    val out = ByteArrayOutputStream()
+    MessagePack.newDefaultPacker(out).use { packer ->
+      packer.packArrayHeader(3)
+      packer.packInt(id)
+      packer.packNil()
+      packer.packMapHeader(3)
+      packer.packString("type")
+      packer.packString("ErrorResponse")
+      packer.packString("error")
+      packer.packString("TASK_STORE_NOT_FOUND")
+      packer.packString("detail")
+      packer.packMapHeader(1)
+      packer.packString("key")
+      packer.packString("job_id")
+    }
+    return out.toByteArray()
+  }
+
+  private fun taskStateStoreResultFrame(id: Int): ByteArray {
+    val out = ByteArrayOutputStream()
+    MessagePack.newDefaultPacker(out).use { packer ->
+      packer.packArrayHeader(3)
+      packer.packInt(id)
+      packer.packMapHeader(2)
+      packer.packString("type")
+      packer.packString("TaskStateStoreResult")
+      packer.packString("value")
+      packer.packString("job-42")
+      packer.packNil()
+    }
+    return out.toByteArray()
+  }
+
   private fun readRequest(fromClient: ByteChannel): RawFrame =
     runBlocking {
       val prefix = fromClient.readByteArray(4)
@@ -193,12 +230,13 @@ class CommsTest {
    */
   private fun roundTrip(
     response: (Int) -> ByteArray,
+    details: StartupDetails = StartupDetails(),
     call: (PublicClient) -> Unit,
   ): Pair<Map<*, *>, Throwable?> {
     val toClient = ByteChannel(autoFlush = true)
     val fromClient = ByteChannel(autoFlush = true)
     val comm = CoordinatorComm(toClient, fromClient)
-    val client = PublicClient(StartupDetails(), CoordinatorClient(comm))
+    val client = PublicClient(details, CoordinatorClient(comm))
 
     val requests = ConcurrentLinkedQueue<RawFrame>()
     val server =
@@ -386,6 +424,90 @@ class CommsTest {
     Assertions.assertInstanceOf(ApiError::class.java, failure)
   }
 
+  private val tiId: UUID = 
UUID.fromString("0199a5d6-1c2e-7c6a-9c1e-7a2f7f0d1e42")
+
+  private fun currentTaskInstance() = StartupDetails().also { it.ti = 
TaskInstance().also { ti -> ti.id = tiId } }
+
+  @Test
+  @DisplayName("taskStateStore.get sends the task instance ID and key and 
unwraps the stored value")
+  @Timeout(value = 30, unit = TimeUnit.SECONDS)
+  fun taskStateStoreGetUnwrapsResult() {
+    var value: Any? = null
+    val (body, failure) = roundTrip(::taskStateStoreResultFrame, 
currentTaskInstance()) { value = it.taskStateStore.get("job_id") }
+
+    Assertions.assertNull(failure, "get should return normally on 
TaskStateStoreResult, got $failure")
+    Assertions.assertEquals("GetTaskStateStore", body["type"])
+    Assertions.assertEquals(tiId.toString(), body["ti_id"])
+    Assertions.assertEquals("job_id", body["key"])
+    Assertions.assertEquals("job-42", value)
+  }
+
+  @Test
+  @DisplayName("taskStateStore.get returns null on TASK_STORE_NOT_FOUND 
instead of raising")
+  @Timeout(value = 30, unit = TimeUnit.SECONDS)
+  fun taskStateStoreGetReturnsNullWhenNotFound() {
+    var value: Any? = "unset"
+    val (_, failure) = roundTrip(::notFoundResponseFrame, 
currentTaskInstance()) { value = it.taskStateStore.get("job_id") }
+
+    Assertions.assertNull(failure, "get should return normally on 
TASK_STORE_NOT_FOUND, got $failure")
+    Assertions.assertNull(value)
+  }
+
+  @Test
+  @DisplayName("taskStateStore.get raises ApiError for any other error 
response")
+  @Timeout(value = 30, unit = TimeUnit.SECONDS)
+  fun taskStateStoreGetRaisesApiErrorOnOtherErrors() {
+    val (_, failure) = roundTrip(::errorResponseFrame, currentTaskInstance()) 
{ it.taskStateStore.get("job_id") }
+
+    Assertions.assertInstanceOf(ApiError::class.java, failure)
+  }
+
+  @Test
+  @DisplayName("taskStateStore.set keeps a null expires_at on the wire so the 
supervisor accepts the request")
+  @Timeout(value = 30, unit = TimeUnit.SECONDS)
+  fun taskStateStoreSetKeepsNullExpiresAtOnTheWire() {
+    val (body, failure) =
+      roundTrip(::okResponseFrame, currentTaskInstance()) { 
it.taskStateStore.set("job_id", 42, TaskStateStore.NEVER_EXPIRE) }
+
+    Assertions.assertNull(failure, "set should return normally on OKResponse, 
got $failure")
+    Assertions.assertEquals("SetTaskStateStore", body["type"])
+    Assertions.assertEquals(tiId.toString(), body["ti_id"])
+    Assertions.assertEquals("job_id", body["key"])
+    Assertions.assertEquals(42L, body["value"])
+    Assertions.assertTrue(body.containsKey("expires_at"), "expires_at must be 
sent even when null: $body")
+    Assertions.assertNull(body["expires_at"])
+  }
+
+  @Test
+  @DisplayName("taskStateStore.set sends expires_at as an ISO-8601 timestamp 
when a retention is given")
+  @Timeout(value = 30, unit = TimeUnit.SECONDS)
+  fun taskStateStoreSetSendsIsoExpiresAt() {
+    val (body, failure) =
+      roundTrip(::okResponseFrame, currentTaskInstance()) { 
it.taskStateStore.set("job_id", 42, Duration.ofHours(6)) }
+
+    Assertions.assertNull(failure, "set should return normally on OKResponse, 
got $failure")
+    val expiresAt = OffsetDateTime.parse(body["expires_at"] as String)
+    Assertions.assertEquals(ZoneOffset.UTC, expiresAt.offset)
+    
Assertions.assertTrue(expiresAt.isAfter(OffsetDateTime.now(ZoneOffset.UTC).plusHours(5)),
 "expires_at $expiresAt")
+  }
+
+  @Test
+  @DisplayName("taskStateStore.delete and clear send the task instance ID and 
accept the supervisor's OK response")
+  @Timeout(value = 30, unit = TimeUnit.SECONDS)
+  fun taskStateStoreDeleteAndClearAcceptOkResponse() {
+    val (deleteBody, deleteFailure) = roundTrip(::okResponseFrame, 
currentTaskInstance()) { it.taskStateStore.delete("job_id") }
+    val (clearBody, clearFailure) = roundTrip(::okResponseFrame, 
currentTaskInstance()) { it.taskStateStore.clear() }
+
+    Assertions.assertNull(deleteFailure, "delete should return normally on 
OKResponse, got $deleteFailure")
+    Assertions.assertEquals("DeleteTaskStateStore", deleteBody["type"])
+    Assertions.assertEquals(tiId.toString(), deleteBody["ti_id"])
+    Assertions.assertEquals("job_id", deleteBody["key"])
+
+    Assertions.assertNull(clearFailure, "clear should return normally on 
OKResponse, got $clearFailure")
+    Assertions.assertEquals("ClearTaskStateStore", clearBody["type"])
+    Assertions.assertEquals(tiId.toString(), clearBody["ti_id"])
+  }
+
   @Test
   @DisplayName("Should fail a pending call when the coordinator socket closes")
   @Timeout(value = 30, unit = TimeUnit.SECONDS)
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 548a099bf81..5ce508b708d 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
@@ -252,6 +252,25 @@ class TaskTest {
           runId: String,
           mapIndex: Int,
         ): Unit = throw UnsupportedOperationException("not used in test")
+
+        override fun getTaskStateStore(
+          tiId: UUID,
+          key: String,
+        ) = throw UnsupportedOperationException("not used in test")
+
+        override fun setTaskStateStore(
+          tiId: UUID,
+          key: String,
+          value: Any,
+          expiresAt: OffsetDateTime?,
+        ): Unit = throw UnsupportedOperationException("not used in test")
+
+        override fun deleteTaskStateStore(
+          tiId: UUID,
+          key: String,
+        ): Unit = throw UnsupportedOperationException("not used in test")
+
+        override fun clearTaskStateStore(tiId: UUID): Unit = throw 
UnsupportedOperationException("not used in test")
       },
     )
 
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 7d8439634aa..9acaa23f6ce 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
@@ -34,6 +34,7 @@ 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.TaskStateStoreResult
 import org.apache.airflow.sdk.execution.comm.VariableResult
 import org.apache.airflow.sdk.execution.comm.XComResult
 import org.junit.jupiter.api.Assertions.assertEquals
@@ -41,6 +42,8 @@ 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 java.time.OffsetDateTime
+import java.util.UUID
 import org.apache.airflow.sdk.execution.Client as Transport
 import org.apache.airflow.sdk.execution.comm.TaskInstance as CommTaskInstance
 
@@ -106,6 +109,25 @@ internal class ArgValuesTest {
           runId: String,
           mapIndex: Int,
         ): Unit = throw NotImplementedError()
+
+        override fun getTaskStateStore(
+          tiId: UUID,
+          key: String,
+        ): TaskStateStoreResult? = throw NotImplementedError()
+
+        override fun setTaskStateStore(
+          tiId: UUID,
+          key: String,
+          value: Any,
+          expiresAt: OffsetDateTime?,
+        ): Unit = throw NotImplementedError()
+
+        override fun deleteTaskStateStore(
+          tiId: UUID,
+          key: String,
+        ): Unit = throw NotImplementedError()
+
+        override fun clearTaskStateStore(tiId: UUID): Unit = throw 
NotImplementedError()
       },
     )
 

Reply via email to