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 56d5432c1d5 Go SDK: require airflow.Context first and stop injecting
by type (#73394)
56d5432c1d5 is described below
commit 56d5432c1d508d96cafaf05c3e4014c4c7567420
Author: PoAn Yang <[email protected]>
AuthorDate: Sun Sep 20 20:14:40 2026 +0800
Go SDK: require airflow.Context first and stop injecting by type (#73394)
Signed-off-by: PoAn Yang <[email protected]>
---
go-sdk/bundle/bundlev1/registry.go | 19 +-
go-sdk/bundle/bundlev1/registry_test.go | 31 +-
go-sdk/bundle/bundlev1/task_test.go | 214 +++---------
.../bundle/concurrentxcom/concurrentxcom.go | 15 +-
.../bundle/concurrentxcom/concurrentxcom_test.go | 11 +-
go-sdk/example/bundle/main.go | 48 ++-
go-sdk/example/bundle/main_test.go | 31 +-
.../bundle/taskflowbinding/taskflowbinding.go | 61 ++--
.../bundle/taskflowbinding/taskflowbinding_test.go | 56 ++--
.../example/bundle/variablewrite/variablewrite.go | 12 +-
go-sdk/pkg/binding/binding.go | 135 +++-----
go-sdk/pkg/binding/binding_test.go | 359 ++++++++++++---------
go-sdk/pkg/execution/integration_test.go | 130 ++------
go-sdk/pkg/execution/task_runner.go | 11 +-
go-sdk/pkg/sdkcontext/keys.go | 20 +-
go-sdk/sdk/context.go | 40 +--
go-sdk/sdk/doc.go | 20 +-
go-sdk/sdk/sdk.go | 4 +-
kubernetes-tests/lang_sdk/go_example/main.go | 13 +-
19 files changed, 506 insertions(+), 724 deletions(-)
diff --git a/go-sdk/bundle/bundlev1/registry.go
b/go-sdk/bundle/bundlev1/registry.go
index b4745249870..329df9c2b30 100644
--- a/go-sdk/bundle/bundlev1/registry.go
+++ b/go-sdk/bundle/bundlev1/registry.go
@@ -32,19 +32,18 @@ type (
// AddTask registers fn as a task, deriving the task_id from
fn's own
// name (so it must match the @task.stub name in the Python
dag).
//
- // fn is an ordinary Go function whose parameters are injected
by type
- // and may appear in any order. Recognised parameters are:
- // - airflow.Context: everything below on one value, plus the
- // identity of the task instance and its Dag run
- // - context.Context: cancelled when the task is asked to stop
- // - *slog.Logger: writes to the task's Airflow log
- // - sdk.Client (or a narrower sdk.VariableClient /
sdk.ConnectionClient /
- // sdk.XComClient): access to Variables, Connections, and
XCom
+ // fn is an ordinary Go function that takes an airflow.Context
first.
+ // The Context is the task's context.Context, cancelled when
the task is
+ // asked to stop. It carries the task's logger, the client for
Variables,
+ // Connections and XCom, and the identity of the task instance
and its
+ // Dag run. Every parameter after the Context is data, filled
from the
+ // arguments of the Python stub's TaskFlow call.
//
// fn must return either error or (result, error): a non-nil
error fails
// the task, and a non-nil first result is pushed as the task's
- // return-value XCom. Passing a non-function, or a function
whose return
- // signature does not match, panics at registration time.
+ // return-value XCom. Passing a non-function, or a function
whose
+ // parameters or return signature do not match, panics at
registration
+ // time.
AddTask(fn any)
// AddTaskWithName is like AddTask but sets task_id explicitly
instead of
diff --git a/go-sdk/bundle/bundlev1/registry_test.go
b/go-sdk/bundle/bundlev1/registry_test.go
index f405c8f5315..c0333963ae5 100644
--- a/go-sdk/bundle/bundlev1/registry_test.go
+++ b/go-sdk/bundle/bundlev1/registry_test.go
@@ -25,22 +25,24 @@ import (
"github.com/stretchr/testify/suite"
+ "github.com/apache/airflow/go-sdk/airflow"
"github.com/apache/airflow/go-sdk/sdk"
)
-func myTask() error { return nil }
-func myTaskWithArgs(ctx context.Context, logger *slog.Logger, client
sdk.Client) error {
- if ctx == nil || logger == nil || client == nil {
- return errors.New("missing required argument")
- }
- return nil
-}
-func errorTask() error { return errors.New("fail") }
+func myTask(airflow.Context) error { return nil }
+
+func myTaskWithArgs(actx airflow.Context, country string, count int) error {
return nil }
+
+func errorTask(airflow.Context) error { return errors.New("fail") }
-func NotErrorRet() int {
+func NotErrorRet(airflow.Context) int {
return 0
}
+// legacyTask declares its context, logger and client as separate parameters
+// instead of taking an airflow.Context.
+func legacyTask(ctx context.Context, logger *slog.Logger, client sdk.Client)
error { return nil }
+
type RegistrySuite struct {
suite.Suite
reg Registry
@@ -113,6 +115,17 @@ func (s *RegistrySuite) TestAddTask_InvalidReturnType() {
)
}
+func (s *RegistrySuite) TestAddTask_MissingAirflowContextPanics() {
+ s.PanicsWithError(
+ `error registering task "legacyTask" for DAG "dag1": task
function `+
+
`github.com/apache/airflow/go-sdk/bundle/bundlev1.legacyTask: parameter 0 is `+
+ `context.Context, but the first parameter must be
airflow.Context`,
+ func() {
+ s.dag.AddTask(legacyTask)
+ },
+ )
+}
+
func (s *RegistrySuite) TestAddTask_ErrorReturnType() {
s.dag.AddTask(errorTask)
_, exists := s.reg.LookupTask("dag1", "errorTask")
diff --git a/go-sdk/bundle/bundlev1/task_test.go
b/go-sdk/bundle/bundlev1/task_test.go
index d3fcf6043f2..fc901b27a3d 100644
--- a/go-sdk/bundle/bundlev1/task_test.go
+++ b/go-sdk/bundle/bundlev1/task_test.go
@@ -24,6 +24,7 @@ import (
"github.com/stretchr/testify/suite"
+ "github.com/apache/airflow/go-sdk/airflow"
"github.com/apache/airflow/go-sdk/pkg/binding"
"github.com/apache/airflow/go-sdk/pkg/logging"
"github.com/apache/airflow/go-sdk/pkg/sdkcontext"
@@ -52,15 +53,15 @@ func (s *TaskSuite) TestReturnValidation() {
errContains string
}{
"no-ret-values": {
- func() {},
+ func(airflow.Context) {},
`func\d+ has 0 return values, must be`,
},
"too-many-ret-values": {
- func() (a, b, c int) { return },
+ func(airflow.Context) (a, b, c int) { return },
`func\d+ has 3 return values, must be`,
},
"invalid-ret": {
- func() (c chan int) { return },
+ func(airflow.Context) (c chan int) { return },
`func\d+ last return value to return error but found
chan`,
},
}
@@ -75,134 +76,60 @@ func (s *TaskSuite) TestReturnValidation() {
}
}
-func (s *TaskSuite) TestArgumentBinding() {
- cases := map[string]struct {
- fn any
- }{
- "no-args": {
- func() error { return nil },
- },
- "context": {
- func(ctx context.Context) error {
- s.Equal("def", ctx.Value("abc"))
- return nil
- },
- },
- "context-and-logger": {
- func(ctx context.Context, logger *slog.Logger) error {
- s.Equal("def", ctx.Value("abc"))
- s.NotNil(logger)
- return nil
- },
- },
- "client": {
- func(client sdk.Client) error {
- s.NotNil(client)
- return nil
- },
- },
- "var-client": {
- func(client sdk.VariableClient) error {
- s.NotNil(client)
- return nil
- },
- },
- "conn-client": {
- func(client sdk.ConnectionClient) error {
- s.NotNil(client)
-
- return nil
- },
- },
- "xcom-client": {
- func(client sdk.XComClient) error {
- s.NotNil(client)
-
- return nil
- },
- },
- }
+// probeKey is an unexported context key used to confirm the live task context
+// (not a freshly built one) backs the airflow.Context a task receives.
+type probeKeyType struct{}
- for name, tt := range cases {
- s.Run(name, func() {
- task, err := NewTaskFunction(tt.fn)
- s.Require().NoError(err)
+var probeKey probeKeyType
- ctx :=
context.WithValue(withTaskClient(context.Background()), "abc", "def")
- logger := slog.New(logging.NewTeeLogger())
- task.Execute(ctx, logger, nil)
- })
+func (s *TaskSuite) TestExecuteBindsAirflowContext() {
+ mapIndex := 3
+ ti := sdk.TaskInstance{
+ DagID: "dag1",
+ RunID: "run1",
+ TaskID: "task1",
+ MapIndex: &mapIndex,
+ TryNumber: 2,
}
-}
+ dagRun := sdk.DagRun{DagID: "dag1", RunID: "run1"}
-// TestClientSubsetInjection checks any subset of sdk.Client is injected, even
-// an unnamed one.
-func (s *TaskSuite) TestClientSubsetInjection() {
- task, err := NewTaskFunction(func(client interface {
- GetVariable(ctx context.Context, key string) (string, error)
- },
- ) error {
- s.NotNil(client)
+ var got airflow.Context
+ task, err := NewTaskFunction(func(actx airflow.Context) error {
+ got = actx
return nil
})
s.Require().NoError(err)
- s.Require().
- NoError(task.Execute(withTaskClient(context.Background()),
slog.New(logging.NewTeeLogger()), nil))
-}
-func (s *TaskSuite) TestNonInjectableParamsAreRejected() {
- cases := map[string]struct {
- fn any
- errContains string
- }{
- "non-client-method": {
- func(x interface{ NotAClientMethod() }) error { return
nil },
- "sdk.Client has no method NotAClientMethod",
- },
- "wrong-signature": {
- func(x interface {
- GetVariable(key string) (string, error)
- },
- ) error {
- return nil
- },
- "method GetVariable is func(context.Context, string)
(string, error) on sdk.Client",
- },
- "func-param": {
- func(cb func()) error { return nil },
- "cannot receive a task argument",
- },
- "context-with-extra-methods": {
- func(x interface {
- context.Context
- TaskInstance() sdk.TaskInstance
- },
- ) error {
- return nil
- },
- "adds methods on top of context.Context",
- },
- }
+ ctx := context.WithValue(
+ withTaskClient(context.Background()),
+ sdkcontext.RuntimeContextKey,
+ sdk.NewTIRunContext(context.Background(), ti, dagRun),
+ )
+ ctx = context.WithValue(ctx, probeKey, "probe-value")
+ logger := slog.New(logging.NewTeeLogger())
+ s.Require().NoError(task.Execute(ctx, logger, nil))
- for name, tt := range cases {
- s.Run(name, func() {
- _, err := NewTaskFunction(tt.fn)
- if s.Assert().Error(err) {
- s.Assert().Contains(err.Error(), "parameter 0")
- s.Assert().Contains(err.Error(), tt.errContains)
- }
- })
- }
+ s.Same(logger, got.Logger())
+ s.Equal(ctx.Value(sdkcontext.SdkClientContextKey), got.Client())
+ s.Equal(ti, got.TaskInstance())
+ s.Equal(dagRun, got.DagRun())
+ s.Equal(
+ "probe-value",
+ got.Value(probeKey),
+ "the Context must be backed by the one passed to Execute",
+ )
}
func (s *TaskSuite) TestExecuteBindsDataParameters() {
var gotCountry string
var gotMeta map[string]any
- task, err := NewTaskFunction(func(log *slog.Logger, country string,
meta map[string]any) error {
- gotCountry = country
- gotMeta = meta
- return nil
- })
+ task, err := NewTaskFunction(
+ func(actx airflow.Context, country string, meta map[string]any)
error {
+ gotCountry = country
+ gotMeta = meta
+ return nil
+ },
+ )
s.Require().NoError(err)
err = task.Execute(
@@ -219,7 +146,7 @@ func (s *TaskSuite) TestExecuteBindsDataParameters() {
}
func (s *TaskSuite) TestExecuteWithoutSpecFailsForDataParameters() {
- task, err := NewTaskFunction(func(country string) error { return nil })
+ task, err := NewTaskFunction(func(actx airflow.Context, country string)
error { return nil })
s.Require().NoError(err)
err = task.Execute(
@@ -231,7 +158,7 @@ func (s *TaskSuite)
TestExecuteWithoutSpecFailsForDataParameters() {
}
func (s *TaskSuite) TestExecuteArityMismatch() {
- task, err := NewTaskFunction(func(country string) error { return nil })
+ task, err := NewTaskFunction(func(actx airflow.Context, country string)
error { return nil })
s.Require().NoError(err)
err = task.Execute(
@@ -248,55 +175,8 @@ func (s *TaskSuite) TestExecuteArityMismatch() {
}
}
-// probeKey is an unexported context key used to confirm the live task context
-// (not a freshly built one) backs the injected sdk.TIRunContext.
-type probeKeyType struct{}
-
-var probeKey probeKeyType
-
-// TestTIRunContextInjection verifies a task declaring sdk.TIRunContext
receives
-// the TaskInstance/DagRun stored on the context, backed by the live task
-// context so it is usable as a context.Context. It must take precedence over
-// the plain context.Context binding, which sdk.TIRunContext also satisfies.
-func (s *TaskSuite) TestTIRunContextInjection() {
- mapIndex := 3
- ti := sdk.TaskInstance{
- DagID: "dag1",
- RunID: "run1",
- TaskID: "task1",
- MapIndex: &mapIndex,
- TryNumber: 2,
- }
- dagRun := sdk.DagRun{DagID: "dag1", RunID: "run1"}
- stored := sdk.NewTIRunContext(context.Background(), ti, dagRun)
-
- var got sdk.TIRunContext
- task, err := NewTaskFunction(func(ctx sdk.TIRunContext) error {
- got = ctx
- return nil
- })
- s.Require().NoError(err)
-
- ctx := context.WithValue(
- withTaskClient(context.Background()),
- sdkcontext.RuntimeContextKey,
- stored,
- )
- ctx = context.WithValue(ctx, probeKey, "probe-value")
- s.Require().NoError(task.Execute(ctx, slog.New(logging.NewTeeLogger()),
nil))
-
- s.Require().NotNil(got, "the task must receive a non-nil TIRunContext")
- s.Equal(ti, got.TaskInstance())
- s.Equal(dagRun, got.DagRun())
- s.Equal(
- "probe-value",
- got.Value(probeKey),
- "the injected context must be backed by the one passed to
Execute",
- )
-}
-
func (s *TaskSuite) TestExecuteRequiresCoordinatorClient() {
- task, err := NewTaskFunction(func() error { return nil })
+ task, err := NewTaskFunction(func(airflow.Context) error { return nil })
s.Require().NoError(err)
err = task.Execute(context.Background(),
slog.New(logging.NewTeeLogger()), nil)
diff --git a/go-sdk/example/bundle/concurrentxcom/concurrentxcom.go
b/go-sdk/example/bundle/concurrentxcom/concurrentxcom.go
index ed2217908ad..bb810ab65a2 100644
--- a/go-sdk/example/bundle/concurrentxcom/concurrentxcom.go
+++ b/go-sdk/example/bundle/concurrentxcom/concurrentxcom.go
@@ -23,11 +23,11 @@ package concurrentxcom
import (
"errors"
"fmt"
- "log/slog"
"reflect"
"sync"
"time"
+ "github.com/apache/airflow/go-sdk/airflow"
"github.com/apache/airflow/go-sdk/sdk"
)
@@ -38,10 +38,11 @@ const (
)
// PullXComsConcurrently pulls a batch of XComs sequentially then concurrently
-// (one goroutine per item), exercising concurrent reads of the injected
+// (one goroutine per item), exercising concurrent reads of the task's
// sdk.Client, and returns both timings.
-func PullXComsConcurrently(ctx sdk.TIRunContext, client sdk.Client, log
*slog.Logger) (any, error) {
- ti := ctx.TaskInstance()
+func PullXComsConcurrently(actx airflow.Context) (any, error) {
+ client := actx.Client()
+ ti := actx.TaskInstance()
// PushXCom needs only the ids off the TaskInstance, not the UUID.
taskInstance := sdk.TaskInstance{
DagID: ti.DagID,
@@ -53,13 +54,13 @@ func PullXComsConcurrently(ctx sdk.TIRunContext, client
sdk.Client, log *slog.Lo
keys := make([]string, numXComs)
for i := range keys {
keys[i] = fmt.Sprintf("item_%d", i)
- if err := client.PushXCom(ctx, taskInstance, keys[i], i); err
!= nil {
+ if err := client.PushXCom(actx, taskInstance, keys[i], i); err
!= nil {
return nil, fmt.Errorf("seeding xcom %s: %w", keys[i],
err)
}
}
pull := func(key string) (any, error) {
- v, err := client.GetXCom(ctx, ti.DagID, ti.RunID, ti.TaskID,
nil, key, nil)
+ v, err := client.GetXCom(actx, ti.DagID, ti.RunID, ti.TaskID,
nil, key, nil)
if err != nil {
return nil, err
}
@@ -106,7 +107,7 @@ func PullXComsConcurrently(ctx sdk.TIRunContext, client
sdk.Client, log *slog.Lo
}
}
- log.InfoContext(ctx, "pulled xcoms concurrently",
+ actx.Logger().InfoContext(actx, "pulled xcoms concurrently",
"num_xcoms", numXComs,
"sequential_ms", sequential.Milliseconds(),
"concurrent_ms", concurrent.Milliseconds(),
diff --git a/go-sdk/example/bundle/concurrentxcom/concurrentxcom_test.go
b/go-sdk/example/bundle/concurrentxcom/concurrentxcom_test.go
index f087fbbb3da..4421fb48b61 100644
--- a/go-sdk/example/bundle/concurrentxcom/concurrentxcom_test.go
+++ b/go-sdk/example/bundle/concurrentxcom/concurrentxcom_test.go
@@ -25,6 +25,7 @@ import (
"github.com/stretchr/testify/assert"
+ "github.com/apache/airflow/go-sdk/airflow"
"github.com/apache/airflow/go-sdk/sdk"
)
@@ -90,17 +91,19 @@ func (m *mockXComClient) GetConnection(ctx context.Context,
connID string) (sdk.
var _ sdk.Client = (*mockXComClient)(nil)
func Test_PullXComsConcurrently(t *testing.T) {
- ctx := sdk.NewTIRunContext(
+ actx := airflow.NewContext(
context.Background(),
- sdk.TaskInstance{
+ slog.Default(),
+ newMockXComClient(),
+ airflow.TaskInstance{
DagID: "concurrent_xcom_dag",
RunID: "run",
TaskID: "pull_xcoms_concurrently",
},
- sdk.DagRun{},
+ airflow.DagRun{},
)
- result, err := PullXComsConcurrently(ctx, newMockXComClient(),
slog.Default())
+ result, err := PullXComsConcurrently(actx)
assert.NoError(t, err)
m, ok := result.(map[string]any)
diff --git a/go-sdk/example/bundle/main.go b/go-sdk/example/bundle/main.go
index e01e29a6447..5422a8fa527 100644
--- a/go-sdk/example/bundle/main.go
+++ b/go-sdk/example/bundle/main.go
@@ -24,12 +24,12 @@ import (
"runtime"
"time"
+ "github.com/apache/airflow/go-sdk/airflow"
v1 "github.com/apache/airflow/go-sdk/bundle/bundlev1"
"github.com/apache/airflow/go-sdk/bundle/bundlev1/bundlev1server"
"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/variablewrite"
- "github.com/apache/airflow/go-sdk/sdk"
)
type myBundle struct{}
@@ -74,17 +74,18 @@ func main() {
}
}
-func extract(ctx sdk.TIRunContext, client sdk.Client, log *slog.Logger) (any,
error) {
+func extract(actx airflow.Context) (any, error) {
+ log := actx.Logger()
log.Info("Hello from task")
- // ctx behaves as a context.Context and also carries the task instance
+ // actx behaves as a context.Context and also carries the task instance
// identifiers and the Dag run's scheduling timestamps. Log every field
the
// runtime context exposes. The fields are namespaced under a "context"
// group (so they serialise as context.ti.* / context.dag_run.* dotted
// keys) to avoid colliding with the reserved task_id/run_id/etc. keys
the
// supervisor strips from its log view.
- ti, dagRun := ctx.TaskInstance(), ctx.DagRun()
- log.InfoContext(ctx, "task runtime context",
+ ti, dagRun := actx.TaskInstance(), actx.DagRun()
+ log.InfoContext(actx, "task runtime context",
slog.Group("context",
slog.Group("ti",
"dag_id", ti.DagID,
@@ -103,13 +104,13 @@ func extract(ctx sdk.TIRunContext, client sdk.Client, log
*slog.Logger) (any, er
),
)
- conn, err := client.GetConnection(ctx, "test_http")
+ conn, err := actx.Client().GetConnection(actx, "test_http")
if err != nil {
- log.ErrorContext(ctx, "unable to get conn", "error", err)
+ log.ErrorContext(actx, "unable to get conn", "error", err)
} else {
// Log only non-sensitive fields; conn.Password and any secrets
in
// conn.Extra must never reach the log stream.
- log.InfoContext(ctx, "got conn",
+ log.InfoContext(actx, "got conn",
"conn_id", conn.ID,
"conn_type", conn.Type,
"host", conn.Host,
@@ -120,8 +121,8 @@ func extract(ctx sdk.TIRunContext, client sdk.Client, log
*slog.Logger) (any, er
// Once per loop,.check if we've been asked to cancel!
select {
- case <-ctx.Done():
- return nil, ctx.Err()
+ case <-actx.Done():
+ return nil, actx.Err()
default:
}
log.Info("After the beep the time will be", "time", time.Now())
@@ -138,25 +139,16 @@ func extract(ctx sdk.TIRunContext, client sdk.Client, log
*slog.Logger) (any, er
}
// transform receives the stub call's literal and XCom arguments.
-func transform(
- ctx sdk.TIRunContext,
- client sdk.VariableClient,
- log *slog.Logger,
- country string,
- extracted map[string]any,
-) error {
- // This function takes a VariableClient and not a Client to make unit
testing it easier. See
- // `./main_test.go` for an example unit of this task fn. Functionally
taking a `sdk.Client` is the same (as
- // Client includes VariableClient) but by using the dedicated type it
can be easier to write unit tests.
- //
- // It also gives a better indication of what features the tasks use
- log.InfoContext(ctx, "Bound TaskFlow arguments",
+// See `./main_test.go` for an example unit test of this task fn.
+func transform(actx airflow.Context, country string, extracted map[string]any)
error {
+ log := actx.Logger()
+ log.InfoContext(actx, "Bound TaskFlow arguments",
"country", country,
"extracted_go_version", extracted["go_version"],
"extracted_timestamp", extracted["timestamp"],
)
key := "my_variable"
- val, err := client.GetVariable(ctx, key)
+ val, err := actx.Client().GetVariable(actx, key)
if err != nil {
return err
}
@@ -169,12 +161,12 @@ func transform(
// task UP_FOR_RETRY -- which only works because the Go SDK now emits a
// RetryTask frame (instead of a terminal FAILED) when ti_context.should_retry
// is set. The retry then runs this task again and it returns nil.
-func load(ctx sdk.TIRunContext, log *slog.Logger) error {
- tryNumber := ctx.TaskInstance().TryNumber
+func load(actx airflow.Context) error {
+ tryNumber := actx.TaskInstance().TryNumber
if tryNumber == 1 {
- log.InfoContext(ctx, "Please fail", "try_number", tryNumber)
+ actx.Logger().InfoContext(actx, "Please fail", "try_number",
tryNumber)
return fmt.Errorf("Please fail")
}
- log.InfoContext(ctx, "Recovered on retry", "try_number", tryNumber)
+ actx.Logger().InfoContext(actx, "Recovered on retry", "try_number",
tryNumber)
return nil
}
diff --git a/go-sdk/example/bundle/main_test.go
b/go-sdk/example/bundle/main_test.go
index 0503495d90e..648163bd4aa 100644
--- a/go-sdk/example/bundle/main_test.go
+++ b/go-sdk/example/bundle/main_test.go
@@ -24,15 +24,17 @@ import (
"github.com/stretchr/testify/assert"
+ "github.com/apache/airflow/go-sdk/airflow"
"github.com/apache/airflow/go-sdk/sdk"
)
// This file serves as an example of how you could write unit tests against
your own Go Tasks.
// An example of how to write a test for a Task function!
-type mockVars struct{}
+// mockVars embeds sdk.Client so it satisfies the whole interface while only
defining
+// the method this test expects. Calling any other method panics on the nil
embedded client.
+type mockVars struct{ sdk.Client }
-// GetVariable implements sdk.VariableClient.
func (m *mockVars) GetVariable(ctx context.Context, key string) (string,
error) {
switch key {
case "my_variable":
@@ -42,27 +44,12 @@ func (m *mockVars) GetVariable(ctx context.Context, key
string) (string, error)
}
}
-// UnmarshalJSONVariable implements sdk.VariableClient.
-func (m *mockVars) UnmarshalJSONVariable(ctx context.Context, key string,
pointer any) error {
- panic("unimplemented")
-}
-
-// SetVariable implements sdk.VariableClient.
-func (m *mockVars) SetVariable(ctx context.Context, key, value, description
string) error {
- panic("unimplemented")
-}
-
-// DeleteVariable implements sdk.VariableClient.
-func (m *mockVars) DeleteVariable(ctx context.Context, key string) error {
- panic("unimplemented")
-}
-
-var _ sdk.VariableClient = (*mockVars)(nil)
-
func Test_transform(t *testing.T) {
- log := slog.Default()
// This is not the best test, but it is a good proof of concept -- you
can just call the function.
- ctx := sdk.NewTIRunContext(context.Background(), sdk.TaskInstance{},
sdk.DagRun{})
- err := transform(ctx, &mockVars{}, log, "uk",
map[string]any{"go_version": "go1.24"})
+ actx := airflow.NewContext(
+ context.Background(), slog.Default(), &mockVars{},
+ airflow.TaskInstance{}, airflow.DagRun{},
+ )
+ err := transform(actx, "uk", map[string]any{"go_version": "go1.24"})
assert.NoError(t, err)
}
diff --git a/go-sdk/example/bundle/taskflowbinding/taskflowbinding.go
b/go-sdk/example/bundle/taskflowbinding/taskflowbinding.go
index 2e2ed740cd7..0b1556fca14 100644
--- a/go-sdk/example/bundle/taskflowbinding/taskflowbinding.go
+++ b/go-sdk/example/bundle/taskflowbinding/taskflowbinding.go
@@ -20,10 +20,9 @@ package taskflowbinding
import (
"fmt"
- "log/slog"
"reflect"
- "github.com/apache/airflow/go-sdk/sdk"
+ "github.com/apache/airflow/go-sdk/airflow"
)
// Config is an object passed through XCom.
@@ -34,9 +33,10 @@ type Config struct {
}
// MakeConfig returns a Config XCom.
-func MakeConfig(log *slog.Logger) (any, error) {
+func MakeConfig(actx airflow.Context) (any, error) {
cfg := Config{Environment: "production", Region: "eu-west-1", Debug:
true}
- log.Info(
+ actx.Logger().InfoContext(
+ actx,
"Pushing config",
"environment",
cfg.Environment,
@@ -49,23 +49,22 @@ func MakeConfig(log *slog.Logger) (any, error) {
}
// MakeNumbers returns an integer-slice XCom.
-func MakeNumbers(log *slog.Logger) (any, error) {
+func MakeNumbers(actx airflow.Context) (any, error) {
numbers := []int{1, 1, 2, 3, 5, 8}
- log.Info("Pushing numbers", "numbers", fmt.Sprint(numbers))
+ actx.Logger().InfoContext(actx, "Pushing numbers", "numbers",
fmt.Sprint(numbers))
return numbers, nil
}
// MakeRegion returns a region XCom.
-func MakeRegion(log *slog.Logger) (any, error) {
+func MakeRegion(actx airflow.Context) (any, error) {
region := "eu-west-1"
- log.Info("Pushing region", "region", region)
+ actx.Logger().InfoContext(actx, "Pushing region", "region", region)
return region, nil
}
// ViaFlatArgs exercises positional binding for literals, defaults, and XComs.
func ViaFlatArgs(
- ctx sdk.TIRunContext,
- log *slog.Logger,
+ actx airflow.Context,
name string,
count int,
ratio float64,
@@ -98,7 +97,7 @@ func ViaFlatArgs(
for _, n := range numbers {
sum += n
}
- log.InfoContext(ctx, "Bound TaskFlow arguments",
+ actx.Logger().InfoContext(actx, "Bound TaskFlow arguments",
"name", name,
"count", count,
"ratio", ratio,
@@ -127,11 +126,7 @@ type ViaStructNoTagsInput struct {
}
// ViaStructNoTags exercises binding without `arg:` tags.
-func ViaStructNoTags(
- ctx sdk.TIRunContext,
- log *slog.Logger,
- input ViaStructNoTagsInput,
-) (any, error) {
+func ViaStructNoTags(actx airflow.Context, input ViaStructNoTagsInput) (any,
error) {
if input.RegionCode != "eu-west-1" || input.Threshold != 0.75 {
return nil, fmt.Errorf(
"struct fields bound incorrectly: region_code=%q
threshold=%v",
@@ -140,7 +135,7 @@ func ViaStructNoTags(
)
}
- log.InfoContext(ctx, "Bound struct (no tags)",
+ actx.Logger().InfoContext(actx, "Bound struct (no tags)",
"region_code", input.RegionCode,
"threshold", input.Threshold,
)
@@ -157,11 +152,7 @@ type ViaStructArgTagInput struct {
}
// ViaStructArgTag exercises explicit field-name binding.
-func ViaStructArgTag(
- ctx sdk.TIRunContext,
- log *slog.Logger,
- input ViaStructArgTagInput,
-) (any, error) {
+func ViaStructArgTag(actx airflow.Context, input ViaStructArgTagInput) (any,
error) {
if input.Region != "eu-west-1" || input.Threshold != 0.75 {
return nil, fmt.Errorf(
"struct fields bound incorrectly: region=%q
threshold=%v",
@@ -170,7 +161,7 @@ func ViaStructArgTag(
)
}
- log.InfoContext(ctx, "Bound struct (arg: tag)",
+ actx.Logger().InfoContext(actx, "Bound struct (arg: tag)",
"region", input.Region,
"threshold", input.Threshold,
)
@@ -188,7 +179,8 @@ type ViaStructUnmatchedArgInput struct {
// ViaStructUnmatchedArg exercises unmatched fields and captured defaults.
func ViaStructUnmatchedArg(
- ctx sdk.TIRunContext, log *slog.Logger, input
ViaStructUnmatchedArgInput,
+ actx airflow.Context,
+ input ViaStructUnmatchedArgInput,
) (any, error) {
if input.Region != "eu-west-1" {
return nil, fmt.Errorf("struct field bound incorrectly:
region=%q", input.Region)
@@ -200,7 +192,7 @@ func ViaStructUnmatchedArg(
)
}
- log.InfoContext(ctx, "Bound struct (unmatched arg)",
+ actx.Logger().InfoContext(actx, "Bound struct (unmatched arg)",
"region", input.Region,
"missing_was_empty", input.Missing == "",
)
@@ -217,9 +209,7 @@ type FlatMapConfig struct {
}
// ViaFlatMap exercises whole-value dict decoding into a struct.
-func ViaFlatMap(
- ctx sdk.TIRunContext, log *slog.Logger, config FlatMapConfig,
-) (any, error) {
+func ViaFlatMap(actx airflow.Context, config FlatMapConfig) (any, error) {
if config.Region != "eu-west-1" || config.Count != 3 {
return nil, fmt.Errorf(
"whole-value map bound incorrectly: region=%q count=%d",
@@ -228,7 +218,7 @@ func ViaFlatMap(
)
}
- log.InfoContext(ctx, "Bound whole map into struct",
+ actx.Logger().InfoContext(actx, "Bound whole map into struct",
"region", config.Region,
"count", config.Count,
)
@@ -241,26 +231,23 @@ type StructMapInput struct {
}
// ViaStructMap exercises dict binding to a struct field.
-func ViaStructMap(
- ctx sdk.TIRunContext, log *slog.Logger, input StructMapInput,
-) (any, error) {
+func ViaStructMap(actx airflow.Context, input StructMapInput) (any, error) {
region, _ := input.Payload["region"].(string)
if region != "eu-west-1" {
return nil, fmt.Errorf("map field bound incorrectly:
payload=%v", input.Payload)
}
- log.InfoContext(ctx, "Bound map onto struct field", "payload",
fmt.Sprint(input.Payload))
+ actx.Logger().
+ InfoContext(actx, "Bound map onto struct field", "payload",
fmt.Sprint(input.Payload))
return map[string]any{"payload": input.Payload}, nil
}
// ViaPlainMap exercises dict decoding into a typed map.
-func ViaPlainMap(
- ctx sdk.TIRunContext, log *slog.Logger, labels map[string]string,
-) (any, error) {
+func ViaPlainMap(actx airflow.Context, labels map[string]string) (any, error) {
if labels["team"] != "data" || labels["tier"] != "gold" {
return nil, fmt.Errorf("plain map bound incorrectly:
labels=%v", labels)
}
- log.InfoContext(ctx, "Bound dict into a plain map", "labels",
fmt.Sprint(labels))
+ actx.Logger().InfoContext(actx, "Bound dict into a plain map",
"labels", fmt.Sprint(labels))
return map[string]any{"team": labels["team"], "tier": labels["tier"]},
nil
}
diff --git a/go-sdk/example/bundle/taskflowbinding/taskflowbinding_test.go
b/go-sdk/example/bundle/taskflowbinding/taskflowbinding_test.go
index 02a3d6c75ad..beb046aa29a 100644
--- a/go-sdk/example/bundle/taskflowbinding/taskflowbinding_test.go
+++ b/go-sdk/example/bundle/taskflowbinding/taskflowbinding_test.go
@@ -18,19 +18,30 @@
package taskflowbinding
import (
- "context"
"log/slog"
"testing"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
+ "github.com/apache/airflow/go-sdk/airflow"
"github.com/apache/airflow/go-sdk/sdk"
)
+// unusedClient satisfies airflow.NewContext, which rejects a nil client.
+// These handlers never call it.
+type unusedClient struct{ sdk.Client }
+
+func testContext(t *testing.T) airflow.Context {
+ t.Helper()
+ return airflow.NewContext(
+ t.Context(), slog.Default(), unusedClient{},
+ airflow.TaskInstance{}, airflow.DagRun{},
+ )
+}
+
func TestViaFlatArgs(t *testing.T) {
- ctx := sdk.NewTIRunContext(context.Background(), sdk.TaskInstance{},
sdk.DagRun{})
- got, err := ViaFlatArgs(ctx, slog.Default(),
+ got, err := ViaFlatArgs(testContext(t),
"summary", 3, 2.5, true,
[]string{"metrics", "hourly"},
Config{Environment: "production", Region: "eu-west-1", Debug:
true},
@@ -46,8 +57,7 @@ func TestViaFlatArgs(t *testing.T) {
}
func TestViaFlatArgsRejectsWrongBinding(t *testing.T) {
- ctx := sdk.NewTIRunContext(context.Background(), sdk.TaskInstance{},
sdk.DagRun{})
- _, err := ViaFlatArgs(ctx, slog.Default(),
+ _, err := ViaFlatArgs(testContext(t),
"summary", 3, 2.5, true,
[]string{"metrics", "hourly"},
Config{},
@@ -58,8 +68,7 @@ func TestViaFlatArgsRejectsWrongBinding(t *testing.T) {
}
func TestViaStructNoTags(t *testing.T) {
- ctx := sdk.NewTIRunContext(context.Background(), sdk.TaskInstance{},
sdk.DagRun{})
- got, err := ViaStructNoTags(ctx, slog.Default(), ViaStructNoTagsInput{
+ got, err := ViaStructNoTags(testContext(t), ViaStructNoTagsInput{
RegionCode: "eu-west-1",
Threshold: 0.75,
})
@@ -71,8 +80,7 @@ func TestViaStructNoTags(t *testing.T) {
}
func TestViaStructNoTagsRejectsWrongBinding(t *testing.T) {
- ctx := sdk.NewTIRunContext(context.Background(), sdk.TaskInstance{},
sdk.DagRun{})
- _, err := ViaStructNoTags(ctx, slog.Default(), ViaStructNoTagsInput{
+ _, err := ViaStructNoTags(testContext(t), ViaStructNoTagsInput{
RegionCode: "wrong-region",
Threshold: 0.75,
})
@@ -80,8 +88,7 @@ func TestViaStructNoTagsRejectsWrongBinding(t *testing.T) {
}
func TestViaStructArgTag(t *testing.T) {
- ctx := sdk.NewTIRunContext(context.Background(), sdk.TaskInstance{},
sdk.DagRun{})
- got, err := ViaStructArgTag(ctx, slog.Default(), ViaStructArgTagInput{
+ got, err := ViaStructArgTag(testContext(t), ViaStructArgTagInput{
Region: "eu-west-1",
Threshold: 0.75,
})
@@ -93,8 +100,7 @@ func TestViaStructArgTag(t *testing.T) {
}
func TestViaStructArgTagRejectsWrongBinding(t *testing.T) {
- ctx := sdk.NewTIRunContext(context.Background(), sdk.TaskInstance{},
sdk.DagRun{})
- _, err := ViaStructArgTag(ctx, slog.Default(), ViaStructArgTagInput{
+ _, err := ViaStructArgTag(testContext(t), ViaStructArgTagInput{
Region: "wrong-region",
Threshold: 0.75,
})
@@ -102,8 +108,7 @@ func TestViaStructArgTagRejectsWrongBinding(t *testing.T) {
}
func TestViaStructUnmatchedArg(t *testing.T) {
- ctx := sdk.NewTIRunContext(context.Background(), sdk.TaskInstance{},
sdk.DagRun{})
- got, err := ViaStructUnmatchedArg(ctx, slog.Default(),
ViaStructUnmatchedArgInput{
+ got, err := ViaStructUnmatchedArg(testContext(t),
ViaStructUnmatchedArgInput{
Region: "eu-west-1",
Missing: "",
})
@@ -115,8 +120,7 @@ func TestViaStructUnmatchedArg(t *testing.T) {
}
func TestViaStructUnmatchedArgRejectsNonZeroMissingField(t *testing.T) {
- ctx := sdk.NewTIRunContext(context.Background(), sdk.TaskInstance{},
sdk.DagRun{})
- _, err := ViaStructUnmatchedArg(ctx, slog.Default(),
ViaStructUnmatchedArgInput{
+ _, err := ViaStructUnmatchedArg(testContext(t),
ViaStructUnmatchedArgInput{
Region: "eu-west-1",
Missing: "unexpected",
})
@@ -124,8 +128,7 @@ func TestViaStructUnmatchedArgRejectsNonZeroMissingField(t
*testing.T) {
}
func TestViaFlatMap(t *testing.T) {
- ctx := sdk.NewTIRunContext(context.Background(), sdk.TaskInstance{},
sdk.DagRun{})
- got, err := ViaFlatMap(ctx, slog.Default(), FlatMapConfig{Region:
"eu-west-1", Count: 3})
+ got, err := ViaFlatMap(testContext(t), FlatMapConfig{Region:
"eu-west-1", Count: 3})
require.NoError(t, err)
summary, ok := got.(map[string]any)
@@ -135,14 +138,12 @@ func TestViaFlatMap(t *testing.T) {
}
func TestViaFlatMapRejectsWrongBinding(t *testing.T) {
- ctx := sdk.NewTIRunContext(context.Background(), sdk.TaskInstance{},
sdk.DagRun{})
- _, err := ViaFlatMap(ctx, slog.Default(), FlatMapConfig{Region:
"wrong-region", Count: 3})
+ _, err := ViaFlatMap(testContext(t), FlatMapConfig{Region:
"wrong-region", Count: 3})
assert.ErrorContains(t, err, "whole-value map bound incorrectly")
}
func TestViaStructMap(t *testing.T) {
- ctx := sdk.NewTIRunContext(context.Background(), sdk.TaskInstance{},
sdk.DagRun{})
- got, err := ViaStructMap(ctx, slog.Default(), StructMapInput{
+ got, err := ViaStructMap(testContext(t), StructMapInput{
Payload: map[string]any{"region": "eu-west-1", "count": 3},
})
require.NoError(t, err)
@@ -153,16 +154,14 @@ func TestViaStructMap(t *testing.T) {
}
func TestViaStructMapRejectsWrongBinding(t *testing.T) {
- ctx := sdk.NewTIRunContext(context.Background(), sdk.TaskInstance{},
sdk.DagRun{})
- _, err := ViaStructMap(ctx, slog.Default(), StructMapInput{
+ _, err := ViaStructMap(testContext(t), StructMapInput{
Payload: map[string]any{"region": "wrong-region"},
})
assert.ErrorContains(t, err, "map field bound incorrectly")
}
func TestViaPlainMap(t *testing.T) {
- ctx := sdk.NewTIRunContext(context.Background(), sdk.TaskInstance{},
sdk.DagRun{})
- got, err := ViaPlainMap(ctx, slog.Default(), map[string]string{
+ got, err := ViaPlainMap(testContext(t), map[string]string{
"team": "data", "tier": "gold",
})
require.NoError(t, err)
@@ -174,7 +173,6 @@ func TestViaPlainMap(t *testing.T) {
}
func TestViaPlainMapRejectsWrongBinding(t *testing.T) {
- ctx := sdk.NewTIRunContext(context.Background(), sdk.TaskInstance{},
sdk.DagRun{})
- _, err := ViaPlainMap(ctx, slog.Default(), map[string]string{"team":
"wrong"})
+ _, err := ViaPlainMap(testContext(t), map[string]string{"team":
"wrong"})
assert.ErrorContains(t, err, "plain map bound incorrectly")
}
diff --git a/go-sdk/example/bundle/variablewrite/variablewrite.go
b/go-sdk/example/bundle/variablewrite/variablewrite.go
index 6094605a45c..44d1647e236 100644
--- a/go-sdk/example/bundle/variablewrite/variablewrite.go
+++ b/go-sdk/example/bundle/variablewrite/variablewrite.go
@@ -22,7 +22,7 @@ package variablewrite
import (
"fmt"
- "github.com/apache/airflow/go-sdk/sdk"
+ "github.com/apache/airflow/go-sdk/airflow"
)
const (
@@ -37,14 +37,16 @@ const (
// WriteAndDeleteVariable stores the current run id under WrittenKey, then
// writes and deletes ScratchKey to exercise the delete path.
-func WriteAndDeleteVariable(ctx sdk.TIRunContext, client sdk.VariableClient)
error {
- if err := client.SetVariable(ctx, WrittenKey, ctx.DagRun().RunID,
WrittenDescription); err != nil {
+func WriteAndDeleteVariable(actx airflow.Context) error {
+ client := actx.Client()
+ runID := actx.DagRun().RunID
+ if err := client.SetVariable(actx, WrittenKey, runID,
WrittenDescription); err != nil {
return fmt.Errorf("setting %s: %w", WrittenKey, err)
}
- if err := client.SetVariable(ctx, ScratchKey, "scratch", ""); err !=
nil {
+ if err := client.SetVariable(actx, ScratchKey, "scratch", ""); err !=
nil {
return fmt.Errorf("setting %s: %w", ScratchKey, err)
}
- if err := client.DeleteVariable(ctx, ScratchKey); err != nil {
+ if err := client.DeleteVariable(actx, ScratchKey); err != nil {
return fmt.Errorf("deleting %s: %w", ScratchKey, err)
}
return nil
diff --git a/go-sdk/pkg/binding/binding.go b/go-sdk/pkg/binding/binding.go
index af18ec554d3..651050e4d6c 100644
--- a/go-sdk/pkg/binding/binding.go
+++ b/go-sdk/pkg/binding/binding.go
@@ -17,8 +17,9 @@
// Package binding resolves TaskFlow arguments for Go task functions.
//
-// Runtime values are injected by type. Other parameters bind positionally,
-// except a sole struct whose fields bind by `arg:` tag or folded Go name.
+// A task function takes an airflow.Context first, and every parameter after
it is data.
+// Data parameters bind positionally, except a sole struct whose fields bind
by `arg:` tag
+// or folded Go name.
// Captured defaults may go unclaimed, and a sole untagged struct can decode
one
// unclaimed argument as a whole value.
//
@@ -71,10 +72,6 @@ type paramKind int
const (
paramAirflowContext paramKind = iota
- paramTIRunContext
- paramContext
- paramLogger
- paramClient
paramData
paramLoneStruct
)
@@ -105,6 +102,13 @@ type Plan struct {
// Analyze validates a task function and builds its binding plan.
func Analyze(fnType reflect.Type, fnName string) (*Plan, error) {
+ if fnType.NumIn() == 0 {
+ return nil, fmt.Errorf(
+ "task function %s: takes no parameters, but the first
parameter must be "+
+ "airflow.Context",
+ fnName,
+ )
+ }
p := &Plan{fnName: fnName, params: make([]paramPlan, fnType.NumIn())}
var dataIdxs []int
for i := range fnType.NumIn() {
@@ -156,41 +160,16 @@ func (p *Plan) Resolve(
client sdk.Client,
args []Arg,
) ([]reflect.Value, error) {
- out := p.resolveInjectables(ctx, logger, client)
+ out := make([]reflect.Value, len(p.params))
+ // Bound to the live task context, so actx.Done() fires on supervisor
shutdown.
+ ti, dagRun := storedRunMetadata(ctx)
+ out[0] = reflect.ValueOf(airflow.NewContext(ctx, logger, client, ti,
dagRun))
if p.loneStruct {
return p.resolveLoneStructParam(ctx, client, args, out)
}
return p.resolveFlatParams(ctx, client, args, out)
}
-func (p *Plan) resolveInjectables(
- ctx context.Context,
- logger *slog.Logger,
- client sdk.Client,
-) []reflect.Value {
- out := make([]reflect.Value, len(p.params))
- for i, plan := range p.params {
- switch plan.kind {
- case paramAirflowContext:
- // Bound to the live task context, so actx.Done() fires
on supervisor shutdown.
- ti, dagRun := storedRunMetadata(ctx)
- out[i] = reflect.ValueOf(airflow.NewContext(ctx,
logger, client, ti, dagRun))
- case paramTIRunContext:
- // Rebuild the stored metadata around the live task
context.
- ti, dagRun := storedRunMetadata(ctx)
- out[i] = reflect.ValueOf(sdk.NewTIRunContext(ctx, ti,
dagRun))
- case paramContext:
- out[i] = reflect.ValueOf(ctx)
- case paramLogger:
- out[i] = reflect.ValueOf(logger)
- case paramClient:
- out[i] = reflect.ValueOf(client)
- case paramData, paramLoneStruct:
- }
- }
- return out
-}
-
// storedRunMetadata reads the task instance and Dag run recorded on the task
context.
func storedRunMetadata(ctx context.Context) (sdk.TaskInstance, sdk.DagRun) {
stored, ok := ctx.Value(sdkcontext.RuntimeContextKey).(sdk.TIRunContext)
@@ -504,32 +483,36 @@ func (p *Plan) decodeArg(
}
func classifyParam(fnName string, in reflect.Type, index int) (paramPlan,
error) {
- switch {
- case isAirflowContext(in):
- // airflow.Context satisfies isContext too, so match it on
identity first.
- return paramPlan{kind: paramAirflowContext, index: index}, nil
- case isTIRunContext(in):
- // TIRunContext also satisfies context.Context, so check it
first.
- return paramPlan{kind: paramTIRunContext, index: index}, nil
- case isContext(in):
- if !contextType.Implements(in) {
+ if index == 0 {
+ if in == airflowContextType {
+ return paramPlan{kind: paramAirflowContext, index:
index}, nil
+ }
+ if in == reflect.PointerTo(airflowContextType) {
return paramPlan{}, fmt.Errorf(
- "task function %s: parameter %d: interface %s
adds methods on top of "+
- "context.Context; declare
airflow.Context or a separate parameter instead",
- fnName, index, in,
+ "task function %s: parameter 0 is %s, but
airflow.Context is taken by value",
+ fnName, in,
)
}
- return paramPlan{kind: paramContext, index: index}, nil
- case isLogger(in):
- return paramPlan{kind: paramLogger, index: index}, nil
- case isClient(in):
- return paramPlan{kind: paramClient, index: index}, nil
+ return paramPlan{}, fmt.Errorf(
+ "task function %s: parameter 0 is %s, but the first
parameter must be airflow.Context",
+ fnName, in,
+ )
+ }
+ // A logger is a pointer to a struct, so it passes isDecodableType.
Without this check
+ // the handler would get a pointer to a zero slog.Logger, and the first
log call would panic.
+ if in.AssignableTo(slogLoggerType) {
+ return paramPlan{}, fmt.Errorf(
+ "task function %s: parameter %d: %s cannot receive a
task argument (the task's "+
+ "logger comes from the Logger method of the
leading airflow.Context)",
+ fnName, index, in,
+ )
}
if in.Kind() == reflect.Interface && in.NumMethod() > 0 {
return paramPlan{}, fmt.Errorf(
- "task function %s: parameter %d: interface %s is not
injectable "+
- "(want airflow.Context, context.Context, or a
subset of sdk.Client): %s",
- fnName, index, in, explainClientMismatch(in),
+ "task function %s: parameter %d: %s cannot receive a
task argument (an interface "+
+ "with methods cannot be decoded, and the task's
context and client come from "+
+ "the leading airflow.Context)",
+ fnName, index, in,
)
}
if !isDecodableType(in) {
@@ -854,51 +837,9 @@ func implementsUnmarshaler(t reflect.Type) bool {
}
var (
- contextType = reflect.TypeFor[context.Context]()
airflowContextType = reflect.TypeFor[airflow.Context]()
- tiRunContextType = reflect.TypeFor[sdk.TIRunContext]()
slogLoggerType = reflect.TypeFor[*slog.Logger]()
- clientType = reflect.TypeFor[sdk.Client]()
jsonUnmarshalerType = reflect.TypeFor[json.Unmarshaler]()
textUnmarshalerType = reflect.TypeFor[encoding.TextUnmarshaler]()
)
-
-// isContext reports whether inType is an interface a plain context.Context
can fill.
-// A struct can implement context.Context too, so it is matched earlier or
bound as data.
-func isContext(inType reflect.Type) bool {
- return inType != nil && inType.Kind() == reflect.Interface &&
- inType.Implements(contextType)
-}
-
-func isAirflowContext(inType reflect.Type) bool {
- return inType == airflowContextType
-}
-
-func isTIRunContext(inType reflect.Type) bool {
- return inType == tiRunContextType
-}
-
-func isLogger(inType reflect.Type) bool {
- return inType != nil && inType.AssignableTo(slogLoggerType)
-}
-
-// isClient reports whether inType is a non-empty subset of sdk.Client.
-func isClient(inType reflect.Type) bool {
- return inType != nil && inType.Kind() == reflect.Interface &&
- inType.NumMethod() > 0 && clientType.Implements(inType)
-}
-
-func explainClientMismatch(in reflect.Type) string {
- for i := range in.NumMethod() {
- m := in.Method(i)
- cm, ok := clientType.MethodByName(m.Name)
- if !ok {
- return fmt.Sprintf("sdk.Client has no method %s",
m.Name)
- }
- if cm.Type != m.Type {
- return fmt.Sprintf("method %s is %s on sdk.Client, not
%s", m.Name, cm.Type, m.Type)
- }
- }
- return "its method set is not a subset of sdk.Client"
-}
diff --git a/go-sdk/pkg/binding/binding_test.go
b/go-sdk/pkg/binding/binding_test.go
index cfbb04858c4..1143c6dab0a 100644
--- a/go-sdk/pkg/binding/binding_test.go
+++ b/go-sdk/pkg/binding/binding_test.go
@@ -109,33 +109,81 @@ func analyze(s *BindingSuite, fn any) *Plan {
func (s *BindingSuite) resolve(fn any, args []Arg, client sdk.Client)
([]reflect.Value, error) {
plan := analyze(s, fn)
- return plan.Resolve(runtimeCtx(), slog.Default(), client, args)
+ values, err := plan.Resolve(runtimeCtx(), slog.Default(), client, args)
+ if err != nil {
+ return nil, err
+ }
+ s.Require().True(values[0].IsValid(), "parameter 0 must be bound")
+ s.Require().IsType(airflow.Context{}, values[0].Interface())
+ return values[1:], nil
}
func (s *BindingSuite) TestAnalyzeClassification() {
plan := analyze(
s,
- func(ctx sdk.TIRunContext, log *slog.Logger, c
sdk.VariableClient, country string, extracted map[string]any) error {
- return nil
- },
+ func(actx airflow.Context, country string, extracted
map[string]any) error { return nil },
)
- s.Equal(2, plan.numData)
+ s.Equal(2, plan.numData, "every parameter after the leading
airflow.Context is data")
- s.Zero(analyze(s, func() error { return nil }).numData)
+ s.Zero(analyze(s, func(actx airflow.Context) error { return nil
}).numData)
s.Equal(
1,
- analyze(s, func(x any) error { return nil }).numData,
+ analyze(s, func(actx airflow.Context, x any) error { return nil
}).numData,
"an `any` parameter is a data parameter",
)
- s.Equal(
- 1,
- analyze(s, func(actx airflow.Context, country string) error {
return nil }).numData,
- "airflow.Context is injected, not a data parameter",
- )
}
-// airflow.Context satisfies context.Context, so classification has to match it
-// on identity ahead of the plain-context case.
+func (s *BindingSuite) TestAnalyzeRequiresLeadingAirflowContext() {
+ type namedContext airflow.Context
+ type embedsContext struct{ airflow.Context }
+
+ cases := map[string]struct {
+ fn any
+ errContains string
+ }{
+ "no-params": {
+ func() error { return nil },
+ "takes no parameters, but the first parameter must be
airflow.Context",
+ },
+ "data-first": {
+ func(country string, actx airflow.Context) error {
return nil },
+ "parameter 0 is string, but the first parameter must be
airflow.Context",
+ },
+ "plain-context-first": {
+ func(ctx context.Context, country string) error {
return nil },
+ "parameter 0 is context.Context, but the first
parameter must be airflow.Context",
+ },
+ "ti-run-context-first": {
+ func(ctx sdk.TIRunContext) error { return nil },
+ "parameter 0 is sdk.TIRunContext, but the first
parameter must be airflow.Context",
+ },
+ "logger-first": {
+ func(log *slog.Logger) error { return nil },
+ "parameter 0 is *slog.Logger, but the first parameter
must be airflow.Context",
+ },
+ "context-by-pointer": {
+ func(actx *airflow.Context) error { return nil },
+ "parameter 0 is *airflow.Context, but airflow.Context
is taken by value",
+ },
+ "named-context-type": {
+ func(actx namedContext) error { return nil },
+ "parameter 0 is binding.namedContext, but the first
parameter must be airflow.Context",
+ },
+ "struct-embedding-context": {
+ func(actx embedsContext) error { return nil },
+ "parameter 0 is binding.embedsContext, but the first
parameter must be airflow.Context",
+ },
+ }
+ for name, tt := range cases {
+ s.Run(name, func() {
+ _, err := Analyze(reflect.TypeOf(tt.fn), "testFn")
+ if s.Assert().Error(err) {
+ s.Assert().Contains(err.Error(), "task function
testFn: "+tt.errContains)
+ }
+ })
+ }
+}
+
func (s *BindingSuite) TestAirflowContextInjection() {
plan := analyze(s, func(actx airflow.Context) error { return nil })
s.Zero(plan.numData)
@@ -175,60 +223,111 @@ func (s *BindingSuite)
TestAirflowContextTracksTaskCancellation() {
s.Require().ErrorIs(actx.Err(), context.Canceled)
}
-// A user struct can implement context.Context too, and must not reach
-// the interface-only check behind the plain-context case.
-func (s *BindingSuite) TestStructImplementingContextIsNotInjectable() {
- type wrappedContext struct{ context.Context }
-
- _, err := Analyze(reflect.TypeOf(func(w wrappedContext) error { return
nil }), "testFn")
- s.Require().NoError(err, "a struct implementing context.Context must
not break analysis")
-}
-
func (s *BindingSuite) TestAnalyzeRejections() {
cases := map[string]struct {
- fn any
- errContains string
+ fn any
+ rejected string
}{
"func-param": {
- func(cb func()) error { return nil },
- "cannot receive a task argument",
+ func(actx airflow.Context, cb func()) error { return
nil },
+ "parameter 1: type func()",
},
"chan-param": {
- func(ch chan int) error { return nil },
- "cannot receive a task argument",
+ func(actx airflow.Context, ch chan int) error { return
nil },
+ "parameter 1: type chan int",
},
"pointer-to-func-param": {
- func(cb *func()) error { return nil },
- "cannot receive a task argument",
+ func(actx airflow.Context, cb *func()) error { return
nil },
+ "parameter 1: type *func()",
},
"slice-of-func-param": {
- func(cbs []func()) error { return nil },
- "cannot receive a task argument",
+ func(actx airflow.Context, cbs []func()) error { return
nil },
+ "parameter 1: type []func()",
},
"map-with-chan-value-param": {
- func(m map[string]chan int) error { return nil },
- "cannot receive a task argument",
+ func(actx airflow.Context, m map[string]chan int) error
{ return nil },
+ "parameter 1: type map[string]chan int",
+ },
+ "chan-param-after-data": {
+ func(actx airflow.Context, country string, ch chan int)
error { return nil },
+ "parameter 2: type chan int",
+ },
+ }
+ for name, tt := range cases {
+ s.Run(name, func() {
+ _, err := Analyze(reflect.TypeOf(tt.fn), "testFn")
+ if s.Assert().Error(err) {
+ s.Assert().Contains(err.Error(), tt.rejected+"
cannot receive a task argument")
+ s.Assert().Contains(err.Error(), "values cannot
be decoded")
+ }
+ })
+ }
+}
+
+func (s *BindingSuite) TestAnalyzeRejectsAirflowSuppliedParams() {
+ type definedLogger *slog.Logger
+
+ const (
+ loggerReason = "logger comes from the Logger method of the
leading airflow.Context"
+ interfaceReason = "an interface with methods cannot be decoded"
+ )
+ cases := map[string]struct {
+ fn any
+ rejected string
+ reason string
+ }{
+ "logger": {
+ func(actx airflow.Context, log *slog.Logger) error {
return nil },
+ "parameter 1: *slog.Logger",
+ loggerReason,
+ },
+ "defined-logger-type": {
+ func(actx airflow.Context, log definedLogger) error {
return nil },
+ "parameter 1: binding.definedLogger",
+ loggerReason,
},
- "non-client-interface": {
- func(x interface{ NotAClientMethod() }) error { return
nil },
- "sdk.Client has no method NotAClientMethod",
+ "logger-after-data": {
+ func(actx airflow.Context, country string, log
*slog.Logger) error { return nil },
+ "parameter 2: *slog.Logger",
+ loggerReason,
},
- "context-with-extra-methods": {
- func(x interface {
- context.Context
- TaskInstance() sdk.TaskInstance
- },
- ) error {
- return nil
- },
- "adds methods on top of context.Context",
+ "client": {
+ func(actx airflow.Context, client sdk.Client) error {
return nil },
+ "parameter 1: sdk.Client",
+ interfaceReason,
+ },
+ "narrow-client": {
+ func(actx airflow.Context, client sdk.VariableClient)
error { return nil },
+ "parameter 1: sdk.VariableClient",
+ interfaceReason,
+ },
+ "client-after-data": {
+ func(actx airflow.Context, country string, client
sdk.Client) error { return nil },
+ "parameter 2: sdk.Client",
+ interfaceReason,
+ },
+ "plain-context": {
+ func(actx airflow.Context, ctx context.Context) error {
return nil },
+ "parameter 1: context.Context",
+ interfaceReason,
+ },
+ "ti-run-context": {
+ func(actx airflow.Context, ctx sdk.TIRunContext) error
{ return nil },
+ "parameter 1: sdk.TIRunContext",
+ interfaceReason,
+ },
+ "other-interface": {
+ func(actx airflow.Context, x interface{
NotAClientMethod() }) error { return nil },
+ "parameter 1: interface { NotAClientMethod() }",
+ interfaceReason,
},
}
for name, tt := range cases {
s.Run(name, func() {
_, err := Analyze(reflect.TypeOf(tt.fn), "testFn")
if s.Assert().Error(err) {
- s.Assert().Contains(err.Error(), tt.errContains)
+ s.Assert().Contains(err.Error(), tt.rejected+"
cannot receive a task argument")
+ s.Assert().Contains(err.Error(), tt.reason)
}
})
}
@@ -249,11 +348,15 @@ type recursiveNode struct {
func (s *BindingSuite) TestAnalyzeAcceptsSelfDecodingAndRecursiveTypes() {
for name, fn := range map[string]any{
- "time.Time": func(when time.Time) error { return nil },
- "slice-of-time": func(when []time.Time) error { return nil
},
- "self-decoding": func(name string, n selfDecodingNode)
error { return nil },
- "self-decoding-sole": func(n selfDecodingNode) error { return
nil },
- "recursive-struct": func(n recursiveNode) error { return nil
},
+ "time.Time": func(actx airflow.Context, when time.Time)
error { return nil },
+ "slice-of-time": func(actx airflow.Context, when []time.Time)
error { return nil },
+ "self-decoding": func(
+ actx airflow.Context, name string, n selfDecodingNode,
+ ) error {
+ return nil
+ },
+ "self-decoding-sole": func(actx airflow.Context, n
selfDecodingNode) error { return nil },
+ "recursive-struct": func(actx airflow.Context, n
recursiveNode) error { return nil },
} {
s.Run(name, func() {
_, err := Analyze(reflect.TypeOf(fn), "testFn")
@@ -262,19 +365,8 @@ func (s *BindingSuite)
TestAnalyzeAcceptsSelfDecodingAndRecursiveTypes() {
}
}
-func (s *BindingSuite) TestNamedClientInterfacesAreInjectable() {
- for name, typ := range map[string]reflect.Type{
- "Client": reflect.TypeFor[sdk.Client](),
- "VariableClient": reflect.TypeFor[sdk.VariableClient](),
- "ConnectionClient": reflect.TypeFor[sdk.ConnectionClient](),
- "XComClient": reflect.TypeFor[sdk.XComClient](),
- } {
- s.True(isClient(typ), "sdk.%s must stay injectable", name)
- }
-}
-
func (s *BindingSuite) TestResolveArityMismatch() {
- fn := func(country string) error { return nil }
+ fn := func(actx airflow.Context, country string) error { return nil }
_, err := s.resolve(fn, nil, &fakeXComClient{})
if s.Assert().Error(err) {
s.Contains(err.Error(), "argument count mismatch")
@@ -283,7 +375,7 @@ func (s *BindingSuite) TestResolveArityMismatch() {
}
_, err = s.resolve(
- func() error { return nil },
+ func(actx airflow.Context) error { return nil },
[]Arg{LiteralArg{Value: "uk"}},
&fakeXComClient{},
)
@@ -293,7 +385,10 @@ func (s *BindingSuite) TestResolveArityMismatch() {
}
func (s *BindingSuite) TestResolveLiterals() {
- fn := func(country string, count int, ratio float64, on bool, tags
[]string, meta map[string]any) error {
+ fn := func(
+ actx airflow.Context,
+ country string, count int, ratio float64, on bool, tags
[]string, meta map[string]any,
+ ) error {
return nil
}
got, err := s.resolve(fn, []Arg{
@@ -314,7 +409,7 @@ func (s *BindingSuite) TestResolveLiterals() {
}
func (s *BindingSuite) TestResolveSelfDecodingLiterals() {
- fn := func(when time.Time, id uuid.UUID, ratio float64) error { return
nil }
+ fn := func(actx airflow.Context, when time.Time, id uuid.UUID, ratio
float64) error { return nil }
got, err := s.resolve(fn, []Arg{
LiteralArg{
Value: "2024-01-02T03:04:05Z",
@@ -333,7 +428,7 @@ func (s *BindingSuite) TestResolveSelfDecodingLiterals() {
}
func (s *BindingSuite) TestResolveTypedMapParam() {
- fn := func(labels map[string]string) error { return nil }
+ fn := func(actx airflow.Context, labels map[string]string) error {
return nil }
got, err := s.resolve(fn, []Arg{
LiteralArg{
Name: "labels",
@@ -345,21 +440,6 @@ func (s *BindingSuite) TestResolveTypedMapParam() {
s.Equal(map[string]string{"team": "data", "tier": "gold"},
got[0].Interface())
}
-func (s *BindingSuite) TestResolveInterleavedInjectables() {
- fn := func(log *slog.Logger, country string, ctx context.Context, meta
map[string]any) error {
- return nil
- }
- got, err := s.resolve(fn, []Arg{
- LiteralArg{Value: "uk", ValueSchema: argSchema("string")},
- LiteralArg{Value: map[string]any{"k": "v"}, ValueSchema:
argSchema("object")},
- }, &fakeXComClient{})
- s.Require().NoError(err)
- s.NotNil(got[0].Interface().(*slog.Logger))
- s.Equal("uk", got[1].Interface())
- s.NotNil(got[2].Interface().(context.Context))
- s.Equal(map[string]any{"k": "v"}, got[3].Interface())
-}
-
func (s *BindingSuite) TestCheckValueTypeMatrix() {
unionType := &genmodels.ArgValueSchema{"type": []any{"string", "null"}}
cases := map[string]struct {
@@ -448,13 +528,14 @@ func (s *BindingSuite) TestCheckValueTypeMatrix() {
}
func (s *BindingSuite) TestResolveTypeMismatchFailsLoudly() {
- fn := func(count int) error { return nil }
+ fn := func(actx airflow.Context, count int) error { return nil }
_, err := s.resolve(
fn,
[]Arg{LiteralArg{Value: "uk", ValueSchema:
argSchema("string")}},
&fakeXComClient{},
)
if s.Assert().Error(err) {
+ s.Contains(err.Error(), "argument 0 (parameter 1)")
s.Contains(
err.Error(),
`the Dag declares JSON-schema type "string" which
cannot bind to Go parameter type int`,
@@ -463,7 +544,7 @@ func (s *BindingSuite) TestResolveTypeMismatchFailsLoudly()
{
}
func (s *BindingSuite) TestResolveLiteralDecodeFailure() {
- fn := func(count int) error { return nil }
+ fn := func(actx airflow.Context, count int) error { return nil }
_, err := s.resolve(fn, []Arg{LiteralArg{Value: "uk"}},
&fakeXComClient{})
if s.Assert().Error(err) {
s.Contains(err.Error(), "decoding literal value into int")
@@ -505,7 +586,7 @@ func (s *BindingSuite) TestResolveXComArgs() {
"probe/return_value": "probe-value",
}}
- fn := func(res extractResult, probe string) error { return nil }
+ fn := func(actx airflow.Context, res extractResult, probe string) error
{ return nil }
got, err := s.resolve(fn, []Arg{
XComArg{TaskID: "extract", ValueSchema: argSchema("object")},
XComArg{TaskID: "probe", ValueSchema: argSchema("string")},
@@ -535,7 +616,7 @@ func (s *BindingSuite) TestResolveXComStrictStructDecode() {
client := &fakeXComClient{values: map[string]any{
"extract/return_value": map[string]any{"go_version": "go1.24",
"renamed_field": 1},
}}
- fn := func(res extractResult) error { return nil }
+ fn := func(actx airflow.Context, res extractResult) error { return nil }
_, err := s.resolve(fn, []Arg{XComArg{TaskID: "extract"}}, client)
if s.Assert().Error(err) {
s.Contains(err.Error(), `decoding xcom from task "extract"`)
@@ -545,7 +626,7 @@ func (s *BindingSuite) TestResolveXComStrictStructDecode() {
func (s *BindingSuite) TestResolveXComPullFailure() {
client := &fakeXComClient{err: sdk.XComNotFound}
- fn := func(res map[string]any) error { return nil }
+ fn := func(actx airflow.Context, res map[string]any) error { return nil
}
_, err := s.resolve(fn, []Arg{XComArg{TaskID: "extract"}}, client)
if s.Assert().Error(err) {
s.Contains(err.Error(), `pulling xcom from task "extract"`)
@@ -554,7 +635,7 @@ func (s *BindingSuite) TestResolveXComPullFailure() {
func (s *BindingSuite) TestResolveMultipleXComPullFailures() {
client := &fakeXComClient{err: sdk.XComNotFound}
- fn := func(a map[string]any, b map[string]any, c map[string]any) error
{ return nil }
+ fn := func(actx airflow.Context, a, b, c map[string]any) error { return
nil }
_, err := s.resolve(fn, []Arg{
XComArg{TaskID: "extract_a"},
XComArg{TaskID: "extract_b"},
@@ -574,7 +655,7 @@ func (s *BindingSuite) TestResolveWholeStructFromXCom() {
"region": "eu-west-1",
},
}}
- fn := func(cfg wholeConfig) error { return nil }
+ fn := func(actx airflow.Context, cfg wholeConfig) error { return nil }
got, err := s.resolve(fn, []Arg{
XComArg{Name: "cfg", TaskID: "make_config", ValueSchema:
argSchema("object")},
}, client)
@@ -587,7 +668,7 @@ func (s *BindingSuite) TestResolveWholeStructFromXCom() {
}
func (s *BindingSuite) TestResolveXComWithoutRuntimeContext() {
- plan := analyze(s, func(res map[string]any) error { return nil })
+ plan := analyze(s, func(actx airflow.Context, res map[string]any) error
{ return nil })
_, err := plan.Resolve(
context.Background(), slog.Default(), &fakeXComClient{},
[]Arg{XComArg{TaskID: "extract"}},
@@ -598,7 +679,7 @@ func (s *BindingSuite)
TestResolveXComWithoutRuntimeContext() {
}
func (s *BindingSuite) TestResolveNullHandling() {
- fn := func(meta map[string]any) error { return nil }
+ fn := func(actx airflow.Context, meta map[string]any) error { return
nil }
got, err := s.resolve(
fn,
[]Arg{LiteralArg{Value: nil, ValueSchema: argSchema("object")}},
@@ -607,7 +688,7 @@ func (s *BindingSuite) TestResolveNullHandling() {
s.Require().NoError(err)
s.Nil(got[0].Interface())
- fnStr := func(country string) error { return nil }
+ fnStr := func(actx airflow.Context, country string) error { return nil }
_, err = s.resolve(fnStr, []Arg{LiteralArg{Value: nil}},
&fakeXComClient{})
if s.Assert().Error(err) {
s.Contains(err.Error(), "not nilable")
@@ -622,7 +703,7 @@ func (fakeArg) Schema() *genmodels.ArgValueSchema { return
nil }
func (fakeArg) sealedArg() {}
func (s *BindingSuite) TestResolveUnsupportedVariant() {
- fn := func(country string) error { return nil }
+ fn := func(actx airflow.Context, country string) error { return nil }
_, err := s.resolve(fn, []Arg{fakeArg{}}, &fakeXComClient{})
if s.Assert().Error(err) {
s.Contains(err.Error(), "unsupported argument binding
binding.fakeArg")
@@ -630,53 +711,39 @@ func (s *BindingSuite) TestResolveUnsupportedVariant() {
}
func (s *BindingSuite) TestResolveNilArg() {
- fn := func(country string) error { return nil }
+ fn := func(actx airflow.Context, country string) error { return nil }
_, err := s.resolve(fn, []Arg{nil}, &fakeXComClient{})
if s.Assert().Error(err) {
s.Contains(err.Error(), "nil argument binding")
}
}
-func (s *BindingSuite) TestResolveTIRunContextRebuild() {
- ti := sdk.TaskInstance{DagID: "dag1", RunID: "run1", TaskID:
"transform"}
- dagRun := sdk.DagRun{DagID: "dag1", RunID: "run1"}
- ctx := context.WithValue(
- runtimeCtx(),
- sdkcontext.RuntimeContextKey,
- sdk.NewTIRunContext(context.Background(), ti, dagRun),
- )
-
- plan := analyze(s, func(rc sdk.TIRunContext, country string) error {
return nil })
- got, err := plan.Resolve(ctx, slog.Default(), &fakeXComClient{}, []Arg{
- LiteralArg{Value: "uk", ValueSchema: argSchema("string")},
- })
- s.Require().NoError(err)
- rc := got[0].Interface().(sdk.TIRunContext)
- s.Equal(ti, rc.TaskInstance())
- s.Equal(dagRun, rc.DagRun())
- s.Equal("uk", got[1].Interface())
-}
-
func (s *BindingSuite) TestAnalyzeLoneStructClassification() {
- plan := analyze(s, func(input simpleInput) error { return nil })
+ plan := analyze(s, func(actx airflow.Context, input simpleInput) error
{ return nil })
s.True(plan.loneStruct, "a sole struct data parameter is resolved by
name at execution")
s.Zero(plan.numData)
- ptrPlan := analyze(s, func(input *simpleInput) error { return nil })
+ ptrPlan := analyze(s, func(actx airflow.Context, input *simpleInput)
error { return nil })
s.True(ptrPlan.loneStruct, "a pointer to a sole struct is detected the
same way")
s.Zero(ptrPlan.numData)
- flatPlan := analyze(s, func(prefix string, cfg wholeConfig) error {
return nil })
+ flatPlan := analyze(
+ s,
+ func(actx airflow.Context, prefix string, cfg wholeConfig)
error { return nil },
+ )
s.False(flatPlan.loneStruct, "a struct alongside another data parameter
is a flat slot")
s.Equal(2, flatPlan.numData)
- scalarPlan := analyze(s, func(name string) error { return nil })
+ scalarPlan := analyze(s, func(actx airflow.Context, name string) error
{ return nil })
s.False(scalarPlan.loneStruct, "a sole non-struct data parameter is
plain positional")
s.Equal(1, scalarPlan.numData)
}
func (s *BindingSuite) TestAnalyzeMultipleStructsAreFlat() {
- plan := analyze(s, func(a wholeConfig, b wholeConfig) error { return
nil })
+ plan := analyze(
+ s,
+ func(actx airflow.Context, a wholeConfig, b wholeConfig) error
{ return nil },
+ )
s.False(plan.loneStruct)
s.Equal(2, plan.numData)
}
@@ -699,23 +766,23 @@ func (s *BindingSuite) TestAnalyzeStructValidation() {
errContains string
}{
"duplicate-arg-names": {
- func(input duplicateArgNames) error { return nil },
+ func(actx airflow.Context, input duplicateArgNames)
error { return nil },
`fields A and B both bind arg name "A"`,
},
"folded-duplicate-arg-names": {
- func(input foldedDuplicateArgNames) error { return nil
},
+ func(actx airflow.Context, input
foldedDuplicateArgNames) error { return nil },
"differ only in case or underscores",
},
"tagged-non-decodable-field": {
- func(input taggedNonDecodableField) error { return nil
},
+ func(actx airflow.Context, input
taggedNonDecodableField) error { return nil },
"cannot receive a task argument",
},
"tagged-struct-not-sole": {
- func(prefix string, input combineInput) error { return
nil },
+ func(actx airflow.Context, prefix string, input
combineInput) error { return nil },
"must be the function's only data parameter",
},
"tagged-struct-trailing": {
- func(input combineInput, suffix string) error { return
nil },
+ func(actx airflow.Context, input combineInput, suffix
string) error { return nil },
"must be the function's only data parameter",
},
}
@@ -730,7 +797,7 @@ func (s *BindingSuite) TestAnalyzeStructValidation() {
}
func (s *BindingSuite) TestResolveStructAllFields() {
- fn := func(input combineInput) error { return nil }
+ fn := func(actx airflow.Context, input combineInput) error { return nil
}
got, err := s.resolve(fn, []Arg{
LiteralArg{Name: "Name", Value: "widget", ValueSchema:
argSchema("string")},
LiteralArg{Name: "count", Value: 7, ValueSchema:
argSchema("integer")},
@@ -743,20 +810,20 @@ func (s *BindingSuite) TestResolveStructAllFields() {
}
func (s *BindingSuite) TestResolveStructXComArg() {
- fn := func(log *slog.Logger, input reportInput) error { return nil }
+ fn := func(actx airflow.Context, input reportInput) error { return nil }
got, err := s.resolve(fn, []Arg{
XComArg{Name: "region", TaskID: "make_region", ValueSchema:
argSchema("string")},
LiteralArg{Name: "Ratio", Value: 0.5, ValueSchema:
argSchema("number")},
}, &fakeXComClient{values: map[string]any{"make_region/return_value":
"east"}})
s.Require().NoError(err)
- input := got[1].Interface().(reportInput)
+ input := got[0].Interface().(reportInput)
s.Equal("east", input.Region, "Region resolves by name despite being
declared after Ratio")
s.Equal(0.5, input.Ratio)
}
func (s *BindingSuite) TestResolveStructSingleClaimedArgBindsByName() {
- fn := func(input simpleInput) error { return nil }
+ fn := func(actx airflow.Context, input simpleInput) error { return nil }
got, err := s.resolve(fn, []Arg{
LiteralArg{Name: "Name", Value: "widget", ValueSchema:
argSchema("string")},
}, &fakeXComClient{})
@@ -765,7 +832,7 @@ func (s *BindingSuite)
TestResolveStructSingleClaimedArgBindsByName() {
}
func (s *BindingSuite) TestResolveStructPointer() {
- fn := func(input *simpleInput) error { return nil }
+ fn := func(actx airflow.Context, input *simpleInput) error { return nil
}
got, err := s.resolve(fn, []Arg{
LiteralArg{Name: "Name", Value: "widget", ValueSchema:
argSchema("string")},
}, &fakeXComClient{})
@@ -776,7 +843,7 @@ func (s *BindingSuite) TestResolveStructPointer() {
}
func (s *BindingSuite) TestResolveLoneStructBothModes() {
- fn := func(cfg wholeConfig) error { return nil }
+ fn := func(actx airflow.Context, cfg wholeConfig) error { return nil }
want := wholeConfig{Environment: "production", Region: "eu-west-1"}
named, err := s.resolve(fn, []Arg{
@@ -802,7 +869,7 @@ type taggedRegionInput struct {
}
func (s *BindingSuite) TestResolveTaggedStructNeverFallsBackToWholeValue() {
- fn := func(input taggedRegionInput) error { return nil }
+ fn := func(actx airflow.Context, input taggedRegionInput) error {
return nil }
_, err := s.resolve(fn, []Arg{
LiteralArg{Name: "region_code", Value: "eu-west-1",
ValueSchema: argSchema("string")},
}, &fakeXComClient{})
@@ -812,7 +879,7 @@ func (s *BindingSuite)
TestResolveTaggedStructNeverFallsBackToWholeValue() {
}
func (s *BindingSuite) TestResolveStructUnclaimedArgFailsLoudly() {
- fn := func(input combineInput) error { return nil }
+ fn := func(actx airflow.Context, input combineInput) error { return nil
}
_, err := s.resolve(fn, []Arg{
LiteralArg{Name: "Name", Value: "widget", ValueSchema:
argSchema("string")},
LiteralArg{Name: "typo", Value: "x", ValueSchema:
argSchema("string")},
@@ -823,7 +890,7 @@ func (s *BindingSuite)
TestResolveStructUnclaimedArgFailsLoudly() {
}
func (s *BindingSuite) TestResolveStructUnclaimedFromDefaultAllowed() {
- fn := func(input combineInput) error { return nil }
+ fn := func(actx airflow.Context, input combineInput) error { return nil
}
got, err := s.resolve(fn, []Arg{
LiteralArg{Name: "Name", Value: "widget", ValueSchema:
argSchema("string")},
LiteralArg{
@@ -838,7 +905,7 @@ func (s *BindingSuite)
TestResolveStructUnclaimedFromDefaultAllowed() {
}
func (s *BindingSuite) TestResolveStructEmptySpecFailsLoudly() {
- fn := func(input simpleInput) error { return nil }
+ fn := func(actx airflow.Context, input simpleInput) error { return nil }
for name, args := range map[string][]Arg{"nil-spec": nil, "empty-spec":
{}} {
s.Run(name, func() {
_, err := s.resolve(fn, args, &fakeXComClient{})
@@ -850,7 +917,7 @@ func (s *BindingSuite)
TestResolveStructEmptySpecFailsLoudly() {
}
func (s *BindingSuite) TestResolveStructOnlyDefaultsZeroValues() {
- fn := func(input twoFieldInput) error { return nil }
+ fn := func(actx airflow.Context, input twoFieldInput) error { return
nil }
got, err := s.resolve(fn, []Arg{
LiteralArg{
Name: "threshold",
@@ -866,7 +933,7 @@ func (s *BindingSuite)
TestResolveStructOnlyDefaultsZeroValues() {
}
func (s *BindingSuite) TestResolveStructUnmatchedFieldZeroValued() {
- fn := func(input twoFieldInput) error { return nil }
+ fn := func(actx airflow.Context, input twoFieldInput) error { return
nil }
got, err := s.resolve(fn, []Arg{
LiteralArg{Name: "Name", Value: "widget", ValueSchema:
argSchema("string")},
}, &fakeXComClient{})
@@ -877,7 +944,7 @@ func (s *BindingSuite)
TestResolveStructUnmatchedFieldZeroValued() {
}
func (s *BindingSuite) TestResolveFlatParamsToleratesCapturedDefaults() {
- fn := func(country string) error { return nil }
+ fn := func(actx airflow.Context, country string) error { return nil }
got, err := s.resolve(fn, []Arg{
LiteralArg{Name: "country", Value: "uk", ValueSchema:
argSchema("string")},
LiteralArg{
@@ -892,7 +959,7 @@ func (s *BindingSuite)
TestResolveFlatParamsToleratesCapturedDefaults() {
}
func (s *BindingSuite) TestResolveFlatParamsBindsDefaultsWhenDeclared() {
- fn := func(country string, verbose bool) error { return nil }
+ fn := func(actx airflow.Context, country string, verbose bool) error {
return nil }
got, err := s.resolve(fn, []Arg{
LiteralArg{Name: "country", Value: "uk", ValueSchema:
argSchema("string")},
LiteralArg{
@@ -908,7 +975,7 @@ func (s *BindingSuite)
TestResolveFlatParamsBindsDefaultsWhenDeclared() {
}
func (s *BindingSuite) TestResolveWholeStructIgnoresCapturedDefaults() {
- fn := func(config wholeConfig) error { return nil }
+ fn := func(actx airflow.Context, config wholeConfig) error { return nil
}
got, err := s.resolve(fn, []Arg{
LiteralArg{
Name: "config",
@@ -932,7 +999,7 @@ type snakeCaseInput struct {
}
func (s *BindingSuite) TestResolveUntaggedFieldsBindSnakeCaseArguments() {
- fn := func(input snakeCaseInput) error { return nil }
+ fn := func(actx airflow.Context, input snakeCaseInput) error { return
nil }
got, err := s.resolve(fn, []Arg{
LiteralArg{Name: "region_code", Value: "eu-west-1",
ValueSchema: argSchema("string")},
LiteralArg{Name: "threshold", Value: 0.75, ValueSchema:
argSchema("number")},
@@ -951,7 +1018,7 @@ type embeddedInput struct {
}
func (s *BindingSuite) TestResolveEmbeddedStructFields() {
- fn := func(input embeddedInput) error { return nil }
+ fn := func(actx airflow.Context, input embeddedInput) error { return
nil }
got, err := s.resolve(fn, []Arg{
LiteralArg{Name: "region", Value: "eu-west-1", ValueSchema:
argSchema("string")},
LiteralArg{Name: "threshold", Value: 0.75, ValueSchema:
argSchema("number")},
@@ -969,7 +1036,7 @@ type money struct {
func (m *money) UnmarshalJSON([]byte) error { m.Amount = "decoded"; return nil
}
func (s *BindingSuite) TestResolveSelfDecodingTypeAgainstStringSchema() {
- fn := func(price money) error { return nil }
+ fn := func(actx airflow.Context, price money) error { return nil }
got, err := s.resolve(fn, []Arg{
LiteralArg{Name: "price", Value: "12.34", ValueSchema:
argSchema("string")},
}, &fakeXComClient{})
@@ -983,7 +1050,7 @@ type callbackConfig struct {
}
func (s *BindingSuite) TestResolveStructWithNonBindableField() {
- fn := func(config callbackConfig) error { return nil }
+ fn := func(actx airflow.Context, config callbackConfig) error { return
nil }
got, err := s.resolve(fn, []Arg{
LiteralArg{
Name: "config",
@@ -998,7 +1065,7 @@ func (s *BindingSuite)
TestResolveStructWithNonBindableField() {
}
func (s *BindingSuite) TestResolveEmptyInterfaceDataParam() {
- fn := func(payload any) error { return nil }
+ fn := func(actx airflow.Context, payload any) error { return nil }
got, err := s.resolve(fn, []Arg{
LiteralArg{
Name: "payload",
diff --git a/go-sdk/pkg/execution/integration_test.go
b/go-sdk/pkg/execution/integration_test.go
index f802c0533f1..abb33a899fd 100644
--- a/go-sdk/pkg/execution/integration_test.go
+++ b/go-sdk/pkg/execution/integration_test.go
@@ -35,7 +35,6 @@ import (
"github.com/apache/airflow/go-sdk/airflow"
"github.com/apache/airflow/go-sdk/bundle/bundlev1"
"github.com/apache/airflow/go-sdk/pkg/execution/genmodels"
- "github.com/apache/airflow/go-sdk/sdk"
)
// assertSucceedTask asserts RunTask produced a terminal SucceedTask body.
@@ -65,15 +64,15 @@ func assertRetryTask(t *testing.T, result any, reasonSubstr
string) {
// --- Test task functions ---
-func failingTask() error {
+func failingTask(airflow.Context) error {
return errors.New("task failed intentionally")
}
-func panicTask() error {
+func panicTask(airflow.Context) error {
panic("something went wrong")
}
-func simpleTask() error {
+func simpleTask(airflow.Context) error {
return nil
}
@@ -200,7 +199,7 @@ func TestTaskRunnerBindsArgs(t *testing.T) {
var gotMeta map[string]any
bundle := buildBundle(t, func(r bundlev1.Registry) {
r.AddDag("test_dag").AddTaskWithName("transform",
- func(log *slog.Logger, country string, meta
map[string]any) error {
+ func(actx airflow.Context, country string, meta
map[string]any) error {
gotCountry = country
gotMeta = meta
return nil
@@ -236,7 +235,7 @@ func TestTaskRunnerArgBindingsArityMismatch(t *testing.T) {
ran := false
bundle := buildBundle(t, func(r bundlev1.Registry) {
r.AddDag("test_dag").AddTaskWithName("transform",
- func(country string, meta map[string]any) error {
+ func(actx airflow.Context, country string, meta
map[string]any) error {
ran = true
return nil
})
@@ -268,7 +267,7 @@ func TestTaskRunnerBindsStructArgs(t *testing.T) {
var got regionInput
bundle := buildBundle(t, func(r bundlev1.Registry) {
r.AddDag("test_dag").AddTaskWithName("transform",
- func(input regionInput) error {
+ func(actx airflow.Context, input regionInput) error {
got = input
return nil
})
@@ -296,7 +295,7 @@ func TestTaskRunnerStructIgnoresUnclaimedDefault(t
*testing.T) {
var got regionInput
bundle := buildBundle(t, func(r bundlev1.Registry) {
r.AddDag("test_dag").AddTaskWithName("transform",
- func(input regionInput) error {
+ func(actx airflow.Context, input regionInput) error {
got = input
return nil
})
@@ -330,7 +329,7 @@ func TestTaskRunnerStructIgnoresUnclaimedDefault(t
*testing.T) {
func TestTaskRunnerArgBindingsTypeMismatch(t *testing.T) {
bundle := buildBundle(t, func(r bundlev1.Registry) {
r.AddDag("test_dag").AddTaskWithName("transform",
- func(count int) error { return nil })
+ func(actx airflow.Context, count int) error { return
nil })
})
details := newStartupDetails(
@@ -354,7 +353,7 @@ func TestTaskRunnerArgBindingsUnknownKind(t *testing.T) {
ran := false
bundle := buildBundle(t, func(r bundlev1.Registry) {
r.AddDag("test_dag").AddTaskWithName("transform",
- func(country string) error {
+ func(actx airflow.Context, country string) error {
ran = true
return nil
})
@@ -377,7 +376,7 @@ func TestTaskRunnerArgBindingsMalformedElement(t
*testing.T) {
ran := false
bundle := buildBundle(t, func(r bundlev1.Registry) {
r.AddDag("test_dag").AddTaskWithName("transform",
- func(country string) error {
+ func(actx airflow.Context, country string) error {
ran = true
return nil
})
@@ -426,7 +425,7 @@ func TestTaskRunnerArgBindingsMissingRequiredFields(t
*testing.T) {
ran := false
bundle := buildBundle(t, func(r bundlev1.Registry) {
r.AddDag("test_dag").AddTaskWithName("transform",
- func(country string) error {
+ func(actx airflow.Context, country
string) error {
ran = true
return nil
})
@@ -447,7 +446,7 @@ func TestTaskRunnerArgBindingsMissingRequiredFields(t
*testing.T) {
func TestTaskRunnerMalformedSpecHonorsShouldRetry(t *testing.T) {
bundle := buildBundle(t, func(r bundlev1.Registry) {
r.AddDag("test_dag").AddTaskWithName("transform",
- func(country string) error { return nil })
+ func(actx airflow.Context, country string) error {
return nil })
})
details := newStartupDetails(
@@ -463,36 +462,18 @@ func TestTaskRunnerMalformedSpecHonorsShouldRetry(t
*testing.T) {
assertRetryTask(t, result, `unknown kind "template"`)
}
-func TestRunTaskHonorsContextCancellation(t *testing.T) {
- bundle := buildBundle(t, func(r bundlev1.Registry) {
- r.AddDag("test_dag").AddTaskWithName("ctxcheck",
- func(ctx context.Context) error { return ctx.Err() })
- })
-
- details := newStartupDetails("ctxcheck")
-
- // A cancelled root context must reach the user task through RunTask's
- // threading; the task surfaces ctx.Err(), which RunTask maps to failed.
- ctx, cancel := context.WithCancel(context.Background())
- cancel()
-
- logger := slog.New(slog.NewTextHandler(io.Discard, nil))
- comm := NewCoordinatorComm(bytes.NewReader(nil), io.Discard, logger)
-
- result := RunTask(ctx, bundle, details, comm, logger)
- assertTaskState(t, result, genmodels.TaskStateStateFailed)
-}
-
-func TestRunTaskInjectsRuntimeContext(t *testing.T) {
+// A handler taking an airflow.Context gets on that one value everything
+// the runtime used to hand over as separate parameters.
+func TestRunTaskInjectsAirflowContext(t *testing.T) {
logical := time.Date(2026, 6, 9, 12, 0, 0, 0, time.UTC)
- start := logical
+ start := logical.Add(-time.Hour)
end := logical.Add(time.Hour)
- var got sdk.TIRunContext
+ var got airflow.Context
bundle := buildBundle(t, func(r bundlev1.Registry) {
r.AddDag("test_dag").AddTaskWithName("ctxgrab",
- func(ctx sdk.TIRunContext) error {
- got = ctx
+ func(actx airflow.Context) error {
+ got = actx
return nil
})
})
@@ -524,11 +505,9 @@ func TestRunTaskInjectsRuntimeContext(t *testing.T) {
result := RunTask(context.Background(), bundle, details, comm, logger)
assertSucceedTask(t, result)
- require.NotNil(
- t,
- got,
- "the task must receive a TIRunContext backed by the live task
context",
- )
+ assert.Same(t, logger, got.Logger(), "the task's logger must arrive on
the Context")
+ assert.NotNil(t, got.Client(), "the coordinator-backed client must
arrive on the Context")
+
ti := got.TaskInstance()
assert.Equal(t, "test_dag", ti.DagID)
assert.Equal(t, "run1", ti.RunID)
@@ -545,58 +524,6 @@ func TestRunTaskInjectsRuntimeContext(t *testing.T) {
assert.Equal(t, start, *dagRun.DataIntervalStart)
require.NotNil(t, dagRun.DataIntervalEnd)
assert.Equal(t, end, *dagRun.DataIntervalEnd)
-}
-
-// A handler taking an airflow.Context gets on that one value everything
-// the runtime used to hand over as separate parameters.
-func TestRunTaskInjectsAirflowContext(t *testing.T) {
- logical := time.Date(2026, 6, 9, 12, 0, 0, 0, time.UTC)
-
- var got airflow.Context
- bundle := buildBundle(t, func(r bundlev1.Registry) {
- r.AddDag("test_dag").AddTaskWithName("ctxgrab",
- func(actx airflow.Context) error {
- got = actx
- return nil
- })
- })
-
- details := &genmodels.StartupDetails{
- TI: genmodels.TaskInstance{
- ID: "550e8400-e29b-41d4-a716-446655440000",
- DagID: "test_dag",
- TaskID: "ctxgrab",
- RunID: "run1",
- TryNumber: 2,
- MapIndex: ptr(-1),
- },
- BundleInfo: genmodels.BundleInfo{Name: "test", Version: "1.0"},
- TIContext: genmodels.TIRunContext{
- DagRun: genmodels.DagRun{LogicalDate: logical},
- },
- }
-
- logger := slog.New(slog.NewTextHandler(io.Discard, nil))
- comm := NewCoordinatorComm(bytes.NewReader(nil), io.Discard, logger)
-
- result := RunTask(context.Background(), bundle, details, comm, logger)
- assertSucceedTask(t, result)
-
- assert.Same(t, logger, got.Logger(), "the task's logger must arrive on
the Context")
- assert.NotNil(t, got.Client(), "the coordinator-backed client must
arrive on the Context")
-
- ti := got.TaskInstance()
- assert.Equal(t, "test_dag", ti.DagID)
- assert.Equal(t, "run1", ti.RunID)
- assert.Equal(t, "ctxgrab", ti.TaskID)
- assert.Equal(t, 2, ti.TryNumber)
- assert.Nil(t, ti.MapIndex, "an unmapped task (map_index -1) must
surface as nil")
-
- dagRun := got.DagRun()
- assert.Equal(t, "test_dag", dagRun.DagID)
- assert.Equal(t, "run1", dagRun.RunID)
- require.NotNil(t, dagRun.LogicalDate)
- assert.Equal(t, logical, *dagRun.LogicalDate)
// A helper taking a plain context.Context recovers the same surface.
recovered, ok := airflow.FromContext(context.Context(got))
@@ -638,11 +565,11 @@ func TestRunTaskAirflowContextHonorsShutdown(t
*testing.T) {
}
func TestRunTaskRuntimeContextMappedIndex(t *testing.T) {
- var got sdk.TIRunContext
+ var got airflow.Context
bundle := buildBundle(t, func(r bundlev1.Registry) {
r.AddDag("test_dag").AddTaskWithName("ctxgrab",
- func(ctx sdk.TIRunContext) error {
- got = ctx
+ func(actx airflow.Context) error {
+ got = actx
return nil
})
})
@@ -762,7 +689,8 @@ func TestServeUsesSupervisorLogLevelEnvironment(t
*testing.T) {
provider := &fakeProvider{
register: func(r bundlev1.Registry) error {
- r.AddDag("dag1").AddTaskWithName("logging", func(logger
*slog.Logger) error {
+ r.AddDag("dag1").AddTaskWithName("logging", func(actx
airflow.Context) error {
+ logger := actx.Logger()
logger.Info("global filtered")
logger.WithGroup("example.child").Debug("namespace debug")
logger.WithGroup("unrelated").Warn("unrelated
filtered")
@@ -839,8 +767,8 @@ func TestServeClientRoundTripEndToEnd(t *testing.T) {
provider := &fakeProvider{
register: func(r bundlev1.Registry) error {
r.AddDag("dag1").AddTaskWithName("getvar",
- func(ctx context.Context, c sdk.Client)
(string, error) {
- v, err := c.GetVariable(ctx, varKey)
+ func(actx airflow.Context) (string, error) {
+ v, err :=
actx.Client().GetVariable(actx, varKey)
if err != nil {
return "", err
}
diff --git a/go-sdk/pkg/execution/task_runner.go
b/go-sdk/pkg/execution/task_runner.go
index e89dcdd526a..ef8bec04b79 100644
--- a/go-sdk/pkg/execution/task_runner.go
+++ b/go-sdk/pkg/execution/task_runner.go
@@ -63,11 +63,12 @@ func RunTask(
client := NewCoordinatorClient(comm)
- // Carries the task runtime context for sdk.TIRunContext injection. The
- // scheduling timestamps live on the nested dag_run object in the
- // supervisor's TIRunContext schema. The base context is a placeholder;
- // bundlev1.Execute rebuilds the value around the live task context when
- // binding the parameter.
+ // runtimeContext carries the task instance and Dag run that binding
puts
+ // on the task's airflow.Context. The scheduling timestamps live on the
+ // nested dag_run object in the supervisor's TIRunContext schema. The
base
+ // context is a placeholder, because binding reads only the task
instance
+ // and Dag run from this value and builds the airflow.Context around the
+ // live task context.
dagRun := details.TIContext.DagRun
runtimeContext := sdk.NewTIRunContext(
context.Background(),
diff --git a/go-sdk/pkg/sdkcontext/keys.go b/go-sdk/pkg/sdkcontext/keys.go
index ed5ebdec637..f907e4facb3 100644
--- a/go-sdk/pkg/sdkcontext/keys.go
+++ b/go-sdk/pkg/sdkcontext/keys.go
@@ -23,17 +23,17 @@ type (
)
var (
- // RuntimeContextKey stores the public, task-facing runtime context
- // (task instance identifiers and the Dag run's scheduling timestamps).
- // The coordinator-mode runtime populates it from StartupDetails; the
- // bundle runtime reads it to inject an sdk.TIRunContext parameter into
- // task functions rather than exposing this key directly. Its value type
- // is sdk.TIRunContext (built over a placeholder base context; the
bundle
- // runtime rebuilds it around the live task context at injection time),
- // but this package does not import sdk to avoid an import cycle.
+ // RuntimeContextKey stores the task instance identifiers and the Dag
run's
+ // scheduling timestamps. The coordinator-mode runtime populates it from
+ // StartupDetails. The bundle runtime reads it to fill the
airflow.Context
+ // a task function takes first, rather than exposing this key directly.
+ // Its value type is sdk.TIRunContext. The base context under it is a
+ // placeholder, because the bundle runtime reads back only the task
+ // instance and Dag run.
RuntimeContextKey = runtimeContextKey{}
- // SdkClientContextKey holds the coordinator-backed sdk.Client injected
into
- // task functions. SDK calls travel over the supervisor comm socket.
+ // SdkClientContextKey holds the coordinator-backed sdk.Client that task
+ // functions reach through airflow.Context. SDK calls travel over the
+ // supervisor comm socket.
SdkClientContextKey = sdkClientContextKey{}
)
diff --git a/go-sdk/sdk/context.go b/go-sdk/sdk/context.go
index c6fb906deca..c6ca1dfe8c3 100644
--- a/go-sdk/sdk/context.go
+++ b/go-sdk/sdk/context.go
@@ -22,30 +22,16 @@ import (
"time"
)
-// TIRunContext is the execution context handed to a task. It behaves as the
-// standard context.Context (cancellation, deadline, request-scoped values) and
-// additionally exposes the identifiers and scheduling timestamps of the task
-// instance that is executing, along with the Dag run it belongs to. It is the
-// Go equivalent of the execution context the Python and Java SDKs expose to
-// task authors.
+// TIRunContext is a context.Context that also exposes the identifiers and
+// scheduling timestamps of the task instance that is executing, along with the
+// Dag run it belongs to.
//
-// The runtime injects it into a task function by parameter type, so declare it
-// as the task's context argument:
-//
-// func myTask(ctx sdk.TIRunContext, log *slog.Logger) error {
-// log.Info("running",
-// "task_id", ctx.TaskInstance().TaskID,
-// "run_id", ctx.DagRun().RunID,
-// )
-// return nil
-// }
-//
-// Because it embeds context.Context it is usable wherever one is expected:
-// pass it straight to client calls, select on ctx.Done(), or hand it to
-// downstream helpers that take a context.Context.
+// The runtime uses it to carry those values on the task's context. Task
+// functions read them from the
[github.com/apache/airflow/go-sdk/airflow.Context]
+// they take first, which exposes the same TaskInstance and DagRun next to the
+// logger and the client.
//
// It is an interface, and only this package implements it.
-// Build one in tests with NewTIRunContext.
//
// The context package's advice against storing a Context in a struct
// (https://pkg.go.dev/context#hdr-Contexts_and_structs) is about domain types
that would
@@ -63,15 +49,13 @@ type TIRunContext interface {
// NewTIRunContext returns a TIRunContext that delegates context behaviour to
// ctx and exposes ti and dagRun. It panics on a nil ctx, mirroring the context
-// package's own constructors. The runtime calls it when binding a task's
-// TIRunContext parameter; in unit tests, use it to hand-build the argument:
-//
-// ctx := sdk.NewTIRunContext(context.Background(),
sdk.TaskInstance{TaskID: "t1"}, sdk.DagRun{})
+// package's own constructors. The runtime calls it to record the task instance
+// and Dag run on the task's context.
func NewTIRunContext(ctx context.Context, ti TaskInstance, dagRun DagRun)
TIRunContext {
if ctx == nil {
- // This cannot happen from the runtime: taskFunction.Execute
always
- // binds the live task context. A nil ctx is a programming
error in
- // the caller, so fail loudly instead of masking it.
+ // This cannot happen from the runtime, which always passes a
non-nil
+ // base context. A nil ctx is a programming error in the
caller, so
+ // fail loudly instead of masking it.
panic("sdk.NewTIRunContext: cannot create TIRunContext from nil
context.Context")
}
return tiRunContext{Context: ctx, ti: ti, dagRun: dagRun}
diff --git a/go-sdk/sdk/doc.go b/go-sdk/sdk/doc.go
index edb9f45b211..332c33cbef5 100644
--- a/go-sdk/sdk/doc.go
+++ b/go-sdk/sdk/doc.go
@@ -19,23 +19,23 @@
Package sdk gives task functions access to the Airflow "model" (Variables,
Connections, and XCom) at run time.
-A task function does not construct a client itself. The runtime inspects the
-function's parameters and injects one by type, so you declare the narrowest
-interface you need and use it:
+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
+gets the client from it:
- func mytask(ctx context.Context, client sdk.Client, log *slog.Logger)
error {
- val, err := client.GetVariable(ctx, "my_variable")
+ func mytask(actx airflow.Context) error {
+ val, err := actx.Client().GetVariable(actx, "my_variable")
if err != nil {
return err
}
- log.Info("got variable", "value", val)
+ actx.Logger().InfoContext(actx, "got variable", "value", val)
return nil
}
-Ask for [Client] for full access, or a narrower interface such as
-[VariableClient] or [ConnectionClient] when the task only reads one kind of
-object. The narrower type documents what the task touches and makes it easy to
-pass a fake in unit tests.
+[Client] combines the narrower [VariableClient], [ConnectionClient] and
+[XComClient]. A helper that only reads one kind of object can take the narrower
+interface. That documents what the helper touches and makes it easy to pass a
+fake in unit tests.
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]
diff --git a/go-sdk/sdk/sdk.go b/go-sdk/sdk/sdk.go
index e378527b01d..f047216108e 100644
--- a/go-sdk/sdk/sdk.go
+++ b/go-sdk/sdk/sdk.go
@@ -121,8 +121,8 @@ type XComClient interface {
}
// Client is the full task-facing API: read/write Variables, read Connections,
-// and read/write XCom. A task that declares an sdk.Client parameter is handed
one
-// by the runtime. If a task needs only one capability, ask for the narrower
+// 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.
type Client interface {
VariableClient
diff --git a/kubernetes-tests/lang_sdk/go_example/main.go
b/kubernetes-tests/lang_sdk/go_example/main.go
index d1de46bf50f..fb00cff815e 100644
--- a/kubernetes-tests/lang_sdk/go_example/main.go
+++ b/kubernetes-tests/lang_sdk/go_example/main.go
@@ -24,13 +24,12 @@ package main
import (
"log"
- "log/slog"
"runtime"
"time"
+ "github.com/apache/airflow/go-sdk/airflow"
v1 "github.com/apache/airflow/go-sdk/bundle/bundlev1"
"github.com/apache/airflow/go-sdk/bundle/bundlev1/bundlev1server"
- "github.com/apache/airflow/go-sdk/sdk"
)
// Must match the dag_id of the Python stub Dag and the Java bundle.
@@ -57,8 +56,8 @@ func main() {
// goExtract returns a map pushed as the task's XCom, mirroring the reference
// example's extract task so the Python downstream can read it.
-func goExtract(ctx sdk.TIRunContext, log *slog.Logger) (any, error) {
- log.InfoContext(ctx, "go_extract running")
+func goExtract(actx airflow.Context) (any, error) {
+ actx.Logger().InfoContext(actx, "go_extract running")
return map[string]any{
"go_version": runtime.Version(),
"timestamp": time.Now().UnixNano(),
@@ -67,11 +66,11 @@ func goExtract(ctx sdk.TIRunContext, log *slog.Logger)
(any, error) {
// goTransform reads the my_variable Airflow variable through the coordinator,
// exercising a GetVariable round-trip over the Execution API.
-func goTransform(ctx sdk.TIRunContext, client sdk.VariableClient, log
*slog.Logger) error {
- val, err := client.GetVariable(ctx, "my_variable")
+func goTransform(actx airflow.Context) error {
+ val, err := actx.Client().GetVariable(actx, "my_variable")
if err != nil {
return err
}
- log.InfoContext(ctx, "go_transform obtained variable", "my_variable",
val)
+ actx.Logger().InfoContext(actx, "go_transform obtained variable",
"my_variable", val)
return nil
}