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