Andrushika commented on code in PR #73420:
URL: https://github.com/apache/airflow/pull/73420#discussion_r4073265013


##########
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:
   Looks like a typed nil still gets through here: for `var p *string`, `value 
== nil` is false, the encoder emits msgpack nil, and the decoded-value check 
accepts nil. 
   
   The frame goes out with `value: null` and the Execution API rejects it with 
“value cannot be null”. 
   Rejecting a top-level nil after decoding, plus a `(*string)(nil)` case in 
the tests, would close it.
   
   



##########
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:
   This case only checks `expires_at` is non-nil, so swapping 
`time.Now().UTC().Add(retention)` for any timestamp still passes. 
   
   Taking `time.Now()` before and after the call and asserting the decoded 
`expires_at` falls in `[before+retention, after+retention]` would lock the 
duration.



-- 
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]

Reply via email to