// Package proxy contains the HTTP handlers for both gpulock listeners: // reverse proxies to Ollama and ComfyUI with GPU lock arbitration in front // of the endpoints that load models. package proxy import ( "bytes" "context" "encoding/json" "errors" "fmt" "log/slog" "net/http" "net/http/httputil" "net/url" "time" "gpu-turnstile/internal/comfy" "gpu-turnstile/internal/lock" "gpu-turnstile/internal/metrics" "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 // 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 LLMWaitTimeout time.Duration UnloadTimeout time.Duration JobTimeout time.Duration WarmModel string } // Server serves both gpulock listeners. type Server struct { cfg Config log *slog.Logger 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() } return &Server{ cfg: cfg, log: log, ollamaProxy: newReverseProxy(ollamaURL, log.With("upstream", "ollama")), comfyProxy: newReverseProxy(comfyURL, log.With("upstream", "comfy")), }, nil } func newReverseProxy(target *url.URL, log *slog.Logger) *httputil.ReverseProxy { return &httputil.ReverseProxy{ 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) } // 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 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 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 captureLimit bytes of the body. type captureWriter struct { http.ResponseWriter status int buf bytes.Buffer } func (w *captureWriter) WriteHeader(code int) { w.status = code w.ResponseWriter.WriteHeader(code) } func (w *captureWriter) Write(p []byte) (int, error) { if w.buf.Len() < captureLimit { 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} 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(), 30*time.Second) 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(), 2*time.Minute) 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() } } }