feat(columbo): retry model calls with per-attempt deadlines (#598)

This commit is contained in:
Antoine Lecompte
2026-08-31 11:28:12 -07:00
committed by GitHub
parent b7629172be
commit 8e9cd64c16
6 changed files with 281 additions and 19 deletions
+5 -1
View File
@@ -99,7 +99,11 @@ as AI-generated with the executed queries listed — so the worst case is a
misleading note that staff are told to verify.
Hard limits per investigation: tool-call budget (`COLUMBO_MAX_TOOL_CALLS`,
16), wall clock (`COLUMBO_TIMEOUT_SECONDS`, 300), tool results truncated to
16), wall clock (`COLUMBO_TIMEOUT_SECONDS`, 600), model calls retried on
transport errors/timeouts/5xx with per-attempt deadlines
(`COLUMBO_MODEL_TIMEOUT_SECONDS` 120 × `COLUMBO_MODEL_ATTEMPTS` 3 — the
response body is buffered per attempt so a mid-body stall retries instead of
killing the run), tool results truncated to
`COLUMBO_TOOL_RESULT_BYTES` with the full payload kept harness-side for jq,
bounded queue + workers, note capped to the embed limit. Model-supplied
query parameters are clamped in the harness before the backend sees them —
+27 -9
View File
@@ -8,6 +8,7 @@ import (
"context"
"encoding/json"
"fmt"
"net/http"
"strings"
"sync"
"time"
@@ -37,14 +38,16 @@ type Investigation struct {
}
type Config struct {
OpenRouterURL string
APIKey string
Model string
TriageModel string
MetricsURL string
LogsURL string
MaxToolCalls int
ToolResultBytes int
OpenRouterURL string
APIKey string
Model string
TriageModel string
MetricsURL string
LogsURL string
MaxToolCalls int
ToolResultBytes int
ModelCallTimeout time.Duration
ModelCallAttempts int
}
type Runner struct {
@@ -280,11 +283,26 @@ func truncateNote(note string) string {
return note[:maxNoteChars] + "…"
}
// chatModel routes requests through the retrying transport: per-ATTEMPT
// deadlines with a couple of retries, because OpenRouter tail latency varies
// wildly across the providers it balances over — a 120s stall on one
// provider took down an otherwise healthy run. The investigation context
// remains the overall budget.
func (r *Runner) chatModel(ctx context.Context, model string) (*openai.ChatModel, error) {
timeout := r.cfg.ModelCallTimeout
if timeout <= 0 {
timeout = 120 * time.Second
}
attempts := r.cfg.ModelCallAttempts
if attempts <= 0 {
attempts = 3
}
return openai.NewChatModel(ctx, &openai.ChatModelConfig{
BaseURL: r.cfg.OpenRouterURL,
APIKey: r.cfg.APIKey,
Model: model,
Timeout: 120 * time.Second,
HTTPClient: &http.Client{
Transport: &retryTransport{base: http.DefaultTransport, attempts: attempts, perAttempt: timeout},
},
})
}
@@ -0,0 +1,98 @@
package agent
import (
"bytes"
"context"
"io"
"net/http"
"time"
"github.com/rs/zerolog"
)
const maxModelResponseBytes = 16 << 20
// retryTransport retries model calls on transport errors, timeouts, and
// retryable statuses (408/429/5xx). Each attempt gets its own deadline and
// the response body is buffered INSIDE the attempt — the failure this exists
// for was a completion stalling mid-body, which is unreachable to a retry
// once RoundTrip has returned a streaming body. Chat completions are small;
// buffering trades streaming (unused) for retryability.
type retryTransport struct {
base http.RoundTripper
attempts int
perAttempt time.Duration
}
func (t *retryTransport) RoundTrip(req *http.Request) (*http.Response, error) {
logger := zerolog.Ctx(req.Context())
var lastErr error
var lastResp *http.Response
for attempt := 1; attempt <= t.attempts; attempt++ {
if attempt > 1 {
if req.Body != nil && req.GetBody == nil {
break
}
select {
case <-req.Context().Done():
return nil, req.Context().Err()
case <-time.After(time.Duration(attempt-1) * time.Second):
}
}
resp, err := t.attempt(req, attempt)
if err != nil {
if req.Context().Err() != nil {
return nil, req.Context().Err()
}
logger.Warn().Err(err).Int("attempt", attempt).Msg("model call attempt failed")
lastErr = err
lastResp = nil
continue
}
if !retryableStatus(resp.StatusCode) {
return resp, nil
}
logger.Warn().Int("attempt", attempt).Int("status", resp.StatusCode).Msg("model call attempt got a retryable status")
lastErr = nil
lastResp = resp
}
if lastResp != nil {
return lastResp, nil
}
return nil, lastErr
}
func (t *retryTransport) attempt(req *http.Request, attempt int) (*http.Response, error) {
ctx := req.Context()
cancel := context.CancelFunc(func() {})
if t.perAttempt > 0 {
ctx, cancel = context.WithTimeout(ctx, t.perAttempt)
}
defer cancel()
attemptReq := req.Clone(ctx)
if attempt > 1 && req.GetBody != nil {
body, err := req.GetBody()
if err != nil {
return nil, err
}
attemptReq.Body = body
}
resp, err := t.base.RoundTrip(attemptReq)
if err != nil {
return nil, err
}
body, err := io.ReadAll(io.LimitReader(resp.Body, maxModelResponseBytes))
_ = resp.Body.Close()
if err != nil {
return nil, err
}
resp.Body = io.NopCloser(bytes.NewReader(body))
return resp, nil
}
func retryableStatus(status int) bool {
return status == http.StatusRequestTimeout || status == http.StatusTooManyRequests || status >= 500
}
@@ -0,0 +1,136 @@
package agent
import (
"context"
"io"
"net/http"
"net/http/httptest"
"strings"
"sync/atomic"
"testing"
"time"
)
func retryClient(attempts int, perAttempt time.Duration) *http.Client {
return &http.Client{Transport: &retryTransport{base: http.DefaultTransport, attempts: attempts, perAttempt: perAttempt}}
}
func postJSON(t *testing.T, client *http.Client, url string) (*http.Response, error) {
t.Helper()
req, err := http.NewRequestWithContext(context.Background(), http.MethodPost, url, strings.NewReader(`{"model":"m"}`))
if err != nil {
t.Fatal(err)
}
return client.Do(req)
}
func TestRetriesOn5xxAndReplaysBody(t *testing.T) {
var calls atomic.Int32
var bodies []string
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
body, _ := io.ReadAll(r.Body)
bodies = append(bodies, string(body))
if calls.Add(1) < 3 {
w.WriteHeader(http.StatusBadGateway)
return
}
_, _ = w.Write([]byte(`{"ok":true}`))
}))
t.Cleanup(srv.Close)
resp, err := postJSON(t, retryClient(3, time.Second), srv.URL)
if err != nil {
t.Fatal(err)
}
defer resp.Body.Close()
if resp.StatusCode != http.StatusOK {
t.Fatalf("status = %d, want 200", resp.StatusCode)
}
body, _ := io.ReadAll(resp.Body)
if string(body) != `{"ok":true}` {
t.Fatalf("body = %q", body)
}
if len(bodies) != 3 || bodies[2] != `{"model":"m"}` {
t.Fatalf("bodies = %q", bodies)
}
}
func TestDoesNotRetryClientErrors(t *testing.T) {
var calls atomic.Int32
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
calls.Add(1)
w.WriteHeader(http.StatusBadRequest)
}))
t.Cleanup(srv.Close)
resp, err := postJSON(t, retryClient(3, time.Second), srv.URL)
if err != nil {
t.Fatal(err)
}
resp.Body.Close()
if resp.StatusCode != http.StatusBadRequest || calls.Load() != 1 {
t.Fatalf("status = %d calls = %d, want one 400", resp.StatusCode, calls.Load())
}
}
func TestRetriesAStalledBody(t *testing.T) {
var calls atomic.Int32
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if calls.Add(1) == 1 {
w.WriteHeader(http.StatusOK)
w.(http.Flusher).Flush()
<-r.Context().Done()
return
}
_, _ = w.Write([]byte(`{"ok":true}`))
}))
t.Cleanup(srv.Close)
resp, err := postJSON(t, retryClient(2, 200*time.Millisecond), srv.URL)
if err != nil {
t.Fatal(err)
}
defer resp.Body.Close()
body, _ := io.ReadAll(resp.Body)
if string(body) != `{"ok":true}` || calls.Load() != 2 {
t.Fatalf("body = %q calls = %d", body, calls.Load())
}
}
func TestExhaustedRetriesReturnTheLastResponse(t *testing.T) {
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
w.WriteHeader(http.StatusServiceUnavailable)
}))
t.Cleanup(srv.Close)
resp, err := postJSON(t, retryClient(2, time.Second), srv.URL)
if err != nil {
t.Fatal(err)
}
resp.Body.Close()
if resp.StatusCode != http.StatusServiceUnavailable {
t.Fatalf("status = %d, want 503", resp.StatusCode)
}
}
func TestParentContextCancelStopsRetries(t *testing.T) {
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
w.WriteHeader(http.StatusServiceUnavailable)
}))
t.Cleanup(srv.Close)
ctx, cancel := context.WithTimeout(context.Background(), 300*time.Millisecond)
defer cancel()
req, err := http.NewRequestWithContext(ctx, http.MethodPost, srv.URL, strings.NewReader("{}"))
if err != nil {
t.Fatal(err)
}
started := time.Now()
_, err = retryClient(5, time.Second).Do(req)
if err == nil {
t.Fatal("expected a context error")
}
if time.Since(started) > 2*time.Second {
t.Fatalf("retries kept running past parent cancellation (%s)", time.Since(started))
}
}
+5 -1
View File
@@ -30,6 +30,8 @@ type Config struct {
MaxToolCalls int
InvestigationTimeout time.Duration
ModelCallTimeout time.Duration
ModelCallAttempts int
ToolResultBytes int
Workers int
@@ -82,7 +84,9 @@ func LoadConfig() Config {
BotURL: envOr("FUTO_BACKUPS_BOT_URL", "http://localhost:3050"),
GrafanaURL: envOr("GRAFANA_URL", "https://grafana.futostatus.com"),
MaxToolCalls: envIntMin("COLUMBO_MAX_TOOL_CALLS", 16, 1),
InvestigationTimeout: time.Duration(envIntMin("COLUMBO_TIMEOUT_SECONDS", 300, 10)) * time.Second,
InvestigationTimeout: time.Duration(envIntMin("COLUMBO_TIMEOUT_SECONDS", 600, 10)) * time.Second,
ModelCallTimeout: time.Duration(envIntMin("COLUMBO_MODEL_TIMEOUT_SECONDS", 120, 10)) * time.Second,
ModelCallAttempts: envIntMin("COLUMBO_MODEL_ATTEMPTS", 3, 1),
ToolResultBytes: envIntMin("COLUMBO_TOOL_RESULT_BYTES", 12288, 512),
Workers: envIntMin("COLUMBO_WORKERS", 2, 1),
OTLPMetricsEndpoint: os.Getenv("OTLP_METRICS_ENDPOINT"),
+10 -8
View File
@@ -49,14 +49,16 @@ func main() {
}
runner := agent.NewRunner(agent.Config{
OpenRouterURL: cfg.OpenRouterURL,
APIKey: cfg.OpenRouterAPIKey,
Model: cfg.Model,
TriageModel: cfg.TriageModel,
MetricsURL: cfg.MetricsURL,
LogsURL: cfg.LogsURL,
MaxToolCalls: cfg.MaxToolCalls,
ToolResultBytes: cfg.ToolResultBytes,
OpenRouterURL: cfg.OpenRouterURL,
APIKey: cfg.OpenRouterAPIKey,
Model: cfg.Model,
TriageModel: cfg.TriageModel,
MetricsURL: cfg.MetricsURL,
LogsURL: cfg.LogsURL,
MaxToolCalls: cfg.MaxToolCalls,
ToolResultBytes: cfg.ToolResultBytes,
ModelCallTimeout: cfg.ModelCallTimeout,
ModelCallAttempts: cfg.ModelCallAttempts,
})
recorder, err := metrics.Setup(cfg.OTLPMetricsEndpoint, cfg.OTLPMetricsURLPath, cfg.OTLPMetricsInterval)
if err != nil {