978
pkg/agent/run.go
Normal file
978
pkg/agent/run.go
Normal file
@@ -0,0 +1,978 @@
|
||||
// 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,
|
||||
},
|
||||
)
|
||||
}
|
||||
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 {
|
||||
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)}},
|
||||
},
|
||||
)
|
||||
|
||||
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()
|
||||
|
||||
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 {
|
||||
toolSpan.RecordError(err)
|
||||
toolSpan.SetStatus(codes.Error, err.Error())
|
||||
toolSpan.End()
|
||||
|
||||
emitHook(agent, func(h RunHooks) { h.OnToolEnd(ctx, agent, tool, result, err) })
|
||||
|
||||
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) {
|
||||
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 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)
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user