Pass-through proxy with timeouts, size limits and a request log (closes #13)
check / check (push) Successful in 2m9s

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
This commit is contained in:
2026-10-03 14:19:01 +00:00
parent fd77e76177
commit 545ce67f44
45 changed files with 4869 additions and 74 deletions
+328
View File
@@ -0,0 +1,328 @@
package proxy_test
import (
"bufio"
"bytes"
"encoding/json"
"io"
"maps"
"net"
"net/http"
"net/http/httptest"
"strings"
"sync"
"testing"
"time"
"sneak.berlin/go/smallwebwaf/internal/config"
"sneak.berlin/go/smallwebwaf/internal/proxy"
"sneak.berlin/go/smallwebwaf/internal/requestlog"
)
const (
// shortTimeout is what a test sets a timeout to, to see it run out.
shortTimeout = 300 * time.Millisecond
// shortTimeoutSetting is shortTimeout as a setting's value.
shortTimeoutSetting = "300ms"
// longTimeoutSetting is a timeout that does not run out in a test.
longTimeoutSetting = "10s"
// waitLimit bounds how long a test waits for what should happen.
waitLimit = 10 * time.Second
// pollInterval is how often a test looks for a log line.
pollInterval = 10 * time.Millisecond
// localhost is where every test server listens, and so the address
// smallwebwaf sees each test's requests come from.
localhost = "127.0.0.1"
)
// The settings the tests set.
const (
clientRequestTimeout = "SWWAF_CLIENT_REQUEST_TIMEOUT"
clientResponseTimeout = "SWWAF_CLIENT_RESPONSE_TIMEOUT"
upstreamRequestTimeout = "SWWAF_UPSTREAM_REQUEST_TIMEOUT"
upstreamResponseTimeout = "SWWAF_UPSTREAM_RESPONSE_TIMEOUT"
requestMaxBytes = "SWWAF_REQUEST_MAX_BYTES"
responseMaxBytes = "SWWAF_RESPONSE_MAX_BYTES"
trustedProxies = "SWWAF_TRUSTED_PROXIES"
)
// output collects what smallwebwaf writes on stdout.
type output struct {
mu sync.Mutex
buf bytes.Buffer
}
// Write adds lines smallwebwaf writes.
func (o *output) Write(p []byte) (int, error) {
o.mu.Lock()
defer o.mu.Unlock()
return o.buf.Write(p)
}
// lines returns every line written so far, decoded.
func (o *output) lines(t *testing.T) []map[string]any {
t.Helper()
o.mu.Lock()
defer o.mu.Unlock()
var lines []map[string]any
for text := range strings.Lines(o.buf.String()) {
var line map[string]any
err := json.Unmarshal([]byte(text), &line)
if err != nil {
t.Fatalf("output line %q is not JSON: %v", text, err)
}
lines = append(lines, line)
}
return lines
}
// logLine is a request log line, as typed fields and as the JSON object
// it was written as.
type logLine struct {
requestlog.Line
fields map[string]any
}
// requestLines waits for count request log lines and returns them.
func (o *output) requestLines(t *testing.T, count int) []logLine {
t.Helper()
deadline := time.Now().Add(waitLimit)
for time.Now().Before(deadline) {
var found []logLine
for _, fields := range o.lines(t) {
if fields["type"] == "request" {
found = append(found, decodeLine(t, fields))
}
}
if len(found) >= count {
return found
}
time.Sleep(pollInterval)
}
t.Fatalf("fewer than %d request log lines after %s", count, waitLimit)
return nil
}
// requestLine waits for the request log line of a test's one request.
func (o *output) requestLine(t *testing.T) logLine {
t.Helper()
return o.requestLines(t, 1)[0]
}
// decodeLine reads a request log line's fields into a logLine.
func decodeLine(t *testing.T, fields map[string]any) logLine {
t.Helper()
encoded, err := json.Marshal(fields)
if err != nil {
t.Fatalf("encode %v: %v", fields, err)
}
line := logLine{fields: fields}
err = json.Unmarshal(encoded, &line.Line)
if err != nil {
t.Fatalf("decode %s: %v", encoded, err)
}
return line
}
// startApp starts app as the app smallwebwaf passes requests to.
func startApp(t *testing.T, app http.HandlerFunc) *httptest.Server {
t.Helper()
server := httptest.NewServer(app)
t.Cleanup(server.Close)
return server
}
// startProxy starts smallwebwaf in front of the app at appURL, with the
// settings in env on top of the defaults, and returns where it listens and
// what it writes.
func startProxy(t *testing.T, appURL string, env map[string]string) (string, *output) {
t.Helper()
settings := map[string]string{"SWWAF_UPSTREAM_URL": appURL}
maps.Copy(settings, env)
cfg, err := config.FromEnvironment(func(name string) (string, bool) {
value, ok := settings[name]
return value, ok
})
if err != nil {
t.Fatalf("settings %v: %v", settings, err)
}
out := &output{}
server := proxy.New(proxy.Params{
Config: cfg,
RequestLog: out,
ProcessLog: requestlog.NewProcessLogger(out),
})
listener, err := (&net.ListenConfig{}).Listen(t.Context(), "tcp", localhost+":0")
if err != nil {
t.Fatalf("listen: %v", err)
}
go func() {
_ = server.Serve(listener)
}()
t.Cleanup(func() {
_ = server.Close()
})
return listener.Addr().String(), out
}
// newClient returns an HTTP client that sends requests as they are made,
// with no compression of its own.
func newClient(t *testing.T) *http.Client {
t.Helper()
transport := &http.Transport{DisableCompression: true}
t.Cleanup(transport.CloseIdleConnections)
return &http.Client{Transport: transport}
}
// answer is a response as a test reads it: the status, the headers, as
// much of the body as arrived, and the error that ended the reading, nil
// when the whole body arrived.
type answer struct {
status int
header http.Header
body []byte
err error
}
// readAnswer reads all of res, and closes its body.
func readAnswer(res *http.Response) answer {
body, err := io.ReadAll(res.Body)
_ = res.Body.Close()
return answer{status: res.StatusCode, header: res.Header, body: body, err: err}
}
// newRequest makes a request for path to smallwebwaf at addr.
func newRequest(t *testing.T, method, addr, path string, body io.Reader) *http.Request {
t.Helper()
req, err := http.NewRequestWithContext(t.Context(), method, "http://"+addr+path, body)
if err != nil {
t.Fatalf("new request: %v", err)
}
return req
}
// do sends req and reads the answer.
func do(t *testing.T, req *http.Request) answer {
t.Helper()
res, err := newClient(t).Do(req)
if err != nil {
t.Fatalf("%s %s: %v", req.Method, req.URL.Path, err)
}
return readAnswer(res)
}
// get sends a GET request for path to smallwebwaf at addr.
func get(t *testing.T, addr, path string) answer {
t.Helper()
return do(t, newRequest(t, http.MethodGet, addr, path, http.NoBody))
}
// dial opens a connection to smallwebwaf at addr, for requests the HTTP
// client cannot make, such as one that stops sending halfway.
func dial(t *testing.T, addr string) net.Conn {
t.Helper()
conn, err := (&net.Dialer{}).DialContext(t.Context(), "tcp", addr)
if err != nil {
t.Fatalf("dial %s: %v", addr, err)
}
t.Cleanup(func() {
_ = conn.Close()
})
return conn
}
// send writes text to conn.
func send(t *testing.T, conn net.Conn, text string) {
t.Helper()
_, err := io.WriteString(conn, text)
if err != nil {
t.Fatalf("send: %v", err)
}
}
// readResponse reads the answer to a request sent on conn.
func readResponse(t *testing.T, conn net.Conn) answer {
t.Helper()
err := conn.SetReadDeadline(time.Now().Add(waitLimit))
if err != nil {
t.Fatalf("set read deadline: %v", err)
}
res, err := http.ReadResponse(bufio.NewReader(conn), nil)
if err != nil {
t.Fatalf("read response: %v", err)
}
return readAnswer(res)
}
// wantLine checks the request log line's status and action.
func wantLine(t *testing.T, line logLine, status int, action string) {
t.Helper()
if line.Status != status || line.Action != action {
t.Errorf("log line has status %d and action %q, want %d and %q",
line.Status, line.Action, status, action)
}
}
// wantStatus checks an answer's status.
func wantStatus(t *testing.T, got answer, status int) {
t.Helper()
if got.status != status {
t.Errorf("status %d, want %d", got.status, status)
}
}
// wantTimedOut checks that what began at start ended once shortTimeout
// had run out, and not much later.
func wantTimedOut(t *testing.T, start time.Time) {
t.Helper()
took := time.Since(start)
if took < shortTimeout || took > shortTimeout+waitLimit/2 {
t.Errorf("took %s, want %s", took, shortTimeout)
}
}