diff --git a/go.sum b/go.sum index b58e609..634f02d 100644 --- a/go.sum +++ b/go.sum @@ -76,16 +76,12 @@ github.com/go-git/gcfg v1.5.1-0.20230307220236-3a3c6141e376 h1:+zs/tPmkDkHx3U66D github.com/go-git/gcfg v1.5.1-0.20230307220236-3a3c6141e376/go.mod h1:an3vInlBmSxCcxctByoQdvwPiA7DTK7jaaFDBTtu0ic= github.com/go-git/go-billy/v5 v5.9.1 h1:8U73XiOTfINdItHVa6z4Gv7ToObcZ6grkqQbLryLCdA= github.com/go-git/go-billy/v5 v5.9.1/go.mod h1:ExsU+jcGwXTBOnyilvAnEM1wug1IxHr4yP2ZXsNRtV0= -github.com/go-git/go-git/v5 v5.19.1 h1:nX27AnaU43/K5bKktKwgBmR9lawoYVe1Ckg0rgzzN00= -github.com/go-git/go-git/v5 v5.19.1/go.mod h1:Pb1v0c7/g8aGQJwx9Us09W85yGoyvSwuhEGMH7zjDKQ= github.com/go-git/go-git/v5 v5.19.2 h1:wkfn7vOlUBu8ivAWKBWisTiwJK4jYHzTF8Ndv1LyGqY= github.com/go-git/go-git/v5 v5.19.2/go.mod h1:QqCBE1EFN5ddFmrliLQ3/ntRCUjZU3EJuwuB/jWEHjk= github.com/go-quicktest/qt v1.101.0 h1:O1K29Txy5P2OK0dGo59b7b0LR6wKfIhttaAhHUyn7eI= github.com/go-quicktest/qt v1.101.0/go.mod h1:14Bz/f7NwaXPtdYEgzsx46kqSxVwTbzVZsDC26tQJow= github.com/go-viper/mapstructure/v2 v2.5.0 h1:vM5IJoUAy3d7zRSVtIwQgBj7BiWtMPfmPEgAXnvj1Ro= github.com/go-viper/mapstructure/v2 v2.5.0/go.mod h1:oJDH3BJKyqBA2TXFhDsKDGDTlndYOZ6rGS0BRZIxGhM= -github.com/gofrs/uuid/v5 v5.4.0 h1:EfbpCTjqMuGyq5ZJwxqzn3Cbr2d0rUZU7v5ycAk/e/0= -github.com/gofrs/uuid/v5 v5.4.0/go.mod h1:CDOjlDMVAtN56jqyRUZh58JT31Tiw7/oQyEXZV+9bD8= github.com/gofrs/uuid/v5 v5.5.0 h1:FkPv6jYQRbZtH3bD8yC7106u+CedTCLF8+t7CLHSZNo= github.com/gofrs/uuid/v5 v5.5.0/go.mod h1:bbAA98EoIlxyRHIVg6ektCSsZ5n8mSbwgEhvhMYlZgg= github.com/golang-jwt/jwt/v5 v5.3.1 h1:kYf81DTWFe7t+1VvL7eS+jKFVWaUnK9cB1qbwn63YCY= @@ -399,8 +395,6 @@ modernc.org/opt v0.2.0/go.mod h1:03fq9lsNfvkYSfxrfUhZCWPk1lm4cq4N+Bh//bEtgns= modernc.org/sortutil v1.2.1 h1:+xyoGf15mM3NMlPDnFqrteY07klSFxLElE2PVuWIJ7w= modernc.org/sortutil v1.2.1/go.mod h1:7ZI3a3REbai7gzCLcotuw9AC4VZVpYMjDzETGsSMqJE= modernc.org/sqlite v1.26.0/go.mod h1:FL3pVXie73rg3Rii6V/u5BoHlSoyeZeIgKZEgHARyCU= -modernc.org/sqlite v1.54.0 h1:JCxR4qwkJvOaqAoYcgDoO25Nc+ROg6EJ2LfBVzdrgog= -modernc.org/sqlite v1.54.0/go.mod h1:4ntCLuNmnH8+GNqjka1wNg7KJd5/Hi5FYp8K+XQ7GZw= modernc.org/sqlite v1.55.0 h1:hIFh0MCH0rGinQ/4KYb5/UbCkRkb+UP+OkLCVWa5MTM= modernc.org/sqlite v1.55.0/go.mod h1:4ntCLuNmnH8+GNqjka1wNg7KJd5/Hi5FYp8K+XQ7GZw= modernc.org/strutil v1.1.3/go.mod h1:MEHNA7PdEnEwLvspRMtWTNnp2nnyvMfkimT1NKNAGbw= diff --git a/internal/assistant/client.go b/internal/assistant/client.go index fb82698..549a402 100644 --- a/internal/assistant/client.go +++ b/internal/assistant/client.go @@ -32,21 +32,22 @@ type ToolExecutor func(context.Context, []ToolCall, func(StreamEvent)) ([]ToolEv // CompletionRequest describes one assistant-owned model completion request. type CompletionRequest struct { - OnEvent func(StreamEvent) `json:"-"` - OnProviderObserve func(context.Context, *CompletionRequest, int) `json:"-"` - OnProviderRequest llm.ProviderRequestHook `json:"-"` - ToolRegistry *tool.Registry `json:"-"` - ExecuteTools ToolExecutor `json:"-"` - SessionID string `json:"session_id"` - SystemPrompt string `json:"system_prompt"` - ThinkingLevel string `json:"thinking_level"` - CWD string `json:"cwd"` - Auth model.RequestAuth `json:"auth"` - Messages []database.MessageEntity `json:"messages"` - Usage model.TokenUsage `json:"usage"` - Model model.Model `json:"model"` - ProviderAttempt int `json:"-"` - DisableTools bool `json:"-"` + OnEvent func(StreamEvent) `json:"-"` + OnProviderObserve func(context.Context, *CompletionRequest, int) `json:"-"` + OnProviderRequest llm.ProviderRequestHook `json:"-"` + ToolRegistry *tool.Registry `json:"-"` + ExecuteTools ToolExecutor `json:"-"` + SessionID string `json:"session_id"` + SystemPrompt string `json:"system_prompt"` + ThinkingLevel string `json:"thinking_level"` + CWD string `json:"cwd"` + Auth model.RequestAuth `json:"auth"` + Messages []database.MessageEntity `json:"messages"` + Usage model.TokenUsage `json:"usage"` + Model model.Model `json:"model"` + ProviderAttempt int `json:"-"` + DisableTools bool `json:"-"` + ToolSideEffectsStarted bool `json:"-"` } // CompletionResult is an assistant-owned provider response plus model-visible side effects. diff --git a/internal/assistant/client_adapter.go b/internal/assistant/client_adapter.go index 8100ae0..8cce3b7 100644 --- a/internal/assistant/client_adapter.go +++ b/internal/assistant/client_adapter.go @@ -62,21 +62,22 @@ func completionRequestFromHookInput(input *llm.HookInput) *CompletionRequest { } return &CompletionRequest{ - OnEvent: nil, - OnProviderObserve: nil, - OnProviderRequest: nil, - ToolRegistry: nil, - ExecuteTools: nil, - SessionID: input.SessionID, - SystemPrompt: "", - ThinkingLevel: input.ThinkingLevel, - CWD: stringFromOptions(input.ProviderOptions, "cwd"), - Auth: requestAuthFromHookInput(input), - Messages: nil, - Usage: model.EmptyTokenUsage(), - Model: modelFromLLMRef(&input.Model), - ProviderAttempt: input.Attempt, - DisableTools: false, + OnEvent: nil, + OnProviderObserve: nil, + OnProviderRequest: nil, + ToolRegistry: nil, + ExecuteTools: nil, + SessionID: input.SessionID, + SystemPrompt: "", + ThinkingLevel: input.ThinkingLevel, + CWD: stringFromOptions(input.ProviderOptions, "cwd"), + Auth: requestAuthFromHookInput(input), + Messages: nil, + Usage: model.EmptyTokenUsage(), + Model: modelFromLLMRef(&input.Model), + ProviderAttempt: input.Attempt, + DisableTools: false, + ToolSideEffectsStarted: false, } } diff --git a/internal/assistant/client_adapter_internal_test.go b/internal/assistant/client_adapter_internal_test.go index e5b65cb..9600484 100644 --- a/internal/assistant/client_adapter_internal_test.go +++ b/internal/assistant/client_adapter_internal_test.go @@ -70,8 +70,9 @@ func TestProviderRequestFromCompletionRequestAdaptsCallbacksAndRequest(t *testin MaxTokens: 50, Reasoning: true, }, - ProviderAttempt: 3, - DisableTools: false, + ProviderAttempt: 3, + DisableTools: false, + ToolSideEffectsStarted: false, } converted := providerRequestFromCompletionRequest(request) diff --git a/internal/assistant/context_build.go b/internal/assistant/context_build.go index 6929a30..9264799 100644 --- a/internal/assistant/context_build.go +++ b/internal/assistant/context_build.go @@ -28,21 +28,22 @@ func (runtime *Runtime) ContextUsage(ctx context.Context, sessionID, cwd string) } request := &CompletionRequest{ - OnEvent: nil, - OnProviderObserve: nil, - OnProviderRequest: nil, - ToolRegistry: registry, - ExecuteTools: nil, - SessionID: sessionID, - SystemPrompt: "", - ThinkingLevel: "", - CWD: cwd, - Auth: model.RequestAuth{Headers: nil, APIKey: "", Error: "", OK: false}, - Messages: nil, - Usage: model.EmptyTokenUsage(), - Model: selectedModel, - ProviderAttempt: 0, - DisableTools: false, + OnEvent: nil, + OnProviderObserve: nil, + OnProviderRequest: nil, + ToolRegistry: registry, + ExecuteTools: nil, + SessionID: sessionID, + SystemPrompt: "", + ThinkingLevel: "", + CWD: cwd, + Auth: model.RequestAuth{Headers: nil, APIKey: "", Error: "", OK: false}, + Messages: nil, + Usage: model.EmptyTokenUsage(), + Model: selectedModel, + ProviderAttempt: 0, + DisableTools: false, + ToolSideEffectsStarted: false, } messages := []database.MessageEntity{} diff --git a/internal/assistant/context_compaction.go b/internal/assistant/context_compaction.go index dd25d4e..0a86916 100644 --- a/internal/assistant/context_compaction.go +++ b/internal/assistant/context_compaction.go @@ -326,21 +326,22 @@ func (runtime *Runtime) summarizeCompaction( ) (string, error) { systemPrompt := compaction.SystemPrompt(plan.PreviousSummary, plan.SplitTurnSummary) request := &CompletionRequest{ - OnEvent: nil, - OnProviderObserve: runtime.emitProviderRequest, - OnProviderRequest: runtime.dispatchProviderRequestHook, - ToolRegistry: tool.NewRegistry(cwd), - ExecuteTools: nil, - DisableTools: true, - SessionID: sessionID, - SystemPrompt: systemPrompt, - ThinkingLevel: thinkingOff, - CWD: cwd, - Auth: auth, - Messages: plan.Messages, - Usage: compactionRequestUsage(selectedModel, systemPrompt, plan.Messages), - Model: *selectedModel, - ProviderAttempt: 0, + OnEvent: nil, + OnProviderObserve: runtime.emitProviderRequest, + OnProviderRequest: runtime.dispatchProviderRequestHook, + ToolRegistry: tool.NewRegistry(cwd), + ExecuteTools: nil, + DisableTools: true, + ToolSideEffectsStarted: false, + SessionID: sessionID, + SystemPrompt: systemPrompt, + ThinkingLevel: thinkingOff, + CWD: cwd, + Auth: auth, + Messages: plan.Messages, + Usage: compactionRequestUsage(selectedModel, systemPrompt, plan.Messages), + Model: *selectedModel, + ProviderAttempt: 0, } result, err := runtime.completeWithRetry(ctx, request, nil) diff --git a/internal/assistant/context_overflow_compaction.go b/internal/assistant/context_overflow_compaction.go index f4d3204..22930d1 100644 --- a/internal/assistant/context_overflow_compaction.go +++ b/internal/assistant/context_overflow_compaction.go @@ -37,7 +37,7 @@ func (runtime *Runtime) completeWithProviderOverflowRecovery( return input.build, input.compactionEntry, result, nil } - if !IsContextWindowError(err) { + if !IsContextWindowError(err) || input.build.Request.ToolSideEffectsStarted { return input.build, input.compactionEntry, nil, err } diff --git a/internal/assistant/context_overflow_compaction_test.go b/internal/assistant/context_overflow_compaction_test.go index 8bf0304..d1d0f88 100644 --- a/internal/assistant/context_overflow_compaction_test.go +++ b/internal/assistant/context_overflow_compaction_test.go @@ -77,6 +77,29 @@ func TestRuntime_ProviderContextOverflowRecoveryScenarios(t *testing.T) { } } +func TestRuntime_ProviderContextOverflowDoesNotRecoverAfterToolExecutionStarts(t *testing.T) { + t.Parallel() + + client := &recordingCompleter{ + complete: func(_ int, request *assistant.CompletionRequest) (*assistant.CompletionResult, error) { + _, err := request.ExecuteTools(context.Background(), nil, nil) + require.NoError(t, err) + + return nil, testContextWindowError() + }, + requests: nil, + disableToolsByCall: nil, + } + runtime := newProviderOverflowRecoveryRuntime(t, client) + + response, _, _, err := runProviderOverflowPrompt(t, runtime, t.Name()) + + require.Nil(t, response) + require.Error(t, err) + assert.True(t, assistant.IsContextWindowError(err)) + assert.Len(t, client.requests, 1) +} + func TestRuntime_ProviderContextOverflowPreservesOriginalErrorWhenNoCompaction(t *testing.T) { t.Parallel() diff --git a/internal/assistant/export_test.go b/internal/assistant/export_test.go index 5c343a4..2d30e05 100644 --- a/internal/assistant/export_test.go +++ b/internal/assistant/export_test.go @@ -131,8 +131,9 @@ func newZeroCompletionRequest(auth model.RequestAuth) *CompletionRequest { MaxTokens: 0, Reasoning: false, }, - ProviderAttempt: 0, - DisableTools: false, + ProviderAttempt: 0, + DisableTools: false, + ToolSideEffectsStarted: false, } } diff --git a/internal/assistant/llm_conversion_internal_test.go b/internal/assistant/llm_conversion_internal_test.go index a13e4a2..726b9c1 100644 --- a/internal/assistant/llm_conversion_internal_test.go +++ b/internal/assistant/llm_conversion_internal_test.go @@ -57,8 +57,9 @@ func TestLLMRequestFromCompletionRequestConvertsAssistantState(t *testing.T) { MaxTokens: 4096, Reasoning: true, }, - ProviderAttempt: 0, - DisableTools: false, + ProviderAttempt: 0, + DisableTools: false, + ToolSideEffectsStarted: false, } converted := llmRequestFromCompletionRequest(request) @@ -92,21 +93,22 @@ func TestLLMRequestFromCompletionRequestNilAndDisabledTools(t *testing.T) { assert.Empty(t, empty.Tools) converted := llmRequestFromCompletionRequest(&CompletionRequest{ - OnEvent: nil, - OnProviderObserve: nil, - OnProviderRequest: nil, - ToolRegistry: tool.NewRegistry(t.TempDir()), - ExecuteTools: nil, - SessionID: "", - SystemPrompt: "", - ThinkingLevel: "", - CWD: "", - Auth: model.RequestAuth{Headers: nil, APIKey: "", Error: "", OK: false}, - Messages: nil, - Usage: model.EmptyTokenUsage(), - Model: emptyTestModel(), - ProviderAttempt: 0, - DisableTools: true, + OnEvent: nil, + OnProviderObserve: nil, + OnProviderRequest: nil, + ToolRegistry: tool.NewRegistry(t.TempDir()), + ExecuteTools: nil, + SessionID: "", + SystemPrompt: "", + ThinkingLevel: "", + CWD: "", + Auth: model.RequestAuth{Headers: nil, APIKey: "", Error: "", OK: false}, + Messages: nil, + Usage: model.EmptyTokenUsage(), + Model: emptyTestModel(), + ProviderAttempt: 0, + DisableTools: true, + ToolSideEffectsStarted: false, }) assert.Empty(t, converted.Tools) assert.True(t, converted.DisableTools) diff --git a/internal/assistant/provider_hooks_internal_test.go b/internal/assistant/provider_hooks_internal_test.go index b70348f..fffaa1c 100644 --- a/internal/assistant/provider_hooks_internal_test.go +++ b/internal/assistant/provider_hooks_internal_test.go @@ -173,8 +173,9 @@ func providerHookTestRequest() *CompletionRequest { MaxTokens: 0, Reasoning: false, }, - ProviderAttempt: 1, - DisableTools: false, + ProviderAttempt: 1, + DisableTools: false, + ToolSideEffectsStarted: false, } } diff --git a/internal/assistant/retry.go b/internal/assistant/retry.go index 0a3842e..8bebc79 100644 --- a/internal/assistant/retry.go +++ b/internal/assistant/retry.go @@ -77,6 +77,14 @@ func defaultRetryConfig() config.RetryConfig { } func retryBackoff(retry config.RetryConfig, onDelay func(time.Duration)) retrylib.BackoffFunc { + return retryBackoffWithOverride(retry, nil, onDelay) +} + +func retryBackoffWithOverride( + retry config.RetryConfig, + override func(time.Duration) time.Duration, + onDelay func(time.Duration), +) retrylib.BackoffFunc { retry = retry.Normalized() backoff := retrylib.NewExponential(retry.BaseDelay) @@ -89,12 +97,28 @@ func retryBackoff(retry config.RetryConfig, onDelay func(time.Duration)) retryli return 0, true } + if override != nil { + delay = override(delay) + } + onDelay(delay) return delay, false }) } +func providerRetryDelay(err error, fallback time.Duration) time.Duration { + var statusErr *provider.StatusError + if errors.As(err, &statusErr) { + providerDelay := min(statusErr.RetryAfter, provider.MaxRetryAfter) + if providerDelay > fallback { + return providerDelay + } + } + + return fallback +} + func maxRetryDelays(retry config.RetryConfig) uint64 { if retry.MaxAttempts <= 1 { return 0 @@ -148,17 +172,24 @@ func IsContextWindowError(err error) bool { } func retryDecisionFromProviderCode(err error) (retry, known bool) { - code, matched := providerErrorCode(err) - if !matched { - return false, false + codes := []string{ + providerErrorContextString(err, provider.ProviderCodeContextKey), + providerErrorContextString(err, provider.ProviderTypeContextKey), + } + if code, matched := providerErrorCode(err); matched { + codes = append(codes, code) } - if nonRetryableProviderCode(code) { - return false, true + for _, code := range codes { + if nonRetryableProviderCode(strings.ToLower(code)) { + return false, true + } } - if retryableProviderCode(code) { - return true, true + for _, code := range codes { + if retryableProviderCode(strings.ToLower(code)) { + return true, true + } } return false, false @@ -182,20 +213,50 @@ func retryableDeadlineExceeded(message string) bool { return !nonRetryableProviderMessage(message) } -func providerErrorCode(err error) (string, bool) { - oopsErr, matched := oops.AsOops(err) - if !matched { - return "", false +func providerFailureCode(err error) (string, bool) { + if value := providerErrorContextString(err, provider.ProviderCodeContextKey); value != "" { + return strings.ToLower(value), true } - codeValue, matched := oopsErr.Code().(string) - if !matched { - return "", false + return providerErrorCode(err) +} + +func providerErrorContextString(err error, key string) string { + for current := err; current != nil; current = errors.Unwrap(current) { + oopsErr, matched := oops.AsOops(current) + if !matched { + continue + } + + if value, ok := oopsErr.Context()[key].(string); ok { + if trimmed := strings.TrimSpace(value); trimmed != "" { + return trimmed + } + } } - code := strings.ToLower(strings.TrimSpace(codeValue)) + return "" +} + +func providerErrorCode(err error) (string, bool) { + for current := err; current != nil; current = errors.Unwrap(current) { + oopsErr, matched := oops.AsOops(current) + if !matched { + continue + } + + codeValue, ok := oopsErr.Code().(string) + if !ok { + continue + } - return code, code != "" + code := strings.ToLower(strings.TrimSpace(codeValue)) + if code != "" && code != "assistant_error" { + return code, true + } + } + + return "", false } func providerErrorStatus(err error) (int, bool) { @@ -228,7 +289,10 @@ func providerErrorStatus(err error) (int, bool) { func retryableProviderCode(code string) bool { switch code { - case "responses_stream_incomplete", + case "server_error", + "internal_error", + "overloaded_error", + "responses_stream_incomplete", "responses_http", "provider_http", "responses_read", @@ -243,6 +307,10 @@ func retryableProviderCode(code string) bool { func nonRetryableProviderCode(code string) bool { switch code { case "context_window_exceeded", + "invalid_request_error", + "authentication_error", + "permission_error", + "insufficient_quota", "openai_chat_decode", "anthropic_decode", "openai_response_decode", @@ -274,6 +342,10 @@ func nonRetryableProviderMessage(message string) bool { } nonRetryable := []string{ + "billing", + "insufficient quota", + "quota exceeded", + "account limit", "invalid api key", "unauthorized", "authentication", @@ -360,6 +432,8 @@ func retryableProviderMessage(message string) bool { "websocket closed", "websocket error", "terminated", + "you can retry your request", + "please try again later", } for _, pattern := range retryable { if strings.Contains(message, pattern) { diff --git a/internal/assistant/retry_internal_test.go b/internal/assistant/retry_internal_test.go index cf03e42..4121d9c 100644 --- a/internal/assistant/retry_internal_test.go +++ b/internal/assistant/retry_internal_test.go @@ -9,6 +9,7 @@ import ( "github.com/stretchr/testify/require" "github.com/omarluq/librecode/internal/config" + "github.com/omarluq/librecode/internal/provider" ) func TestRetryBackoffUsesCappedExponentialDelays(t *testing.T) { @@ -37,6 +38,59 @@ func TestRetryBackoffUsesCappedExponentialDelays(t *testing.T) { assert.Equal(t, []time.Duration{10 * time.Millisecond, 15 * time.Millisecond, 15 * time.Millisecond}, delays) } +func TestProviderRetryDelayHonorsBoundedProviderDelay(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + retryAfter time.Duration + fallback time.Duration + want time.Duration + }{ + {name: "longer provider delay", retryAfter: time.Minute, fallback: 2 * time.Second, want: time.Minute}, + {name: "longer fallback", retryAfter: time.Minute, fallback: 2 * time.Minute, want: 2 * time.Minute}, + {name: "provider delay capped", retryAfter: time.Hour, fallback: 2 * time.Second, want: provider.MaxRetryAfter}, + } + + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + t.Parallel() + + err := newProviderStatusError(test.retryAfter) + + assert.Equal(t, test.want, providerRetryDelay(err, test.fallback)) + }) + } + + err := newProviderStatusError(time.Minute) + delays := []time.Duration{} + backoff := retryBackoffWithOverride(config.RetryConfig{ + BaseDelay: 10 * time.Millisecond, + MaxDelay: 15 * time.Millisecond, + MaxAttempts: 2, + Enabled: true, + }, func(delay time.Duration) time.Duration { + return providerRetryDelay(err, delay) + }, func(delay time.Duration) { + delays = append(delays, delay) + }) + + delay, stop := backoff.Next() + require.False(t, stop) + assert.Equal(t, time.Minute, delay) + assert.Equal(t, []time.Duration{time.Minute}, delays) +} + +func newProviderStatusError(retryAfter time.Duration) *provider.StatusError { + return &provider.StatusError{ + Details: nil, + RequestShape: nil, + RequestID: "", + RetryAfter: retryAfter, + Status: 0, + } +} + func TestShouldRetryModelErrorTreatsHTTP2StreamErrorsAsTransient(t *testing.T) { t.Parallel() diff --git a/internal/assistant/retry_test.go b/internal/assistant/retry_test.go index 12f5363..21ccd27 100644 --- a/internal/assistant/retry_test.go +++ b/internal/assistant/retry_test.go @@ -64,6 +64,16 @@ func TestShouldRetryModelError(t *testing.T) { err: errors.New("provider is overloaded, please try again"), want: true, }, + { + name: "explicit provider retry guidance", + err: errors.New("an error occurred while processing your request; you can retry your request"), + want: true, + }, + { + name: "quota remains non-retryable despite retry guidance", + err: errors.New("quota exceeded; you can retry your request"), + want: false, + }, { name: "canceled context", err: context.Canceled, @@ -95,6 +105,29 @@ func TestShouldRetryModelErrorHandlesResponsesStreamFailures(t *testing.T) { Errorf("provider stream closed before completion"), want: true, }, + { + name: "response failed with transient provider code", + err: oops.In("assistant").Code("responses_failed"). + With("provider_code", "server_error"). + Errorf("processing failed"), + want: true, + }, + { + name: "response failed with transient type and unknown code", + err: oops.In("assistant").Code("responses_failed"). + With("provider_type", "server_error"). + With("provider_code", "backend_specific"). + Errorf("processing failed"), + want: true, + }, + { + name: "non-retryable type overrides transient code", + err: oops.In("assistant").Code("responses_failed"). + With("provider_type", "invalid_request_error"). + With("provider_code", "server_error"). + Errorf("invalid request"), + want: false, + }, { name: "response failed without retryable details", err: oops.In("assistant").Code("responses_failed").Errorf("invalid prompt"), diff --git a/internal/assistant/runtime_events.go b/internal/assistant/runtime_events.go index fb0431c..1cdbada 100644 --- a/internal/assistant/runtime_events.go +++ b/internal/assistant/runtime_events.go @@ -2,9 +2,11 @@ package assistant import ( "context" + "log/slog" "github.com/omarluq/librecode/internal/assistant/lifecyclepayload" "github.com/omarluq/librecode/internal/extension" + "github.com/omarluq/librecode/internal/provider" ) func (runtime *Runtime) emitProviderRequest(ctx context.Context, request *CompletionRequest, attempt int) { @@ -72,6 +74,29 @@ func (runtime *Runtime) emitProviderError(ctx context.Context, request *Completi return } + if runtime.logger != nil { + attributes := []any{ + slog.String("provider", request.Model.Provider), + slog.String("model", request.Model.ID), + slog.String("api", request.Model.API), + slog.Int("attempt", attempt), + slog.String("error", err.Error()), + } + if code, ok := providerFailureCode(err); ok { + attributes = append(attributes, slog.String("provider_code", code)) + } + + if requestID := providerErrorContextString(err, provider.ProviderRequestIDContextKey); requestID != "" { + attributes = append(attributes, slog.String("provider_request_id", requestID)) + } + + if responseID := providerErrorContextString(err, provider.ProviderResponseIDContextKey); responseID != "" { + attributes = append(attributes, slog.String("provider_response_id", responseID)) + } + + runtime.logger.WarnContext(ctx, "provider request failed", attributes...) + } + payload := lifecyclepayload.ProviderErrorPayload(&lifecyclepayload.ProviderError{ Err: err, API: request.Model.API, diff --git a/internal/assistant/runtime_model.go b/internal/assistant/runtime_model.go index 623b345..83b24d8 100644 --- a/internal/assistant/runtime_model.go +++ b/internal/assistant/runtime_model.go @@ -164,23 +164,36 @@ type modelCompletionRequestInput struct { } func (runtime *Runtime) modelCompletionRequest(input *modelCompletionRequestInput) *CompletionRequest { - return &CompletionRequest{ - OnEvent: input.onEvent, - OnProviderObserve: runtime.emitProviderRequest, - OnProviderRequest: runtime.dispatchProviderRequestHook, - ToolRegistry: input.registry, - ExecuteTools: runtime.executeProviderToolCalls(input.registry), - SessionID: input.sessionID, - SystemPrompt: input.systemPrompt, - ThinkingLevel: runtime.thinkingLevel(), - CWD: input.cwd, - Auth: input.auth, - Messages: input.messages, - Usage: input.usage, - Model: *input.selectedModel, - ProviderAttempt: 0, - DisableTools: false, + request := &CompletionRequest{ + OnEvent: input.onEvent, + OnProviderObserve: runtime.emitProviderRequest, + OnProviderRequest: runtime.dispatchProviderRequestHook, + ToolRegistry: input.registry, + ExecuteTools: nil, + SessionID: input.sessionID, + SystemPrompt: input.systemPrompt, + ThinkingLevel: runtime.thinkingLevel(), + CWD: input.cwd, + Auth: input.auth, + Messages: input.messages, + Usage: input.usage, + Model: *input.selectedModel, + ProviderAttempt: 0, + DisableTools: false, + ToolSideEffectsStarted: false, } + executeTools := runtime.executeProviderToolCalls(input.registry) + request.ExecuteTools = func( + ctx context.Context, + calls []ToolCall, + onEvent func(StreamEvent), + ) ([]ToolEvent, error) { + request.ToolSideEffectsStarted = true + + return executeTools(ctx, calls, onEvent) + } + + return request } func (runtime *Runtime) completeWithRetry( @@ -199,7 +212,9 @@ func (runtime *Runtime) completeWithRetry( var retryErr error - backoff := retryBackoff(retry, func(delay time.Duration) { + backoff := retryBackoffWithOverride(retry, func(delay time.Duration) time.Duration { + return providerRetryDelay(retryErr, delay) + }, func(delay time.Duration) { retryEvent := RetryEvent{ Kind: RetryEventStart, Error: "", @@ -256,7 +271,7 @@ func (runtime *Runtime) retryAttempt( return result, nil } - if !ShouldRetryModelError(err) { + if request.ToolSideEffectsStarted || !ShouldRetryModelError(err) { return nil, err } diff --git a/internal/assistant/runtime_test.go b/internal/assistant/runtime_test.go index 88cbc02..8898617 100644 --- a/internal/assistant/runtime_test.go +++ b/internal/assistant/runtime_test.go @@ -298,6 +298,18 @@ func TestRuntime_PromptRetriesTransientModelErrors(t *testing.T) { assert.Equal(t, assistant.RetryEventEnd, retryEvents[1].Kind) } +func TestRuntime_PromptDoesNotRetryAfterToolExecutionStarts(t *testing.T) { + t.Parallel() + + client := &sideEffectFailureCompleter{attempts: 0} + runtime, _ := newTestRuntimeWithClient(t, client) + + _, err := runtime.Prompt(context.Background(), newRuntimePromptRequest(testRuntimeCWD, "do not replay", "")) + + require.Error(t, err) + assert.Equal(t, 1, client.attempts) +} + func TestRuntime_PromptPersistsEmptyProviderResponse(t *testing.T) { t.Parallel() @@ -768,6 +780,10 @@ type retryCompleter struct { failuresRemaining int } +type sideEffectFailureCompleter struct { + attempts int +} + type failingProgressCompleter struct { emit func(*assistant.CompletionRequest) err error @@ -777,6 +793,20 @@ type emptyCompleter struct { attempts int } +func (client *sideEffectFailureCompleter) Complete( + ctx context.Context, + request *assistant.CompletionRequest, +) (*assistant.CompletionResult, error) { + client.attempts++ + + _, err := request.ExecuteTools(ctx, nil, nil) + if err != nil { + return nil, oops.In("assistant_test").Code("execute_tools").Wrapf(err, "execute tools") + } + + return nil, errors.New("provider is temporarily unavailable") +} + func (client *capturingCompleter) Complete( ctx context.Context, request *assistant.CompletionRequest, diff --git a/internal/assistant/tool_schema_cache_internal_test.go b/internal/assistant/tool_schema_cache_internal_test.go index 251311e..575ac5d 100644 --- a/internal/assistant/tool_schema_cache_internal_test.go +++ b/internal/assistant/tool_schema_cache_internal_test.go @@ -227,7 +227,8 @@ func newSchemaEstimateRequest(t *testing.T, api string, disableTools bool) *Comp MaxTokens: 0, Reasoning: false, }, - ProviderAttempt: 0, - DisableTools: disableTools, + ProviderAttempt: 0, + DisableTools: disableTools, + ToolSideEffectsStarted: false, } } diff --git a/internal/provider/client_internal_test.go b/internal/provider/client_internal_test.go index cd91e99..3baf09b 100644 --- a/internal/provider/client_internal_test.go +++ b/internal/provider/client_internal_test.go @@ -4,6 +4,7 @@ import ( "strings" "testing" + "github.com/samber/oops" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" @@ -76,6 +77,42 @@ func TestParseSSEResultExtractsToolCallFromOutputItems(t *testing.T) { assert.Equal(t, "README.md", testutil.ToolArgumentFields(result.ToolCalls[0].Arguments)["path"]) } +func TestParseSSEResultPreservesFailureMetadata(t *testing.T) { + t.Parallel() + + stream := "data: " + `{"type":"response.failed","request_id":"req_123",` + + `"error":{"type":"server_error","code":"internal_error"},` + + `"response":{"id":"resp_123","error":{"message":"please retry"}}}` + "\n" + + _, err := parseSSEResult(strings.NewReader(stream), nil) + require.Error(t, err) + + coded, matched := oops.AsOops(err) + require.True(t, matched) + assert.Equal(t, "responses_failed", coded.Code()) + assert.Equal(t, "server_error", coded.Context()[ProviderTypeContextKey]) + assert.Equal(t, "internal_error", coded.Context()[ProviderCodeContextKey]) + assert.Equal(t, "resp_123", coded.Context()[ProviderResponseIDContextKey]) + assert.Equal(t, "req_123", coded.Context()[ProviderRequestIDContextKey]) + assert.Contains(t, err.Error(), "please retry") +} + +func TestParseSSEResultMergesSplitFailureMetadata(t *testing.T) { + t.Parallel() + + stream := "data: " + `{"type":"response.failed","error":{"type":"server_error","code":"backend_specific"},` + + `"response":{"id":"resp_123","error":{"message":"please retry"}}}` + "\n" + + _, err := parseSSEResult(strings.NewReader(stream), nil) + require.Error(t, err) + + coded, matched := oops.AsOops(err) + require.True(t, matched) + assert.Equal(t, "server_error", coded.Context()[ProviderTypeContextKey]) + assert.Equal(t, "backend_specific", coded.Context()[ProviderCodeContextKey]) + assert.Contains(t, err.Error(), "please retry") +} + func TestParseSSEResultFailureCases(t *testing.T) { t.Parallel() diff --git a/internal/provider/errors.go b/internal/provider/errors.go index ab2cb45..22e07c6 100644 --- a/internal/provider/errors.go +++ b/internal/provider/errors.go @@ -5,6 +5,7 @@ import ( "encoding/json" "fmt" "strings" + "time" "unicode/utf8" "github.com/samber/oops" @@ -13,12 +14,22 @@ import ( const ( providerBodyPreviewBytes = 4096 - providerCodeContextKey = "provider_code" - providerParamContextKey = "provider_param" - providerTypeContextKey = "provider_type" - providerBodyPreviewContextKey = "body_preview" - providerBodyTruncatedContextKey = "body_truncated" - providerRequestShapeContextKey = "request_shape" + // ProviderCodeContextKey identifies a provider error code in structured error context. + ProviderCodeContextKey = "provider_code" + // ProviderTypeContextKey identifies a provider error type in structured error context. + ProviderTypeContextKey = "provider_type" + // ProviderRequestIDContextKey identifies a provider request ID in structured error context. + ProviderRequestIDContextKey = "provider_request_id" + // ProviderResponseIDContextKey identifies a provider response ID in structured error context. + ProviderResponseIDContextKey = "provider_response_id" + // ProviderParamContextKey identifies the provider request parameter associated with an error. + ProviderParamContextKey = "provider_param" + // ProviderBodyPreviewContextKey identifies a bounded provider response-body preview. + ProviderBodyPreviewContextKey = "body_preview" + // ProviderBodyTruncatedContextKey reports whether a provider response-body preview was truncated. + ProviderBodyTruncatedContextKey = "body_truncated" + // ProviderRequestShapeContextKey identifies safe request-shape diagnostics. + ProviderRequestShapeContextKey = "request_shape" ) type providerError struct { @@ -42,6 +53,8 @@ type StatusError struct { Details *ErrorDetails err error RequestShape *RequestShape + RequestID string + RetryAfter time.Duration Status int } @@ -62,11 +75,23 @@ func (err *StatusError) Unwrap() error { } func providerStatusError(status int, content []byte, requestShape *RequestShape) error { + return providerStatusErrorWithRetryAfter(status, content, requestShape, 0, "") +} + +func providerStatusErrorWithRetryAfter( + status int, + content []byte, + requestShape *RequestShape, + retryAfter time.Duration, + requestID string, +) error { details := providerErrorDetailsFromBytes(content) statusErr := &StatusError{ Details: &details, err: nil, RequestShape: requestShape, + RequestID: requestID, + RetryAfter: retryAfter, Status: status, } statusErr.err = providerStatusOops(statusErr) @@ -78,28 +103,32 @@ func providerStatusOops(statusErr *StatusError) error { message := providerStatusMessage(statusErr.Status, statusErr.Details) builder := oops.In("provider").Code("provider_status").With("status", statusErr.Status) + if statusErr.RequestID != "" { + builder = builder.With(ProviderRequestIDContextKey, statusErr.RequestID) + } + if statusErr.Details.Type != "" { - builder = builder.With(providerTypeContextKey, statusErr.Details.Type) + builder = builder.With(ProviderTypeContextKey, statusErr.Details.Type) } if statusErr.Details.Code != "" { - builder = builder.With(providerCodeContextKey, statusErr.Details.Code) + builder = builder.With(ProviderCodeContextKey, statusErr.Details.Code) } if statusErr.Details.Param != "" { - builder = builder.With(providerParamContextKey, statusErr.Details.Param) + builder = builder.With(ProviderParamContextKey, statusErr.Details.Param) } if statusErr.Details.BodyPreview != "" { - builder = builder.With(providerBodyPreviewContextKey, statusErr.Details.BodyPreview) + builder = builder.With(ProviderBodyPreviewContextKey, statusErr.Details.BodyPreview) } if statusErr.Details.BodyTruncated { - builder = builder.With(providerBodyTruncatedContextKey, true) + builder = builder.With(ProviderBodyTruncatedContextKey, true) } if !statusErr.RequestShape.empty() { - builder = builder.With(providerRequestShapeContextKey, statusErr.RequestShape.Payload()) + builder = builder.With(ProviderRequestShapeContextKey, statusErr.RequestShape.Payload()) } return builder.Errorf("%s", message) @@ -122,7 +151,7 @@ func providerErrorToOops(code string, providerError *providerError) error { return oops.In("provider"). Code(code). With(jsonTypeKey, providerError.Type). - With(providerCodeContextKey, providerError.Code). + With(ProviderCodeContextKey, providerError.Code). Errorf("%s", message) } diff --git a/internal/provider/errors_internal_test.go b/internal/provider/errors_internal_test.go index 88b7196..09d0a18 100644 --- a/internal/provider/errors_internal_test.go +++ b/internal/provider/errors_internal_test.go @@ -3,6 +3,7 @@ package provider import ( "strings" "testing" + "time" "unicode/utf8" "github.com/samber/oops" @@ -29,6 +30,17 @@ func TestProviderStatusErrorUsesStructuredMessage(t *testing.T) { assert.Equal(t, "provider", coded.Domain()) } +func TestProviderStatusErrorPreservesRetryAfter(t *testing.T) { + t.Parallel() + + err := providerStatusErrorWithRetryAfter(503, nil, nil, 3*time.Second, "req_123") + + var statusErr *StatusError + require.ErrorAs(t, err, &statusErr) + assert.Equal(t, 3*time.Second, statusErr.RetryAfter) + assert.Equal(t, "req_123", statusErr.RequestID) +} + func TestProviderStatusErrorFallsBackToHTTPStatus(t *testing.T) { t.Parallel() @@ -70,9 +82,9 @@ func TestProviderStatusErrorIncludesStructuredDetails(t *testing.T) { require.ErrorAs(t, err, &coded) context := coded.Context() assert.Equal(t, 400, context["status"]) - assert.Equal(t, "invalid_request_error", context[providerTypeContextKey]) - assert.Equal(t, "unknown_parameter", context[providerCodeContextKey]) - assert.Equal(t, "input[2].content", context[providerParamContextKey]) + assert.Equal(t, "invalid_request_error", context[ProviderTypeContextKey]) + assert.Equal(t, "unknown_parameter", context[ProviderCodeContextKey]) + assert.Equal(t, "input[2].content", context[ProviderParamContextKey]) assert.Equal(t, map[string]any{ "has_include": false, "has_parallel_tool_calls": false, @@ -81,8 +93,8 @@ func TestProviderStatusErrorIncludesStructuredDetails(t *testing.T) { "input_count": 3, "key_count": 1, "keys": []string{jsonInputKey}, - }, context[providerRequestShapeContextKey]) - assert.Contains(t, context[providerBodyPreviewContextKey], "unknown parameter") + }, context[ProviderRequestShapeContextKey]) + assert.Contains(t, context[ProviderBodyPreviewContextKey], "unknown parameter") } func TestProviderStatusErrorBoundsBodyPreview(t *testing.T) { @@ -105,11 +117,11 @@ func TestProviderStatusErrorBoundsBodyPreview(t *testing.T) { var coded oops.OopsError require.ErrorAs(t, err, &coded) context := coded.Context() - preview, ok := context[providerBodyPreviewContextKey].(string) + preview, ok := context[ProviderBodyPreviewContextKey].(string) require.True(t, ok) assert.LessOrEqual(t, len(preview), providerBodyPreviewBytes) assert.True(t, utf8.ValidString(preview)) - assertIsTrue(t, context[providerBodyTruncatedContextKey]) + assertIsTrue(t, context[ProviderBodyTruncatedContextKey]) }) } } diff --git a/internal/provider/http.go b/internal/provider/http.go index fd295e6..35d8bad 100644 --- a/internal/provider/http.go +++ b/internal/provider/http.go @@ -8,7 +8,9 @@ import ( "maps" "net/http" "net/url" + "strconv" "strings" + "time" "github.com/samber/oops" @@ -17,6 +19,9 @@ import ( ) const ( + // MaxRetryAfter bounds provider-directed waits to avoid unbounded or overflowing delays. + MaxRetryAfter = 5 * time.Minute + providerResponseLimitBytes int64 = 16 * units.MiB codexAccountIDHeader = "chatgpt-account-id" codexOriginatorHeader = "originator" @@ -45,7 +50,13 @@ func (client *HTTPCompletionClient) requestProviderStream( return nil, oops.In("provider").Code("provider_error_read").Wrapf(readErr, "read provider error") } - return nil, providerStatusError(response.StatusCode, content, providerRequestShape(payload)) + return nil, providerStatusErrorWithRetryAfter( + response.StatusCode, + content, + providerRequestShape(payload), + providerRetryAfter(response.Header, time.Now()), + providerRequestID(response.Header), + ) } return parse(response.Body) @@ -76,6 +87,38 @@ func readProviderBody(reader io.Reader) ([]byte, error) { return body, providerWrap(err, "read provider response") } +func providerRequestID(header http.Header) string { + return firstNonEmptyString(header.Get("X-Request-Id"), header.Get("Request-Id")) +} + +// providerRetryAfter returns a bounded provider-requested delay. +func providerRetryAfter(header http.Header, now time.Time) time.Duration { + millisecondsValue := strings.TrimSpace(header.Get("Retry-After-Ms")) + if milliseconds, err := strconv.ParseInt(millisecondsValue, 10, 64); err == nil && milliseconds > 0 { + return boundedRetryAfter(milliseconds, time.Millisecond) + } + + value := strings.TrimSpace(header.Get("Retry-After")) + if seconds, err := strconv.ParseInt(value, 10, 64); err == nil && seconds > 0 { + return boundedRetryAfter(seconds, time.Second) + } + + if date, err := http.ParseTime(value); err == nil && date.After(now) { + return min(date.Sub(now), MaxRetryAfter) + } + + return 0 +} + +func boundedRetryAfter(value int64, unit time.Duration) time.Duration { + maximum := int64(MaxRetryAfter / unit) + if value >= maximum { + return MaxRetryAfter + } + + return time.Duration(value) * unit +} + func closeBody(body io.Closer) { if err := body.Close(); err != nil { return diff --git a/internal/provider/http_internal_test.go b/internal/provider/http_internal_test.go index 2a54e6e..068d2be 100644 --- a/internal/provider/http_internal_test.go +++ b/internal/provider/http_internal_test.go @@ -1,8 +1,12 @@ package provider import ( + "math" + "net/http" + "strconv" "strings" "testing" + "time" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" @@ -10,6 +14,48 @@ import ( "github.com/omarluq/librecode/internal/llm" ) +func TestProviderRetryAfterIsBounded(t *testing.T) { + t.Parallel() + + const retryAfterHeader = "Retry-After" + + now := time.Date(2026, time.July, 29, 12, 0, 0, 0, time.UTC) + tests := []struct { + header http.Header + name string + want time.Duration + }{ + { + name: "milliseconds overflow", + header: http.Header{"Retry-After-Ms": []string{strconv.FormatInt(math.MaxInt64, 10)}}, + want: MaxRetryAfter, + }, + { + name: "seconds overflow", + header: http.Header{retryAfterHeader: []string{strconv.FormatInt(math.MaxInt64, 10)}}, + want: MaxRetryAfter, + }, + { + name: "date beyond maximum", + header: http.Header{retryAfterHeader: []string{now.Add(time.Hour).Format(http.TimeFormat)}}, + want: MaxRetryAfter, + }, + { + name: "short delay", + header: http.Header{retryAfterHeader: []string{"3"}}, + want: 3 * time.Second, + }, + } + + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + t.Parallel() + + assert.Equal(t, test.want, providerRetryAfter(test.header, now)) + }) + } +} + func TestReadProviderBodyRejectsBodiesAboveLimit(t *testing.T) { t.Parallel() diff --git a/internal/provider/openai_responses_sse.go b/internal/provider/openai_responses_sse.go index 6eac03b..12fbb5e 100644 --- a/internal/provider/openai_responses_sse.go +++ b/internal/provider/openai_responses_sse.go @@ -1,6 +1,7 @@ package provider import ( + "encoding/json" "errors" "io" "strings" @@ -298,17 +299,84 @@ func providerResultFromSSEFinalResponse(accumulator *sseAccumulator, fallbackTex func sseProviderError(code string, event map[string]any, fallback string) error { message := fallback - if response, ok := event["response"].(map[string]any); ok { - if responseMessage := sseErrorMessage(response); responseMessage != "" { - message = responseMessage - } + response := map[string]any(nil) + if value, ok := event["response"].(map[string]any); ok { + response = value + } + + if responseMessage := sseErrorMessage(response); responseMessage != "" { + message = responseMessage } if eventMessage := sseErrorMessage(event); eventMessage != "" { message = eventMessage } - return oops.In("provider").Code(code).Errorf("%s", message) + builder := oops.In("provider").Code(code) + + errorDetails := sseErrorDetails(response, event) + if errorDetails.Type != "" { + builder = builder.With(ProviderTypeContextKey, errorDetails.Type) + } + + if errorDetails.Code != "" { + builder = builder.With(ProviderCodeContextKey, errorDetails.Code) + } + + responseID := firstNonEmptyString(stringValue(response["id"]), stringValue(event["response_id"])) + if responseID != "" { + builder = builder.With(ProviderResponseIDContextKey, responseID) + } + + requestID := firstNonEmptyString(stringValue(event["request_id"]), stringValue(response["request_id"])) + if requestID != "" { + builder = builder.With(ProviderRequestIDContextKey, requestID) + } + + return builder.Errorf("%s", message) +} + +func sseErrorDetails(objects ...map[string]any) ErrorDetails { + merged := emptyProviderErrorDetails() + + for _, object := range objects { + details := sseObjectErrorDetails(object) + merged.Message = firstNonEmptyString(merged.Message, details.Message) + merged.Type = firstNonEmptyString(merged.Type, details.Type) + merged.Code = firstNonEmptyString(merged.Code, details.Code) + merged.Param = firstNonEmptyString(merged.Param, details.Param) + } + + return merged +} + +func sseObjectErrorDetails(object map[string]any) ErrorDetails { + if object == nil { + return emptyProviderErrorDetails() + } + + if nested, ok := object["error"].(map[string]any); ok { + if details := providerErrorDetailsFromMap(nested); providerErrorDetailsPresent(&details) { + return details + } + } + + return providerErrorDetailsFromMap(object) +} + +func providerErrorDetailsPresent(details *ErrorDetails) bool { + return details != nil && (details.Message != "" || details.Type != "" || details.Code != "") +} + +func providerErrorDetailsFromMap(object map[string]any) ErrorDetails { + content, err := json.Marshal(object) + if err != nil { + return emptyProviderErrorDetails() + } + + details, _ := providerErrorDetailsFromJSON(content) + + return details } func sseErrorMessage(object map[string]any) string { diff --git a/internal/terminal/async_events.go b/internal/terminal/async_events.go index 8b32eb0..550dad2 100644 --- a/internal/terminal/async_events.go +++ b/internal/terminal/async_events.go @@ -366,6 +366,14 @@ func (app *App) handlePromptLifecycleEvent(ctx context.Context, payload *asyncEv return true case asyncEventPromptRetry: app.emitPromptRetryExtensionEvent(ctx, payload) + + if payload.Provider == string(assistant.RetryEventStart) { + app.streamingText = "" + app.streamingThinkingText = "" + app.resetStreamingBlocks() + app.streamedToolEvents = 0 + } + app.setStatus(payload.Text) return true diff --git a/internal/terminal/async_events_internal_test.go b/internal/terminal/async_events_internal_test.go index 0cd9773..fb4faa6 100644 --- a/internal/terminal/async_events_internal_test.go +++ b/internal/terminal/async_events_internal_test.go @@ -12,6 +12,7 @@ import ( "github.com/omarluq/librecode/internal/assistant" "github.com/omarluq/librecode/internal/model" + "github.com/omarluq/librecode/internal/tool" "github.com/omarluq/librecode/internal/transcript" ) @@ -25,6 +26,7 @@ const ( asyncTestToolStart = "async-bash" asyncTestCompact = "async context auto-compacted" asyncTestIgnored = "async-ignored" + asyncTestPartial = "partial" ) type promptHandlerCase struct { @@ -519,10 +521,35 @@ func promptLifecycleEventCases() []promptLifecycleCase { { name: "prompt retry", payload: asyncTestEvent(asyncEventPromptRetry, string(assistant.RetryEventStart), "retrying", 1), - setup: func(*App) {}, + setup: func(app *App) { + app.streamingText = asyncTestPartial + app.streamingThinkingText = "thought" + app.transcript.Streaming.Blocks = []chatMessage{{ + Role: transcript.RoleAssistant, + Content: asyncTestPartial, + CreatedAt: time.Time{}, + }} + app.runningToolBlocks = []runningToolBlock{{ + StartedAt: time.Time{}, + Call: assistant.ToolCallEvent{ + ArgumentsJSON: "", + ID: "", + ParentCallID: "", + Name: testToolRead, + Arguments: tool.EmptyArguments(), + Sequence: 0, + }, + }} + app.streamedToolEvents = 1 + }, assert: func(t *testing.T, app *App) { t.Helper() assert.Equal(t, "retrying", app.statusMessage) + assert.Empty(t, app.streamingText) + assert.Empty(t, app.streamingThinkingText) + assert.Empty(t, app.transcript.Streaming.Blocks) + assert.Empty(t, app.runningToolBlocks) + assert.Zero(t, app.streamedToolEvents) }, wantHandled: true, },