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 99c9012db77 Go SDK: introduce the `airflow` package with 
`airflow.Context` (#73277)
99c9012db77 is described below

commit 99c9012db775fc08813504292143fd5bcac1c864
Author: Jason(Zhe-You) Liu <[email protected]>
AuthorDate: Sat Sep 19 11:27:01 2026 +0800

    Go SDK: introduce the `airflow` package with `airflow.Context` (#73277)
    
    * Introduce the airflow package with airflow.Context
    
    Every task handler takes an airflow.Context first: a struct embedding
    context.Context that also exposes Logger(), Client(), TaskInstance() and
    DagRun(). What Airflow supplies a task arrives as a method on that one
    value instead of as a parameter of its own, so a handler can no longer
    declare no context at all and get its logger separately from the context
    it logs against.
    
    NewContext is the only way to build one, and FromContext recovers the same
    surface inside a helper typed as a plain context.Context.
    
    pkg/binding matches the type on identity ahead of the TIRunContext and
    plain-context cases, and binds it to the live task context so actx.Done()
    fires on supervisor shutdown. isContext is narrowed to interfaces along
    the way, because a struct implementing context.Context used to reach
    contextType.Implements(in), which panics when in is not an interface.
    
    Correct the sdk.TIRunContext doc comment. The context package's advice
    against storing a Context in a struct is about domain types carrying a
    request-scoped context in a field, not about a purpose-built context type.
    
    * Assert which select branch fired in the shutdown test
    
    Both branches returned a non-nil error, so RunTask reported FAILED either
    way and the test passed whether or not actx.Done() fired. Record the
    branch and the context error, and assert on both.
    
    * Reject a nil logger or client in airflow.NewContext
    
    Logger() fell back to slog.Default() while Client() returned nil, so a
    zero Context logged fine but panicked on a client call. Validate both in
    the constructor instead and drop the fallback, so the two accessors
    behave the same.
---
 go-sdk/airflow/context.go                | 117 +++++++++++++++++++++
 go-sdk/airflow/context_test.go           | 171 +++++++++++++++++++++++++++++++
 go-sdk/airflow/doc.go                    |  41 ++++++++
 go-sdk/bundle/bundlev1/registry.go       |   2 +
 go-sdk/pkg/binding/binding.go            |  48 ++++++---
 go-sdk/pkg/binding/binding_test.go       |  57 +++++++++++
 go-sdk/pkg/execution/integration_test.go |  91 ++++++++++++++++
 go-sdk/sdk/context.go                    |  13 ++-
 8 files changed, 522 insertions(+), 18 deletions(-)

diff --git a/go-sdk/airflow/context.go b/go-sdk/airflow/context.go
new file mode 100644
index 00000000000..193e5aa14c2
--- /dev/null
+++ b/go-sdk/airflow/context.go
@@ -0,0 +1,117 @@
+// Licensed to the Apache Software Foundation (ASF) under one
+// or more contributor license agreements.  See the NOTICE file
+// distributed with this work for additional information
+// regarding copyright ownership.  The ASF licenses this file
+// to you under the Apache License, Version 2.0 (the
+// "License"); you may not use this file except in compliance
+// with the License.  You may obtain a copy of the License at
+//
+//   http://www.apache.org/licenses/LICENSE-2.0
+//
+// Unless required by applicable law or agreed to in writing,
+// software distributed under the License is distributed on an
+// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
+// KIND, either express or implied.  See the License for the
+// specific language governing permissions and limitations
+// under the License.
+
+package airflow
+
+import (
+       "context"
+       "log/slog"
+
+       "github.com/apache/airflow/go-sdk/sdk"
+)
+
+type (
+       // TaskInstance identifies the task instance a handler is running for.
+       TaskInstance = sdk.TaskInstance
+
+       // DagRun identifies the Dag run the task instance belongs to,
+       // and carries its scheduling timestamps.
+       DagRun = sdk.DagRun
+)
+
+// Context is the first parameter of every task handler.
+// It is a context.Context bound to the running task, so actx.Done() fires 
when the supervisor
+// asks the task to stop, and it exposes what Airflow gives the task: 
[Context.Logger],
+// [Context.Client], [Context.TaskInstance] and [Context.DagRun].
+//
+// [NewContext] builds one. The zero Context is not usable.
+type Context struct {
+       context.Context
+
+       values taskValues
+}
+
+// taskValues is what NewContext stores on the context chain so FromContext 
can recover it.
+type taskValues struct {
+       logger *slog.Logger
+       client sdk.Client
+       ti     TaskInstance
+       dagRun DagRun
+}
+
+type contextKey struct{}
+
+// NewContext returns a [Context] backed by ctx.
+// It panics if ctx, logger or client is nil, so every accessor on the 
returned Context is safe.
+//
+// The runtime calls it when it binds a handler's first parameter.
+// Call it directly to unit-test a handler:
+//
+//     actx := airflow.NewContext(
+//             t.Context(), slog.Default(), fakeClient,
+//             airflow.TaskInstance{DagID: "py_etl", TaskID: "transform", 
TryNumber: 1},
+//             airflow.DagRun{DagID: "py_etl", RunID: "run1"},
+//     )
+//     require.NoError(t, transform(actx, "US"))
+func NewContext(
+       ctx context.Context,
+       logger *slog.Logger,
+       client sdk.Client,
+       ti TaskInstance,
+       dagRun DagRun,
+) Context {
+       switch {
+       case ctx == nil:
+               panic("airflow.NewContext: nil context.Context")
+       case logger == nil:
+               panic("airflow.NewContext: nil logger")
+       case client == nil:
+               panic("airflow.NewContext: nil client")
+       }
+       values := taskValues{logger: logger, client: client, ti: ti, dagRun: 
dagRun}
+       return Context{Context: context.WithValue(ctx, contextKey{}, values), 
values: values}
+}
+
+// FromContext recovers the Airflow surface inside a helper typed as a plain 
context.Context,
+// reporting whether ctx carries one.
+//
+// It succeeds for the [Context] a handler was given and for any context 
derived from it.
+// The returned Context keeps ctx, so a deadline or cancellation added on the 
way down applies.
+func FromContext(ctx context.Context) (Context, bool) {
+       if ctx == nil {
+               return Context{}, false
+       }
+       values, ok := ctx.Value(contextKey{}).(taskValues)
+       if !ok {
+               return Context{}, false
+       }
+       return Context{Context: ctx, values: values}, true
+}
+
+// Logger writes to the task's Airflow log.
+// The logger takes a context, so pass the same Context: 
actx.Logger().InfoContext(actx, "msg").
+func (c Context) Logger() *slog.Logger { return c.values.logger }
+
+// Client reads Airflow Variables, Connections and XCom.
+// Its calls take a context, so pass the same Context: 
actx.Client().GetVariable(actx, "name").
+func (c Context) Client() sdk.Client { return c.values.client }
+
+// TaskInstance identifies the task instance that is executing.
+func (c Context) TaskInstance() TaskInstance { return c.values.ti }
+
+// DagRun identifies the Dag run the task instance belongs to.
+func (c Context) DagRun() DagRun { return c.values.dagRun }
diff --git a/go-sdk/airflow/context_test.go b/go-sdk/airflow/context_test.go
new file mode 100644
index 00000000000..2e6b4504685
--- /dev/null
+++ b/go-sdk/airflow/context_test.go
@@ -0,0 +1,171 @@
+// Licensed to the Apache Software Foundation (ASF) under one
+// or more contributor license agreements.  See the NOTICE file
+// distributed with this work for additional information
+// regarding copyright ownership.  The ASF licenses this file
+// to you under the Apache License, Version 2.0 (the
+// "License"); you may not use this file except in compliance
+// with the License.  You may obtain a copy of the License at
+//
+//   http://www.apache.org/licenses/LICENSE-2.0
+//
+// Unless required by applicable law or agreed to in writing,
+// software distributed under the License is distributed on an
+// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
+// KIND, either express or implied.  See the License for the
+// specific language governing permissions and limitations
+// under the License.
+
+package airflow
+
+import (
+       "context"
+       "io"
+       "log/slog"
+       "testing"
+       "time"
+
+       "github.com/stretchr/testify/assert"
+       "github.com/stretchr/testify/require"
+
+       "github.com/apache/airflow/go-sdk/sdk"
+)
+
+type probeKey struct{}
+
+// fakeClient satisfies sdk.Client without implementing any of it.
+// These tests only check which value the accessor hands back.
+type fakeClient struct {
+       sdk.Client
+}
+
+func testValues() (*slog.Logger, *fakeClient, TaskInstance, DagRun) {
+       logger := slog.New(slog.NewTextHandler(io.Discard, nil))
+       client := &fakeClient{}
+       ti := TaskInstance{DagID: "py_etl", RunID: "run1", TaskID: "transform", 
TryNumber: 2}
+       dagRun := DagRun{DagID: "py_etl", RunID: "run1"}
+       return logger, client, ti, dagRun
+}
+
+func TestNewContextAccessors(t *testing.T) {
+       logger, client, ti, dagRun := testValues()
+
+       base := context.WithValue(context.Background(), probeKey{}, 
"probe-value")
+       actx := NewContext(base, logger, client, ti, dagRun)
+
+       assert.Same(t, logger, actx.Logger())
+       assert.Same(t, client, actx.Client())
+       assert.Equal(t, ti, actx.TaskInstance())
+       assert.Equal(t, dagRun, actx.DagRun())
+       assert.Equal(
+               t,
+               "probe-value",
+               actx.Value(probeKey{}),
+               "context behaviour must delegate to the base context",
+       )
+}
+
+// Fail at construction rather than hand out a Context whose accessors return 
nil.
+func TestNewContextRejectsNilArgs(t *testing.T) {
+       logger, client, ti, dagRun := testValues()
+
+       cases := map[string]func(){
+               "nil context": func() { NewContext(nil, logger, client, ti, 
dagRun) },
+               "nil logger":  func() { NewContext(context.Background(), nil, 
client, ti, dagRun) },
+               "nil client":  func() { NewContext(context.Background(), 
logger, nil, ti, dagRun) },
+       }
+       for name, build := range cases {
+               t.Run(name, func(t *testing.T) {
+                       assert.Panics(t, build)
+               })
+       }
+}
+
+// execution.Serve traps the supervisor's SIGINT/SIGTERM into the context
+// the runtime binds, so a cooperative handler sees it on actx.Done().
+func TestContextDoneFollowsBaseCancellation(t *testing.T) {
+       logger, client, ti, dagRun := testValues()
+
+       base, cancel := context.WithCancel(context.Background())
+       actx := NewContext(base, logger, client, ti, dagRun)
+
+       select {
+       case <-actx.Done():
+               t.Fatal("actx must not be done before the base context is 
cancelled")
+       default:
+       }
+
+       cancel()
+
+       select {
+       case <-actx.Done():
+       case <-time.After(time.Second):
+               t.Fatal("actx.Done() must fire when the base context is 
cancelled")
+       }
+       assert.ErrorIs(t, actx.Err(), context.Canceled)
+}
+
+func TestFromContextOnTaskContext(t *testing.T) {
+       logger, client, ti, dagRun := testValues()
+       actx := NewContext(context.Background(), logger, client, ti, dagRun)
+
+       got, ok := FromContext(actx)
+       require.True(t, ok, "the Context a handler is given must carry itself")
+       assert.Same(t, logger, got.Logger())
+       assert.Same(t, client, got.Client())
+       assert.Equal(t, ti, got.TaskInstance())
+       assert.Equal(t, dagRun, got.DagRun())
+}
+
+// A helper recovers the surface whatever the handler wrapped the context in.
+func TestFromContextOnDerivedContext(t *testing.T) {
+       logger, client, ti, dagRun := testValues()
+       actx := NewContext(context.Background(), logger, client, ti, dagRun)
+
+       cases := map[string]context.Context{
+               "plain context.Context": context.Context(actx),
+               "context.WithValue":     context.WithValue(actx, probeKey{}, 
"probe-value"),
+               "context.WithoutCancel": context.WithoutCancel(actx),
+       }
+       for name, ctx := range cases {
+               t.Run(name, func(t *testing.T) {
+                       got, ok := FromContext(ctx)
+                       require.True(t, ok)
+                       assert.Same(t, logger, got.Logger())
+                       assert.Same(t, client, got.Client())
+                       assert.Equal(t, ti, got.TaskInstance())
+                       assert.Equal(t, dagRun, got.DagRun())
+               })
+       }
+}
+
+// The recovered Context keeps the caller's context, so a cancellation added
+// on the way down applies to it.
+func TestFromContextKeepsCallerCancellation(t *testing.T) {
+       logger, client, ti, dagRun := testValues()
+       actx := NewContext(context.Background(), logger, client, ti, dagRun)
+
+       inner, cancel := context.WithCancel(actx)
+       got, ok := FromContext(inner)
+       require.True(t, ok)
+
+       cancel()
+
+       select {
+       case <-got.Done():
+       case <-time.After(time.Second):
+               t.Fatal("the recovered Context must honour the caller's 
cancellation")
+       }
+       assert.NoError(t, actx.Err(), "cancelling a derived context must not 
cancel the task")
+}
+
+// A context that never passed through NewContext must report false
+// rather than hand back a Context that looks usable.
+func TestFromContextOnPlainContext(t *testing.T) {
+       got, ok := FromContext(context.Background())
+       assert.False(t, ok)
+       assert.Equal(t, Context{}, got)
+
+       got, ok = FromContext(nil)
+       assert.False(t, ok)
+       assert.Equal(t, Context{}, got)
+}
diff --git a/go-sdk/airflow/doc.go b/go-sdk/airflow/doc.go
new file mode 100644
index 00000000000..76f13fd31c1
--- /dev/null
+++ b/go-sdk/airflow/doc.go
@@ -0,0 +1,41 @@
+// Licensed to the Apache Software Foundation (ASF) under one
+// or more contributor license agreements.  See the NOTICE file
+// distributed with this work for additional information
+// regarding copyright ownership.  The ASF licenses this file
+// to you under the Apache License, Version 2.0 (the
+// "License"); you may not use this file except in compliance
+// with the License.  You may obtain a copy of the License at
+//
+//   http://www.apache.org/licenses/LICENSE-2.0
+//
+// Unless required by applicable law or agreed to in writing,
+// software distributed under the License is distributed on an
+// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
+// KIND, either express or implied.  See the License for the
+// specific language governing permissions and limitations
+// under the License.
+
+/*
+Package airflow is what a Dag author imports to write the Go body of a task.
+
+Every task handler takes a [Context] first.
+Everything Airflow gives the task arrives as a method on it: the logger, the 
client for
+Variables, Connections and XCom, and the identity of the task instance and its 
Dag run.
+
+       func transform(actx airflow.Context, country string) error {
+               actx.Logger().InfoContext(actx, "transforming", "country", 
country)
+
+               threshold, err := actx.Client().GetVariable(actx, 
"etl_threshold")
+               if err != nil {
+                       return err
+               }
+               return writeRows(actx, country, threshold)
+       }
+
+[Context] is a context.Context, so pass it to a client call or to 
http.NewRequestWithContext,
+and select on actx.Done(), which fires when the supervisor asks the task to 
stop.
+Cleanup that must outlive that cancellation runs under 
context.WithoutCancel(actx).
+
+A helper typed as a plain context.Context recovers the same surface with 
[FromContext].
+*/
+package airflow
diff --git a/go-sdk/bundle/bundlev1/registry.go 
b/go-sdk/bundle/bundlev1/registry.go
index 563cbca3566..b4745249870 100644
--- a/go-sdk/bundle/bundlev1/registry.go
+++ b/go-sdk/bundle/bundlev1/registry.go
@@ -34,6 +34,8 @@ type (
                //
                // 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 /
diff --git a/go-sdk/pkg/binding/binding.go b/go-sdk/pkg/binding/binding.go
index 02b3018dc24..af18ec554d3 100644
--- a/go-sdk/pkg/binding/binding.go
+++ b/go-sdk/pkg/binding/binding.go
@@ -37,6 +37,7 @@ import (
        "strings"
        "sync"
 
+       "github.com/apache/airflow/go-sdk/airflow"
        "github.com/apache/airflow/go-sdk/pkg/execution/genmodels"
        "github.com/apache/airflow/go-sdk/pkg/sdkcontext"
        "github.com/apache/airflow/go-sdk/sdk"
@@ -69,7 +70,8 @@ func (LiteralArg) sealedArg() {}
 type paramKind int
 
 const (
-       paramTIRunContext paramKind = iota
+       paramAirflowContext paramKind = iota
+       paramTIRunContext
        paramContext
        paramLogger
        paramClient
@@ -169,13 +171,13 @@ func (p *Plan) resolveInjectables(
        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.
-                       var ti sdk.TaskInstance
-                       var dagRun sdk.DagRun
-                       if stored, ok := 
ctx.Value(sdkcontext.RuntimeContextKey).(sdk.TIRunContext); ok {
-                               ti, dagRun = stored.TaskInstance(), 
stored.DagRun()
-                       }
+                       ti, dagRun := storedRunMetadata(ctx)
                        out[i] = reflect.ValueOf(sdk.NewTIRunContext(ctx, ti, 
dagRun))
                case paramContext:
                        out[i] = reflect.ValueOf(ctx)
@@ -189,6 +191,15 @@ func (p *Plan) resolveInjectables(
        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)
+       if !ok {
+               return sdk.TaskInstance{}, sdk.DagRun{}
+       }
+       return stored.TaskInstance(), stored.DagRun()
+}
+
 func (p *Plan) resolveFlatParams(
        ctx context.Context,
        c sdk.XComClient,
@@ -494,6 +505,9 @@ 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
@@ -501,7 +515,7 @@ func classifyParam(fnName string, in reflect.Type, index 
int) (paramPlan, error)
                if !contextType.Implements(in) {
                        return paramPlan{}, fmt.Errorf(
                                "task function %s: parameter %d: interface %s 
adds methods on top of "+
-                                       "context.Context; declare 
sdk.TIRunContext or a separate parameter instead",
+                                       "context.Context; declare 
airflow.Context or a separate parameter instead",
                                fnName, index, in,
                        )
                }
@@ -514,7 +528,7 @@ func classifyParam(fnName string, in reflect.Type, index 
int) (paramPlan, error)
        if in.Kind() == reflect.Interface && in.NumMethod() > 0 {
                return paramPlan{}, fmt.Errorf(
                        "task function %s: parameter %d: interface %s is not 
injectable "+
-                               "(want context.Context, sdk.TIRunContext, or a 
subset of sdk.Client): %s",
+                               "(want airflow.Context, context.Context, or a 
subset of sdk.Client): %s",
                        fnName, index, in, explainClientMismatch(in),
                )
        }
@@ -840,17 +854,25 @@ func implementsUnmarshaler(t reflect.Type) bool {
 }
 
 var (
-       contextType      = reflect.TypeFor[context.Context]()
-       tiRunContextType = reflect.TypeFor[sdk.TIRunContext]()
-       slogLoggerType   = reflect.TypeFor[*slog.Logger]()
-       clientType       = reflect.TypeFor[sdk.Client]()
+       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.Implements(contextType)
+       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 {
diff --git a/go-sdk/pkg/binding/binding_test.go 
b/go-sdk/pkg/binding/binding_test.go
index 8f94c9a8c11..cfbb04858c4 100644
--- a/go-sdk/pkg/binding/binding_test.go
+++ b/go-sdk/pkg/binding/binding_test.go
@@ -19,6 +19,7 @@ package binding
 
 import (
        "context"
+       "io"
        "log/slog"
        "reflect"
        "sync"
@@ -28,6 +29,7 @@ import (
        "github.com/google/uuid"
        "github.com/stretchr/testify/suite"
 
+       "github.com/apache/airflow/go-sdk/airflow"
        "github.com/apache/airflow/go-sdk/pkg/execution/genmodels"
        "github.com/apache/airflow/go-sdk/pkg/sdkcontext"
        "github.com/apache/airflow/go-sdk/sdk"
@@ -125,6 +127,61 @@ func (s *BindingSuite) TestAnalyzeClassification() {
                analyze(s, func(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) TestAirflowContextInjection() {
+       plan := analyze(s, func(actx airflow.Context) error { return nil })
+       s.Zero(plan.numData)
+
+       logger := slog.New(slog.NewTextHandler(io.Discard, nil))
+       client := &fakeXComClient{}
+       values, err := plan.Resolve(runtimeCtx(), logger, client, nil)
+       s.Require().NoError(err)
+       s.Require().Len(values, 1)
+
+       actx, ok := values[0].Interface().(airflow.Context)
+       s.Require().True(ok, "parameter 0 must be bound to an airflow.Context")
+       s.Same(logger, actx.Logger())
+       s.Same(client, actx.Client())
+       s.Equal("dag1", actx.TaskInstance().DagID)
+       s.Equal("transform", actx.TaskInstance().TaskID)
+       s.Equal("run1", actx.DagRun().RunID)
+
+       // A helper taking a plain context.Context gets the same surface back.
+       recovered, ok := airflow.FromContext(context.Context(actx))
+       s.Require().True(ok)
+       s.Equal(actx.TaskInstance(), recovered.TaskInstance())
+}
+
+// The Context is bound to the live task context, not a placeholder.
+func (s *BindingSuite) TestAirflowContextTracksTaskCancellation() {
+       plan := analyze(s, func(actx airflow.Context) error { return nil })
+
+       ctx, cancel := context.WithCancel(runtimeCtx())
+       values, err := plan.Resolve(ctx, slog.Default(), &fakeXComClient{}, nil)
+       s.Require().NoError(err)
+       actx, ok := values[0].Interface().(airflow.Context)
+       s.Require().True(ok)
+
+       s.Require().NoError(actx.Err())
+       cancel()
+       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() {
diff --git a/go-sdk/pkg/execution/integration_test.go 
b/go-sdk/pkg/execution/integration_test.go
index 482165d887d..f802c0533f1 100644
--- a/go-sdk/pkg/execution/integration_test.go
+++ b/go-sdk/pkg/execution/integration_test.go
@@ -32,6 +32,7 @@ import (
        "github.com/stretchr/testify/assert"
        "github.com/stretchr/testify/require"
 
+       "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"
@@ -546,6 +547,96 @@ func TestRunTaskInjectsRuntimeContext(t *testing.T) {
        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))
+       require.True(t, ok)
+       assert.Equal(t, ti, recovered.TaskInstance())
+}
+
+// Serve traps SIGINT/SIGTERM into the context it hands RunTask, so a
+// supervisor shutdown reaches the handler on actx.Done().
+func TestRunTaskAirflowContextHonorsShutdown(t *testing.T) {
+       var sawDone bool
+       var sawErr error
+       bundle := buildBundle(t, func(r bundlev1.Registry) {
+               r.AddDag("test_dag").AddTaskWithName("ctxcheck",
+                       func(actx airflow.Context) error {
+                               select {
+                               case <-actx.Done():
+                                       sawDone = true
+                               default:
+                               }
+                               sawErr = actx.Err()
+                               return sawErr
+                       })
+       })
+
+       details := newStartupDetails("ctxcheck")
+
+       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)
+
+       assert.True(t, sawDone, "actx.Done() must fire on a cancelled task 
context")
+       assert.ErrorIs(t, sawErr, context.Canceled)
+       assertTaskState(t, result, genmodels.TaskStateStateFailed)
+}
+
 func TestRunTaskRuntimeContextMappedIndex(t *testing.T) {
        var got sdk.TIRunContext
        bundle := buildBundle(t, func(r bundlev1.Registry) {
diff --git a/go-sdk/sdk/context.go b/go-sdk/sdk/context.go
index afd0dd37050..c6fb906deca 100644
--- a/go-sdk/sdk/context.go
+++ b/go-sdk/sdk/context.go
@@ -44,11 +44,14 @@ import (
 // pass it straight to client calls, select on ctx.Done(), or hand it to
 // downstream helpers that take a context.Context.
 //
-// It is an interface rather than a struct holding a context.Context, which
-// the context package advises against 
(https://pkg.go.dev/context#hdr-Contexts_and_structs):
-// the runtime constructs a fresh value around the live task context for each
-// invocation, and task code cannot end up with a half-initialised value. Build
-// one in tests with NewTIRunContext.
+// 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
+// carry a request-scoped context in a field, not about a purpose-built 
context type.
+// That is why [github.com/apache/airflow/go-sdk/airflow.Context], the value a 
task handler
+// takes first, is a struct embedding context.Context.
 type TIRunContext interface {
        context.Context
 

Reply via email to