diff --git a/lint/drain_run_stream_before_release.go b/lint/drain_run_stream_before_release.go new file mode 100644 index 0000000000..784bb15035 --- /dev/null +++ b/lint/drain_run_stream_before_release.go @@ -0,0 +1,283 @@ +package main + +import ( + "go/ast" + "go/token" + "go/types" + "strings" + + "github.com/dgageot/rubocop-go/cop" + "github.com/dgageot/rubocop-go/prog" + "golang.org/x/tools/go/cfg" +) + +// DrainRunStreamBeforeRelease checks locally consumed RunStream channels. A +// cancellation request is not a join: callers must drain before returning and +// releasing turn ownership. It checks return paths, not cancellation or the +// ordering of Unlock/finish calls. Channels never received locally, ownership +// passed to helpers, and channel variable reuse are outside its analysis. +// The program pass resolves imports; the file runner only supplies partial types. +var DrainRunStreamBeforeRelease = &prog.Func{ + Meta: cop.Meta{ + Name: "Lint/DrainRunStreamBeforeRelease", + Description: "drain RunStream channels before returning, including cancellation and error paths", + Severity: cop.Warning, + }, + Run: func(p *prog.Pass) { + for _, pkg := range p.Program.Packages { + for _, file := range pkg.Syntax { + if strings.HasSuffix(p.Program.Fset.Position(file.Pos()).Filename, "_test.go") { + continue + } + pass := &cop.Pass{FileSet: p.Program.Fset, File: file, Info: pkg.TypesInfo, Package: pkg.Types} + ast.Inspect(file, func(n ast.Node) bool { + var body *ast.BlockStmt + switch fn := n.(type) { + case *ast.FuncDecl: + body = fn.Body + case *ast.FuncLit: + body = fn.Body + default: + return true + } + for _, call := range checkRunStreamDrains(pass, body) { + p.ReportAtf(call.Pos(), call.End(), "RunStream can outlive this consumer; cancel and drain the channel before returning or releasing turn ownership") + } + return true + }) + } + } + }, +} + +type runStreamSource struct { + call *ast.CallExpr + bind ast.Node + object types.Object +} + +func checkRunStreamDrains(p *cop.Pass, body *ast.BlockStmt) []*ast.CallExpr { + if body == nil || p.Info == nil { + return nil + } + var sources []runStreamSource + inspectStreamBody(body, func(n ast.Node) { + switch n := n.(type) { + case *ast.AssignStmt: + if len(n.Lhs) == 1 && len(n.Rhs) == 1 { + if call, ok := n.Rhs[0].(*ast.CallExpr); ok && isRunStreamCall(p, call) { + sources = append(sources, runStreamSource{call, n, streamObject(p, n.Lhs[0])}) + } + } + case *ast.ValueSpec: + if len(n.Names) == 1 && len(n.Values) == 1 { + if call, ok := n.Values[0].(*ast.CallExpr); ok && isRunStreamCall(p, call) { + sources = append(sources, runStreamSource{call, n, p.Info.ObjectOf(n.Names[0])}) + } + } + case *ast.RangeStmt: + if call, ok := n.X.(*ast.CallExpr); ok && isRunStreamCall(p, call) { + sources = append(sources, runStreamSource{call: call, bind: call}) + } + } + }) + if len(sources) == 0 { + return nil + } + graph := cfg.New(body, func(*ast.CallExpr) bool { return true }) + var offenses []*ast.CallExpr + for _, source := range sources { + if streamCanEscape(p, graph, body, source) { + offenses = append(offenses, source.call) + } + } + return offenses +} + +func isRunStreamCall(p *cop.Pass, call *ast.CallExpr) bool { + sel, ok := call.Fun.(*ast.SelectorExpr) + if !ok || (sel.Sel.Name != "RunStream" && sel.Sel.Name != "runStream") { + return false + } + t := p.Info.TypeOf(call) + if t == nil { + return false + } + ch, ok := t.Underlying().(*types.Chan) + if !ok { + return false + } + event, ok := types.Unalias(ch.Elem()).(*types.Named) + return ok && event.Obj().Name() == "Event" && event.Obj().Pkg() != nil && + event.Obj().Pkg().Path() == "github.com/docker/docker-agent/pkg/runtime" +} + +func streamObject(p *cop.Pass, expr ast.Expr) types.Object { + id, ok := ast.Unparen(expr).(*ast.Ident) + if !ok || id.Name == "_" { + return nil + } + return p.Info.ObjectOf(id) +} + +func (s runStreamSource) matches(p *cop.Pass, expr ast.Expr) bool { + return expr == s.call || (s.object != nil && streamObject(p, expr) == s.object) +} + +func inspectStreamBody(body ast.Node, visit func(ast.Node)) { + ast.Inspect(body, func(n ast.Node) bool { + if _, nested := n.(*ast.FuncLit); nested { + return false + } + visit(n) + return true + }) +} + +type streamDrainState struct { + block *cfg.Block + pending bool + deferred streamDrainKind + nilled bool +} + +func streamCanEscape(p *cop.Pass, graph *cfg.CFG, body *ast.BlockStmt, source runStreamSource) bool { + // Comma-ok receives establish closure only on the !ok branch, not merely + // because a select received an event or the context was cancelled. + closedChecks := map[types.Object]bool{} + consumed := false + inspectStreamBody(body, func(n ast.Node) { + switch n := n.(type) { + case *ast.RangeStmt: + consumed = consumed || source.matches(p, n.X) + case *ast.UnaryExpr: + consumed = consumed || (n.Op == token.ARROW && source.matches(p, n.X)) + case *ast.AssignStmt: + if len(n.Lhs) == 2 && len(n.Rhs) == 1 { + if recv, ok := n.Rhs[0].(*ast.UnaryExpr); ok && recv.Op == token.ARROW && source.matches(p, recv.X) { + if object := streamObject(p, n.Lhs[1]); object != nil { + closedChecks[object] = true + } + } + } + } + }) + if !consumed { + return false // Ownership may be passed to another consumer. + } + + queue := []streamDrainState{{block: graph.Blocks[0]}} + seen := map[streamDrainState]bool{} + for len(queue) > 0 { + state := queue[0] + queue = queue[1:] + if seen[state] { + continue + } + seen[state] = true + for _, node := range state.block.Nodes { + if node == source.bind { + state.pending = true + state.nilled = false + if state.deferred != streamDrainCaptured { + state.deferred = streamDrainNone + } + } + if assign, ok := node.(*ast.AssignStmt); ok { + for i, lhs := range assign.Lhs { + if source.matches(p, lhs) && i < len(assign.Rhs) { + if id, ok := assign.Rhs[i].(*ast.Ident); ok && id.Name == "nil" { + state.nilled = true + } + } + } + } + if d, ok := node.(*ast.DeferStmt); ok { + if kind := deferredStreamDrain(p, d, source); kind != streamDrainNone { + state.deferred = kind + } + } + if _, ok := node.(*ast.ReturnStmt); ok && state.pending && (state.deferred == streamDrainNone || (state.deferred == streamDrainCaptured && state.nilled)) { + return true + } + } + for edge, next := range state.block.Succs { + out := state + out.block = next + if loop, ok := state.block.Stmt.(*ast.RangeStmt); ok && state.block.Kind == cfg.KindRangeLoop && source.matches(p, loop.X) && edge == 1 { + out.pending = false + } + if len(state.block.Succs) == 2 && len(state.block.Nodes) > 0 { + cond, _ := state.block.Nodes[len(state.block.Nodes)-1].(ast.Expr) + if closedStreamBranch(p, cond, closedChecks, edge) { + out.pending = false + } + // RunStream returns a non-nil channel. A for ch != nil loop + // exits after its comma-ok receive observes closure. + if state.pending && !state.nilled && streamNonNilBranch(p, cond, source, edge) { + continue + } + } + queue = append(queue, out) + } + } + return false +} + +func closedStreamBranch(p *cop.Pass, cond ast.Expr, checks map[types.Object]bool, edge int) bool { + if cond == nil { + return false + } + if not, ok := cond.(*ast.UnaryExpr); ok && not.Op == token.NOT { + return checks[streamObject(p, not.X)] && edge == 0 + } + return checks[streamObject(p, cond)] && edge == 1 +} + +func streamNonNilBranch(p *cop.Pass, cond ast.Expr, source runStreamSource, edge int) bool { + binary, ok := cond.(*ast.BinaryExpr) + if !ok || !source.matches(p, binary.X) { + return false + } + id, ok := binary.Y.(*ast.Ident) + if !ok || id.Name != "nil" { + return false + } + return (binary.Op == token.NEQ && edge == 1) || (binary.Op == token.EQL && edge == 0) +} + +type streamDrainKind int + +const ( + streamDrainNone streamDrainKind = iota + streamDrainCaptured + streamDrainSnapshot +) + +func deferredStreamDrain(p *cop.Pass, d *ast.DeferStmt, source runStreamSource) streamDrainKind { + fn, ok := d.Call.Fun.(*ast.FuncLit) + if !ok { + return streamDrainNone + } + matches := source.matches + kind := streamDrainCaptured + if len(d.Call.Args) == 1 && source.matches(p, d.Call.Args[0]) && len(fn.Type.Params.List) == 1 && len(fn.Type.Params.List[0].Names) == 1 { + kind = streamDrainSnapshot + param := p.Info.ObjectOf(fn.Type.Params.List[0].Names[0]) + matches = func(p *cop.Pass, expr ast.Expr) bool { return streamObject(p, expr) == param } + } + for _, stmt := range fn.Body.List { + switch stmt := stmt.(type) { + case *ast.ExprStmt: // e.g. cancel() before draining + continue + case *ast.RangeStmt: + if matches(p, stmt.X) && len(stmt.Body.List) == 0 { + return kind + } + return streamDrainNone + default: + return streamDrainNone + } + } + return streamDrainNone +} diff --git a/lint/drain_run_stream_before_release_test.go b/lint/drain_run_stream_before_release_test.go new file mode 100644 index 0000000000..c4cdc0530e --- /dev/null +++ b/lint/drain_run_stream_before_release_test.go @@ -0,0 +1,120 @@ +package main + +import ( + "go/ast" + "go/parser" + "go/token" + "go/types" + "testing" + + "github.com/dgageot/rubocop-go/cop" + "github.com/dgageot/rubocop-go/coptest" + "github.com/dgageot/rubocop-go/prog" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "golang.org/x/tools/go/packages" +) + +func TestDrainRunStreamBeforeRelease(t *testing.T) { + t.Parallel() + for _, tc := range []struct { + name string + body string + want int + }{ + {name: "return on cancellation", body: `events := r.RunStream(); for range events { if stop { return } }`, want: 1}, + {name: "return on error", body: `events := r.RunStream(); for event := range events { switch event { case nil: return } }`, want: 1}, + {name: "direct range", body: `for range r.RunStream() { return }`, want: 1}, + {name: "var declaration", body: `var events = r.RunStream(); for range events { return }`, want: 1}, + {name: "private runStream", body: `events := r.runStream(); for range events { return }`, want: 1}, + {name: "break abandons channel", body: `events := r.RunStream(); for range events { break }`, want: 1}, + {name: "labeled break", body: `events := r.RunStream(); outer: for range events { for { break outer } }`, want: 1}, + {name: "goto skips drain", body: `events := r.RunStream(); for range events { goto out }; out: return`, want: 1}, + {name: "conditional drain", body: `events := r.RunStream(); for range events { break }; if stop { for range events {} }`, want: 1}, + {name: "conditional defer", body: `events := r.RunStream(); if stop { defer func(){ for range events {} }() }; for range events { return }`, want: 1}, + {name: "unrelated channel defer", body: `events := r.RunStream(); defer func(){ for range other {} }(); for range events { return }`, want: 1}, + {name: "deferred drain can return early", body: `events := r.RunStream(); defer func(){ for range events { return } }(); for range events { return }`, want: 1}, + {name: "closure return does not drain owner", body: `events := r.RunStream(); _ = func(){ for range events {} }; for range events { return }`, want: 1}, + {name: "shadowed stream", body: `events := r.RunStream(); for range events { events := other; for range events {}; return }`, want: 1}, + {name: "select cancellation", body: `events := r.RunStream(); for events != nil { select { case _, ok := <-events: if !ok { events = nil; continue }; if stop { return }; case <-other: return } }`, want: 1}, + {name: "nil on cancellation abandons stream", body: `events := r.RunStream(); for events != nil { select { case _, ok := <-events: if !ok { events = nil }; case <-other: events = nil } }`, want: 1}, + {name: "ignored ok does not prove closure", body: `events := r.RunStream(); event, _ := <-events; _ = event; if len(other) == 0 { for range events {} }`, want: 1}, + {name: "captured drain loses nilled channel", body: `events := r.RunStream(); defer func(){ for range events {} }(); for events != nil { select { case <-events: events = nil } }`, want: 1}, + {name: "snapshot drain keeps nilled channel", body: `events := r.RunStream(); defer func(ch <-chan Event){ for range ch {} }(events); for events != nil { select { case <-events: events = nil } }`}, + {name: "capture registered before assignment", body: `var events <-chan Event; defer func(){ for range events {} }(); events = r.RunStream(); for range events { return }`}, + {name: "snapshot registered before assignment", body: `var events <-chan Event; defer func(ch <-chan Event){ for range ch {} }(events); events = r.RunStream(); for range events { return }`, want: 1}, + {name: "receive is not a drain", body: `events := r.RunStream(); <-events`, want: 1}, + {name: "complete range", body: `events := r.RunStream(); for range events { if stop { continue } }`}, + {name: "complete direct range", body: `for range r.RunStream() {}`}, + {name: "break then drain", body: `events := r.RunStream(); for range events { if stop { break } }; for range events {}`}, + {name: "nested switch break", body: `events := r.RunStream(); for event := range events { switch event { case nil: break } }`}, + {name: "nested loop break", body: `events := r.RunStream(); for range events { for { break } }`}, + {name: "nested function return", body: `events := r.RunStream(); for range events { func(){ return }() }`}, + {name: "deferred drain", body: `events := r.RunStream(); defer func(){ cancel(); for range events {} }(); for range events { return }`}, + {name: "deferred snapshot", body: `events := r.RunStream(); defer func(ch <-chan Event){ cancel(); for range ch {} }(events); for range events { return }`}, + {name: "select until closed", body: `events := r.RunStream(); for events != nil { select { case _, ok := <-events: if !ok { events = nil; continue } } }`}, + {name: "return only after closed", body: `events := r.RunStream(); for { _, ok := <-events; if !ok { return } }`}, + {name: "positive ok branch", body: `events := r.RunStream(); for { if _, ok := <-events; ok { continue } else { return } }`}, + {name: "ownership passed to helper", body: `events := r.RunStream(); consume(events)`}, + {name: "unrelated RunStream type", body: `events := otherRuntime{}.RunStream(); for range events { return }`}, + } { + t.Run(tc.name, func(t *testing.T) { + t.Parallel() + offenses := runDrainCop(t, "sample.go", tc.body) + assert.Len(t, offenses, tc.want) + for _, offense := range offenses { + assert.Equal(t, "Lint/DrainRunStreamBeforeRelease", offense.CopName) + assert.Contains(t, offense.Message, "cancel and drain") + } + }) + } +} + +func TestDrainRunStreamBeforeReleaseSkipsTests(t *testing.T) { + t.Parallel() + assert.Empty(t, runDrainCop(t, "sample_test.go", `for range r.RunStream() { return }`)) +} + +func runDrainCop(t *testing.T, filename, body string) []cop.Offense { + t.Helper() + src := `package runtime + type Event interface{ GetAgentName() string } + type Runtime interface{ RunStream() <-chan Event; runStream() <-chan Event } + type otherRuntime struct{} + func (otherRuntime) RunStream() <-chan int { return nil } + func cancel() {} + func consume(<-chan Event) {} + func run(r Runtime, other <-chan Event, stop bool) { ` + body + ` } + ` + fset := token.NewFileSet() + file, err := parser.ParseFile(fset, filename, src, parser.ParseComments) + require.NoError(t, err) + info := &types.Info{ + Types: make(map[ast.Expr]types.TypeAndValue), + Defs: make(map[*ast.Ident]types.Object), + Uses: make(map[*ast.Ident]types.Object), + } + var config types.Config + pkg, err := config.Check("github.com/docker/docker-agent/pkg/runtime", fset, []*ast.File{file}, info) + require.NoError(t, err) + pass := &prog.Pass{Cop: DrainRunStreamBeforeRelease, Program: &prog.Program{Fset: fset, Packages: []*packages.Package{{Syntax: []*ast.File{file}, Types: pkg, TypesInfo: info}}}} + DrainRunStreamBeforeRelease.Check(pass) + return pass.Offenses() +} + +func TestDrainRunStreamBeforeReleaseImportedRuntime(t *testing.T) { + t.Parallel() + offenses := coptest.RunProgram(t, DrainRunStreamBeforeRelease, coptest.ProgramFiles{ + "go.mod": "module github.com/docker/docker-agent\n\ngo 1.27\n", + "pkg/runtime/runtime.go": `package runtime + type Event interface{ GetAgentName() string } + type Runtime interface{ RunStream() <-chan Event } + `, + "consumer/consumer.go": `package consumer + import rt "github.com/docker/docker-agent/pkg/runtime" + func run(r rt.Runtime) { events := r.RunStream(); for range events { return } } + `, + }) + require.Len(t, offenses, 1) + assert.Contains(t, offenses[0].Pos.Filename, "consumer.go") +} diff --git a/lint/main.go b/lint/main.go index 86fa1d4618..76d0527995 100644 --- a/lint/main.go +++ b/lint/main.go @@ -55,6 +55,7 @@ var programCops = []prog.Cop{ SessionStateAccessors, StreamCloseSafety, ExclusiveStreamLease, + DrainRunStreamBeforeRelease, rubocops.NewLintContextConnectivity(), } diff --git a/pkg/a2a/adapter.go b/pkg/a2a/adapter.go index 295d726e7f..3928c58927 100644 --- a/pkg/a2a/adapter.go +++ b/pkg/a2a/adapter.go @@ -2,6 +2,7 @@ package a2a import ( "cmp" + "context" "errors" "fmt" "iter" @@ -127,8 +128,14 @@ func runDockerAgent(ctx agent.InvocationContext, t *team.Team, agentName string, return } - // Run the agent and collect events - eventsChan := rt.RunStream(ctx, sess) + // Early iterator exits must cancel and join the runtime. + streamCtx, cancel := context.WithCancel(ctx) + eventsChan := rt.RunStream(streamCtx, sess) + defer func() { + cancel() + for range eventsChan { + } + }() // Track accumulated content for chunked responses var contentBuilder strings.Builder @@ -152,7 +159,7 @@ func runDockerAgent(ctx agent.InvocationContext, t *team.Team, agentName string, for event := range eventsChan { if ctx.Ended() { - slog.Debug("Invocation ended, stopping agent", "agent", agentName) + slog.DebugContext(ctx, "Invocation ended, stopping agent", "agent", agentName) return } diff --git a/pkg/a2a/adapter_drain_test.go b/pkg/a2a/adapter_drain_test.go new file mode 100644 index 0000000000..831339bf22 --- /dev/null +++ b/pkg/a2a/adapter_drain_test.go @@ -0,0 +1,114 @@ +package a2a + +import ( + "context" + "sync" + "sync/atomic" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/docker/docker-agent/pkg/chat" + "github.com/docker/docker-agent/pkg/model/provider/base" + "github.com/docker/docker-agent/pkg/modelsdev" + "github.com/docker/docker-agent/pkg/servesafety" + "github.com/docker/docker-agent/pkg/session" + "github.com/docker/docker-agent/pkg/tools" +) + +// stallingStream yields its chunks and then parks until the model call is +// cancelled, standing in for a provider that only stops on cancellation. +type stallingStream struct { + ctx context.Context //nolint:containedctx // the model-call context, inspected after the adapter returns + chunks []string + closed atomic.Bool +} + +func (s *stallingStream) Recv() (chat.MessageStreamResponse, error) { + if len(s.chunks) > 0 { + chunk := s.chunks[0] + s.chunks = s.chunks[1:] + return chat.MessageStreamResponse{ + Choices: []chat.MessageStreamChoice{{Delta: chat.MessageDelta{Content: chunk}}}, + }, nil + } + <-s.ctx.Done() + return chat.MessageStreamResponse{}, s.ctx.Err() +} + +func (s *stallingStream) Close() { s.closed.Store(true) } + +type stallingProvider struct { + chunks []string + + mu sync.Mutex + stream *stallingStream +} + +func (p *stallingProvider) ID() modelsdev.ID { return modelsdev.NewID("test", "mock-model") } + +func (p *stallingProvider) CreateChatCompletionStream(ctx context.Context, _ []chat.Message, _ []tools.Tool) (chat.MessageStream, error) { + p.mu.Lock() + defer p.mu.Unlock() + p.stream = &stallingStream{ctx: ctx, chunks: p.chunks} + return p.stream, nil +} + +func (p *stallingProvider) BaseConfig() base.Config { return base.Config{} } + +func (p *stallingProvider) MaxTokens() int { return 0 } + +func (p *stallingProvider) current() *stallingStream { + p.mu.Lock() + defer p.mu.Unlock() + return p.stream +} + +// When the ADK consumer stops early, the adapter must cancel the runtime it +// started and wait for the stream to close before returning, rather than +// leaving the model call running in the background. +func TestRunDockerAgent_EarlyStopCancelsAndDrainsRuntime(t *testing.T) { + t.Parallel() + + for _, tc := range []struct { + name string + chunks []string + stop func(ctx *fakeInvocationContext) bool + }{ + { + name: "consumer breaks", + chunks: []string{"first"}, + stop: func(*fakeInvocationContext) bool { return true }, + }, + { + name: "invocation ended", + chunks: []string{"first", "second"}, + stop: func(ctx *fakeInvocationContext) bool { ctx.EndInvocation(); return false }, + }, + } { + t.Run(tc.name, func(t *testing.T) { + t.Parallel() + + prov := &stallingProvider{chunks: tc.chunks} + tm, root := newTeamWithProvider(prov) + store := session.NewInMemorySessionStore() + ctx := newFakeInvocationContext(t.Context(), "a2a-ctx-drain-"+tc.name, "Hi") + + var yielded int + for _, err := range runDockerAgent(ctx, tm, root.Name(), root, store, servesafety.Resolved{Policy: session.SafetyPolicyRestricted}, testWorkspaceRoot) { + require.NoError(t, err) + yielded++ + if tc.stop(ctx) { + break + } + } + assert.Equal(t, 1, yielded) + + stream := prov.current() + require.NotNil(t, stream, "the model must have been called") + require.ErrorIs(t, stream.ctx.Err(), context.Canceled, "the adapter must cancel the runtime it abandons") + assert.True(t, stream.closed.Load(), "the adapter must drain the runtime stream before returning") + }) + } +} diff --git a/pkg/acp/agent.go b/pkg/acp/agent.go index 354dec7e39..a396a62df5 100644 --- a/pkg/acp/agent.go +++ b/pkg/acp/agent.go @@ -702,7 +702,14 @@ func (a *Agent) runAgent(ctx context.Context, acpSess *Session) error { slog.DebugContext(ctx, "Failed to emit available commands", "error", err) } - eventsChan := acpSess.rt.RunStream(ctx, acpSess.sess) + runCtx, cancel := context.WithCancel(ctx) + eventsChan := acpSess.rt.RunStream(runCtx, acpSess.sess) + // Cancel on handler errors too, before waiting for runtime teardown. + defer func() { + cancel() + for range eventsChan { + } + }() toolCallArgs := map[string]string{} for event := range eventsChan { diff --git a/pkg/acp/runagent_drain_test.go b/pkg/acp/runagent_drain_test.go new file mode 100644 index 0000000000..f5c300d180 --- /dev/null +++ b/pkg/acp/runagent_drain_test.go @@ -0,0 +1,83 @@ +package acp + +import ( + "context" + "testing" + "testing/synctest" + + acpsdk "github.com/coder/acp-go-sdk" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/docker/docker-agent/pkg/runtime" + "github.com/docker/docker-agent/pkg/session" +) + +type drainingPromptRuntime struct { + fakeRuntime + + started chan struct{} + release chan struct{} + stopped chan struct{} + fail bool +} + +func (r *drainingPromptRuntime) RunStream(ctx context.Context, _ *session.Session) <-chan runtime.Event { + events := make(chan runtime.Event) + close(r.started) + go func() { + defer close(events) + defer close(r.stopped) + if r.fail { + events <- &runtime.ToolCallResponseEvent{ToolCallID: "missing-start"} + } + <-ctx.Done() + events <- runtime.Warning("teardown started", "root") + <-r.release + for range 256 { + events <- runtime.Warning("teardown still running", "root") + } + }() + return events +} + +func TestPrompt_DrainsBeforeReleasingTurn(t *testing.T) { + t.Parallel() + for _, fail := range []bool{false, true} { + name := "cancel" + if fail { + name = "handler error" + } + t.Run(name, func(t *testing.T) { + synctest.Test(t, func(t *testing.T) { + rt := &drainingPromptRuntime{started: make(chan struct{}), release: make(chan struct{}), stopped: make(chan struct{}), fail: fail} + agent, sess, _ := newPromptTestAgent(t, rt) + done := promptAsync(agent, t.Context(), promptRequest("first")) + <-rt.started + if !fail { + require.NoError(t, agent.Cancel(t.Context(), acpsdk.CancelNotification{SessionId: testSessionID})) + } + synctest.Wait() + select { + case <-sess.turns: + t.Error("turn token released before runtime teardown") + default: + } + close(rt.release) + synctest.Wait() + result := <-done + if fail { + require.ErrorContains(t, result.err, "missing tool call arguments") + } else { + require.NoError(t, result.err) + assert.Equal(t, acpsdk.StopReasonCancelled, result.response.StopReason) + } + select { + case <-rt.stopped: + default: + t.Error("teardown events were not drained") + } + }) + }) + } +} diff --git a/pkg/cli/runner.go b/pkg/cli/runner.go index ea15f98902..5f028a549a 100644 --- a/pkg/cli/runner.go +++ b/pkg/cli/runner.go @@ -118,23 +118,32 @@ func Run(ctx context.Context, out *Printer, cfg Config, rt runtime.Runtime, sess sess.AddMessage(userMsg) sess.AddAttachedFile(attachedPath) + // Stop only this turn; subsequent user messages still need the parent context. + runCtx, stopRun := context.WithCancel(ctx) + events := rt.RunStream(runCtx, sess) + defer func() { + stopRun() + for range events { + } + }() + if cfg.OutputJSON { - for event := range rt.RunStream(ctx, sess) { + for event := range events { switch e := event.(type) { case *runtime.ToolCallConfirmationEvent: // JSON mode has no user at stdin — reject unconditionally. // A confirmation event under AutoApprove means a // preempt-yolo hook overrode --yolo; the safe answer is // still Reject (the hook said Ask, not Approve). - rt.Resume(ctx, runtime.ResumeReject("")) + rt.Resume(runCtx, runtime.ResumeReject("")) case *runtime.ElicitationRequestEvent: - _ = rt.ResumeElicitation(ctx, "decline", nil, e.ElicitationID) + _ = rt.ResumeElicitation(runCtx, "decline", nil, e.ElicitationID) case *runtime.MaxIterationsReachedEvent: switch handleMaxIterationsAutoApprove(cfg.AutoApprove, &autoExtensions, e.MaxIterations) { case maxIterContinue: - rt.Resume(ctx, runtime.ResumeApprove()) + rt.Resume(runCtx, runtime.ResumeApprove()) default: // maxIterStop or maxIterPrompt (no interactive prompt in JSON mode) - rt.Resume(ctx, runtime.ResumeReject("")) + rt.Resume(runCtx, runtime.ResumeReject("")) return nil } case *runtime.ErrorEvent: @@ -154,7 +163,7 @@ func Run(ctx context.Context, out *Printer, cfg Config, rt runtime.Runtime, sess firstLoop := true lastAgent := rt.CurrentAgentName(ctx) var lastConfirmedToolCallID string - for event := range rt.RunStream(ctx, sess) { + for event := range events { agentName := event.GetAgentName() if agentName != "" && (firstLoop || lastAgent != agentName) { if !firstLoop { @@ -178,15 +187,15 @@ func Run(ctx context.Context, out *Printer, cfg Config, rt runtime.Runtime, sess lastConfirmedToolCallID = e.ToolCall.ID // Store the ID to avoid duplicate printing switch result { case ConfirmationApprove: - rt.Resume(ctx, runtime.ResumeApprove()) + rt.Resume(runCtx, runtime.ResumeApprove()) case ConfirmationApproveBalanced: sess.SetSafetyPolicy(session.SafetyPolicyBalanced) - rt.Resume(ctx, runtime.ResumeApproveBalanced()) + rt.Resume(runCtx, runtime.ResumeApproveBalanced()) case ConfirmationApproveSession: sess.SetSafetyPolicy(session.SafetyPolicyAutonomous) - rt.Resume(ctx, runtime.ResumeApproveAutonomous()) + rt.Resume(runCtx, runtime.ResumeApproveAutonomous()) case ConfirmationReject: - rt.Resume(ctx, runtime.ResumeReject("")) + rt.Resume(runCtx, runtime.ResumeReject("")) lastConfirmedToolCallID = "" // Clear on reject since tool won't execute case ConfirmationAbort: // Stop the agent loop immediately @@ -228,20 +237,20 @@ func Run(ctx context.Context, out *Printer, cfg Config, rt runtime.Runtime, sess case *runtime.MaxIterationsReachedEvent: switch handleMaxIterationsAutoApprove(cfg.AutoApprove, &autoExtensions, e.MaxIterations) { case maxIterContinue: - rt.Resume(ctx, runtime.ResumeApprove()) + rt.Resume(runCtx, runtime.ResumeApprove()) case maxIterStop: - rt.Resume(ctx, runtime.ResumeReject("")) + rt.Resume(runCtx, runtime.ResumeReject("")) return nil case maxIterPrompt: result := out.PromptMaxIterationsContinue(ctx, e.MaxIterations) switch result { case ConfirmationApprove: - rt.Resume(ctx, runtime.ResumeApprove()) + rt.Resume(runCtx, runtime.ResumeApprove()) case ConfirmationReject: - rt.Resume(ctx, runtime.ResumeReject("")) + rt.Resume(runCtx, runtime.ResumeReject("")) return nil case ConfirmationAbort: - rt.Resume(ctx, runtime.ResumeReject("")) + rt.Resume(runCtx, runtime.ResumeReject("")) return nil } } @@ -250,7 +259,7 @@ func Run(ctx context.Context, out *Printer, cfg Config, rt runtime.Runtime, sess if !ok || serverURL == "" { // Keep draining after declining forms so follow-up events cannot stall the turn. slog.WarnContext(ctx, "Declining elicitation without form support in CLI mode", "message", e.Message) - _ = rt.ResumeElicitation(ctx, "decline", nil, e.ElicitationID) + _ = rt.ResumeElicitation(runCtx, "decline", nil, e.ElicitationID) continue } @@ -262,9 +271,9 @@ func Run(ctx context.Context, out *Printer, cfg Config, rt runtime.Runtime, sess switch result { case ConfirmationApprove: - _ = rt.ResumeElicitation(ctx, "accept", nil, e.ElicitationID) + _ = rt.ResumeElicitation(runCtx, "accept", nil, e.ElicitationID) case ConfirmationReject: - _ = rt.ResumeElicitation(ctx, "decline", nil, e.ElicitationID) + _ = rt.ResumeElicitation(runCtx, "decline", nil, e.ElicitationID) return errors.New("OAuth authorization rejected by user") } } diff --git a/pkg/cli/runner_drain_test.go b/pkg/cli/runner_drain_test.go new file mode 100644 index 0000000000..626fdc680d --- /dev/null +++ b/pkg/cli/runner_drain_test.go @@ -0,0 +1,128 @@ +package cli + +import ( + "bytes" + "context" + "strings" + "testing" + + "gotest.tools/v3/assert" + + "github.com/docker/docker-agent/pkg/runtime" + "github.com/docker/docker-agent/pkg/session" +) + +// stallingRunStream emits the trigger events, waits for the run context to be +// cancelled, then floods more events than any buffer holds before closing. +// done is closed only once every event has been consumed, so Run can only +// observe it closed if it cancelled the producer and drained the stream. +func stallingRunStream(triggers []runtime.Event, done chan struct{}) func(context.Context, *session.Session) <-chan runtime.Event { + return func(ctx context.Context, _ *session.Session) <-chan runtime.Event { + ch := make(chan runtime.Event) + go func() { + defer close(ch) + defer close(done) + for _, e := range triggers { + ch <- e + } + <-ctx.Done() + for range 256 { + ch <- runtime.Warning("trailing", "test") + } + }() + return ch + } +} + +func repeatMaxIterEvents(n int) []runtime.Event { + events := make([]runtime.Event, n) + for i := range events { + events[i] = maxIterEvent(10) + } + return events +} + +// Every early return out of a turn (error, safety cap) must cancel the +// runtime stream and drain it before the turn ends. +func TestRunEarlyReturnCancelsAndDrainsStream(t *testing.T) { + t.Parallel() + + for _, tc := range []struct { + name string + cfg Config + triggers []runtime.Event + wantErr string + }{ + { + name: "json error", + cfg: Config{OutputJSON: true}, + triggers: []runtime.Event{runtime.Error("model failed")}, + wantErr: "model failed", + }, + { + name: "json max iterations cap", + cfg: Config{OutputJSON: true, AutoApprove: true}, + triggers: repeatMaxIterEvents(maxAutoExtensions + 1), + }, + { + name: "text max iterations cap", + cfg: Config{AutoApprove: true}, + triggers: repeatMaxIterEvents(maxAutoExtensions + 1), + }, + } { + t.Run(tc.name, func(t *testing.T) { + t.Parallel() + + done := make(chan struct{}) + rt := &mockRuntime{runStreamFn: stallingRunStream(tc.triggers, done)} + + var buf bytes.Buffer + err := Run(t.Context(), NewPrinter(&buf), tc.cfg, rt, session.New(), []string{"hello"}) + if tc.wantErr == "" { + assert.NilError(t, err) + } else { + assert.ErrorContains(t, err, tc.wantErr) + } + + select { + case <-done: + default: + t.Fatal("the turn ended before cancelling and draining the runtime stream") + } + assert.Check(t, !strings.Contains(buf.String(), "trailing"), "drained events must not be printed: %q", buf.String()) + }) + } +} + +// A turn that ends early must only cancel its own stream: later user messages +// still run with a live context. +func TestRunEarlyReturnKeepsFollowUpTurnsAlive(t *testing.T) { + t.Parallel() + + var turns int + var secondTurnCtxErr error + rt := &mockRuntime{ + runStreamFn: func(ctx context.Context, _ *session.Session) <-chan runtime.Event { + turns++ + ch := make(chan runtime.Event, maxAutoExtensions+1) + if turns == 1 { + for _, e := range repeatMaxIterEvents(maxAutoExtensions + 1) { + ch <- e + } + } else { + secondTurnCtxErr = ctx.Err() + ch <- runtime.AgentChoice("test", "sess", "second turn answer") + } + close(ch) + return ch + }, + } + + var buf bytes.Buffer + err := Run(t.Context(), NewPrinter(&buf), Config{AutoApprove: true}, rt, session.New(), []string{"one", "two"}) + assert.NilError(t, err) + + assert.Equal(t, turns, 2) + assert.NilError(t, secondTurnCtxErr) + assert.Check(t, strings.Contains(buf.String(), "second turn answer"), "the follow-up turn must run: %q", buf.String()) +} diff --git a/pkg/runtime/loop.go b/pkg/runtime/loop.go index e3a4d9df0f..0c4a7feacc 100644 --- a/pkg/runtime/loop.go +++ b/pkg/runtime/loop.go @@ -1152,7 +1152,14 @@ func (r *LocalRuntime) runTurn( // messages. This is a convenience wrapper around RunStream for non-streaming // callers. func (r *LocalRuntime) Run(ctx context.Context, sess *session.Session) ([]session.Message, error) { + ctx, cancel := context.WithCancel(ctx) events := r.RunStream(ctx, sess) + // Cancel before draining so error returns cannot strand the producer. + defer func() { + cancel() + for range events { + } + }() for event := range events { if errEvent, ok := event.(*ErrorEvent); ok { return nil, fmt.Errorf("%s", errEvent.Error) diff --git a/pkg/runtime/remote_runtime.go b/pkg/runtime/remote_runtime.go index b264bf09cd..16947b9a13 100644 --- a/pkg/runtime/remote_runtime.go +++ b/pkg/runtime/remote_runtime.go @@ -313,7 +313,14 @@ func (r *RemoteRuntime) RunStream(ctx context.Context, sess *session.Session) <- // Run starts the agent's interaction loop and returns the final messages func (r *RemoteRuntime) Run(ctx context.Context, sess *session.Session) ([]session.Message, error) { + ctx, cancel := context.WithCancel(ctx) eventsChan := r.RunStream(ctx, sess) + // Cancel before draining so error returns cannot strand the forwarder. + defer func() { + cancel() + for range eventsChan { + } + }() for event := range eventsChan { if errEvent, ok := event.(*ErrorEvent); ok { diff --git a/pkg/runtime/run_drain_test.go b/pkg/runtime/run_drain_test.go new file mode 100644 index 0000000000..6ef5049d70 --- /dev/null +++ b/pkg/runtime/run_drain_test.go @@ -0,0 +1,107 @@ +package runtime + +import ( + "context" + "errors" + "testing" + "testing/synctest" + + "github.com/stretchr/testify/require" + + "github.com/docker/docker-agent/pkg/agent" + "github.com/docker/docker-agent/pkg/api" + "github.com/docker/docker-agent/pkg/config/latest" + "github.com/docker/docker-agent/pkg/session" + "github.com/docker/docker-agent/pkg/team" +) + +// stallingObserver stands in for teardown work (e.g. persistence) that only +// completes once the run context is cancelled: OnEvent parks on StreamStopped +// until ctx is done, so the channel cannot close before the consumer cancels. +type stallingObserver struct{ done chan struct{} } + +func (stallingObserver) OnRunStart(context.Context, *session.Session) {} + +func (o stallingObserver) OnEvent(ctx context.Context, _ *session.Session, event Event) { + if _, ok := event.(*StreamStoppedEvent); ok { + <-ctx.Done() + close(o.done) + } +} + +// Run returning on the first ErrorEvent must not abandon the stream: it has +// to cancel the runtime it started and wait for the channel to close, so no +// teardown work outlives the call. +func TestLocalRuntime_RunCancelsAndDrainsStreamOnError(t *testing.T) { + t.Parallel() + + synctest.Test(t, func(t *testing.T) { + prov := &failingProvider{id: "test/failing", err: errors.New("401 unauthorized")} + root := agent.New("root", "test", agent.WithModel(prov)) + obs := stallingObserver{done: make(chan struct{})} + rt, err := NewLocalRuntime(t.Context(), team.New(team.WithAgents(root)), + WithSessionCompaction(false), + WithModelStore(mockModelStore{}), + WithEventObserver(obs), + ) + require.NoError(t, err) + + _, err = rt.Run(t.Context(), session.New(session.WithUserMessage("hi"))) + require.ErrorContains(t, err, "401 unauthorized") + + select { + case <-obs.done: + default: + t.Fatal("Run returned before cancelling and draining its stream") + } + }) +} + +// stallingRemoteClient serves a stream that fails, waits for the caller to +// cancel, then floods more events than the runtime buffer holds. Only a +// consumer that cancels and drains lets the producer finish. +type stallingRemoteClient struct { + stubRemoteClient + + done chan struct{} +} + +func (c *stallingRemoteClient) RunAgent(ctx context.Context, _, _ string, _ []api.Message, _ string) (<-chan Event, error) { + ch := make(chan Event) + go func() { + defer close(ch) + defer close(c.done) + ch <- Error("remote failure") + <-ctx.Done() + for range 2 * defaultEventChannelCapacity { + ch <- Warning("trailing", "test") + } + }() + return ch, nil +} + +func (c *stallingRemoteClient) RunAgentWithAgentName(ctx context.Context, sessionID, agentFile, _ string, msgs []api.Message, model string) (<-chan Event, error) { + return c.RunAgent(ctx, sessionID, agentFile, msgs, model) +} + +func TestRemoteRuntime_RunCancelsAndDrainsStreamOnError(t *testing.T) { + t.Parallel() + + synctest.Test(t, func(t *testing.T) { + client := &stallingRemoteClient{ + stubRemoteClient: stubRemoteClient{cfg: &latest.Config{Agents: latest.Agents{{Name: "test"}}}}, + done: make(chan struct{}), + } + rt, err := NewRemoteRuntime(client) + require.NoError(t, err) + + _, err = rt.Run(t.Context(), session.New(session.WithUserMessage("hi"))) + require.EqualError(t, err, "remote failure") + + select { + case <-client.done: + default: + t.Fatal("Run returned before cancelling and draining its stream") + } + }) +} diff --git a/pkg/runtime/runtime.go b/pkg/runtime/runtime.go index ba2a51b1f0..a99c55762b 100644 --- a/pkg/runtime/runtime.go +++ b/pkg/runtime/runtime.go @@ -77,7 +77,9 @@ type Runtime interface { EmitAgentInfo(ctx context.Context, events EventSink) // ResetStartupInfo resets the startup info emission flag, allowing re-emission ResetStartupInfo() - // RunStream starts the agent's interaction loop and returns a channel of events + // RunStream starts the agent's interaction loop and returns a channel of events. + // Consumers must drain it to closure before releasing turn ownership, even + // after cancellation. Cancel the supplied context before draining on errors. RunStream(ctx context.Context, sess *session.Session) <-chan Event // Run starts the agent's interaction loop and returns the final messages Run(ctx context.Context, sess *session.Session) ([]session.Message, error) diff --git a/pkg/server/session_manager.go b/pkg/server/session_manager.go index c8d0e27b59..f13b8ef6ef 100644 --- a/pkg/server/session_manager.go +++ b/pkg/server/session_manager.go @@ -1167,6 +1167,12 @@ func (sm *SessionManager) RunSession(ctx context.Context, sessionID, agentFilena } stream := runtimeSession.runtime.RunStream(streamCtx, sess) + // Teardown must finish before the deferred streaming.Unlock. + defer func(events <-chan runtime.Event) { + cancel() + for range events { + } + }(stream) for stream != nil { select { case event, ok := <-titleEvents: diff --git a/pkg/server/session_manager_drain_test.go b/pkg/server/session_manager_drain_test.go new file mode 100644 index 0000000000..2d3e9b10da --- /dev/null +++ b/pkg/server/session_manager_drain_test.go @@ -0,0 +1,88 @@ +package server + +import ( + "context" + "testing" + "testing/synctest" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/docker/docker-agent/pkg/api" + "github.com/docker/docker-agent/pkg/runtime" + "github.com/docker/docker-agent/pkg/session" +) + +type drainingSessionRuntime struct { + fakeRuntime + + started chan struct{} + release chan struct{} + stopped chan struct{} +} + +func (r *drainingSessionRuntime) RunStream(ctx context.Context, _ *session.Session) <-chan runtime.Event { + events := make(chan runtime.Event) + close(r.started) + go func() { + defer close(events) + defer close(r.stopped) + <-ctx.Done() + events <- runtime.Warning("teardown started", "root") + <-r.release + for range 256 { + events <- runtime.Warning("teardown still running", "root") + } + }() + return events +} + +func TestRunSession_CancellationDrainsBeforeUnlock(t *testing.T) { + t.Parallel() + synctest.Test(t, func(t *testing.T) { + ctx, cancel := context.WithCancel(t.Context()) + defer cancel() + rt := &drainingSessionRuntime{started: make(chan struct{}), release: make(chan struct{}), stopped: make(chan struct{})} + sess := session.New() + sm := newTestSessionManager(t, sess, rt) + output, err := sm.RunSession(ctx, sess.ID, "agent", "root", []api.Message{{Content: "first"}}, "") + require.NoError(t, err) + outputClosed := make(chan struct{}) + go func() { + for range output { + } + close(outputClosed) + }() + <-rt.started + cancel() + synctest.Wait() + rs, ok := sm.runtimeSessions.Load(sess.ID) + require.True(t, ok) + unlocked := rs.streaming.TryLock() + if unlocked { + rs.streaming.Unlock() + } + assert.False(t, unlocked, "turn ownership released before runtime teardown") + if !unlocked { + _, err = sm.RunSession(t.Context(), sess.ID, "agent", "root", []api.Message{{Content: "second"}}, "") + require.ErrorIs(t, err, ErrSessionBusy) + } + select { + case <-outputClosed: + t.Error("forwarder returned before runtime teardown") + default: + } + close(rt.release) + synctest.Wait() + select { + case <-rt.stopped: + default: + t.Error("teardown events were not drained") + } + select { + case <-outputClosed: + default: + t.Error("forwarder did not finish") + } + }) +}