Files
smallwebwaf/internal/proxy/passthrough_test.go
T
sneak 545ce67f44
check / check (push) Successful in 2m9s
Pass-through proxy with timeouts, size limits and a request log (closes #13)
The repo's first code, with the layout the prompts policies ask for:
Makefile, script/ entrypoints, a Dockerfile whose lint and test phases
gate the build, the Gitea workflow, the canonical dotfiles and
REPO_POLICIES.md. smallwebwaf passes each request to the app through
httputil.ReverseProxy within the four timeouts and two size limits,
works out the client's address behind trusted proxies, and writes one
JSON line per request. The tests run against real local servers.
SPEC.md now says what Go's HTTP server does before smallwebwaf sees a
request; make fmt only rewraps EVALUATION.md.

Model: opus-5-5
2026-10-03 14:19:01 +00:00

420 lines
11 KiB
Go

package proxy_test
import (
"bufio"
"bytes"
"errors"
"io"
"net"
"net/http"
"slices"
"strings"
"sync/atomic"
"testing"
"time"
"sneak.berlin/go/smallwebwaf/internal/config"
"sneak.berlin/go/smallwebwaf/internal/proxy"
"sneak.berlin/go/smallwebwaf/internal/requestlog"
)
// A request target with an escaped slash and space in its path, and a
// query with a parameter ReverseProxy cannot parse.
const (
rawPath = "/some%2Fpath/with%20space"
rawQuery = "b=2&a=1&bad=%zz;x"
)
// chunkSize is the size of each part of a body a test sends in parts.
const chunkSize = 1 << 10
var errNotStreamed = errors.New("the first part never reached the app")
// appSaw is what the app received.
type appSaw struct {
method string
target string
header http.Header
body []byte
}
func TestPassesRequestAndAnswerUnchanged(t *testing.T) {
t.Parallel()
requestBody := bytes.Repeat([]byte("request body "), 8000)
answerBody := bytes.Repeat([]byte("answer body "), 8000)
saw := make(chan appSaw, 1)
app := startApp(t, func(w http.ResponseWriter, r *http.Request) {
body, _ := io.ReadAll(r.Body)
saw <- appSaw{r.Method, r.RequestURI, r.Header.Clone(), body}
w.Header().Set("X-App", "yes")
w.Header().Add("Set-Cookie", "a=1")
w.Header().Add("Set-Cookie", "b=2")
w.WriteHeader(http.StatusTeapot)
_, _ = w.Write(answerBody)
})
addr, out := startProxy(t, app.URL, nil)
req := newRequest(t, http.MethodPatch, addr, rawPath+"?"+rawQuery,
bytes.NewReader(requestBody))
req.Header.Add("X-Test", "one")
req.Header.Add("X-Test", "two")
req.Header.Set("User-Agent", "test-agent")
got := do(t, req)
wantAppSaw(t, <-saw, requestBody)
wantAnswer(t, got, answerBody)
line := out.requestLine(t)
wantLine(t, line, http.StatusTeapot, requestlog.ActionForward)
wantRequestFields(t, line, addr, len(requestBody), len(answerBody))
}
// wantAppSaw checks that the app received the test's request unchanged.
func wantAppSaw(t *testing.T, saw appSaw, body []byte) {
t.Helper()
if saw.method != http.MethodPatch || saw.target != rawPath+"?"+rawQuery {
t.Errorf("app saw %s %s, want %s %s", saw.method, saw.target,
http.MethodPatch, rawPath+"?"+rawQuery)
}
if !slices.Equal(saw.header.Values("X-Test"), []string{"one", "two"}) {
t.Errorf("app saw X-Test %q", saw.header.Values("X-Test"))
}
if saw.header.Get("User-Agent") != "test-agent" {
t.Errorf("app saw User-Agent %q", saw.header.Get("User-Agent"))
}
if !bytes.Equal(saw.body, body) {
t.Errorf("app saw a body of %d bytes, want the %d sent",
len(saw.body), len(body))
}
}
// wantAnswer checks that the client received the app's answer unchanged.
func wantAnswer(t *testing.T, got answer, body []byte) {
t.Helper()
wantStatus(t, got, http.StatusTeapot)
if got.header.Get("X-App") != "yes" {
t.Errorf("client got X-App %q", got.header.Get("X-App"))
}
if !slices.Equal(got.header.Values("Set-Cookie"), []string{"a=1", "b=2"}) {
t.Errorf("client got Set-Cookie %q", got.header.Values("Set-Cookie"))
}
if got.err != nil || !bytes.Equal(got.body, body) {
t.Errorf("client got %d bytes (%v), want the %d the app sent",
len(got.body), got.err, len(body))
}
}
// wantRequestFields checks the log line's fields about the request.
func wantRequestFields(t *testing.T, line logLine, host string, sent, received int) {
t.Helper()
want := requestlog.Line{
Type: "request", Time: line.Time, ClientIP: localhost, PeerIP: localhost,
Method: http.MethodPatch, Host: host, Path: rawPath, Query: rawQuery,
Protocol: "HTTP/1.1", Status: http.StatusTeapot,
UpstreamStatus: http.StatusTeapot, RequestBytes: int64(sent),
ResponseBytes: int64(received), UserAgent: "test-agent",
Action: requestlog.ActionForward, DurationTotal: line.DurationTotal,
DurationUpstreamTotal: line.DurationUpstreamTotal,
}
if line.Line != want {
t.Errorf("log line\n%+v\nwant\n%+v", line.Line, want)
}
_, err := time.Parse(time.RFC3339, line.Time)
if err != nil || line.DurationTotal <= 0 || line.DurationUpstreamTotal <= 0 {
t.Errorf("log line has time %q and durations %v and %v",
line.Time, line.DurationTotal, line.DurationUpstreamTotal)
}
}
func TestStreamsTheRequestBody(t *testing.T) {
t.Parallel()
chunk := bytes.Repeat([]byte("x"), chunkSize)
firstArrived := make(chan struct{})
app := startApp(t, func(w http.ResponseWriter, r *http.Request) {
first := make([]byte, len(chunk))
_, err := io.ReadFull(r.Body, first)
if err != nil {
return
}
close(firstArrived)
rest, _ := io.ReadAll(r.Body)
_, _ = w.Write(rest)
})
addr, _ := startProxy(t, app.URL, nil)
body, writer := io.Pipe()
go func() {
_, _ = writer.Write(chunk)
select {
case <-firstArrived:
_, _ = writer.Write(chunk)
_ = writer.Close()
case <-time.After(waitLimit):
_ = writer.CloseWithError(errNotStreamed)
}
}()
got := do(t, newRequest(t, http.MethodPost, addr, "/upload", body))
if got.err != nil || !bytes.Equal(got.body, chunk) {
t.Errorf("app read %d bytes after the first part (%v), want %d",
len(got.body), got.err, len(chunk))
}
}
func TestStreamsTheAnswerBody(t *testing.T) {
t.Parallel()
chunk := bytes.Repeat([]byte("y"), chunkSize)
firstArrived := make(chan struct{})
app := startApp(t, func(w http.ResponseWriter, _ *http.Request) {
_, _ = w.Write(chunk)
_ = http.NewResponseController(w).Flush()
select {
case <-firstArrived:
_, _ = w.Write(chunk)
case <-time.After(waitLimit):
}
})
addr, _ := startProxy(t, app.URL, nil)
req := newRequest(t, http.MethodGet, addr, "/download", http.NoBody)
res, err := newClient(t).Do(req)
if err != nil {
t.Fatalf("request: %v", err)
}
first := make([]byte, len(chunk))
_, err = io.ReadFull(res.Body, first)
close(firstArrived)
got := readAnswer(res)
if err != nil || got.err != nil || !bytes.Equal(got.body, chunk) {
t.Errorf("client read %d bytes after the first part (%v, %v), want %d",
len(got.body), err, got.err, len(chunk))
}
}
func TestUpgradedConnectionOutlastsTheTimeouts(t *testing.T) {
t.Parallel()
app := startApp(t, echoAfterUpgrade)
addr, out := startProxy(t, app.URL, map[string]string{
clientRequestTimeout: shortTimeoutSetting,
clientResponseTimeout: shortTimeoutSetting,
upstreamRequestTimeout: shortTimeoutSetting,
upstreamResponseTimeout: shortTimeoutSetting,
})
conn := dial(t, addr)
send(t, conn, "GET /socket HTTP/1.1\r\nHost: app\r\n"+
"Connection: Upgrade\r\nUpgrade: websocket\r\n\r\n")
reader := bufio.NewReader(conn)
res, err := http.ReadResponse(reader, nil)
if err != nil {
t.Fatalf("read the answer to the upgrade: %v", err)
}
_ = res.Body.Close()
if res.StatusCode != http.StatusSwitchingProtocols {
t.Fatalf("status %d, want %d", res.StatusCode, http.StatusSwitchingProtocols)
}
// Wait past every timeout, then use the connection.
time.Sleep(3 * shortTimeout)
send(t, conn, "still here\n")
echoed, err := reader.ReadString('\n')
if err != nil || echoed != "still here\n" {
t.Errorf("echo %q (%v), want %q", echoed, err, "still here\n")
}
_ = conn.Close()
wantLine(t, out.requestLine(t), http.StatusSwitchingProtocols,
requestlog.ActionForward)
}
// echoAfterUpgrade is an app that switches protocols on request, and then
// sends back each line it receives.
func echoAfterUpgrade(w http.ResponseWriter, r *http.Request) {
if r.Header.Get("Upgrade") != "websocket" {
http.Error(w, "not an upgrade", http.StatusBadRequest)
return
}
conn, buffered, err := http.NewResponseController(w).Hijack()
if err != nil {
return
}
defer func() {
_ = conn.Close()
}()
_, _ = buffered.WriteString("HTTP/1.1 101 Switching Protocols\r\n" +
"Connection: Upgrade\r\nUpgrade: websocket\r\n\r\n")
_ = buffered.Flush()
for {
line, err := buffered.ReadString('\n')
if err != nil {
return
}
_, _ = buffered.WriteString(line)
_ = buffered.Flush()
}
}
func TestServerHasTheFixedLimits(t *testing.T) {
t.Parallel()
cfg, err := config.FromEnvironment(func(string) (string, bool) { return "", false })
if err != nil {
t.Fatalf("default settings: %v", err)
}
server := proxy.New(proxy.Params{
Config: cfg,
RequestLog: io.Discard,
ProcessLog: requestlog.NewProcessLogger(io.Discard),
})
if server.Addr != ":8080" || server.MaxHeaderBytes != 32<<10 ||
server.IdleTimeout != 2*time.Minute || server.ReadHeaderTimeout != time.Minute {
t.Errorf("server listens on %q with header limit %d, idle time %s and "+
"header timeout %s", server.Addr, server.MaxHeaderBytes,
server.IdleTimeout, server.ReadHeaderTimeout)
}
}
func TestRefusesHeadersOver32KiB(t *testing.T) {
t.Parallel()
var calls atomic.Int32
app := startApp(t, func(http.ResponseWriter, *http.Request) {
calls.Add(1)
})
addr, _ := startProxy(t, app.URL, nil)
for _, tc := range []struct {
headerSize int
want int
}{
{headerSize: 30 << 10, want: http.StatusOK},
{headerSize: 40 << 10, want: http.StatusRequestHeaderFieldsTooLarge},
} {
req := newRequest(t, http.MethodGet, addr, "/", http.NoBody)
req.Header.Set("X-Large", strings.Repeat("a", tc.headerSize))
wantStatus(t, do(t, req), tc.want)
}
if calls.Load() != 1 {
t.Errorf("the app was called %d times, want once", calls.Load())
}
}
func TestAnswers502WhenTheAppCannotBeReached(t *testing.T) {
t.Parallel()
listener, err := (&net.ListenConfig{}).Listen(t.Context(), "tcp", localhost+":0")
if err != nil {
t.Fatalf("listen: %v", err)
}
closedAddr := listener.Addr().String()
_ = listener.Close()
addr, out := startProxy(t, "http://"+closedAddr, nil)
wantStatus(t, get(t, addr, "/"), http.StatusBadGateway)
wantLine(t, out.requestLine(t), http.StatusBadGateway,
requestlog.ActionUpstreamError)
logged := slices.ContainsFunc(out.lines(t), func(line map[string]any) bool {
return line["type"] == "process" && line["msg"] == "request to the app failed"
})
if !logged {
t.Errorf("no process line says the request to the app failed")
}
}
func TestLogsAnAnswerThatBrokeOff(t *testing.T) {
t.Parallel()
app := startApp(t, func(w http.ResponseWriter, _ *http.Request) {
_, _ = io.WriteString(w, "the first part")
_ = http.NewResponseController(w).Flush()
panic(http.ErrAbortHandler) // drops the connection mid-answer
})
addr, out := startProxy(t, app.URL, nil)
got := get(t, addr, "/")
if string(got.body) != "the first part" || !errors.Is(got.err, io.ErrUnexpectedEOF) {
t.Errorf("client read %q (%v), want the first part cut off", got.body, got.err)
}
wantLine(t, out.requestLine(t), http.StatusOK, requestlog.ActionUpstreamError)
}
func TestLogsAClientThatWentAway(t *testing.T) {
t.Parallel()
arrived := make(chan struct{})
app := startApp(t, func(_ http.ResponseWriter, r *http.Request) {
close(arrived)
<-r.Context().Done()
})
addr, out := startProxy(t, app.URL, nil)
conn := dial(t, addr)
send(t, conn, "GET /slow HTTP/1.1\r\nHost: app\r\n\r\n")
select {
case <-arrived:
case <-time.After(waitLimit):
t.Fatal("the request never reached the app")
}
_ = conn.Close()
line := out.requestLine(t)
if !line.Aborted || line.Status != 0 || line.Action != requestlog.ActionForward {
t.Errorf("log line has aborted %v, status %d and action %q, "+
"want true, 0 and %q",
line.Aborted, line.Status, line.Action, requestlog.ActionForward)
}
}