Files
probo/pkg/agent/run.go
2026-03-14 12:28:26 +01:00

1181 lines
30 KiB
Go

// Copyright (c) 2026 Probo Inc <hello@getprobo.com>.
//
// Permission to use, copy, modify, and/or distribute this software for any
// purpose with or without fee is hereby granted, provided that the above
// copyright notice and this permission notice appear in all copies.
//
// THE SOFTWARE IS PROVIDED "AS IS" AND THE AUTHOR DISCLAIMS ALL WARRANTIES WITH
// REGARD TO THIS SOFTWARE INCLUDING ALL IMPLIED WARRANTIES OF MERCHANTABILITY
// AND FITNESS. IN NO EVENT SHALL THE AUTHOR BE LIABLE FOR ANY SPECIAL, DIRECT,
// INDIRECT, OR CONSEQUENTIAL DAMAGES OR ANY DAMAGES WHATSOEVER RESULTING FROM
// LOSS OF USE, DATA OR PROFITS, WHETHER IN AN ACTION OF CONTRACT, NEGLIGENCE OR
// OTHER TORTIOUS ACTION, ARISING OUT OF OR IN CONNECTION WITH THE USE OR
// PERFORMANCE OF THIS SOFTWARE.
package agent
import (
"context"
"errors"
"fmt"
"sync"
"go.gearno.de/kit/log"
"go.opentelemetry.io/otel"
"go.opentelemetry.io/otel/attribute"
"go.opentelemetry.io/otel/codes"
"go.opentelemetry.io/otel/trace"
"go.probo.inc/probo/pkg/llm"
)
const tracerName = "go.probo.inc/probo/pkg/agent"
type (
CallLLMFunc func(ctx context.Context, agent *Agent, req *llm.ChatCompletionRequest) (*llm.ChatCompletionResponse, error)
runOpts struct {
callLLM CallLLMFunc
onEvent func(ctx context.Context, ev StreamEvent)
skipInputGuardrails bool
skipSessionLoad bool
initialUsage llm.Usage
initialTurns int
}
loopState struct {
agent *Agent
toolMap map[string]ToolDescriptor
toolDefs []llm.Tool
messages []llm.Message
inputMessages []llm.Message
systemPrompt string
totalUsage llm.Usage
turns int
toolUsedInRun bool
tracer trace.Tracer
runSpan trace.Span
opts runOpts
logger *log.Logger
}
parallelToolEntry struct {
result ToolResult
err error
}
)
func noopEvent(_ context.Context, _ StreamEvent) {}
func blockingCallLLM(ctx context.Context, agent *Agent, req *llm.ChatCompletionRequest) (*llm.ChatCompletionResponse, error) {
return agent.client.ChatCompletion(ctx, req)
}
func (a *Agent) Run(ctx context.Context, messages []llm.Message) (*Result, error) {
return coreLoop(
ctx,
a,
messages,
runOpts{
callLLM: blockingCallLLM,
onEvent: noopEvent,
},
)
}
func (s *loopState) resolveAgentTools(ctx context.Context) error {
tools, toolMap, err := s.agent.resolveTools(ctx)
if err != nil {
return err
}
toolDefs := make([]llm.Tool, len(tools))
for i, t := range tools {
toolDefs[i] = t.Definition()
}
s.toolMap = toolMap
s.toolDefs = toolDefs
return nil
}
func (s *loopState) finishRun(ctx context.Context, result *Result, err error) (*Result, error) {
defer s.runSpan.End()
defer func() {
emitHook(s.agent, func(h RunHooks) { h.OnRunEnd(ctx, s.agent, result, err) })
}()
if err != nil {
s.runSpan.RecordError(err)
s.runSpan.SetStatus(codes.Error, err.Error())
s.opts.onEvent(ctx, StreamEvent{Type: StreamEventError, Agent: s.agent, Err: err})
s.logger.ErrorCtx(
ctx,
"agent run failed",
log.String("agent", s.agent.name),
log.Int("turns", s.turns),
log.Error(err),
)
return result, err
}
if s.agent.session != nil {
if saveErr := s.agent.session.Save(ctx, s.agent.sessionID, result.Messages); saveErr != nil {
err = fmt.Errorf("cannot save session: %w", saveErr)
result = nil
s.runSpan.RecordError(err)
s.runSpan.SetStatus(codes.Error, err.Error())
s.opts.onEvent(ctx, StreamEvent{Type: StreamEventError, Agent: s.agent, Err: err})
s.logger.ErrorCtx(
ctx,
"cannot save session",
log.String("agent", s.agent.name),
log.Error(err),
)
return result, err
}
}
s.runSpan.SetAttributes(
attribute.Int("agent.turns", result.Turns),
attribute.Int("agent.usage.input_tokens", result.Usage.InputTokens),
attribute.Int("agent.usage.output_tokens", result.Usage.OutputTokens),
)
s.logger.InfoCtx(
ctx,
"agent run completed",
log.String("agent", s.agent.name),
log.Int("turns", result.Turns),
log.Int("input_tokens", result.Usage.InputTokens),
log.Int("output_tokens", result.Usage.OutputTokens),
)
s.opts.onEvent(ctx, StreamEvent{Type: StreamEventComplete, Agent: s.agent, Result: result})
return result, err
}
func (s *loopState) applyHandoff(ctx context.Context, handoffTarget *Handoff) error {
emitHook(s.agent, func(h RunHooks) { h.OnHandoff(ctx, s.agent, handoffTarget.Agent) })
emitAgentHook(handoffTarget.Agent, func(h AgentHooks) { h.OnHandoff(ctx, handoffTarget.Agent, s.agent) })
s.opts.onEvent(ctx, StreamEvent{Type: StreamEventHandoff, Agent: handoffTarget.Agent})
s.logger.InfoCtx(
ctx,
"agent handoff",
log.String("from", s.agent.name),
log.String("to", handoffTarget.Agent.name),
)
if handoffTarget.InputFilter != nil {
s.messages = handoffTarget.InputFilter(
HandoffInputData{
InputHistory: s.inputMessages,
NewItems: s.messages[len(s.inputMessages):],
},
)
}
s.agent = handoffTarget.Agent
s.logger = s.agent.logger
if err := s.resolveAgentTools(ctx); err != nil {
return err
}
emitAgentHook(s.agent, func(h AgentHooks) { h.OnStart(ctx, s.agent) })
s.opts.onEvent(ctx, StreamEvent{Type: StreamEventAgentStart, Agent: s.agent})
s.systemPrompt = s.agent.buildSystemPrompt(ctx)
s.toolUsedInRun = false
return nil
}
func coreLoop(ctx context.Context, startAgent *Agent, inputMessages []llm.Message, opts runOpts) (*Result, error) {
s := &loopState{
agent: startAgent,
inputMessages: inputMessages,
systemPrompt: startAgent.buildSystemPrompt(ctx),
totalUsage: opts.initialUsage,
turns: opts.initialTurns,
tracer: otel.GetTracerProvider().Tracer(tracerName),
opts: opts,
logger: startAgent.logger,
}
s.messages = make([]llm.Message, 0, len(inputMessages))
if !opts.skipSessionLoad && s.agent.session != nil {
prev, err := s.agent.session.Load(ctx, s.agent.sessionID)
if err != nil {
return nil, fmt.Errorf("cannot load session: %w", err)
}
s.logger.InfoCtx(
ctx,
"session loaded",
log.String("session_id", s.agent.sessionID),
log.Int("message_count", len(prev)),
)
s.messages = append(s.messages, prev...)
}
s.messages = append(s.messages, inputMessages...)
if !opts.skipInputGuardrails {
if err := runInputGuardrails(ctx, s.agent, s.messages); err != nil {
opts.onEvent(ctx, StreamEvent{Type: StreamEventError, Agent: s.agent, Err: err})
s.logger.ErrorCtx(
ctx,
"input guardrail tripped",
log.Error(err),
)
return nil, err
}
}
emitHook(s.agent, func(h RunHooks) { h.OnRunStart(ctx, s.agent, s.messages) })
emitAgentHook(s.agent, func(h AgentHooks) { h.OnStart(ctx, s.agent) })
opts.onEvent(ctx, StreamEvent{Type: StreamEventAgentStart, Agent: s.agent})
ctx, s.runSpan = s.tracer.Start(
ctx,
fmt.Sprintf("agent.run %s", s.agent.name),
trace.WithSpanKind(trace.SpanKindInternal),
trace.WithAttributes(
attribute.String("agent.name", s.agent.name),
attribute.String("agent.model", s.agent.model),
),
)
if err := s.resolveAgentTools(ctx); err != nil {
s.runSpan.End()
return nil, err
}
s.logger.InfoCtx(
ctx,
"agent run started",
log.String("model", s.agent.model),
log.Int("max_turns", s.agent.maxTurns),
log.Int("tool_count", len(s.toolDefs)),
)
for {
if err := ctx.Err(); err != nil {
return s.finishRun(ctx, nil, fmt.Errorf("cannot complete: %w", err))
}
if s.turns >= s.agent.maxTurns {
return s.finishRun(ctx, nil, &MaxTurnsExceededError{MaxTurns: s.agent.maxTurns})
}
fullMessages := buildFullMessages(s.systemPrompt, s.messages)
responseFormat := s.agent.responseFormat
if responseFormat == nil && s.agent.outputType != nil {
responseFormat = s.agent.outputType.responseFormat()
}
toolChoice := s.agent.modelSettings.ToolChoice
if s.toolUsedInRun && s.agent.resetToolChoice && toolChoice != nil {
toolChoice = nil
}
req := &llm.ChatCompletionRequest{
Model: s.agent.model,
Messages: fullMessages,
Tools: s.toolDefs,
Temperature: s.agent.modelSettings.Temperature,
TopP: s.agent.modelSettings.TopP,
FrequencyPenalty: s.agent.modelSettings.FrequencyPenalty,
PresencePenalty: s.agent.modelSettings.PresencePenalty,
MaxTokens: s.agent.modelSettings.MaxTokens,
ToolChoice: toolChoice,
ParallelToolCalls: s.agent.modelSettings.ParallelToolCalls,
ResponseFormat: responseFormat,
}
s.logger.InfoCtx(
ctx,
"calling LLM",
log.Int("turn", s.turns+1),
log.Int("message_count", len(fullMessages)),
)
resp, err := callLLMWithHooks(ctx, s.agent, req, opts)
if err != nil {
return s.finishRun(ctx, nil, fmt.Errorf("cannot call LLM: %w", err))
}
s.totalUsage = s.totalUsage.Add(resp.Usage)
s.turns++
s.logger.InfoCtx(
ctx,
"LLM response received",
log.Int("turn", s.turns),
log.String("finish_reason", string(resp.FinishReason)),
log.Int("input_tokens", resp.Usage.InputTokens),
log.Int("output_tokens", resp.Usage.OutputTokens),
)
s.messages = append(s.messages, resp.Message)
switch resp.FinishReason {
case llm.FinishReasonStop, llm.FinishReasonLength:
if err := runOutputGuardrails(ctx, s.agent, resp.Message); err != nil {
return s.finishRun(ctx, nil, err)
}
result := &Result{
Messages: s.messages,
Usage: s.totalUsage,
Turns: s.turns,
LastAgent: s.agent,
}
emitAgentHook(s.agent, func(h AgentHooks) { h.OnEnd(ctx, s.agent, resp.Message.Text()) })
opts.onEvent(ctx, StreamEvent{Type: StreamEventAgentEnd, Agent: s.agent})
return s.finishRun(ctx, result, nil)
case llm.FinishReasonToolCalls:
s.toolUsedInRun = true
s.logger.InfoCtx(
ctx,
"dispatching tool calls",
log.Int("count", len(resp.Message.ToolCalls)),
)
handoffTarget, toolResults, toolMsgs, err := dispatchToolCalls(
ctx,
s.tracer,
s.agent,
resp.Message.ToolCalls,
s.toolMap,
opts.onEvent,
s.logger,
)
s.messages = append(s.messages, toolMsgs...)
if err != nil {
if nae, ok := errors.AsType[*needsApprovalError](err); ok {
s.logger.InfoCtx(
ctx,
"run interrupted, approval required",
log.Int("pending_count", len(nae.pendingApprovals)),
)
msgsCopy := make([]llm.Message, len(s.messages))
copy(msgsCopy, s.messages)
return s.finishRun(
ctx,
nil,
&InterruptedError{
ToolCalls: nae.allToolCalls,
PendingApprovals: nae.pendingApprovals,
Agent: s.agent,
Messages: msgsCopy,
Usage: s.totalUsage,
Turns: s.turns,
},
)
}
if nie, ok := errors.AsType[*nestedInterruptionError](err); ok {
s.logger.InfoCtx(
ctx,
"nested agent interrupted, approval required",
log.Int("pending_count", len(nie.inner.PendingApprovals)),
log.String("nested_agent", nie.inner.Agent.name),
)
msgsCopy := make([]llm.Message, len(s.messages))
copy(msgsCopy, s.messages)
return s.finishRun(
ctx,
nil,
&InterruptedError{
ToolCalls: nie.inner.ToolCalls,
PendingApprovals: nie.inner.PendingApprovals,
Agent: nie.inner.Agent,
Messages: nie.inner.Messages,
Usage: nie.inner.Usage,
Turns: nie.inner.Turns,
outerState: &outerLoopState{
agent: s.agent,
messages: msgsCopy,
usage: s.totalUsage,
turns: s.turns,
allToolCalls: nie.allToolCalls,
toolCallID: nie.toolCallID,
completedCalls: nie.completedCalls,
innerInterrupt: nie.inner,
},
},
)
}
return s.finishRun(ctx, nil, err)
}
finalOutput, isFinal, behaviorErr := s.agent.toolUseBehavior(ctx, toolResults)
if behaviorErr != nil {
return s.finishRun(ctx, nil, fmt.Errorf("cannot evaluate tool use behavior: %w", behaviorErr))
}
if isFinal {
s.messages = append(
s.messages, llm.Message{
Role: llm.RoleAssistant,
Parts: []llm.Part{llm.TextPart{Text: finalOutput}},
},
)
result := &Result{
Messages: s.messages,
Usage: s.totalUsage,
Turns: s.turns,
LastAgent: s.agent,
}
emitAgentHook(s.agent, func(h AgentHooks) { h.OnEnd(ctx, s.agent, finalOutput) })
opts.onEvent(ctx, StreamEvent{Type: StreamEventAgentEnd, Agent: s.agent})
return s.finishRun(ctx, result, nil)
}
if handoffTarget != nil {
if handoffErr := s.applyHandoff(ctx, handoffTarget); handoffErr != nil {
return s.finishRun(ctx, nil, handoffErr)
}
}
case llm.FinishReasonContentFilter:
return s.finishRun(ctx, nil, fmt.Errorf("cannot complete: content was filtered by the provider"))
default:
return s.finishRun(ctx, nil, fmt.Errorf("cannot complete: unexpected finish reason %q", resp.FinishReason))
}
}
}
func buildFullMessages(systemPrompt string, messages []llm.Message) []llm.Message {
fullMessages := make([]llm.Message, 0, len(messages)+1)
if systemPrompt != "" {
fullMessages = append(
fullMessages,
llm.Message{
Role: llm.RoleSystem,
Parts: []llm.Part{llm.TextPart{Text: systemPrompt}},
},
)
}
return append(fullMessages, messages...)
}
func callLLMWithHooks(
ctx context.Context,
agent *Agent,
req *llm.ChatCompletionRequest,
opts runOpts,
) (*llm.ChatCompletionResponse, error) {
emitHook(agent, func(h RunHooks) { h.OnLLMStart(ctx, agent, req.Messages) })
emitAgentHook(agent, func(h AgentHooks) { h.OnLLMStart(ctx, agent, req.Messages) })
resp, err := opts.callLLM(ctx, agent, req)
if err != nil {
emitHook(agent, func(h RunHooks) { h.OnLLMEnd(ctx, agent, nil, err) })
emitAgentHook(agent, func(h AgentHooks) { h.OnLLMEnd(ctx, agent, nil, err) })
return nil, err
}
emitHook(agent, func(h RunHooks) { h.OnLLMEnd(ctx, agent, resp, nil) })
emitAgentHook(agent, func(h AgentHooks) { h.OnLLMEnd(ctx, agent, resp, nil) })
return resp, nil
}
func dispatchToolCalls(
ctx context.Context,
tracer trace.Tracer,
agent *Agent,
toolCalls []llm.ToolCall,
toolMap map[string]ToolDescriptor,
onEvent func(context.Context, StreamEvent),
logger *log.Logger,
) (*Handoff, []ToolCallResult, []llm.Message, error) {
if err := checkApproval(ctx, agent, toolCalls); err != nil {
return nil, nil, nil, err
}
return executeToolCalls(ctx, tracer, agent, toolCalls, toolMap, onEvent, logger)
}
func executeToolCalls(
ctx context.Context,
tracer trace.Tracer,
agent *Agent,
toolCalls []llm.ToolCall,
toolMap map[string]ToolDescriptor,
onEvent func(context.Context, StreamEvent),
logger *log.Logger,
) (*Handoff, []ToolCallResult, []llm.Message, error) {
descriptors := make([]ToolDescriptor, len(toolCalls))
handoffIdx := -1
for i, tc := range toolCalls {
desc, ok := toolMap[tc.Function.Name]
if !ok {
return nil, nil, nil, fmt.Errorf("cannot dispatch tool call: unknown tool %q", tc.Function.Name)
}
descriptors[i] = desc
if _, isHandoff := desc.(*handoffToolAdapter); isHandoff && handoffIdx == -1 {
handoffIdx = i
}
}
if handoffIdx >= 0 {
return executeWithHandoff(ctx, tracer, agent, toolCalls, descriptors, handoffIdx, onEvent, logger)
}
tools := make([]Tool, len(descriptors))
for i, d := range descriptors {
tools[i] = d.(Tool)
}
results, msgs, err := executeParallel(ctx, tracer, agent, toolCalls, tools, onEvent, logger)
return nil, results, msgs, err
}
func executeWithHandoff(
ctx context.Context,
tracer trace.Tracer,
agent *Agent,
toolCalls []llm.ToolCall,
descriptors []ToolDescriptor,
handoffIdx int,
onEvent func(context.Context, StreamEvent),
logger *log.Logger,
) (*Handoff, []ToolCallResult, []llm.Message, error) {
var (
results []ToolCallResult
msgs []llm.Message
)
for i := range handoffIdx {
logger.InfoCtx(
ctx,
"executing tool before handoff",
log.String("tool", toolCalls[i].Function.Name),
)
tr, err := executeSingleTool(ctx, tracer, agent, toolCalls[i], descriptors[i].(Tool), onEvent, logger)
if err != nil {
if ie, ok := errors.AsType[*InterruptedError](err); ok {
var completed []completedCall
for j := range results {
completed = append(
completed,
completedCall{
toolCallID: toolCalls[j].ID,
result: results[j].Result,
},
)
}
return nil, nil, msgs, &nestedInterruptionError{
inner: ie,
toolCallID: toolCalls[i].ID,
allToolCalls: toolCalls,
completedCalls: completed,
}
}
return nil, nil, msgs, err
}
msgs = append(
msgs,
llm.Message{
Role: llm.RoleTool,
ToolCallID: toolCalls[i].ID,
Parts: []llm.Part{llm.TextPart{Text: tr.Content}},
},
)
results = append(
results,
ToolCallResult{
ToolName: toolCalls[i].Function.Name,
Arguments: toolCalls[i].Function.Arguments,
Result: tr,
},
)
}
ht := descriptors[handoffIdx].(*handoffToolAdapter)
if ht.handoff.OnHandoff != nil {
if err := ht.handoff.OnHandoff(ctx); err != nil {
return nil, nil, msgs, fmt.Errorf("cannot execute handoff callback for %q: %w", ht.handoff.Agent.name, err)
}
}
msgs = append(
msgs,
llm.Message{
Role: llm.RoleTool,
ToolCallID: toolCalls[handoffIdx].ID,
Parts: []llm.Part{llm.TextPart{Text: fmt.Sprintf("Transferred to %s", ht.handoff.Agent.name)}},
},
)
for i := handoffIdx + 1; i < len(toolCalls); i++ {
msgs = append(
msgs,
llm.Message{
Role: llm.RoleTool,
ToolCallID: toolCalls[i].ID,
Parts: []llm.Part{
llm.TextPart{Text: "Tool call was not executed because a handoff occurred."},
},
},
)
}
return ht.handoff, results, msgs, nil
}
func executeParallel(
ctx context.Context,
tracer trace.Tracer,
agent *Agent,
toolCalls []llm.ToolCall,
tools []Tool,
onEvent func(context.Context, StreamEvent),
logger *log.Logger,
) ([]ToolCallResult, []llm.Message, error) {
entries := make([]parallelToolEntry, len(toolCalls))
var wg sync.WaitGroup
wg.Add(len(toolCalls))
for i := range toolCalls {
go func(idx int, tc llm.ToolCall, tool Tool) {
defer wg.Done()
tr, err := executeSingleTool(ctx, tracer, agent, tc, tool, onEvent, logger)
if err != nil {
entries[idx] = parallelToolEntry{err: err}
return
}
entries[idx] = parallelToolEntry{result: tr}
}(i, toolCalls[i], tools[i])
}
wg.Wait()
for i, entry := range entries {
var ie *InterruptedError
if entry.err != nil && errors.As(entry.err, &ie) {
var completed []completedCall
for j, other := range entries {
if j == i {
continue
}
if other.err != nil {
completed = append(
completed,
completedCall{
toolCallID: toolCalls[j].ID,
result: ToolResult{
Content: fmt.Sprintf("Error: %s", other.err.Error()),
IsError: true,
},
},
)
continue
}
completed = append(
completed,
completedCall{
toolCallID: toolCalls[j].ID,
result: other.result,
},
)
}
return nil, nil, &nestedInterruptionError{
inner: ie,
toolCallID: toolCalls[i].ID,
allToolCalls: toolCalls,
completedCalls: completed,
}
}
}
var (
results []ToolCallResult
msgs []llm.Message
)
for i, tc := range toolCalls {
entry := entries[i]
if entry.err != nil {
msgs = append(
msgs,
llm.Message{
Role: llm.RoleTool,
ToolCallID: tc.ID,
Parts: []llm.Part{
llm.TextPart{
Text: fmt.Sprintf("Error: %s", entry.err.Error()),
},
},
},
)
results = append(
results,
ToolCallResult{
ToolName: tc.Function.Name,
Arguments: tc.Function.Arguments,
Result: ToolResult{
Content: fmt.Sprintf("Error: %s", entry.err.Error()),
IsError: true,
},
},
)
continue
}
msgs = append(
msgs,
llm.Message{
Role: llm.RoleTool,
ToolCallID: tc.ID,
Parts: []llm.Part{llm.TextPart{Text: entry.result.Content}},
},
)
results = append(
results,
ToolCallResult{
ToolName: tc.Function.Name,
Arguments: tc.Function.Arguments,
Result: entry.result,
},
)
}
return results, msgs, nil
}
func executeSingleTool(
ctx context.Context,
tracer trace.Tracer,
agent *Agent,
tc llm.ToolCall,
tool Tool,
onEvent func(context.Context, StreamEvent),
logger *log.Logger,
) (ToolResult, error) {
onEvent(ctx, StreamEvent{Type: StreamEventToolStart, Agent: agent, Tool: tool})
emitHook(agent, func(h RunHooks) { h.OnToolStart(ctx, agent, tool, tc.Function.Arguments) })
emitAgentHook(agent, func(h AgentHooks) { h.OnToolStart(ctx, agent, tool) })
toolCtx, toolSpan := tracer.Start(
ctx,
fmt.Sprintf("agent.tool %s", tool.Name()),
trace.WithSpanKind(trace.SpanKindInternal),
trace.WithAttributes(
attribute.String("tool.name", tool.Name()),
),
)
logger.InfoCtx(
ctx,
"executing tool",
log.String("tool", tool.Name()),
)
result, err := tool.Execute(toolCtx, tc.Function.Arguments)
if err != nil {
if _, ok := errors.AsType[*InterruptedError](err); ok {
toolSpan.SetAttributes(attribute.Bool("tool.interrupted", true))
toolSpan.End()
onEvent(ctx, StreamEvent{Type: StreamEventToolEnd, Agent: agent, Tool: tool, Err: err})
emitHook(agent, func(h RunHooks) { h.OnToolEnd(ctx, agent, tool, ToolResult{}, err) })
emitAgentHook(agent, func(h AgentHooks) { h.OnToolEnd(ctx, agent, tool, ToolResult{}) })
return ToolResult{}, err
}
toolSpan.RecordError(err)
toolSpan.SetStatus(codes.Error, err.Error())
toolSpan.End()
onEvent(ctx, StreamEvent{Type: StreamEventToolEnd, Agent: agent, Tool: tool, Err: err})
emitHook(agent, func(h RunHooks) { h.OnToolEnd(ctx, agent, tool, result, err) })
emitAgentHook(agent, func(h AgentHooks) { h.OnToolEnd(ctx, agent, tool, result) })
logger.ErrorCtx(
ctx,
"tool execution failed",
log.String("tool", tool.Name()),
log.Error(err),
)
return ToolResult{}, fmt.Errorf("cannot execute tool %q: %w", tool.Name(), err)
}
toolSpan.SetAttributes(attribute.Bool("tool.is_error", result.IsError))
toolSpan.End()
onEvent(ctx, StreamEvent{Type: StreamEventToolEnd, Agent: agent, Tool: tool, ToolResult: &result})
emitHook(agent, func(h RunHooks) { h.OnToolEnd(ctx, agent, tool, result, nil) })
emitAgentHook(agent, func(h AgentHooks) { h.OnToolEnd(ctx, agent, tool, result) })
logger.InfoCtx(
ctx,
"tool execution completed",
log.String("tool", tool.Name()),
log.Bool("is_error", result.IsError),
)
return result, nil
}
func checkApproval(ctx context.Context, a *Agent, toolCalls []llm.ToolCall) error {
if a.approval == nil {
return nil
}
var pending []llm.ToolCall
for _, tc := range toolCalls {
if a.approval.requiresApproval(ctx, tc) {
pending = append(pending, tc)
}
}
if len(pending) > 0 {
return &needsApprovalError{
allToolCalls: toolCalls,
pendingApprovals: pending,
}
}
return nil
}
func runInputGuardrails(ctx context.Context, agent *Agent, messages []llm.Message) error {
for _, g := range agent.inputGuardrails {
result, err := g.Check(ctx, messages)
if err != nil {
return fmt.Errorf("cannot run input guardrail %q: %w", g.Name(), err)
}
if result != nil && result.Tripwire {
emitHook(agent, func(h RunHooks) { h.OnGuardrailTripped(ctx, agent, g.Name(), result) })
return &InputGuardrailTrippedError{
Guardrail: g.Name(),
Message: result.Message,
}
}
}
return nil
}
func runOutputGuardrails(ctx context.Context, agent *Agent, message llm.Message) error {
for _, g := range agent.outputGuardrails {
result, err := g.Check(ctx, message)
if err != nil {
return fmt.Errorf("cannot run output guardrail %q: %w", g.Name(), err)
}
if result != nil && result.Tripwire {
emitHook(agent, func(h RunHooks) { h.OnGuardrailTripped(ctx, agent, g.Name(), result) })
return &OutputGuardrailTrippedError{
Guardrail: g.Name(),
Message: result.Message,
}
}
}
return nil
}
// Resume continues an interrupted run after human approval decisions have
// been collected. It executes or denies each pending tool call according to
// the provided ResumeInput, then re-enters the agent loop. Input guardrails
// are not re-evaluated because the messages were already validated in the
// original Run call.
func Resume(ctx context.Context, interrupted *InterruptedError, input ResumeInput) (*Result, error) {
if interrupted.outerState != nil {
return resumeNested(ctx, interrupted, input)
}
tracer := otel.GetTracerProvider().Tracer(tracerName)
logger := interrupted.Agent.logger
agent := interrupted.Agent
messages := interrupted.Messages
logger.InfoCtx(
ctx,
"resuming interrupted run",
log.String("agent", agent.name),
log.Int("pending_approvals", len(interrupted.PendingApprovals)),
log.Int("total_tool_calls", len(interrupted.ToolCalls)),
)
_, toolMap, err := agent.resolveTools(ctx)
if err != nil {
return nil, fmt.Errorf("cannot resolve tools for resume: %w", err)
}
pendingSet := make(map[string]struct{}, len(interrupted.PendingApprovals))
for _, tc := range interrupted.PendingApprovals {
pendingSet[tc.ID] = struct{}{}
}
var handoffTarget *Handoff
for _, tc := range interrupted.ToolCalls {
if _, needsApproval := pendingSet[tc.ID]; needsApproval {
approval, ok := input.Approvals[tc.ID]
if !ok || !approval.Approved {
reason := "Tool call was denied by human review."
if ok && approval.Message != "" {
reason = approval.Message
}
logger.InfoCtx(
ctx,
"tool call denied",
log.String("tool", tc.Function.Name),
log.String("tool_call_id", tc.ID),
log.String("reason", reason),
)
messages = append(
messages,
llm.Message{
Role: llm.RoleTool,
ToolCallID: tc.ID,
Parts: []llm.Part{
llm.TextPart{Text: reason},
},
},
)
continue
}
logger.InfoCtx(
ctx,
"tool call approved",
log.String("tool", tc.Function.Name),
log.String("tool_call_id", tc.ID),
)
}
desc, toolOK := toolMap[tc.Function.Name]
if !toolOK {
return nil, fmt.Errorf("unknown tool %q", tc.Function.Name)
}
if ht, ok := desc.(*handoffToolAdapter); ok {
if ht.handoff.OnHandoff != nil {
if err := ht.handoff.OnHandoff(ctx); err != nil {
return nil, fmt.Errorf("cannot execute handoff callback for %q: %w", ht.handoff.Agent.name, err)
}
}
messages = append(
messages,
llm.Message{
Role: llm.RoleTool,
ToolCallID: tc.ID,
Parts: []llm.Part{
llm.TextPart{
Text: fmt.Sprintf("Transferred to %s", ht.handoff.Agent.name),
},
},
},
)
handoffTarget = ht.handoff
break
}
tr, execErr := executeSingleTool(ctx, tracer, agent, tc, desc.(Tool), noopEvent, logger)
if execErr != nil {
return nil, execErr
}
messages = append(
messages,
llm.Message{
Role: llm.RoleTool,
ToolCallID: tc.ID,
Parts: []llm.Part{
llm.TextPart{
Text: tr.Content,
},
},
},
)
}
resumeAgent := agent
if handoffTarget != nil {
logger.InfoCtx(
ctx,
"resume handoff",
log.String("from", agent.name),
log.String("to", handoffTarget.Agent.name),
)
emitHook(agent, func(h RunHooks) { h.OnHandoff(ctx, agent, handoffTarget.Agent) })
emitAgentHook(handoffTarget.Agent, func(h AgentHooks) { h.OnHandoff(ctx, handoffTarget.Agent, agent) })
if handoffTarget.InputFilter != nil {
filtered := handoffTarget.InputFilter(
HandoffInputData{
InputHistory: interrupted.Messages,
NewItems: messages[len(interrupted.Messages):],
},
)
messages = filtered
}
resumeAgent = handoffTarget.Agent
}
return coreLoop(
ctx,
resumeAgent,
messages,
runOpts{
callLLM: blockingCallLLM,
onEvent: noopEvent,
skipInputGuardrails: true,
skipSessionLoad: true,
initialUsage: interrupted.Usage,
initialTurns: interrupted.Turns,
},
)
}
func resumeNested(ctx context.Context, interrupted *InterruptedError, input ResumeInput) (*Result, error) {
outer := interrupted.outerState
logger := outer.agent.logger
logger.InfoCtx(
ctx,
"resuming nested agent interruption",
log.String("outer_agent", outer.agent.name),
log.String("inner_agent", interrupted.Agent.name),
)
innerResult, err := Resume(ctx, outer.innerInterrupt, input)
if err != nil {
var innerIE *InterruptedError
if errors.As(err, &innerIE) {
return nil, &InterruptedError{
ToolCalls: innerIE.ToolCalls,
PendingApprovals: innerIE.PendingApprovals,
Agent: innerIE.Agent,
Messages: innerIE.Messages,
Usage: innerIE.Usage,
Turns: innerIE.Turns,
outerState: &outerLoopState{
agent: outer.agent,
messages: outer.messages,
usage: outer.usage,
turns: outer.turns,
allToolCalls: outer.allToolCalls,
toolCallID: outer.toolCallID,
completedCalls: outer.completedCalls,
innerInterrupt: innerIE,
},
}
}
return nil, fmt.Errorf("cannot resume nested agent: %w", err)
}
completedMap := make(map[string]ToolResult, len(outer.completedCalls))
for _, cc := range outer.completedCalls {
completedMap[cc.toolCallID] = cc.result
}
messages := make([]llm.Message, len(outer.messages))
copy(messages, outer.messages)
for _, tc := range outer.allToolCalls {
var content string
if tc.ID == outer.toolCallID {
content = innerResult.FinalMessage().Text()
} else if cr, ok := completedMap[tc.ID]; ok {
content = cr.Content
} else {
content = "Error: tool execution was interrupted"
}
messages = append(
messages,
llm.Message{
Role: llm.RoleTool,
ToolCallID: tc.ID,
Parts: []llm.Part{llm.TextPart{Text: content}},
},
)
}
return coreLoop(
ctx,
outer.agent,
messages,
runOpts{
callLLM: blockingCallLLM,
onEvent: noopEvent,
skipInputGuardrails: true,
skipSessionLoad: true,
initialUsage: outer.usage,
initialTurns: outer.turns,
},
)
}
func emitHook(agent *Agent, fn func(RunHooks)) {
for _, h := range agent.hooks {
fn(h)
}
}
func emitAgentHook(agent *Agent, fn func(AgentHooks)) {
if agent.agentHooks != nil {
fn(agent.agentHooks)
}
}