package proxy import ( "fmt" "io" "net" "net/http" "net/http/httptest" "strings" "testing" "time" "github.com/Riotpiaole/homelab-frontend/internal/config" ) // TestConnectTimeout verifies that connections fail at the configured timeout. func TestConnectTimeout(t *testing.T) { // Use a port that's unlikely to have anything listening on it // This will cause the connection to hang/timeout cfg := &config.Config{ Routes: map[string]*config.Route{ "timeout-route": { Name: "timeout-route", Upstream: config.Upstream{ Address: "127.0.0.1:1", ConnectTimeout: 100 * time.Millisecond, 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() start := time.Now() resp, err := http.Get(server.URL + "/test") elapsed := time.Since(start) // Should fail quickly (within a reasonable tolerance of the connect timeout) if elapsed > 500*time.Millisecond { t.Errorf("connect timeout took too long: %.1fs (expected ~0.1s)", elapsed.Seconds()) } if err == nil && resp.StatusCode != http.StatusBadGateway { resp.Body.Close() t.Errorf("expected error or 502, got status %d", resp.StatusCode) } } // TestReadTimeout verifies that a stalled upstream times out with a 5xx response. func TestReadTimeout(t *testing.T) { // Create an upstream that accepts but never sends data listener, err := net.Listen("tcp", "127.0.0.1:0") if err != nil { t.Fatalf("failed to create listener: %v", err) } defer listener.Close() go func() { for { conn, err := listener.Accept() if err != nil { return } // Accept but never respond - this will trigger the read timeout go func() { time.Sleep(10 * time.Second) conn.Close() }() } }() upstreamAddr := listener.Addr().String() cfg := &config.Config{ Routes: map[string]*config.Route{ "timeout-route": { Name: "timeout-route", Upstream: config.Upstream{ Address: upstreamAddr, ConnectTimeout: 5 * time.Second, ReadTimeout: 100 * time.Millisecond, WriteTimeout: 5 * time.Second, MaxBodySize: 1024 * 1024, AuthRequired: false, }, }, }, } handler := New(cfg) defer handler.Close() server := httptest.NewServer(handler) defer server.Close() start := time.Now() resp, err := http.Get(server.URL + "/test") if err != nil { t.Fatalf("request failed: %v", err) } elapsed := time.Since(start) // Should timeout around the read timeout (with some tolerance) if elapsed < 50*time.Millisecond || elapsed > 500*time.Millisecond { t.Logf("read timeout took %.1fs (expected ~0.1s)", elapsed.Seconds()) } if resp.StatusCode != http.StatusBadGateway { t.Errorf("expected 502 on timeout, got %d", resp.StatusCode) } resp.Body.Close() } // TestLongStreamNotTruncated verifies that a stream with activity within the window // is not cut off by the read timeout. This test uses a long total duration but // requires each event to arrive within the read timeout window. func TestLongStreamNotTruncated(t *testing.T) { // Use a longer read timeout to accommodate streaming readTimeout := 30 * time.Second eventGap := 100 * time.Millisecond eventCount := 10 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) for i := 0; i < eventCount; i++ { if i > 0 { time.Sleep(eventGap) } fmt.Fprintf(w, "data: token%d\n\n", i) if err := rc.Flush(); err != nil { return } } 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{ "stream-route": { Name: "stream-route", Upstream: config.Upstream{ Address: upstreamAddr, ConnectTimeout: 5 * time.Second, ReadTimeout: readTimeout, 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 + "/stream") if err != nil { t.Fatalf("request failed: %v", err) } defer resp.Body.Close() // The total time should be: eventGap * (eventCount - 1) = 100ms * 9 = 900ms // This is longer than readTimeout (500ms), but should NOT be cut because // the timeout is for inactivity (idle time between reads), not total duration. // The test verifies we get all events despite the long total duration. totalTime := time.Duration(eventCount-1) * eventGap start := time.Now() body, err := io.ReadAll(resp.Body) elapsed := time.Since(start) if err != nil { t.Fatalf("failed to read response: %v", err) } responseStr := string(body) // Verify we got all events for i := 0; i < eventCount; i++ { expectedToken := fmt.Sprintf("token%d", i) if !strings.Contains(responseStr, expectedToken) { t.Errorf("expected token %s in response, but not found", expectedToken) } } // Verify we got the DONE sentinel if !strings.Contains(responseStr, "[DONE]") { t.Errorf("expected [DONE] sentinel in response") } _ = totalTime // indicate we're aware it's used conceptually _ = elapsed } // TestConfiguredTimeoutValues verifies that timeout configuration is used. func TestConfiguredTimeoutValues(t *testing.T) { // Chat route configuration chatCfg := &config.Config{ Routes: map[string]*config.Route{ "chat": { Name: "chat", Upstream: config.Upstream{ Address: "127.0.0.1:8000", ConnectTimeout: 10 * time.Second, ReadTimeout: 1 * time.Hour, WriteTimeout: 1 * time.Hour, MaxBodySize: 10 * 1024 * 1024, AuthRequired: false, }, }, }, } // Embeddings route configuration embeddingsCfg := &config.Config{ Routes: map[string]*config.Route{ "embeddings": { Name: "embeddings", Upstream: config.Upstream{ Address: "127.0.0.1:8001", ConnectTimeout: 10 * time.Second, ReadTimeout: 10 * time.Minute, WriteTimeout: 10 * time.Minute, MaxBodySize: 50 * 1024 * 1024, AuthRequired: false, }, }, }, } chatHandler := New(chatCfg) defer chatHandler.Close() embHandler := New(embeddingsCfg) defer embHandler.Close() // Verify chat timeouts chatRoute := chatHandler.routes["chat"] if chatRoute.Upstream.ConnectTimeout != 10*time.Second { t.Errorf("chat connect timeout: expected 10s, got %v", chatRoute.Upstream.ConnectTimeout) } if chatRoute.Upstream.ReadTimeout != 1*time.Hour { t.Errorf("chat read timeout: expected 1h, got %v", chatRoute.Upstream.ReadTimeout) } // Verify embeddings timeouts embRoute := embHandler.routes["embeddings"] if embRoute.Upstream.ConnectTimeout != 10*time.Second { t.Errorf("embeddings connect timeout: expected 10s, got %v", embRoute.Upstream.ConnectTimeout) } if embRoute.Upstream.ReadTimeout != 10*time.Minute { t.Errorf("embeddings read timeout: expected 10m, got %v", embRoute.Upstream.ReadTimeout) } }