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