Files
smallwebwaf/internal/proxy/passthrough_test.go
T
clawbot b8edfcec37
check / check (push) Failing after 2s
Spend less of make test writing the image and waiting (closes #56)
The test phase spends most of its time compiling with the race
detector from an empty build cache; then come writing the test image
and the internal/proxy tests.

- Go's build cache is on a tmpfs in the test phase, so its 137 MB are
  no longer written into the test image.
- TestUpgradedConnectionOutlastsTheTimeouts waits until just past
  shortTimeout after the request was sent, rather than 7.5 s after the
  upgrade, so it ends with the other timing tests.
- shortTimeout is written as waitLimit / 2, as its comment says it is.

Model: opus-5-5
2026-10-04 08:19:19 +00:00

422 lines
11 KiB
Go

package proxy_test
import (
"bufio"
"bytes"
"errors"
"io"
"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")
sent := time.Now()
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)
}
// Every timeout started by the time smallwebwaf read the request:
// wait until just past shortTimeout after it was sent, then use the
// connection.
time.Sleep(time.Until(sent.Add(shortTimeout + 100*time.Millisecond)))
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 != 28<<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)
// size counts every byte of the request: the request line, the
// headers and the blank line that ends them.
const (
start = "GET / HTTP/1.1\r\nHost: app\r\nX-Large: "
end = "\r\n\r\n"
)
for _, tc := range []struct {
size int
want int
}{
{size: 32 << 10, want: http.StatusOK},
{size: 32<<10 + 1, want: http.StatusRequestHeaderFieldsTooLarge},
} {
conn := dial(t, addr)
send(t, conn, start+strings.Repeat("a", tc.size-len(start)-len(end))+end)
wantStatus(t, readResponse(t, conn), 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()
// No test can listen on port 1: listening on port 0 gets one from 32768 up.
addr, out := startProxy(t, "http://"+localhost+":1", 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)
}
}