Files

683 lines
20 KiB
Go

// 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"
"gpu-turnstile/internal/supervise"
)
// 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
// ComfySup, when non-nil, is the managed ComfyUI process: any ComfyUI
// request starts it on demand, /prompt additionally waits for
// readiness before taking the GPU lock. nil = unmanaged upstream.
ComfySup *supervise.Process
// 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
// LLMBusyMode is "wait" (default) or "reject". In reject mode an LLM
// request that arrives while an image job is active or pending is
// answered immediately with LLMBusyStatus and a Retry-After header
// (BusyRetryAfter seconds) instead of waiting for the lock. In wait
// mode the Retry-After header is sent when LLMWaitTimeout expires.
LLMBusyMode string
LLMBusyStatus int
BusyRetryAfter int
// 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
busyMode string
busyStatus int
busyRetryAfter int
ollamaProxy *httputil.ReverseProxy
comfyProxy *httputil.ReverseProxy
}
// New builds a Server, validating the upstream URLs. At least one of
// OllamaURL / ComfyURL must be set; an empty URL disables that consumer —
// its handler is then never served, its client may be nil, and the image
// job flow skips the Ollama unload/warm steps.
func New(cfg Config) (*Server, error) {
if cfg.OllamaURL == "" && cfg.ComfyURL == "" {
return nil, fmt.Errorf("at least one of OllamaURL or ComfyURL is required")
}
var ollamaURL, comfyURL *url.URL
if cfg.OllamaURL != "" {
u, err := url.Parse(cfg.OllamaURL)
if err != nil || u.Scheme == "" || u.Host == "" {
return nil, fmt.Errorf("invalid OLLAMA_URL %q", cfg.OllamaURL)
}
ollamaURL = u
}
if cfg.ComfyURL != "" {
u, err := url.Parse(cfg.ComfyURL)
if err != nil || u.Scheme == "" || u.Host == "" {
return nil, fmt.Errorf("invalid COMFY_URL %q", cfg.ComfyURL)
}
comfyURL = u
}
log := cfg.Log
if log == nil {
log = slog.Default()
}
if cfg.UnloadPollInterval > 0 && cfg.Ollama != nil {
cfg.Ollama.PollInterval = cfg.UnloadPollInterval
}
if cfg.HistoryPollInterval > 0 && cfg.Comfy != nil {
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
}
busyMode := "wait"
if cfg.LLMBusyMode == "reject" {
busyMode = "reject"
}
busyStatus := cfg.LLMBusyStatus
if busyStatus == 0 {
busyStatus = http.StatusServiceUnavailable
}
busyRetryAfter := cfg.BusyRetryAfter
if busyRetryAfter <= 0 {
busyRetryAfter = 30
}
retry := &retryTransport{
base: http.DefaultTransport,
initial: backoffInitial,
max: backoffMax,
log: log,
}
s := &Server{
cfg: cfg,
log: log,
logWriter: cfg.LogWriter,
freeTimeout: freeTimeout,
warmTimeout: warmTimeout,
captureLimit: captureLimit,
backoffInitial: backoffInitial,
backoffMax: backoffMax,
busyMode: busyMode,
busyStatus: busyStatus,
busyRetryAfter: busyRetryAfter,
}
if ollamaURL != nil {
s.ollamaProxy = newReverseProxy(ollamaURL, retry, log.With("upstream", "ollama"))
}
if comfyURL != nil {
s.comfyProxy = newReverseProxy(comfyURL, retry, log.With("upstream", "comfy"))
}
return s, 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 listener for Ollama-compatible clients.
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()
if s.busyMode == "reject" {
if !s.cfg.Lock.TryAcquireLLM() {
s.cfg.Metrics.ObserveLockWait("llm", time.Since(start).Seconds())
msg := "GPU busy: image job active or queued"
if holder := s.cfg.Lock.External(); holder != "" {
msg = "GPU busy: " + holder
}
s.log.Info("llm request rejected; GPU busy",
"path", r.URL.Path, "status", s.busyStatus)
w.Header().Set("Retry-After", strconv.Itoa(s.busyRetryAfter))
http.Error(w, msg, s.busyStatus)
return
}
s.cfg.Metrics.ObserveLockWait("llm", time.Since(start).Seconds())
defer s.cfg.Lock.ReleaseLLM()
s.ollamaProxy.ServeHTTP(w, r)
return
}
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 {
w.Header().Set("Retry-After", strconv.Itoa(s.busyRetryAfter))
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 listener for ComfyUI clients.
func (s *Server) ComfyHandler() http.Handler {
return s.logRequests("comfy", 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 r.Method == http.MethodPost && r.URL.Path == "/prompt" {
s.handlePrompt(w, r)
return
}
if s.cfg.ComfySup != nil {
// Any other ComfyUI request also wakes the managed server; the
// retry backoff bridges the time it needs to come up. While a
// foreign process holds the GPU we refuse to spawn it — the
// request gets the busy answer instead of fighting for VRAM.
if holder := s.cfg.Lock.External(); holder != "" && !s.cfg.ComfySup.Running() {
w.Header().Set("Retry-After", strconv.Itoa(s.busyRetryAfter))
http.Error(w, "GPU busy: "+holder, http.StatusServiceUnavailable)
return
}
if err := s.cfg.ComfySup.EnsureRunning(); err != nil {
http.Error(w, fmt.Sprintf("cannot start ComfyUI: %v", err), http.StatusBadGateway)
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")
// Normally the managed server is brought up *before* taking the GPU
// lock: torch can take a minute to load, and LLM traffic should keep
// flowing in the meantime. While a foreign process holds the GPU, LLM
// traffic is blocked anyway and a fresh ComfyUI would fight it for
// VRAM — so the lock comes first in that case.
comfyFirst := s.cfg.ComfySup != nil && s.cfg.Lock.External() == ""
if comfyFirst {
if err := s.cfg.ComfySup.EnsureRunning(); err != nil {
http.Error(w, fmt.Sprintf("cannot start ComfyUI: %v", err), http.StatusBadGateway)
return
}
if err := s.cfg.ComfySup.WaitReady(r.Context()); err != nil {
http.Error(w, fmt.Sprintf("ComfyUI did not become ready: %v", err), http.StatusBadGateway)
return
}
}
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")
if s.cfg.ComfySup != nil && !comfyFirst {
if err := s.cfg.ComfySup.EnsureRunning(); err != nil {
s.cfg.Lock.ReleaseImage()
http.Error(w, fmt.Sprintf("cannot start ComfyUI: %v", err), http.StatusBadGateway)
return
}
if err := s.cfg.ComfySup.WaitReady(r.Context()); err != nil {
s.cfg.Lock.ReleaseImage()
http.Error(w, fmt.Sprintf("ComfyUI did not become ready: %v", err), http.StatusBadGateway)
return
}
}
if s.cfg.Ollama != nil {
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.ComfySup != nil {
// The idle clock starts when the job ends, not when it began.
s.cfg.ComfySup.NoteActivity()
}
if s.cfg.Ollama != nil && 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()
}
}
}