Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
26 changes: 26 additions & 0 deletions workflow/agentworkflow/hosting.go
Original file line number Diff line number Diff line change
Expand Up @@ -429,6 +429,11 @@ func (h *hostExecutor) runAgentAndDispatch(wctx *workflow.Context, messages []*m
}
resp.Update(update)
}
// Match .NET hosting: if cancellation arrives after the last streamed update
// but before the turn is finalized, do not emit the aggregated completion.
if err := wctx.Err(); err != nil {
return err
}
Comment thread
michelle-clayton-work marked this conversation as resolved.
resp.Coalesce()

// Stamp this hosting executor's name on every aggregated message,
Expand All @@ -441,11 +446,17 @@ func (h *hostExecutor) runAgentAndDispatch(wctx *workflow.Context, messages []*m
}

if h.cfg.EmitResponseEvents {
if err := wctx.Err(); err != nil {
return err
}
if err := wctx.YieldOutput(&resp); err != nil {
return err
}
}

if err := wctx.Err(); err != nil {
return err
}
if err := h.dispatchRequests(wctx, resp.Messages); err != nil {
return err
}
Expand All @@ -455,11 +466,17 @@ func (h *hostExecutor) runAgentAndDispatch(wctx *workflow.Context, messages []*m
// workflow can cause invalid request errors when the receiving agent uses an
// API that does not accept those output-only item types as input.
if forwardableMessages := filterForwardableMessages(resp.Messages); len(forwardableMessages) > 0 {
if err := wctx.Err(); err != nil {
return err
}
if err := wctx.SendMessage("", forwardableMessages); err != nil {
return err
}
}

if err := wctx.Err(); err != nil {
return err
}
if err := h.releasePendingTurnIfReady(wctx); err != nil {
return err
}
Expand Down Expand Up @@ -529,6 +546,9 @@ func (h *hostExecutor) releasePendingTurnIfReady(wctx *workflow.Context) error {
}
// Forward a fresh TurnToken stamped with the resolved EmitEvents value so
// downstream executors observe the effective per-turn setting.
if err := wctx.Err(); err != nil {
return err
}
return wctx.SendMessage("", workflow.TurnToken{EmitEvents: emit})
}

Expand Down Expand Up @@ -653,6 +673,9 @@ func dispatchTrackedRequests[T any](wctx *workflow.Context, order []string, requ
if !ok {
continue
}
if err := wctx.Err(); err != nil {
return err
}
delete(requests, id)
added, err := dispatcher.TrackRequest(wctx, request)
if err != nil {
Expand All @@ -664,6 +687,9 @@ func dispatchTrackedRequests[T any](wctx *workflow.Context, order []string, requ
}

for _, request := range dispatches {
if err := wctx.Err(); err != nil {
return err
}
if err := dispatcher.DispatchRequest(wctx, request); err != nil {
return err
}
Expand Down
89 changes: 89 additions & 0 deletions workflow/agentworkflow/hosting_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -173,6 +173,24 @@ func newContentAgent(updates ...*agent.ResponseUpdate) *agent.Agent {
)
}

func newCancelOnCompletionAgent(cancel context.CancelFunc) *agent.Agent {
run := func(ctx context.Context, _ []*message.Message, _ ...agent.Option) iter.Seq2[*agent.ResponseUpdate, error] {
return func(yield func(*agent.ResponseUpdate, error) bool) {
if !yield(&agent.ResponseUpdate{
Role: message.RoleAssistant,
Contents: []message.Content{&message.TextContent{Text: "done"}},
}, nil) {
return
}
cancel()
}
}
return agent.New(
agent.ProviderConfig{ProviderName: "cancel-on-completion", Run: run},
agent.Config{ID: testAgentID, Name: testAgentName},
)
}

func newNamedNoopAgent(id string, name string) *agent.Agent {
run := func(_ context.Context, _ []*message.Message, _ ...agent.Option) iter.Seq2[*agent.ResponseUpdate, error] {
return func(yield func(*agent.ResponseUpdate, error) bool) {
Expand Down Expand Up @@ -470,6 +488,77 @@ func TestHostedAgent_EmitsResponseIfConfigured(t *testing.T) {
}
}

func TestHostedAgent_CancellationSuppressesFinalResponseEmission(t *testing.T) {
ctx, cancel := context.WithCancel(t.Context())
defer cancel()

binding := agentworkflow.New(newCancelOnCompletionAgent(cancel), agentworkflow.Config{
EmitResponseEvents: true,
ForwardIncomingMessages: new(false),
})
executor, err := binding.CreateInstance("")
if err != nil {
t.Fatalf("CreateInstance: %v", err)
}

state := map[string]any{}
var yielded []any
var sent []any
wctx := &workflow.Context{
Context: ctx,
AddEvent: func(workflow.Event) error {
return nil
},
SendMessage: func(_ string, msg any) error {
sent = append(sent, msg)
return nil
},
YieldOutput: func(output any) error {
yielded = append(yielded, output)
return nil
},
ReadState: func(key string, _ string) (any, error) {
return state[key], nil
},
ReadOrInitState: func(key string, _ string, init func(context.Context, string, string) (any, error)) (any, error) {
if value, ok := state[key]; ok {
return value, nil
}
value, err := init(ctx, key, "")
if err != nil {
return nil, err
}
state[key] = value
return value, nil
},
QueueStateUpdate: func(key string, _ string, value any) error {
if value == nil {
delete(state, key)
return nil
}
state[key] = value
return nil
},
}

if _, err := executor.Execute(wctx, &message.Message{
Role: message.RoleUser,
Contents: []message.Content{&message.TextContent{Text: "go"}},
}); err != nil {
t.Fatalf("buffer message: %v", err)
}

if _, err := executor.Execute(wctx, workflow.TurnToken{}); !errors.Is(err, context.Canceled) {
t.Fatalf("turn error = %v, want context canceled", err)
}
if len(yielded) != 0 {
t.Fatalf("yielded outputs = %d, want 0 after cancellation", len(yielded))
}
if len(sent) != 0 {
t.Fatalf("sent messages = %d, want 0 after cancellation", len(sent))
}
}

// TestHostedAgent_ReassignsRolesIfConfigured verifies that the agent receives
// only RoleUser / self-authored RoleAssistant messages by default, and that
// disabling reassignment causes a RoleCheckAgent to surface an error event
Expand Down
Loading