mirror of
https://github.com/immich-app/yucca.git
synced 2026-09-30 13:33:00 +08:00
feat(columbo): retry model calls with per-attempt deadlines (#598)
This commit is contained in:
+5
-1
@@ -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 —
|
||||
|
||||
@@ -8,6 +8,7 @@ import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
@@ -45,6 +46,8 @@ type Config struct {
|
||||
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))
|
||||
}
|
||||
}
|
||||
@@ -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"),
|
||||
|
||||
@@ -57,6 +57,8 @@ func main() {
|
||||
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 {
|
||||
|
||||
Reference in New Issue
Block a user