// Package proxy contains the HTTP handlers for both gpu-turnstile listeners: // reverse proxies to Ollama and ComfyUI with GPU lock arbitration in front // of the endpoints that load models. package proxy import ( "bufio" "bytes" "context" "crypto/tls" "encoding/json" "errors" "fmt" "io" "log/slog" "net" "net/http" "net/http/httputil" "net/url" "os" "strconv" "strings" "time" "gpu-turnstile/internal/comfy" "gpu-turnstile/internal/lock" "gpu-turnstile/internal/metrics" "gpu-turnstile/internal/ollama" ) // 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 { OllamaURL string ComfyURL string Lock *lock.Lock Ollama *ollama.Client Comfy *comfy.Client Metrics *metrics.Metrics Log *slog.Logger // LogColor enables ANSI colors in per-request log lines. Ignored when // the log level is above INFO (request lines are not emitted at all). LogColor bool // LogWriter receives the colored per-request lines; nil means stderr. LogWriter io.Writer LLMWaitTimeout time.Duration 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 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 logWriter io.Writer freeTimeout time.Duration warmTimeout time.Duration captureLimit int64 backoffInitial time.Duration backoffMax time.Duration ollamaProxy *httputil.ReverseProxy comfyProxy *httputil.ReverseProxy } // New builds a Server, validating the upstream URLs. func New(cfg Config) (*Server, error) { ollamaURL, err := url.Parse(cfg.OllamaURL) if err != nil || ollamaURL.Scheme == "" || ollamaURL.Host == "" { return nil, fmt.Errorf("invalid OLLAMA_URL %q", cfg.OllamaURL) } comfyURL, err := url.Parse(cfg.ComfyURL) if err != nil || comfyURL.Scheme == "" || comfyURL.Host == "" { return nil, fmt.Errorf("invalid COMFY_URL %q", cfg.ComfyURL) } log := cfg.Log 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 } 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, logWriter: cfg.LogWriter, 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 } // retryTransport retries requests whose failure means the upstream never // saw them — any dial-phase error (connection refused, dial timeout, DNS // failure), TLS handshake errors — plus 5xx responses when the request // body can be replayed (GETs and requests with GetBody set). The wait // doubles from initial up to max between attempts. The loop runs until // the request succeeds, fails in a non-retryable way, or the client's // context is cancelled. type retryTransport struct { base http.RoundTripper initial time.Duration max time.Duration log *slog.Logger } // shouldRetry reports whether a RoundTrip error means the request never // reached the upstream application and is therefore safe to send again. func shouldRetry(err error) bool { // Dial-phase failures: refused, timeout, unreachable, DNS (wrapped). var opErr *net.OpError if errors.As(err, &opErr) && opErr.Op == "dial" { return true } var dnsErr *net.DNSError if errors.As(err, &dnsErr) { return true } // TLS handshake failures: the HTTP request was never written. var recordErr tls.RecordHeaderError if errors.As(err, &recordErr) { return true } var certErr *tls.CertificateVerificationError if errors.As(err, &certErr) { return true } var alertErr tls.AlertError if errors.As(err, &alertErr) { return true } return false } // replayable reports whether the request body can be sent again. Bodies // streamed from the client (GetBody == nil) cannot, so 5xx responses to // POSTs are not retried: the upstream may have partially processed them, // and re-sending could duplicate work (e.g. a second ComfyUI prompt). func replayable(req *http.Request) bool { return req.Body == nil || req.Body == http.NoBody || req.GetBody != nil } func (t *retryTransport) logAt(level slog.Level, msg string, args ...any) { if t.log != nil && t.log.Enabled(context.Background(), level) { t.log.Log(context.Background(), level, msg, args...) } } func (t *retryTransport) RoundTrip(req *http.Request) (*http.Response, error) { wait := t.initial attempt := 0 for { resp, err := t.base.RoundTrip(req) switch { case err != nil && !shouldRetry(err): return nil, err case err == nil && (resp.StatusCode < 500 || !replayable(req)): return resp, nil } // Retryable failure: a transport error, or a 5xx response. var reason string if err != nil { reason = err.Error() } else { reason = resp.Status io.Copy(io.Discard, resp.Body) resp.Body.Close() if req.GetBody != nil { if body, berr := req.GetBody(); berr == nil { req.Body = body } } } attempt++ // One WARN per outage episode; subsequent attempts at INFO. level := slog.LevelInfo if attempt == 1 { level = slog.LevelWarn } t.logAt(level, "upstream unavailable; retrying with backoff", "path", req.URL.Path, "reason", reason, "attempt", attempt, "retry_in", wait) select { case <-req.Context().Done(): if err != nil { return nil, err } return nil, req.Context().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() }, // Flush after every write so NDJSON/SSE streams and websocket // upgrades pass through unbuffered. FlushInterval: -1, ErrorHandler: func(w http.ResponseWriter, r *http.Request, err error) { log.Warn("upstream error", "path", r.URL.Path, "err", err) http.Error(w, "upstream unavailable", http.StatusBadGateway) }, } } func (s *Server) writeHealthz(w http.ResponseWriter) { state, n, pending := s.cfg.Lock.Snapshot() w.Header().Set("Content-Type", "application/json") json.NewEncoder(w).Encode(map[string]any{ "state": state, "llm_inflight": n, "image_pending": pending, }) } func (s *Server) writeMetrics(w http.ResponseWriter) { state, n, pending := s.cfg.Lock.Snapshot() w.Header().Set("Content-Type", "text/plain; version=0.0.4") s.cfg.Metrics.Render(w, string(state), n, pending) } // ANSI colors for per-request log lines. const ( ansiReset = "\x1b[0m" ansiCyan = "\x1b[36m" ansiGreen = "\x1b[32m" ansiYellow = "\x1b[33m" ansiRed = "\x1b[31m" ) // statusRecorder remembers the response status while passing everything // through, including streaming flushes and websocket hijacks. type statusRecorder struct { http.ResponseWriter status int } func (r *statusRecorder) WriteHeader(code int) { r.status = code r.ResponseWriter.WriteHeader(code) } func (r *statusRecorder) Flush() { if f, ok := r.ResponseWriter.(http.Flusher); ok { f.Flush() } } func (r *statusRecorder) Hijack() (net.Conn, *bufio.ReadWriter, error) { h, ok := r.ResponseWriter.(http.Hijacker) if !ok { return nil, nil, errors.New("response writer does not support hijacking") } return h.Hijack() } func (r *statusRecorder) Unwrap() http.ResponseWriter { return r.ResponseWriter } // reqLine emits one request log line. With color enabled slog cannot be // used: its text handler escapes the ANSI sequences, so the line is // written to stderr directly in the same key=value shape. Without color // it is a plain slog INFO line. func (s *Server) reqLine(log *slog.Logger, code, line string, attrs ...any) { if !s.cfg.LogColor { log.Info(line, attrs...) return } var sb strings.Builder sb.WriteString("time=" + time.Now().Format("2006-01-02T15:04:05.000Z07:00") + " level=INFO ") sb.WriteString(code + line + ansiReset) for i := 0; i+1 < len(attrs); i += 2 { fmt.Fprintf(&sb, " %v=%v", attrs[i], attrs[i+1]) } sb.WriteByte('\n') w := s.logWriter if w == nil { w = os.Stderr } io.WriteString(w, sb.String()) } // logRequests logs one line per incoming request and one per completed // response at INFO level, colored when enabled: cyan "-->" for incoming, // green/yellow/red "<--" for responses by status class. At log levels // above INFO it is a pass-through. func (s *Server) logRequests(listener string, next http.Handler) http.Handler { log := s.log.With("listener", listener) return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { if !log.Enabled(r.Context(), slog.LevelInfo) { next.ServeHTTP(w, r) return } start := time.Now() s.reqLine(log, ansiCyan, "--> "+r.Method+" "+r.URL.RequestURI(), "listener", listener, "remote", r.RemoteAddr) rec := &statusRecorder{ResponseWriter: w, status: http.StatusOK} next.ServeHTTP(rec, r) code := ansiGreen switch { case rec.status >= 500: code = ansiRed case rec.status >= 400: code = ansiYellow } s.reqLine(log, code, "<-- "+strconv.Itoa(rec.status)+" "+r.Method+" "+r.URL.RequestURI(), "listener", listener, "ms", time.Since(start).Milliseconds()) }) } // llmPaths are the Ollama endpoints that load models into VRAM and therefore // take the LLM lock. Everything else passes through unlocked. var llmPaths = map[string]bool{ "/api/generate": true, "/api/chat": true, "/api/embed": true, "/api/embeddings": true, "/v1/chat/completions": true, "/v1/completions": true, "/v1/embeddings": true, } func isLLMRequest(r *http.Request) bool { return r.Method == http.MethodPost && llmPaths[r.URL.Path] } // OllamaHandler serves the Ollama-facing listener. func (s *Server) OllamaHandler() http.Handler { return s.logRequests("ollama", http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { switch r.URL.Path { case "/healthz": s.writeHealthz(w) return case "/metrics": s.writeMetrics(w) return } if !isLLMRequest(r) { s.ollamaProxy.ServeHTTP(w, r) return } start := time.Now() wctx, cancel := context.WithTimeout(r.Context(), s.cfg.LLMWaitTimeout) err := s.cfg.Lock.AcquireLLM(wctx) cancel() s.cfg.Metrics.ObserveLockWait("llm", time.Since(start).Seconds()) if err != nil { if errors.Is(err, context.DeadlineExceeded) && r.Context().Err() == nil { http.Error(w, "GPU busy: timed out waiting for the lock", http.StatusServiceUnavailable) } return } defer s.cfg.Lock.ReleaseLLM() s.ollamaProxy.ServeHTTP(w, r) })) } // ComfyHandler serves the ComfyUI-facing listener. func (s *Server) ComfyHandler() http.Handler { return s.logRequests("comfy", http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { if r.URL.Path == "/healthz" { s.writeHealthz(w) return } if r.Method == http.MethodPost && r.URL.Path == "/prompt" { s.handlePrompt(w, r) return } s.comfyProxy.ServeHTTP(w, r) })) } // captureWriter passes the response through unchanged while recording the // 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) { w.status = code w.ResponseWriter.WriteHeader(code) } func (w *captureWriter) Write(p []byte) (int, error) { if int64(w.buf.Len()) < w.limit { w.buf.Write(p) } return w.ResponseWriter.Write(p) } func (w *captureWriter) Flush() { if f, ok := w.ResponseWriter.(http.Flusher); ok { f.Flush() } } func (w *captureWriter) Unwrap() http.ResponseWriter { return w.ResponseWriter } // handlePrompt implements the image job flow from the spec: acquire the // image lock, unload Ollama, forward to ComfyUI, then track the job in the // background and free VRAM before releasing the lock. func (s *Server) handlePrompt(w http.ResponseWriter, r *http.Request) { log := s.log.With("op", "image") start := time.Now() if err := s.cfg.Lock.AcquireImage(r.Context()); err != nil { if errors.Is(err, context.DeadlineExceeded) { http.Error(w, "GPU busy: timed out waiting for the lock", http.StatusServiceUnavailable) } return } s.cfg.Metrics.ObserveLockWait("image", time.Since(start).Seconds()) log.Info("image lock acquired") uctx, ucancel := context.WithTimeout(r.Context(), s.cfg.UnloadTimeout) elapsed, uerr := s.cfg.Ollama.UnloadAll(uctx) ucancel() s.cfg.Metrics.ObserveUnload(elapsed.Seconds()) switch { case r.Context().Err() != nil: s.cfg.Lock.ReleaseImage() return case uerr != nil: // Degrade, don't fail the user's request on a misbehaving neighbour. log.Warn("ollama unload incomplete; continuing", "err", uerr) default: log.Info("ollama models unloaded", "seconds", elapsed.Seconds()) } cw := &captureWriter{ResponseWriter: w, status: http.StatusOK, limit: s.captureLimit} s.comfyProxy.ServeHTTP(cw, r) var accepted struct { PromptID string `json:"prompt_id"` } if cw.status == http.StatusOK { _ = json.Unmarshal(cw.buf.Bytes(), &accepted) } if accepted.PromptID == "" { log.Info("prompt not accepted; releasing image lock", "status", cw.status) s.cfg.Lock.ReleaseImage() return } s.cfg.Metrics.IncImageJobs() go s.finishImageJob(accepted.PromptID) } // finishImageJob runs after the prompt has been accepted by ComfyUI: wait // for the job to finish, free ComfyUI's models, release the lock, and // optionally warm the chat model. func (s *Server) finishImageJob(promptID string) { log := s.log.With("prompt_id", promptID) log.Info("image job running") ctx, cancel := context.WithTimeout(context.Background(), s.cfg.JobTimeout) err := s.cfg.Comfy.WaitJob(ctx, promptID) cancel() if err != nil { log.Warn("image job did not complete cleanly; releasing lock anyway", "err", err) } else { log.Info("image job completed") } 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) } freeCancel() s.cfg.Lock.ReleaseImage() log.Info("image lock released") if s.cfg.WarmModel != "" { if state, _, _ := s.cfg.Lock.Snapshot(); state == lock.StateIdle { 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 { log.Info("warm model reloaded", "model", s.cfg.WarmModel) } wcancel() } } }