diff --git a/agent/agent.go b/agent/agent.go index bca9e444..abb913ef 100644 --- a/agent/agent.go +++ b/agent/agent.go @@ -3,6 +3,7 @@ package agent import ( "context" "errors" + "math" "slices" "time" @@ -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 @@ -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]. @@ -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. diff --git a/agent/event_clone.go b/agent/event_clone.go index adb95f29..21016bb4 100644 --- a/agent/event_clone.go +++ b/agent/event_clone.go @@ -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 } } @@ -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 = ¬ice + } return value case MessageCommitted: diff --git a/agent/event_validate.go b/agent/event_validate.go index c45ed7fb..aaaf92b5 100644 --- a/agent/event_validate.go +++ b/agent/event_validate.go @@ -143,6 +143,7 @@ func validModelStreamEventType(eventType ai.StreamEventType) bool { ai.StreamToolCallStart, ai.StreamToolCallDelta, ai.StreamToolCallEnd, + ai.StreamRetry, ai.StreamMessageEnd: return true default: diff --git a/agent/run.go b/agent/run.go index a85f74a7..b239c0a9 100644 --- a/agent/run.go +++ b/agent/run.go @@ -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 } @@ -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, @@ -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) { @@ -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 @@ -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. diff --git a/agent/stream_recovery_test.go b/agent/stream_recovery_test.go new file mode 100644 index 00000000..7378c30d --- /dev/null +++ b/agent/stream_recovery_test.go @@ -0,0 +1,225 @@ +package agent_test + +import ( + "context" + "errors" + "io" + "sync/atomic" + "testing" + "time" + + "github.com/rsbin1178/pips/agent" + "github.com/rsbin1178/pips/ai" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +// interruptedStreamModel streams partial output and then fails for its first +// failures attempts, so a re-issue can be observed to succeed on a later one. +type interruptedStreamModel struct { + failures int32 + silent bool // fail before any delta, leaving produced false + err error + calls atomic.Int32 +} + +func (m *interruptedStreamModel) Generate(context.Context, ai.Request) (*ai.Response, error) { + return nil, errors.New("interruptedStreamModel: Generate is not scripted") +} + +func (m *interruptedStreamModel) Stream(_ context.Context, _ ai.Request) ai.Stream { + attempt := m.calls.Add(1) + + return func(yield func(ai.StreamEvent, error) bool) { + if !yield(ai.StreamEvent{Type: ai.StreamMessageStart, ID: "resp"}, nil) { + return + } + + if attempt <= m.failures { + if !m.silent && !yield(ai.StreamEvent{Type: ai.StreamTextDelta, Text: "partial"}, nil) { + return + } + + yield(ai.StreamEvent{}, m.err) + + return + } + + yield(ai.StreamEvent{Type: ai.StreamTextDelta, Text: "final"}, nil) + yield(ai.StreamEvent{Type: ai.StreamMessageEnd, FinishReason: ai.FinishStop}, nil) + } +} + +func (m *interruptedStreamModel) Provider() ai.Provider { return ai.Provider("interrupted") } +func (m *interruptedStreamModel) ModelID() string { return "interrupted-1" } +func (m *interruptedStreamModel) Capabilities() ai.Capabilities { return ai.Capabilities{Text: true} } + +// collectStream drains one agent stream, returning the non-delta event types and +// the last retry notice seen. +func collectStream(t *testing.T, a *agent.Agent, sess *agent.Session) ([]agent.EventType, *ai.RetryNotice, error) { + t.Helper() + + var ( + types []agent.EventType + notice *ai.RetryNotice + ) + + for ev, err := range a.Stream(t.Context(), sess, ai.UserText("hi")) { + if err != nil { + return types, notice, err + } + + require.NoError(t, ev.Validate()) + + switch payload := ev.Payload().(type) { + case agent.ModelStreamEvent: + if payload.Event.Type == ai.StreamRetry { + notice = payload.Event.Retry + } + + continue + case agent.CandidateDiscarded: + types = append(types, ev.Type()) + + continue + } + + types = append(types, ev.Type()) + } + + return types, notice, nil +} + +func TestStreamRecoveryReissuesTurnAfterPartialOutput(t *testing.T) { + t.Parallel() + + model := &interruptedStreamModel{failures: 1, err: io.ErrUnexpectedEOF} + + a, err := agent.New(model, agent.WithStreamRecovery(1, time.Millisecond, time.Millisecond)) + require.NoError(t, err) + + sess := agent.NewSession() + + types, notice, err := collectStream(t, a, sess) + require.NoError(t, err) + + // The partial candidate is retracted before the re-issue, and only the + // successful attempt reaches the session. + assert.Equal(t, []agent.EventType{ + agent.EventRunStarted, + agent.EventTurnStarted, + agent.EventCandidateDiscarded, + agent.EventMessageCommitted, + agent.EventTurnCompleted, + agent.EventRunCompleted, + }, types) + + require.NotNil(t, notice) + assert.Equal(t, 1, notice.Attempt, "the notice counts retries, not tries") + assert.Equal(t, 1, notice.MaxRetries) + assert.Equal(t, time.Millisecond, notice.Delay) + assert.Equal(t, "stream ended early", notice.Reason) + + assert.Equal(t, int32(2), model.calls.Load()) + + msgs := sess.Messages() + require.Len(t, msgs, 2) + + assistant, ok := msgs[1].(ai.AssistantMessage) + require.True(t, ok) + text, ok := assistant.Parts[0].(ai.TextPart) + require.True(t, ok) + assert.Equal(t, "final", text.Text) +} + +func TestStreamRecoveryIsOffByDefault(t *testing.T) { + t.Parallel() + + model := &interruptedStreamModel{failures: 1, err: io.ErrUnexpectedEOF} + + a, err := agent.New(model) + require.NoError(t, err) + + _, notice, err := collectStream(t, a, agent.NewSession()) + require.ErrorIs(t, err, io.ErrUnexpectedEOF) + assert.Nil(t, notice) + assert.Equal(t, int32(1), model.calls.Load()) +} + +func TestStreamRecoveryGivesUpAfterBoundedAttempts(t *testing.T) { + t.Parallel() + + model := &interruptedStreamModel{failures: 5, err: io.ErrUnexpectedEOF} + + a, err := agent.New(model, agent.WithStreamRecovery(2, time.Millisecond, time.Millisecond)) + require.NoError(t, err) + + _, notice, err := collectStream(t, a, agent.NewSession()) + require.ErrorIs(t, err, io.ErrUnexpectedEOF) + require.NotNil(t, notice) + assert.Equal(t, 2, notice.Attempt) + assert.Equal(t, 2, notice.MaxRetries) + assert.Equal(t, int32(3), model.calls.Load(), "one try plus two re-issues") +} + +// TestStreamRecoveryBackoffGrowsThenStops pins the wait schedule: every +// re-issue waits longer than the one before it, up to the ceiling. +func TestStreamRecoveryBackoffGrowsThenStops(t *testing.T) { + t.Parallel() + + model := &interruptedStreamModel{failures: 4, err: io.ErrUnexpectedEOF} + + a, err := agent.New(model, agent.WithStreamRecovery(3, time.Millisecond, 3*time.Millisecond)) + require.NoError(t, err) + + var delays []time.Duration + + for ev, streamErr := range a.Stream(t.Context(), agent.NewSession(), ai.UserText("hi")) { + if streamErr != nil { + break + } + + if payload, ok := ev.Payload().(agent.ModelStreamEvent); ok && payload.Event.Retry != nil { + delays = append(delays, payload.Event.Retry.Delay) + } + } + + assert.Equal(t, + []time.Duration{time.Millisecond, 2 * time.Millisecond, 3 * time.Millisecond}, + delays, + "the wait doubles per re-issue and stops at the ceiling", + ) +} + +func TestStreamRecoveryRequiresObservedOutput(t *testing.T) { + t.Parallel() + + model := &interruptedStreamModel{failures: 5, silent: true, err: io.ErrUnexpectedEOF} + + a, err := agent.New(model, agent.WithStreamRecovery(2, time.Millisecond, time.Millisecond)) + require.NoError(t, err) + + // A request that produced nothing belongs to the model middleware, which + // has its own budget; the loop must not stack a second budget on top. + _, notice, err := collectStream(t, a, agent.NewSession()) + require.ErrorIs(t, err, io.ErrUnexpectedEOF) + assert.Nil(t, notice) + assert.Equal(t, int32(1), model.calls.Load()) +} + +func TestStreamRecoverySkipsNonRetryableFailure(t *testing.T) { + t.Parallel() + + model := &interruptedStreamModel{ + failures: 5, + err: ai.NewError(ai.ProviderOpenAI, 400, "invalid request"), + } + + a, err := agent.New(model, agent.WithStreamRecovery(2, time.Millisecond, time.Millisecond)) + require.NoError(t, err) + + _, notice, err := collectStream(t, a, agent.NewSession()) + require.Error(t, err) + assert.Nil(t, notice) + assert.Equal(t, int32(1), model.calls.Load()) +} diff --git a/ai/errors.go b/ai/errors.go index bb9dbc4a..3aa7a245 100644 --- a/ai/errors.go +++ b/ai/errors.go @@ -2,8 +2,11 @@ package ai import ( "context" + "crypto/tls" + "crypto/x509" "errors" "fmt" + "io" "net/http" "time" ) @@ -133,6 +136,11 @@ func IsRetryable(err error) bool { case errors.Is(err, context.Canceled), errors.Is(err, context.DeadlineExceeded): // The caller's context is gone; retrying under it cannot succeed. return false + case isCertificateFailure(err): + // A certificate the client cannot verify is not transient: the caller + // has to fix it, so report it on the first attempt. Transient TLS + // conditions (handshake timeouts, resets) stay retryable. + return false case errors.Is(err, ErrRateLimited), errors.Is(err, ErrOverloaded): return true case errors.Is(err, ErrAuth), errors.Is(err, ErrInvalidRequest), errors.Is(err, ErrUnsupported): @@ -149,3 +157,62 @@ func IsRetryable(err error) bool { // such as connection resets or timeouts. return true } + +// isCertificateFailure reports whether err is a certificate validation failure +// rather than a transient TLS condition. +func isCertificateFailure(err error) bool { + var verification *tls.CertificateVerificationError + if errors.As(err, &verification) { + return true + } + + var ( + unknownAuthority x509.UnknownAuthorityError + hostname x509.HostnameError + invalid x509.CertificateInvalidError + ) + + return errors.As(err, &unknownAuthority) || + errors.As(err, &hostname) || + errors.As(err, &invalid) +} + +// NewRetryNotice builds the notice for the retry that is about to start after +// err. attempt is its 1-based ordinal and budget is the number of retries +// available. The reason is classified provider-neutrally, so the same failure +// renders the same way whoever re-issues the request: the retry middleware or +// an agent loop replaying a turn. +func NewRetryNotice(attempt, budget int, delay time.Duration, err error) *RetryNotice { + if attempt < 1 { + attempt = 1 + } + + if budget < attempt { + budget = attempt + } + + return &RetryNotice{ + Attempt: attempt, + MaxRetries: budget, + Delay: delay, + Reason: retryReason(err), + } +} + +// retryReason renders the failure class a frontend shows next to the wait. It +// stays a short phrase: the provider's own wording belongs to the terminal +// error, which still carries the full chain. +func retryReason(err error) string { + switch { + case errors.Is(err, ErrRateLimited): + return "rate limited" + case errors.Is(err, ErrOverloaded): + return "provider overloaded" + case errors.Is(err, io.ErrUnexpectedEOF), errors.Is(err, io.EOF): + return "stream ended early" + case errors.Is(err, context.DeadlineExceeded): + return "request timed out" + default: + return "connection error" + } +} diff --git a/ai/errors_test.go b/ai/errors_test.go index 183558e6..1957abb0 100644 --- a/ai/errors_test.go +++ b/ai/errors_test.go @@ -2,6 +2,8 @@ package ai_test import ( "context" + "crypto/tls" + "crypto/x509" "errors" "fmt" "net" @@ -111,6 +113,23 @@ func TestIsRetryable(t *testing.T) { }, {"plain transport error", errors.New("connection reset by peer"), true}, {"wrapped rate limit", fmt.Errorf("call: %w", ai.NewError(ai.ProviderOpenAI, 429, "x")), true}, + { + // A certificate the client cannot verify is the user's to fix, so + // it is reported on the first attempt instead of being replayed. + "certificate verification", + &url.Error{Op: "Post", URL: "https://example.invalid", Err: &tls.CertificateVerificationError{}}, + false, + }, + { + "unknown authority", + fmt.Errorf("dial: %w", x509.UnknownAuthorityError{}), + false, + }, + { + "certificate hostname mismatch", + x509.HostnameError{Certificate: &x509.Certificate{}, Host: "example.invalid"}, + false, + }, } for _, tc := range cases { t.Run(tc.name, func(t *testing.T) { diff --git a/ai/middleware/retry/retry.go b/ai/middleware/retry/retry.go index b5e7719f..38f25652 100644 --- a/ai/middleware/retry/retry.go +++ b/ai/middleware/retry/retry.go @@ -5,6 +5,7 @@ // Retry-After delay is honored when present. Streaming is retried only before // the first event is produced — once output has been observed, replaying the // request could duplicate content, so the original error is surfaced instead. +// Each streaming re-attempt is announced with one [ai.StreamRetry] event. package retry import ( @@ -122,6 +123,10 @@ func (m *model) Generate(ctx context.Context, req ai.Request) (*ai.Response, err // Stream retries only before any event is produced. Once the first event // arrives the stream is passed through verbatim. +// +// Every re-attempt is announced with one [ai.StreamRetry] event carrying the +// attempt that is about to start and the backoff before it, so a consumer can +// report the wait instead of looking stalled. func (m *model) Stream(ctx context.Context, req ai.Request) ai.Stream { return func(yield func(ai.StreamEvent, error) bool) { for attempt := range m.cfg.maxAttempts { @@ -137,7 +142,17 @@ func (m *model) Stream(ctx context.Context, req ai.Request) ai.Stream { return } - if sleepErr := m.cfg.sleep(ctx, m.backoff(attempt, err)); sleepErr != nil { + delay := m.backoff(attempt, err) + // attempt is zero-based, so this is retry attempt+1 of the + // maxAttempts-1 replays the configuration allows. + if !yield(ai.StreamEvent{ + Type: ai.StreamRetry, + Retry: ai.NewRetryNotice(attempt+1, m.cfg.maxAttempts-1, delay, err), + }, nil) { + return + } + + if sleepErr := m.cfg.sleep(ctx, delay); sleepErr != nil { yield(ai.StreamEvent{}, sleepErr) return } diff --git a/ai/middleware/retry/retry_test.go b/ai/middleware/retry/retry_test.go index f8183404..c6938fae 100644 --- a/ai/middleware/retry/retry_test.go +++ b/ai/middleware/retry/retry_test.go @@ -3,6 +3,7 @@ package retry_test import ( "context" "errors" + "io" "sync/atomic" "testing" "time" @@ -161,6 +162,66 @@ func TestStreamRetriesBeforeFirstEvent(t *testing.T) { assert.Equal(t, int32(2), base.calls.Load()) } +func TestStreamAnnouncesRetryBeforeReattempt(t *testing.T) { + t.Parallel() + + base := &scriptedModel{ + responses: []*ai.Response{nil, {Message: ai.AssistantText("hello")}}, + errs: []error{ai.NewError(ai.ProviderOpenAI, 503, "overloaded"), nil}, + } + + var delays []time.Duration + + model := retry.New( + retry.WithMaxAttempts(3), + retry.WithJitter(func() float64 { return 1 }), + noSleep(&delays), + )(base) + + var notices []ai.RetryNotice + + for ev, err := range model.Stream(t.Context(), ai.Request{}) { + require.NoError(t, err) + + if ev.Type == ai.StreamRetry { + require.NotNil(t, ev.Retry) + + notices = append(notices, *ev.Retry) + } + } + + require.Len(t, notices, 1) + assert.Equal(t, 1, notices[0].Attempt, "the notice counts retries, not tries") + assert.Equal(t, 2, notices[0].MaxRetries) + assert.Equal(t, "provider overloaded", notices[0].Reason) + require.Len(t, delays, 1) + assert.Equal(t, delays[0], notices[0].Delay, "the advertised wait is the wait performed") +} + +func TestStreamRetryReasonNamesEarlyStreamEnd(t *testing.T) { + t.Parallel() + + base := &scriptedModel{ + responses: []*ai.Response{nil, {Message: ai.AssistantText("hello")}}, + errs: []error{io.ErrUnexpectedEOF, nil}, + } + + var notices []ai.RetryNotice + + model := retry.New(retry.WithMaxAttempts(2), noSleep(new([]time.Duration)))(base) + + for ev, err := range model.Stream(t.Context(), ai.Request{}) { + require.NoError(t, err) + + if ev.Retry != nil { + notices = append(notices, *ev.Retry) + } + } + + require.Len(t, notices, 1) + assert.Equal(t, "stream ended early", notices[0].Reason) +} + // midStreamModel yields one event then fails, to prove no replay after output. type midStreamModel struct { calls atomic.Int32 diff --git a/ai/stream.go b/ai/stream.go index cc405908..89558e79 100644 --- a/ai/stream.go +++ b/ai/stream.go @@ -3,6 +3,7 @@ package ai import ( "iter" "strings" + "time" ) // Stream is a sequence of streaming events. Iterate with range; breaking out @@ -23,6 +24,8 @@ type StreamEventType string // Stream event types, in the order a well-formed stream produces them: // one message_start; any interleaving of text_delta, reasoning_delta, and // tool_call_start/tool_call_delta/tool_call_end groups; one message_end. +// A retry event may appear anywhere before message_end: it reports that the +// request is being re-attempted, and it is not model output. const ( // StreamMessageStart opens the response; it carries Provider, ID, and // Model. @@ -46,6 +49,11 @@ const ( // StreamMessageEnd closes the response; it carries FinishReason and, // when the provider reports it, Usage. It may also carry Grounding. StreamMessageEnd StreamEventType = "message_end" + // StreamRetry reports that a retryable failure ended the attempt in + // flight and a new attempt is starting. It carries [StreamEvent.Retry] + // and no model output; consumers keep whatever they already received and + // report the wait instead of treating the failure as terminal. + StreamRetry StreamEventType = "retry" ) // StreamEvent is one normalized increment of a streaming response. Only the @@ -87,12 +95,33 @@ type StreamEvent struct { Usage *Usage Grounding *GroundingMetadata Warnings []Warning + + // Retry is set on retry events and nil for every other type. + Retry *RetryNotice +} + +// RetryNotice describes an attempt that is about to start after a retryable +// failure. It is display metadata: a consumer may render it, ignore it, or +// store it, and none of that changes the response being assembled. +// +// The counters follow the convention the other agent frontends use: Attempt is +// the retry ordinal — 1 for the first replay — and MaxRetries is the retry +// budget, so "1/10" reads as the first of at most ten replays. +type RetryNotice struct { + // Attempt is the 1-based ordinal of the retry that is about to start. + Attempt int + // MaxRetries is the budget of retries available, so Attempt <= MaxRetries. + MaxRetries int + // Delay is the backoff before the next attempt starts. + Delay time.Duration + // Reason is a short, provider-neutral cause, safe to display as-is. + Reason string } // Collect drains a stream and assembles the complete [Response], preserving // part order (text, reasoning, and tool calls appear where they occurred). // On mid-stream failure it returns the partial response together with the -// error. +// error. Retry events are ignored: they report a wait, not output. func Collect(stream Stream) (*Response, error) { acc := newAccumulator() @@ -175,6 +204,9 @@ func (a *accumulator) add(ev StreamEvent) { if ev.Citation != nil { a.resp.Citations = append(a.resp.Citations, *ev.Citation) } + case StreamRetry: + // A retry notice carries no output and must not close an open text or + // reasoning part, so the accumulated response is unchanged. case StreamMessageEnd: a.finish(ev) } diff --git a/internal/coding/agent_draft.go b/internal/coding/agent_draft.go index e124f4af..871bda3a 100644 --- a/internal/coding/agent_draft.go +++ b/internal/coding/agent_draft.go @@ -413,6 +413,7 @@ func (r *Runtime) runAgentDraftProposal( agent.WithName(agentDraftAgentName), agent.WithMaxTurns(agentDraftAgentMaxTurns), agent.WithMaxTokens(agentDraftAgentMaxTokens), + streamRecoveryOption(), agent.WithParallelTools(1), agent.WithToolTimeout(r.opts.ToolTimeout), agent.WithStopWhen(func(agent.RunInfo) bool { return collector.called() }), diff --git a/internal/coding/child_control_scope.go b/internal/coding/child_control_scope.go index d9c657e9..83c3788e 100644 --- a/internal/coding/child_control_scope.go +++ b/internal/coding/child_control_scope.go @@ -384,6 +384,7 @@ func newChildControlScope( agent.WithName("subagent/"+scope.plan.Identity.ID), agent.WithMaxTurns(scope.plan.Limits.MaxTurns), agent.WithMaxTokens(scope.plan.Limits.MaxTokens), + streamRecoveryOption(), agent.WithParallelTools(1), agent.WithStopWhen(scope.guard.stopWhen), agent.WithToolTimeout(factory.toolTimeout), diff --git a/internal/coding/event.go b/internal/coding/event.go index ad40c0ee..34172295 100644 --- a/internal/coding/event.go +++ b/internal/coding/event.go @@ -71,6 +71,7 @@ const ( EventMessageCommitted EventType = "message.committed" EventMessageDelta EventType = "message.delta" EventMessageDiscarded EventType = "message.discarded" + EventModelRetry EventType = "model.retry" EventToolStarted EventType = "tool.started" EventToolUpdated EventType = "tool.updated" EventToolCompleted EventType = "tool.completed" @@ -314,6 +315,20 @@ type MessageDiscarded struct { Turn int `json:"turn"` } +// ModelRetry reports that a model request is being re-attempted after a +// retryable failure. It is progress, not output: nothing is added to the +// transcript, and the attempt in flight replaces the one that failed. Attempt +// is the 1-based ordinal of the retry about to start, MaxRetries is the retry +// budget available, DelayMillis is the backoff before it, and Reason is a +// short, content-free cause. +type ModelRetry struct { + Turn int `json:"turn,omitempty"` + Attempt int `json:"attempt"` + MaxRetries int `json:"max_retries"` + DelayMillis int64 `json:"delay_ms"` + Reason string `json:"reason,omitempty"` +} + // ToolCall is the bounded model request projected into tool lifecycle events. type ToolCall struct { ID string `json:"id"` @@ -603,6 +618,7 @@ func (TurnCompleted) eventPayload() {} func (MessageCommitted) eventPayload() {} func (MessageDelta) eventPayload() {} func (MessageDiscarded) eventPayload() {} +func (ModelRetry) eventPayload() {} func (ToolStarted) eventPayload() {} func (ToolUpdated) eventPayload() {} func (ToolCompleted) eventPayload() {} @@ -679,7 +695,7 @@ func validateEnvelopeIDs(event Event) error { return invalidEvent("%s requires only an interaction id", event.Type) } case EventRunStarted, EventRunCompleted, EventRunInterrupted, EventTurnStarted, EventTurnCompleted, - EventMessageCommitted, EventMessageDelta, EventMessageDiscarded, + EventMessageCommitted, EventMessageDelta, EventMessageDiscarded, EventModelRetry, EventToolStarted, EventToolUpdated, EventToolCompleted, EventSubagentCreated, EventSubagentStarted, EventSubagentProgress, EventSubagentCompleted, EventSubagentFailed, @@ -845,6 +861,14 @@ func validatePayload(eventType EventType, payload EventPayload) error { if value.Turn < 1 { return payloadInvalid(eventType, payload, errors.New("turn must be positive")) } + case ModelRetry: + if eventType != EventModelRetry { + return payloadMismatch(eventType, payload) + } + + if err := validateModelRetryPayload(value); err != nil { + return payloadInvalid(eventType, payload, err) + } case ToolStarted: if eventType != EventToolStarted { return payloadMismatch(eventType, payload) @@ -2079,6 +2103,26 @@ func validateMessageDelta(delta MessageDelta) error { return nil } +// validateModelRetryPayload bounds the retry notice: an attempt always names a +// reachable attempt number, and the reason stays a short display string that +// the disclosure projection can keep as-is. +func validateModelRetryPayload(retry ModelRetry) error { + switch { + case retry.Turn < 0: + return errors.New("retry turn must not be negative") + case retry.Attempt < 1: + return errors.New("retry attempt must be positive") + case retry.MaxRetries < retry.Attempt: + return errors.New("retry budget must cover the attempt") + case retry.DelayMillis < 0: + return errors.New("retry delay must not be negative") + case !validBoundedText(retry.Reason, maxDiagnosticMessage, true): + return errors.New("invalid retry reason") + } + + return nil +} + func validateToolCall(call ToolCall) error { if validateEventID("tool call id", call.ID, true) != nil || !validIdentifierText(call.Name, maxEventIDBytes, false) || len(call.Arguments) > maxEventTextBytes { diff --git a/internal/coding/event_clone.go b/internal/coding/event_clone.go index 626d4f57..7d41e169 100644 --- a/internal/coding/event_clone.go +++ b/internal/coding/event_clone.go @@ -63,6 +63,8 @@ func cloneEventPayload(payload EventPayload) EventPayload { return cloneMessageDelta(value) case MessageDiscarded: return value + case ModelRetry: + return value case ToolStarted: value.Call = cloneToolCall(value.Call) return value @@ -233,6 +235,34 @@ func toolCallFromAI(call ai.ToolCallPart) ToolCall { return ToolCall{ID: call.ID, Name: call.Name, Arguments: slices.Clone(call.Args)} } +// modelRetryFromAI projects one retry notice. The notice is progress, not +// output: it carries no content, so the projection only has to keep the retry +// arithmetic inside the range the event contract accepts. +func modelRetryFromAI(turn int, notice *ai.RetryNotice) ModelRetry { + retry := ModelRetry{Turn: turn} + + if notice != nil { + retry.Attempt = notice.Attempt + retry.MaxRetries = notice.MaxRetries + retry.DelayMillis = notice.Delay.Milliseconds() + retry.Reason = notice.Reason + } + + if retry.Attempt < 1 { + retry.Attempt = 1 + } + + if retry.MaxRetries < retry.Attempt { + retry.MaxRetries = retry.Attempt + } + + if retry.DelayMillis < 0 { + retry.DelayMillis = 0 + } + + return retry +} + func toolResultMessage(result ai.ToolResultPart) ai.ToolMessage { cloned, err := ai.CloneParts([]ai.Part{result}) if err != nil { diff --git a/internal/coding/event_codec.go b/internal/coding/event_codec.go index 752b2fb4..76ad9ca5 100644 --- a/internal/coding/event_codec.go +++ b/internal/coding/event_codec.go @@ -164,6 +164,8 @@ func decodeEventPayload(eventType EventType, data []byte) (EventPayload, error) return decodePayload[MessageDelta](data) case EventMessageDiscarded: return decodePayload[MessageDiscarded](data) + case EventModelRetry: + return decodePayload[ModelRetry](data) case EventToolStarted: return decodePayload[ToolStarted](data) case EventToolUpdated: diff --git a/internal/coding/event_writer.go b/internal/coding/event_writer.go index f8e567b0..eae89470 100644 --- a/internal/coding/event_writer.go +++ b/internal/coding/event_writer.go @@ -139,8 +139,7 @@ func (p *agentProjector) project(event agent.Event) (Event, error) { eventType = EventTurnStarted payload = TurnStarted{Turn: source.Turn} case agent.ModelStreamEvent: - eventType = EventMessageDelta - payload = messageDeltaFromAI(source.Event) + eventType, payload = modelStreamProjection(source.Turn, source.Event) case agent.MessageCommitted: eventType = EventMessageCommitted payload = MessageCommitted{Message: cloneMessage(source.Message)} @@ -194,3 +193,14 @@ func (p *agentProjector) project(event agent.Event) (Event, error) { payload, ) } + +// modelStreamProjection maps one normalized stream event to its Coding event. +// Deltas build the provisional draft; a retry notice reports progress, so it +// becomes its own event and never enters that draft. +func modelStreamProjection(turn int, event ai.StreamEvent) (EventType, EventPayload) { + if event.Type == ai.StreamRetry { + return EventModelRetry, modelRetryFromAI(turn, event.Retry) + } + + return EventMessageDelta, messageDeltaFromAI(event) +} diff --git a/internal/coding/model/retry.go b/internal/coding/model/retry.go index dc3af5dd..9fec427a 100644 --- a/internal/coding/model/retry.go +++ b/internal/coding/model/retry.go @@ -6,7 +6,13 @@ import ( "github.com/rsbin1178/pips/ai/middleware/retry" ) -const codingModelMaxRetries = 5 +// codingModelMaxRetries is the retry budget for one model request in a Coding +// run: ten replays of a failure that produced nothing, following the default +// the other agent frontends ship (Claude Code retries transient failures up to +// ten times; Grok Build's live retry state reports a budget in the same range). +// The backoff doubles from 500ms with full jitter, capped at 30s, and a +// provider Retry-After wins when it is longer. +const codingModelMaxRetries = 10 // withCodingModel wraps an adapter with the coding middleware chain: the // resolved capability declaration is applied on top of whatever the adapter diff --git a/internal/coding/model/retry_test.go b/internal/coding/model/retry_test.go index 70f883c7..64e9a624 100644 --- a/internal/coding/model/retry_test.go +++ b/internal/coding/model/retry_test.go @@ -3,6 +3,7 @@ package model import ( "context" "errors" + "fmt" "sync/atomic" "testing" "time" @@ -13,47 +14,47 @@ import ( "github.com/stretchr/testify/require" ) -func TestWithCodingRetryUsesFiveRetryBudget(t *testing.T) { +func TestWithCodingRetryUsesTenRetryBudget(t *testing.T) { t.Parallel() - t.Run("sixth attempt succeeds", func(t *testing.T) { + // The retry budget follows the default the other agent frontends ship; pin + // the number so a change is deliberate. + assert.Equal(t, 10, codingModelMaxRetries) + + transient := func(failures int) []error { + errs := make([]error, 0, failures) + for index := 1; index <= failures; index++ { + errs = append(errs, fmt.Errorf("transient %d", index)) + } + + return errs + } + + t.Run("tenth retry succeeds", func(t *testing.T) { t.Parallel() - base := newRetryModel( - errors.New("transient 1"), - errors.New("transient 2"), - errors.New("transient 3"), - errors.New("transient 4"), - errors.New("transient 5"), - nil, - ) + // One initial attempt plus ten replays: the eleventh call answers. + base := newRetryModel(append(transient(codingModelMaxRetries), nil)...) model := withCodingRetry(base, noRetrySleep()) response, err := ai.Collect(model.Stream(t.Context(), ai.Request{})) require.NoError(t, err) assert.Equal(t, "ok", response.Text()) - assert.Equal(t, int32(6), base.calls.Load()) + assert.Equal(t, int32(codingModelMaxRetries+1), base.calls.Load()) }) - t.Run("six failures exhaust budget", func(t *testing.T) { + t.Run("budget exhaustion surfaces the last failure", func(t *testing.T) { t.Parallel() finalErr := errors.New("final transport failure") - base := newRetryModel( - errors.New("transient 1"), - errors.New("transient 2"), - errors.New("transient 3"), - errors.New("transient 4"), - errors.New("transient 5"), - finalErr, - ) + base := newRetryModel(append(transient(codingModelMaxRetries), finalErr)...) model := withCodingRetry(base, noRetrySleep()) _, err := ai.Collect(model.Stream(t.Context(), ai.Request{})) require.ErrorIs(t, err, finalErr) - assert.Equal(t, int32(6), base.calls.Load()) + assert.Equal(t, int32(codingModelMaxRetries+1), base.calls.Load()) }) } diff --git a/internal/coding/reducer.go b/internal/coding/reducer.go index d73753e1..87a98bf9 100644 --- a/internal/coding/reducer.go +++ b/internal/coding/reducer.go @@ -82,6 +82,19 @@ type RunState struct { Usage TokenUsage `json:"usage"` } +// RetryState is the live retry notice for the run in flight: a model request +// that failed and is starting a new attempt. A frontend renders the wait from +// it instead of appearing stalled. Attempt is the 1-based retry ordinal and +// MaxRetries the budget it counts against; Deadline is when the next attempt +// starts, and it is zero when the producer reported no delay. +type RetryState struct { + Active bool `json:"active,omitempty"` + Attempt int `json:"attempt,omitempty"` + MaxRetries int `json:"max_retries,omitempty"` + Reason string `json:"reason,omitempty"` + Deadline time.Time `json:"deadline,omitzero"` +} + // CandidateIdentity binds provisional deltas and their eventual commit or // discard to one Agent run and turn. Text equality is never ownership. type CandidateIdentity struct { @@ -259,8 +272,11 @@ type State struct { LastError *RuntimeError `json:"last_error,omitempty"` Tree SessionTree `json:"tree"` Compaction CompactionState `json:"compaction"` - ContextTokens int `json:"context_tokens"` - Tasks tasklist.Snapshot `json:"tasks"` + // Retry is the live retry notice for the run in flight; the next event of + // any other kind clears it, because any of them ends the wait. + Retry RetryState `json:"retry,omitzero"` + ContextTokens int `json:"context_tokens"` + Tasks tasklist.Snapshot `json:"tasks"` activeRuns map[string]int openTurns map[string]int @@ -456,6 +472,12 @@ func (state *State) reduce(event Event) error { return protocolError("session changed from %q to %q", state.SessionID, event.SessionID) } + // A retry notice stands only until the next event: whatever follows ends + // the wait it announced. + if event.Type != EventModelRetry { + state.Retry = RetryState{} + } + if err := state.apply(event); err != nil { return err } @@ -472,6 +494,9 @@ func (state *State) reduce(event Event) error { func (state *State) reduceDeltas(events []Event) error { folded := make([]MessageDelta, 0, len(events)) + // The folded run is streaming output, which ends any announced wait. + state.Retry = RetryState{} + for _, event := range events { if event.Sequence != state.Sequence+1 { return protocolError("sequence %d follows %d", event.Sequence, state.Sequence) @@ -751,6 +776,22 @@ func (state *State) apply(event Event) error { } state.Draft = nil state.DraftCandidate = CandidateIdentity{} + case ModelRetry: + if _, err := state.activeRun(event.RunID); err != nil { + return err + } + + if state.openTurn(event.RunID) == 0 { + return protocolError("model retry emitted outside an active turn") + } + + state.Retry = RetryState{ + Active: true, + Attempt: payload.Attempt, + MaxRetries: payload.MaxRetries, + Reason: payload.Reason, + Deadline: event.Time.Add(time.Duration(payload.DelayMillis) * time.Millisecond), + } case ToolStarted: if _, err := state.activeRun(event.RunID); err != nil { return err diff --git a/internal/coding/runtime_context_budget.go b/internal/coding/runtime_context_budget.go index 205dc0bd..869badd7 100644 --- a/internal/coding/runtime_context_budget.go +++ b/internal/coding/runtime_context_budget.go @@ -65,7 +65,9 @@ func markObservedModelOutput(current *interaction, event agent.Event) { case agent.MessageCommitted, agent.ToolStarted: current.modelResponseObserved = true case agent.ModelStreamEvent: - if value.Event.Type != ai.StreamMessageStart { + // A retry notice reports a wait rather than model output, so it must + // not mark the run as having produced a response. + if value.Event.Type != ai.StreamMessageStart && value.Event.Type != ai.StreamRetry { current.modelResponseObserved = true } } diff --git a/internal/coding/runtime_interaction.go b/internal/coding/runtime_interaction.go index de424cfc..e33320af 100644 --- a/internal/coding/runtime_interaction.go +++ b/internal/coding/runtime_interaction.go @@ -977,6 +977,7 @@ func (r *Runtime) openInteraction( agentOptions := append( composed.AgentOptions(), agent.WithMaxTurns(maxTurns), + streamRecoveryOption(), agent.WithStopWhen(func(info agent.RunInfo) bool { return current.hookStopRequestedNow() || failureGuard.stopWhen(info) }), diff --git a/internal/coding/stream_recovery.go b/internal/coding/stream_recovery.go new file mode 100644 index 00000000..d83d4a9c --- /dev/null +++ b/internal/coding/stream_recovery.go @@ -0,0 +1,27 @@ +package coding + +import ( + "time" + + "github.com/rsbin1178/pips/agent" +) + +// Stream recovery bounds for every Coding Agent run. A turn whose model stream +// failed after it had already streamed output is re-issued with a growing +// backoff, mirroring the budget the model middleware spends on failures that +// produced nothing: ten re-issues (eleven calls at most for one turn) cover a +// three-minute outage, and a longer one is reported rather than replayed +// forever. Nothing from a failed attempt is committed, so a re-issue cannot +// duplicate content or repeat a tool call. +const ( + streamRecoveryAttempts = 10 + streamRecoveryBaseDelay = 2 * time.Second + streamRecoveryMaxDelay = 30 * time.Second +) + +// streamRecoveryOption is the shared recovery policy for the interaction loop, +// subagents, team workers, and the internal draft/proposal agents, so every +// Coding run survives an interrupted stream the same way. +func streamRecoveryOption() agent.Option { + return agent.WithStreamRecovery(streamRecoveryAttempts, streamRecoveryBaseDelay, streamRecoveryMaxDelay) +} diff --git a/internal/coding/stream_recovery_runtime_test.go b/internal/coding/stream_recovery_runtime_test.go new file mode 100644 index 00000000..a9b00113 --- /dev/null +++ b/internal/coding/stream_recovery_runtime_test.go @@ -0,0 +1,110 @@ +package coding + +import ( + "context" + "errors" + "io" + "sync" + "testing" + + "github.com/rsbin1178/pips/ai" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +// interruptedRuntimeModel streams a partial answer and then fails once, so the +// Coding recovery path can be observed end to end: the second attempt answers. +type interruptedRuntimeModel struct { + mu sync.Mutex + calls int + failure error + reply string +} + +func (m *interruptedRuntimeModel) Generate(context.Context, ai.Request) (*ai.Response, error) { + return nil, errors.New("interruptedRuntimeModel: Generate is not scripted") +} + +func (m *interruptedRuntimeModel) Stream(_ context.Context, _ ai.Request) ai.Stream { + m.mu.Lock() + m.calls++ + attempt := m.calls + m.mu.Unlock() + + return func(yield func(ai.StreamEvent, error) bool) { + if !yield(ai.StreamEvent{Type: ai.StreamMessageStart, ID: "resp-1"}, nil) { + return + } + + if attempt == 1 { + if !yield(ai.StreamEvent{Type: ai.StreamTextDelta, Text: "half"}, nil) { + return + } + + yield(ai.StreamEvent{}, m.failure) + + return + } + + yield(ai.StreamEvent{Type: ai.StreamTextDelta, Text: m.reply}, nil) + yield(ai.StreamEvent{Type: ai.StreamMessageEnd, FinishReason: ai.FinishStop}, nil) + } +} + +func (m *interruptedRuntimeModel) attempts() int { + m.mu.Lock() + defer m.mu.Unlock() + + return m.calls +} + +func (m *interruptedRuntimeModel) Provider() ai.Provider { return ai.ProviderOpenAI } +func (m *interruptedRuntimeModel) ModelID() string { return "interrupted-runtime" } +func (m *interruptedRuntimeModel) Capabilities() ai.Capabilities { + return ai.Capabilities{Text: true, Tools: true} +} + +// TestRuntimeReissuesInterruptedStreamAndReportsTheWait covers the failure this +// change was written for: the provider truncates the response body, the turn is +// re-issued, and the frontend is told about the wait instead of appearing slow. +func TestRuntimeReissuesInterruptedStreamAndReportsTheWait(t *testing.T) { + t.Parallel() + + model := &interruptedRuntimeModel{failure: io.ErrUnexpectedEOF, reply: "recovered"} + + runtime := openTestRuntime(t, model) + events := collectRuntimeEvents(t, runtime.Prompt(t.Context(), ai.UserText("hello"))) + assertEventSequence(t, events) + + var notices []ModelRetry + + for _, event := range events { + if retry, ok := event.Payload.(ModelRetry); ok { + notices = append(notices, retry) + } + } + + require.Len(t, notices, 1) + assert.Equal(t, 1, notices[0].Attempt, "the notice counts retries, not tries") + assert.Equal(t, streamRecoveryAttempts, notices[0].MaxRetries) + assert.Equal(t, "stream ended early", notices[0].Reason) + assert.Contains(t, eventTypes(events), EventMessageDiscarded, + "the partial candidate is retracted before the re-issue") + + // The notice is live state rather than transcript, and the reducer keeps it + // only for the duration of the wait (see the reducer test in + // stream_retry_test.go). + snapshot := runtime.Snapshot() + assert.Equal(t, InteractionSucceeded, snapshot.Interaction.Outcome) + assert.False(t, snapshot.Retry.Active, "the wait ends with the next attempt") + assert.Equal(t, 2, model.attempts()) + + require.Len(t, snapshot.Transcript, 2) + assistant, ok := snapshot.Transcript[1].(ai.AssistantMessage) + require.True(t, ok) + text, ok := assistant.Parts[0].(ai.TextPart) + require.True(t, ok) + assert.Equal(t, "recovered", text.Text) + + require.NoError(t, runtime.Close(t.Context())) +} diff --git a/internal/coding/stream_retry_test.go b/internal/coding/stream_retry_test.go new file mode 100644 index 00000000..42caa154 --- /dev/null +++ b/internal/coding/stream_retry_test.go @@ -0,0 +1,184 @@ +package coding + +import ( + "io" + "strings" + "testing" + "time" + + "github.com/rsbin1178/pips/agent" + "github.com/rsbin1178/pips/ai" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +// reduceAll applies events in order, assigning the sequences a live stream +// would carry. +func reduceAll(t *testing.T, events []Event) State { + t.Helper() + + var state State + + for index, event := range events { + event.Sequence = uint64(index + 1) + + next, err := Reduce(state, event) + require.NoError(t, err) + + state = next + } + + return state +} + +// retryEvents opens a run whose first attempt streamed one delta, which is the +// state a retry notice arrives in. +func retryEvents(extra ...Event) []Event { + events := []Event{ + newSessionEvent(EventSessionOpened, SessionOpened{Provider: ai.ProviderOpenAI, ModelID: "gpt-test"}), + newInteractionEvent(EventInteractionStarted, InteractionStarted{}), + newStatusEvent(EventStatusChanged, StatusChanged{Phase: PhaseRunning}), + newTestEvent(EventRunStarted, RunStarted{Agent: "coding"}), + newTestEvent(EventTurnStarted, TurnStarted{Turn: 1}), + newTestEvent(EventMessageDelta, MessageDelta{Kind: ai.StreamTextDelta, Text: "partial"}), + } + + return append(events, extra...) +} + +func TestReduceModelRetryTracksWaitUntilTheNextEvent(t *testing.T) { + t.Parallel() + + notice := ModelRetry{Turn: 1, Attempt: 2, MaxRetries: 3, DelayMillis: 4_000, Reason: "stream ended early"} + + waiting := reduceAll(t, retryEvents(newTestEvent(EventModelRetry, notice))) + require.True(t, waiting.Retry.Active) + assert.Equal(t, 2, waiting.Retry.Attempt) + assert.Equal(t, 3, waiting.Retry.MaxRetries) + assert.Equal(t, "stream ended early", waiting.Retry.Reason) + assert.Equal(t, eventTestTime.Add(4*time.Second), waiting.Retry.Deadline, + "the deadline anchors the countdown a frontend renders") + + // A re-issued turn discards the output that streamed before the failure, + // and the first delta of the new attempt ends the wait. + resumed := reduceAll(t, retryEvents( + newTestEvent(EventMessageDiscarded, MessageDiscarded{Turn: 1}), + newTestEvent(EventModelRetry, notice), + newTestEvent(EventMessageDelta, MessageDelta{Kind: ai.StreamTextDelta, Text: "retried"}), + )) + assert.False(t, resumed.Retry.Active) + assert.Zero(t, resumed.Retry.Deadline) + require.Len(t, resumed.Draft, 1) + assert.Equal(t, "retried", resumed.Draft[0].Text) +} + +func TestReduceBatchEndsTheWaitWhenADeltaFollows(t *testing.T) { + t.Parallel() + + prefix := retryEvents() + state := reduceAll(t, prefix) + + notice := newTestEvent(EventModelRetry, ModelRetry{Turn: 1, Attempt: 2, MaxRetries: 3, Reason: "connection error"}) + notice.Sequence = uint64(len(prefix) + 1) + + waiting, err := ReduceBatch(state, []Event{notice}) + require.NoError(t, err) + require.True(t, waiting.Retry.Active) + + delta := newTestEvent(EventMessageDelta, MessageDelta{Kind: ai.StreamTextDelta, Text: "again"}) + delta.Sequence = notice.Sequence + 1 + + // The live path reduces one frame per batch, which folds delta runs; the + // folded run is output, so it must end the wait too. + resumed, err := ReduceBatch(state, []Event{notice, delta}) + require.NoError(t, err) + assert.False(t, resumed.Retry.Active) + require.NotEmpty(t, resumed.Draft) + assert.Equal(t, "again", resumed.Draft[len(resumed.Draft)-1].Text) +} + +func TestReduceModelRetryRequiresAnActiveTurn(t *testing.T) { + t.Parallel() + + // A retry notice belongs to an open turn, because a stream cannot have an + // attempt to retry without one. + events := []Event{ + newSessionEvent(EventSessionOpened, SessionOpened{Provider: ai.ProviderOpenAI, ModelID: "gpt-test"}), + newInteractionEvent(EventInteractionStarted, InteractionStarted{}), + newStatusEvent(EventStatusChanged, StatusChanged{Phase: PhaseRunning}), + newTestEvent(EventRunStarted, RunStarted{Agent: "coding"}), + newTestEvent(EventModelRetry, ModelRetry{Turn: 1, Attempt: 1, MaxRetries: 1}), + } + + prefix := events[:len(events)-1] + state := reduceAll(t, prefix) + + last := events[len(events)-1] + last.Sequence = uint64(len(events)) + + _, err := Reduce(state, last) + require.ErrorIs(t, err, ErrEventProtocol) +} + +func TestModelRetryEventRoundTripsAndRejectsMalformedPayloads(t *testing.T) { + t.Parallel() + + event := newTestEvent(EventModelRetry, ModelRetry{ + Turn: 1, Attempt: 2, MaxRetries: 6, DelayMillis: 1_500, Reason: "connection error", + }) + require.NoError(t, ValidateEvent(event)) + + encoded, err := MarshalEvent(event) + require.NoError(t, err) + + decoded, err := UnmarshalEvent(encoded) + require.NoError(t, err) + assert.Equal(t, EventModelRetry, decoded.Type) + assert.Equal(t, event.Payload, decoded.Payload) + + bad := []ModelRetry{ + {Turn: 1, Attempt: 0, MaxRetries: 1}, // no attempt number + {Turn: 1, Attempt: 3, MaxRetries: 2}, // unreachable attempt + {Turn: 1, Attempt: 1, MaxRetries: 1, DelayMillis: -1}, + {Turn: -1, Attempt: 1, MaxRetries: 1}, // negative turn + {Turn: 1, Attempt: 1, MaxRetries: 1, Reason: strings.Repeat("x", maxDiagnosticMessage+1)}, + } + for _, payload := range bad { + require.ErrorIs(t, ValidateEvent(newTestEvent(EventModelRetry, payload)), ErrInvalidEvent) + } + + assert.ErrorIs(t, ValidateEvent(newTestEvent(EventMessageDelta, ModelRetry{Attempt: 1, MaxRetries: 1})), + ErrInvalidEvent, "the payload must not ride another type") +} + +func TestAgentProjectorMapsModelStreamRetryToModelRetryEvent(t *testing.T) { + t.Parallel() + + writer, err := newEventWriter("session-1", func() time.Time { return eventTestTime }) + require.NoError(t, err) + + projector, err := newAgentProjector(writer, "interaction-1") + require.NoError(t, err) + + event, err := agent.NewEvent( + agent.RunMetadata{RunID: "run-1", Agent: "coding"}, + eventTestTime, + agent.ModelStreamEvent{Turn: 1, Event: ai.StreamEvent{ + Type: ai.StreamRetry, + Retry: ai.NewRetryNotice(2, 6, 4*time.Second, io.ErrUnexpectedEOF), + }}, + ) + require.NoError(t, err) + + projected, err := projector.project(event) + require.NoError(t, err) + assert.Equal(t, EventModelRetry, projected.Type) + + retry, ok := projected.Payload.(ModelRetry) + require.True(t, ok) + assert.Equal(t, 1, retry.Turn) + assert.Equal(t, 2, retry.Attempt) + assert.Equal(t, 6, retry.MaxRetries) + assert.Equal(t, int64(4_000), retry.DelayMillis) + assert.Equal(t, "stream ended early", retry.Reason) +} diff --git a/internal/coding/team_proposal_agent.go b/internal/coding/team_proposal_agent.go index 753b8b1d..6738e42d 100644 --- a/internal/coding/team_proposal_agent.go +++ b/internal/coding/team_proposal_agent.go @@ -426,6 +426,7 @@ func (r *Runtime) runTeamProposalAgent( agent.WithName(teamProposalAgentName), agent.WithMaxTurns(teamProposalAgentMaxTurns), agent.WithMaxTokens(teamProposalAgentMaxTokens), + streamRecoveryOption(), agent.WithParallelTools(1), agent.WithToolTimeout(r.opts.ToolTimeout), agent.WithStopWhen(func(agent.RunInfo) bool { return collector.called() }), diff --git a/internal/coding/tui/activity.go b/internal/coding/tui/activity.go index ee3be71b..e34dfc62 100644 --- a/internal/coding/tui/activity.go +++ b/internal/coding/tui/activity.go @@ -29,6 +29,8 @@ const ( activityApproval activityRecovery activityCompacting + activityRetrying + activityWaiting activityPaused activityInterrupting ) @@ -43,10 +45,18 @@ const ( activityLabelApproval = "Waiting for approval…" activityLabelRecovery = "Waiting for recovery…" activityLabelCompacting = "Compacting context…" + activityLabelRetrying = "Retrying…" + activityLabelWaiting = "Waiting for the model…" activityLabelPaused = "Paused…" activityLabelInterrupting = "Interrupting…" ) +// modelStallThreshold is how long an open turn may go without carrying any +// model progress before the activity row says so. A gateway that buffers a +// response looks identical to a dead one without this signal; twenty seconds is +// the threshold the other agent frontends use. +const modelStallThreshold = 20 * time.Second + type activityStatus struct { kind activityKind label string @@ -59,6 +69,22 @@ type activityContext struct { hasBridge bool isCanceling bool isResolvingAllowedApproval bool + // now anchors a retry countdown. A zero time renders the notice without one. + now time.Time + // lastModelEventAt is when the parent stream last carried model progress. + // Together with now it tells a slow model from a silent one. + lastModelEventAt time.Time +} + +// modelIdle is how long an open turn has gone without model progress. Zero +// means there is nothing to report: no turn is live, or the caller does not +// track arrival times at all (a test, or a non-interactive frontend). +func (context activityContext) modelIdle() time.Duration { + if context.now.IsZero() || context.lastModelEventAt.IsZero() || !hasOpenTurn(context.state.Runs) { + return 0 + } + + return context.now.Sub(context.lastModelEventAt) } var activitySpinner = spinner.Spinner{ @@ -293,9 +319,40 @@ func resolveBlockingActivity(context activityContext) (activityStatus, bool) { }, true } + if context.state.Retry.Active { + return activityStatus{ + kind: activityRetrying, + label: activityLabelRetrying, + detail: retryActivityDetail( + context.state.Retry, + context.now, + ), + }, true + } + return activityStatus{}, false } +// retryActivityDetail spells out where the retry stands, how long the wait is, +// and why the request is being re-attempted, so a viewer can tell a slow retry +// from a stall. The counter reads "retry 2/10": the second replay of a budget +// of ten, which is how the other agent frontends count them. +func retryActivityDetail(retry coding.RetryState, now time.Time) string { + detail := fmt.Sprintf("retry %d/%d", retry.Attempt, retry.MaxRetries) + + if !now.IsZero() { + if remaining := retry.Deadline.Sub(now).Round(time.Second); remaining >= time.Second { + detail += " · in " + remaining.String() + } + } + + if retry.Reason != "" { + detail += " · " + retry.Reason + } + + return detail +} + func resolveProgressActivity(context activityContext) (activityStatus, bool) { if activities := runningToolActivities(context.state.Tools); len(activities) > 0 { if len(activities) == 1 { @@ -320,6 +377,14 @@ func resolveProgressActivity(context activityContext) (activityStatus, bool) { }, true } + if idle := context.modelIdle(); idle >= modelStallThreshold { + return activityStatus{ + kind: activityWaiting, + label: activityLabelWaiting, + detail: "no data for " + idle.Round(time.Second).String(), + }, true + } + if hasDraftKind(context.state.Draft, ai.StreamTextDelta) { return activityStatus{ kind: activityResponding, label: activityLabelResponding, @@ -467,6 +532,10 @@ func activityColor(kind activityKind, palette colorPalette) color.Color { return palette.idle case activityApproval, activityRecovery, activityPaused, activityInterrupting: return palette.warning + case activityRetrying: + return palette.warning + case activityWaiting: + return palette.warning case activityUnknown: return palette.muted default: diff --git a/internal/coding/tui/activity_test.go b/internal/coding/tui/activity_test.go index f12430b8..ccf85b6b 100644 --- a/internal/coding/tui/activity_test.go +++ b/internal/coding/tui/activity_test.go @@ -5,6 +5,7 @@ import ( "iter" "strings" "testing" + "time" tea "charm.land/bubbletea/v2" "github.com/charmbracelet/x/ansi" @@ -160,6 +161,72 @@ func TestResolveActivity(t *testing.T) { name: "compaction", context: activityContext{state: withCompaction(running)}, kind: activityCompacting, label: activityLabelCompacting, visible: true, }, + { + name: "retrying a stream counts down the wait", + context: activityContext{ + state: withRetry(running, coding.RetryState{ + Active: true, Attempt: 2, MaxRetries: 6, Reason: "stream ended early", + Deadline: retryTestTime.Add(4 * time.Second), + }), + now: retryTestTime, + }, + kind: activityRetrying, + label: activityLabelRetrying, detail: "retry 2/6 · in 4s · stream ended early", + visible: true, + }, + { + name: "retrying without a clock shows the attempt only", + context: activityContext{state: withRetry(running, coding.RetryState{ + Active: true, Attempt: 1, MaxRetries: 2, + })}, + kind: activityRetrying, label: activityLabelRetrying, detail: "retry 1/2", + visible: true, + }, + { + name: "waiting when the model goes quiet", + context: activityContext{ + state: running, + now: retryTestTime, + lastModelEventAt: retryTestTime.Add(-25 * time.Second), + }, + kind: activityWaiting, label: activityLabelWaiting, detail: "no data for 25s", + visible: true, + }, + { + name: "recent model data is not a stall", + context: activityContext{ + state: running, + now: retryTestTime, + lastModelEventAt: retryTestTime.Add(-5 * time.Second), + }, + kind: activityThinking, label: activityLabelThinking, visible: true, + }, + { + name: "a running tool outranks a quiet model", + context: activityContext{ + state: withTools(running, coding.ToolState{ + Call: coding.ToolCall{ + ID: "call-1", Name: "shell", Arguments: ai.JSON(`{"command":"make test"}`), + }, + Status: coding.ToolStatusRunning, + }), + now: retryTestTime, + lastModelEventAt: retryTestTime.Add(-90 * time.Second), + }, + kind: activityTool, label: "Running…", detail: "make test", visible: true, + }, + { + name: "no open turn means no stall report", + context: activityContext{ + state: withDraft( + working, + coding.MessageDelta{Kind: ai.StreamTextDelta, Text: "done"}, + ), + now: retryTestTime, + lastModelEventAt: retryTestTime.Add(-90 * time.Second), + }, + kind: activityResponding, label: activityLabelResponding, visible: true, + }, { name: "interrupting takes priority", context: activityContext{ state: withCompaction(running), isCanceling: true, @@ -608,6 +675,15 @@ func withCompaction(state coding.State) coding.State { return state } +// retryTestTime anchors a retry countdown so the rendered detail is exact. +var retryTestTime = time.Date(2026, time.July, 21, 8, 30, 0, 0, time.UTC) + +func withRetry(state coding.State, retry coding.RetryState) coding.State { + state.Retry = retry + + return state +} + func TestActivityLabelsRemainSingleLine(t *testing.T) { t.Parallel() @@ -618,6 +694,8 @@ func TestActivityLabelsRemainSingleLine(t *testing.T) { activityLabelApproval, activityLabelRecovery, activityLabelCompacting, + activityLabelRetrying, + activityLabelWaiting, activityLabelInterrupting, } { assert.False(t, strings.ContainsAny(label, "\r\n")) diff --git a/internal/coding/tui/model.go b/internal/coding/tui/model.go index b25a45b4..ca42439d 100644 --- a/internal/coding/tui/model.go +++ b/internal/coding/tui/model.go @@ -191,7 +191,11 @@ type Model struct { // observedSequence is the sequence of the last parent-Session event // delivered to this Model. It is not the projection's sequence: an adopted // Runtime snapshot can legitimately be ahead of the delivered records. - observedSequence uint64 + observedSequence uint64 + // lastModelEventAt is when the parent stream last carried model progress: a + // new turn, a streaming delta, or a retry notice. The activity row uses it + // to report a silent model instead of looking merely slow. + lastModelEventAt time.Time starting bool cancelStart bool waiting bool @@ -2008,6 +2012,8 @@ func (m *Model) activityStatus() (activityStatus, bool) { hasBridge: m.bridge != nil, isCanceling: m.canceling, isResolvingAllowedApproval: m.mainApprovalExecutionStarting(), + now: time.Now(), + lastModelEventAt: m.lastModelEventAt, }) } @@ -2789,6 +2795,25 @@ func (m *Model) streamBatchRefresh(batch []streamItem) []tea.Cmd { return append(refresh, m.loadPlanViewIfNeeded()) } +// modelProgressObserved reports whether a frame carried anything that shows the +// model is working: a new interaction, run, or turn, a streaming delta, or a +// retry notice. Tool lifecycle events deliberately do not count, because a +// running tool owns the activity row on its own. +func modelProgressObserved(events []coding.Event) bool { + for _, event := range events { + switch event.Type { + case coding.EventInteractionStarted, coding.EventRunStarted, coding.EventTurnStarted, + coding.EventMessageDelta, coding.EventModelRetry: + return true + default: + // Other events are either not model progress or own a row of their + // own; new event types stay inert here. + } + } + + return false +} + // advanceParentState applies one frame of parent-Session events to m.state and // reports whether the projection can still be rendered. func (m *Model) advanceParentState(items []streamItem) bool { @@ -2803,6 +2828,10 @@ func (m *Model) advanceParentState(items []streamItem) bool { return true } + if modelProgressObserved(events) { + m.lastModelEventAt = time.Now() + } + if m.runtimeState && m.controller != nil { if err := m.checkParentSequence(events); err != nil { m.failStream(err)