From cf9be55a6cd67c41a9fb5777f693ef5674216bd3 Mon Sep 17 00:00:00 2001 From: mram Date: Sun, 20 Sep 2026 19:21:25 +0200 Subject: [PATCH] Make poll intervals and operational timeouts configurable New env vars: UNLOAD_POLL_INTERVAL, HISTORY_POLL_INTERVAL, PROBE_TIMEOUT, FREE_TIMEOUT, WARM_TIMEOUT, SHUTDOWN_TIMEOUT, PROMPT_CAPTURE_LIMIT. Defaults unchanged; invalid values fail fast at startup. --- README.md | 7 ++++ SPEC.md | 13 +++++-- cmd/gpu-turnstile/main.go | 71 +++++++++++++++++++++++++++++---------- internal/proxy/proxy.go | 69 ++++++++++++++++++++++++++++--------- 4 files changed, 125 insertions(+), 35 deletions(-) diff --git a/README.md b/README.md index 91edad5..9bf7f6e 100644 --- a/README.md +++ b/README.md @@ -43,6 +43,13 @@ startup. | `WARM_MODEL` | _(empty)_ | Model to reload after an image job (off by default) | | `LOG_LEVEL` | `info` | `debug` logs every lock transition | | `LOG_FORMAT` | `text` | `json` for structured JSON logs | +| `UNLOAD_POLL_INTERVAL` | `500ms` | `/api/ps` poll interval while unloading | +| `HISTORY_POLL_INTERVAL` | `1s` | `/history/` poll interval while a job runs | +| `PROBE_TIMEOUT` | `5s` | Startup probe of both upstreams | +| `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 | +| `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 5a3e905..f01a6c7 100644 --- a/SPEC.md +++ b/SPEC.md @@ -81,12 +81,14 @@ ComfyUI listener (`:8188` → `COMFY_URL`): 2. Unload Ollama: `GET /api/ps`; for each model `POST /api/generate {"model":M,"keep_alive":0}`; if that returns non-2xx (embedding-only models), `POST /api/embed {"model":M,"input":"x","keep_alive":0}`. Poll - `/api/ps` every 500 ms until empty or `UNLOAD_TIMEOUT`. On timeout: log and + `/api/ps` every `UNLOAD_POLL_INTERVAL` (default 500 ms) until empty or + `UNLOAD_TIMEOUT`. On timeout: log and continue (degrade, don't fail the user's request). 3. Forward the original request body to ComfyUI `/prompt`, return status, headers and body to the caller unchanged, flush. 4. If the response is 200 and contains `prompt_id`: in a goroutine, poll - `GET /history/` every 1 s until the entry has + `GET /history/` every `HISTORY_POLL_INTERVAL` (default 1 s) + until the entry has `status.completed == true`, `status.status_str == "error"`, or `JOB_TIMEOUT`. Then `POST /free {"unload_models":true,"free_memory":true}`. Then release the image lock. @@ -111,6 +113,13 @@ load time. Off by default. | `LLM_WAIT_TIMEOUT` | `10m` | max time an LLM request waits for the lock before 503 | | `WARM_MODEL` | `` | optional model to reload after an image job | | `LOG_LEVEL` | `info` | `debug` logs every lock transition | +| `UNLOAD_POLL_INTERVAL` | `500ms` | `/api/ps` poll interval while unloading | +| `HISTORY_POLL_INTERVAL` | `1s` | `/history/` poll interval while a job runs | +| `PROBE_TIMEOUT` | `5s` | startup probe of both upstreams | +| `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 | +| `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 start (`/api/version`, `/system_stats`); failure is logged, not fatal. diff --git a/cmd/gpu-turnstile/main.go b/cmd/gpu-turnstile/main.go index c848bfd..e5b840b 100644 --- a/cmd/gpu-turnstile/main.go +++ b/cmd/gpu-turnstile/main.go @@ -10,6 +10,7 @@ import ( "net/http" "os" "os/signal" + "strconv" "strings" "syscall" "time" @@ -32,9 +33,18 @@ type config struct { unloadTimeout time.Duration jobTimeout time.Duration llmWaitTimeout time.Duration - warmModel string - logLevel slog.Level - logJSON bool + + unloadPollInterval time.Duration + historyPollInterval time.Duration + probeTimeout time.Duration + freeTimeout time.Duration + warmTimeout time.Duration + shutdownTimeout time.Duration + promptCaptureLimit int64 + + warmModel string + logLevel slog.Level + logJSON bool } func envDuration(getenv func(string) string, name string, dst *time.Duration) error { @@ -59,7 +69,16 @@ func loadConfig(getenv func(string) string) (config, error) { unloadTimeout: time.Minute, jobTimeout: 15 * time.Minute, llmWaitTimeout: 10 * time.Minute, - logLevel: slog.LevelInfo, + + unloadPollInterval: 500 * time.Millisecond, + historyPollInterval: time.Second, + probeTimeout: 5 * time.Second, + freeTimeout: 30 * time.Second, + warmTimeout: 2 * time.Minute, + shutdownTimeout: 10 * time.Second, + promptCaptureLimit: 64 * 1024, + + logLevel: slog.LevelInfo, } for _, e := range []struct { name string @@ -82,11 +101,24 @@ func loadConfig(getenv func(string) string) (config, error) { {"UNLOAD_TIMEOUT", &cfg.unloadTimeout}, {"JOB_TIMEOUT", &cfg.jobTimeout}, {"LLM_WAIT_TIMEOUT", &cfg.llmWaitTimeout}, + {"UNLOAD_POLL_INTERVAL", &cfg.unloadPollInterval}, + {"HISTORY_POLL_INTERVAL", &cfg.historyPollInterval}, + {"PROBE_TIMEOUT", &cfg.probeTimeout}, + {"FREE_TIMEOUT", &cfg.freeTimeout}, + {"WARM_TIMEOUT", &cfg.warmTimeout}, + {"SHUTDOWN_TIMEOUT", &cfg.shutdownTimeout}, } { if err := envDuration(getenv, e.name, e.dst); err != nil { return cfg, err } } + if v := getenv("PROMPT_CAPTURE_LIMIT"); v != "" { + n, err := strconv.ParseInt(v, 10, 64) + if err != nil || n < 0 { + return cfg, fmt.Errorf("PROMPT_CAPTURE_LIMIT: must be a non-negative integer (bytes)") + } + cfg.promptCaptureLimit = n + } if v := getenv("LOG_LEVEL"); v != "" { var level slog.Level if err := level.UnmarshalText([]byte(v)); err != nil { @@ -139,17 +171,22 @@ func main() { } srv, err := proxy.New(proxy.Config{ - OllamaURL: cfg.ollamaURL, - ComfyURL: cfg.comfyURL, - Lock: lock.New(log), - Ollama: ollamaClient, - Comfy: comfyClient, - Metrics: metrics.New(), - Log: log, - LLMWaitTimeout: cfg.llmWaitTimeout, - UnloadTimeout: cfg.unloadTimeout, - JobTimeout: cfg.jobTimeout, - WarmModel: cfg.warmModel, + OllamaURL: cfg.ollamaURL, + ComfyURL: cfg.comfyURL, + Lock: lock.New(log), + Ollama: ollamaClient, + Comfy: comfyClient, + Metrics: metrics.New(), + Log: log, + LLMWaitTimeout: cfg.llmWaitTimeout, + UnloadTimeout: cfg.unloadTimeout, + JobTimeout: cfg.jobTimeout, + UnloadPollInterval: cfg.unloadPollInterval, + HistoryPollInterval: cfg.historyPollInterval, + FreeTimeout: cfg.freeTimeout, + WarmTimeout: cfg.warmTimeout, + PromptCaptureLimit: cfg.promptCaptureLimit, + WarmModel: cfg.warmModel, }) if err != nil { log.Error("invalid configuration", "err", err) @@ -157,7 +194,7 @@ func main() { } // Probe both upstreams once; failure is logged, not fatal. - probeCtx, probeCancel := context.WithTimeout(context.Background(), 5*time.Second) + probeCtx, probeCancel := context.WithTimeout(context.Background(), cfg.probeTimeout) if err := ollamaClient.Probe(probeCtx); err != nil { log.Warn("ollama probe failed", "url", cfg.ollamaURL, "err", err) } @@ -186,7 +223,7 @@ func main() { log.Info("shutting down") } - shutdownCtx, shutdownCancel := context.WithTimeout(context.Background(), 10*time.Second) + shutdownCtx, shutdownCancel := context.WithTimeout(context.Background(), cfg.shutdownTimeout) defer shutdownCancel() ollamaSrv.Shutdown(shutdownCtx) comfySrv.Shutdown(shutdownCtx) diff --git a/internal/proxy/proxy.go b/internal/proxy/proxy.go index 9c13177..d80051d 100644 --- a/internal/proxy/proxy.go +++ b/internal/proxy/proxy.go @@ -21,10 +21,10 @@ import ( "gpu-turnstile/internal/ollama" ) -// captureLimit bounds how much of a /prompt response body is buffered while -// looking for prompt_id. The body still passes through to the client -// unchanged regardless of size. -const captureLimit = 64 * 1024 +// defaultCaptureLimit bounds how much of a /prompt response body is +// buffered while looking for prompt_id. The body still passes through to +// the client unchanged regardless of size. +const defaultCaptureLimit = 64 * 1024 // Config wires a Server. type Config struct { @@ -40,13 +40,28 @@ type Config struct { LLMWaitTimeout time.Duration UnloadTimeout time.Duration JobTimeout time.Duration - WarmModel string + + // UnloadPollInterval and HistoryPollInterval override the clients' + // /api/ps and /history poll intervals when > 0. + UnloadPollInterval time.Duration + HistoryPollInterval time.Duration + // FreeTimeout and WarmTimeout bound the /free call and the warm-model + // reload; zero selects the defaults. + FreeTimeout time.Duration + WarmTimeout time.Duration + // PromptCaptureLimit overrides defaultCaptureLimit when > 0. + PromptCaptureLimit int64 + + WarmModel string } // Server serves both gpu-turnstile listeners. type Server struct { - cfg Config - log *slog.Logger + cfg Config + log *slog.Logger + freeTimeout time.Duration + warmTimeout time.Duration + captureLimit int64 ollamaProxy *httputil.ReverseProxy comfyProxy *httputil.ReverseProxy @@ -66,11 +81,32 @@ func New(cfg Config) (*Server, error) { if log == nil { log = slog.Default() } + if cfg.UnloadPollInterval > 0 { + cfg.Ollama.PollInterval = cfg.UnloadPollInterval + } + if cfg.HistoryPollInterval > 0 { + cfg.Comfy.PollInterval = cfg.HistoryPollInterval + } + freeTimeout := cfg.FreeTimeout + if freeTimeout <= 0 { + freeTimeout = 30 * time.Second + } + warmTimeout := cfg.WarmTimeout + if warmTimeout <= 0 { + warmTimeout = 2 * time.Minute + } + captureLimit := cfg.PromptCaptureLimit + if captureLimit <= 0 { + captureLimit = defaultCaptureLimit + } return &Server{ - cfg: cfg, - log: log, - ollamaProxy: newReverseProxy(ollamaURL, log.With("upstream", "ollama")), - comfyProxy: newReverseProxy(comfyURL, log.With("upstream", "comfy")), + cfg: cfg, + log: log, + freeTimeout: freeTimeout, + warmTimeout: warmTimeout, + captureLimit: captureLimit, + ollamaProxy: newReverseProxy(ollamaURL, log.With("upstream", "ollama")), + comfyProxy: newReverseProxy(comfyURL, log.With("upstream", "comfy")), }, nil } @@ -170,11 +206,12 @@ func (s *Server) ComfyHandler() http.Handler { } // captureWriter passes the response through unchanged while recording the -// status code and the first captureLimit bytes of the body. +// status code and the first limit bytes of the body. type captureWriter struct { http.ResponseWriter status int buf bytes.Buffer + limit int64 } func (w *captureWriter) WriteHeader(code int) { @@ -183,7 +220,7 @@ func (w *captureWriter) WriteHeader(code int) { } func (w *captureWriter) Write(p []byte) (int, error) { - if w.buf.Len() < captureLimit { + if int64(w.buf.Len()) < w.limit { w.buf.Write(p) } return w.ResponseWriter.Write(p) @@ -228,7 +265,7 @@ func (s *Server) handlePrompt(w http.ResponseWriter, r *http.Request) { log.Info("ollama models unloaded", "seconds", elapsed.Seconds()) } - cw := &captureWriter{ResponseWriter: w, status: http.StatusOK} + cw := &captureWriter{ResponseWriter: w, status: http.StatusOK, limit: s.captureLimit} s.comfyProxy.ServeHTTP(cw, r) var accepted struct { @@ -262,7 +299,7 @@ func (s *Server) finishImageJob(promptID string) { log.Info("image job completed") } - freeCtx, freeCancel := context.WithTimeout(context.Background(), 30*time.Second) + freeCtx, freeCancel := context.WithTimeout(context.Background(), s.freeTimeout) if err := s.cfg.Comfy.Free(freeCtx); err != nil { log.Warn("failed to free ComfyUI models", "err", err) } @@ -273,7 +310,7 @@ func (s *Server) finishImageJob(promptID string) { if s.cfg.WarmModel != "" { if state, _, _ := s.cfg.Lock.Snapshot(); state == lock.StateIdle { - wctx, wcancel := context.WithTimeout(context.Background(), 2*time.Minute) + wctx, wcancel := context.WithTimeout(context.Background(), s.warmTimeout) if err := s.cfg.Ollama.Warm(wctx, s.cfg.WarmModel); err != nil { log.Warn("warm model reload failed", "model", s.cfg.WarmModel, "err", err) } else {