cool9850311 commented on code in PR #73420:
URL: https://github.com/apache/airflow/pull/73420#discussion_r4073708806
##########
go-sdk/pkg/execution/client.go:
##########
@@ -253,3 +315,181 @@ func (c *CoordinatorClient) PushXCom(
_, err := c.comm.Communicate(ctx, msg)
return err
}
+
+// GetTaskState requests a task state value from the supervisor.
+func (c *CoordinatorClient) GetTaskState(ctx context.Context, key string)
(any, error) {
+ resp, err := c.comm.Communicate(
+ ctx,
+ genmodels.GetTaskStateStore{TIID: c.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
+}
+
+// UnmarshalJSONTaskState gets a task state value and unmarshals it into
pointer.
+func (c *CoordinatorClient) UnmarshalJSONTaskState(
+ ctx context.Context,
+ key string,
+ pointer any,
+) error {
+ val, err := c.GetTaskState(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)
+}
+
+// SetTaskState asks the supervisor to store a task state value, expiring it
+// according to the deployment's default retention.
+func (c *CoordinatorClient) SetTaskState(ctx context.Context, key string,
value any) error {
+ expiry, err := resolveDefaultExpiry(time.Now())
+ if err != nil {
+ return err
+ }
+ return c.sendSetTaskState(ctx, key, value, expiry)
+}
+
+// SetTaskStateWithRetention stores a task state value with a caller-chosen
lifetime.
+func (c *CoordinatorClient) SetTaskStateWithRetention(
+ ctx context.Context,
+ key string,
+ value any,
+ retention time.Duration,
+) error {
+ var expiry any
+ switch {
+ // Checked before any arithmetic: adding NeverExpire overflows.
+ case retention == sdk.NeverExpire:
+ expiry = nil
+ case retention <= 0:
+ return fmt.Errorf(
+ "task state retention must be positive or
sdk.NeverExpire, got %s: "+
+ "use SetTaskState to follow the deployment
default, or DeleteTaskState to drop key %q",
+ retention, key,
+ )
+ default:
+ expiry = time.Now().UTC().Add(retention)
+ }
+ return c.sendSetTaskState(ctx, key, value, expiry)
+}
+
+func (c *CoordinatorClient) sendSetTaskState(
+ ctx context.Context,
+ key string,
+ value any,
+ expiry any,
+) error {
+ if value == nil {
Review Comment:
Fixed in 51c5029ff1: the nil check now runs on the decoded value, so a typed
nil like `(*string)(nil)` is rejected before sending. Added typed-nil cases to
`TestCoordinatorClientSetTaskStateRejectsNilValue`.
##########
go-sdk/pkg/execution/client_test.go:
##########
@@ -452,6 +457,501 @@ func TestCoordinatorClientGetXComMapIndex(t *testing.T) {
}
}
+// Only an absent value may fall back; a malformed one must fail as Python's
+// TaskStateStoreAccessor.set does rather than silently use the shipped
default.
+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 falls back", unset: true, want:
now.UTC().AddDate(0, 0, 30)},
+ {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"},
+ {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"},
+ }
+
+ for _, tc := range tests {
+ t.Run(tc.name, func(t *testing.T) {
+ // Setenv registers the restore; only a truly unset
variable falls back.
+ 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 TestCoordinatorClientSetTaskStateRejectsMisconfiguredRetention(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.SetTaskState(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 TestCoordinatorClientSetTaskStateRejectsNonJSONValues(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.SetTaskState(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 TestCoordinatorClientSetTaskStateAcceptsJSONShapes(t *testing.T) {
+ 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.SetTaskState(context.Background(), "job_id", value))
+ assert.NotZero(t, requestBuf.Len())
+ })
+ }
+}
+
+func TestCoordinatorClientGetTaskState(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.GetTaskState(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 TestCoordinatorClientGetTaskStateNotFound(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.GetTaskState(context.Background(), "missing")
+ require.Error(t, err)
+ assert.ErrorIs(t, err, sdk.TaskStateNotFound)
+ assert.Contains(t, err.Error(), "missing")
+}
+
+func TestCoordinatorClientGetTaskStateErrorPassThrough(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.GetTaskState(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 TestCoordinatorClientUnmarshalJSONTaskState(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.UnmarshalJSONTaskState(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.UnmarshalJSONTaskState(context.Background(),
"missing", &got)
+ assert.ErrorIs(t, err, sdk.TaskStateNotFound)
+ })
+}
+
+// expires_at is sent even when null: the supervisor requires the field.
+func TestCoordinatorClientSetTaskState(t *testing.T) {
+ tests := []struct {
+ name string
+ retentionDays string
+ wantExpiresNil bool
+ }{
+ {name: "deployment retention is applied", retentionDays: "7"},
+ {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)
+
+ require.NoError(t,
client.SetTaskState(context.Background(), "job_id", "abc123"))
+
+ 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 {
+ assert.NotNil(t, sentMap["expires_at"])
+ }
+ })
+ }
+}
+
+func TestCoordinatorClientSetTaskStateWithRetention(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)
+
+ err := client.SetTaskStateWithRetention(
+ context.Background(), "job_id", "abc123",
tc.retention,
+ )
+ 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 {
+ assert.NotNil(t, sentMap["expires_at"])
Review Comment:
Fixed in 51c5029ff1: the test now brackets the call with `time.Now()` and
asserts `expires_at` falls within `[before+retention, after+retention]`.
Applied the same to the default-retention case in
`TestCoordinatorClientSetTaskState` (7 days), which had the same `NotNil`-only
check.
--
This is an automated message from the Apache Git Service.
To respond to the message, please log on to GitHub and use the
URL above to go to the specific comment.
To unsubscribe, e-mail: [email protected]
For queries about this service, please contact Infrastructure at:
[email protected]