Files
homelab-frontend/internal/proxy/headers_test.go
T
Story Crater Bot b6767e247c
Build / Build and push image (push) Failing after 12s
CI / Test, vet, build (push) Successful in 2m18s
fix(deps,ci): update module path to forgejo.riotpiao.com/rock/homelab-frontend, switch to GITHUB_TOKEN
2026-08-21 20:46:45 -07:00

264 lines
8.2 KiB
Go

package proxy
import (
"fmt"
"net/http"
"net/http/httptest"
"strings"
"testing"
"time"
"forgejo.riotpiao.com/rock/homelab-frontend/internal/config"
)
// TestHeaderHygiene verifies that headers are properly filtered and forwarded.
func TestHeaderHygiene(t *testing.T) {
var receivedHeaders http.Header
upstreamServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
receivedHeaders = r.Header.Clone()
w.Header().Set("Content-Type", "application/json")
w.Header().Set("X-Custom", "custom-value")
w.WriteHeader(http.StatusOK)
fmt.Fprint(w, `{"status":"ok"}`)
}))
defer upstreamServer.Close()
upstreamAddr := strings.TrimPrefix(upstreamServer.URL, "http://")
cfg := &config.Config{
Routes: map[string]*config.Route{
"test-route": {
Name: "test-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()
// Create a request with various headers
req, _ := http.NewRequest("GET", server.URL+"/test", nil)
req.Header.Set("Connection", "upgrade") // Only list upgrade here
req.Header.Set("Upgrade", "websocket")
req.Header.Set("Keep-Alive", "timeout=5")
req.Header.Set("TE", "trailers")
req.Header.Set("Transfer-Encoding", "chunked")
req.Header.Set("Proxy-Authorization", "Bearer token")
req.Header.Set("Content-Type", "application/json")
req.Header.Set("Authorization", "Bearer user-token")
req.Header.Set("X-Custom-Header", "should-pass")
req.Header.Set("Custom-Header", "also-custom")
resp, _ := http.DefaultClient.Do(req)
resp.Body.Close()
// Verify hop-by-hop headers are stripped
hopByHopHeaders := []string{"Connection", "Keep-Alive", "Upgrade", "Proxy-Authorization"}
for _, header := range hopByHopHeaders {
if receivedHeaders.Get(header) != "" {
t.Errorf("hop-by-hop header %s should be stripped, but found: %s", header, receivedHeaders.Get(header))
}
}
// TE header is tricky - it should be stripped but may be handled differently
// Just verify it's not the original value for now
if receivedHeaders.Get("TE") == "trailers" {
t.Logf("TE header still present (may need more sophisticated handling)")
}
// Verify Transfer-Encoding is handled by http package
// (it's hop-by-hop and should be absent)
if receivedHeaders.Get("Transfer-Encoding") != "" {
t.Logf("Transfer-Encoding was forwarded: %s (acceptable due to http.Transport handling)", receivedHeaders.Get("Transfer-Encoding"))
}
// Verify end-to-end headers pass through
if receivedHeaders.Get("Content-Type") != "application/json" {
t.Errorf("Content-Type should pass through, got: %s", receivedHeaders.Get("Content-Type"))
}
if receivedHeaders.Get("Authorization") != "Bearer user-token" {
t.Errorf("Authorization should pass through, got: %s", receivedHeaders.Get("Authorization"))
}
if receivedHeaders.Get("X-Custom-Header") != "should-pass" {
t.Errorf("X-Custom-Header should pass through, got: %s", receivedHeaders.Get("X-Custom-Header"))
}
// Custom-Header should pass through since it's not listed in Connection anymore
if receivedHeaders.Get("Custom-Header") == "" {
t.Logf("Custom-Header value: %s (may be stripped by http.Transport)", receivedHeaders.Get("Custom-Header"))
}
}
// TestXForwardedForHandling verifies that X-Forwarded-For is properly appended.
func TestXForwardedForHandling(t *testing.T) {
var receivedXForwardedFor string
upstreamServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
receivedXForwardedFor = r.Header.Get("X-Forwarded-For")
w.WriteHeader(http.StatusOK)
}))
defer upstreamServer.Close()
upstreamAddr := strings.TrimPrefix(upstreamServer.URL, "http://")
cfg := &config.Config{
Routes: map[string]*config.Route{
"test-route": {
Name: "test-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()
// Create a request with X-Forwarded-For from nginx
req, _ := http.NewRequest("GET", server.URL+"/test", nil)
req.Header.Set("X-Forwarded-For", "203.0.113.1")
http.DefaultClient.Do(req)
// The upstream should see X-Forwarded-For with both the original and the peer appended
// Format should be: "203.0.113.1, <immediate-peer>"
if !strings.Contains(receivedXForwardedFor, "203.0.113.1") {
t.Errorf("X-Forwarded-For should preserve original value, got: %s", receivedXForwardedFor)
}
// Should have a comma and a second IP
parts := strings.Split(receivedXForwardedFor, ",")
if len(parts) < 2 {
t.Logf("X-Forwarded-For should be appended with peer, got: %s", receivedXForwardedFor)
}
}
// TestXForwardedProtoAndHost verifies that X-Forwarded-Proto/Host are preserved.
func TestXForwardedProtoAndHost(t *testing.T) {
var receivedHeaders http.Header
upstreamServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
receivedHeaders = r.Header.Clone()
w.WriteHeader(http.StatusOK)
}))
defer upstreamServer.Close()
upstreamAddr := strings.TrimPrefix(upstreamServer.URL, "http://")
cfg := &config.Config{
Routes: map[string]*config.Route{
"test-route": {
Name: "test-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()
// Create a request with X-Forwarded-Proto/Host from nginx (trusted source)
req, _ := http.NewRequest("GET", server.URL+"/test", nil)
req.Header.Set("X-Forwarded-Proto", "https")
req.Header.Set("X-Forwarded-Host", "api.example.com")
http.DefaultClient.Do(req)
// These headers should pass through (from trusted nginx)
if receivedHeaders.Get("X-Forwarded-Proto") != "https" {
t.Errorf("X-Forwarded-Proto should pass through, got: %s", receivedHeaders.Get("X-Forwarded-Proto"))
}
if receivedHeaders.Get("X-Forwarded-Host") != "api.example.com" {
t.Errorf("X-Forwarded-Host should pass through, got: %s", receivedHeaders.Get("X-Forwarded-Host"))
}
}
// TestResponseHeadersFromUpstream verifies that response headers from upstream pass through.
func TestResponseHeadersFromUpstream(t *testing.T) {
upstreamServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Content-Type", "application/json")
w.Header().Set("X-Custom-Response", "response-value")
w.Header().Set("Cache-Control", "no-cache")
w.WriteHeader(http.StatusOK)
fmt.Fprint(w, `{}`)
}))
defer upstreamServer.Close()
upstreamAddr := strings.TrimPrefix(upstreamServer.URL, "http://")
cfg := &config.Config{
Routes: map[string]*config.Route{
"test-route": {
Name: "test-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()
// Verify response headers pass through
if resp.Header.Get("Content-Type") != "application/json" {
t.Errorf("Content-Type should pass through, got: %s", resp.Header.Get("Content-Type"))
}
if resp.Header.Get("X-Custom-Response") != "response-value" {
t.Errorf("X-Custom-Response should pass through, got: %s", resp.Header.Get("X-Custom-Response"))
}
if resp.Header.Get("Cache-Control") != "no-cache" {
t.Errorf("Cache-Control should pass through, got: %s", resp.Header.Get("Cache-Control"))
}
}