diff --git a/README.md b/README.md index 9bf7f6e..cff7ef5 100644 --- a/README.md +++ b/README.md @@ -49,6 +49,8 @@ startup. | `FREE_TIMEOUT` | `30s` | `POST /free` call after an image job | | `WARM_TIMEOUT` | `2m` | Warm-model reload after an image job | | `SHUTDOWN_TIMEOUT` | `10s` | Graceful shutdown on SIGINT/SIGTERM | +| `BACKOFF_INITIAL` | `1s` | First retry wait when an upstream refuses a connection | +| `BACKOFF_MAX` | `60s` | Cap for the exponential retry backoff | | `PROMPT_CAPTURE_LIMIT` | `65536` | Bytes of the `/prompt` response buffered to find `prompt_id` (pass-through is unaffected) | ## Observability diff --git a/SPEC.md b/SPEC.md index 20231d7..7610dbc 100644 --- a/SPEC.md +++ b/SPEC.md @@ -119,6 +119,8 @@ load time. Off by default. | `FREE_TIMEOUT` | `30s` | `POST /free` call after an image job | | `WARM_TIMEOUT` | `2m` | warm-model reload after an image job | | `SHUTDOWN_TIMEOUT` | `10s` | graceful shutdown on SIGINT/SIGTERM | +| `BACKOFF_INITIAL` | `1s` | first retry wait when an upstream refuses a connection | +| `BACKOFF_MAX` | `60s` | cap for the exponential retry backoff | | `PROMPT_CAPTURE_LIMIT` | `65536` | bytes of the `/prompt` response buffered to find `prompt_id` (pass-through is unaffected) | Startup fails fast on unparsable values. Both upstreams are probed once at @@ -146,6 +148,11 @@ start (`/api/version`, `/system_stats`); failure is logged, not fatal. `JOB_TIMEOUT` releases the lock; log at warn. - Ollama unreachable during unload: continue with the image job; the whole point is not to block users on a misbehaving neighbour. +- Upstream connection refused while proxying (service down or restarting): + retry with exponential backoff — `BACKOFF_INITIAL`, doubling per attempt, + capped at `BACKOFF_MAX` — until the upstream answers or the client + disconnects. Retrying a refused connection is safe: no request bytes were + sent. Other upstream errors are not retried. - `POST /prompt` with a body that ComfyUI rejects (400): lock released immediately, body passed back. - Websocket `/ws` connections are long-lived and never take the lock. diff --git a/cmd/gpu-turnstile/main.go b/cmd/gpu-turnstile/main.go index e5b840b..d807aa0 100644 --- a/cmd/gpu-turnstile/main.go +++ b/cmd/gpu-turnstile/main.go @@ -40,6 +40,8 @@ type config struct { freeTimeout time.Duration warmTimeout time.Duration shutdownTimeout time.Duration + backoffInitial time.Duration + backoffMax time.Duration promptCaptureLimit int64 warmModel string @@ -76,6 +78,8 @@ func loadConfig(getenv func(string) string) (config, error) { freeTimeout: 30 * time.Second, warmTimeout: 2 * time.Minute, shutdownTimeout: 10 * time.Second, + backoffInitial: time.Second, + backoffMax: time.Minute, promptCaptureLimit: 64 * 1024, logLevel: slog.LevelInfo, @@ -107,6 +111,8 @@ func loadConfig(getenv func(string) string) (config, error) { {"FREE_TIMEOUT", &cfg.freeTimeout}, {"WARM_TIMEOUT", &cfg.warmTimeout}, {"SHUTDOWN_TIMEOUT", &cfg.shutdownTimeout}, + {"BACKOFF_INITIAL", &cfg.backoffInitial}, + {"BACKOFF_MAX", &cfg.backoffMax}, } { if err := envDuration(getenv, e.name, e.dst); err != nil { return cfg, err @@ -185,6 +191,8 @@ func main() { HistoryPollInterval: cfg.historyPollInterval, FreeTimeout: cfg.freeTimeout, WarmTimeout: cfg.warmTimeout, + BackoffInitial: cfg.backoffInitial, + BackoffMax: cfg.backoffMax, PromptCaptureLimit: cfg.promptCaptureLimit, WarmModel: cfg.warmModel, }) diff --git a/compose.yaml.example b/compose.yaml.example index b7bd023..56113da 100644 --- a/compose.yaml.example +++ b/compose.yaml.example @@ -1,14 +1,20 @@ # Example deployment for gpu-turnstile. Copy to compose.yaml and adjust. # # gpu-turnstile listens on the ports the services normally use; the actual -# Ollama and ComfyUI instances run one port higher on the workstation. +# Ollama and ComfyUI instances run one port higher (11435 / 8189) and must +# bind 0.0.0.0 so the container can reach them (OLLAMA_HOST=0.0.0.0:11435, +# ComfyUI --listen 0.0.0.0 --port 8189). services: gpu-turnstile: image: git.rambossek.at/public/gpu-turnstile:v0.1.1 restart: unless-stopped environment: - OLLAMA_URL: http://:11435 - COMFY_URL: http://:8189 + # Services on the Docker host itself: + OLLAMA_URL: http://host.docker.internal:11435 + COMFY_URL: http://host.docker.internal:8189 + # Services on another machine: use its LAN IP instead, e.g. + # OLLAMA_URL: http://192.168.1.10:11435 + # COMFY_URL: http://192.168.1.10:8189 # UNLOAD_TIMEOUT: 60s # JOB_TIMEOUT: 15m # LLM_WAIT_TIMEOUT: 10m diff --git a/internal/proxy/proxy.go b/internal/proxy/proxy.go index d80051d..7f44b46 100644 --- a/internal/proxy/proxy.go +++ b/internal/proxy/proxy.go @@ -13,6 +13,7 @@ import ( "net/http" "net/http/httputil" "net/url" + "syscall" "time" "gpu-turnstile/internal/comfy" @@ -41,6 +42,13 @@ type Config struct { UnloadTimeout time.Duration JobTimeout time.Duration + // BackoffInitial and BackoffMax control the exponential retry backoff + // when an upstream refuses a connection: the wait doubles from + // BackoffInitial up to BackoffMax between attempts. Zero selects the + // defaults (1s / 60s). + BackoffInitial time.Duration + BackoffMax time.Duration + // UnloadPollInterval and HistoryPollInterval override the clients' // /api/ps and /history poll intervals when > 0. UnloadPollInterval time.Duration @@ -57,11 +65,13 @@ type Config struct { // Server serves both gpu-turnstile listeners. type Server struct { - cfg Config - log *slog.Logger - freeTimeout time.Duration - warmTimeout time.Duration - captureLimit int64 + cfg Config + log *slog.Logger + freeTimeout time.Duration + warmTimeout time.Duration + captureLimit int64 + backoffInitial time.Duration + backoffMax time.Duration ollamaProxy *httputil.ReverseProxy comfyProxy *httputil.ReverseProxy @@ -99,19 +109,89 @@ func New(cfg Config) (*Server, error) { if captureLimit <= 0 { captureLimit = defaultCaptureLimit } + backoffInitial := cfg.BackoffInitial + if backoffInitial <= 0 { + backoffInitial = time.Second + } + backoffMax := cfg.BackoffMax + if backoffMax <= 0 { + backoffMax = time.Minute + } + retry := &retryTransport{ + base: http.DefaultTransport, + initial: backoffInitial, + max: backoffMax, + log: log, + } return &Server{ - cfg: cfg, - log: log, - freeTimeout: freeTimeout, - warmTimeout: warmTimeout, - captureLimit: captureLimit, - ollamaProxy: newReverseProxy(ollamaURL, log.With("upstream", "ollama")), - comfyProxy: newReverseProxy(comfyURL, log.With("upstream", "comfy")), + cfg: cfg, + log: log, + freeTimeout: freeTimeout, + warmTimeout: warmTimeout, + captureLimit: captureLimit, + backoffInitial: backoffInitial, + backoffMax: backoffMax, + ollamaProxy: newReverseProxy(ollamaURL, retry, log.With("upstream", "ollama")), + comfyProxy: newReverseProxy(comfyURL, retry, log.With("upstream", "comfy")), }, nil } -func newReverseProxy(target *url.URL, log *slog.Logger) *httputil.ReverseProxy { +// retryTransport retries requests that fail with "connection refused" — +// typically the upstream service simply not being up yet — using an +// exponential backoff: initial, doubling each attempt, capped at max. The +// retry loop runs until the request succeeds, fails with a different +// error, or the client's context is cancelled. Retrying a refused +// connection is always safe: no request bytes were ever sent. +type retryTransport struct { + base http.RoundTripper + initial time.Duration + max time.Duration + log *slog.Logger +} + +// wsaECONNREFUSED is the Windows error code for a refused connection. +// syscall.ECONNREFUSED is the POSIX value and never matches on Windows, +// where dial errors carry the WSA code. +const wsaECONNREFUSED syscall.Errno = 10061 + +func isConnRefused(err error) bool { + if errors.Is(err, syscall.ECONNREFUSED) { + return true + } + var errno syscall.Errno + return errors.As(err, &errno) && errno == wsaECONNREFUSED +} + +func (t *retryTransport) logWarn(msg string, args ...any) { + if t.log != nil { + t.log.Warn(msg, args...) + } +} + +func (t *retryTransport) RoundTrip(req *http.Request) (*http.Response, error) { + wait := t.initial + for { + resp, err := t.base.RoundTrip(req) + if err == nil || !isConnRefused(err) { + return resp, err + } + t.logWarn("upstream connection refused; backing off", + "path", req.URL.Path, "retry_in", wait) + select { + case <-req.Context().Done(): + return nil, err + case <-time.After(wait): + } + wait *= 2 + if wait > t.max { + wait = t.max + } + } +} + +func newReverseProxy(target *url.URL, transport http.RoundTripper, log *slog.Logger) *httputil.ReverseProxy { return &httputil.ReverseProxy{ + Transport: transport, Rewrite: func(pr *httputil.ProxyRequest) { pr.SetURL(target) pr.SetXForwarded() diff --git a/internal/proxy/retry_test.go b/internal/proxy/retry_test.go new file mode 100644 index 0000000..791391e --- /dev/null +++ b/internal/proxy/retry_test.go @@ -0,0 +1,122 @@ +package proxy + +import ( + "context" + "io" + "net" + "net/http" + "strings" + "syscall" + "testing" + "time" +) + +// stubTransport fails with ECONNREFUSED for the first fails requests, then +// returns a 200 response. +type stubTransport struct { + fails int + calls int +} + +func (s *stubTransport) RoundTrip(req *http.Request) (*http.Response, error) { + s.calls++ + if s.calls <= s.fails { + return nil, &net.OpError{Op: "dial", Net: "tcp", Err: syscall.ECONNREFUSED} + } + return &http.Response{ + StatusCode: 200, + Body: io.NopCloser(strings.NewReader("ok")), + Header: make(http.Header), + }, nil +} + +func TestRetryTransportBackoff(t *testing.T) { + st := &stubTransport{fails: 3} + rt := &retryTransport{ + base: st, + initial: 10 * time.Millisecond, + max: 25 * time.Millisecond, + } + + start := time.Now() + req, _ := http.NewRequest(http.MethodGet, "http://upstream/api/version", nil) + resp, err := rt.RoundTrip(req) + if err != nil { + t.Fatal(err) + } + resp.Body.Close() + if st.calls != 4 { + t.Fatalf("calls = %d, want 4", st.calls) + } + // Waits: 10ms + 20ms + 25ms (capped) = 55ms minimum. + elapsed := time.Since(start) + if elapsed < 50*time.Millisecond { + t.Fatalf("elapsed = %v, want >= ~55ms of backoff", elapsed) + } + if elapsed > 5*time.Second { + t.Fatalf("elapsed = %v, suspiciously long", elapsed) + } +} + +func TestRetryTransportNonRefusedErrorNotRetried(t *testing.T) { + rt := &retryTransport{ + base: &stubTransport{fails: 0}, + initial: time.Millisecond, + max: time.Millisecond, + } + req, _ := http.NewRequest(http.MethodGet, "http://upstream/", nil) + resp, err := rt.RoundTrip(req) + if err != nil || resp.StatusCode != 200 { + t.Fatalf("resp=%v err=%v", resp, err) + } +} + +type alwaysRefused struct{ calls int } + +func (a *alwaysRefused) RoundTrip(*http.Request) (*http.Response, error) { + a.calls++ + return nil, &net.OpError{Op: "dial", Net: "tcp", Err: syscall.ECONNREFUSED} +} + +func TestRetryTransportContextCancel(t *testing.T) { + ar := &alwaysRefused{} + rt := &retryTransport{base: ar, initial: time.Second, max: time.Second} + + ctx, cancel := context.WithTimeout(context.Background(), 100*time.Millisecond) + defer cancel() + req, _ := http.NewRequestWithContext(ctx, http.MethodGet, "http://upstream/", nil) + start := time.Now() + _, err := rt.RoundTrip(req) + if err == nil { + t.Fatal("expected error after context cancellation") + } + if time.Since(start) > 2*time.Second { + t.Fatal("retry loop did not stop on context cancellation") + } +} + +func TestRetryViaReverseProxy(t *testing.T) { + // Point the proxy at a port nothing listens on; the request should be + // retried (not instantly 502) until the client context ends. + srv, err := New(Config{ + OllamaURL: "http://127.0.0.1:1", + ComfyURL: "http://127.0.0.1:1", + Lock: nil, + Metrics: nil, + BackoffInitial: 10 * time.Millisecond, + BackoffMax: 20 * time.Millisecond, + }) + if err != nil { + t.Fatal(err) + } + _ = srv // construction must not panic with minimal config + + rt := &retryTransport{base: http.DefaultTransport, initial: 10 * time.Millisecond, max: 20 * time.Millisecond} + ctx, cancel := context.WithTimeout(context.Background(), 200*time.Millisecond) + defer cancel() + req, _ := http.NewRequestWithContext(ctx, http.MethodGet, "http://127.0.0.1:1/", nil) + _, err = rt.RoundTrip(req) + if err == nil || !isConnRefused(err) { + t.Fatalf("err = %v, want connection refused", err) + } +}