Backoff now covers all dial-phase errors (refused, timeout, DNS), TLS handshake failures, and 5xx responses with replayable bodies; streamed POSTs are never replayed to avoid duplicate work. First retry of an episode logs at WARN, subsequent attempts at INFO. LOGLEVEL (LOG_LEVEL kept as alias) now defaults to warn: startup logs version plus every setting; INFO adds one line per incoming request and per response with status/duration, ANSI-colored in text mode (bypasses slog's escaping so colors render in docker compose logs); NO_COLOR or LOG_FORMAT=json disables colors.
550 lines
16 KiB
Go
550 lines
16 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"
|
|
)
|
|
|
|
// 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
|
|
|
|
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
|
|
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,
|
|
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')
|
|
os.Stderr.WriteString(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()
|
|
}
|
|
}
|
|
}
|