package proxy import ( "bufio" "fmt" "io" "net/http" "net/http/httptest" "runtime" "strings" "sync" "testing" "time" "forgejo.riotpiao.com/rock/homelab-frontend/internal/config" ) // TestSSEUnbuffered verifies that SSE events stream to the client without buffering. func TestSSEUnbuffered(t *testing.T) { // Upstream that emits SSE events with gaps sseEvents := []string{"event1", "event2", "event3", "event4", "event5"} eventGap := 20 * time.Millisecond upstreamServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { w.Header().Set("Content-Type", "text/event-stream") w.Header().Set("Cache-Control", "no-cache") w.Header().Set("Connection", "keep-alive") w.WriteHeader(http.StatusOK) rc := http.NewResponseController(w) for i, event := range sseEvents { if i > 0 { time.Sleep(eventGap) } fmt.Fprintf(w, "data: %s\n\n", event) if err := rc.Flush(); err != nil { return } } fmt.Fprintf(w, "data: [DONE]\n\n") _ = rc.Flush() })) defer upstreamServer.Close() upstreamAddr := strings.TrimPrefix(upstreamServer.URL, "http://") cfg := &config.Config{ Routes: map[string]*config.Route{ "sse-route": { Name: "sse-route", Upstream: config.Upstream{ Address: upstreamAddr, ConnectTimeout: 5 * time.Second, ReadTimeout: 10 * time.Second, WriteTimeout: 5 * time.Second, MaxBodySize: 1024 * 1024, AuthRequired: false, }, }, }, } handler := New(cfg) defer handler.Close() server := httptest.NewServer(handler) defer server.Close() // Connect to the proxy resp, err := http.Get(server.URL + "/sse") if err != nil { t.Fatalf("request failed: %v", err) } defer resp.Body.Close() // Verify headers if resp.Header.Get("Content-Type") != "text/event-stream" { t.Errorf("expected Content-Type: text/event-stream, got %s", resp.Header.Get("Content-Type")) } if resp.Header.Get("Cache-Control") != "no-cache" { t.Errorf("expected Cache-Control: no-cache, got %s", resp.Header.Get("Cache-Control")) } // Read events and measure timing reader := bufio.NewReader(resp.Body) eventTimes := make([]time.Time, 0, len(sseEvents)) observedEvents := make([]string, 0, len(sseEvents)) for { line, err := reader.ReadString('\n') if err != nil { if err == io.EOF { break } t.Fatalf("read failed: %v", err) } line = strings.TrimSpace(line) if strings.HasPrefix(line, "data: ") { event := strings.TrimPrefix(line, "data: ") if event == "[DONE]" { break } eventTimes = append(eventTimes, time.Now()) observedEvents = append(observedEvents, event) } } // Verify we got all events if len(observedEvents) != len(sseEvents) { t.Errorf("expected %d events, got %d", len(sseEvents), len(observedEvents)) } // Verify events match for i, expected := range sseEvents { if i < len(observedEvents) && observedEvents[i] != expected { t.Errorf("event %d: expected %s, got %s", i, expected, observedEvents[i]) } } // Verify timing gaps between events are reasonable // The gaps should be approximately eventGap (allowing for some overhead) for i := 1; i < len(eventTimes); i++ { gap := eventTimes[i].Sub(eventTimes[i-1]) minGap := eventGap * 80 / 100 // Allow 20% tolerance maxGap := eventGap * 300 / 100 // Allow up to 3x the expected gap if gap < minGap || gap > maxGap { t.Logf("event gap %d: %.1fms (expected ~%.1fms)", i, gap.Seconds()*1000, eventGap.Seconds()*1000) } } } // TestChunkedUnbuffered verifies that chunked responses stream without buffering. func TestChunkedUnbuffered(t *testing.T) { chunks := []string{"chunk1\n", "chunk2\n", "chunk3\n"} chunkGap := 20 * time.Millisecond upstreamServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { w.Header().Set("Content-Type", "text/plain") w.WriteHeader(http.StatusOK) rc := http.NewResponseController(w) for i, chunk := range chunks { if i > 0 { time.Sleep(chunkGap) } fmt.Fprint(w, chunk) if err := rc.Flush(); err != nil { return } } })) defer upstreamServer.Close() upstreamAddr := strings.TrimPrefix(upstreamServer.URL, "http://") cfg := &config.Config{ Routes: map[string]*config.Route{ "chunked-route": { Name: "chunked-route", Upstream: config.Upstream{ Address: upstreamAddr, ConnectTimeout: 5 * time.Second, ReadTimeout: 10 * time.Second, WriteTimeout: 5 * time.Second, MaxBodySize: 1024 * 1024, AuthRequired: false, }, }, }, } handler := New(cfg) defer handler.Close() server := httptest.NewServer(handler) defer server.Close() resp, err := http.Get(server.URL + "/chunked") if err != nil { t.Fatalf("request failed: %v", err) } defer resp.Body.Close() // Read chunks and verify they arrive before all are sent reader := bufio.NewReader(resp.Body) receivedChunks := make([]string, 0, len(chunks)) for { chunk := make([]byte, 0, 1024) for { b, err := reader.ReadByte() if err != nil { if err == io.EOF { break } t.Fatalf("read failed: %v", err) } chunk = append(chunk, b) if b == '\n' { break } } if len(chunk) > 0 { receivedChunks = append(receivedChunks, string(chunk)) } if len(receivedChunks) >= len(chunks) { break } } // Verify chunks match if len(receivedChunks) != len(chunks) { t.Errorf("expected %d chunks, got %d", len(chunks), len(receivedChunks)) } for i, expected := range chunks { if i < len(receivedChunks) && strings.TrimSpace(receivedChunks[i]) != strings.TrimSpace(expected) { t.Errorf("chunk %d: expected %s, got %s", i, strings.TrimSpace(expected), strings.TrimSpace(receivedChunks[i])) } } } // TestHeadersBeforeBody verifies that response headers reach the client before the body. func TestHeadersBeforeBody(t *testing.T) { headersSent := make(chan bool, 1) upstreamServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { w.Header().Set("X-Custom-Header", "test-value") w.Header().Set("Content-Type", "text/plain") w.WriteHeader(http.StatusOK) headersSent <- true // Simulate slow body send time.Sleep(100 * time.Millisecond) fmt.Fprint(w, "body content") })) defer upstreamServer.Close() upstreamAddr := strings.TrimPrefix(upstreamServer.URL, "http://") cfg := &config.Config{ Routes: map[string]*config.Route{ "headers-route": { Name: "headers-route", Upstream: config.Upstream{ Address: upstreamAddr, ConnectTimeout: 5 * time.Second, ReadTimeout: 5 * time.Second, WriteTimeout: 5 * time.Second, MaxBodySize: 1024 * 1024, AuthRequired: false, }, }, }, } handler := New(cfg) defer handler.Close() server := httptest.NewServer(handler) defer server.Close() resp, err := http.Get(server.URL + "/headers") if err != nil { t.Fatalf("request failed: %v", err) } // Headers should be immediately available if resp.Header.Get("X-Custom-Header") != "test-value" { t.Errorf("custom header not received before body") } resp.Body.Close() } // TestResponseHeadersPassThrough verifies that various response headers survive the proxy. func TestResponseHeadersPassThrough(t *testing.T) { testHeaders := map[string]string{ "Content-Type": "application/json", "Cache-Control": "no-cache, no-store", "X-Custom": "custom-value", } upstreamServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { for k, v := range testHeaders { w.Header().Set(k, v) } w.WriteHeader(http.StatusOK) fmt.Fprint(w, "test") })) defer upstreamServer.Close() upstreamAddr := strings.TrimPrefix(upstreamServer.URL, "http://") cfg := &config.Config{ Routes: map[string]*config.Route{ "headers-route": { Name: "headers-route", Upstream: config.Upstream{ Address: upstreamAddr, ConnectTimeout: 5 * time.Second, ReadTimeout: 5 * time.Second, WriteTimeout: 5 * time.Second, MaxBodySize: 1024 * 1024, AuthRequired: false, }, }, }, } handler := New(cfg) defer handler.Close() server := httptest.NewServer(handler) defer server.Close() resp, err := http.Get(server.URL + "/test") if err != nil { t.Fatalf("request failed: %v", err) } defer resp.Body.Close() for k, v := range testHeaders { if resp.Header.Get(k) != v { t.Errorf("header %s: expected %s, got %s", k, v, resp.Header.Get(k)) } } } // TestDONESentinel verifies that the [DONE] sentinel reaches the client. func TestDONESentinel(t *testing.T) { upstreamServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { w.Header().Set("Content-Type", "text/event-stream") w.WriteHeader(http.StatusOK) rc := http.NewResponseController(w) fmt.Fprint(w, "data: token1\n\n") _ = rc.Flush() fmt.Fprint(w, "data: [DONE]\n\n") _ = rc.Flush() })) defer upstreamServer.Close() upstreamAddr := strings.TrimPrefix(upstreamServer.URL, "http://") cfg := &config.Config{ Routes: map[string]*config.Route{ "sse-route": { Name: "sse-route", Upstream: config.Upstream{ Address: upstreamAddr, ConnectTimeout: 5 * time.Second, ReadTimeout: 5 * time.Second, WriteTimeout: 5 * time.Second, MaxBodySize: 1024 * 1024, AuthRequired: false, }, }, }, } handler := New(cfg) defer handler.Close() server := httptest.NewServer(handler) defer server.Close() resp, err := http.Get(server.URL + "/sse") if err != nil { t.Fatalf("request failed: %v", err) } defer resp.Body.Close() reader := bufio.NewReader(resp.Body) foundDONE := false for { line, err := reader.ReadString('\n') if err != nil { if err == io.EOF { break } t.Fatalf("read failed: %v", err) } line = strings.TrimSpace(line) if strings.Contains(line, "[DONE]") { foundDONE = true break } } if !foundDONE { t.Errorf("expected [DONE] sentinel, not found") } } // TestNoFullBuffering verifies that the response is not fully buffered in memory. func TestNoFullBuffering(t *testing.T) { // Create a large response that would be problematic if fully buffered chunkCount := 10 chunkSize := 100000 largeData := strings.Repeat("x", chunkCount*chunkSize) upstreamServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { w.Header().Set("Content-Type", "text/plain") w.WriteHeader(http.StatusOK) rc := http.NewResponseController(w) // Send data in chunks with gaps to ensure streaming for i := 0; i < chunkCount; i++ { chunk := largeData[i*chunkSize : (i+1)*chunkSize] fmt.Fprint(w, chunk) if err := rc.Flush(); err != nil { return } if i < chunkCount-1 { time.Sleep(10 * time.Millisecond) } } })) defer upstreamServer.Close() upstreamAddr := strings.TrimPrefix(upstreamServer.URL, "http://") cfg := &config.Config{ Routes: map[string]*config.Route{ "large-route": { Name: "large-route", Upstream: config.Upstream{ Address: upstreamAddr, ConnectTimeout: 5 * time.Second, ReadTimeout: 10 * time.Second, WriteTimeout: 5 * time.Second, MaxBodySize: 100 * 1024 * 1024, AuthRequired: false, }, }, }, } handler := New(cfg) defer handler.Close() server := httptest.NewServer(handler) defer server.Close() resp, err := http.Get(server.URL + "/large") if err != nil { t.Fatalf("request failed: %v", err) } defer resp.Body.Close() // Read the response in chunks to verify streaming totalRead := 0 readChunkSize := 8192 for { buf := make([]byte, readChunkSize) n, err := resp.Body.Read(buf) if n > 0 { totalRead += n } if err != nil { if err == io.EOF { break } t.Fatalf("read failed: %v", err) } } if totalRead != len(largeData) { t.Errorf("expected to read %d bytes, got %d", len(largeData), totalRead) } } // TestTCPBackpressure verifies that TCP backpressure is respected during streaming. // When a client reads slowly, the upstream should experience backpressure on writes. func TestTCPBackpressure(t *testing.T) { // Track when upstream started writing and when each write completed var writeTimes []time.Time writesMu := sync.Mutex{} upstreamServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { w.Header().Set("Content-Type", "text/event-stream") w.Header().Set("X-Accel-Buffering", "no") // Issue #33: disable buffering w.WriteHeader(http.StatusOK) rc := http.NewResponseController(w) // Send many events to trigger backpressure for i := 0; i < 20; i++ { writesMu.Lock() writeTimes = append(writeTimes, time.Now()) writesMu.Unlock() fmt.Fprintf(w, "data: event%d\n\n", i) if err := rc.Flush(); err != nil { return } } })) defer upstreamServer.Close() upstreamAddr := strings.TrimPrefix(upstreamServer.URL, "http://") cfg := &config.Config{ Routes: map[string]*config.Route{ "backpressure-route": { Name: "backpressure-route", Upstream: config.Upstream{ Address: upstreamAddr, ConnectTimeout: 5 * time.Second, ReadTimeout: 10 * time.Second, WriteTimeout: 5 * time.Second, MaxBodySize: 1024 * 1024, AuthRequired: false, }, }, }, } handler := New(cfg) defer handler.Close() server := httptest.NewServer(handler) defer server.Close() resp, err := http.Get(server.URL + "/backpressure") if err != nil { t.Fatalf("request failed: %v", err) } defer resp.Body.Close() // Verify X-Accel-Buffering header is passed through if resp.Header.Get("X-Accel-Buffering") != "no" { t.Errorf("X-Accel-Buffering header not propagated, got: %s", resp.Header.Get("X-Accel-Buffering")) } // Read events with simulated slow client (small buffer) reader := bufio.NewReader(resp.Body) readStart := time.Now() eventCount := 0 for { line, err := reader.ReadString('\n') if err != nil { if err == io.EOF { break } t.Fatalf("read failed: %v", err) } if strings.HasPrefix(strings.TrimSpace(line), "data:") { eventCount++ // Simulate slow client by adding delay time.Sleep(5 * time.Millisecond) } } // Verify we got all events if eventCount != 20 { t.Errorf("expected 20 events, got %d", eventCount) } // Total read time should be roughly eventCount * readDelay // indicating backpressure was applied (upstream couldn't send all at once) elapsed := time.Since(readStart) expectedMin := time.Duration(20*5) * time.Millisecond if elapsed < expectedMin { t.Logf("backpressure test: elapsed=%.0fms (expected ~%.0fms)", elapsed.Seconds()*1000, expectedMin.Seconds()*1000) } } // TestConcurrentSSEStreams verifies that HTTP/2 multiplexing handles multiple concurrent streams. // Issue #32: Multiple LLM requests should not block each other. func TestConcurrentSSEStreams(t *testing.T) { upstreamServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { w.Header().Set("Content-Type", "text/event-stream") w.Header().Set("X-Accel-Buffering", "no") w.WriteHeader(http.StatusOK) rc := http.NewResponseController(w) // Each request sends unique identifier reqID := r.URL.Query().Get("id") for i := 0; i < 5; i++ { fmt.Fprintf(w, "data: [%s] event %d\n\n", reqID, i) if err := rc.Flush(); err != nil { return } time.Sleep(10 * time.Millisecond) } })) defer upstreamServer.Close() upstreamAddr := strings.TrimPrefix(upstreamServer.URL, "http://") cfg := &config.Config{ Routes: map[string]*config.Route{ "concurrent-route": { Name: "concurrent-route", Upstream: config.Upstream{ Address: upstreamAddr, ConnectTimeout: 5 * time.Second, ReadTimeout: 10 * time.Second, WriteTimeout: 5 * time.Second, MaxBodySize: 1024 * 1024, AuthRequired: false, }, }, }, } handler := New(cfg) defer handler.Close() server := httptest.NewServer(handler) defer server.Close() // Launch multiple concurrent requests var wg sync.WaitGroup results := make(map[string][]string) resultsMu := sync.Mutex{} for id := 0; id < 3; id++ { wg.Add(1) go func(streamID int) { defer wg.Done() url := fmt.Sprintf("%s/concurrent?id=stream%d", server.URL, streamID) resp, err := http.Get(url) if err != nil { t.Errorf("request failed: %v", err) return } defer resp.Body.Close() reader := bufio.NewReader(resp.Body) var events []string for { line, err := reader.ReadString('\n') if err != nil { if err == io.EOF { break } t.Errorf("read failed: %v", err) return } line = strings.TrimSpace(line) if strings.HasPrefix(line, "data:") { events = append(events, line) } } resultsMu.Lock() results[fmt.Sprintf("stream%d", streamID)] = events resultsMu.Unlock() }(id) } wg.Wait() // Verify all streams got their events for i := 0; i < 3; i++ { key := fmt.Sprintf("stream%d", i) events, ok := results[key] if !ok { t.Errorf("stream%d: no results", i) continue } if len(events) != 5 { t.Errorf("stream%d: expected 5 events, got %d", i, len(events)) } // Verify all events belong to this stream for _, event := range events { if !strings.Contains(event, key) { t.Errorf("stream%d: event from wrong stream: %s", i, event) } } } } // TestClientDisconnectCancelsUpstream verifies that when a client closes mid-stream, // the upstream request context is cancelled immediately and no goroutines are leaked. func TestClientDisconnectCancelsUpstream(t *testing.T) { contextCancelledAt := time.Time{} contextCancelledMu := sync.Mutex{} upstreamRequestedAt := time.Time{} upstreamServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { upstreamRequestedAt = time.Now() w.Header().Set("Content-Type", "text/event-stream") w.WriteHeader(http.StatusOK) rc := http.NewResponseController(w) // Send events until context is cancelled for i := 0; i < 100; i++ { select { case <-r.Context().Done(): contextCancelledMu.Lock() contextCancelledAt = time.Now() contextCancelledMu.Unlock() return default: } fmt.Fprintf(w, "data: event%d\n\n", i) if err := rc.Flush(); err != nil { contextCancelledMu.Lock() contextCancelledAt = time.Now() contextCancelledMu.Unlock() return } time.Sleep(50 * time.Millisecond) } })) defer upstreamServer.Close() upstreamAddr := strings.TrimPrefix(upstreamServer.URL, "http://") cfg := &config.Config{ Routes: map[string]*config.Route{ "disconnect-route": { Name: "disconnect-route", Upstream: config.Upstream{ Address: upstreamAddr, ConnectTimeout: 5 * time.Second, ReadTimeout: 30 * time.Second, WriteTimeout: 30 * time.Second, MaxBodySize: 1024 * 1024, AuthRequired: false, }, }, }, } handler := New(cfg) defer handler.Close() server := httptest.NewServer(handler) defer server.Close() // Baseline goroutine count baselineGoroutines := runtime.NumGoroutine() // Make a request with a custom HTTP client that allows us to close the connection client := &http.Client{ Timeout: 30 * time.Second, } req, err := http.NewRequest("GET", server.URL+"/disconnect", nil) if err != nil { t.Fatalf("request creation failed: %v", err) } resp, err := client.Do(req) if err != nil { t.Fatalf("request failed: %v", err) } defer resp.Body.Close() // Read a few events reader := bufio.NewReader(resp.Body) for i := 0; i < 2; i++ { line, err := reader.ReadString('\n') if err != nil { t.Fatalf("read failed: %v", err) } if !strings.Contains(line, "data:") { i-- // skip non-data lines } } // Close the response body (simulating client disconnect) resp.Body.Close() // Wait a bit for cancellation to propagate time.Sleep(200 * time.Millisecond) // Verify context was cancelled contextCancelledMu.Lock() cancelled := !contextCancelledAt.IsZero() cancelDelay := time.Duration(0) if cancelled { cancelDelay = contextCancelledAt.Sub(upstreamRequestedAt) } contextCancelledMu.Unlock() if !cancelled { t.Errorf("expected upstream context to be cancelled, but it was not") } // Verify cancellation happened quickly (within 1s) if cancelDelay > 1*time.Second { t.Errorf("context cancellation took %.2fs (expected < 1s)", cancelDelay.Seconds()) } // Wait a bit for goroutines to clean up time.Sleep(100 * time.Millisecond) // Check for goroutine leaks finalGoroutines := runtime.NumGoroutine() if finalGoroutines > baselineGoroutines+5 { t.Errorf("possible goroutine leak: baseline=%d, final=%d", baselineGoroutines, finalGoroutines) } }