diff --git a/internal/providers/providerio/providerio.go b/internal/providers/providerio/providerio.go index e1adeb5cd..6b809607c 100644 --- a/internal/providers/providerio/providerio.go +++ b/internal/providers/providerio/providerio.go @@ -14,6 +14,7 @@ import ( "runtime" "strconv" "strings" + "sync" "time" "github.com/Gitlawb/zero/internal/trace" @@ -89,6 +90,22 @@ func ContentStallTimeout(idleTimeout time.Duration) time.Duration { // entirely (streams may then hang until the HTTP/transport layer gives up). const streamIdleTimeoutEnv = "ZERO_STREAM_IDLE_TIMEOUT" +// bareSecondsDuration converts a positive second count into a time.Duration. +// time.Duration is int64 nanoseconds, so a count that passes strconv.Atoi can +// still overflow when multiplied by time.Second and wrap to 0, which would +// disable a timeout that should have stayed at its default. ok is false when +// the product does not fit. +func bareSecondsDuration(secs int) (time.Duration, bool) { + if secs <= 0 { + return 0, false + } + const maxSecs = int64(math.MaxInt64) / int64(time.Second) + if int64(secs) > maxSecs { + return 0, false + } + return time.Duration(secs) * time.Second, true +} + // ResolveStreamIdleTimeout selects the effective stream idle timeout. Precedence: // an explicit positive option (e.g. set by a test) wins; otherwise the // ZERO_STREAM_IDLE_TIMEOUT env override if set and valid; otherwise @@ -105,15 +122,59 @@ func ResolveStreamIdleTimeout(option time.Duration) time.Duration { if d, err := time.ParseDuration(raw); err == nil && d > 0 { return d } - if secs, err := strconv.Atoi(raw); err == nil && secs > 0 { - return time.Duration(secs) * time.Second + if secs, err := strconv.Atoi(raw); err == nil { + if d, ok := bareSecondsDuration(secs); ok { + return d + } } - // Unparseable / non-positive: fall through to the default rather than - // silently disabling the watchdog on a typo. + // Unparseable / non-positive / overflow: fall through to the default + // rather than silently disabling the watchdog. } return DefaultStreamIdleTimeout } +// DefaultResponseHeaderTimeout is how long the shared HTTP transport waits for +// a response header after the request is written. 120s, not 60s: a slow cloud +// proxy (e.g. ollama `*:cloud`) can withhold its 200 response header until the +// upstream model emits a first token, so a 60s cap risked aborting a +// legitimately-slow-but-alive request. 120s still bounds a truly dead reused +// connection (which never responds) while tolerating slow header delivery; slow +// first tokens after the header are covered by the idle + content-stall +// watchdogs. +const DefaultResponseHeaderTimeout = 120 * time.Second + +// responseHeaderTimeoutEnv is the global override for the response header +// timeout. It accepts the same forms as ZERO_STREAM_IDLE_TIMEOUT: a Go duration +// ("5m", "300s", "90s") or a bare number of seconds ("300"). A value of "0", +// "off", "none", or "disabled" removes the limit entirely (a connection that +// never answers may then wait until the request context ends). Useful when a +// local model server needs longer than DefaultResponseHeaderTimeout to produce +// the first byte, for example a cold model load on a throttled Ollama. +const responseHeaderTimeoutEnv = "ZERO_RESPONSE_HEADER_TIMEOUT" + +// ResolveResponseHeaderTimeout selects the effective response header timeout: +// the ZERO_RESPONSE_HEADER_TIMEOUT env override if set and valid, otherwise +// DefaultResponseHeaderTimeout. A returned value <= 0 means no limit. +func ResolveResponseHeaderTimeout() time.Duration { + if raw := strings.TrimSpace(os.Getenv(responseHeaderTimeoutEnv)); raw != "" { + switch strings.ToLower(raw) { + case "0", "off", "none", "disabled": + return 0 + } + if d, err := time.ParseDuration(raw); err == nil && d > 0 { + return d + } + if secs, err := strconv.Atoi(raw); err == nil { + if d, ok := bareSecondsDuration(secs); ok { + return d + } + } + // Unparseable / non-positive / overflow: fall through to the default + // rather than silently removing the limit. + } + return DefaultResponseHeaderTimeout +} + // NormalizeBaseURL trims trailing slashes and validates an HTTP API base URL. func NormalizeBaseURL(baseURL string, defaultBaseURL string, label string) (string, error) { baseURL = strings.TrimSpace(baseURL) @@ -127,7 +188,18 @@ func NormalizeBaseURL(baseURL string, defaultBaseURL string, label string) (stri return baseURL, nil } -// sharedHTTPClient is the process-wide client used when a provider supplies none. +// sharedHTTPClients holds one stall-hardened client per resolved response-header +// timeout. The timeout is read when HTTPClient is called, not at package init: +// a test (or a process that sets ZERO_RESPONSE_HEADER_TIMEOUT before building +// a provider) must see that value on the transport it actually dials. The +// transport field is never rewritten after creation, so concurrent requests +// keep the conn pool of the timeout they resolved. +var sharedHTTPClients = struct { + mu sync.Mutex + byTimeout map[time.Duration]*http.Client +}{} + +// stallHardenedClient builds the client used when a provider supplies none. // It tunes the default transport to defeat the stale-pooled-connection hang: Go // keeps idle keep-alive connections in a pool, and a later request can reuse one // the server/NAT has silently dropped. Because the model call is a POST (non- @@ -159,15 +231,11 @@ func NormalizeBaseURL(baseURL string, defaultBaseURL string, label string) (stri // to the minutes-long stalls this avoids — and this doesn't touch // Linux/Windows, where the underlying OS doesn't keep dead/degraded // pooled connections around as long. -var sharedHTTPClient = func() *http.Client { +func stallHardenedClient(headerTimeout time.Duration) *http.Client { transport := http.DefaultTransport.(*http.Transport).Clone() - // 120s, not 60s: a slow cloud proxy (e.g. ollama `*:cloud`) can withhold its - // 200 response header until the upstream model emits a first token, so a 60s cap - // risked aborting a legitimately-slow-but-alive request. 120s still bounds a - // truly dead reused connection (which never responds) while tolerating slow - // header delivery; slow first tokens after the header are covered by the idle + - // content-stall watchdogs. - transport.ResponseHeaderTimeout = 120 * time.Second + // DefaultResponseHeaderTimeout (120s) unless ZERO_RESPONSE_HEADER_TIMEOUT + // overrides it; see the constant for why the default is 120s and not 60s. + transport.ResponseHeaderTimeout = headerTimeout transport.IdleConnTimeout = 30 * time.Second // Periodically close idle connections to prevent stale HTTP/2 // connections from causing PROTOCOL_ERROR on the next request. @@ -183,7 +251,22 @@ var sharedHTTPClient = func() *http.Client { } transport.DisableKeepAlives = runtime.GOOS == "darwin" return &http.Client{Transport: transport} -}() +} + +func sharedStallClient() *http.Client { + timeout := ResolveResponseHeaderTimeout() + sharedHTTPClients.mu.Lock() + defer sharedHTTPClients.mu.Unlock() + if sharedHTTPClients.byTimeout == nil { + sharedHTTPClients.byTimeout = make(map[time.Duration]*http.Client) + } + if client := sharedHTTPClients.byTimeout[timeout]; client != nil { + return client + } + client := stallHardenedClient(timeout) + sharedHTTPClients.byTimeout[timeout] = client + return client +} const defaultIdleConnCloseInterval = 30 * time.Second @@ -221,7 +304,7 @@ func HTTPClient(client *http.Client) *http.Client { if client != nil { return client } - return sharedHTTPClient + return sharedStallClient() } // SendEvent writes a provider event without blocking cancellation cleanup. diff --git a/internal/providers/providerio/providerio_test.go b/internal/providers/providerio/providerio_test.go index 284226587..bfb45e0ed 100644 --- a/internal/providers/providerio/providerio_test.go +++ b/internal/providers/providerio/providerio_test.go @@ -340,6 +340,10 @@ func TestStreamTimeoutMessage(t *testing.T) { // response-header wait + shorter idle-conn reuse) that defeats the macOS stale- // pooled-connection hang; an explicit client is returned untouched. func TestHTTPClientReturnsStallHardenedSharedClient(t *testing.T) { + // Package init must not bake the timeout. An ambient + // ZERO_RESPONSE_HEADER_TIMEOUT would otherwise make this assertion + // depend on whatever the process inherited. + t.Setenv("ZERO_RESPONSE_HEADER_TIMEOUT", "") got := HTTPClient(nil) if got == nil { t.Fatal("HTTPClient(nil) returned nil") @@ -374,6 +378,66 @@ func TestHTTPClientReturnsStallHardenedSharedClient(t *testing.T) { } } +// The shared transport must enforce the timeout ResolveResponseHeaderTimeout +// returns. Checking the resolver alone stays green if the client still carries +// the value captured at init. +func TestHTTPClientTransportUsesResolvedHeaderTimeout(t *testing.T) { + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + delay, err := time.ParseDuration(r.URL.Query().Get("delay")) + if err != nil { + http.Error(w, "bad delay", http.StatusBadRequest) + return + } + time.Sleep(delay) + w.WriteHeader(http.StatusNoContent) + })) + t.Cleanup(srv.Close) + + t.Setenv("ZERO_RESPONSE_HEADER_TIMEOUT", "200ms") + short := HTTPClient(nil) + tr, ok := short.Transport.(*http.Transport) + if !ok { + t.Fatalf("transport type = %T, want *http.Transport", short.Transport) + } + if tr.ResponseHeaderTimeout != ResolveResponseHeaderTimeout() || tr.ResponseHeaderTimeout != 200*time.Millisecond { + t.Fatalf("ResponseHeaderTimeout = %v, resolver = %v, want 200ms", tr.ResponseHeaderTimeout, ResolveResponseHeaderTimeout()) + } + started := time.Now() + resp, err := short.Get(srv.URL + "?delay=1s") + if resp != nil { + _ = resp.Body.Close() + } + if err == nil { + t.Fatal("slow header succeeded; the transport did not apply the 200ms timeout") + } + if !strings.Contains(err.Error(), "timeout awaiting response headers") { + t.Fatalf("error = %v, want the transport's response-header timeout", err) + } + if elapsed := time.Since(started); elapsed >= 800*time.Millisecond { + t.Fatalf("header wait took %v; the transport waited on the handler instead of the resolved timeout", elapsed) + } + + t.Setenv("ZERO_RESPONSE_HEADER_TIMEOUT", "2s") + long := HTTPClient(nil) + if long == short { + t.Fatal("a new resolved timeout reused the transport that still has the previous timeout") + } + if got := long.Transport.(*http.Transport).ResponseHeaderTimeout; got != 2*time.Second { + t.Fatalf("ResponseHeaderTimeout = %v, want 2s", got) + } + resp, err = long.Get(srv.URL + "?delay=50ms") + if err != nil { + t.Fatalf("header within the resolved timeout failed: %v", err) + } + _ = resp.Body.Close() + if resp.StatusCode != http.StatusNoContent { + t.Fatalf("status = %d, want 204", resp.StatusCode) + } + if again := HTTPClient(nil); again != long { + t.Fatal("the same resolved timeout must keep one shared client") + } +} + // startIdleConnCloser periodically closes idle pooled connections so stale HTTP/2 // connections are not reused across long idle periods. func TestStartIdleConnCloserClosesIdleConnections(t *testing.T) { diff --git a/internal/providers/providerio/response_header_timeout_resolve_test.go b/internal/providers/providerio/response_header_timeout_resolve_test.go new file mode 100644 index 000000000..a79b9f49f --- /dev/null +++ b/internal/providers/providerio/response_header_timeout_resolve_test.go @@ -0,0 +1,73 @@ +package providerio + +import ( + "testing" + "time" +) + +func TestResolveResponseHeaderTimeout(t *testing.T) { + const env = "ZERO_RESPONSE_HEADER_TIMEOUT" + + t.Run("default when env is unset or empty", func(t *testing.T) { + t.Setenv(env, "") + if got := ResolveResponseHeaderTimeout(); got != DefaultResponseHeaderTimeout { + t.Fatalf("got %v, want default %v", got, DefaultResponseHeaderTimeout) + } + }) + + t.Run("default keeps the value that was previously hardcoded", func(t *testing.T) { + if DefaultResponseHeaderTimeout != 120*time.Second { + t.Fatalf("default is %v; this override must not change the 120s default", DefaultResponseHeaderTimeout) + } + }) + + t.Run("env Go duration", func(t *testing.T) { + t.Setenv(env, "240s") + if got := ResolveResponseHeaderTimeout(); got != 240*time.Second { + t.Fatalf("got %v, want 240s", got) + } + t.Setenv(env, "5m") + if got := ResolveResponseHeaderTimeout(); got != 5*time.Minute { + t.Fatalf("got %v, want 5m", got) + } + }) + + t.Run("env bare seconds", func(t *testing.T) { + t.Setenv(env, "300") + if got := ResolveResponseHeaderTimeout(); got != 300*time.Second { + t.Fatalf("got %v, want 300s", got) + } + }) + + t.Run("env value is trimmed", func(t *testing.T) { + t.Setenv(env, " 90s ") + if got := ResolveResponseHeaderTimeout(); got != 90*time.Second { + t.Fatalf("got %v, want 90s", got) + } + }) + + t.Run("env removes the limit", func(t *testing.T) { + for _, value := range []string{"0", "off", "none", "disabled", "OFF", "Disabled"} { + t.Setenv(env, value) + if got := ResolveResponseHeaderTimeout(); got != 0 { + t.Fatalf("%q: got %v, want 0 (no limit)", value, got) + } + } + }) + + t.Run("bare seconds that overflow time.Duration keep the default", func(t *testing.T) { + t.Setenv(env, "36028797018963968") + if got := ResolveResponseHeaderTimeout(); got != DefaultResponseHeaderTimeout { + t.Fatalf("got %v, want default %v (overflow must not become a zero timeout)", got, DefaultResponseHeaderTimeout) + } + }) + + t.Run("invalid env falls back to default, not unlimited", func(t *testing.T) { + for _, value := range []string{"banana", "-5s", "-1", "1.5x"} { + t.Setenv(env, value) + if got := ResolveResponseHeaderTimeout(); got != DefaultResponseHeaderTimeout { + t.Fatalf("%q: got %v, want default %v (a typo must not remove the limit)", value, got, DefaultResponseHeaderTimeout) + } + } + }) +} diff --git a/internal/providers/providerio/stream_idle_resolve_test.go b/internal/providers/providerio/stream_idle_resolve_test.go index 13702b1fd..a240bac8e 100644 --- a/internal/providers/providerio/stream_idle_resolve_test.go +++ b/internal/providers/providerio/stream_idle_resolve_test.go @@ -49,6 +49,13 @@ func TestResolveStreamIdleTimeout(t *testing.T) { } }) + t.Run("bare seconds that overflow time.Duration keep the default", func(t *testing.T) { + t.Setenv(env, "36028797018963968") + if got := ResolveStreamIdleTimeout(0); got != DefaultStreamIdleTimeout { + t.Fatalf("got %v, want default %v (overflow must not disable the watchdog)", got, DefaultStreamIdleTimeout) + } + }) + t.Run("invalid env falls back to default, not disabled", func(t *testing.T) { t.Setenv(env, "banana") if got := ResolveStreamIdleTimeout(0); got != DefaultStreamIdleTimeout {