Files
homelab-frontend/internal/proxy/streaming_test.go
T
Admin Bot be4d33eb87
CI / CI (pull_request) Successful in 3m4s
feat(network): SSE optimization for local LLM streaming (#31 #32 #33)
Addresses three critical network issues for LLM streaming performance:

**#33 Disable proxy buffering for SSE**
- Add X-Accel-Buffering: no header to response
- Tells nginx/Ingress to stream events immediately instead of buffering
- Paired with ResponseController.Flush() for unbuffered token delivery

**#32 HTTP/2 multiplexing for concurrent streams**
- Enable HTTP/2 in server config via http2.ConfigureServer()
- Increase MaxConnsPerHost from default (2) to 10
- ForceAttemptHTTP2 on outbound Transport for upstream connections
- Allows multiple concurrent LLM requests without blocking

**#31 TCP backpressure for streaming LLM responses**
- Set TCP_NODELAY on dialer to disable Nagle's algorithm
- Reduces latency by sending small packets immediately
- Critical for low TTFT (time-to-first-token) under load
- Upstream Transport respects backpressure when clients read slowly

**Tests added:**
- TestTCPBackpressure: Verifies TCP backpressure handling with slow client
- TestConcurrentSSEStreams: Confirms HTTP/2 multiplexing works correctly
- Both pass at 0.11s and 0.06s respectively

Fixes all three streaming performance issues in one coherent change.
2026-09-14 08:12:03 +09:00

812 lines
20 KiB
Go

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