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
105 changes: 85 additions & 20 deletions agent/agent.go
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@ package agent
import (
"context"
"errors"
"math"
"slices"
"time"

Expand All @@ -25,7 +26,10 @@ const (
// run-scoped state lives in the [Session] passed to each run.
//
// Resilience is layered at the model, not the agent: wrap the model with
// [ai.Chain] and the ai middleware packages before passing it in.
// [ai.Chain] and the ai middleware packages before passing it in. The one
// exception is [WithStreamRecovery]: only the loop can retract output a
// frontend has already rendered, so a turn whose stream broke after producing
// output is re-issued here rather than inside the middleware.
type Agent struct {
model ai.LanguageModel
tools *toolbox
Expand All @@ -46,25 +50,70 @@ const (
)

type config struct {
name string
system string
tools []Tool
maxTurns int
maxTokens int
toolTimeout time.Duration
parallelTools int
steeringMode QueueMode
followUpMode QueueMode
stopWhen func(RunInfo) bool
beforeTool gate
afterTool func(context.Context, ToolResultInfo) *ToolResultOverride
prepareTurn func(context.Context, RunInfo) TurnUpdate
candidate func(context.Context, CandidateAnswerInfo) CandidateAnswerDecision
transform func(context.Context, []ai.Message) ([]ai.Message, error)
onEvent func(context.Context, Event)
requestFn func(*ai.Request)
inputGuards []inputGuardrail
outputGuards []outputGuardrail
name string
system string
tools []Tool
maxTurns int
maxTokens int
toolTimeout time.Duration
parallelTools int
steeringMode QueueMode
followUpMode QueueMode
stopWhen func(RunInfo) bool
streamRecovery streamRecovery
beforeTool gate
afterTool func(context.Context, ToolResultInfo) *ToolResultOverride
prepareTurn func(context.Context, RunInfo) TurnUpdate
candidate func(context.Context, CandidateAnswerInfo) CandidateAnswerDecision
transform func(context.Context, []ai.Message) ([]ai.Message, error)
onEvent func(context.Context, Event)
requestFn func(*ai.Request)
inputGuards []inputGuardrail
outputGuards []outputGuardrail
}

// Stream recovery limits. The backoff doubles per re-issue up to the ceiling,
// which mirrors the retry middleware's shape one layer up.
const (
// DefaultStreamRecoveryBase is the backoff before the first re-issue.
DefaultStreamRecoveryBase = time.Second
// DefaultStreamRecoveryMax caps one backoff delay.
DefaultStreamRecoveryMax = 30 * time.Second
)

// streamRecovery bounds the re-issue of a turn whose model stream failed after
// it had already streamed output.
type streamRecovery struct {
attempts int
base time.Duration
max time.Duration
}

func newStreamRecovery(attempts int, base, ceiling time.Duration) streamRecovery {
if attempts <= 0 {
return streamRecovery{}
}

recovery := streamRecovery{attempts: attempts, base: base, max: ceiling}
if recovery.base <= 0 {
recovery.base = DefaultStreamRecoveryBase
}

if recovery.max < recovery.base {
recovery.max = DefaultStreamRecoveryMax
}

return recovery
}

// backoff returns the delay before re-issue number attempt (zero-based).
func (r streamRecovery) backoff(attempt int) time.Duration {
delay := float64(r.base) * math.Pow(2, float64(attempt))
if delay > float64(r.max) {
delay = float64(r.max)
}

return time.Duration(delay)
}

// Option configures an [Agent].
Expand Down Expand Up @@ -119,6 +168,22 @@ func WithStopWhen(cond func(RunInfo) bool) Option {
return func(c *config) { c.stopWhen = cond }
}

// WithStreamRecovery re-issues a turn whose model stream failed after it had
// already streamed output, up to attempts extra tries with exponential backoff
// (base, capped by maxDelay; zero values select [DefaultStreamRecoveryBase] and
// [DefaultStreamRecoveryMax]). Zero attempts — the default — surfaces the
// failure instead.
//
// The middleware owns failures that produced nothing; this option covers the
// ones that arrived after output, where replaying from the middleware would
// duplicate what the frontend already rendered. Before each re-issue the loop
// emits [CandidateDiscarded] so a consumer drops the partial output, which
// makes the retry safe: nothing from the failed attempt was committed and no
// tool ran.
func WithStreamRecovery(attempts int, base, maxDelay time.Duration) Option {
return func(c *config) { c.streamRecovery = newStreamRecovery(attempts, base, maxDelay) }
}

// WithBeforeTool installs a gate consulted before each tool call executes.
// Gates run serially in call order on the run's goroutine. See
// [ToolDecisionAction] for the available verdicts.
Expand Down
15 changes: 11 additions & 4 deletions agent/event_clone.go
Original file line number Diff line number Diff line change
Expand Up @@ -20,7 +20,7 @@ func cloneEvent(event Event) Event {
case RunStarted, TurnStarted, CandidateDiscarded, TurnCompleted, RunCompleted:
return event
case ModelStreamEvent:
if payload.Event.Usage == nil {
if payload.Event.Usage == nil && payload.Event.Retry == nil {
return event
}
}
Expand All @@ -33,12 +33,19 @@ func cloneEvent(event Event) Event {
func cloneEventPayload(payload EventPayload) EventPayload {
switch value := payload.(type) {
case ModelStreamEvent:
if value.Event.Usage == nil {
if value.Event.Usage == nil && value.Event.Retry == nil {
return payload
}

usage := *value.Event.Usage
value.Event.Usage = &usage
if value.Event.Usage != nil {
usage := *value.Event.Usage
value.Event.Usage = &usage
}

if value.Event.Retry != nil {
notice := *value.Event.Retry
value.Event.Retry = &notice
}

return value
case MessageCommitted:
Expand Down
1 change: 1 addition & 0 deletions agent/event_validate.go
Original file line number Diff line number Diff line change
Expand Up @@ -143,6 +143,7 @@ func validModelStreamEventType(eventType ai.StreamEventType) bool {
ai.StreamToolCallStart,
ai.StreamToolCallDelta,
ai.StreamToolCallEnd,
ai.StreamRetry,
ai.StreamMessageEnd:
return true
default:
Expand Down
93 changes: 84 additions & 9 deletions agent/run.go
Original file line number Diff line number Diff line change
Expand Up @@ -176,9 +176,7 @@ func (r *run) turn(ctx context.Context, turn int) (result *RunResult, next bool,

previousResponse := r.lastResp

resp, stopped, err := r.agent.callModel(
ctx, r.model, turnTools, requestUpdate, turn, msgs, r.emit, r.streaming,
)
resp, stopped, err := r.callTurn(ctx, turn, turnTools, requestUpdate, msgs)
if stopped {
return nil, false, nil
}
Expand Down Expand Up @@ -556,9 +554,81 @@ func (r *run) finish(stop StopReason, turns int, pending []ai.ToolCallPart) (*Ru
}, nil
}

// callTurn performs one turn's model call, re-issuing it when a stream that had
// already produced output failed with a retryable error. Every re-issue first
// discards the provisional candidate so a frontend drops the partial output.
// Nothing from the failed attempt reached the session and no tool ran, so a
// re-issue can neither duplicate content nor repeat an effect.
func (r *run) callTurn(
ctx context.Context,
turn int,
turnTools *toolbox,
requestUpdate *runModelRequest,
msgs []ai.Message,
) (resp *ai.Response, stopped bool, err error) {
recovery := r.agent.cfg.streamRecovery

for attempt := 0; ; attempt++ {
var produced bool

resp, produced, stopped, err = r.agent.callModel(
ctx, r.model, turnTools, requestUpdate, turn, msgs, r.emit, r.streaming,
)
if stopped || err == nil {
return resp, stopped, err
}

if !produced || attempt >= recovery.attempts || !ai.IsRetryable(err) || ctx.Err() != nil {
return resp, false, err
}

delay := recovery.backoff(attempt)
if !r.emit(CandidateDiscarded{Turn: turn}) ||
!r.emit(retryNotice(turn, attempt, recovery.attempts, delay, err)) {
return nil, true, nil
}

if sleepErr := sleepContext(ctx, delay); sleepErr != nil {
return nil, false, sleepErr
}
}
}

// retryNotice announces the re-issue that is about to start. The loop reports
// the wait on the same stream channel the retry middleware uses, so a frontend
// renders one kind of notice whichever layer replays the request. Attempt is
// the 1-based ordinal of this re-issue against the loop's re-issue budget.
func retryNotice(turn, attempt, recoveryAttempts int, delay time.Duration, err error) ModelStreamEvent {
return ModelStreamEvent{
Turn: turn,
Event: ai.StreamEvent{
Type: ai.StreamRetry,
Retry: ai.NewRetryNotice(attempt+1, recoveryAttempts, delay, err),
},
}
}

// sleepContext waits for delay or the context, whichever comes first.
func sleepContext(ctx context.Context, delay time.Duration) error {
if delay <= 0 {
return ctx.Err()
}

timer := time.NewTimer(delay)
defer timer.Stop()

select {
case <-ctx.Done():
return ctx.Err()
case <-timer.C:
return nil
}
}

// callModel performs one model call. When streaming, deltas tee through emit
// while [ai.Collect] folds them into the completed response; stopped reports
// that the consumer quit mid-stream.
// while [ai.Collect] folds them into the completed response; produced reports
// that the call streamed output before it ended, and stopped reports that the
// consumer quit mid-stream.
func (a *Agent) callModel(
ctx context.Context,
model ai.LanguageModel,
Expand All @@ -568,12 +638,12 @@ func (a *Agent) callModel(
msgs []ai.Message,
emit emitFunc,
streaming bool,
) (resp *ai.Response, stopped bool, err error) {
) (resp *ai.Response, produced, stopped bool, err error) {
req := a.requestWithTools(msgs, tools, update)

if !streaming {
resp, err = model.Generate(ctx, req)
return resp, false, err
return resp, false, false, err
}

resp, err = ai.Collect(func(yield func(ai.StreamEvent, error) bool) {
Expand All @@ -588,6 +658,11 @@ func (a *Agent) callModel(
return
}

// A retry notice reports a wait, so it is not output.
if ev.Type != ai.StreamMessageStart && ev.Type != ai.StreamRetry {
produced = true
}

if !emit(ModelStreamEvent{Turn: turn, Event: ev}) {
stopped = true
return
Expand All @@ -600,10 +675,10 @@ func (a *Agent) callModel(
})

if stopped {
return nil, true, nil
return nil, produced, true, nil
}

return resp, false, err
return resp, produced, false, err
}

// shouldStop checks the configured stop conditions after a completed turn.
Expand Down
Loading
Loading