gpulock: GPU arbitration proxy for Ollama + ComfyUI

Implements SPEC.md: two listeners, one writer-preferring two-mode lock,
Ollama unload before image jobs, ComfyUI history polling + VRAM free,
optional model warm-up, healthz/metrics endpoints, streaming-safe
reverse proxies, Dockerfile and Gitea Actions CI.
This commit is contained in:
mram
2026-09-20 18:05:20 +02:00
commit 065c294a96
14 changed files with 2027 additions and 0 deletions
+138
View File
@@ -0,0 +1,138 @@
// Package comfy is a minimal client for the ComfyUI endpoints gpulock needs
// after a prompt has been accepted: polling /history and freeing VRAM.
package comfy
import (
"bytes"
"context"
"encoding/json"
"fmt"
"io"
"log/slog"
"net/http"
"net/url"
"time"
)
// Client talks to a ComfyUI upstream.
type Client struct {
base string
hc *http.Client
log *slog.Logger
// PollInterval is how often /history/<id> is polled. Defaults to 1s;
// tests can shorten it.
PollInterval time.Duration
}
// New validates baseURL and returns a Client.
func New(baseURL string, log *slog.Logger) (*Client, error) {
u, err := url.Parse(baseURL)
if err != nil || u.Scheme == "" || u.Host == "" {
return nil, fmt.Errorf("invalid ComfyUI URL %q", baseURL)
}
return &Client{
base: u.String(),
hc: &http.Client{Timeout: 30 * time.Second},
log: log,
PollInterval: time.Second,
}, nil
}
// Probe checks that ComfyUI answers on /system_stats.
func (c *Client) Probe(ctx context.Context) error {
req, err := http.NewRequestWithContext(ctx, http.MethodGet, c.base+"/system_stats", nil)
if err != nil {
return err
}
resp, err := c.hc.Do(req)
if err != nil {
return err
}
defer resp.Body.Close()
io.Copy(io.Discard, resp.Body)
if resp.StatusCode/100 != 2 {
return fmt.Errorf("/system_stats returned %s", resp.Status)
}
return nil
}
type historyEntry struct {
Status struct {
Completed bool `json:"completed"`
StatusStr string `json:"status_str"`
} `json:"status"`
}
// jobDone reports whether the history entry for promptID shows a finished
// job (completed or errored).
func (c *Client) jobDone(ctx context.Context, promptID string) (bool, error) {
req, err := http.NewRequestWithContext(ctx, http.MethodGet, c.base+"/history/"+promptID, nil)
if err != nil {
return false, err
}
resp, err := c.hc.Do(req)
if err != nil {
return false, err
}
defer resp.Body.Close()
if resp.StatusCode/100 != 2 {
io.Copy(io.Discard, resp.Body)
return false, fmt.Errorf("/history/%s returned %s", promptID, resp.Status)
}
var history map[string]historyEntry
if err := json.NewDecoder(resp.Body).Decode(&history); err != nil {
return false, err
}
entry, ok := history[promptID]
if !ok {
return false, nil
}
return entry.Status.Completed || entry.Status.StatusStr == "error", nil
}
// WaitJob polls until the job for promptID is finished. The context should
// carry the job deadline; transient polling errors are logged and retried,
// so a returned error is always from the context.
func (c *Client) WaitJob(ctx context.Context, promptID string) error {
for {
done, err := c.jobDone(ctx, promptID)
if err != nil {
if ctx.Err() != nil {
return ctx.Err()
}
c.log.Warn("failed to poll job history", "prompt_id", promptID, "err", err)
}
if done {
return nil
}
select {
case <-ctx.Done():
return ctx.Err()
case <-time.After(c.PollInterval):
}
}
}
// Free asks ComfyUI to unload its models and free VRAM.
func (c *Client) Free(ctx context.Context) error {
buf, err := json.Marshal(map[string]bool{"unload_models": true, "free_memory": true})
if err != nil {
return err
}
req, err := http.NewRequestWithContext(ctx, http.MethodPost, c.base+"/free", bytes.NewReader(buf))
if err != nil {
return err
}
req.Header.Set("Content-Type", "application/json")
resp, err := c.hc.Do(req)
if err != nil {
return err
}
defer resp.Body.Close()
io.Copy(io.Discard, resp.Body)
if resp.StatusCode/100 != 2 {
return fmt.Errorf("/free returned %s", resp.Status)
}
return nil
}
+157
View File
@@ -0,0 +1,157 @@
// Package lock implements the two-mode GPU arbitration lock: any number of
// concurrent LLM requests ("readers") or exactly one image job ("writer"),
// with image jobs taking priority over newly arriving LLM requests.
package lock
import (
"context"
"log/slog"
"sync"
)
// State is the current GPU occupancy state.
type State string
const (
StateIdle State = "idle"
StateLLM State = "llm"
StateImage State = "image"
)
type imageWaiter struct{ id uint64 }
// Lock is a writer-preferring two-mode lock. The zero value is not usable;
// construct with New.
type Lock struct {
mu sync.Mutex
change chan struct{} // closed and replaced on every state change
n int // LLM requests in flight
imageActive bool // an image job holds the GPU
imageQ []imageWaiter
nextID uint64
log *slog.Logger
}
// New returns a ready-to-use Lock. log may be nil; if set, every state
// transition is logged at debug level.
func New(log *slog.Logger) *Lock {
return &Lock{change: make(chan struct{}), log: log}
}
// broadcast wakes all waiters. Call with mu held.
func (l *Lock) broadcast() {
close(l.change)
l.change = make(chan struct{})
}
func (l *Lock) logTransition(msg string, args ...any) {
if l.log != nil {
l.log.Debug(msg, args...)
}
}
// AcquireLLM blocks until no image job is active or pending, then registers
// one in-flight LLM request. Returns ctx.Err() if the context is cancelled
// while waiting; no state is changed in that case.
func (l *Lock) AcquireLLM(ctx context.Context) error {
l.mu.Lock()
for l.imageActive || len(l.imageQ) > 0 {
ch := l.change
l.mu.Unlock()
select {
case <-ctx.Done():
return ctx.Err()
case <-ch:
}
l.mu.Lock()
}
l.n++
n := l.n
l.mu.Unlock()
l.logTransition("lock transition", "state", StateLLM, "llm_inflight", n)
return nil
}
// ReleaseLLM marks one LLM request as finished.
func (l *Lock) ReleaseLLM() {
l.mu.Lock()
l.n--
n := l.n
if l.n == 0 {
l.broadcast()
}
l.mu.Unlock()
if n == 0 {
l.logTransition("lock transition", "state", StateIdle, "llm_inflight", 0)
}
}
// AcquireImage queues the caller FIFO behind other image jobs, blocks new
// LLM requests immediately, and waits for in-flight LLM requests to drain
// before granting exclusive GPU access. Returns ctx.Err() if the context is
// cancelled while waiting; the caller is removed from the queue.
func (l *Lock) AcquireImage(ctx context.Context) error {
l.mu.Lock()
l.nextID++
w := imageWaiter{id: l.nextID}
l.imageQ = append(l.imageQ, w)
pending := len(l.imageQ)
l.broadcast() // new LLM requests must block from now on
l.mu.Unlock()
l.logTransition("image job pending", "image_queue", pending)
for {
l.mu.Lock()
if l.imageQ[0].id == w.id && l.n == 0 && !l.imageActive {
l.imageQ = l.imageQ[1:]
l.imageActive = true
l.mu.Unlock()
l.logTransition("lock transition", "state", StateImage)
return nil
}
ch := l.change
l.mu.Unlock()
select {
case <-ctx.Done():
l.mu.Lock()
for i, q := range l.imageQ {
if q.id == w.id {
l.imageQ = append(l.imageQ[:i], l.imageQ[i+1:]...)
break
}
}
l.broadcast()
l.mu.Unlock()
l.logTransition("image job wait cancelled", "image_queue", len(l.imageQ))
return ctx.Err()
case <-ch:
}
}
}
// ReleaseImage frees the GPU after an image job.
func (l *Lock) ReleaseImage() {
l.mu.Lock()
l.imageActive = false
l.broadcast()
l.mu.Unlock()
l.logTransition("lock transition", "state", StateIdle)
}
// Snapshot reports the current state, the number of in-flight LLM requests,
// and whether an image job is active or waiting.
func (l *Lock) Snapshot() (state State, llmInflight int, imagePending bool) {
l.mu.Lock()
defer l.mu.Unlock()
switch {
case l.imageActive:
state = StateImage
case l.n > 0:
state = StateLLM
default:
state = StateIdle
}
return state, l.n, l.imageActive || len(l.imageQ) > 0
}
+228
View File
@@ -0,0 +1,228 @@
package lock
import (
"context"
"sync"
"testing"
"time"
)
func waitFor(t *testing.T, ch <-chan struct{}, what string) {
t.Helper()
select {
case <-ch:
case <-time.After(2 * time.Second):
t.Fatalf("timed out waiting for %s", what)
}
}
func assertBlocked(t *testing.T, ch <-chan struct{}, what string) {
t.Helper()
select {
case <-ch:
t.Fatalf("%s should still be blocked", what)
case <-time.After(50 * time.Millisecond):
}
}
func TestLLMConcurrency(t *testing.T) {
lk := New(nil)
ctx := context.Background()
for i := 0; i < 3; i++ {
if err := lk.AcquireLLM(ctx); err != nil {
t.Fatal(err)
}
}
state, n, pending := lk.Snapshot()
if state != StateLLM || n != 3 || pending {
t.Fatalf("got state=%s n=%d pending=%v", state, n, pending)
}
lk.ReleaseLLM()
lk.ReleaseLLM()
if state, _, _ := lk.Snapshot(); state != StateLLM {
t.Fatalf("state=%s, want llm", state)
}
lk.ReleaseLLM()
if state, n, _ := lk.Snapshot(); state != StateIdle || n != 0 {
t.Fatalf("got state=%s n=%d, want idle", state, n)
}
}
func TestImageWaitsForLLMDrainAndBlocksNewLLM(t *testing.T) {
lk := New(nil)
ctx := context.Background()
if err := lk.AcquireLLM(ctx); err != nil {
t.Fatal(err)
}
if err := lk.AcquireLLM(ctx); err != nil {
t.Fatal(err)
}
imageAcquired := make(chan struct{})
go func() {
if err := lk.AcquireImage(ctx); err != nil {
t.Error(err)
}
close(imageAcquired)
}()
assertBlocked(t, imageAcquired, "image acquire while LLMs in flight")
// While an image job is pending, new LLM requests must block even
// though LLM concurrency is otherwise allowed.
llmAcquired := make(chan struct{})
go func() {
if err := lk.AcquireLLM(ctx); err != nil {
t.Error(err)
}
close(llmAcquired)
}()
assertBlocked(t, llmAcquired, "LLM acquire while image pending")
lk.ReleaseLLM()
assertBlocked(t, imageAcquired, "image acquire with one LLM still in flight")
lk.ReleaseLLM()
waitFor(t, imageAcquired, "image acquire after drain")
assertBlocked(t, llmAcquired, "LLM acquire while image active")
state, _, pending := lk.Snapshot()
if state != StateImage || !pending {
t.Fatalf("got state=%s pending=%v", state, pending)
}
lk.ReleaseImage()
waitFor(t, llmAcquired, "LLM acquire after image release")
lk.ReleaseLLM()
}
func TestImageFIFO(t *testing.T) {
lk := New(nil)
ctx := context.Background()
if err := lk.AcquireImage(ctx); err != nil {
t.Fatal(err)
}
order := make(chan int, 2)
for i, id := range []int{1, 2} {
go func() {
if err := lk.AcquireImage(ctx); err != nil {
t.Error(err)
return
}
order <- id
}()
// Ensure this waiter is queued before the next one starts.
for deadline := time.Now().Add(2 * time.Second); ; {
lk.mu.Lock()
qlen := len(lk.imageQ)
lk.mu.Unlock()
if qlen == i+1 {
break
}
if time.Now().After(deadline) {
t.Fatalf("image waiter %d never queued", id)
}
time.Sleep(time.Millisecond)
}
}
lk.ReleaseImage()
if got := <-order; got != 1 {
t.Fatalf("first image job = %d, want 1", got)
}
select {
case got := <-order:
t.Fatalf("second image job %d acquired while first active", got)
case <-time.After(50 * time.Millisecond):
}
lk.ReleaseImage()
if got := <-order; got != 2 {
t.Fatalf("second image job = %d, want 2", got)
}
lk.ReleaseImage()
}
func TestContextCancelRemovesLLMWaiter(t *testing.T) {
lk := New(nil)
if err := lk.AcquireImage(context.Background()); err != nil {
t.Fatal(err)
}
ctx, cancel := context.WithCancel(context.Background())
done := make(chan error, 1)
go func() { done <- lk.AcquireLLM(ctx) }()
time.Sleep(20 * time.Millisecond)
cancel()
if err := <-done; err == nil {
t.Fatal("expected context error")
}
lk.ReleaseImage()
if state, n, _ := lk.Snapshot(); state != StateIdle || n != 0 {
t.Fatalf("got state=%s n=%d, want idle", state, n)
}
}
func TestContextCancelRemovesImageWaiter(t *testing.T) {
lk := New(nil)
ctx := context.Background()
if err := lk.AcquireLLM(ctx); err != nil {
t.Fatal(err)
}
wctx, cancel := context.WithCancel(ctx)
done := make(chan error, 1)
go func() { done <- lk.AcquireImage(wctx) }()
time.Sleep(20 * time.Millisecond)
cancel()
if err := <-done; err == nil {
t.Fatal("expected context error")
}
lk.ReleaseLLM()
// The cancelled waiter must be gone: a new LLM request acquires
// immediately instead of blocking behind it.
if err := lk.AcquireLLM(ctx); err != nil {
t.Fatal(err)
}
if _, _, pending := lk.Snapshot(); pending {
t.Fatal("image_pending=true after waiter was cancelled")
}
lk.ReleaseLLM()
}
// TestRace hammers the lock from many goroutines; run with -race.
func TestRace(t *testing.T) {
lk := New(nil)
var wg sync.WaitGroup
for i := 0; i < 8; i++ {
wg.Add(1)
go func(i int) {
defer wg.Done()
for j := 0; j < 25; j++ {
if i%4 == 3 {
if err := lk.AcquireImage(context.Background()); err != nil {
t.Error(err)
return
}
lk.Snapshot()
lk.ReleaseImage()
} else {
if err := lk.AcquireLLM(context.Background()); err != nil {
t.Error(err)
return
}
lk.Snapshot()
lk.ReleaseLLM()
}
}
}(i)
}
wg.Wait()
state, n, pending := lk.Snapshot()
if state != StateIdle || n != 0 || pending {
t.Fatalf("leaked lock state: state=%s n=%d pending=%v", state, n, pending)
}
}
+121
View File
@@ -0,0 +1,121 @@
// Package metrics provides the Prometheus text exposition for gpulock
// without any external dependencies.
package metrics
import (
"fmt"
"io"
"math"
"sync/atomic"
)
var buckets = []float64{0.05, 0.1, 0.25, 0.5, 1, 2.5, 5, 10, 30, 60, 120, 300, 600}
// histogram is a fixed-bucket cumulative histogram safe for concurrent use.
type histogram struct {
counts []atomic.Uint64 // len(buckets)+1; last slot is +Inf
sum atomic.Uint64 // math.Float64bits
count atomic.Uint64
}
func newHistogram() *histogram {
return &histogram{counts: make([]atomic.Uint64, len(buckets)+1)}
}
func (h *histogram) observe(v float64) {
for i, b := range buckets {
if v <= b {
h.counts[i].Add(1)
}
}
h.counts[len(buckets)].Add(1)
h.sum.Add(math.Float64bits(v))
h.count.Add(1)
}
func (h *histogram) write(w io.Writer, name, label string) {
for i, b := range buckets {
fmt.Fprintf(w, "%s_bucket{%sle=\"%g\"} %d\n", name, label, b, h.counts[i].Load())
}
fmt.Fprintf(w, "%s_bucket{%sle=\"+Inf\"} %d\n", name, label, h.counts[len(buckets)].Load())
fmt.Fprintf(w, "%s_sum{%s} %g\n", name, trimLabel(label), math.Float64frombits(h.sum.Load()))
fmt.Fprintf(w, "%s_count{%s} %d\n", name, trimLabel(label), h.count.Load())
}
// trimLabel converts a `kind="llm",`-style prefix into the label set used for
// _sum/_count lines (same labels, no trailing comma).
func trimLabel(label string) string {
if len(label) > 0 && label[len(label)-1] == ',' {
return label[:len(label)-1]
}
return label
}
// Metrics holds all gpulock metric values.
type Metrics struct {
llmWait *histogram
imageWait *histogram
unload *histogram
imageJobs atomic.Uint64
}
// New returns a zeroed Metrics.
func New() *Metrics {
return &Metrics{
llmWait: newHistogram(),
imageWait: newHistogram(),
unload: newHistogram(),
}
}
// ObserveLockWait records how long a lock acquisition of the given kind
// ("llm" or "image") waited.
func (m *Metrics) ObserveLockWait(kind string, seconds float64) {
if kind == "image" {
m.imageWait.observe(seconds)
} else {
m.llmWait.observe(seconds)
}
}
// ObserveUnload records how long an Ollama unload took.
func (m *Metrics) ObserveUnload(seconds float64) { m.unload.observe(seconds) }
// IncImageJobs counts one accepted image job.
func (m *Metrics) IncImageJobs() { m.imageJobs.Add(1) }
// Render writes the Prometheus text exposition for the given lock snapshot.
func (m *Metrics) Render(w io.Writer, state string, llmInflight int, imagePending bool) {
fmt.Fprint(w, `# HELP gpulock_state Current GPU state (1 for the active state).
# TYPE gpulock_state gauge
`)
for _, s := range []string{"idle", "llm", "image"} {
v := 0
if s == state {
v = 1
}
fmt.Fprintf(w, "gpulock_state{state=%q} %d\n", s, v)
}
pending := 0
if imagePending {
pending = 1
}
fmt.Fprintf(w, `# HELP gpulock_llm_inflight LLM requests currently in flight.
# TYPE gpulock_llm_inflight gauge
gpulock_llm_inflight %d
# HELP gpulock_image_pending Whether an image job is active or waiting.
# TYPE gpulock_image_pending gauge
gpulock_image_pending %d
# HELP gpulock_image_jobs_total Image jobs accepted by ComfyUI.
# TYPE gpulock_image_jobs_total counter
gpulock_image_jobs_total %d
# HELP gpulock_lock_wait_seconds Time spent waiting to acquire the GPU lock.
# TYPE gpulock_lock_wait_seconds histogram
`, llmInflight, pending, m.imageJobs.Load())
m.llmWait.write(w, "gpulock_lock_wait_seconds", `kind="llm",`)
m.imageWait.write(w, "gpulock_lock_wait_seconds", `kind="image",`)
fmt.Fprint(w, `# HELP gpulock_unload_seconds Time spent unloading Ollama models before an image job.
# TYPE gpulock_unload_seconds histogram
`)
m.unload.write(w, "gpulock_unload_seconds", "")
}
+180
View File
@@ -0,0 +1,180 @@
// Package ollama is a minimal client for the Ollama management endpoints
// gpulock needs: listing loaded models, unloading them, and warming a model.
package ollama
import (
"bytes"
"context"
"encoding/json"
"fmt"
"io"
"log/slog"
"net/http"
"net/url"
"time"
)
// Client talks to an Ollama upstream.
type Client struct {
base string
hc *http.Client
log *slog.Logger
// PollInterval is how often /api/ps is re-checked while waiting for
// models to unload. Defaults to 500ms; tests can shorten it.
PollInterval time.Duration
}
// New validates baseURL and returns a Client.
func New(baseURL string, log *slog.Logger) (*Client, error) {
u, err := url.Parse(baseURL)
if err != nil || u.Scheme == "" || u.Host == "" {
return nil, fmt.Errorf("invalid Ollama URL %q", baseURL)
}
return &Client{
base: u.String(),
hc: &http.Client{Timeout: 30 * time.Second},
log: log,
PollInterval: 500 * time.Millisecond,
}, nil
}
// Probe checks that Ollama answers on /api/version.
func (c *Client) Probe(ctx context.Context) error {
req, err := http.NewRequestWithContext(ctx, http.MethodGet, c.base+"/api/version", nil)
if err != nil {
return err
}
resp, err := c.hc.Do(req)
if err != nil {
return err
}
defer resp.Body.Close()
io.Copy(io.Discard, resp.Body)
if resp.StatusCode/100 != 2 {
return fmt.Errorf("/api/version returned %s", resp.Status)
}
return nil
}
type psResponse struct {
Models []struct {
Name string `json:"name"`
Model string `json:"model"`
} `json:"models"`
}
// LoadedModels returns the names of models currently held in memory.
func (c *Client) LoadedModels(ctx context.Context) ([]string, error) {
req, err := http.NewRequestWithContext(ctx, http.MethodGet, c.base+"/api/ps", nil)
if err != nil {
return nil, err
}
resp, err := c.hc.Do(req)
if err != nil {
return nil, err
}
defer resp.Body.Close()
if resp.StatusCode/100 != 2 {
io.Copy(io.Discard, resp.Body)
return nil, fmt.Errorf("/api/ps returned %s", resp.Status)
}
var ps psResponse
if err := json.NewDecoder(resp.Body).Decode(&ps); err != nil {
return nil, err
}
models := make([]string, 0, len(ps.Models))
for _, m := range ps.Models {
if m.Name != "" {
models = append(models, m.Name)
} else {
models = append(models, m.Model)
}
}
return models, nil
}
func (c *Client) post(ctx context.Context, path string, body any) (int, error) {
buf, err := json.Marshal(body)
if err != nil {
return 0, err
}
req, err := http.NewRequestWithContext(ctx, http.MethodPost, c.base+path, bytes.NewReader(buf))
if err != nil {
return 0, err
}
req.Header.Set("Content-Type", "application/json")
resp, err := c.hc.Do(req)
if err != nil {
return 0, err
}
defer resp.Body.Close()
io.Copy(io.Discard, resp.Body)
return resp.StatusCode, nil
}
// unloadModel asks Ollama to evict one model. Embedding-only models reject
// /api/generate, so fall back to /api/embed on non-2xx.
func (c *Client) unloadModel(ctx context.Context, model string) error {
status, err := c.post(ctx, "/api/generate", map[string]any{"model": model, "keep_alive": 0})
if err == nil && status/100 == 2 {
return nil
}
status, err2 := c.post(ctx, "/api/embed", map[string]any{"model": model, "input": "x", "keep_alive": 0})
if err2 != nil {
return fmt.Errorf("generate: %v; embed: %v", err, err2)
}
if status/100 != 2 {
return fmt.Errorf("embed unload returned status %d", status)
}
return nil
}
// UnloadAll evicts every loaded model and waits until /api/ps is empty. The
// context should carry the unload deadline; when it expires UnloadAll returns
// ctx.Err() so the caller can degrade instead of failing. It returns the
// elapsed wall time.
func (c *Client) UnloadAll(ctx context.Context) (time.Duration, error) {
start := time.Now()
models, err := c.LoadedModels(ctx)
if err != nil {
return time.Since(start), fmt.Errorf("list models: %w", err)
}
for _, m := range models {
if err := c.unloadModel(ctx, m); err != nil {
c.log.Warn("failed to unload model", "model", m, "err", err)
}
}
for {
models, err := c.LoadedModels(ctx)
if err == nil && len(models) == 0 {
return time.Since(start), nil
}
if err != nil {
c.log.Warn("failed to poll /api/ps", "err", err)
}
select {
case <-ctx.Done():
return time.Since(start), ctx.Err()
case <-time.After(c.PollInterval):
}
}
}
// Warm reloads a model and pins it in memory (keep_alive=-1) so the next
// chat request does not pay the load time.
func (c *Client) Warm(ctx context.Context, model string) error {
status, err := c.post(ctx, "/api/generate", map[string]any{
"model": model,
"prompt": "",
"stream": false,
"keep_alive": -1,
})
if err != nil {
return err
}
if status/100 != 2 {
return fmt.Errorf("warm generate returned status %d", status)
}
return nil
}
+285
View File
@@ -0,0 +1,285 @@
// 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()
}
}
}
+347
View File
@@ -0,0 +1,347 @@
package proxy
import (
"bufio"
"fmt"
"io"
"net/http"
"net/http/httptest"
"strings"
"sync"
"testing"
"time"
"gpu-turnstile/internal/comfy"
"gpu-turnstile/internal/lock"
"gpu-turnstile/internal/metrics"
"gpu-turnstile/internal/ollama"
)
// recorder collects upstream call events in order.
type recorder struct {
mu sync.Mutex
events []string
}
func (r *recorder) add(e string) {
r.mu.Lock()
r.events = append(r.events, e)
r.mu.Unlock()
}
func (r *recorder) index(e string) int {
r.mu.Lock()
defer r.mu.Unlock()
for i, ev := range r.events {
if ev == e {
return i
}
}
return -1
}
// fakes wires up fake Ollama and ComfyUI upstreams plus the gpulock server.
type fakes struct {
rec *recorder
ollama *httptest.Server
comfy *httptest.Server
server *httptest.Server // Ollama-facing gpulock listener
comfySrv *httptest.Server // ComfyUI-facing gpulock listener
freeCh chan struct{}
chatCh chan struct{}
historyMu sync.Mutex
history string
}
func newFakes(t *testing.T) *fakes {
t.Helper()
f := &fakes{rec: &recorder{}, freeCh: make(chan struct{}), chatCh: make(chan struct{})}
var psCalls int
var psMu sync.Mutex
f.ollama = httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
switch r.URL.Path {
case "/api/version":
io.WriteString(w, `{"version":"0.0.0"}`)
case "/api/ps":
psMu.Lock()
psCalls++
c := psCalls
psMu.Unlock()
f.rec.add("ps")
if c == 1 {
io.WriteString(w, `{"models":[{"name":"chat-model"}]}`)
} else {
io.WriteString(w, `{"models":[]}`)
}
case "/api/generate":
f.rec.add("unload")
io.WriteString(w, `{}`)
case "/api/chat":
f.rec.add("chat")
close(f.chatCh)
w.Header().Set("Content-Type", "application/x-ndjson")
io.WriteString(w, `{"done":true}`+"\n")
default:
http.NotFound(w, r)
}
}))
t.Cleanup(f.ollama.Close)
f.comfy = httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
switch {
case r.URL.Path == "/system_stats":
io.WriteString(w, `{}`)
case r.URL.Path == "/prompt":
f.rec.add("prompt")
w.Header().Set("Content-Type", "application/json")
io.WriteString(w, `{"prompt_id":"p1"}`)
case r.URL.Path == "/history/p1":
f.rec.add("history")
f.historyMu.Lock()
h := f.history
f.historyMu.Unlock()
w.Header().Set("Content-Type", "application/json")
if h == "" {
io.WriteString(w, `{}`)
} else {
fmt.Fprintf(w, `{"p1":{"status":{"completed":%s,"status_str":"success"}}}`, h)
}
case r.URL.Path == "/free":
f.rec.add("free")
close(f.freeCh)
io.WriteString(w, `{}`)
default:
http.NotFound(w, r)
}
}))
t.Cleanup(f.comfy.Close)
ollamaClient, err := ollama.New(f.ollama.URL, nil)
if err != nil {
t.Fatal(err)
}
ollamaClient.PollInterval = 5 * time.Millisecond
comfyClient, err := comfy.New(f.comfy.URL, nil)
if err != nil {
t.Fatal(err)
}
comfyClient.PollInterval = 5 * time.Millisecond
srv, err := New(Config{
OllamaURL: f.ollama.URL,
ComfyURL: f.comfy.URL,
Lock: lock.New(nil),
Ollama: ollamaClient,
Comfy: comfyClient,
Metrics: metrics.New(),
LLMWaitTimeout: 2 * time.Second,
UnloadTimeout: 2 * time.Second,
JobTimeout: 2 * time.Second,
})
if err != nil {
t.Fatal(err)
}
f.server = httptest.NewServer(srv.OllamaHandler())
t.Cleanup(f.server.Close)
f.comfySrv = httptest.NewServer(srv.ComfyHandler())
t.Cleanup(f.comfySrv.Close)
return f
}
func (f *fakes) completeJob() {
f.historyMu.Lock()
f.history = "true"
f.historyMu.Unlock()
}
func TestImageJobSequenceAndLLMBlocked(t *testing.T) {
f := newFakes(t)
// Start the image job. The response returns as soon as ComfyUI has
// accepted the prompt; history polling and /free run in the background.
resp, err := http.Post(f.comfySrv.URL+"/prompt", "application/json", strings.NewReader(`{"workflow":{}}`))
if err != nil {
t.Fatal(err)
}
body, _ := io.ReadAll(resp.Body)
resp.Body.Close()
if resp.StatusCode != 200 || !strings.Contains(string(body), `"prompt_id":"p1"`) {
t.Fatalf("prompt response = %d %s", resp.StatusCode, body)
}
// While the image job is running (history not yet complete), an LLM
// request must be held.
chatDone := make(chan struct{})
go func() {
defer close(chatDone)
resp, err := http.Post(f.server.URL+"/api/chat", "application/json", strings.NewReader(`{}`))
if err == nil {
io.Copy(io.Discard, resp.Body)
resp.Body.Close()
}
}()
select {
case <-f.chatCh:
t.Fatal("chat reached Ollama while image job still running")
case <-time.After(100 * time.Millisecond):
}
// Finish the job; the background poller then calls /free, releases the
// lock, and the chat request goes through.
f.completeJob()
select {
case <-chatDone:
case <-time.After(3 * time.Second):
t.Fatal("chat request never completed")
}
// Assert the call sequence for one image job.
for _, pair := range [][2]string{
{"ps", "unload"},
{"unload", "prompt"},
{"prompt", "history"},
{"history", "free"},
{"free", "chat"},
} {
a, b := f.rec.index(pair[0]), f.rec.index(pair[1])
if a < 0 || b < 0 || a >= b {
f.rec.mu.Lock()
t.Fatalf("expected %s before %s; events: %v", pair[0], pair[1], f.rec.events)
}
}
}
func TestPromptRejectedReleasesLock(t *testing.T) {
f := newFakes(t)
// Make ComfyUI reject the prompt.
f.comfy.Config.Handler = http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if r.URL.Path == "/prompt" {
f.rec.add("prompt")
http.Error(w, "bad workflow", http.StatusBadRequest)
return
}
http.NotFound(w, r)
})
resp, err := http.Post(f.comfySrv.URL+"/prompt", "application/json", strings.NewReader(`{}`))
if err != nil {
t.Fatal(err)
}
body, _ := io.ReadAll(resp.Body)
resp.Body.Close()
if resp.StatusCode != 400 || !strings.Contains(string(body), "bad workflow") {
t.Fatalf("prompt response = %d %s", resp.StatusCode, body)
}
// The lock must already be free: a chat request goes straight through.
resp, err = http.Post(f.server.URL+"/api/chat", "application/json", strings.NewReader(`{}`))
if err != nil {
t.Fatal(err)
}
resp.Body.Close()
if resp.StatusCode != 200 {
t.Fatalf("chat status = %d", resp.StatusCode)
}
}
func TestStreamingNotBuffered(t *testing.T) {
gate := make(chan struct{})
upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Content-Type", "application/x-ndjson")
fl := w.(http.Flusher)
io.WriteString(w, `{"chunk":1}`+"\n")
fl.Flush()
<-gate // hold chunk 2 back until the client has seen chunk 1
io.WriteString(w, `{"chunk":2}`+"\n")
fl.Flush()
}))
defer upstream.Close()
srv, err := New(Config{
OllamaURL: upstream.URL,
ComfyURL: "http://127.0.0.1:1",
Lock: lock.New(nil),
Metrics: metrics.New(),
LLMWaitTimeout: time.Second,
})
if err != nil {
t.Fatal(err)
}
front := httptest.NewServer(srv.OllamaHandler())
defer front.Close()
resp, err := http.Post(front.URL+"/api/chat", "application/json", strings.NewReader(`{}`))
if err != nil {
t.Fatal(err)
}
defer resp.Body.Close()
// If gpulock buffered the stream, this read would never complete while
// the gate is closed.
line, err := bufio.NewReader(resp.Body).ReadString('\n')
if err != nil {
t.Fatal(err)
}
if !strings.Contains(line, `"chunk":1`) {
t.Fatalf("first line = %q", line)
}
close(gate)
rest, err := io.ReadAll(resp.Body)
if err != nil {
t.Fatal(err)
}
if !strings.Contains(string(rest), `"chunk":2`) {
t.Fatalf("rest = %q", rest)
}
}
func TestHealthzAndMetrics(t *testing.T) {
f := newFakes(t)
for _, base := range []string{f.server.URL, f.comfySrv.URL} {
resp, err := http.Get(base + "/healthz")
if err != nil {
t.Fatal(err)
}
body, _ := io.ReadAll(resp.Body)
resp.Body.Close()
if resp.StatusCode != 200 || !strings.Contains(string(body), `"state":"idle"`) {
t.Fatalf("healthz = %d %s", resp.StatusCode, body)
}
}
resp, err := http.Get(f.server.URL + "/metrics")
if err != nil {
t.Fatal(err)
}
body, _ := io.ReadAll(resp.Body)
resp.Body.Close()
text := string(body)
for _, want := range []string{
`gpulock_state{state="idle"} 1`,
"gpulock_llm_inflight 0",
"gpulock_image_jobs_total 0",
`gpulock_lock_wait_seconds_bucket{kind="llm",le="+Inf"} 0`,
"gpulock_unload_seconds_count{} 0",
} {
if !strings.Contains(text, want) {
t.Fatalf("metrics missing %q", want)
}
}
}
func TestPassThroughNoLock(t *testing.T) {
f := newFakes(t)
resp, err := http.Get(f.server.URL + "/api/version")
if err != nil {
t.Fatal(err)
}
body, _ := io.ReadAll(resp.Body)
resp.Body.Close()
if resp.StatusCode != 200 || !strings.Contains(string(body), "0.0.0") {
t.Fatalf("pass-through = %d %s", resp.StatusCode, body)
}
}