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) } }