Files
smallwebwaf/internal/proxy/requestlog_test.go
T
clawbot 7e3aa6a5bb
check / check (push) Waiting to run
Log the rest of the request log's fields (closes #79)
Each request log line now has the fields "Request log" in SPEC.md lists
whose features are built: instance (SWWAF_INSTANCE_NAME), scheme,
request_id (a trusted proxy's X-Request-ID or a new one, sent on to the
app), forwarded_for, client_group, content_type, content_length, the
headers SWWAF_LOG_REQUEST_HEADERS names, has_authorization, has_cookie,
websocket, response_content_type, cache_control, location, counts and
the timings. Authorization, Cookie and Set-Cookie values are never
logged. An entry of SWWAF_LOG_REQUEST_HEADERS that is not a header name
stops the start.

Deviation: counts has request totals only.
Deviation: SWWAF_INSTANCE_NAME is on request lines only.

Model: opus-5-5
2026-10-06 12:46:43 +00:00

338 lines
11 KiB
Go

package proxy_test
import (
"io"
"maps"
"math"
"net/http"
"reflect"
"slices"
"strings"
"testing"
"time"
"sneak.berlin/go/smallwebwaf/internal/proxy"
"sneak.berlin/go/smallwebwaf/internal/ratelimit"
"sneak.berlin/go/smallwebwaf/internal/requestlog"
)
const (
// requestIDHeader carries the request's id.
requestIDHeader = "X-Request-ID"
// instance is the SWWAF_INSTANCE_NAME a test sets.
instance = "fsn1app1/gitea"
// ipv6Client is a client on IPv6, and ipv6Group the netblock the rate
// limits count it as.
ipv6Client = "2001:db8::7"
ipv6Group = "2001:db8::/64"
)
func TestLogLineHasEachFieldWhereItApplies(t *testing.T) {
t.Parallel()
received := make(chan string, 2) // the request ids the app received
app := startApp(t, func(w http.ResponseWriter, r *http.Request) {
received <- r.Header.Get(requestIDHeader)
_, _ = io.Copy(io.Discard, r.Body)
if r.URL.Path != "/full" {
w.WriteHeader(http.StatusNoContent)
return
}
w.Header().Set("Content-Type", "text/html")
w.Header().Set("Cache-Control", "no-store")
w.Header().Set("Location", "/elsewhere")
w.WriteHeader(http.StatusFound)
_, _ = io.WriteString(w, "moved")
})
addr, out := startProxy(t, app.URL, map[string]string{
trustedProxies: trustLocalhost,
rateLimitExemptNets: localhost,
instanceName: instance,
logRequestHeaders: "Accept,x-custom,Authorization,cookie,SET-COOKIE",
})
// This request comes from ipv6Client through a trusted proxy, with a
// body and each header the log line looks at, and is answered with a
// redirect.
conn := dial(t, addr)
send(t, conn, "POST /full HTTP/1.1\r\nHost: "+appHost+"\r\n"+
forwardedFor+": 198.51.100.7, "+ipv6Client+"\r\n"+
forwardedProto+": "+secure+"\r\n"+requestIDHeader+": from-traefik\r\n"+
"Content-Type: application/x-www-form-urlencoded\r\nContent-Length: 3\r\n"+
"Accept: text/html\r\nX-Custom: one\r\nX-Custom: two\r\n"+
"Authorization: Bearer secret-token\r\nCookie: session=secret-cookie\r\n"+
"Set-Cookie: secret-set-cookie\r\n\r\na=b")
wantStatus(t, readResponse(t, conn), http.StatusFound)
// A request's log line can come after its answer: each is waited for
// before the next request, so that the lines are in order.
full := out.requestLines(t, 1)[0]
// This one comes from 127.0.0.1, which the rate limits do not count,
// with a body whose length it does not announce and no header the log
// line looks at, and is answered with 204 and no header.
conn = dial(t, addr)
send(t, conn, "POST /bare HTTP/1.1\r\nHost: "+appHost+"\r\n"+
"Transfer-Encoding: chunked\r\n\r\n0\r\n\r\n")
wantStatus(t, readResponse(t, conn), http.StatusNoContent)
bare := out.requestLines(t, 2)[1]
wantFullLine(t, full)
wantBareLine(t, bare)
for _, line := range []logLine{full, bare} {
got := <-received
if got != line.RequestID {
t.Errorf("the app received request id %q, the log line has %q",
got, line.RequestID)
}
}
if strings.Contains(out.text(), "secret") {
t.Errorf("a value of Authorization, Cookie or Set-Cookie is logged:\n%s",
out.text())
}
}
// wantFullLine checks the log line of the request with every header the
// line looks at. Its timings are checked by TestTimingsAreInOrder.
func wantFullLine(t *testing.T, line logLine) {
t.Helper()
headers := map[string]string{"accept": "text/html", "x-custom": "one, two"}
want := withTimings(line, requestlog.Line{
Type: requestType, Time: line.Time, Instance: instance,
ClientIP: ipv6Client, Method: http.MethodPost, Scheme: secure,
Host: appHost, Path: "/full", Protocol: protocol,
Status: http.StatusFound, RequestBytes: 3, ResponseBytes: 5,
RequestID: "from-traefik", PeerIP: localhost,
ForwardedFor: "198.51.100.7, " + ipv6Client, ClientGroup: ipv6Group,
ContentType: "application/x-www-form-urlencoded", ContentLength: 3,
RequestHeaders: headers, HasAuthorization: true, HasCookie: true,
ResponseContentType: "text/html", UpstreamStatus: http.StatusFound,
CacheControl: "no-store", Location: "/elsewhere",
Action: requestlog.ActionForward,
Counts: ratelimit.Counts{Minute: 1, Hour: 1, Day: 1},
})
if !reflect.DeepEqual(line.Line, want) {
t.Errorf("log line\n%+v\nwant\n%+v", line.Line, want)
}
}
// wantBareLine checks the log line of the request with none of them, and
// that the fields that do not apply to it are left out.
func wantBareLine(t *testing.T, line logLine) {
t.Helper()
want := withTimings(line, requestlog.Line{
Type: requestType, Time: line.Time, Instance: instance,
ClientIP: localhost, Method: http.MethodPost, Scheme: plain,
Host: appHost, Path: "/bare", Protocol: protocol,
Status: http.StatusNoContent, RequestID: line.RequestID,
PeerIP: localhost, ClientGroup: localhost + "/32",
UpstreamStatus: http.StatusNoContent, Action: requestlog.ActionForward,
})
if !reflect.DeepEqual(line.Line, want) || line.RequestID == "" {
t.Errorf("log line\n%+v\nwant\n%+v, with a request id", line.Line, want)
}
for _, name := range []string{
"forwarded_for", "content_type", "content_length", "request_headers",
"has_authorization", "has_cookie", "websocket", "response_content_type",
"cache_control", "location", "counts",
} {
_, present := line.fields[name]
if present {
t.Errorf("log line has %s, which does not apply", name)
}
}
}
// withTimings returns want with the timings of line.
func withTimings(line logLine, want requestlog.Line) requestlog.Line {
want.DurationTotal = line.DurationTotal
want.DurationChecks = line.DurationChecks
want.DurationUpstreamConnect = line.DurationUpstreamConnect
want.DurationUpstreamFirstByte = line.DurationUpstreamFirstByte
want.DurationUpstreamTotal = line.DurationUpstreamTotal
return want
}
func TestRequestIDAndSchemeComeOnlyFromATrustedProxy(t *testing.T) {
t.Parallel()
const sentID = "from-traefik"
sent := http.Header{requestIDHeader: {sentID}, forwardedProto: {secure}}
trusted := map[string]string{trustedProxies: trustLocalhost}
for _, tc := range []struct {
name string
env map[string]string
header http.Header
// wantID is the request id logged, "" for a new one.
wantID, wantScheme string
}{
{"a trusted proxy's are kept", trusted, sent, sentID, secure},
{"without them, the id is new and the scheme http", trusted, nil, "", plain},
{"another peer's are replaced", nil, sent, "", plain},
} {
t.Run(tc.name, func(t *testing.T) {
t.Parallel()
received := make(chan string, 2)
app := startApp(t, func(_ http.ResponseWriter, r *http.Request) {
received <- r.Header.Get(requestIDHeader)
})
addr, out := startProxy(t, app.URL, tc.env)
// Two requests, so that two new ids can be told apart.
ids := make([]string, 0, 2)
for i := range 2 {
req := newRequest(t, http.MethodGet, addr, "/", http.NoBody)
maps.Copy(req.Header, tc.header)
wantStatus(t, do(t, req), http.StatusOK)
line := out.requestLines(t, i+1)[i]
ids = append(ids, line.RequestID)
got := <-received
if line.RequestID != got || line.Scheme != tc.wantScheme {
t.Errorf("log line has request_id %q and scheme %q, and the "+
"app received id %q; want the same id and scheme %q",
line.RequestID, line.Scheme, got, tc.wantScheme)
}
}
switch {
case tc.wantID != "" && (ids[0] != tc.wantID || ids[1] != tc.wantID):
t.Errorf("request ids %q, want %q", ids, tc.wantID)
case tc.wantID == "" && (slices.Contains(ids, sentID) ||
slices.Contains(ids, "") || ids[0] == ids[1]):
t.Errorf("request ids %q, want two new ones", ids)
}
})
}
}
func TestTimingsAreInOrder(t *testing.T) {
t.Parallel()
const denied = "192.0.2.50" // in SWWAF_DENY_NETS
app := startApp(t, func(w http.ResponseWriter, _ *http.Request) {
// The pauses set the times apart; a hold-up of the test only
// lengthens them.
time.Sleep(time.Millisecond)
w.WriteHeader(http.StatusOK)
_ = http.NewResponseController(w).Flush()
time.Sleep(time.Millisecond)
_, _ = io.WriteString(w, "done")
})
addr, out := startProxy(t, app.URL, map[string]string{
trustedProxies: trustLocalhost,
denyNets: denied,
})
// Each log line is waited for before the next request, so that the
// lines are in order.
wantStatus(t, get(t, addr, "/"), http.StatusOK)
forwarded := out.requestLines(t, 1)[0]
req := newRequest(t, http.MethodGet, addr, "/", http.NoBody)
req.Header.Set(forwardedFor, denied)
wantStatus(t, do(t, req), http.StatusForbidden)
refused := out.requestLines(t, 2)[1]
wantStatus(t, get(t, addr, proxy.HealthPath), http.StatusOK)
health := out.requestLines(t, 3)[2]
// A request passed to the app has every timing; one refused, none of
// the app's; the health check, which runs no check, only the total.
wantTimings(t, forwarded, "duration_total", "duration_checks",
"duration_upstream_connect", "duration_upstream_first_byte",
"duration_upstream_total")
wantTimings(t, refused, "duration_total", "duration_checks")
wantTimings(t, health, "duration_total")
if t.Failed() {
return
}
// In whole microseconds, as they are logged, so that the sum below is
// exact.
total := microseconds(forwarded.DurationTotal)
checks := microseconds(*forwarded.DurationChecks)
connect := microseconds(*forwarded.DurationUpstreamConnect)
firstByte := microseconds(*forwarded.DurationUpstreamFirstByte)
upstream := microseconds(*forwarded.DurationUpstreamTotal)
// The checks end before the request is handed to the app, and the
// connection comes before the answer, which the app ends after a
// pause.
if checks+upstream > total || connect >= firstByte || firstByte >= upstream {
t.Errorf("timings in microseconds: total %d, checks %d, connect %d, "+
"first byte %d, upstream total %d", total, checks, connect, firstByte,
upstream)
}
if *refused.DurationChecks > refused.DurationTotal {
t.Errorf("refused request's checks took %v of %v milliseconds",
*refused.DurationChecks, refused.DurationTotal)
}
}
// wantTimings checks that the timings named are the only ones line has.
func wantTimings(t *testing.T, line logLine, want ...string) {
t.Helper()
var got []string
for name := range line.fields {
if strings.HasPrefix(name, "duration_") {
got = append(got, name)
}
}
slices.Sort(got)
slices.Sort(want)
if !slices.Equal(got, want) {
t.Errorf("log line of %s has timings %v, want %v", line.Path, got, want)
}
}
// microseconds is a timing in whole microseconds.
func microseconds(milliseconds float64) int64 {
return int64(math.Round(milliseconds * 1000))
}
func TestLogsAnUpgradedConnection(t *testing.T) {
t.Parallel()
app := startApp(t, echoAfterUpgrade)
addr, out := startProxy(t, app.URL, nil)
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")
wantStatus(t, readResponse(t, conn), http.StatusSwitchingProtocols)
_ = conn.Close()
line := out.requestLine(t)
if line.fields["websocket"] != true {
t.Errorf("log line has websocket %v, want true", line.fields["websocket"])
}
}