2086 lines
64 KiB
Go
2086 lines
64 KiB
Go
/*
|
|
* Copyright 2026 CloudWeGo Authors
|
|
*
|
|
* Licensed 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 adk
|
|
|
|
import (
|
|
"context"
|
|
"errors"
|
|
"fmt"
|
|
"sync/atomic"
|
|
"testing"
|
|
"time"
|
|
|
|
"github.com/stretchr/testify/assert"
|
|
"github.com/stretchr/testify/require"
|
|
|
|
"github.com/cloudwego/eino/components/model"
|
|
"github.com/cloudwego/eino/components/tool"
|
|
"github.com/cloudwego/eino/compose"
|
|
"github.com/cloudwego/eino/schema"
|
|
)
|
|
|
|
// --- helpers shared across edge-case tests ---
|
|
|
|
// blockingChatModel blocks until unblockCh is closed, then returns a fixed response.
|
|
type blockingChatModel struct {
|
|
unblockCh chan struct{}
|
|
response *schema.Message
|
|
started chan struct{}
|
|
callCount int32
|
|
}
|
|
|
|
func newBlockingChatModel(response *schema.Message) *blockingChatModel {
|
|
return &blockingChatModel{
|
|
unblockCh: make(chan struct{}),
|
|
response: response,
|
|
started: make(chan struct{}, 1),
|
|
}
|
|
}
|
|
|
|
func (m *blockingChatModel) Generate(ctx context.Context, _ []*schema.Message, _ ...model.Option) (*schema.Message, error) {
|
|
atomic.AddInt32(&m.callCount, 1)
|
|
select {
|
|
case m.started <- struct{}{}:
|
|
default:
|
|
}
|
|
<-m.unblockCh
|
|
return m.response, nil
|
|
}
|
|
|
|
func (m *blockingChatModel) Stream(ctx context.Context, _ []*schema.Message, _ ...model.Option) (*schema.StreamReader[*schema.Message], error) {
|
|
atomic.AddInt32(&m.callCount, 1)
|
|
select {
|
|
case m.started <- struct{}{}:
|
|
default:
|
|
}
|
|
<-m.unblockCh
|
|
return schema.StreamReaderFromArray([]*schema.Message{m.response}), nil
|
|
}
|
|
|
|
func (m *blockingChatModel) BindTools(_ []*schema.ToolInfo) error { return nil }
|
|
|
|
// errorChatModel returns an error from Generate/Stream.
|
|
type errorChatModel struct {
|
|
err error
|
|
started chan struct{}
|
|
}
|
|
|
|
func (m *errorChatModel) Generate(_ context.Context, _ []*schema.Message, _ ...model.Option) (*schema.Message, error) {
|
|
if m.started != nil {
|
|
select {
|
|
case m.started <- struct{}{}:
|
|
default:
|
|
}
|
|
}
|
|
return nil, m.err
|
|
}
|
|
|
|
func (m *errorChatModel) Stream(_ context.Context, _ []*schema.Message, _ ...model.Option) (*schema.StreamReader[*schema.Message], error) {
|
|
return nil, m.err
|
|
}
|
|
|
|
func (m *errorChatModel) BindTools(_ []*schema.ToolInfo) error { return nil }
|
|
|
|
// plainResponseModel returns immediately with a fixed text response (no tool calls).
|
|
type plainResponseModel struct {
|
|
text string
|
|
}
|
|
|
|
func (m *plainResponseModel) Generate(_ context.Context, _ []*schema.Message, _ ...model.Option) (*schema.Message, error) {
|
|
return schema.AssistantMessage(m.text, nil), nil
|
|
}
|
|
|
|
func (m *plainResponseModel) Stream(_ context.Context, _ []*schema.Message, _ ...model.Option) (*schema.StreamReader[*schema.Message], error) {
|
|
return schema.StreamReaderFromArray([]*schema.Message{schema.AssistantMessage(m.text, nil)}), nil
|
|
}
|
|
|
|
func (m *plainResponseModel) BindTools(_ []*schema.ToolInfo) error { return nil }
|
|
|
|
// blockingTool blocks until unblockCh is closed.
|
|
type blockingTool struct {
|
|
name string
|
|
unblockCh chan struct{}
|
|
started chan struct{}
|
|
callCount int32
|
|
}
|
|
|
|
func newBlockingTool(name string) *blockingTool {
|
|
return &blockingTool{
|
|
name: name,
|
|
unblockCh: make(chan struct{}),
|
|
started: make(chan struct{}, 4),
|
|
}
|
|
}
|
|
|
|
func (t *blockingTool) Info(_ context.Context) (*schema.ToolInfo, error) {
|
|
return &schema.ToolInfo{
|
|
Name: t.name,
|
|
Desc: "blocking tool",
|
|
ParamsOneOf: schema.NewParamsOneOfByParams(map[string]*schema.ParameterInfo{
|
|
"input": {Type: "string"},
|
|
}),
|
|
}, nil
|
|
}
|
|
|
|
func (t *blockingTool) InvokableRun(_ context.Context, _ string, _ ...tool.Option) (string, error) {
|
|
atomic.AddInt32(&t.callCount, 1)
|
|
select {
|
|
case t.started <- struct{}{}:
|
|
default:
|
|
}
|
|
<-t.unblockCh
|
|
return "result", nil
|
|
}
|
|
|
|
func toolCallMsg(calls ...schema.ToolCall) *schema.Message {
|
|
return &schema.Message{Role: schema.Assistant, ToolCalls: calls}
|
|
}
|
|
|
|
func toolCall(id, name, args string) schema.ToolCall {
|
|
return schema.ToolCall{ID: id, Type: "function", Function: schema.FunctionCall{Name: name, Arguments: args}}
|
|
}
|
|
|
|
func drainEvents(iter *AsyncIterator[*AgentEvent]) ([]*AgentEvent, bool) {
|
|
var events []*AgentEvent
|
|
hasCancelError := false
|
|
for {
|
|
e, ok := iter.Next()
|
|
if !ok {
|
|
break
|
|
}
|
|
events = append(events, e)
|
|
var ce *CancelError
|
|
if e.Err != nil && errors.As(e.Err, &ce) {
|
|
hasCancelError = true
|
|
}
|
|
}
|
|
return events, hasCancelError
|
|
}
|
|
|
|
// --- tests ---
|
|
|
|
// TestWithCancel_BeforeExecutionStarts verifies that a cancel issued before
|
|
// the graph begins executing still produces a CancelError without invoking
|
|
// the model or tools.
|
|
func TestWithCancel_BeforeExecutionStarts(t *testing.T) {
|
|
ctx := context.Background()
|
|
|
|
blk := newBlockingChatModel(toolCallMsg(toolCall("c1", "bt", `{"input":"x"}`)))
|
|
bt := newBlockingTool("bt")
|
|
|
|
agent, err := NewChatModelAgent(ctx, &ChatModelAgentConfig{
|
|
Name: "TestAgent",
|
|
Description: "test",
|
|
Model: blk,
|
|
ToolsConfig: ToolsConfig{
|
|
ToolsNodeConfig: compose.ToolsNodeConfig{Tools: []tool.BaseTool{bt}},
|
|
},
|
|
})
|
|
assert.NoError(t, err)
|
|
|
|
cancelOpt, cancelFn := WithCancel()
|
|
|
|
// Extract the cancelContext so we can wait for cancelChan to close,
|
|
// ensuring the cancel is fully registered before Run starts.
|
|
cc := getCommonOptions(nil, cancelOpt).cancelCtx
|
|
|
|
// Call cancel BEFORE calling agent.Run.
|
|
// The cancelFunc must succeed (not hang) even though execution hasn't started.
|
|
cancelDone := make(chan error, 1)
|
|
go func() {
|
|
handle, _ := cancelFn()
|
|
cancelDone <- handle.Wait()
|
|
}()
|
|
|
|
// Wait for cancelChan to close so the pre-execution check in runFunc
|
|
// deterministically sees shouldCancel()=true (eliminates goroutine scheduling race).
|
|
<-cc.cancelChan
|
|
|
|
// Now start the run — it should see shouldCancel()=true and emit CancelError immediately.
|
|
iter := agent.Run(ctx, &AgentInput{Messages: []Message{schema.UserMessage("hi")}}, cancelOpt)
|
|
|
|
_, hasCancelError := drainEvents(iter)
|
|
assert.True(t, hasCancelError, "expected CancelError when cancel precedes execution")
|
|
|
|
// cancelFn must have already returned (or return quickly now that doneChan is closed).
|
|
select {
|
|
case cancelErr := <-cancelDone:
|
|
// Either nil (cancel handled) or ErrExecutionEnded is acceptable
|
|
// depending on exact timing; what matters is it didn't hang.
|
|
_ = cancelErr
|
|
case <-time.After(3 * time.Second):
|
|
t.Fatal("cancelFn blocked indefinitely after pre-start cancel")
|
|
}
|
|
|
|
// Model and tool must not have been invoked.
|
|
assert.Equal(t, int32(0), atomic.LoadInt32(&bt.callCount), "tool must not be called")
|
|
}
|
|
|
|
// TestWithCancel_AfterCompletion verifies cancelFn returns ErrExecutionEnded
|
|
// when called after a normal run finishes.
|
|
func TestWithCancel_AfterCompletion(t *testing.T) {
|
|
ctx := context.Background()
|
|
|
|
agent, err := NewChatModelAgent(ctx, &ChatModelAgentConfig{
|
|
Name: "TestAgent",
|
|
Description: "test",
|
|
Model: &plainResponseModel{text: "done"},
|
|
})
|
|
require.NoError(t, err)
|
|
|
|
cancelOpt, cancelFn := WithCancel()
|
|
iter := agent.Run(ctx, &AgentInput{Messages: []Message{schema.UserMessage("hi")}}, cancelOpt)
|
|
|
|
// Drain all events so the run completes.
|
|
for {
|
|
_, ok := iter.Next()
|
|
if !ok {
|
|
break
|
|
}
|
|
}
|
|
|
|
handle, _ := cancelFn()
|
|
cancelErr := handle.Wait()
|
|
assert.ErrorIs(t, cancelErr, ErrExecutionEnded)
|
|
}
|
|
|
|
// TestWithCancel_DerivedAgentToolCancelContextMarkedDoneAfterRun verifies that
|
|
// an explicitly derived AgentTool child cancel context is owned by the child run,
|
|
// even when the Go context also carries the parent cancel context.
|
|
func TestWithCancel_DerivedAgentToolCancelContextMarkedDoneAfterRun(t *testing.T) {
|
|
ctx := context.Background()
|
|
|
|
agent, err := NewChatModelAgent(ctx, &ChatModelAgentConfig{
|
|
Name: "ChildAgent",
|
|
Description: "test child agent",
|
|
Model: &plainResponseModel{text: "done"},
|
|
})
|
|
require.NoError(t, err)
|
|
|
|
parent := newCancelContext()
|
|
parentCtx := withCancelContext(ctx, parent)
|
|
child := parent.deriveAgentToolCancelContext(parentCtx)
|
|
|
|
childOpt := WrapImplSpecificOptFn(func(o *options) {
|
|
o.cancelCtx = child
|
|
})
|
|
iter := agent.Run(parentCtx, &AgentInput{Messages: []Message{schema.UserMessage("hi")}}, childOpt)
|
|
for {
|
|
_, ok := iter.Next()
|
|
if !ok {
|
|
break
|
|
}
|
|
}
|
|
|
|
select {
|
|
case <-child.doneChan:
|
|
case <-time.After(time.Second):
|
|
t.Fatal("derived AgentTool cancel context was not marked done after child run completion")
|
|
}
|
|
}
|
|
|
|
// TestWithCancel_AfterBusinessInterrupt verifies cancelFn returns ErrExecutionEnded
|
|
// when called after the agent has been interrupted by business logic.
|
|
func TestWithCancel_AfterBusinessInterrupt(t *testing.T) {
|
|
ctx := context.Background()
|
|
|
|
// Use a model that triggers a compose.Interrupt so the agent stops with an interrupt.
|
|
interruptModel := &interruptingChatModel{}
|
|
|
|
agent, err := NewChatModelAgent(ctx, &ChatModelAgentConfig{
|
|
Name: "TestAgent",
|
|
Description: "test",
|
|
Model: interruptModel,
|
|
})
|
|
require.NoError(t, err)
|
|
|
|
store := newCancelTestStore()
|
|
runner := NewRunner(ctx, RunnerConfig{
|
|
Agent: agent,
|
|
CheckPointStore: store,
|
|
})
|
|
|
|
cancelOpt, cancelFn := WithCancel()
|
|
iter := runner.Run(ctx, []Message{schema.UserMessage("hi")}, cancelOpt, WithCheckPointID("biz-interrupt-1"))
|
|
|
|
// Drain — expect an interrupt action event, no cancel error.
|
|
var gotInterrupt bool
|
|
for {
|
|
e, ok := iter.Next()
|
|
if !ok {
|
|
break
|
|
}
|
|
if e.Action != nil && e.Action.Interrupted != nil {
|
|
gotInterrupt = true
|
|
}
|
|
}
|
|
assert.True(t, gotInterrupt, "expected business interrupt event")
|
|
|
|
handle, _ := cancelFn()
|
|
cancelErr := handle.Wait()
|
|
assert.ErrorIs(t, cancelErr, ErrExecutionEnded)
|
|
}
|
|
|
|
// TestWithCancel_AfterError verifies cancelFn returns ErrExecutionEnded
|
|
// when called after the agent errors out.
|
|
func TestWithCancel_AfterError(t *testing.T) {
|
|
ctx := context.Background()
|
|
|
|
modelErr := errors.New("model exploded")
|
|
agent, err := NewChatModelAgent(ctx, &ChatModelAgentConfig{
|
|
Name: "TestAgent",
|
|
Description: "test",
|
|
Model: &errorChatModel{err: modelErr},
|
|
})
|
|
require.NoError(t, err)
|
|
|
|
cancelOpt, cancelFn := WithCancel()
|
|
iter := agent.Run(ctx, &AgentInput{Messages: []Message{schema.UserMessage("hi")}}, cancelOpt)
|
|
|
|
for {
|
|
_, ok := iter.Next()
|
|
if !ok {
|
|
break
|
|
}
|
|
}
|
|
|
|
handle, _ := cancelFn()
|
|
cancelErr := handle.Wait()
|
|
assert.ErrorIs(t, cancelErr, ErrExecutionEnded)
|
|
}
|
|
|
|
// TestWithCancel_TimeoutEscalation tests that WithAgentCancelTimeout causes the
|
|
// cancel to escalate to immediate when the safe-point hasn't fired yet, and
|
|
// that the resulting CancelError has Escalated=true.
|
|
//
|
|
// Strategy: use CancelAfterChatModel mode. The model blocks (never completes),
|
|
// so the safe-point can't fire naturally. After the timeout, escalateToImmediate
|
|
// closes immediateChan which aborts the model stream via cancelMonitoredModel
|
|
// and causes a CancelError — no compose graph-interrupt races involved.
|
|
func TestWithCancel_TimeoutEscalation(t *testing.T) {
|
|
ctx := context.Background()
|
|
|
|
blk := newBlockingChatModel(schema.AssistantMessage("hello", nil))
|
|
|
|
agent, err := NewChatModelAgent(ctx, &ChatModelAgentConfig{
|
|
Name: "TestAgent",
|
|
Description: "test",
|
|
Model: blk,
|
|
})
|
|
require.NoError(t, err)
|
|
|
|
runner := NewRunner(ctx, RunnerConfig{
|
|
Agent: agent,
|
|
EnableStreaming: true, // use streaming so cancelMonitoredModel.Stream is exercised
|
|
})
|
|
|
|
timeout := 300 * time.Millisecond
|
|
// CancelAfterChatModel + timeout: safe-point can't fire (model never finishes),
|
|
// so after 300ms the timeout goroutine escalates to immediate.
|
|
cancelOpt, cancelFn := WithCancel()
|
|
iter := runner.Run(ctx, []Message{schema.UserMessage("go")}, cancelOpt)
|
|
|
|
select {
|
|
case <-blk.started:
|
|
case <-time.After(5 * time.Second):
|
|
t.Fatal("model did not start")
|
|
}
|
|
|
|
// Fire cancelFn; it will wait for escalation to complete.
|
|
start := time.Now()
|
|
handle, _ := cancelFn(WithAgentCancelMode(CancelAfterChatModel), WithAgentCancelTimeout(timeout))
|
|
cancelErr := handle.Wait()
|
|
elapsed := time.Since(start)
|
|
|
|
assert.ErrorIs(t, cancelErr, ErrCancelTimeout, "cancel should return ErrCancelTimeout after timeout escalation")
|
|
assert.True(t, elapsed >= timeout, "should wait at least the timeout duration, elapsed=%v", elapsed)
|
|
assert.True(t, elapsed < 3*time.Second, "should complete shortly after timeout, elapsed=%v", elapsed)
|
|
|
|
var cancelError *CancelError
|
|
for {
|
|
e, ok := iter.Next()
|
|
if !ok {
|
|
break
|
|
}
|
|
var ce *CancelError
|
|
if e.Err != nil && errors.As(e.Err, &ce) {
|
|
cancelError = ce
|
|
}
|
|
}
|
|
if assert.NotNil(t, cancelError, "expected CancelError after timeout escalation") {
|
|
assert.True(t, cancelError.Info.Escalated, "CancelError should report Escalated=true")
|
|
assert.True(t, cancelError.Info.Timeout, "CancelError should report Timeout=true")
|
|
}
|
|
}
|
|
|
|
// TestWithCancel_AfterChatModel_WithTools verifies CancelAfterChatModel fires
|
|
// when the model returns tool calls (the safe-point is on the tool-calls path).
|
|
func TestWithCancel_AfterChatModel_WithTools(t *testing.T) {
|
|
ctx := context.Background()
|
|
|
|
blk := newBlockingChatModel(toolCallMsg(toolCall("c1", "bt", `{"input":"x"}`)))
|
|
bt := newBlockingTool("bt")
|
|
|
|
agent, err := NewChatModelAgent(ctx, &ChatModelAgentConfig{
|
|
Name: "TestAgent",
|
|
Description: "test",
|
|
Model: blk,
|
|
ToolsConfig: ToolsConfig{
|
|
ToolsNodeConfig: compose.ToolsNodeConfig{Tools: []tool.BaseTool{bt}},
|
|
},
|
|
})
|
|
require.NoError(t, err)
|
|
|
|
cancelOpt, cancelFn := WithCancel()
|
|
iter := agent.Run(ctx, &AgentInput{Messages: []Message{schema.UserMessage("hi")}}, cancelOpt)
|
|
|
|
select {
|
|
case <-blk.started:
|
|
case <-time.After(5 * time.Second):
|
|
t.Fatal("model did not start")
|
|
}
|
|
|
|
cancelDone := make(chan error, 1)
|
|
go func() {
|
|
handle, _ := cancelFn(WithAgentCancelMode(CancelAfterChatModel))
|
|
cancelDone <- handle.Wait()
|
|
}()
|
|
|
|
time.Sleep(20 * time.Millisecond)
|
|
|
|
close(blk.unblockCh)
|
|
|
|
cancelErr := <-cancelDone
|
|
assert.NoError(t, cancelErr)
|
|
|
|
_, hasCancelError := drainEvents(iter)
|
|
assert.True(t, hasCancelError, "CancelError expected after model returns tool calls")
|
|
}
|
|
|
|
// TestWithCancel_CancelImmediate_StreamAborted verifies that CancelImmediate
|
|
// during model execution surfaces CancelError and completes quickly.
|
|
// Uses blockingChatModel which blocks in Stream(), keeping the agent's run
|
|
// function alive so the cancel context stays in stateRunning.
|
|
func TestWithCancel_CancelImmediate_StreamAborted(t *testing.T) {
|
|
ctx := context.Background()
|
|
|
|
blk := newBlockingChatModel(schema.AssistantMessage("hello", nil))
|
|
|
|
agent, err := NewChatModelAgent(ctx, &ChatModelAgentConfig{
|
|
Name: "TestAgent",
|
|
Description: "test",
|
|
Model: blk,
|
|
})
|
|
require.NoError(t, err)
|
|
|
|
runner := NewRunner(ctx, RunnerConfig{
|
|
Agent: agent,
|
|
EnableStreaming: true,
|
|
})
|
|
|
|
cancelOpt, cancelFn := WithCancel()
|
|
iter := runner.Run(ctx, []Message{schema.UserMessage("hi")}, cancelOpt)
|
|
|
|
select {
|
|
case <-blk.started:
|
|
case <-time.After(5 * time.Second):
|
|
t.Fatal("model did not start")
|
|
}
|
|
time.Sleep(50 * time.Millisecond)
|
|
|
|
start := time.Now()
|
|
handle, _ := cancelFn()
|
|
cancelErr := handle.Wait()
|
|
assert.NoError(t, cancelErr)
|
|
elapsed := time.Since(start)
|
|
assert.True(t, elapsed < 2*time.Second, "cancel should complete quickly, elapsed=%v", elapsed)
|
|
|
|
var foundCancelError bool
|
|
for {
|
|
e, ok := iter.Next()
|
|
if !ok {
|
|
break
|
|
}
|
|
if e.Action != nil && e.Action.Interrupted != nil {
|
|
foundCancelError = true
|
|
}
|
|
var ce *CancelError
|
|
if e.Err != nil && errors.As(e.Err, &ce) {
|
|
foundCancelError = true
|
|
}
|
|
}
|
|
assert.True(t, foundCancelError, "expected CancelError in event stream")
|
|
}
|
|
|
|
// TestWithCancel_MultipleToolsConcurrent verifies that CancelAfterToolCalls
|
|
// waits for ALL concurrent tool calls to complete before cancelling.
|
|
func TestWithCancel_MultipleToolsConcurrent(t *testing.T) {
|
|
ctx := context.Background()
|
|
|
|
bt1 := newBlockingTool("tool1")
|
|
bt2 := newBlockingTool("tool2")
|
|
|
|
// Model calls both tools in one response.
|
|
modelResp := toolCallMsg(
|
|
toolCall("c1", "tool1", `{"input":"a"}`),
|
|
toolCall("c2", "tool2", `{"input":"b"}`),
|
|
)
|
|
modelWithTools := &simpleChatModel{response: modelResp}
|
|
|
|
agent, err := NewChatModelAgent(ctx, &ChatModelAgentConfig{
|
|
Name: "TestAgent",
|
|
Description: "test",
|
|
Model: modelWithTools,
|
|
ToolsConfig: ToolsConfig{
|
|
ToolsNodeConfig: compose.ToolsNodeConfig{Tools: []tool.BaseTool{bt1, bt2}},
|
|
},
|
|
})
|
|
assert.NoError(t, err)
|
|
|
|
cancelOpt, cancelFn := WithCancel()
|
|
iter := agent.Run(ctx, &AgentInput{Messages: []Message{schema.UserMessage("go")}}, cancelOpt)
|
|
|
|
// Wait for both tools to start.
|
|
for i := 0; i < 2; i++ {
|
|
select {
|
|
case <-bt1.started:
|
|
case <-bt2.started:
|
|
case <-time.After(5 * time.Second):
|
|
t.Fatal("tools did not start")
|
|
}
|
|
}
|
|
|
|
// Request cancel after tool calls while both are still blocking.
|
|
cancelDone := make(chan error, 1)
|
|
go func() {
|
|
handle, _ := cancelFn(WithAgentCancelMode(CancelAfterToolCalls))
|
|
cancelDone <- handle.Wait()
|
|
}()
|
|
|
|
// Unblock both tools — cancel should fire only after both complete.
|
|
time.Sleep(50 * time.Millisecond)
|
|
close(bt1.unblockCh)
|
|
time.Sleep(50 * time.Millisecond)
|
|
close(bt2.unblockCh)
|
|
|
|
cancelErr := <-cancelDone
|
|
assert.NoError(t, cancelErr)
|
|
|
|
assert.Equal(t, int32(1), atomic.LoadInt32(&bt1.callCount), "tool1 should complete")
|
|
assert.Equal(t, int32(1), atomic.LoadInt32(&bt2.callCount), "tool2 should complete")
|
|
|
|
_, hasCancelError := drainEvents(iter)
|
|
assert.True(t, hasCancelError, "expected CancelError after concurrent tools completed")
|
|
}
|
|
|
|
// TestWithCancel_GraphInterruptRaceBeforeSet verifies that a CancelImmediate
|
|
// issued before setGraphInterruptFunc is called still results in cancellation.
|
|
// This exercises the retroactive-fire path in setGraphInterruptFunc.
|
|
func TestWithCancel_GraphInterruptRaceBeforeSet(t *testing.T) {
|
|
ctx := context.Background()
|
|
|
|
blk := newBlockingChatModel(schema.AssistantMessage("hi", nil))
|
|
|
|
agent, err := NewChatModelAgent(ctx, &ChatModelAgentConfig{
|
|
Name: "TestAgent",
|
|
Description: "test",
|
|
Model: blk,
|
|
})
|
|
require.NoError(t, err)
|
|
|
|
cancelOpt, cancelFn := WithCancel()
|
|
|
|
// Cancel immediately before run starts.
|
|
go func() {
|
|
handle, _ := cancelFn()
|
|
_ = handle.Wait()
|
|
}()
|
|
|
|
iter := agent.Run(ctx, &AgentInput{Messages: []Message{schema.UserMessage("hi")}}, cancelOpt)
|
|
|
|
done := make(chan struct{})
|
|
go func() {
|
|
defer close(done)
|
|
drainEvents(iter)
|
|
}()
|
|
|
|
select {
|
|
case <-done:
|
|
case <-time.After(5 * time.Second):
|
|
t.Fatal("iteration did not complete after pre-start CancelImmediate")
|
|
}
|
|
}
|
|
|
|
// TestWithCancel_NoCheckpointStore verifies cancel completes and does not panic
|
|
// when no checkpoint store is configured.
|
|
func TestWithCancel_NoCheckpointStore(t *testing.T) {
|
|
ctx := context.Background()
|
|
|
|
blk := newBlockingChatModel(schema.AssistantMessage("hi", nil))
|
|
|
|
agent, err := NewChatModelAgent(ctx, &ChatModelAgentConfig{
|
|
Name: "TestAgent",
|
|
Description: "test",
|
|
Model: blk,
|
|
})
|
|
require.NoError(t, err)
|
|
|
|
runner := NewRunner(ctx, RunnerConfig{
|
|
Agent: agent,
|
|
// No CheckPointStore set.
|
|
})
|
|
|
|
cancelOpt, cancelFn := WithCancel()
|
|
iter := runner.Run(ctx, []Message{schema.UserMessage("hi")}, cancelOpt)
|
|
|
|
select {
|
|
case <-blk.started:
|
|
case <-time.After(5 * time.Second):
|
|
t.Fatal("model did not start")
|
|
}
|
|
time.Sleep(30 * time.Millisecond)
|
|
|
|
handle, _ := cancelFn()
|
|
cancelErr := handle.Wait()
|
|
assert.NoError(t, cancelErr)
|
|
|
|
var ce *CancelError
|
|
for {
|
|
e, ok := iter.Next()
|
|
if !ok {
|
|
break
|
|
}
|
|
if e.Err != nil && errors.As(e.Err, &ce) {
|
|
break
|
|
}
|
|
}
|
|
assert.NotNil(t, ce, "expected CancelError even without checkpoint store")
|
|
}
|
|
|
|
// TestWithCancel_ModelError verifies that a model error marks the cancelCtx as
|
|
// done so that a subsequent cancelFn call returns ErrExecutionEnded.
|
|
func TestWithCancel_ModelError(t *testing.T) {
|
|
ctx := context.Background()
|
|
|
|
modelErr := errors.New("model failed")
|
|
agent, err := NewChatModelAgent(ctx, &ChatModelAgentConfig{
|
|
Name: "TestAgent",
|
|
Description: "test",
|
|
Model: &errorChatModel{err: modelErr},
|
|
})
|
|
require.NoError(t, err)
|
|
|
|
cancelOpt, cancelFn := WithCancel()
|
|
iter := agent.Run(ctx, &AgentInput{Messages: []Message{schema.UserMessage("hi")}}, cancelOpt)
|
|
|
|
var gotModelErr bool
|
|
for {
|
|
e, ok := iter.Next()
|
|
if !ok {
|
|
break
|
|
}
|
|
if e.Err != nil && !errors.As(e.Err, new(*CancelError)) {
|
|
gotModelErr = true
|
|
}
|
|
}
|
|
assert.True(t, gotModelErr, "expected non-cancel error event from model failure")
|
|
|
|
handle, _ := cancelFn()
|
|
cancelErr := handle.Wait()
|
|
assert.ErrorIs(t, cancelErr, ErrExecutionEnded, "cancelFn should return ErrExecutionEnded after model error")
|
|
}
|
|
|
|
// TestWithCancel_Resume_SafePoint covers CancelAfterChatModel and
|
|
// CancelAfterToolCalls on a Resume path.
|
|
func TestWithCancel_Resume_SafePoint(t *testing.T) {
|
|
ctx := context.Background()
|
|
|
|
// --- phase 1: run to get a checkpoint via CancelImmediate ---
|
|
blk := newBlockingChatModel(toolCallMsg(toolCall("c1", "bt", `{"input":"x"}`)))
|
|
bt := newSlowTool("bt", 50*time.Millisecond, "result")
|
|
|
|
agent1, err := NewChatModelAgent(ctx, &ChatModelAgentConfig{
|
|
Name: "TestAgent",
|
|
Description: "test",
|
|
Model: blk,
|
|
ToolsConfig: ToolsConfig{
|
|
ToolsNodeConfig: compose.ToolsNodeConfig{Tools: []tool.BaseTool{bt}},
|
|
},
|
|
})
|
|
assert.NoError(t, err)
|
|
|
|
store := newCancelTestStore()
|
|
runner1 := NewRunner(ctx, RunnerConfig{
|
|
Agent: agent1,
|
|
CheckPointStore: store,
|
|
})
|
|
|
|
cancelOpt1, cancelFn1 := WithCancel()
|
|
iter1 := runner1.Run(ctx, []Message{schema.UserMessage("hi")}, cancelOpt1, WithCheckPointID("resume-sp-1"))
|
|
|
|
select {
|
|
case <-blk.started:
|
|
case <-time.After(5 * time.Second):
|
|
t.Fatal("model did not start in phase 1")
|
|
}
|
|
_, _ = cancelFn1()
|
|
drainEvents(iter1)
|
|
|
|
// --- phase 2: resume, cancel after chat model ---
|
|
resumeModel := newBlockingChatModel(toolCallMsg(toolCall("c1", "bt", `{"input":"x"}`)))
|
|
|
|
bt2 := newSlowTool("bt", 50*time.Millisecond, "result")
|
|
agent2, err := NewChatModelAgent(ctx, &ChatModelAgentConfig{
|
|
Name: "TestAgent",
|
|
Description: "test",
|
|
Model: resumeModel,
|
|
ToolsConfig: ToolsConfig{
|
|
ToolsNodeConfig: compose.ToolsNodeConfig{Tools: []tool.BaseTool{bt2}},
|
|
},
|
|
})
|
|
assert.NoError(t, err)
|
|
|
|
runner2 := NewRunner(ctx, RunnerConfig{
|
|
Agent: agent2,
|
|
CheckPointStore: store,
|
|
})
|
|
|
|
cancelOpt2, cancelFn2 := WithCancel()
|
|
resumeIter, err := runner2.Resume(ctx, "resume-sp-1", cancelOpt2)
|
|
require.NoError(t, err)
|
|
|
|
select {
|
|
case <-resumeModel.started:
|
|
case <-time.After(5 * time.Second):
|
|
t.Fatal("model did not start in phase 2")
|
|
}
|
|
|
|
cancelDone := make(chan error, 1)
|
|
go func() {
|
|
handle, _ := cancelFn2(WithAgentCancelMode(CancelAfterChatModel))
|
|
cancelDone <- handle.Wait()
|
|
}()
|
|
|
|
time.Sleep(50 * time.Millisecond)
|
|
|
|
close(resumeModel.unblockCh)
|
|
|
|
cancelErr := <-cancelDone
|
|
assert.NoError(t, cancelErr)
|
|
|
|
_, hasCancelError := drainEvents(resumeIter)
|
|
assert.True(t, hasCancelError, "CancelError expected after resumed model returns tool calls")
|
|
}
|
|
|
|
// callbackTool is a tool that calls onCall when invoked.
|
|
type callbackTool struct {
|
|
name string
|
|
onCall func()
|
|
}
|
|
|
|
func (t *callbackTool) Info(_ context.Context) (*schema.ToolInfo, error) {
|
|
return &schema.ToolInfo{
|
|
Name: t.name,
|
|
Desc: "callback tool",
|
|
ParamsOneOf: schema.NewParamsOneOfByParams(map[string]*schema.ParameterInfo{
|
|
"input": {Type: "string"},
|
|
}),
|
|
}, nil
|
|
}
|
|
|
|
func (t *callbackTool) InvokableRun(_ context.Context, _ string, _ ...tool.Option) (string, error) {
|
|
if t.onCall != nil {
|
|
t.onCall()
|
|
}
|
|
return "ok", nil
|
|
}
|
|
|
|
// interruptingChatModel returns a compose.Interrupt error to simulate a
|
|
// business interrupt during execution.
|
|
type interruptingChatModel struct{}
|
|
|
|
func (m *interruptingChatModel) Generate(ctx context.Context, _ []*schema.Message, _ ...model.Option) (*schema.Message, error) {
|
|
return nil, compose.Interrupt(ctx, "test interrupt")
|
|
}
|
|
|
|
func (m *interruptingChatModel) Stream(ctx context.Context, _ []*schema.Message, _ ...model.Option) (*schema.StreamReader[*schema.Message], error) {
|
|
return nil, compose.Interrupt(ctx, "test interrupt")
|
|
}
|
|
|
|
func (m *interruptingChatModel) BindTools(_ []*schema.ToolInfo) error { return nil }
|
|
|
|
// TestWithCancel_TargetedResume_CancelImmediate cancels an agent via CancelImmediate,
|
|
// extracts InterruptContexts from the resulting CancelError, and uses them
|
|
// for targeted resumption via Runner.ResumeWithParams.
|
|
func TestWithCancel_TargetedResume_CancelImmediate(t *testing.T) {
|
|
ctx := context.Background()
|
|
|
|
blk := newBlockingChatModel(toolCallMsg(toolCall("c1", "st", `{"input":"x"}`)))
|
|
st := newSlowTool("st", 50*time.Millisecond, "result")
|
|
|
|
agent, err := NewChatModelAgent(ctx, &ChatModelAgentConfig{
|
|
Name: "TestAgent",
|
|
Description: "test",
|
|
Model: blk,
|
|
ToolsConfig: ToolsConfig{
|
|
ToolsNodeConfig: compose.ToolsNodeConfig{Tools: []tool.BaseTool{st}},
|
|
},
|
|
})
|
|
require.NoError(t, err)
|
|
|
|
store := newCancelTestStore()
|
|
runner := NewRunner(ctx, RunnerConfig{
|
|
Agent: agent,
|
|
CheckPointStore: store,
|
|
})
|
|
|
|
cancelOpt, cancelFn := WithCancel()
|
|
iter := runner.Run(ctx, []Message{schema.UserMessage("go")}, cancelOpt, WithCheckPointID("targeted-imm-1"))
|
|
|
|
select {
|
|
case <-blk.started:
|
|
case <-time.After(5 * time.Second):
|
|
t.Fatal("model did not start")
|
|
}
|
|
|
|
handle, _ := cancelFn() // CancelImmediate (default)
|
|
cancelErr := handle.Wait()
|
|
assert.NoError(t, cancelErr)
|
|
|
|
var cancelError *CancelError
|
|
for {
|
|
e, ok := iter.Next()
|
|
if !ok {
|
|
break
|
|
}
|
|
var ce *CancelError
|
|
if e.Err != nil && errors.As(e.Err, &ce) {
|
|
cancelError = ce
|
|
}
|
|
}
|
|
|
|
require.NotNil(t, cancelError, "expected CancelError")
|
|
require.NotEmpty(t, cancelError.InterruptContexts, "CancelError should have InterruptContexts for targeted resume")
|
|
|
|
// --- resume with targeted params ---
|
|
targets := make(map[string]any)
|
|
for _, ic := range cancelError.InterruptContexts {
|
|
targets[ic.ID] = nil
|
|
}
|
|
|
|
resumeModel := &plainResponseModel{text: "resumed"}
|
|
agent2, err := NewChatModelAgent(ctx, &ChatModelAgentConfig{
|
|
Name: "TestAgent",
|
|
Description: "test",
|
|
Model: resumeModel,
|
|
ToolsConfig: ToolsConfig{
|
|
ToolsNodeConfig: compose.ToolsNodeConfig{Tools: []tool.BaseTool{st}},
|
|
},
|
|
})
|
|
require.NoError(t, err)
|
|
|
|
runner2 := NewRunner(ctx, RunnerConfig{
|
|
Agent: agent2,
|
|
CheckPointStore: store,
|
|
})
|
|
|
|
resumeIter, err := runner2.ResumeWithParams(ctx, "targeted-imm-1", &ResumeParams{Targets: targets})
|
|
require.NoError(t, err)
|
|
|
|
var gotOutput bool
|
|
for {
|
|
e, ok := resumeIter.Next()
|
|
if !ok {
|
|
break
|
|
}
|
|
if e.Err != nil {
|
|
t.Fatalf("unexpected error during targeted resume: %v", e.Err)
|
|
}
|
|
if e.Output != nil && e.Output.MessageOutput != nil {
|
|
gotOutput = true
|
|
}
|
|
}
|
|
assert.True(t, gotOutput, "targeted resume should produce output")
|
|
}
|
|
|
|
// TestWithCancel_TargetedResume_SafePoint cancels an agent via CancelAfterChatModel
|
|
// (safe-point) and verifies that InterruptContexts are populated on the CancelError
|
|
// and that targeted resume via ResumeWithParams succeeds.
|
|
// Since safe-point cancels now use compose.Interrupt, compose saves checkpoint data,
|
|
// making the cancel fully resumable.
|
|
func TestWithCancel_TargetedResume_SafePoint(t *testing.T) {
|
|
ctx := context.Background()
|
|
|
|
// The model returns a tool call so the react graph routes to toolPreHandle,
|
|
// which detects CancelAfterChatModel and fires compose.Interrupt.
|
|
blk := newBlockingChatModel(toolCallMsg(toolCall("c1", "st", `{"input":"x"}`)))
|
|
st := newSlowTool("st", 0, "result")
|
|
|
|
agent, err := NewChatModelAgent(ctx, &ChatModelAgentConfig{
|
|
Name: "TestAgent",
|
|
Description: "test",
|
|
Model: blk,
|
|
ToolsConfig: ToolsConfig{
|
|
ToolsNodeConfig: compose.ToolsNodeConfig{Tools: []tool.BaseTool{st}},
|
|
},
|
|
})
|
|
require.NoError(t, err)
|
|
|
|
store := newCancelTestStore()
|
|
runner := NewRunner(ctx, RunnerConfig{
|
|
Agent: agent,
|
|
CheckPointStore: store,
|
|
})
|
|
|
|
cancelOpt, cancelFn := WithCancel()
|
|
iter := runner.Run(ctx, []Message{schema.UserMessage("go")}, cancelOpt, WithCheckPointID("targeted-sp-1"))
|
|
|
|
select {
|
|
case <-blk.started:
|
|
case <-time.After(5 * time.Second):
|
|
t.Fatal("model did not start")
|
|
}
|
|
|
|
// Start cancelFn in background so the CAS happens before the model unblocks.
|
|
cancelDone := make(chan error, 1)
|
|
go func() {
|
|
handle, _ := cancelFn(WithAgentCancelMode(CancelAfterChatModel))
|
|
cancelDone <- handle.Wait()
|
|
}()
|
|
time.Sleep(50 * time.Millisecond)
|
|
close(blk.unblockCh)
|
|
|
|
cancelErr := <-cancelDone
|
|
assert.NoError(t, cancelErr)
|
|
|
|
var cancelError *CancelError
|
|
for {
|
|
e, ok := iter.Next()
|
|
if !ok {
|
|
break
|
|
}
|
|
var ce *CancelError
|
|
if e.Err != nil && errors.As(e.Err, &ce) {
|
|
cancelError = ce
|
|
}
|
|
}
|
|
|
|
require.NotNil(t, cancelError, "expected CancelError")
|
|
require.NotEmpty(t, cancelError.InterruptContexts, "CancelError should have InterruptContexts for targeted resume")
|
|
|
|
// --- resume with targeted params ---
|
|
targets := make(map[string]any)
|
|
for _, ic := range cancelError.InterruptContexts {
|
|
targets[ic.ID] = nil
|
|
}
|
|
|
|
resumeModel := &plainResponseModel{text: "resumed"}
|
|
agent2, err := NewChatModelAgent(ctx, &ChatModelAgentConfig{
|
|
Name: "TestAgent",
|
|
Description: "test",
|
|
Model: resumeModel,
|
|
ToolsConfig: ToolsConfig{
|
|
ToolsNodeConfig: compose.ToolsNodeConfig{Tools: []tool.BaseTool{st}},
|
|
},
|
|
})
|
|
require.NoError(t, err)
|
|
|
|
runner2 := NewRunner(ctx, RunnerConfig{
|
|
Agent: agent2,
|
|
CheckPointStore: store,
|
|
})
|
|
|
|
resumeIter, err := runner2.ResumeWithParams(ctx, "targeted-sp-1", &ResumeParams{Targets: targets})
|
|
require.NoError(t, err)
|
|
|
|
var gotOutput bool
|
|
for {
|
|
e, ok := resumeIter.Next()
|
|
if !ok {
|
|
break
|
|
}
|
|
if e.Err != nil {
|
|
t.Fatalf("unexpected error during targeted resume: %v", e.Err)
|
|
}
|
|
if e.Output != nil && e.Output.MessageOutput != nil {
|
|
gotOutput = true
|
|
}
|
|
}
|
|
assert.True(t, gotOutput, "targeted resume should produce output")
|
|
}
|
|
|
|
// TestWithCancel_Resume_CancelAfterChatModel_MessagePreserved tests both the
|
|
// ReAct (with-tools) and noTools paths to ensure that when a
|
|
// CancelAfterChatModel safe-point fires and the run is later resumed, the
|
|
// original Message returned by the chat model is preserved through the
|
|
// StatefulInterrupt checkpoint.
|
|
//
|
|
// For the ReAct path: the model returns a tool-call message. On resume the
|
|
// cancelCheck node must return that same message so the branch routes to the
|
|
// ToolNode and the tool actually executes.
|
|
//
|
|
// For the noTools path: the model returns a plain text message. On resume the
|
|
// cancel-check lambda must return that same message as the chain output.
|
|
func TestWithCancel_Resume_CancelAfterChatModel_MessagePreserved(t *testing.T) {
|
|
t.Run("react_path_tool_call_preserved", func(t *testing.T) {
|
|
ctx := context.Background()
|
|
|
|
// Phase-2 model returns no tool calls so the graph ends.
|
|
// We track whether the tool actually executes on resume.
|
|
toolExecuted := make(chan struct{}, 1)
|
|
st := &callbackTool{
|
|
name: "my_tool",
|
|
onCall: func() {
|
|
select {
|
|
case toolExecuted <- struct{}{}:
|
|
default:
|
|
}
|
|
},
|
|
}
|
|
|
|
// Phase-1 model returns a tool call.
|
|
blk := newBlockingChatModel(toolCallMsg(toolCall("c1", "my_tool", `{"input":"x"}`)))
|
|
|
|
agent1, err := NewChatModelAgent(ctx, &ChatModelAgentConfig{
|
|
Name: "TestAgent",
|
|
Description: "test",
|
|
Model: blk,
|
|
ToolsConfig: ToolsConfig{
|
|
ToolsNodeConfig: compose.ToolsNodeConfig{Tools: []tool.BaseTool{st}},
|
|
},
|
|
})
|
|
require.NoError(t, err)
|
|
|
|
store := newCancelTestStore()
|
|
runner1 := NewRunner(ctx, RunnerConfig{
|
|
Agent: agent1,
|
|
CheckPointStore: store,
|
|
})
|
|
|
|
cancelOpt1, cancelFn1 := WithCancel()
|
|
iter1 := runner1.Run(ctx, []Message{schema.UserMessage("hi")},
|
|
cancelOpt1, WithCheckPointID("react-msg-preserved-1"))
|
|
|
|
select {
|
|
case <-blk.started:
|
|
case <-time.After(5 * time.Second):
|
|
t.Fatal("model did not start in phase 1")
|
|
}
|
|
|
|
cancelDone := make(chan error, 1)
|
|
go func() {
|
|
handle, _ := cancelFn1(WithAgentCancelMode(CancelAfterChatModel))
|
|
cancelDone <- handle.Wait()
|
|
}()
|
|
time.Sleep(50 * time.Millisecond)
|
|
close(blk.unblockCh)
|
|
|
|
cancelErr := <-cancelDone
|
|
assert.NoError(t, cancelErr)
|
|
|
|
_, hasCancelError := drainEvents(iter1)
|
|
assert.True(t, hasCancelError, "expected CancelError from phase 1")
|
|
|
|
// Phase 2: resume. The model for phase-2 returns plain text (no tool
|
|
// calls) so the react graph ends after one iteration. But first the
|
|
// tool from the checkpoint must execute.
|
|
resumeModel := &plainResponseModel{text: "done"}
|
|
agent2, err := NewChatModelAgent(ctx, &ChatModelAgentConfig{
|
|
Name: "TestAgent",
|
|
Description: "test",
|
|
Model: resumeModel,
|
|
ToolsConfig: ToolsConfig{
|
|
ToolsNodeConfig: compose.ToolsNodeConfig{Tools: []tool.BaseTool{st}},
|
|
},
|
|
})
|
|
require.NoError(t, err)
|
|
|
|
runner2 := NewRunner(ctx, RunnerConfig{
|
|
Agent: agent2,
|
|
CheckPointStore: store,
|
|
})
|
|
|
|
resumeIter, err := runner2.Resume(ctx, "react-msg-preserved-1")
|
|
require.NoError(t, err)
|
|
|
|
for {
|
|
e, ok := resumeIter.Next()
|
|
if !ok {
|
|
break
|
|
}
|
|
if e.Err != nil {
|
|
t.Fatalf("unexpected error during resume: %v", e.Err)
|
|
}
|
|
}
|
|
|
|
// The key assertion: the tool must have been called during resume,
|
|
// which can only happen if the tool-call message was preserved.
|
|
select {
|
|
case <-toolExecuted:
|
|
// success
|
|
default:
|
|
t.Fatal("tool was not executed on resume — the tool-call message was lost")
|
|
}
|
|
})
|
|
|
|
}
|
|
|
|
// TestHandleRunFuncError_AlreadyHandled_NoDuplicate verifies that when
|
|
// markCancelHandled() was already claimed by a sub-agent's handleRunFuncError,
|
|
// the sequential workflow's checkCancel does not emit a second CancelError.
|
|
//
|
|
// Setup: sequential[cma1, cma2] with CancelAfterToolCalls. cma1 has tools,
|
|
// cancel fires while tool is running. After tool completes, the safe-point
|
|
// fires in cma1's handleRunFuncError (claiming markCancelHandled). The
|
|
// sequential workflow's checkCancel at the transition point should find
|
|
// markCancelHandled returns false and skip — producing exactly 1 CancelError.
|
|
func TestHandleRunFuncError_AlreadyHandled_NoDuplicate(t *testing.T) {
|
|
ctx := context.Background()
|
|
|
|
bt := newBlockingTool("bt")
|
|
|
|
// cma1: model returns a tool call immediately, tool blocks until unblocked
|
|
cma1Model := newBlockingChatModel(toolCallMsg(toolCall("c1", "bt", `{"input":"x"}`)))
|
|
close(cma1Model.unblockCh) // model returns immediately
|
|
|
|
agent1, err := NewChatModelAgent(ctx, &ChatModelAgentConfig{
|
|
Name: "agent1", Description: "first", Instruction: "test",
|
|
Model: cma1Model,
|
|
ToolsConfig: ToolsConfig{
|
|
ToolsNodeConfig: compose.ToolsNodeConfig{Tools: []tool.BaseTool{bt}},
|
|
},
|
|
})
|
|
require.NoError(t, err)
|
|
|
|
agent2Model := &plainResponseModel{text: "agent2-response"}
|
|
agent2, err := NewChatModelAgent(ctx, &ChatModelAgentConfig{
|
|
Name: "agent2", Description: "second", Instruction: "test",
|
|
Model: agent2Model,
|
|
})
|
|
require.NoError(t, err)
|
|
|
|
seqAgent, err := NewSequentialAgent(ctx, &SequentialAgentConfig{
|
|
Name: "seq", Description: "sequential", SubAgents: []Agent{agent1, agent2},
|
|
})
|
|
require.NoError(t, err)
|
|
|
|
runner := NewRunner(ctx, RunnerConfig{
|
|
Agent: seqAgent, EnableStreaming: false,
|
|
})
|
|
|
|
cancelOpt, cancelFn := WithCancel()
|
|
iter := runner.Run(ctx, []Message{schema.UserMessage("test")}, cancelOpt)
|
|
|
|
// Wait for tool to start
|
|
select {
|
|
case <-bt.started:
|
|
case <-time.After(5 * time.Second):
|
|
t.Fatal("Tool did not start")
|
|
}
|
|
|
|
// Cancel while tool is still running (in goroutine because cancelFn blocks
|
|
// until execution finishes), then unblock tool so safe-point fires
|
|
go func() {
|
|
handle, _ := cancelFn(WithAgentCancelMode(CancelAfterToolCalls))
|
|
_ = handle.Wait()
|
|
}()
|
|
|
|
// Give cancel time to register, then unblock tool
|
|
time.Sleep(50 * time.Millisecond)
|
|
close(bt.unblockCh)
|
|
|
|
cancelCount := 0
|
|
for {
|
|
event, ok := iter.Next()
|
|
if !ok {
|
|
break
|
|
}
|
|
var ce *CancelError
|
|
if event.Err != nil && errors.As(event.Err, &ce) {
|
|
cancelCount++
|
|
}
|
|
}
|
|
|
|
assert.Equal(t, 1, cancelCount, "Should have exactly one CancelError, no duplicate from handleRunFuncError + checkCancel")
|
|
}
|
|
|
|
func TestWithCancel_CancelAfterChatModel_NestedAgentTool(t *testing.T) {
|
|
ctx := context.Background()
|
|
|
|
subAgentModel := newBlockingChatModel(toolCallMsg(toolCall("c1", "sub_tool", `{"input":"x"}`)))
|
|
subAgentModelStarted := subAgentModel.started
|
|
subTool := newBlockingTool("sub_tool")
|
|
|
|
subAgent, err := NewChatModelAgent(ctx, &ChatModelAgentConfig{
|
|
Name: "sub_agent",
|
|
Description: "test sub agent",
|
|
Instruction: "you are a sub agent",
|
|
Model: subAgentModel,
|
|
ToolsConfig: ToolsConfig{
|
|
ToolsNodeConfig: compose.ToolsNodeConfig{Tools: []tool.BaseTool{subTool}},
|
|
},
|
|
})
|
|
require.NoError(t, err)
|
|
|
|
supervisorModel := &simpleChatModel{
|
|
response: &schema.Message{
|
|
Role: schema.Assistant,
|
|
ToolCalls: []schema.ToolCall{{
|
|
ID: "call_1", Type: "function",
|
|
Function: schema.FunctionCall{
|
|
Name: TransferToAgentToolName,
|
|
Arguments: `{"agent_name": "sub_agent"}`,
|
|
},
|
|
}},
|
|
},
|
|
}
|
|
|
|
supervisorAgent, err := NewChatModelAgent(ctx, &ChatModelAgentConfig{
|
|
Name: "supervisor",
|
|
Description: "supervisor agent (equivalent to DeepAgent)",
|
|
Instruction: "you are a supervisor",
|
|
Model: supervisorModel,
|
|
})
|
|
require.NoError(t, err)
|
|
|
|
agentWithSubAgents, err := SetSubAgents(ctx, supervisorAgent, []Agent{subAgent})
|
|
require.NoError(t, err)
|
|
|
|
runner := NewRunner(ctx, RunnerConfig{
|
|
Agent: agentWithSubAgents,
|
|
EnableStreaming: false,
|
|
})
|
|
|
|
cancelOpt, cancelFn := WithCancel()
|
|
iter := runner.Run(ctx, []Message{schema.UserMessage("test")}, cancelOpt)
|
|
|
|
select {
|
|
case <-subAgentModelStarted:
|
|
case <-time.After(10 * time.Second):
|
|
t.Fatal("Sub-agent model did not start")
|
|
}
|
|
|
|
time.Sleep(50 * time.Millisecond)
|
|
|
|
cancelDone := make(chan error, 1)
|
|
go func() {
|
|
handle, _ := cancelFn(WithAgentCancelMode(CancelAfterChatModel), WithRecursive())
|
|
cancelDone <- handle.Wait()
|
|
}()
|
|
|
|
time.Sleep(100 * time.Millisecond)
|
|
close(subAgentModel.unblockCh)
|
|
|
|
cancelErr := <-cancelDone
|
|
assert.NoError(t, cancelErr)
|
|
|
|
hasCancelError := false
|
|
for {
|
|
event, ok := iter.Next()
|
|
if !ok {
|
|
break
|
|
}
|
|
var ce *CancelError
|
|
if event.Err != nil && errors.As(event.Err, &ce) {
|
|
hasCancelError = true
|
|
}
|
|
}
|
|
|
|
assert.True(t, hasCancelError, "CancelError expected from nested agent tool with tools")
|
|
}
|
|
|
|
// slowStreamingTool implements StreamableTool (but NOT InvokableTool), streaming
|
|
// chunks slowly so CancelImmediate can fire mid-stream.
|
|
type slowStreamingTool struct {
|
|
name string
|
|
chunkInterval time.Duration
|
|
chunks []string
|
|
started chan struct{}
|
|
gate chan struct{} // if non-nil, blocks after first chunk until closed
|
|
}
|
|
|
|
func (t *slowStreamingTool) Info(_ context.Context) (*schema.ToolInfo, error) {
|
|
return &schema.ToolInfo{
|
|
Name: t.name,
|
|
Desc: "slow streaming tool",
|
|
ParamsOneOf: schema.NewParamsOneOfByParams(map[string]*schema.ParameterInfo{
|
|
"input": {Type: "string"},
|
|
}),
|
|
}, nil
|
|
}
|
|
|
|
func (t *slowStreamingTool) StreamableRun(_ context.Context, _ string, _ ...tool.Option) (*schema.StreamReader[string], error) {
|
|
r, w := schema.Pipe[string](1)
|
|
go func() {
|
|
defer w.Close()
|
|
select {
|
|
case t.started <- struct{}{}:
|
|
default:
|
|
}
|
|
for i, chunk := range t.chunks {
|
|
time.Sleep(t.chunkInterval)
|
|
if closed := w.Send(chunk, nil); closed {
|
|
return
|
|
}
|
|
// After the second chunk, block on gate so the caller can
|
|
// issue a cancel while the tool is deterministically still streaming.
|
|
// We wait until chunk index 1 (second chunk) so that the framework
|
|
// has time to receive the first chunk and forward the streaming
|
|
// event to the iterator, ensuring ErrStreamCanceled is observable.
|
|
if i == 1 && t.gate != nil {
|
|
<-t.gate
|
|
}
|
|
}
|
|
}()
|
|
return r, nil
|
|
}
|
|
|
|
// toolCallStreamModel returns a tool-call message on the first Stream call,
|
|
// then a plain text response on subsequent calls.
|
|
type toolCallStreamModel struct {
|
|
callCount int32
|
|
}
|
|
|
|
func (m *toolCallStreamModel) Generate(_ context.Context, _ []*schema.Message, _ ...model.Option) (*schema.Message, error) {
|
|
if atomic.AddInt32(&m.callCount, 1) == 1 {
|
|
return toolCallMsg(toolCall("c1", "slow_tool", `{"input":"x"}`)), nil
|
|
}
|
|
return schema.AssistantMessage("done", nil), nil
|
|
}
|
|
|
|
func (m *toolCallStreamModel) Stream(ctx context.Context, input []*schema.Message, opts ...model.Option) (*schema.StreamReader[*schema.Message], error) {
|
|
msg, err := m.Generate(ctx, input, opts...)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
return schema.StreamReaderFromArray([]*schema.Message{msg}), nil
|
|
}
|
|
|
|
func (m *toolCallStreamModel) BindTools(_ []*schema.ToolInfo) error { return nil }
|
|
|
|
// TestWithCancel_CancelImmediate_StreamableToolAborted verifies that CancelImmediate
|
|
// during StreamableTool streaming surfaces ErrStreamCanceled on the tool's
|
|
// MessageStream.Recv(), just like it does for ChatModel streaming.
|
|
func TestWithCancel_CancelImmediate_StreamableToolAborted(t *testing.T) {
|
|
ctx := context.Background()
|
|
|
|
tcm := &toolCallStreamModel{}
|
|
gate := make(chan struct{})
|
|
st := &slowStreamingTool{
|
|
name: "slow_tool",
|
|
chunkInterval: 100 * time.Millisecond,
|
|
chunks: []string{"a", "b", "c", "d", "e", "f", "g", "h", "i", "j"},
|
|
started: make(chan struct{}, 1),
|
|
gate: gate,
|
|
}
|
|
|
|
agent, err := NewChatModelAgent(ctx, &ChatModelAgentConfig{
|
|
Name: "TestAgent",
|
|
Description: "test",
|
|
Model: tcm,
|
|
ToolsConfig: ToolsConfig{
|
|
ToolsNodeConfig: compose.ToolsNodeConfig{Tools: []tool.BaseTool{st}},
|
|
},
|
|
})
|
|
require.NoError(t, err)
|
|
|
|
runner := NewRunner(ctx, RunnerConfig{
|
|
Agent: agent,
|
|
EnableStreaming: true,
|
|
})
|
|
|
|
cancelOpt, cancelFn := WithCancel()
|
|
iter := runner.Run(ctx, []Message{schema.UserMessage("hi")}, cancelOpt)
|
|
|
|
// Wait for the tool to start streaming and send its first chunk.
|
|
// The tool then blocks on the gate, guaranteeing the execution is
|
|
// still in progress when we issue the cancel.
|
|
select {
|
|
case <-st.started:
|
|
case <-time.After(5 * time.Second):
|
|
t.Fatal("tool did not start streaming")
|
|
}
|
|
|
|
// Drain events in a separate goroutine so we can issue the cancel
|
|
// from the main goroutine after confirming the tool stream event
|
|
// has been received.
|
|
type result struct {
|
|
foundStreamCanceled bool
|
|
foundCancelError bool
|
|
}
|
|
resultCh := make(chan result, 1)
|
|
toolStreamReady := make(chan struct{})
|
|
go func() {
|
|
var r result
|
|
for {
|
|
e, ok := iter.Next()
|
|
if !ok {
|
|
break
|
|
}
|
|
|
|
// ErrStreamCanceled appears on the tool's MessageStream.Recv()
|
|
if e.Output != nil && e.Output.MessageOutput != nil && e.Output.MessageOutput.IsStreaming &&
|
|
e.Output.MessageOutput.Role == schema.Tool {
|
|
// Signal that the tool stream event has been received.
|
|
close(toolStreamReady)
|
|
stream := e.Output.MessageOutput.MessageStream
|
|
for {
|
|
_, recvErr := stream.Recv()
|
|
if recvErr != nil {
|
|
if errors.Is(recvErr, ErrStreamCanceled) {
|
|
r.foundStreamCanceled = true
|
|
}
|
|
break
|
|
}
|
|
}
|
|
}
|
|
|
|
if e.Action != nil && e.Action.Interrupted != nil {
|
|
r.foundCancelError = true
|
|
}
|
|
var ce *CancelError
|
|
if e.Err != nil && errors.As(e.Err, &ce) {
|
|
r.foundCancelError = true
|
|
}
|
|
}
|
|
resultCh <- r
|
|
}()
|
|
|
|
// Wait for the iterator goroutine to receive the tool streaming event.
|
|
// At this point the tool goroutine is blocked on the gate, and the
|
|
// iterator goroutine is blocked on stream.Recv(), so the execution is
|
|
// guaranteed to still be in progress.
|
|
select {
|
|
case <-toolStreamReady:
|
|
case <-time.After(5 * time.Second):
|
|
t.Fatal("tool stream event was not received by the iterator")
|
|
}
|
|
|
|
// Issue cancel while the tool goroutine is blocked on gate.
|
|
// wrapStreamWithCancelMonitoring detects immediateChan and sends
|
|
// ErrStreamCanceled to the consumer side. We do NOT close gate here —
|
|
// keeping the tool goroutine blocked ensures the graph interrupt (timeout=0)
|
|
// wins the race against normal completion. Close gate in defer for cleanup.
|
|
defer close(gate)
|
|
handle, _ := cancelFn()
|
|
cancelErr := handle.Wait()
|
|
|
|
r := <-resultCh
|
|
|
|
if errors.Is(cancelErr, ErrExecutionEnded) {
|
|
// On slower runtimes (e.g. Go 1.19 CI), the execution can complete
|
|
// before the cancel signal is delivered — this is a valid race outcome.
|
|
t.Log("cancel raced with completion (ErrExecutionEnded) — skipping cancel assertions")
|
|
return
|
|
}
|
|
assert.NoError(t, cancelErr)
|
|
assert.True(t, r.foundStreamCanceled, "expected ErrStreamCanceled on tool's MessageStream.Recv()")
|
|
assert.True(t, r.foundCancelError, "expected CancelError in event stream")
|
|
}
|
|
|
|
// TestWithCancel_CancelImmediate_NestedAgentTool_ResumeFromToolsNode verifies that
|
|
// when a nested ChatModelAgent (wrapped as an AgentTool inside an outer ChatModelAgent)
|
|
// is canceled via CancelImmediate and then resumed with Runner.Resume (no params),
|
|
// the outer agent resumes from the ToolsNode rather than restarting from the beginning.
|
|
//
|
|
// Regression test: previously, the outer ChatModelAgent would restart from its Init/ChatModel
|
|
// node instead of resuming from the ToolsNode, causing the outer model to be called again
|
|
// with the original user message before the AgentTool and inner ChatModelAgent were resumed.
|
|
func TestWithCancel_CancelImmediate_NestedAgentTool_ResumeFromToolsNode(t *testing.T) {
|
|
for _, tc := range []struct {
|
|
name string
|
|
enableStreaming bool
|
|
innerHasTools bool
|
|
recursive bool
|
|
}{
|
|
{"Invoke_InnerNoTools_NonRecursive", false, false, false},
|
|
{"Stream_InnerNoTools_NonRecursive", true, false, false},
|
|
{"Invoke_InnerWithTools_NonRecursive", false, true, false},
|
|
{"Stream_InnerWithTools_NonRecursive", true, true, false},
|
|
{"Invoke_InnerNoTools_Recursive", false, false, true},
|
|
{"Stream_InnerNoTools_Recursive", true, false, true},
|
|
{"Invoke_InnerWithTools_Recursive", false, true, true},
|
|
{"Stream_InnerWithTools_Recursive", true, true, true},
|
|
} {
|
|
t.Run(tc.name, func(t *testing.T) {
|
|
ctx := context.Background()
|
|
|
|
// --- inner agent: its model blocks so we can cancel mid-execution ---
|
|
var innerTools []tool.BaseTool
|
|
var innerModelResp *schema.Message
|
|
if tc.innerHasTools {
|
|
innerModelResp = toolCallMsg(toolCall("ic1", "inner_tool", `{"input":"x"}`))
|
|
innerTools = []tool.BaseTool{newBlockingTool("inner_tool")}
|
|
} else {
|
|
innerModelResp = &schema.Message{Role: schema.Assistant, Content: "inner agent done"}
|
|
}
|
|
innerModel := newBlockingChatModel(innerModelResp)
|
|
t.Cleanup(func() {
|
|
close(innerModel.unblockCh)
|
|
})
|
|
|
|
innerCfg := &ChatModelAgentConfig{
|
|
Name: "InnerAgent",
|
|
Description: "inner agent that blocks",
|
|
Instruction: "you are an inner agent",
|
|
Model: innerModel,
|
|
}
|
|
if len(innerTools) > 0 {
|
|
innerCfg.ToolsConfig = ToolsConfig{
|
|
ToolsNodeConfig: compose.ToolsNodeConfig{Tools: innerTools},
|
|
}
|
|
}
|
|
innerAgent, err := NewChatModelAgent(ctx, innerCfg)
|
|
require.NoError(t, err)
|
|
|
|
// --- outer agent: counting model ---
|
|
// Call 1: returns a tool call that invokes InnerAgent.
|
|
// Call 2 (only needed on resume): returns a plain final answer.
|
|
outerModelCallCount := int32(0)
|
|
outerModel := &countingChatModel{
|
|
callCount: &outerModelCallCount,
|
|
responses: []*schema.Message{
|
|
toolCallMsg(toolCall("c1", "InnerAgent", `{"request":"do something"}`)),
|
|
schema.AssistantMessage("outer completed", nil),
|
|
},
|
|
}
|
|
|
|
outerAgent, err := NewChatModelAgent(ctx, &ChatModelAgentConfig{
|
|
Name: "OuterAgent",
|
|
Description: "outer agent with nested agent tool",
|
|
Instruction: "you are an outer agent",
|
|
Model: outerModel,
|
|
ToolsConfig: ToolsConfig{
|
|
ToolsNodeConfig: compose.ToolsNodeConfig{
|
|
Tools: []tool.BaseTool{NewAgentTool(ctx, innerAgent)},
|
|
},
|
|
},
|
|
})
|
|
require.NoError(t, err)
|
|
|
|
store := newCancelTestStore()
|
|
checkpointID := "cancel-nested-resume-" + tc.name
|
|
|
|
runner1 := NewRunner(ctx, RunnerConfig{
|
|
Agent: outerAgent,
|
|
EnableStreaming: tc.enableStreaming,
|
|
CheckPointStore: store,
|
|
})
|
|
|
|
// --- phase 1: run and cancel while inner agent model is blocked ---
|
|
cancelOpt, cancelFn := WithCancel()
|
|
iter := runner1.Run(ctx, []Message{schema.UserMessage("go")}, cancelOpt, WithCheckPointID(checkpointID))
|
|
|
|
// Wait for inner model to start (meaning outer model already returned tool call).
|
|
select {
|
|
case <-innerModel.started:
|
|
case <-time.After(5 * time.Second):
|
|
t.Fatal("inner model did not start")
|
|
}
|
|
|
|
// At this point outerModel should have been called exactly once.
|
|
assert.Equal(t, int32(1), atomic.LoadInt32(&outerModelCallCount),
|
|
"outer model should have been called once before cancel")
|
|
|
|
// Cancel immediately. Recursive cases additionally propagate the cancel
|
|
// request into the AgentTool's internal ChatModelAgent.
|
|
var handle *CancelHandle
|
|
if tc.recursive {
|
|
handle, _ = cancelFn(WithRecursive())
|
|
} else {
|
|
handle, _ = cancelFn()
|
|
}
|
|
cancelErr := handle.Wait()
|
|
assert.NoError(t, cancelErr)
|
|
|
|
_, hasCancelError := drainEvents(iter)
|
|
assert.True(t, hasCancelError, "expected CancelError from canceled nested agent tool")
|
|
|
|
// --- phase 2: resume with Runner.Resume (no ResumeWithParams, no interrupt ID) ---
|
|
// Build fresh agents for resume. Recursive cancel should resume the
|
|
// inner ChatModelAgent inside AgentTool before the top-level
|
|
// ChatModelAgent produces the final answer.
|
|
resumeFirstModelCall := make(chan string, 5)
|
|
resumeOuterModelCallCount := int32(0)
|
|
resumeOuterModel := &countingChatModel{
|
|
callCount: &resumeOuterModelCallCount,
|
|
callCh: resumeFirstModelCall,
|
|
callLabel: "outer",
|
|
responses: []*schema.Message{
|
|
schema.AssistantMessage("outer completed after resume", nil),
|
|
},
|
|
}
|
|
|
|
resumeInnerModelCallCount := int32(0)
|
|
resumeInnerResponses := []*schema.Message{schema.AssistantMessage("inner agent done after resume", nil)}
|
|
if len(innerTools) > 0 {
|
|
resumeInnerResponses = []*schema.Message{
|
|
toolCallMsg(toolCall("ic1", "inner_tool", `{"input":"x"}`)),
|
|
schema.AssistantMessage("inner agent done after resume", nil),
|
|
}
|
|
}
|
|
resumeInnerModel := &countingChatModel{
|
|
callCount: &resumeInnerModelCallCount,
|
|
callCh: resumeFirstModelCall,
|
|
callLabel: "inner",
|
|
responses: resumeInnerResponses,
|
|
}
|
|
resumeInnerCfg := &ChatModelAgentConfig{
|
|
Name: "InnerAgent",
|
|
Description: "inner agent that returns immediately on resume",
|
|
Instruction: "you are an inner agent",
|
|
Model: resumeInnerModel,
|
|
}
|
|
if len(innerTools) > 0 {
|
|
resumeInnerCfg.ToolsConfig = ToolsConfig{
|
|
ToolsNodeConfig: compose.ToolsNodeConfig{
|
|
Tools: []tool.BaseTool{newSlowTool("inner_tool", 0, "inner tool result")},
|
|
},
|
|
}
|
|
}
|
|
resumeInnerAgent, err := NewChatModelAgent(ctx, resumeInnerCfg)
|
|
require.NoError(t, err)
|
|
|
|
resumeOuterAgent, err := NewChatModelAgent(ctx, &ChatModelAgentConfig{
|
|
Name: "OuterAgent",
|
|
Description: "outer agent with nested agent tool",
|
|
Instruction: "you are an outer agent",
|
|
Model: resumeOuterModel,
|
|
ToolsConfig: ToolsConfig{
|
|
ToolsNodeConfig: compose.ToolsNodeConfig{
|
|
Tools: []tool.BaseTool{NewAgentTool(ctx, resumeInnerAgent)},
|
|
},
|
|
},
|
|
})
|
|
require.NoError(t, err)
|
|
|
|
runner2 := NewRunner(ctx, RunnerConfig{
|
|
Agent: resumeOuterAgent,
|
|
EnableStreaming: tc.enableStreaming,
|
|
CheckPointStore: store,
|
|
})
|
|
|
|
resumeIter, err := runner2.Resume(ctx, checkpointID)
|
|
require.NoError(t, err)
|
|
|
|
select {
|
|
case firstModel := <-resumeFirstModelCall:
|
|
if tc.recursive {
|
|
assert.Equal(t, "inner", firstModel,
|
|
"recursive cancel should resume the AgentTool/internal ChatModelAgent first")
|
|
} else {
|
|
assert.Contains(t, []string{"outer", "inner"}, firstModel,
|
|
"non-recursive cancel does not define whether a root or already-persisted inner checkpoint resumes first")
|
|
}
|
|
case <-time.After(5 * time.Second):
|
|
t.Fatal("no model call observed during resume")
|
|
}
|
|
|
|
var resumeEvents []*AgentEvent
|
|
for {
|
|
event, ok := resumeIter.Next()
|
|
if !ok {
|
|
break
|
|
}
|
|
if event.Err != nil {
|
|
t.Fatalf("unexpected error during resume: %v", event.Err)
|
|
}
|
|
resumeEvents = append(resumeEvents, event)
|
|
}
|
|
|
|
// The outer model should have been called exactly once during resume
|
|
// (to produce the final answer after receiving tool results).
|
|
// If it was called with the original user message (restarting from scratch),
|
|
// the counting model would either exceed its response list or the call count
|
|
// would be wrong.
|
|
assert.Equal(t, int32(1), atomic.LoadInt32(&resumeOuterModelCallCount),
|
|
"outer model should be called exactly once during resume (for final answer after tool results), "+
|
|
"not restarted from the beginning")
|
|
|
|
// Verify we got the completion output.
|
|
var gotOutput bool
|
|
for _, event := range resumeEvents {
|
|
content, err := messageOutputContent(event)
|
|
require.NoError(t, err)
|
|
if content == "outer completed after resume" {
|
|
gotOutput = true
|
|
}
|
|
}
|
|
assert.True(t, gotOutput, "should get final output from resumed outer agent")
|
|
})
|
|
}
|
|
}
|
|
|
|
func TestWithCancel_CancelImmediate_RecursiveAgentTool_ResumeDeepestAgentTool(t *testing.T) {
|
|
ctx := context.Background()
|
|
|
|
leafModel := newBlockingChatModel(schema.AssistantMessage("leaf done", nil))
|
|
t.Cleanup(func() {
|
|
close(leafModel.unblockCh)
|
|
})
|
|
leafAgent, err := NewChatModelAgent(ctx, &ChatModelAgentConfig{
|
|
Name: "LeafAgent",
|
|
Description: "leaf agent that blocks",
|
|
Instruction: "you are a leaf agent",
|
|
Model: leafModel,
|
|
})
|
|
require.NoError(t, err)
|
|
|
|
middleModelCallCount := int32(0)
|
|
middleModel := &countingChatModel{
|
|
callCount: &middleModelCallCount,
|
|
responses: []*schema.Message{
|
|
toolCallMsg(toolCall("middle-leaf", "LeafAgent", `{"request":"leaf work"}`)),
|
|
schema.AssistantMessage("middle done", nil),
|
|
},
|
|
}
|
|
middleAgent, err := NewChatModelAgent(ctx, &ChatModelAgentConfig{
|
|
Name: "MiddleAgent",
|
|
Description: "middle agent with an agent tool",
|
|
Instruction: "you are a middle agent",
|
|
Model: middleModel,
|
|
ToolsConfig: ToolsConfig{
|
|
ToolsNodeConfig: compose.ToolsNodeConfig{
|
|
Tools: []tool.BaseTool{NewAgentTool(ctx, leafAgent)},
|
|
},
|
|
},
|
|
})
|
|
require.NoError(t, err)
|
|
|
|
outerModelCallCount := int32(0)
|
|
outerModel := &countingChatModel{
|
|
callCount: &outerModelCallCount,
|
|
responses: []*schema.Message{
|
|
toolCallMsg(toolCall("outer-middle", "MiddleAgent", `{"request":"middle work"}`)),
|
|
schema.AssistantMessage("outer done", nil),
|
|
},
|
|
}
|
|
outerAgent, err := NewChatModelAgent(ctx, &ChatModelAgentConfig{
|
|
Name: "OuterAgent",
|
|
Description: "outer agent with recursive agent tool nesting",
|
|
Instruction: "you are an outer agent",
|
|
Model: outerModel,
|
|
ToolsConfig: ToolsConfig{
|
|
ToolsNodeConfig: compose.ToolsNodeConfig{
|
|
Tools: []tool.BaseTool{NewAgentTool(ctx, middleAgent)},
|
|
},
|
|
},
|
|
})
|
|
require.NoError(t, err)
|
|
|
|
store := newCancelTestStore()
|
|
checkpointID := "cancel-recursive-agent-tool-resume"
|
|
runner1 := NewRunner(ctx, RunnerConfig{Agent: outerAgent, CheckPointStore: store})
|
|
|
|
cancelOpt, cancelFn := WithCancel()
|
|
iter := runner1.Run(ctx, []Message{schema.UserMessage("go")}, cancelOpt, WithCheckPointID(checkpointID))
|
|
|
|
select {
|
|
case <-leafModel.started:
|
|
case <-time.After(5 * time.Second):
|
|
t.Fatal("leaf model did not start")
|
|
}
|
|
|
|
handle, _ := cancelFn(WithRecursive())
|
|
require.NoError(t, handle.Wait())
|
|
_, hasCancelError := drainEvents(iter)
|
|
assert.True(t, hasCancelError, "expected CancelError from recursive nested agent tool")
|
|
|
|
firstModelCall := make(chan string, 8)
|
|
resumeLeafModelCallCount := int32(0)
|
|
resumeLeafModel := &countingChatModel{
|
|
callCount: &resumeLeafModelCallCount,
|
|
callCh: firstModelCall,
|
|
callLabel: "leaf",
|
|
responses: []*schema.Message{schema.AssistantMessage("leaf done after resume", nil)},
|
|
}
|
|
resumeLeafAgent, err := NewChatModelAgent(ctx, &ChatModelAgentConfig{
|
|
Name: "LeafAgent",
|
|
Description: "leaf agent that returns on resume",
|
|
Instruction: "you are a leaf agent",
|
|
Model: resumeLeafModel,
|
|
})
|
|
require.NoError(t, err)
|
|
|
|
resumeMiddleModelCallCount := int32(0)
|
|
resumeMiddleModel := &countingChatModel{
|
|
callCount: &resumeMiddleModelCallCount,
|
|
callCh: firstModelCall,
|
|
callLabel: "middle",
|
|
responses: []*schema.Message{schema.AssistantMessage("middle done after resume", nil)},
|
|
}
|
|
resumeMiddleAgent, err := NewChatModelAgent(ctx, &ChatModelAgentConfig{
|
|
Name: "MiddleAgent",
|
|
Description: "middle agent with an agent tool",
|
|
Instruction: "you are a middle agent",
|
|
Model: resumeMiddleModel,
|
|
ToolsConfig: ToolsConfig{
|
|
ToolsNodeConfig: compose.ToolsNodeConfig{
|
|
Tools: []tool.BaseTool{NewAgentTool(ctx, resumeLeafAgent)},
|
|
},
|
|
},
|
|
})
|
|
require.NoError(t, err)
|
|
|
|
resumeOuterModelCallCount := int32(0)
|
|
resumeOuterModel := &countingChatModel{
|
|
callCount: &resumeOuterModelCallCount,
|
|
callCh: firstModelCall,
|
|
callLabel: "outer",
|
|
responses: []*schema.Message{schema.AssistantMessage("outer done after resume", nil)},
|
|
}
|
|
resumeOuterAgent, err := NewChatModelAgent(ctx, &ChatModelAgentConfig{
|
|
Name: "OuterAgent",
|
|
Description: "outer agent with recursive agent tool nesting",
|
|
Instruction: "you are an outer agent",
|
|
Model: resumeOuterModel,
|
|
ToolsConfig: ToolsConfig{
|
|
ToolsNodeConfig: compose.ToolsNodeConfig{
|
|
Tools: []tool.BaseTool{NewAgentTool(ctx, resumeMiddleAgent)},
|
|
},
|
|
},
|
|
})
|
|
require.NoError(t, err)
|
|
|
|
runner2 := NewRunner(ctx, RunnerConfig{Agent: resumeOuterAgent, CheckPointStore: store})
|
|
resumeIter, err := runner2.Resume(ctx, checkpointID)
|
|
require.NoError(t, err)
|
|
|
|
select {
|
|
case first := <-firstModelCall:
|
|
assert.Equal(t, "leaf", first, "recursive AgentTool nesting should resume the deepest internal agent first")
|
|
case <-time.After(5 * time.Second):
|
|
t.Fatal("no model call observed during resume")
|
|
}
|
|
|
|
resumeEvents, hasResumeCancelError := drainEvents(resumeIter)
|
|
require.False(t, hasResumeCancelError, "resume should complete without another CancelError")
|
|
assert.NotEmpty(t, resumeEvents)
|
|
assert.Equal(t, int32(1), atomic.LoadInt32(&resumeLeafModelCallCount))
|
|
assert.Equal(t, int32(1), atomic.LoadInt32(&resumeMiddleModelCallCount))
|
|
assert.Equal(t, int32(1), atomic.LoadInt32(&resumeOuterModelCallCount))
|
|
}
|
|
|
|
func TestWithCancel_CancelImmediate_ConcurrentAgentTools_ResumeWithoutRestart(t *testing.T) {
|
|
ctx := context.Background()
|
|
|
|
innerAModel := newBlockingChatModel(schema.AssistantMessage("inner A done", nil))
|
|
innerBModel := newBlockingChatModel(schema.AssistantMessage("inner B done", nil))
|
|
t.Cleanup(func() {
|
|
close(innerAModel.unblockCh)
|
|
close(innerBModel.unblockCh)
|
|
})
|
|
|
|
innerAAgent, err := NewChatModelAgent(ctx, &ChatModelAgentConfig{
|
|
Name: "InnerAgentA",
|
|
Description: "inner agent A",
|
|
Instruction: "you are inner agent A",
|
|
Model: innerAModel,
|
|
})
|
|
require.NoError(t, err)
|
|
innerBAgent, err := NewChatModelAgent(ctx, &ChatModelAgentConfig{
|
|
Name: "InnerAgentB",
|
|
Description: "inner agent B",
|
|
Instruction: "you are inner agent B",
|
|
Model: innerBModel,
|
|
})
|
|
require.NoError(t, err)
|
|
|
|
outerModelCallCount := int32(0)
|
|
outerModel := &countingChatModel{
|
|
callCount: &outerModelCallCount,
|
|
responses: []*schema.Message{
|
|
toolCallMsg(
|
|
toolCall("outer-a", "InnerAgentA", `{"request":"work A"}`),
|
|
toolCall("outer-b", "InnerAgentB", `{"request":"work B"}`),
|
|
),
|
|
schema.AssistantMessage("outer done", nil),
|
|
},
|
|
}
|
|
outerAgent, err := NewChatModelAgent(ctx, &ChatModelAgentConfig{
|
|
Name: "OuterAgent",
|
|
Description: "outer agent with concurrent agent tools",
|
|
Instruction: "you are an outer agent",
|
|
Model: outerModel,
|
|
ToolsConfig: ToolsConfig{
|
|
ToolsNodeConfig: compose.ToolsNodeConfig{
|
|
Tools: []tool.BaseTool{
|
|
NewAgentTool(ctx, innerAAgent),
|
|
NewAgentTool(ctx, innerBAgent),
|
|
},
|
|
},
|
|
},
|
|
})
|
|
require.NoError(t, err)
|
|
|
|
store := newCancelTestStore()
|
|
checkpointID := "cancel-concurrent-agent-tools-resume"
|
|
runner1 := NewRunner(ctx, RunnerConfig{Agent: outerAgent, CheckPointStore: store})
|
|
|
|
cancelOpt, cancelFn := WithCancel()
|
|
iter := runner1.Run(ctx, []Message{schema.UserMessage("go")}, cancelOpt, WithCheckPointID(checkpointID))
|
|
|
|
for _, started := range []chan struct{}{innerAModel.started, innerBModel.started} {
|
|
select {
|
|
case <-started:
|
|
case <-time.After(5 * time.Second):
|
|
t.Fatal("both concurrent inner models should start before cancel")
|
|
}
|
|
}
|
|
|
|
handle, _ := cancelFn(WithRecursive())
|
|
require.NoError(t, handle.Wait())
|
|
_, hasCancelError := drainEvents(iter)
|
|
assert.True(t, hasCancelError, "expected CancelError from concurrent agent tools")
|
|
|
|
firstModelCall := make(chan string, 8)
|
|
resumeInnerAModelCallCount := int32(0)
|
|
resumeInnerAModel := &countingChatModel{
|
|
callCount: &resumeInnerAModelCallCount,
|
|
callCh: firstModelCall,
|
|
callLabel: "innerA",
|
|
responses: []*schema.Message{schema.AssistantMessage("inner A done after resume", nil)},
|
|
}
|
|
resumeInnerBModelCallCount := int32(0)
|
|
resumeInnerBModel := &countingChatModel{
|
|
callCount: &resumeInnerBModelCallCount,
|
|
callCh: firstModelCall,
|
|
callLabel: "innerB",
|
|
responses: []*schema.Message{schema.AssistantMessage("inner B done after resume", nil)},
|
|
}
|
|
|
|
resumeInnerAAgent, err := NewChatModelAgent(ctx, &ChatModelAgentConfig{
|
|
Name: "InnerAgentA",
|
|
Description: "inner agent A",
|
|
Instruction: "you are inner agent A",
|
|
Model: resumeInnerAModel,
|
|
})
|
|
require.NoError(t, err)
|
|
resumeInnerBAgent, err := NewChatModelAgent(ctx, &ChatModelAgentConfig{
|
|
Name: "InnerAgentB",
|
|
Description: "inner agent B",
|
|
Instruction: "you are inner agent B",
|
|
Model: resumeInnerBModel,
|
|
})
|
|
require.NoError(t, err)
|
|
|
|
resumeOuterModelCallCount := int32(0)
|
|
resumeOuterModel := &countingChatModel{
|
|
callCount: &resumeOuterModelCallCount,
|
|
callCh: firstModelCall,
|
|
callLabel: "outer",
|
|
responses: []*schema.Message{schema.AssistantMessage("outer done after resume", nil)},
|
|
}
|
|
resumeOuterAgent, err := NewChatModelAgent(ctx, &ChatModelAgentConfig{
|
|
Name: "OuterAgent",
|
|
Description: "outer agent with concurrent agent tools",
|
|
Instruction: "you are an outer agent",
|
|
Model: resumeOuterModel,
|
|
ToolsConfig: ToolsConfig{
|
|
ToolsNodeConfig: compose.ToolsNodeConfig{
|
|
Tools: []tool.BaseTool{
|
|
NewAgentTool(ctx, resumeInnerAAgent),
|
|
NewAgentTool(ctx, resumeInnerBAgent),
|
|
},
|
|
},
|
|
},
|
|
})
|
|
require.NoError(t, err)
|
|
|
|
runner2 := NewRunner(ctx, RunnerConfig{Agent: resumeOuterAgent, CheckPointStore: store})
|
|
resumeIter, err := runner2.Resume(ctx, checkpointID)
|
|
require.NoError(t, err)
|
|
|
|
select {
|
|
case first := <-firstModelCall:
|
|
assert.Contains(t, []string{"innerA", "innerB"}, first,
|
|
"concurrent AgentTools should resume an internal agent before the outer model")
|
|
case <-time.After(5 * time.Second):
|
|
t.Fatal("no model call observed during resume")
|
|
}
|
|
|
|
resumeEvents, hasResumeCancelError := drainEvents(resumeIter)
|
|
require.False(t, hasResumeCancelError, "resume should complete without another CancelError")
|
|
assert.NotEmpty(t, resumeEvents)
|
|
assert.Equal(t, int32(1), atomic.LoadInt32(&resumeInnerAModelCallCount))
|
|
assert.Equal(t, int32(1), atomic.LoadInt32(&resumeInnerBModelCallCount))
|
|
assert.Equal(t, int32(1), atomic.LoadInt32(&resumeOuterModelCallCount))
|
|
}
|
|
|
|
// countingChatModel is a chat model that counts calls and records inputs.
|
|
// It returns responses from a fixed slice, indexed by call count.
|
|
type countingChatModel struct {
|
|
callCount *int32
|
|
inputsCh chan []*schema.Message // optional: receives a copy of each input
|
|
callCh chan string // optional: receives callLabel when Generate is called
|
|
callLabel string
|
|
responses []*schema.Message
|
|
}
|
|
|
|
func (m *countingChatModel) Generate(_ context.Context, input []*schema.Message, _ ...model.Option) (*schema.Message, error) {
|
|
idx := int(atomic.AddInt32(m.callCount, 1)) - 1
|
|
if m.callCh != nil {
|
|
select {
|
|
case m.callCh <- m.callLabel:
|
|
default:
|
|
}
|
|
}
|
|
if m.inputsCh != nil {
|
|
cp := make([]*schema.Message, len(input))
|
|
copy(cp, input)
|
|
select {
|
|
case m.inputsCh <- cp:
|
|
default:
|
|
}
|
|
}
|
|
if idx >= len(m.responses) {
|
|
return nil, fmt.Errorf("countingChatModel: call %d exceeds %d responses (outer model was called too many times - possible restart from beginning)", idx+1, len(m.responses))
|
|
}
|
|
return m.responses[idx], nil
|
|
}
|
|
|
|
func (m *countingChatModel) Stream(_ context.Context, input []*schema.Message, opts ...model.Option) (*schema.StreamReader[*schema.Message], error) {
|
|
msg, err := m.Generate(context.Background(), input, opts...)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
return schema.StreamReaderFromArray([]*schema.Message{msg}), nil
|
|
}
|
|
|
|
func (m *countingChatModel) BindTools(_ []*schema.ToolInfo) error { return nil }
|
|
|
|
func messageOutputContent(event *AgentEvent) (string, error) {
|
|
if event.Output == nil || event.Output.MessageOutput == nil {
|
|
return "", nil
|
|
}
|
|
mo := event.Output.MessageOutput
|
|
if mo.IsStreaming {
|
|
msg, err := schema.ConcatMessageStream(mo.MessageStream)
|
|
if err != nil {
|
|
return "", err
|
|
}
|
|
if msg == nil {
|
|
return "", nil
|
|
}
|
|
return msg.Content, nil
|
|
}
|
|
if mo.Message == nil {
|
|
return "", nil
|
|
}
|
|
return mo.Message.Content, nil
|
|
}
|