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
 }

Reply via email to