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 8bb642be2d5 Add task state store support to the Go SDK task client
(#73420)
8bb642be2d5 is described below
commit 8bb642be2d500479fb0a9718de1837abb2777ecf
Author: Cheng-Yi Shih <[email protected]>
AuthorDate: Tue Oct 6 23:36:32 2026 +0800
Add task state store support to the Go SDK task client (#73420)
* Add task state store support to the Go SDK task client
A Go task that submits a long-running external job had no way to remember
it across an attempt boundary, so a worker crash meant the retry submitted
the job a second time - duplicate work and duplicate cost. Python has
exposed this store since 3.3 as context["task_state_store"], and the
Language SDK specification lists task-state-store as a capability the Go
SDK still declared unsupported.
The coordinator protocol was already in place - the supervisor forwards the
four store messages and their Go models are generated from the schema
snapshot - so only the task-facing surface was missing.
Retention needs the deployment's [state_store] default_retention_days, which
a language SDK runtime cannot read for itself. The supervisor now resolves
it
and passes it in the environment at launch, the same way it already does for
the task log levels, so a key written from Go expires exactly when one
written from Python does instead of silently outliving it. A deployment that
misconfigures that value fails the write rather than quietly substituting a
different lifetime, and a value the store cannot hold is refused before it
is
sent; both mirror what the Python accessor does, so the two SDKs fail the
same way on the same input.
* Validate task state values against the encoder's actual output
The value check mirrored the encoder's field rules by walking the Go value
with reflection, and it drifted from them: a byte array slipped through and
reached the supervisor as bytes, while a struct field tagged json:"-" was
rejected even though the encoder never sends it. Checking what the frame
encoder actually emits makes the check agree with the wire by construction,
and sharing one encoder constructor keeps the two from drifting apart.
The e2e retention test could not fail: the Go fallback and the config
default
are both 30 days, so a missing propagation looked identical to a working
one.
The retention now comes from airflow.cfg at 7 days. It cannot be set as an
environment variable, because the supervisor copies the whole worker
environment into the Go subprocess, which would carry the value across even
without the propagation under test.
A unit test that passed with or without the fallback it claimed to cover is
removed, and comments are trimmed to the reasons behind the code.
* Reject typed nil task state values and pin the sent expiry in tests
A nil pointer stored in an interface is not equal to nil, so the old check
let it through; it still encodes to null, which the Execution API rejects.
Checking the decoded value catches every form of nil the encoder emits.
The retention tests only checked that an expiry was sent, so a wrong
retention would still pass. They now require the expiry to fall within the
retention added to the time around the call.
* Fall back to the default retention only when the variable is absent
The supervisor reads the retention with a config fallback that applies only
to a missing key, so an empty setting reaches the runtime as an empty value.
Python raises on it while Go quietly retained keys for the fallback period,
which is the drift the comment here already warned against.
* Group the task state store behind one store interface
The flat methods made the store the odd one out: Python reaches it as
context["task_state_store"], and the Java SDK's own PR groups it as
client.taskStateStore, while Go spelled the resource in every method name.
Retention becomes an option rather than a second method, which is how the
other SDKs spell it - a keyword argument in Python, a default parameter in
Java - and leaves one Set for callers to find.
The decode helper cannot be called UnmarshalJSON: go vet enforces
encoding/json's signature on any method of that name.
The 30-day fallback is gone. The supervisor always resolves the deployment's
retention, so a missing or unusable value is a misconfiguration; guessing a
lifetime nobody configured is how the e2e test came to pass without the
propagation it was meant to prove. A day count that wraps the timestamp is
rejected too, because it would store a key that is already expired.
---
.../authoring-and-scheduling/language-sdks/go.rst | 99 +++-
.../tests/airflow_e2e_tests/conftest.py | 7 +
.../tests/airflow_e2e_tests/constants.py | 3 +
.../airflow_e2e_tests/e2e_test_utils/clients.py | 7 +
.../go_sdk_tests/test_go_sdk_task_state.py | 114 ++++
go-sdk/README.md | 2 +-
go-sdk/capabilities.yaml | 4 +-
.../cmd/airflow-go-pack/pack_integration_test.go | 3 +
go-sdk/dags/go_examples.py | 19 +-
.../bundle/concurrentxcom/concurrentxcom_test.go | 4 +
go-sdk/example/bundle/main.go | 3 +
go-sdk/example/bundle/taskstate/taskstate.go | 99 ++++
go-sdk/pkg/execution/client.go | 277 +++++++++-
go-sdk/pkg/execution/client_test.go | 583 ++++++++++++++++++++-
go-sdk/pkg/execution/frames.go | 14 +-
go-sdk/pkg/execution/integration_test.go | 49 ++
go-sdk/pkg/execution/task_runner.go | 2 +-
go-sdk/sdk/doc.go | 8 +-
go-sdk/sdk/errors.go | 6 +
go-sdk/sdk/sdk.go | 89 +++-
.../src/airflow/sdk/coordinators/_subprocess.py | 1 +
.../tests/task_sdk/coordinators/test_subprocess.py | 3 +
22 files changed, 1362 insertions(+), 34 deletions(-)
diff --git a/airflow-core/docs/authoring-and-scheduling/language-sdks/go.rst
b/airflow-core/docs/authoring-and-scheduling/language-sdks/go.rst
index 25a5e3c14f6..7648c141531 100644
--- a/airflow-core/docs/authoring-and-scheduling/language-sdks/go.rst
+++ b/airflow-core/docs/authoring-and-scheduling/language-sdks/go.rst
@@ -273,7 +273,7 @@ and lets a test pass a fake.
The ``sdk.Client`` surface
~~~~~~~~~~~~~~~~~~~~~~~~~~~~
-``actx.Client()`` returns an ``sdk.Client``, which composes three smaller
interfaces, so a helper can depend
+``actx.Client()`` returns an ``sdk.Client``, which composes four smaller
interfaces, so a helper can depend
on just one:
* ``VariableClient`` - ``GetVariable`` (returns the Variable as a string),
``UnmarshalJSONVariable``
@@ -282,6 +282,9 @@ on just one:
``Host``, ``Port``, ``Login``, ``Password``, ``Path``, ``Extra`` (a
``map[string]any``), plus a
``GetURI()`` helper.
* ``XComClient`` - ``GetXCom`` to read an upstream task's XCom and
``PushXCom`` to publish one.
+* ``TaskStateStoreClient`` - ``TaskStateStore``, returning the store for this
task instance:
+ ``Get``, ``UnmarshalJSONValue`` (decodes a JSON value into a pointer you
provide), ``Set``,
+ ``Delete``, and ``Clear``. See :ref:`go-sdk/task-state-store`.
``GetXCom`` returns the stored value as an ``any``; see :ref:`go-sdk/types`
for how the stored JSON maps to
Go types.
@@ -305,8 +308,98 @@ before storing it.
precedence over the stored value when the Variable is read back. Calling
``SetVariable`` with an empty
description clears any existing description.
-Not-found lookups return sentinel errors - ``VariableNotFound``,
``ConnectionNotFound``, ``XComNotFound`` -
-so you can branch on a missing value with ``errors.Is`` rather than parsing an
error string.
+Not-found lookups return sentinel errors - ``VariableNotFound``,
``ConnectionNotFound``, ``XComNotFound``,
+``TaskStateNotFound`` - so you can branch on a missing value with
``errors.Is`` rather than parsing an error
+string.
+
+.. _go-sdk/task-state-store:
+
+The task state store
+~~~~~~~~~~~~~~~~~~~~~~
+
+``actx.Client().TaskStateStore()`` returns a persistent key/value store
private to one task instance,
+and the Go SDK's entry point to durable execution. It is the same store the
Python SDK exposes as
+``context["task_state_store"]``; see :doc:`/core-concepts/task-state-store`
for the concept and its
+configuration.
+
+The store is scoped to ``dag_id``, ``run_id``, ``task_id``, and ``map_index``.
It deliberately does *not*
+include ``try_number``, so a value written by one attempt is still readable by
the next one: a task that
+records an external job ID or its own progress can resume after a worker crash
or a retry instead of
+redoing the work. The Execution API confines every call to the task instance
the caller is running as, so
+there is no way to address another task's store - pass results between tasks
with XCom instead.
+
+The usual shape is to look for a checkpoint first and only do the expensive
work when it is missing:
+
+.. code-block:: go
+
+ import (
+ "errors"
+
+ "github.com/apache/airflow/go-sdk/airflow"
+ "github.com/apache/airflow/go-sdk/sdk"
+ )
+
+ func runSparkJob(actx airflow.Context) error {
+ store := actx.Client().TaskStateStore()
+
+ var jobID string
+ stored, err := store.Get(actx, "job_id")
+ switch {
+ case errors.Is(err, sdk.TaskStateNotFound):
+ // First attempt: submit the job and remember its ID before doing
anything else.
+ if jobID, err = sparkClient.SubmitJob(actx); err != nil {
+ return err
+ }
+ if err := store.Set(actx, "job_id", jobID,
sdk.WithRetention(sdk.NeverExpire)); err != nil {
+ return err
+ }
+ case err != nil:
+ return err
+ default:
+ // Get returns an any; the value was stored by this task as a
string.
+ jobID = stored.(string)
+ actx.Logger().InfoContext(actx, "reattaching to job submitted by
an earlier attempt", "job_id", jobID)
+ }
+
+ return sparkClient.WaitForCompletion(actx, jobID)
+ }
+
+``value`` must not be nil and must be JSON-representable - a string, number,
bool, slice, map, or a struct,
+which is stored as an object built from its exported fields and their ``json``
tags. A custom
+``MarshalJSON`` is not called, so a type that relies on one is stored as the
shape of its fields, and a
+struct with no exported fields is stored as ``{}``. Read a scalar back with
``Get``, which returns it as an ``any`` (the
+numeric caveat in :ref:`go-sdk/types` applies here too); for an object or
array,
+``UnmarshalJSONValue`` decodes it straight into a pointer you provide.
+
+A value the store cannot hold is rejected before it is sent, so you get an
error naming the problem rather
+than a round trip that fails on the server. The one that catches people out is
``time.Time``, which JSON has
+no spelling for - store ``value.Format(time.RFC3339)`` and parse it back with
``time.Parse``. Non-finite
+floats and ``[]byte`` are refused for the same reason. This mirrors the Python
SDK, where the same values
+fail Pydantic validation before the write leaves the worker.
+
+Keys expire, so retention is part of writing a value:
+
+* ``Set`` without options uses the deployment's ``[state_store]
default_retention_days`` (30 days by
+ default). The Go runtime cannot read Airflow's config, so the supervisor
resolves that value and passes
+ it in the environment when it launches the bundle. A deployment that sets it
to something unusable - a
+ negative number, or a value that is not a whole number of days - fails the
write, exactly as it does for
+ a Python task, rather than quietly substituting a different lifetime.
+* ``sdk.WithRetention`` takes an explicit, positive ``time.Duration``, or
``sdk.NeverExpire`` for a key
+ that is skipped by garbage collection entirely. A zero or negative retention
is rejected rather than
+ given a meaning of its own: to follow the deployment default omit the
option, and to drop a key call
+ ``Delete``.
+
+``Delete`` removes one key (deleting a key that does not exist is not an
error) and ``Clear`` removes
+every key stored for this task instance.
+
+.. note::
+
+ The Go SDK does not implement the worker-side state backend (``[workers]
state_store_backend``), which
+ offloads large values to external storage and records only a reference
marker in the database. If a
+ deployment configures one, a Go task reading a key that was written through
that backend receives the raw
+ reference marker rather than the original value, and a Go task writing a key
stores the whole value in the
+ database instead of offloading it. This is the same behaviour as a Python
worker that does not have the
+ backend configured.
.. _go-sdk/runtime-context:
diff --git a/airflow-e2e-tests/tests/airflow_e2e_tests/conftest.py
b/airflow-e2e-tests/tests/airflow_e2e_tests/conftest.py
index 1282326cc9a..9814cb5147f 100644
--- a/airflow-e2e-tests/tests/airflow_e2e_tests/conftest.py
+++ b/airflow-e2e-tests/tests/airflow_e2e_tests/conftest.py
@@ -45,6 +45,7 @@ from airflow_e2e_tests.constants import (
GO_SDK_DAGS_PATH,
GO_SDK_EXAMPLE_BUNDLE_PKG,
GO_SDK_ROOT_PATH,
+ GO_SDK_STATE_STORE_RETENTION_DAYS,
JAVA_COMPOSE_PATH,
JAVA_DOCKERFILE_PATH,
JAVA_SDK_EXAMPLE_DAGS_PATH,
@@ -673,6 +674,12 @@ def _setup_go_sdk_integration(dot_env_file, tmp_dir):
)
os.environ["ENV_FILE_PATH"] = str(dot_env_file)
+ # Config file, not an env var: the supervisor copies the worker
environment into the Go
+ # subprocess, so an env var would reach Go even without the propagation
under test.
+ (tmp_dir / "config" / "airflow.cfg").write_text(
+ f"[state_store]\ndefault_retention_days =
{GO_SDK_STATE_STORE_RETENTION_DAYS}\n"
+ )
+
def _setup_openlineage_integration(dot_env_file, tmp_dir, compose_file_names):
"""Set up the openlineage E2E test mode.
diff --git a/airflow-e2e-tests/tests/airflow_e2e_tests/constants.py
b/airflow-e2e-tests/tests/airflow_e2e_tests/constants.py
index 4a1a2f7be38..47c698c5d0a 100644
--- a/airflow-e2e-tests/tests/airflow_e2e_tests/constants.py
+++ b/airflow-e2e-tests/tests/airflow_e2e_tests/constants.py
@@ -95,6 +95,9 @@ GO_SDK_BUNDLE_NAME = "example_dags"
# Where airflow-go-pack writes the packed bundle inside the repo (go-sdk/bin
is gitignored).
GO_SDK_BIN_PATH = GO_SDK_ROOT_PATH / "bin"
GO_COMPOSE_PATH = AIRFLOW_ROOT_PATH / "airflow-e2e-tests" / "docker" / "go.yml"
+# Far from the supervisor's own 30-day config fallback, so the e2e test can
tell a
+# propagated value from it.
+GO_SDK_STATE_STORE_RETENTION_DAYS = 7
# Go toolchain image used to build the bundle in the containerized path (i.e.
unless
# LANG_SDK_NATIVE_TOOLCHAIN is set); must satisfy go-sdk/go.mod's toolchain.
# The Alpine variant is ~7x smaller than the Debian one and is safe here
because the
diff --git
a/airflow-e2e-tests/tests/airflow_e2e_tests/e2e_test_utils/clients.py
b/airflow-e2e-tests/tests/airflow_e2e_tests/e2e_test_utils/clients.py
index d56cecae3d7..f555cd11542 100644
--- a/airflow-e2e-tests/tests/airflow_e2e_tests/e2e_test_utils/clients.py
+++ b/airflow-e2e-tests/tests/airflow_e2e_tests/e2e_test_utils/clients.py
@@ -163,6 +163,13 @@ class AirflowClient:
"""Get an Airflow Variable via API."""
return self._make_request(method="GET", endpoint=f"variables/{key}")
+ def get_task_state_store(self, dag_id: str, run_id: str, task_id: str,
key: str):
+ """Get a single task state store entry via API."""
+ return self._make_request(
+ method="GET",
+
endpoint=f"dags/{dag_id}/dagRuns/{run_id}/taskInstances/{task_id}/state-store/{key}",
+ )
+
def trigger_dag_and_wait(self, dag_id: str, json=None):
"""Trigger a DAG and wait for it to complete."""
self.un_pause_dag(dag_id)
diff --git
a/airflow-e2e-tests/tests/airflow_e2e_tests/go_sdk_tests/test_go_sdk_task_state.py
b/airflow-e2e-tests/tests/airflow_e2e_tests/go_sdk_tests/test_go_sdk_task_state.py
new file mode 100644
index 00000000000..37cf349a3f7
--- /dev/null
+++
b/airflow-e2e-tests/tests/airflow_e2e_tests/go_sdk_tests/test_go_sdk_task_state.py
@@ -0,0 +1,114 @@
+# 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.
+"""E2E test for the Go SDK ``task_state_dag`` example
(``go-sdk/example/bundle/taskstate``).
+
+Assertions read the API-visible store rather than the task's XCom summary, so
the data is
+proven to have reached the database.
+"""
+
+from __future__ import annotations
+
+from dataclasses import dataclass
+from datetime import datetime, timedelta, timezone
+from http import HTTPStatus
+
+import pytest
+import requests
+
+from airflow_e2e_tests.constants import GO_SDK_STATE_STORE_RETENTION_DAYS
+from airflow_e2e_tests.e2e_test_utils.clients import AirflowClient
+
+_GO_TASK_TIMEOUT = 300
+
+_DAG_ID = "task_state_dag"
+_TASK_ID = "roundtrip_task_state"
+
+
+@dataclass(frozen=True)
+class _CompletedRun:
+ """The single ``task_state_dag`` run shared across this module's tests."""
+
+ client: AirflowClient
+ run_id: str
+ state: str
+ ti_states: dict[str, str]
+
+ def state_store(self, key: str):
+ return self.client.get_task_state_store(dag_id=_DAG_ID,
run_id=self.run_id, task_id=_TASK_ID, key=key)
+
+
[email protected](scope="module")
+def completed_run() -> _CompletedRun:
+ client = AirflowClient()
+ resp = client.trigger_dag(_DAG_ID, json={"logical_date":
datetime.now(timezone.utc).isoformat()})
+ run_id = resp["dag_run_id"]
+ state = client.wait_for_dag_run(dag_id=_DAG_ID, run_id=run_id,
timeout=_GO_TASK_TIMEOUT)
+ ti_resp = client.get_task_instances(dag_id=_DAG_ID, run_id=run_id)
+ ti_states = {ti["task_id"]: ti.get("state") for ti in
ti_resp.get("task_instances", [])}
+ return _CompletedRun(client=client, run_id=run_id, state=state,
ti_states=ti_states)
+
+
+def _parse_timestamp(value: str) -> datetime:
+ return datetime.fromisoformat(value.replace("Z", "+00:00"))
+
+
+def test_task_succeeded(completed_run: _CompletedRun):
+ assert completed_run.state == "success", (
+ f"expected the run to succeed; got {completed_run.state!r}. task
states: {completed_run.ti_states}"
+ )
+ assert completed_run.ti_states.get(_TASK_ID) == "success",
completed_run.ti_states
+
+
+def test_value_written_by_go_task_is_readable(completed_run: _CompletedRun):
+ entry = completed_run.state_store("go_e2e_run_id")
+ assert entry.get("value") == completed_run.run_id, (
+ f"go_e2e_run_id should hold this run's id {completed_run.run_id!r},
got {entry!r}"
+ )
+
+
+def test_never_expire_key_has_null_expiry(completed_run: _CompletedRun):
+ entry = completed_run.state_store("go_e2e_run_id")
+ # A null expiry on the wire is the end-to-end proof that sdk.NeverExpire
reached the database.
+ assert entry.get("expires_at") is None, entry
+
+
+def test_structured_value_roundtrips(completed_run: _CompletedRun):
+ entry = completed_run.state_store("go_e2e_counter")
+ value = entry.get("value")
+ assert isinstance(value, dict), entry
+ assert value.get("processed") == 3, entry
+ assert isinstance(value.get("processed"), int), entry
+ assert value.get("cursor") == "abc-123", entry
+ assert isinstance(value.get("cursor"), str), entry
+
+
+def test_deleted_key_is_gone(completed_run: _CompletedRun):
+ with pytest.raises(requests.HTTPError) as excinfo:
+ completed_run.state_store("go_e2e_scratch")
+ assert excinfo.value.response.status_code == HTTPStatus.NOT_FOUND
+
+
+def test_default_retention_applied(completed_run: _CompletedRun):
+ entry = completed_run.state_store("go_e2e_retained")
+ assert entry.get("expires_at") is not None, entry
+ gap = _parse_timestamp(entry["expires_at"]) -
_parse_timestamp(entry["updated_at"])
+ expected = timedelta(days=GO_SDK_STATE_STORE_RETENTION_DAYS)
+ assert abs(gap - expected) <= timedelta(minutes=5), (
+ f"expected ~{GO_SDK_STATE_STORE_RETENTION_DAYS} days, got {gap /
timedelta(days=1):.1f} days; "
+ "if ~30, the supervisor fell back to its own config default, so it did
not see "
+ f"[state_store] default_retention_days. entry: {entry!r}"
+ )
diff --git a/go-sdk/README.md b/go-sdk/README.md
index 1ecd423310d..ad344112abe 100644
--- a/go-sdk/README.md
+++ b/go-sdk/README.md
@@ -349,7 +349,7 @@ prek hook regenerate it.
| capability: `variable-read-write` | MUST | ✓ | 3.4 | |
| capability: `self-contained-bundle` | MUST | ✓ | 3.3 | AFBNDL01 native
binary via airflow-go-pack |
| 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.4 | |
| 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/go-sdk/capabilities.yaml b/go-sdk/capabilities.yaml
index e038f7395d2..cabf5d45871 100644
--- a/go-sdk/capabilities.yaml
+++ b/go-sdk/capabilities.yaml
@@ -85,8 +85,8 @@ 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.4"
asset-state-store:
supported: false
note: "no task-facing state-store API yet"
diff --git a/go-sdk/cmd/airflow-go-pack/pack_integration_test.go
b/go-sdk/cmd/airflow-go-pack/pack_integration_test.go
index e1cc29fe68e..7f5c6b4a373 100644
--- a/go-sdk/cmd/airflow-go-pack/pack_integration_test.go
+++ b/go-sdk/cmd/airflow-go-pack/pack_integration_test.go
@@ -154,6 +154,9 @@ dags:
- "extract"
- "transform"
- "load"
+ task_state_dag:
+ tasks:
+ - "roundtrip_task_state"
taskflow_binding_dag:
tasks:
- "make_config"
diff --git a/go-sdk/dags/go_examples.py b/go-sdk/dags/go_examples.py
index c87fb3a682f..0344e5e2455 100644
--- a/go-sdk/dags/go_examples.py
+++ b/go-sdk/dags/go_examples.py
@@ -17,12 +17,13 @@
"""
Python stub Dags mirroring the Go SDK example bundle
(``go-sdk/example/bundle``).
-Four Dags, all backed by the same Go bundle: ``simple_dag`` (extract/transform/
+Five Dags, all backed by the same Go bundle: ``simple_dag`` (extract/transform/
load, below), ``concurrent_xcom_dag`` (one ``pull_xcoms_concurrently`` task
timing sequential vs goroutine XCom pulls), ``taskflow_binding_dag`` (one
task per shape of the TaskFlow argument-binding surface; see its Dag function
-below), and ``variable_write_dag`` (one ``write_and_delete_variable`` task that
-writes and deletes Airflow Variables).
+below), ``variable_write_dag`` (one ``write_and_delete_variable`` task that
+writes and deletes Airflow Variables), and ``task_state_dag`` (one
+``roundtrip_task_state`` task that round-trips the task state store).
``simple_dag`` sandwiches the Go tasks between two native Python tasks so the
run exercises XCom across the language boundary, the same way
@@ -237,3 +238,15 @@ def variable_write_dag():
variable_write_dag()
+
+
[email protected](queue="golang")
+def roundtrip_task_state(): ...
+
+
+@dag(dag_id="task_state_dag")
+def task_state_dag():
+ roundtrip_task_state()
+
+
+task_state_dag()
diff --git a/go-sdk/example/bundle/concurrentxcom/concurrentxcom_test.go
b/go-sdk/example/bundle/concurrentxcom/concurrentxcom_test.go
index d3a8a7deea8..088eaaa96e0 100644
--- a/go-sdk/example/bundle/concurrentxcom/concurrentxcom_test.go
+++ b/go-sdk/example/bundle/concurrentxcom/concurrentxcom_test.go
@@ -88,6 +88,10 @@ func (m *mockXComClient) GetConnection(ctx context.Context,
connID string) (sdk.
panic("unimplemented")
}
+func (m *mockXComClient) TaskStateStore() sdk.TaskStateStore {
+ panic("unimplemented")
+}
+
var _ sdk.Client = (*mockXComClient)(nil)
func Test_PullXComsConcurrently(t *testing.T) {
diff --git a/go-sdk/example/bundle/main.go b/go-sdk/example/bundle/main.go
index 00c346d420d..25e74681cd2 100644
--- a/go-sdk/example/bundle/main.go
+++ b/go-sdk/example/bundle/main.go
@@ -27,6 +27,7 @@ import (
"github.com/apache/airflow/go-sdk/airflow"
"github.com/apache/airflow/go-sdk/example/bundle/concurrentxcom"
"github.com/apache/airflow/go-sdk/example/bundle/taskflowbinding"
+ "github.com/apache/airflow/go-sdk/example/bundle/taskstate"
"github.com/apache/airflow/go-sdk/example/bundle/variablewrite"
)
@@ -83,6 +84,8 @@ func main() {
"write_and_delete_variable",
variablewrite.WriteAndDeleteVariable,
),
+
+ airflow.TaskHandler("task_state_dag", "roundtrip_task_state",
taskstate.RoundtripTaskState),
)
if err := bundle.Serve(); err != nil {
diff --git a/go-sdk/example/bundle/taskstate/taskstate.go
b/go-sdk/example/bundle/taskstate/taskstate.go
new file mode 100644
index 00000000000..c91ff9c6f0d
--- /dev/null
+++ b/go-sdk/example/bundle/taskstate/taskstate.go
@@ -0,0 +1,99 @@
+// Licensed to the Apache Software Foundation (ASF) under one
+// or more contributor license agreements. See the NOTICE file
+// distributed with this work for additional information
+// regarding copyright ownership. The ASF licenses this file
+// to you under the Apache License, Version 2.0 (the
+// "License"); you may not use this file except in compliance
+// with the License. You may obtain a copy of the License at
+//
+// http://www.apache.org/licenses/LICENSE-2.0
+//
+// Unless required by applicable law or agreed to in writing,
+// software distributed under the License is distributed on an
+// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
+// KIND, either express or implied. See the License for the
+// specific language governing permissions and limitations
+// under the License.
+
+// Package taskstate holds the roundtrip_task_state task used by the Go SDK
+// task state e2e test.
+package taskstate
+
+import (
+ "errors"
+ "fmt"
+
+ "github.com/apache/airflow/go-sdk/airflow"
+ "github.com/apache/airflow/go-sdk/sdk"
+)
+
+const (
+ // RunIDKey holds the run id; the e2e test asserts it has no expiry.
+ RunIDKey = "go_e2e_run_id"
+ // CounterKey holds a map value the e2e test reads back as an object.
+ CounterKey = "go_e2e_counter"
+ // RetainedKey uses the default retention; the e2e test asserts it
expires.
+ RetainedKey = "go_e2e_retained"
+ // ScratchKey is deleted by the task; the e2e test asserts it is gone.
+ ScratchKey = "go_e2e_scratch"
+)
+
+type counter struct {
+ Processed int `json:"processed"`
+ Cursor string `json:"cursor"`
+}
+
+// RoundtripTaskState exercises every task state store operation except Clear.
+func RoundtripTaskState(actx airflow.Context) (any, error) {
+ store := actx.Client().TaskStateStore()
+ runID := actx.DagRun().RunID
+
+ if err := store.Set(actx, RunIDKey, runID,
sdk.WithRetention(sdk.NeverExpire)); err != nil {
+ return nil, fmt.Errorf("setting %s: %w", RunIDKey, err)
+ }
+ want := counter{Processed: 3, Cursor: "abc-123"}
+ stored := map[string]any{"processed": want.Processed, "cursor":
want.Cursor}
+ if err := store.Set(actx, CounterKey, stored); err != nil {
+ return nil, fmt.Errorf("setting %s: %w", CounterKey, err)
+ }
+ if err := store.Set(actx, RetainedKey, "retained"); err != nil {
+ return nil, fmt.Errorf("setting %s: %w", RetainedKey, err)
+ }
+ if err := store.Set(actx, ScratchKey, "scratch"); err != nil {
+ return nil, fmt.Errorf("setting %s: %w", ScratchKey, err)
+ }
+ if err := store.Delete(actx, ScratchKey); err != nil {
+ return nil, fmt.Errorf("deleting %s: %w", ScratchKey, err)
+ }
+
+ readBack, err := store.Get(actx, RunIDKey)
+ if err != nil {
+ return nil, fmt.Errorf("getting %s: %w", RunIDKey, err)
+ }
+ if readBack != runID {
+ return nil, fmt.Errorf("getting %s: got %v, want run id %q",
RunIDKey, readBack, runID)
+ }
+
+ var got counter
+ if err := store.UnmarshalJSONValue(actx, CounterKey, &got); err != nil {
+ return nil, fmt.Errorf("decoding %s: %w", CounterKey, err)
+ }
+ if got != want {
+ return nil, fmt.Errorf("decoding %s: got %+v, want %+v",
CounterKey, got, want)
+ }
+
+ if _, err := store.Get(actx, ScratchKey); !errors.Is(err,
sdk.TaskStateNotFound) {
+ if err == nil {
+ return nil, fmt.Errorf("getting %s: key survived its
delete", ScratchKey)
+ }
+ return nil, fmt.Errorf("getting %s: want TaskStateNotFound,
got: %w", ScratchKey, err)
+ }
+
+ // No Clear: the e2e test reads these keys after the task finishes.
+ return map[string]any{
+ "run_id": runID,
+ "processed": got.Processed,
+ "cursor": got.Cursor,
+ "scratch_deleted": true,
+ }, nil
+}
diff --git a/go-sdk/pkg/execution/client.go b/go-sdk/pkg/execution/client.go
index 1f09da1e1ba..610a5c53a5c 100644
--- a/go-sdk/pkg/execution/client.go
+++ b/go-sdk/pkg/execution/client.go
@@ -18,12 +18,18 @@
package execution
import (
+ "bytes"
"context"
"encoding/json"
"errors"
"fmt"
+ "math"
"os"
+ "strconv"
"strings"
+ "time"
+
+ "github.com/vmihailenco/msgpack/v5"
"github.com/apache/airflow/go-sdk/pkg/execution/genmodels"
"github.com/apache/airflow/go-sdk/sdk"
@@ -36,8 +42,13 @@ const (
errCodeVariableNotFound = "VARIABLE_NOT_FOUND"
errCodeConnectionNotFound = "CONNECTION_NOT_FOUND"
errCodeXComNotFound = "XCOM_NOT_FOUND"
+ errCodeTaskStoreNotFound = "TASK_STORE_NOT_FOUND"
)
+// A language SDK runtime cannot read Airflow config, so the supervisor passes
+// this setting at launch
(task-sdk/src/airflow/sdk/coordinators/_subprocess.py).
+const defaultRetentionDaysEnv = "AIRFLOW__STATE_STORE__DEFAULT_RETENTION_DAYS"
+
// translateAPIError converts a supervisor *APIError whose Err field matches
// code into a sentinel-wrapped error. Any other error - including a
// *APIError with a different code - is returned unchanged so callers can keep
@@ -57,15 +68,83 @@ func translateAPIError(err error, code string, sentinel
error, key string) error
// over the comm socket using msgpack-framed IPC instead of HTTP.
type CoordinatorClient struct {
comm *CoordinatorComm
+ // Bound at construction, not per call: the Execution API scopes the
task
+ // state store to the caller's own task instance ("ti:self").
+ tiID string
}
var _ sdk.Client = (*CoordinatorClient)(nil)
// NewCoordinatorClient creates a new client backed by the comm socket.
-func NewCoordinatorClient(comm *CoordinatorComm) *CoordinatorClient {
+func NewCoordinatorClient(comm *CoordinatorComm, tiID string)
*CoordinatorClient {
return &CoordinatorClient{
comm: comm,
+ tiID: tiID,
+ }
+}
+
+// resolveDefaultExpiry returns nil ("never expires") for a retention of 0. The
+// supervisor always passes the setting, so an absent or malformed value is a
+// misconfiguration and fails the write rather than silently retaining the key
+// for a period nobody configured.
+func resolveDefaultExpiry(now time.Time) (any, error) {
+ raw, ok := os.LookupEnv(defaultRetentionDaysEnv)
+ if !ok {
+ return nil, fmt.Errorf(
+ "%s is not set; it carries the deployment's %q key in
%q section to the runtime",
+ defaultRetentionDaysEnv,
+ "default_retention_days",
+ "state_store",
+ )
+ }
+ days, err := parseRetentionDays(raw)
+ if err != nil {
+ return nil, err
+ }
+ if days == 0 {
+ return nil, nil
+ }
+ expiry := now.UTC().AddDate(0, 0, days)
+ // A day count big enough to wrap the timestamp would store a key that
is
+ // already expired, losing it on the next cleanup; Python raises
+ // OverflowError on the same setting.
+ if !expiry.After(now.UTC()) {
+ return nil, fmt.Errorf(
+ "a retention of %d days overflows the expiry timestamp.
Please check %q key in %q section",
+ days,
+ "default_retention_days",
+ "state_store",
+ )
+ }
+ return expiry, nil
+}
+
+// parseRetentionDays accepts "7.0" because Python's getint does.
+func parseRetentionDays(raw string) (int, error) {
+ days, err := strconv.Atoi(raw)
+ if err != nil {
+ f, floatErr := strconv.ParseFloat(raw, 64)
+ // A float outside int64 range (an infinity included) converts
with an
+ // implementation-defined result, so it is rejected rather than
turned
+ // into whatever day count this platform happens to produce.
+ if floatErr != nil || f != math.Trunc(f) || f >= math.MaxInt64
|| f <= math.MinInt64 {
+ return 0, fmt.Errorf(
+ "failed to convert value to int. Please check
%q key in %q section. Current value: %q",
+ "default_retention_days",
+ "state_store",
+ raw,
+ )
+ }
+ days = int(f)
+ }
+ if days < 0 {
+ return 0, fmt.Errorf(
+ "[state_store] default_retention_days must be >= 0, got
%d. "+
+ "Set to 0 to disable expiry.",
+ days,
+ )
}
+ return days, nil
}
// GetVariable requests a variable value from the supervisor.
@@ -281,3 +360,199 @@ func (c *CoordinatorClient) skipDownstreamTasks(ctx
context.Context, taskIDs []s
_, err := c.comm.Communicate(ctx, genmodels.SkipDownstreamTasks{Tasks:
taskIDs})
return err
}
+
+// TaskStateStore returns the task state store scoped to this task instance.
+func (c *CoordinatorClient) TaskStateStore() sdk.TaskStateStore {
+ return taskStateStore{client: c}
+}
+
+// taskStateStore serves sdk.TaskStateStore over the coordinator comm.
+type taskStateStore struct {
+ client *CoordinatorClient
+}
+
+// Get asks the supervisor for a task state value.
+func (s taskStateStore) Get(ctx context.Context, key string) (any, error) {
+ resp, err := s.client.comm.Communicate(
+ ctx,
+ genmodels.GetTaskStateStore{TIID: s.client.tiID, Key: key},
+ )
+ if err != nil {
+ return nil, translateAPIError(err, errCodeTaskStoreNotFound,
sdk.TaskStateNotFound, key)
+ }
+
+ var result genmodels.TaskStateStoreResult
+ if err := decodeBody(resp, &result); err != nil {
+ return nil, fmt.Errorf("decoding task state result: %w", err)
+ }
+
+ return result.Value, nil
+}
+
+// UnmarshalJSONValue gets a task state value and unmarshals it into pointer.
+func (s taskStateStore) UnmarshalJSONValue(
+ ctx context.Context,
+ key string,
+ pointer any,
+) error {
+ val, err := s.Get(ctx, key)
+ if err != nil {
+ return err
+ }
+ // The value arrives already decoded from msgpack, not as JSON text, so
it
+ // is re-marshaled before encoding/json can fill a typed pointer.
+ b, err := json.Marshal(val)
+ if err != nil {
+ return fmt.Errorf("marshaling task state value: %w", err)
+ }
+ return json.Unmarshal(b, pointer)
+}
+
+// Set asks the supervisor to store a task state value.
+func (s taskStateStore) Set(
+ ctx context.Context,
+ key string,
+ value any,
+ opts ...sdk.SetOption,
+) error {
+ var options sdk.SetOptions
+ for _, opt := range opts {
+ opt(&options)
+ }
+
+ // The caller's value is checked first so a misconfigured deployment
cannot
+ // mask a programming error in the task.
+ if err := validateJSONRepresentable(value); err != nil {
+ return fmt.Errorf("cannot set task state key %q: %w", key, err)
+ }
+
+ expiry, err := resolveExpiry(options.Retention, time.Now())
+ if err != nil {
+ return fmt.Errorf("cannot set task state key %q: %w", key, err)
+ }
+
+ // TODO: warn when the serialized value exceeds the deployment's
+ // [state_store] max_value_storage_bytes, matching Python's
+ // airflow.sdk.execution_time.context task store setter.
+
+ _, err = s.client.comm.Communicate(ctx, genmodels.SetTaskStateStore{
+ TIID: s.client.tiID,
+ Key: key,
+ Value: value,
+ ExpiresAt: expiry,
+ })
+ return err
+}
+
+// Delete asks the supervisor to delete a task state value.
+func (s taskStateStore) Delete(ctx context.Context, key string) error {
+ _, err := s.client.comm.Communicate(
+ ctx,
+ genmodels.DeleteTaskStateStore{TIID: s.client.tiID, Key: key},
+ )
+ return err
+}
+
+// Clear asks the supervisor to delete every task state value for this task
+// instance.
+func (s taskStateStore) Clear(ctx context.Context) error {
+ _, err := s.client.comm.Communicate(
+ ctx,
+ genmodels.ClearTaskStateStore{TIID: s.client.tiID},
+ )
+ return err
+}
+
+// resolveExpiry turns a retention into the wire expires_at. A nil retention -
+// no sdk.WithRetention - follows the deployment default.
+func resolveExpiry(retention *time.Duration, now time.Time) (any, error) {
+ if retention == nil {
+ return resolveDefaultExpiry(now)
+ }
+ switch r := *retention; {
+ // Checked before any arithmetic: adding NeverExpire overflows.
+ case r == sdk.NeverExpire:
+ return nil, nil
+ case r <= 0:
+ return nil, fmt.Errorf(
+ "retention must be positive or sdk.NeverExpire, got %s:
omit "+
+ "sdk.WithRetention to follow the deployment
default, or call Delete to drop the key",
+ r,
+ )
+ default:
+ return now.UTC().Add(r), nil
+ }
+}
+
+// validateJSONRepresentable checks what the frame encoder actually emits, not
+// the Go value: a reflection walk has to mirror the encoder's field rules
(tags,
+// "-", omitempty, embedding, marshalers) and misjudges values wherever it
drifts.
+func validateJSONRepresentable(value any) error {
+ var buf bytes.Buffer
+ if err := newFrameEncoder(&buf).Encode(value); err != nil {
+ return fmt.Errorf("%T is not JSON representable: %w", value,
err)
+ }
+ dec := msgpack.NewDecoder(&buf)
+ // The default map decoder fails opaquely on non-string keys.
+ dec.SetMapDecoder(func(d *msgpack.Decoder) (any, error) {
+ return d.DecodeUntypedMap()
+ })
+ decoded, err := dec.DecodeInterface()
+ if err != nil {
+ return fmt.Errorf("%T is not JSON representable: %w", value,
err)
+ }
+ // Checked after decoding: a typed nil such as a nil *string is not ==
nil,
+ // yet it still encodes to null, which the Execution API rejects.
+ if decoded == nil {
+ return errors.New("value must not be nil")
+ }
+ return validateDecodedValue(decoded)
+}
+
+func validateDecodedValue(v any) error {
+ switch v := v.(type) {
+ case nil, string, bool,
+ int, int8, int16, int32, int64,
+ uint, uint8, uint16, uint32, uint64:
+ return nil
+ case float32:
+ return checkFinite(float64(v))
+ case float64:
+ return checkFinite(v)
+ case time.Time:
+ return fmt.Errorf(
+ "time.Time is not JSON representable; store
value.Format(time.RFC3339) " +
+ "and parse it back with time.Parse",
+ )
+ case []byte:
+ return fmt.Errorf(
+ "[]byte is not JSON representable; encode it, for
example with base64.StdEncoding.EncodeToString",
+ )
+ case []any:
+ for _, elem := range v {
+ if err := validateDecodedValue(elem); err != nil {
+ return err
+ }
+ }
+ return nil
+ case map[any]any:
+ for k, elem := range v {
+ if _, ok := k.(string); !ok {
+ return fmt.Errorf("map keys must be strings,
got %T", k)
+ }
+ if err := validateDecodedValue(elem); err != nil {
+ return err
+ }
+ }
+ return nil
+ default:
+ return fmt.Errorf("%T is not JSON representable", v)
+ }
+}
+
+func checkFinite(f float64) error {
+ if math.IsNaN(f) || math.IsInf(f, 0) {
+ return fmt.Errorf("value must be a finite number; NaN and Inf
are not JSON representable")
+ }
+ return nil
+}
diff --git a/go-sdk/pkg/execution/client_test.go
b/go-sdk/pkg/execution/client_test.go
index da15844c621..bd8dfd3f6e7 100644
--- a/go-sdk/pkg/execution/client_test.go
+++ b/go-sdk/pkg/execution/client_test.go
@@ -23,7 +23,10 @@ import (
"errors"
"io"
"log/slog"
+ "math"
+ "os"
"testing"
+ "time"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
@@ -31,6 +34,8 @@ import (
"github.com/apache/airflow/go-sdk/sdk"
)
+const testTIID = "0199e0e5-1b2c-7c3d-8e4f-5a6b7c8d9e0f"
+
// TestCoordinatorClientGetVariableEnvOverride verifies that an
// AIRFLOW_VAR_<UPPER(key)> environment override short-circuits the comm
// socket.
@@ -42,7 +47,7 @@ func TestCoordinatorClientGetVariableEnvOverride(t
*testing.T) {
// to assert *no* IO occurred by failing if anything is read or written.
logger := slog.New(slog.NewTextHandler(io.Discard, nil))
comm := NewCoordinatorComm(assertNoReadReader{t: t},
assertNoWriteWriter{t: t}, logger)
- client := NewCoordinatorClient(comm)
+ client := NewCoordinatorClient(comm, testTIID)
val, err := client.GetVariable(context.Background(), "my_key")
require.NoError(t, err)
@@ -65,7 +70,7 @@ func TestCoordinatorClientGetVariableNoEnvOverride(t
*testing.T) {
var requestBuf bytes.Buffer
logger := slog.New(slog.NewTextHandler(io.Discard, nil))
comm := NewCoordinatorComm(&responseBuf, &requestBuf, logger)
- client := NewCoordinatorClient(comm)
+ client := NewCoordinatorClient(comm, testTIID)
val, err := client.GetVariable(context.Background(), "my_key")
require.NoError(t, err)
@@ -124,7 +129,7 @@ func TestCoordinatorClientErrorTranslation(t *testing.T) {
logger := slog.New(slog.NewTextHandler(io.Discard, nil))
comm := NewCoordinatorComm(&responseBuf, io.Discard,
logger)
- client := NewCoordinatorClient(comm)
+ client := NewCoordinatorClient(comm, testTIID)
err := tc.call(client)
require.Error(t, err)
@@ -150,7 +155,7 @@ func TestCoordinatorClientErrorPassThrough(t *testing.T) {
logger := slog.New(slog.NewTextHandler(io.Discard, nil))
comm := NewCoordinatorComm(&responseBuf, io.Discard, logger)
- client := NewCoordinatorClient(comm)
+ client := NewCoordinatorClient(comm, testTIID)
_, err := client.GetVariable(context.Background(), "any_key")
require.Error(t, err)
@@ -187,7 +192,7 @@ func TestCoordinatorClientSetVariable(t *testing.T) {
var requestBuf bytes.Buffer
logger := slog.New(slog.NewTextHandler(io.Discard, nil))
comm := NewCoordinatorComm(&responseBuf, &requestBuf,
logger)
- client := NewCoordinatorClient(comm)
+ client := NewCoordinatorClient(comm, testTIID)
require.NoError(
t,
@@ -221,7 +226,7 @@ func TestCoordinatorClientDeleteVariable(t *testing.T) {
var requestBuf bytes.Buffer
logger := slog.New(slog.NewTextHandler(io.Discard, nil))
comm := NewCoordinatorComm(&responseBuf, &requestBuf, logger)
- client := NewCoordinatorClient(comm)
+ client := NewCoordinatorClient(comm, testTIID)
require.NoError(t, client.DeleteVariable(context.Background(),
"my_key"))
@@ -265,7 +270,7 @@ func TestCoordinatorClientVariableWriteErrors(t *testing.T)
{
logger := slog.New(slog.NewTextHandler(io.Discard, nil))
comm := NewCoordinatorComm(&responseBuf, io.Discard,
logger)
- client := NewCoordinatorClient(comm)
+ client := NewCoordinatorClient(comm, testTIID)
var apiErr *APIError
require.ErrorAs(t, tc.call(client), &apiErr)
@@ -291,7 +296,7 @@ func
TestCoordinatorClientGetConnectionPreservesEmptyCredentials(t *testing.T) {
logger := slog.New(slog.NewTextHandler(io.Discard, nil))
comm := NewCoordinatorComm(&responseBuf, io.Discard, logger)
- client := NewCoordinatorClient(comm)
+ client := NewCoordinatorClient(comm, testTIID)
conn, err := client.GetConnection(context.Background(), "c")
require.NoError(t, err)
@@ -314,7 +319,7 @@ func TestCoordinatorClientGetConnectionAbsentCredentials(t
*testing.T) {
logger := slog.New(slog.NewTextHandler(io.Discard, nil))
comm := NewCoordinatorComm(&responseBuf, io.Discard, logger)
- client := NewCoordinatorClient(comm)
+ client := NewCoordinatorClient(comm, testTIID)
conn, err := client.GetConnection(context.Background(), "c")
require.NoError(t, err)
@@ -364,7 +369,7 @@ func TestCoordinatorClientPushXComMapIndex(t *testing.T) {
var requestBuf bytes.Buffer
logger := slog.New(slog.NewTextHandler(io.Discard, nil))
comm := NewCoordinatorComm(&responseBuf, &requestBuf,
logger)
- client := NewCoordinatorClient(comm)
+ client := NewCoordinatorClient(comm, testTIID)
ti := sdk.TaskInstance{
DagID: "d",
@@ -432,7 +437,7 @@ func TestCoordinatorClientGetXComMapIndex(t *testing.T) {
var requestBuf bytes.Buffer
logger := slog.New(slog.NewTextHandler(io.Discard, nil))
comm := NewCoordinatorComm(&responseBuf, &requestBuf,
logger)
- client := NewCoordinatorClient(comm)
+ client := NewCoordinatorClient(comm, testTIID)
_, err := client.GetXCom(context.Background(), "d",
"r", "t", tc.mapIndex, "k", nil)
require.NoError(t, err)
@@ -452,6 +457,552 @@ func TestCoordinatorClientGetXComMapIndex(t *testing.T) {
}
}
+// The supervisor always passes the setting, so an absent or malformed value
must
+// fail as Python's TaskStateStoreAccessor.set does rather than be guessed at.
+func TestResolveDefaultExpiry(t *testing.T) {
+ now := time.Date(2026, 6, 9, 12, 0, 0, 0, time.UTC)
+
+ tests := []struct {
+ name string
+ env string
+ unset bool
+ want any
+ wantErr string
+ }{
+ {name: "unset is an error", unset: true, wantErr: "is not set"},
+ {name: "honours supervisor value", env: "7", want:
now.UTC().AddDate(0, 0, 7)},
+ {name: "zero days never expires", env: "0", want: nil},
+ // Python's config parser accepts a whole-number float spelling.
+ {name: "whole float accepted", env: "7.0", want:
now.UTC().AddDate(0, 0, 7)},
+ {name: "unparsable is an error", env: "abc", wantErr: "failed
to convert value to int"},
+ // conf.get returns "" for an empty setting, which must not
read as unset.
+ {name: "empty is an error", env: "", wantErr: "failed to
convert value to int"},
+ {name: "fractional is an error", env: "7.5", wantErr: "failed
to convert value to int"},
+ {name: "negative is an error", env: "-1", wantErr: "must be >=
0, got -1"},
+ {
+ name: "out of int64 range is an error",
+ env: "1e30",
+ wantErr: "failed to convert value to int",
+ },
+ {
+ name: "day count that wraps the timestamp is an
error",
+ env: "9223372036854775807",
+ wantErr: "overflows the expiry timestamp",
+ },
+ }
+
+ for _, tc := range tests {
+ t.Run(tc.name, func(t *testing.T) {
+ // Setenv registers the restore, so unsetting it here
is undone after.
+ t.Setenv(defaultRetentionDaysEnv, tc.env)
+ if tc.unset {
+ require.NoError(t,
os.Unsetenv(defaultRetentionDaysEnv))
+ }
+
+ got, err := resolveDefaultExpiry(now)
+ if tc.wantErr != "" {
+ require.ErrorContains(t, err, tc.wantErr)
+ assert.Nil(t, got)
+ return
+ }
+ require.NoError(t, err)
+ if tc.want == nil {
+ assert.Nil(t, got, "a nil expiry must be
untyped so msgpack encodes null")
+ return
+ }
+ assert.Equal(t, tc.want, got)
+ })
+ }
+}
+
+func TestTaskStateStoreSetRejectsMisconfiguredRetention(t *testing.T) {
+ t.Setenv(defaultRetentionDaysEnv, "-1")
+
+ var requestBuf bytes.Buffer
+ logger := slog.New(slog.NewTextHandler(io.Discard, nil))
+ client := NewCoordinatorClient(
+ NewCoordinatorComm(&bytes.Buffer{}, &requestBuf, logger),
+ testTIID,
+ )
+
+ err := client.TaskStateStore().Set(context.Background(), "job_id",
"app_001")
+
+ require.ErrorContains(t, err, "must be >= 0, got -1")
+ assert.Zero(t, requestBuf.Len(), "a rejected write must not reach the
supervisor")
+}
+
+// Mirrors Python's test_set_datetime_raises_validation_error.
+func TestTaskStateStoreSetRejectsNonJSONValues(t *testing.T) {
+ tests := []struct {
+ name string
+ value any
+ wantErr string
+ }{
+ {
+ name: "datetime",
+ value: time.Date(2026, 5, 15, 0, 0, 0, 0, time.UTC),
+ wantErr: "time.Time is not JSON representable",
+ },
+ {
+ name: "datetime nested in a map",
+ value: map[string]any{"watermark": time.Date(2026, 5,
15, 0, 0, 0, 0, time.UTC)},
+ wantErr: "time.Time is not JSON representable",
+ },
+ {name: "NaN", value: math.NaN(), wantErr: "finite number"},
+ {name: "Inf", value: math.Inf(1), wantErr: "finite number"},
+ {name: "byte slice", value: []byte("raw"), wantErr: "[]byte is
not JSON representable"},
+ {name: "byte array", value: [16]byte{}, wantErr: "[]byte is not
JSON representable"},
+ {
+ name: "non-string map key",
+ value: map[int]string{1: "a"},
+ wantErr: "map keys must be strings",
+ },
+ }
+
+ for _, tc := range tests {
+ t.Run(tc.name, func(t *testing.T) {
+ var requestBuf bytes.Buffer
+ logger := slog.New(slog.NewTextHandler(io.Discard, nil))
+ client := NewCoordinatorClient(
+ NewCoordinatorComm(&bytes.Buffer{},
&requestBuf, logger), testTIID,
+ )
+
+ err :=
client.TaskStateStore().Set(context.Background(), "job_id", tc.value)
+
+ require.ErrorContains(t, err, tc.wantErr)
+ assert.Zero(t, requestBuf.Len(), "a rejected write must
not reach the supervisor")
+ })
+ }
+}
+
+func TestTaskStateStoreSetAcceptsJSONShapes(t *testing.T) {
+ // A default-retention write needs the setting the supervisor passes.
+ t.Setenv(defaultRetentionDaysEnv, "30")
+
+ type checkpoint struct {
+ Processed int `msgpack:"processed"`
+ Cursors []string `msgpack:"cursors"`
+ }
+ type skippedTime struct {
+ When time.Time `json:"-"`
+ Name string `json:"name"`
+ }
+ values := map[string]any{
+ "struct": checkpoint{Processed: 3, Cursors:
[]string{"a"}},
+ "struct skipping time.Time": skippedTime{When: time.Now(),
Name: "x"},
+ "nested": map[string]any{"rows": []any{1,
"two", 3.5, true, nil}},
+ "scalar": "plain",
+ }
+
+ for name, value := range values {
+ t.Run(name, func(t *testing.T) {
+ responsePayload := encodeResponseFrame(t, 0, nil, nil)
+ var responseBuf bytes.Buffer
+ require.NoError(t, writeFrame(&responseBuf,
responsePayload))
+
+ var requestBuf bytes.Buffer
+ logger := slog.New(slog.NewTextHandler(io.Discard, nil))
+ client := NewCoordinatorClient(
+ NewCoordinatorComm(&responseBuf, &requestBuf,
logger),
+ testTIID,
+ )
+
+ require.NoError(t,
client.TaskStateStore().Set(context.Background(), "job_id", value))
+ assert.NotZero(t, requestBuf.Len())
+ })
+ }
+}
+
+func TestTaskStateStoreGet(t *testing.T) {
+ tests := []struct {
+ name string
+ value any
+ }{
+ {name: "scalar value", value: "abc123"},
+ {name: "structured value", value: map[string]any{"cursor":
"abc", "done": true}},
+ }
+
+ for _, tc := range tests {
+ t.Run(tc.name, func(t *testing.T) {
+ responsePayload := encodeResponseFrame(t, 0,
map[string]any{
+ "type": "TaskStateStoreResult",
+ "value": tc.value,
+ }, nil)
+ var responseBuf bytes.Buffer
+ require.NoError(t, writeFrame(&responseBuf,
responsePayload))
+
+ var requestBuf bytes.Buffer
+ logger := slog.New(slog.NewTextHandler(io.Discard, nil))
+ comm := NewCoordinatorComm(&responseBuf, &requestBuf,
logger)
+ client := NewCoordinatorClient(comm, testTIID)
+
+ got, err :=
client.TaskStateStore().Get(context.Background(), "job_id")
+ require.NoError(t, err)
+ assert.Equal(t, tc.value, got)
+
+ sent, err := readFrame(&requestBuf)
+ require.NoError(t, err)
+ assert.Equal(t, map[string]any{
+ "type": "GetTaskStateStore",
+ "ti_id": testTIID,
+ "key": "job_id",
+ }, rawToMap(t, sent.Body))
+ })
+ }
+}
+
+func TestTaskStateStoreGetNotFound(t *testing.T) {
+ responsePayload := encodeResponseFrame(t, 0, nil, map[string]any{
+ "type": "ErrorResponse",
+ "error": "TASK_STORE_NOT_FOUND",
+ "detail": map[string]any{"msg": "no such key"},
+ })
+ var responseBuf bytes.Buffer
+ require.NoError(t, writeFrame(&responseBuf, responsePayload))
+
+ logger := slog.New(slog.NewTextHandler(io.Discard, nil))
+ comm := NewCoordinatorComm(&responseBuf, io.Discard, logger)
+ client := NewCoordinatorClient(comm, testTIID)
+
+ _, err := client.TaskStateStore().Get(context.Background(), "missing")
+ require.Error(t, err)
+ assert.ErrorIs(t, err, sdk.TaskStateNotFound)
+ assert.Contains(t, err.Error(), "missing")
+}
+
+func TestTaskStateStoreGetErrorPassThrough(t *testing.T) {
+ responsePayload := encodeResponseFrame(t, 0, nil, map[string]any{
+ "type": "ErrorResponse",
+ "error": "API_SERVER_ERROR",
+ "detail": map[string]any{"msg": "boom"},
+ })
+ var responseBuf bytes.Buffer
+ require.NoError(t, writeFrame(&responseBuf, responsePayload))
+
+ logger := slog.New(slog.NewTextHandler(io.Discard, nil))
+ comm := NewCoordinatorComm(&responseBuf, io.Discard, logger)
+ client := NewCoordinatorClient(comm, testTIID)
+
+ _, err := client.TaskStateStore().Get(context.Background(), "job_id")
+ require.Error(t, err)
+ assert.False(t, errors.Is(err, sdk.TaskStateNotFound),
+ "generic supervisor errors must not be translated to
TaskStateNotFound")
+ var apiErr *ApiError
+ require.True(t, errors.As(err, &apiErr))
+ assert.Equal(t, "API_SERVER_ERROR", apiErr.Err)
+}
+
+func TestTaskStateStoreUnmarshalJSONValue(t *testing.T) {
+ type checkpoint struct {
+ Cursor string `json:"cursor"`
+ Done bool `json:"done"`
+ }
+
+ t.Run("decodes into a struct", func(t *testing.T) {
+ responsePayload := encodeResponseFrame(t, 0, map[string]any{
+ "type": "TaskStateStoreResult",
+ "value": map[string]any{"cursor": "abc", "done": true},
+ }, nil)
+ var responseBuf bytes.Buffer
+ require.NoError(t, writeFrame(&responseBuf, responsePayload))
+
+ logger := slog.New(slog.NewTextHandler(io.Discard, nil))
+ comm := NewCoordinatorComm(&responseBuf, io.Discard, logger)
+ client := NewCoordinatorClient(comm, testTIID)
+
+ var got checkpoint
+ require.NoError(
+ t,
+
client.TaskStateStore().UnmarshalJSONValue(context.Background(), "job_id",
&got),
+ )
+ assert.Equal(t, checkpoint{Cursor: "abc", Done: true}, got)
+ })
+
+ t.Run("propagates not found", func(t *testing.T) {
+ responsePayload := encodeResponseFrame(t, 0, nil,
map[string]any{
+ "type": "ErrorResponse",
+ "error": "TASK_STORE_NOT_FOUND",
+ "detail": map[string]any{"msg": "no such key"},
+ })
+ var responseBuf bytes.Buffer
+ require.NoError(t, writeFrame(&responseBuf, responsePayload))
+
+ logger := slog.New(slog.NewTextHandler(io.Discard, nil))
+ comm := NewCoordinatorComm(&responseBuf, io.Discard, logger)
+ client := NewCoordinatorClient(comm, testTIID)
+
+ var got checkpoint
+ err :=
client.TaskStateStore().UnmarshalJSONValue(context.Background(), "missing",
&got)
+ assert.ErrorIs(t, err, sdk.TaskStateNotFound)
+ })
+}
+
+// expires_at is sent even when null: the supervisor requires the field.
+func TestTaskStateStoreSet(t *testing.T) {
+ tests := []struct {
+ name string
+ retentionDays string
+ wantRetention time.Duration
+ wantExpiresNil bool
+ }{
+ {
+ name: "deployment retention is applied",
+ retentionDays: "7",
+ wantRetention: 7 * 24 * time.Hour,
+ },
+ {name: "zero retention sends null", retentionDays: "0",
wantExpiresNil: true},
+ }
+
+ for _, tc := range tests {
+ t.Run(tc.name, func(t *testing.T) {
+ t.Setenv(defaultRetentionDaysEnv, tc.retentionDays)
+
+ responsePayload := encodeResponseFrame(t, 0,
map[string]any{"type": "OKResponse"}, nil)
+ var responseBuf bytes.Buffer
+ require.NoError(t, writeFrame(&responseBuf,
responsePayload))
+
+ var requestBuf bytes.Buffer
+ logger := slog.New(slog.NewTextHandler(io.Discard, nil))
+ comm := NewCoordinatorComm(&responseBuf, &requestBuf,
logger)
+ client := NewCoordinatorClient(comm, testTIID)
+
+ before := time.Now()
+ require.NoError(
+ t,
+
client.TaskStateStore().Set(context.Background(), "job_id", "abc123"),
+ )
+ after := time.Now()
+
+ sent, err := readFrame(&requestBuf)
+ require.NoError(t, err)
+ sentMap := rawToMap(t, sent.Body)
+ assert.Equal(t, "SetTaskStateStore", sentMap["type"])
+ assert.Equal(t, testTIID, sentMap["ti_id"])
+ assert.Equal(t, "job_id", sentMap["key"])
+ assert.Equal(t, "abc123", sentMap["value"])
+ require.Contains(t, sentMap, "expires_at",
+ "expires_at must be present even when null")
+ if tc.wantExpiresNil {
+ assert.Nil(t, sentMap["expires_at"])
+ } else {
+ got, ok := sentMap["expires_at"].(time.Time)
+ require.True(t, ok, "expires_at must be a
timestamp, got %T", sentMap["expires_at"])
+ assert.WithinRange(t, got,
before.Add(tc.wantRetention), after.Add(tc.wantRetention))
+ }
+ })
+ }
+}
+
+func TestTaskStateStoreSetWithRetention(t *testing.T) {
+ tests := []struct {
+ name string
+ retention time.Duration
+ wantErr bool
+ wantExpiresNil bool
+ }{
+ {name: "positive retention is sent", retention: time.Hour},
+ {name: "NeverExpire sends null", retention: sdk.NeverExpire,
wantExpiresNil: true},
+ {name: "zero retention is rejected", retention: 0, wantErr:
true},
+ {name: "negative retention is rejected", retention: -time.Hour,
wantErr: true},
+ }
+
+ for _, tc := range tests {
+ t.Run(tc.name, func(t *testing.T) {
+ responsePayload := encodeResponseFrame(t, 0,
map[string]any{"type": "OKResponse"}, nil)
+ var responseBuf bytes.Buffer
+ require.NoError(t, writeFrame(&responseBuf,
responsePayload))
+
+ var requestBuf bytes.Buffer
+ logger := slog.New(slog.NewTextHandler(io.Discard, nil))
+ comm := NewCoordinatorComm(&responseBuf, &requestBuf,
logger)
+ client := NewCoordinatorClient(comm, testTIID)
+
+ before := time.Now()
+ err := client.TaskStateStore().Set(
+ context.Background(), "job_id", "abc123",
sdk.WithRetention(tc.retention),
+ )
+ after := time.Now()
+ if tc.wantErr {
+ require.Error(t, err)
+ assert.Zero(t, requestBuf.Len(), "a rejected
retention must send no frame")
+ return
+ }
+ require.NoError(t, err)
+
+ sent, err := readFrame(&requestBuf)
+ require.NoError(t, err)
+ sentMap := rawToMap(t, sent.Body)
+ assert.Equal(t, "SetTaskStateStore", sentMap["type"])
+ assert.Equal(t, testTIID, sentMap["ti_id"])
+ require.Contains(t, sentMap, "expires_at")
+ if tc.wantExpiresNil {
+ assert.Nil(t, sentMap["expires_at"])
+ } else {
+ got, ok := sentMap["expires_at"].(time.Time)
+ require.True(t, ok, "expires_at must be a
timestamp, got %T", sentMap["expires_at"])
+ assert.WithinRange(t, got,
before.Add(tc.retention), after.Add(tc.retention))
+ }
+ })
+ }
+}
+
+func TestTaskStateStoreSetRejectsNilValue(t *testing.T) {
+ tests := []struct {
+ name string
+ call func(client *CoordinatorClient) error
+ }{
+ {
+ name: "Set",
+ call: func(client *CoordinatorClient) error {
+ return
client.TaskStateStore().Set(context.Background(), "job_id", nil)
+ },
+ },
+ {
+ name: "Set with retention",
+ call: func(client *CoordinatorClient) error {
+ return client.TaskStateStore().Set(
+ context.Background(), "job_id", nil,
sdk.WithRetention(time.Hour),
+ )
+ },
+ },
+ {
+ name: "Set typed nil",
+ call: func(client *CoordinatorClient) error {
+ return
client.TaskStateStore().Set(context.Background(), "job_id", (*string)(nil))
+ },
+ },
+ {
+ name: "Set with retention, typed nil",
+ call: func(client *CoordinatorClient) error {
+ return client.TaskStateStore().Set(
+ context.Background(), "job_id",
(*string)(nil), sdk.WithRetention(time.Hour),
+ )
+ },
+ },
+ }
+
+ for _, tc := range tests {
+ t.Run(tc.name, func(t *testing.T) {
+ var requestBuf bytes.Buffer
+ logger := slog.New(slog.NewTextHandler(io.Discard, nil))
+ comm := NewCoordinatorComm(bytes.NewReader(nil),
&requestBuf, logger)
+ client := NewCoordinatorClient(comm, testTIID)
+
+ require.Error(t, tc.call(client))
+ assert.Zero(t, requestBuf.Len(), "a nil value must send
no frame")
+ })
+ }
+}
+
+func TestTaskStateStoreDelete(t *testing.T) {
+ responsePayload := encodeResponseFrame(
+ t,
+ 0,
+ map[string]any{"type": "OKResponse", "ok": true},
+ nil,
+ )
+ var responseBuf bytes.Buffer
+ require.NoError(t, writeFrame(&responseBuf, responsePayload))
+
+ var requestBuf bytes.Buffer
+ logger := slog.New(slog.NewTextHandler(io.Discard, nil))
+ comm := NewCoordinatorComm(&responseBuf, &requestBuf, logger)
+ client := NewCoordinatorClient(comm, testTIID)
+
+ require.NoError(t, client.TaskStateStore().Delete(context.Background(),
"job_id"))
+
+ sent, err := readFrame(&requestBuf)
+ require.NoError(t, err)
+ assert.Equal(t, map[string]any{
+ "type": "DeleteTaskStateStore",
+ "ti_id": testTIID,
+ "key": "job_id",
+ }, rawToMap(t, sent.Body))
+}
+
+func TestTaskStateStoreClear(t *testing.T) {
+ responsePayload := encodeResponseFrame(
+ t,
+ 0,
+ map[string]any{"type": "OKResponse", "ok": true},
+ nil,
+ )
+ var responseBuf bytes.Buffer
+ require.NoError(t, writeFrame(&responseBuf, responsePayload))
+
+ var requestBuf bytes.Buffer
+ logger := slog.New(slog.NewTextHandler(io.Discard, nil))
+ comm := NewCoordinatorComm(&responseBuf, &requestBuf, logger)
+ client := NewCoordinatorClient(comm, testTIID)
+
+ require.NoError(t, client.TaskStateStore().Clear(context.Background()))
+
+ sent, err := readFrame(&requestBuf)
+ require.NoError(t, err)
+ sentMap := rawToMap(t, sent.Body)
+ assert.Equal(t, map[string]any{
+ "type": "ClearTaskStateStore",
+ "ti_id": testTIID,
+ }, sentMap)
+ assert.NotContains(t, sentMap, "key")
+}
+
+func TestTaskStateStoreWriteErrors(t *testing.T) {
+ // A default-retention write needs the setting the supervisor passes.
+ t.Setenv(defaultRetentionDaysEnv, "30")
+
+ tests := []struct {
+ name string
+ call func(client *CoordinatorClient) error
+ }{
+ {
+ name: "Set",
+ call: func(client *CoordinatorClient) error {
+ return
client.TaskStateStore().Set(context.Background(), "job_id", "v")
+ },
+ },
+ {
+ name: "Set with retention",
+ call: func(client *CoordinatorClient) error {
+ return client.TaskStateStore().Set(
+ context.Background(), "job_id", "v",
sdk.WithRetention(time.Hour),
+ )
+ },
+ },
+ {
+ name: "Delete",
+ call: func(client *CoordinatorClient) error {
+ return
client.TaskStateStore().Delete(context.Background(), "job_id")
+ },
+ },
+ {
+ name: "Clear",
+ call: func(client *CoordinatorClient) error {
+ return
client.TaskStateStore().Clear(context.Background())
+ },
+ },
+ }
+ for _, tc := range tests {
+ t.Run(tc.name, func(t *testing.T) {
+ responsePayload := encodeResponseFrame(t, 0, nil,
map[string]any{
+ "type": "ErrorResponse",
+ "error": "API_SERVER_ERROR",
+ "detail": map[string]any{"status_code": 403},
+ })
+ var responseBuf bytes.Buffer
+ require.NoError(t, writeFrame(&responseBuf,
responsePayload))
+
+ logger := slog.New(slog.NewTextHandler(io.Discard, nil))
+ comm := NewCoordinatorComm(&responseBuf, io.Discard,
logger)
+ client := NewCoordinatorClient(comm, testTIID)
+
+ var apiErr *ApiError
+ require.ErrorAs(t, tc.call(client), &apiErr)
+ assert.Equal(t, "API_SERVER_ERROR", apiErr.Err)
+ })
+ }
+}
+
// assertNoReadReader fails the test on any Read call.
type assertNoReadReader struct{ t *testing.T }
@@ -496,7 +1047,10 @@ func TestCoordinatorClientSkipDownstreamTasks(t
*testing.T) {
var requestBuf bytes.Buffer
logger := slog.New(slog.NewTextHandler(io.Discard, nil))
- client :=
NewCoordinatorClient(NewCoordinatorComm(&responseBuf, &requestBuf, logger))
+ client := NewCoordinatorClient(
+ NewCoordinatorComm(&responseBuf, &requestBuf,
logger),
+ testTIID,
+ )
err := client.skipDownstreamTasks(context.Background(),
[]string{"load", "report"})
if tc.wantErr {
@@ -547,7 +1101,10 @@ func TestCoordinatorClientDeleteXCom(t *testing.T) {
var requestBuf bytes.Buffer
logger := slog.New(slog.NewTextHandler(io.Discard, nil))
- client :=
NewCoordinatorClient(NewCoordinatorComm(&responseBuf, &requestBuf, logger))
+ client := NewCoordinatorClient(
+ NewCoordinatorComm(&responseBuf, &requestBuf,
logger),
+ testTIID,
+ )
ti := sdk.TaskInstance{DagID: "d", RunID: "r", TaskID:
"t", MapIndex: tc.mapIndex}
err := client.deleteXCom(context.Background(), ti,
"skipmixin_key")
diff --git a/go-sdk/pkg/execution/frames.go b/go-sdk/pkg/execution/frames.go
index e1124c1b80d..c1f5e73052c 100644
--- a/go-sdk/pkg/execution/frames.go
+++ b/go-sdk/pkg/execution/frames.go
@@ -60,10 +60,7 @@ func encodeRequest(id int64, body any) ([]byte, error) {
body = genmodels.EnsureType(body)
var buf bytes.Buffer
- enc := msgpack.NewEncoder(&buf)
- enc.UseCompactInts(true)
- // Use JSON field names for user values; explicit msgpack tags still
win.
- enc.SetCustomStructTag("json")
+ enc := newFrameEncoder(&buf)
if err := enc.EncodeArrayLen(2); err != nil {
return nil, err
@@ -77,6 +74,15 @@ func encodeRequest(id int64, body any) ([]byte, error) {
return buf.Bytes(), nil
}
+// newFrameEncoder builds the encoder every outbound frame is written with.
+// Use JSON field names for user values; explicit msgpack tags still win.
+func newFrameEncoder(w io.Writer) *msgpack.Encoder {
+ enc := msgpack.NewEncoder(w)
+ enc.UseCompactInts(true)
+ enc.SetCustomStructTag("json")
+ return enc
+}
+
// writeFrame writes a length-prefixed msgpack payload to the writer.
// Format: [4-byte big-endian length][payload bytes].
//
diff --git a/go-sdk/pkg/execution/integration_test.go
b/go-sdk/pkg/execution/integration_test.go
index edcc07145fa..664563511d4 100644
--- a/go-sdk/pkg/execution/integration_test.go
+++ b/go-sdk/pkg/execution/integration_test.go
@@ -37,6 +37,7 @@ import (
"github.com/apache/airflow/go-sdk/internal/contexttest"
"github.com/apache/airflow/go-sdk/pkg/binding"
"github.com/apache/airflow/go-sdk/pkg/execution/genmodels"
+ "github.com/apache/airflow/go-sdk/sdk"
)
// assertSucceedTask asserts RunTask produced a terminal SucceedTask body.
@@ -553,6 +554,54 @@ func TestRunTaskInjectsAirflowContext(t *testing.T) {
assert.Equal(t, end, *dagRun.DataIntervalEnd)
}
+// Guards against the runtime binding the task state store to an empty task
instance id.
+func TestRunTaskBindsTaskStateStoreClient(t *testing.T) {
+ const tiID = "0199e0e5-1b2c-7c3d-8e4f-5a6b7c8d9e0f"
+
+ // A default-retention write needs the setting the supervisor passes.
+ t.Setenv(defaultRetentionDaysEnv, "30")
+
+ var got sdk.TaskStateStoreClient
+ bundle := buildBundle(t, func(r testBundle) {
+ r.AddDag("test_dag").AddTaskWithName("statestore",
+ func(actx contexttest.Context) error {
+ got = actx.Client()
+ return actx.Client().TaskStateStore().Set(actx,
"job_id", "abc123")
+ })
+ })
+
+ details := &genmodels.StartupDetails{
+ TI: genmodels.TaskInstance{
+ ID: tiID,
+ DagID: "test_dag",
+ TaskID: "statestore",
+ RunID: "run1",
+ MapIndex: ptr(-1),
+ },
+ BundleInfo: genmodels.BundleInfo{Name: "test", Version: "1.0"},
+ }
+
+ responsePayload := encodeResponseFrame(t, 0, map[string]any{"type":
"OKResponse"}, nil)
+ var responseBuf bytes.Buffer
+ require.NoError(t, writeFrame(&responseBuf, responsePayload))
+
+ var requestBuf bytes.Buffer
+ logger := slog.New(slog.NewTextHandler(io.Discard, nil))
+ comm := NewCoordinatorComm(&responseBuf, &requestBuf, logger)
+
+ result := RunTask(context.Background(), bundle, details, comm, logger)
+ assertSucceedTask(t, result)
+
+ require.NotNil(t, got, "the task must reach a coordinator-backed task
state store")
+
+ sent, err := readFrame(&requestBuf)
+ require.NoError(t, err)
+ sentMap := rawToMap(t, sent.Body)
+ assert.Equal(t, "SetTaskStateStore", sentMap["type"])
+ assert.Equal(t, tiID, sentMap["ti_id"],
+ "the runtime must bind the store to the started task instance")
+}
+
// Serve traps SIGINT/SIGTERM into the context it hands RunTask, so a
// supervisor shutdown reaches the handler on actx.Done().
func TestRunTaskAirflowContextHonorsShutdown(t *testing.T) {
diff --git a/go-sdk/pkg/execution/task_runner.go
b/go-sdk/pkg/execution/task_runner.go
index 509e9d0e7ba..45bd228ff5a 100644
--- a/go-sdk/pkg/execution/task_runner.go
+++ b/go-sdk/pkg/execution/task_runner.go
@@ -63,7 +63,7 @@ func RunTask(
}
}
- client := NewCoordinatorClient(comm)
+ client := NewCoordinatorClient(comm, details.TI.ID)
// runtimeContext carries the task instance and Dag run that binding
puts
// on the task's airflow.Context. The scheduling timestamps live on the
diff --git a/go-sdk/sdk/doc.go b/go-sdk/sdk/doc.go
index 332c33cbef5..a2961861357 100644
--- a/go-sdk/sdk/doc.go
+++ b/go-sdk/sdk/doc.go
@@ -17,7 +17,7 @@
/*
Package sdk gives task functions access to the Airflow "model" (Variables,
-Connections, and XCom) at run time.
+Connections, XCom, and the per-task-instance state store) at run time.
A task function does not construct a client itself. It takes an
[github.com/apache/airflow/go-sdk/airflow.Context] as its first parameter and
@@ -37,9 +37,13 @@ gets the client from it:
interface. That documents what the helper touches and makes it easy to pass a
fake in unit tests.
+[TaskStateStoreClient] keeps values across attempts of the same task instance,
+so a task can resume its progress after a retry.
+
To publish a result, return a value from the task function: the runtime pushes
it as the task's return-value XCom, so most tasks never call [XComClient]
directly. Lookups that miss return a wrapped sentinel error
([VariableNotFound],
-[ConnectionNotFound], [XComNotFound]) you can test for with errors.Is.
+[ConnectionNotFound], [XComNotFound], [TaskStateNotFound]) you can test for
with
+errors.Is.
*/
package sdk
diff --git a/go-sdk/sdk/errors.go b/go-sdk/sdk/errors.go
index cba5fdccfa6..d6b42fbfbfd 100644
--- a/go-sdk/sdk/errors.go
+++ b/go-sdk/sdk/errors.go
@@ -38,3 +38,9 @@ var ConnectionNotFound = errors.New("connection not found")
//
// See the “GetXCom“ method of [XComClient] for an example
var XComNotFound = errors.New("xcom not found")
+
+// TaskStateNotFound is an error value used to signal that a task state value
could not be found (and that
+// there were no communication issues with the API server).
+//
+// See the “Get“ method of [TaskStateStore] for an example
+var TaskStateNotFound = errors.New("task state not found")
diff --git a/go-sdk/sdk/sdk.go b/go-sdk/sdk/sdk.go
index fdfa677268c..75d170c3eb8 100644
--- a/go-sdk/sdk/sdk.go
+++ b/go-sdk/sdk/sdk.go
@@ -17,7 +17,11 @@
package sdk
-import "context"
+import (
+ "context"
+ "math"
+ "time"
+)
const (
// VariableEnvPrefix is the environment-variable prefix used as a local
@@ -35,6 +39,11 @@ const (
XComReturnValueKey = "return_value"
)
+// NeverExpire, passed to [WithRetention], exempts a task state key from
expiry:
+// it is kept until deleted or until its Dag run is removed. It is the Go
+// equivalent of the Python SDK's “NEVER_EXPIRE“.
+const NeverExpire = time.Duration(math.MaxInt64)
+
// VariableClient reads, writes, and deletes Airflow Variables.
//
// Go has no function overloading, so the "give me the raw string" and
@@ -120,12 +129,84 @@ type XComClient interface {
PushXCom(ctx context.Context, ti TaskInstance, key string, value any)
error
}
+// TaskStateStoreClient exposes the task state store of the running task
+// instance.
+type TaskStateStoreClient interface {
+ // TaskStateStore returns the store scoped to this task instance.
+ TaskStateStore() TaskStateStore
+}
+
+// TaskStateStore reads and writes a key/value store private to the running
task
+// instance. The store is keyed by dag_id, run_id, task_id, and map_index but
+// not try_number, so a value written by one attempt is readable by the next —
a
+// task can record progress and resume after a retry.
+type TaskStateStore interface {
+ // Get returns the value stored under key for this task instance.
+ //
+ // If the key is not found error will be a wrapped
``TaskStateNotFound``:
+ //
+ // store := client.TaskStateStore()
+ // val, err := store.Get(ctx, "checkpoint")
+ // if errors.Is(err, TaskStateNotFound) {
+ // // Handle not found, set default,
return custom error etc
+ // } else {
+ // // Other errors here, such as transport
timeouts etc.
+ // }
+ Get(ctx context.Context, key string) (any, error)
+
+ // UnmarshalJSONValue fetches a task state value and unmarshals it into
+ // pointer via json.Unmarshal. Use it for values stored as JSON objects
or
+ // arrays; pointer must be a non-nil pointer.
+ //
+ // The name keeps the UnmarshalJSON prefix of [VariableClient] without
+ // colliding with encoding/json's UnmarshalJSON, whose signature go vet
+ // enforces on any method of that name.
+ UnmarshalJSONValue(ctx context.Context, key string, pointer any) error
+
+ // Set stores value under key, creating or replacing the entry. Without
+ // [WithRetention] the key expires per the deployment's
+ // “[state_store] default_retention_days“; if that setting is invalid or
+ // absent the write fails.
+ //
+ // value must be non-nil and built from strings, numbers, bools, slices,
+ // string-keyed maps, and structs. A time.Time, a byte slice or array,
or a
+ // non-finite float is rejected before it is sent.
+ Set(ctx context.Context, key string, value any, opts ...SetOption) error
+
+ // Delete removes the value stored under key. Deleting a key that does
not
+ // exist is not an error.
+ Delete(ctx context.Context, key string) error
+
+ // Clear removes every key stored for this task instance.
+ Clear(ctx context.Context) error
+}
+
+// SetOptions carries the options a [TaskStateStore.Set] call was given.
+type SetOptions struct {
+ // Retention is nil unless the caller passed [WithRetention], which is
what
+ // tells Set to follow the deployment default instead.
+ Retention *time.Duration
+}
+
+// SetOption overrides a default of [TaskStateStore.Set].
+type SetOption func(*SetOptions)
+
+// WithRetention sets how long a key is kept, counted from the write. It must
be
+// positive or [NeverExpire]; zero or negative is rejected.
+func WithRetention(retention time.Duration) SetOption {
+ return func(opts *SetOptions) {
+ opts.Retention = &retention
+ }
+}
+
// Client is the full task-facing API: read/write Variables, read Connections,
-// and read/write XCom. A task gets one from its airflow.Context by calling
-// actx.Client(). A helper that needs only one capability can take the narrower
-// VariableClient, ConnectionClient, or XComClient instead.
+// read/write XCom, and read/write the task state store. A task gets one from
+// its airflow.Context by calling actx.Client(). A helper that needs only one
+// capability can take the narrower VariableClient, ConnectionClient,
+// XComClient, or TaskStateStoreClient instead.
type Client interface {
VariableClient
ConnectionClient
XComClient
+ TaskStateStoreClient
}
diff --git a/task-sdk/src/airflow/sdk/coordinators/_subprocess.py
b/task-sdk/src/airflow/sdk/coordinators/_subprocess.py
index feed0da740b..5a3e4720f51 100644
--- a/task-sdk/src/airflow/sdk/coordinators/_subprocess.py
+++ b/task-sdk/src/airflow/sdk/coordinators/_subprocess.py
@@ -197,6 +197,7 @@ def _build_runtime_env() -> dict[str, str]:
"AIRFLOW__API__BASE_URL": conf.get("api", "base_url", fallback="/"),
"AIRFLOW__OPERATORS__DEFAULT_DEFERRABLE":
str(conf.getboolean("operators", "default_deferrable")),
"AIRFLOW__TRIGGERER__QUEUES_ENABLED": str(conf.getboolean("triggerer",
"queues_enabled")),
+ "AIRFLOW__STATE_STORE__DEFAULT_RETENTION_DAYS":
conf.get("state_store", "default_retention_days"),
}
diff --git a/task-sdk/tests/task_sdk/coordinators/test_subprocess.py
b/task-sdk/tests/task_sdk/coordinators/test_subprocess.py
index 219bcdd8b51..e150a0122d5 100644
--- a/task-sdk/tests/task_sdk/coordinators/test_subprocess.py
+++ b/task-sdk/tests/task_sdk/coordinators/test_subprocess.py
@@ -877,6 +877,7 @@ class TestPopenActivitySubprocessStart:
{("api", "base_url"): None},
{
"AIRFLOW__API__BASE_URL": "/",
+ "AIRFLOW__STATE_STORE__DEFAULT_RETENTION_DAYS": "30",
"AIRFLOW__OPERATORS__DEFAULT_DEFERRABLE": "False",
"AIRFLOW__TRIGGERER__QUEUES_ENABLED": "False",
},
@@ -887,11 +888,13 @@ class TestPopenActivitySubprocessStart:
("api", "base_url"): "https://airflow.example.com/sub/",
("operators", "default_deferrable"): "true",
("triggerer", "queues_enabled"): "1",
+ ("state_store", "default_retention_days"): "7",
},
{
"AIRFLOW__API__BASE_URL":
"https://airflow.example.com/sub/",
"AIRFLOW__OPERATORS__DEFAULT_DEFERRABLE": "True",
"AIRFLOW__TRIGGERER__QUEUES_ENABLED": "True",
+ "AIRFLOW__STATE_STORE__DEFAULT_RETENTION_DAYS": "7",
},
id="set",
),