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

Milestone 1, the repo's first code. smallwebwaf passes each request to the app and the answer back unchanged, streaming bodies and WebSocket upgrades, within four timeouts (client and app, request and response) and two size limits, and writes one JSON line per request to stdout. Every setting has an SWWAF_ name and a default, and an invalid value stops the start. The repo gets the standard layout: script/ entrypoints, make targets that call them, a Dockerfile that runs the checks, and the Gitea workflow.

Disclosure: SPEC.md changed. Go's server reads the request line and headers before smallwebwaf sees the request, so slow headers are closed without an answer, and neither slow nor oversized headers get a log line.
Disclosure: standard library only.

Model: opus-5-5
This commit was merged in pull request #39.
This commit is contained in:
2026-10-03 17:24:34 +02:00
parent fd77e76177
commit d76715b0df
47 changed files with 4917 additions and 74 deletions
+171
View File
@@ -0,0 +1,171 @@
package proxy
import (
"errors"
"io"
"net/http"
"sync/atomic"
"sneak.berlin/go/smallwebwaf/internal/requestlog"
)
// errResponseTooLarge ends an app's response body that is longer than
// SWWAF_RESPONSE_MAX_BYTES.
var errResponseTooLarge = errors.New(
"the response body is over SWWAF_RESPONSE_MAX_BYTES")
// requestBody is the client's request body on its way to the app. The
// transport reads it on a goroutine of its own.
type requestBody struct {
// body is the client's body, ending in an *http.MaxBytesError past
// SWWAF_REQUEST_MAX_BYTES.
body io.ReadCloser
rq *request
// waiting is true while a Read waits for the client to send more.
waiting atomic.Bool
// received is true once the client has sent the whole body.
received atomic.Bool
// bytes is how much of the body has been read.
bytes atomic.Int64
}
// Read reads from the client's body.
func (b *requestBody) Read(p []byte) (int, error) {
b.waiting.Store(true)
n, err := b.body.Read(p)
b.waiting.Store(false)
b.bytes.Add(int64(n))
var tooLarge *http.MaxBytesError
switch {
case errors.Is(err, io.EOF):
b.received.Store(true)
b.rq.bodyReceived()
case errors.As(err, &tooLarge):
b.rq.refuse(refusal{
status: http.StatusRequestEntityTooLarge,
action: requestlog.ActionTooLarge,
})
}
return n, err
}
// Close closes the client's body.
func (b *requestBody) Close() error {
return b.body.Close()
}
// responseBody is the app's response body on its way to the client.
type responseBody struct {
// body is the app's body, ending in an *http.MaxBytesError past
// SWWAF_RESPONSE_MAX_BYTES.
body io.ReadCloser
rq *request
}
// Read reads from the app's body.
func (b *responseBody) Read(p []byte) (int, error) {
n, err := b.body.Read(p)
if err == nil {
return n, nil
}
var tooLarge *http.MaxBytesError
switch {
case errors.Is(err, io.EOF):
b.rq.responseReceived()
case errors.As(err, &tooLarge):
b.rq.refuse(refusal{
status: http.StatusBadGateway,
action: requestlog.ActionTooLarge,
})
return n, errResponseTooLarge
case b.rq.in.Context().Err() == nil:
// The answer broke off, not because the client went away. If a
// timeout cut it, that refusal came first and is the one kept.
b.rq.refuse(refusal{
status: http.StatusBadGateway,
action: requestlog.ActionUpstreamError,
})
}
return n, err
}
// Close closes the app's body.
func (b *responseBody) Close() error {
return b.body.Close()
}
// limitBody returns body, cut off with an *http.MaxBytesError after
// maxBytes, or unchanged if maxBytes is zero, which is off.
func limitBody(body io.ReadCloser, maxBytes int64) io.ReadCloser {
if maxBytes == 0 {
return body
}
// Without a ResponseWriter, MaxBytesReader only counts and cuts off.
return http.MaxBytesReader(nil, body, maxBytes)
}
// responseWriter is the response to the client. It notes the status and
// size for the log line, and the first error writing to the client.
type responseWriter struct {
http.ResponseWriter
// status is the final status sent, or zero before one is.
status int
bytes int64
err error
}
// WriteHeader sends the status and headers. An informational 1xx status
// is passed on and the final status still comes later.
func (w *responseWriter) WriteHeader(status int) {
if status >= http.StatusOK && w.status == 0 {
w.status = status
}
w.ResponseWriter.WriteHeader(status)
}
// Write sends part of the body.
func (w *responseWriter) Write(p []byte) (int, error) {
if w.status == 0 {
w.status = http.StatusOK
}
n, err := w.ResponseWriter.Write(p)
w.bytes += int64(n)
w.noteError(err)
return n, err
}
// FlushError sends what has been written so far.
// http.ResponseController calls it, as ReverseProxy does after each
// write.
func (w *responseWriter) FlushError() error {
err := http.NewResponseController(w.ResponseWriter).Flush()
w.noteError(err)
return err
}
// Unwrap lets http.ResponseController reach net/http's own
// ResponseWriter, which is how ReverseProxy takes over the connection of
// an upgraded request.
func (w *responseWriter) Unwrap() http.ResponseWriter {
return w.ResponseWriter
}
// noteError keeps the first error writing to the client.
func (w *responseWriter) noteError(err error) {
if w.err == nil {
w.err = err
}
}
+101
View File
@@ -0,0 +1,101 @@
package proxy
import (
"net/http"
"net/netip"
"slices"
"strings"
)
// peerAddress is the address of the request's TCP peer, normally traefik.
func peerAddress(r *http.Request) netip.Addr {
addrPort, err := netip.ParseAddrPort(r.RemoteAddr)
if err != nil {
return netip.Addr{}
}
return addrPort.Addr().Unmap()
}
// clientAddress works out who the client is. A peer outside the trusted
// proxies is the client, and what it says in X-Forwarded-For is ignored.
// For a peer inside them, X-Forwarded-For is read from the right, and the
// first address outside them is the client; if every address in it is
// inside, the leftmost is, and with no header, the peer. An entry that is
// not an address ends the reading, since nothing to its left can be
// believed.
func clientAddress(
peer netip.Addr, forwardedFor []string, trusted []netip.Prefix,
) netip.Addr {
client := peer
if !isInside(peer, trusted) {
return client
}
entries := strings.Split(strings.Join(forwardedFor, ","), ",")
for _, entry := range slices.Backward(entries) {
addr, err := netip.ParseAddr(strings.TrimSpace(entry))
if err != nil {
break
}
client = addr.Unmap()
if !isInside(client, trusted) {
break
}
}
return client
}
// isInside reports whether addr is in one of the netblocks.
func isInside(addr netip.Addr, netblocks []netip.Prefix) bool {
return slices.ContainsFunc(netblocks, func(netblock netip.Prefix) bool {
return netblock.Contains(addr)
})
}
// setForwardedHeaders sets the headers in which the app learns about the
// client, so that it sees what it would see from traefik directly. A
// trusted proxy's forwarded headers pass on, with the proxy's own address
// added to X-Forwarded-For. Those of any other peer are its own claims and
// are replaced: X-Forwarded-For names the peer, X-Forwarded-Host the host
// it asked for, and X-Forwarded-Proto plain http, which is how it reached
// smallwebwaf.
func setForwardedHeaders(in, out *http.Request, peer netip.Addr, trusted bool) {
forwardedFor := peer.String()
if trusted {
// ReverseProxy removes these from out before Rewrite.
for _, name := range []string{"Forwarded", "X-Forwarded-Host", "X-Forwarded-Proto"} {
values, ok := in.Header[name]
if ok {
out.Header[name] = values
}
}
prior := in.Header.Values("X-Forwarded-For")
if len(prior) > 0 {
forwardedFor = strings.Join(prior, ", ") + ", " + forwardedFor
}
out.Header.Set("X-Forwarded-For", forwardedFor)
return
}
// ReverseProxy has removed Forwarded and the three set below; these
// are the other headers in which traefik tells the app about the
// client and its request.
for _, name := range []string{
"X-Forwarded-Port", "X-Forwarded-Server", "X-Forwarded-Uri",
"X-Forwarded-Method", "X-Forwarded-Prefix", "X-Forwarded-Tls-Client-Cert",
"X-Forwarded-Tls-Client-Cert-Info", "X-Real-Ip",
} {
out.Header.Del(name)
}
out.Header.Set("X-Forwarded-For", forwardedFor)
out.Header.Set("X-Forwarded-Host", in.Host)
out.Header.Set("X-Forwarded-Proto", "http")
}
+161
View File
@@ -0,0 +1,161 @@
package proxy_test
import (
"encoding/json"
"net/http"
"testing"
)
const (
// trustLocalhost trusts the address every test connects from, and a
// network for proxies in front of it.
trustLocalhost = localhost + "/32,10.0.0.0/8"
// appHost is the host every test asks for.
appHost = "app.example"
// client is the client's address, as a proxy names it.
client = "203.0.113.9"
// forwardedFor is the header that lists the client and its proxies.
forwardedFor = "X-Forwarded-For"
// secure is the scheme a client reached traefik with.
secure = "https"
)
// appHeaders is what the app tells about the headers it received.
type appHeaders struct {
Host string `json:"host"`
ForwardedFor string `json:"forwardedFor"`
ForwardedHost string `json:"forwardedHost"`
ForwardedProto string `json:"forwardedProto"`
RealIP string `json:"realIp"`
}
// clientAddressCase is a request and what smallwebwaf makes of it.
type clientAddressCase struct {
name string
env map[string]string
header http.Header
wantClient string
wantApp appHeaders
}
func TestClientAddressAndForwardedHeaders(t *testing.T) {
t.Parallel()
for _, tc := range clientAddressCases() {
t.Run(tc.name, func(t *testing.T) {
t.Parallel()
got, line := requestWithHeaders(t, tc.env, tc.header)
tc.wantApp.Host = appHost
if got != tc.wantApp {
t.Errorf("app received %+v, want %+v", got, tc.wantApp)
}
if line.ClientIP != tc.wantClient || line.PeerIP != localhost {
t.Errorf("log line has client_ip %q and peer_ip %q, want %q and %q",
line.ClientIP, line.PeerIP, tc.wantClient, localhost)
}
})
}
}
// clientAddressCases are the requests TestClientAddressAndForwardedHeaders
// sends, from 127.0.0.1, which the default trusted proxies leave out.
func clientAddressCases() []clientAddressCase {
trusted := map[string]string{trustedProxies: trustLocalhost}
forged := http.Header{
forwardedFor: {client},
"X-Forwarded-Host": {"forged.example"},
"X-Forwarded-Proto": {secure},
"X-Real-Ip": {client},
}
replaced := appHeaders{
ForwardedFor: localhost, ForwardedHost: appHost, ForwardedProto: "http",
}
return []clientAddressCase{{
name: "a peer outside the trusted proxies is the client, " +
"and its forwarded headers are replaced",
header: forged, wantClient: localhost, wantApp: replaced,
}, {
name: "set but empty, the trusted proxies trust nothing",
env: map[string]string{trustedProxies: ""},
header: forged, wantClient: localhost, wantApp: replaced,
}, {
name: "behind a trusted peer, the client is the first address " +
"outside the trusted proxies from the right",
env: trusted,
header: http.Header{
forwardedFor: {"198.51.100.7, " + client + ", 10.0.0.2"},
"X-Forwarded-Host": {appHost},
"X-Forwarded-Proto": {secure},
"X-Real-Ip": {client},
},
wantClient: client,
wantApp: appHeaders{
ForwardedFor: "198.51.100.7, " + client + ", 10.0.0.2, " + localhost,
ForwardedHost: appHost, ForwardedProto: secure, RealIP: client,
},
}, {
name: "when every address is a trusted proxy, the leftmost is the client",
env: trusted,
header: http.Header{forwardedFor: {"10.0.0.5, 10.0.0.2"}},
wantClient: "10.0.0.5",
wantApp: appHeaders{ForwardedFor: "10.0.0.5, 10.0.0.2, " + localhost},
}, {
name: "with no header, a trusted peer is the client",
env: trusted,
wantClient: localhost,
wantApp: appHeaders{ForwardedFor: localhost},
}, {
name: "an entry that is not an address ends the reading",
env: trusted,
header: http.Header{forwardedFor: {client + ", unknown, 10.0.0.2"}},
wantClient: "10.0.0.2",
wantApp: appHeaders{
ForwardedFor: client + ", unknown, 10.0.0.2, " + localhost,
},
}, {
name: "several header lines are read as one list",
env: trusted,
header: http.Header{forwardedFor: {"2001:db8::7", "10.0.0.2"}},
wantClient: "2001:db8::7",
wantApp: appHeaders{ForwardedFor: "2001:db8::7, 10.0.0.2, " + localhost},
}}
}
// requestWithHeaders sends a request for appHost with header through
// smallwebwaf, with the settings in env, and returns the headers the app
// received and the request's log line.
func requestWithHeaders(
t *testing.T, env map[string]string, header http.Header,
) (appHeaders, logLine) {
t.Helper()
app := startApp(t, func(w http.ResponseWriter, r *http.Request) {
_ = json.NewEncoder(w).Encode(appHeaders{
Host: r.Host,
ForwardedFor: r.Header.Get(forwardedFor),
ForwardedHost: r.Header.Get("X-Forwarded-Host"),
ForwardedProto: r.Header.Get("X-Forwarded-Proto"),
RealIP: r.Header.Get("X-Real-Ip"),
})
})
addr, out := startProxy(t, app.URL, env)
req := newRequest(t, http.MethodGet, addr, "/", http.NoBody)
req.Host = appHost
req.Header = header.Clone()
answered := do(t, req)
var got appHeaders
err := json.Unmarshal(answered.body, &got)
if err != nil {
t.Fatalf("decode the app's answer %q: %v", answered.body, err)
}
return got, out.requestLine(t)
}
+146
View File
@@ -0,0 +1,146 @@
package proxy_test
import (
"bytes"
"errors"
"io"
"net/http"
"strconv"
"sync/atomic"
"testing"
"sneak.berlin/go/smallwebwaf/internal/requestlog"
)
// sizeLimit is the size limit the tests set, 1K as a setting.
const (
sizeLimit = 1 << 10
sizeLimitSetting = "1K"
)
func TestRequestBodyLimit(t *testing.T) {
t.Parallel()
for _, tc := range []struct {
name string
size int
// announced sends the size in Content-Length; otherwise the body
// is sent in chunks with no length given.
announced bool
want int
action string
// refusedBeforeApp is a refusal before anything reaches the app.
// A body over the limit with no length given has already partly
// reached the app when it is refused.
refusedBeforeApp bool
}{
{"announced, over the limit", 2 * sizeLimit, true,
http.StatusRequestEntityTooLarge, requestlog.ActionTooLarge, true},
{"announced, at the limit", sizeLimit, true,
http.StatusOK, requestlog.ActionForward, false},
{"not announced, over the limit", 4 * sizeLimit, false,
http.StatusRequestEntityTooLarge, requestlog.ActionTooLarge, false},
{"not announced, at the limit", sizeLimit, false,
http.StatusOK, requestlog.ActionForward, false},
} {
t.Run(tc.name, func(t *testing.T) {
t.Parallel()
var calls atomic.Int32
app := startApp(t, func(_ http.ResponseWriter, r *http.Request) {
calls.Add(1)
_, _ = io.Copy(io.Discard, r.Body)
})
addr, out := startProxy(t, app.URL, map[string]string{
requestMaxBytes: sizeLimitSetting,
})
var body io.Reader = bytes.NewReader(make([]byte, tc.size))
if !tc.announced {
body = io.MultiReader(body) // hides the length
}
wantStatus(t, do(t, newRequest(t, http.MethodPost, addr, "/upload", body)),
tc.want)
wantLine(t, out.requestLine(t), tc.want, tc.action)
if tc.refusedBeforeApp && calls.Load() != 0 {
t.Errorf("the app was called %d times, want never", calls.Load())
}
})
}
}
func TestResponseBodyLimit(t *testing.T) {
t.Parallel()
for _, tc := range []struct {
name string
size int
// announced sends the size in Content-Length; otherwise the body
// is sent in chunks with no length given.
announced bool
want int
action string
// received is how much of a body the client gets, and cutOff
// whether the connection is then cut.
received int
cutOff bool
}{
{"announced, over the limit", 2 * sizeLimit, true, http.StatusBadGateway,
requestlog.ActionTooLarge, len("Bad Gateway\n"), false},
{"announced, at the limit", sizeLimit, true, http.StatusOK,
requestlog.ActionForward, sizeLimit, false},
{"not announced, over the limit", 4 * sizeLimit, false, http.StatusOK,
requestlog.ActionTooLarge, sizeLimit, true},
{"not announced, at the limit", sizeLimit, false, http.StatusOK,
requestlog.ActionForward, sizeLimit, false},
} {
t.Run(tc.name, func(t *testing.T) {
t.Parallel()
app := startApp(t, func(w http.ResponseWriter, _ *http.Request) {
answerWithSize(w, tc.size, tc.announced)
})
addr, out := startProxy(t, app.URL, map[string]string{
responseMaxBytes: sizeLimitSetting,
})
got := get(t, addr, "/download")
wantStatus(t, got, tc.want)
if len(got.body) != tc.received ||
errors.Is(got.err, io.ErrUnexpectedEOF) != tc.cutOff {
t.Errorf("client got %d bytes (%v), want %d",
len(got.body), got.err, tc.received)
}
line := out.requestLine(t)
wantLine(t, line, tc.want, tc.action)
if line.UpstreamStatus != http.StatusOK {
t.Errorf("log line has upstream_status %d", line.UpstreamStatus)
}
})
}
}
// answerWithSize answers with a body of size bytes, announced in
// Content-Length or sent in chunks with no length given.
func answerWithSize(w http.ResponseWriter, size int, announced bool) {
body := make([]byte, size)
if announced {
w.Header().Set("Content-Length", strconv.Itoa(size))
_, _ = w.Write(body)
return
}
// Sending part of it before the end keeps Go's server from working
// out the length.
_, _ = w.Write(body[:size/2])
_ = http.NewResponseController(w).Flush()
_, _ = w.Write(body[size/2:])
}
+426
View File
@@ -0,0 +1,426 @@
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 != 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()
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)
}
}
+105
View File
@@ -0,0 +1,105 @@
// Package proxy passes each request to the app and the app's answer back,
// unchanged, within the size and time limits, and writes one request log
// line for each request.
package proxy
import (
"io"
"log"
"log/slog"
"net/http"
"time"
"sneak.berlin/go/smallwebwaf/internal/config"
)
// The request line and headers a client may send, and how long a
// kept-open client connection may wait for its next request, are fixed
// rather than settings. The limit on the request line and headers is
// 32 KiB, but Go's server reads 4 KiB past its MaxHeaderBytes before it
// refuses, so MaxHeaderBytes is set 4 KiB lower. The idle time is longer
// than the 90 seconds after which traefik closes a connection it is not
// using, so traefik never sends a request on a connection smallwebwaf is
// closing.
const (
requestHeaderMaxBytes = 32<<10 - 4<<10
clientIdleTimeout = 120 * time.Second
)
// How smallwebwaf keeps connections to the app open between requests.
const (
appIdleConns = 100
appIdleConnTimeout = 90 * time.Second
)
// Params are what New needs.
type Params struct {
Config *config.Config
// RequestLog receives one JSON line per request.
RequestLog io.Writer
// ProcessLog receives the process's own messages.
ProcessLog *slog.Logger
}
// New returns the server smallwebwaf runs: each request it reads passes
// through the proxy. Go's server itself refuses headers over 32 KiB, with
// 431, closes a connection idle for 120 seconds, and applies
// SWWAF_CLIENT_REQUEST_TIMEOUT while the headers arrive; the proxy
// applies the timeouts and size limits from then on.
func New(params Params) *http.Server {
errorLog := slog.NewLogLogger(params.ProcessLog.Handler(), slog.LevelWarn)
return &http.Server{
Addr: params.Config.ListenAddr,
Handler: &handler{
config: params.Config,
requestLog: params.RequestLog,
processLog: params.ProcessLog,
errorLog: errorLog,
transport: newTransport(),
},
ReadHeaderTimeout: params.Config.ClientRequestTimeout,
IdleTimeout: clientIdleTimeout,
MaxHeaderBytes: requestHeaderMaxBytes,
ErrorLog: errorLog,
}
}
// handler is the proxy. It holds what every request shares; what belongs
// to one request is in a request.
type handler struct {
config *config.Config
requestLog io.Writer
processLog *slog.Logger
errorLog *log.Logger
transport http.RoundTripper
}
// newTransport returns what carries requests to the app. It never goes
// through a proxy named in the environment, and leaves the app's answers
// compressed or not as the app sent them.
func newTransport() *http.Transport {
return &http.Transport{
MaxIdleConns: appIdleConns,
MaxIdleConnsPerHost: appIdleConns,
IdleConnTimeout: appIdleConnTimeout,
DisableCompression: true,
}
}
// ServeHTTP handles one request: it works out the client, runs the
// checks, passes the request to the app and the answer back within the
// limits, and writes the request's log line.
func (h *handler) ServeHTTP(w http.ResponseWriter, r *http.Request) {
rq := h.newRequest(w, r)
defer rq.finish()
refused := rq.check()
if refused != nil {
rq.answer(*refused)
return
}
rq.forward(r.Context())
}
+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)
}
}
+452
View File
@@ -0,0 +1,452 @@
package proxy
import (
"context"
"errors"
"net/http"
"net/http/httptrace"
"net/http/httputil"
"net/netip"
"os"
"sync"
"sync/atomic"
"time"
"sneak.berlin/go/smallwebwaf/internal/requestlog"
)
// flushAfterEachWrite has ReverseProxy pass on each part of the app's
// answer as soon as it arrives.
const flushAfterEachWrite time.Duration = -1
// refusal is smallwebwaf refusing a request, or refusing to go on with it:
// the status the client is answered if the response has not started yet,
// and the action the log line names.
type refusal struct {
status int
action string
}
// request is one request on its way through smallwebwaf, from the moment
// its headers have been read to its log line.
type request struct {
h *handler
in *http.Request
// rc sets the deadlines of the connection to the client.
rc *http.ResponseController
out *responseWriter
body *requestBody // nil for a request without a body
line requestlog.Line
peer netip.Addr
peerTrusted bool
start time.Time
// upstreamStart is when the request was handed to the app.
upstreamStart time.Time
// cancel ends the request to the app.
cancel context.CancelFunc
// refused is the first refusal, from whichever goroutine meets it.
refused atomic.Pointer[refusal]
// complete is true once the app's whole answer has been passed on.
complete bool
// mu guards what follows. The timeouts run on goroutines of their
// own, and the transport starts and stops them from its own; once
// timersStopped is set, none of them acts any more.
mu sync.Mutex
timersStopped bool
clientRequestTimer *time.Timer
upstreamRequestTimer *time.Timer
upstreamResponseTimer *time.Timer
// requestSent is when the app had been sent the whole request.
requestSent time.Time
}
// newRequest starts handling r: it notes the time and works out the
// client.
func (h *handler) newRequest(w http.ResponseWriter, r *http.Request) *request {
start := time.Now()
peer := peerAddress(r)
trusted := h.config.TrustedProxies
client := clientAddress(peer, r.Header.Values("X-Forwarded-For"), trusted)
rq := &request{
h: h,
in: r,
rc: http.NewResponseController(w),
out: &responseWriter{ResponseWriter: w},
peer: peer,
peerTrusted: isInside(peer, trusted),
start: start,
line: requestlog.Line{
Time: requestlog.FormatTime(start),
ClientIP: client.String(),
PeerIP: peer.String(),
Method: r.Method,
Host: r.Host,
Path: r.URL.EscapedPath(),
Query: r.URL.RawQuery,
Protocol: r.Proto,
Referer: r.Referer(),
UserAgent: r.UserAgent(),
Action: requestlog.ActionForward,
},
}
if r.Body != http.NoBody {
rq.body = &requestBody{body: limitBody(r.Body, h.config.RequestMaxBytes), rq: rq}
}
return rq
}
// check is the one place where a request can be refused once its client
// is known, before its body is read or anything reaches the app; the rate
// limits and country lists of milestone 2 go here. It returns nil to let
// the request through.
func (rq *request) check() *refusal {
maxBytes := rq.h.config.RequestMaxBytes
if maxBytes > 0 && rq.in.ContentLength > maxBytes {
return &refusal{
status: http.StatusRequestEntityTooLarge,
action: requestlog.ActionTooLarge,
}
}
return nil
}
// forward passes the request to the app and the app's answer back. ctx
// is the request's own context.
func (rq *request) forward(ctx context.Context) {
ctx, cancel := context.WithCancel(ctx)
defer cancel()
rq.cancel = cancel
ctx = httptrace.WithClientTrace(ctx, &httptrace.ClientTrace{
WroteRequest: rq.wroteRequest,
})
out := rq.in.WithContext(ctx)
if rq.body != nil {
out.Body = rq.body
}
reverseProxy := &httputil.ReverseProxy{
Rewrite: rq.rewrite,
Transport: rq.h.transport,
FlushInterval: flushAfterEachWrite,
ErrorLog: rq.h.errorLog,
ModifyResponse: rq.modifyResponse,
ErrorHandler: rq.answerError,
}
rq.startRequestTimers()
rq.upstreamStart = time.Now()
reverseProxy.ServeHTTP(rq.out, out)
}
// rewrite makes the request the app receives: the client's request,
// unchanged, sent to SWWAF_UPSTREAM_URL, with the forwarded headers set.
func (rq *request) rewrite(pr *httputil.ProxyRequest) {
upstream := rq.h.config.UpstreamURL
pr.Out.URL.Scheme = upstream.Scheme
pr.Out.URL.Host = upstream.Host
// ReverseProxy drops query parameters it cannot parse; the app gets
// the query as the client sent it.
pr.Out.URL.RawQuery = pr.In.URL.RawQuery
setForwardedHeaders(pr.In, pr.Out, rq.peer, rq.peerTrusted)
}
// modifyResponse looks at the app's answer before ReverseProxy passes it
// on.
func (rq *request) modifyResponse(res *http.Response) error {
rq.line.UpstreamStatus = res.StatusCode
if res.StatusCode == http.StatusSwitchingProtocols {
// An upgraded connection, such as a WebSocket, is not cut by the
// timeouts. ReverseProxy writes this answer straight to the
// connection it takes over, not through rq.out.
rq.stopTimers()
rq.out.status = res.StatusCode
return nil
}
maxBytes := rq.h.config.ResponseMaxBytes
if maxBytes > 0 && res.Body != http.NoBody && res.ContentLength > maxBytes {
rq.refuse(refusal{status: http.StatusBadGateway, action: requestlog.ActionTooLarge})
return errResponseTooLarge
}
res.Body = &responseBody{body: limitBody(res.Body, maxBytes), rq: rq}
rq.startClientResponseTimeout()
return nil
}
// answerError is ReverseProxy's ErrorHandler: the request could not be
// passed to the app, or the app's answer cannot be passed on.
func (rq *request) answerError(_ http.ResponseWriter, _ *http.Request, err error) {
refused := rq.refused.Load()
if refused == nil {
if rq.in.Context().Err() != nil {
return // the client has gone, and there is no one to answer
}
rq.h.processLog.Warn("request to the app failed", "error", err.Error())
refused = &refusal{
status: http.StatusBadGateway,
action: requestlog.ActionUpstreamError,
}
}
rq.answer(*refused)
}
// answer sends smallwebwaf's own answer, unless the response has already
// started, and records the refusal for the log line.
func (rq *request) answer(r refusal) {
rq.refused.CompareAndSwap(nil, &r)
if rq.out.status != 0 {
return // too late to answer: the connection can only be cut
}
// A client found too slow is read no more; any other may go on
// sending until its time is up, so that Go's server can read the
// rest of the body and end the request cleanly.
deadline := rq.clientRequestDeadline()
if r.status == http.StatusRequestTimeout {
deadline = time.Now()
}
rq.stopReadingBody(deadline)
timeout := rq.h.config.ClientResponseTimeout
if timeout > 0 {
_ = rq.rc.SetWriteDeadline(time.Now().Add(timeout))
}
http.Error(rq.out, http.StatusText(r.status), r.status)
}
// refuse records r, unless an earlier refusal was, and ends the request
// to the app.
func (rq *request) refuse(r refusal) {
rq.refused.CompareAndSwap(nil, &r)
rq.cancel()
}
// finish ends the request's timeouts and writes its log line.
func (rq *request) finish() {
rq.stopTimers()
refused := rq.refused.Load()
if refused == nil {
rq.stopReadingBody(rq.clientRequestDeadline())
}
line := &rq.line
line.Status = rq.out.status
line.ResponseBytes = rq.out.bytes
if rq.body != nil {
line.RequestBytes = rq.body.bytes.Load()
}
switch {
case refused != nil:
line.Action = refused.action
case errors.Is(rq.out.err, os.ErrDeadlineExceeded):
// The client took longer than SWWAF_CLIENT_RESPONSE_TIMEOUT to
// take the response.
line.Action = requestlog.ActionTimedOut
case !rq.complete && (rq.out.err != nil || rq.in.Context().Err() != nil):
line.Aborted = true
}
now := time.Now()
line.DurationTotal = requestlog.Milliseconds(now.Sub(rq.start))
if !rq.upstreamStart.IsZero() {
line.DurationUpstreamTotal = requestlog.Milliseconds(now.Sub(rq.upstreamStart))
}
err := requestlog.Write(rq.h.requestLog, line)
if err != nil {
rq.h.processLog.Error("writing the request log failed", "error", err.Error())
}
}
// clientRequestDeadline is when the client must have sent its whole
// request, or zero when SWWAF_CLIENT_REQUEST_TIMEOUT is off.
func (rq *request) clientRequestDeadline() time.Time {
timeout := rq.h.config.ClientRequestTimeout
if timeout == 0 {
return time.Time{}
}
return rq.start.Add(timeout)
}
// stopReadingBody ends, at deadline, the reading of a client body that has
// not arrived whole: Go's server then reads no more of it, and closes the
// connection after the answer.
func (rq *request) stopReadingBody(deadline time.Time) {
if rq.body == nil || rq.body.received.Load() {
return
}
_ = rq.rc.SetReadDeadline(deadline)
}
// startRequestTimers starts the timeouts that run while the request goes
// to the app: SWWAF_CLIENT_REQUEST_TIMEOUT until the client has sent its
// whole body, and SWWAF_UPSTREAM_REQUEST_TIMEOUT until the app has been
// sent the whole request.
func (rq *request) startRequestTimers() {
rq.mu.Lock()
defer rq.mu.Unlock()
if rq.body != nil && rq.h.config.ClientRequestTimeout > 0 {
rq.clientRequestTimer = time.AfterFunc(
time.Until(rq.clientRequestDeadline()), rq.requestTimedOut)
}
timeout := rq.h.config.UpstreamRequestTimeout
if timeout > 0 {
rq.upstreamRequestTimer = time.AfterFunc(timeout, rq.requestTimedOut)
}
}
// requestTimedOut is called when a request timeout runs out while the
// request is still on its way to the app. The answer names the side
// smallwebwaf was waiting on at that moment: 408 when it was waiting for
// the client to send more of its body, 504 when it was waiting for the
// app to be reached or to take what it had.
func (rq *request) requestTimedOut() {
rq.mu.Lock()
defer rq.mu.Unlock()
if rq.timersStopped {
return
}
if rq.body == nil || !rq.body.waiting.Load() {
rq.refuse(refusal{
status: http.StatusGatewayTimeout,
action: requestlog.ActionTimedOut,
})
return
}
rq.refuse(refusal{
status: http.StatusRequestTimeout,
action: requestlog.ActionTimedOut,
})
// The transport gives up on the app only once its Read of the
// client's body returns, so that Read is ended now. The lock keeps
// this from reaching the connection after the request is handled.
_ = rq.rc.SetReadDeadline(time.Now())
}
// bodyReceived is called once the client has sent its whole body.
func (rq *request) bodyReceived() {
rq.mu.Lock()
defer rq.mu.Unlock()
stopTimer(rq.clientRequestTimer)
}
// wroteRequest is called once the app has been sent the whole request:
// the request timeouts end and SWWAF_UPSTREAM_RESPONSE_TIMEOUT starts.
func (rq *request) wroteRequest(info httptrace.WroteRequestInfo) {
if info.Err != nil {
return // the transport gives up, or tries again
}
rq.mu.Lock()
defer rq.mu.Unlock()
if rq.timersStopped {
return
}
stopTimer(rq.clientRequestTimer)
stopTimer(rq.upstreamRequestTimer)
rq.requestSent = time.Now()
timeout := rq.h.config.UpstreamResponseTimeout
if timeout > 0 {
rq.upstreamResponseTimer = time.AfterFunc(timeout, rq.responseTimedOut)
}
}
// responseTimedOut is called when SWWAF_UPSTREAM_RESPONSE_TIMEOUT runs out
// before the app has sent its whole answer.
func (rq *request) responseTimedOut() {
rq.mu.Lock()
defer rq.mu.Unlock()
if !rq.timersStopped {
rq.refuse(refusal{
status: http.StatusGatewayTimeout,
action: requestlog.ActionTimedOut,
})
}
}
// responseReceived is called once the app has sent its whole answer.
func (rq *request) responseReceived() {
rq.complete = true
rq.stopTimers()
}
// startClientResponseTimeout sets SWWAF_CLIENT_RESPONSE_TIMEOUT on the
// connection to the client: the response must reach the client within it
// of the end of the request, or of now if the app answers before it has
// the whole request.
func (rq *request) startClientResponseTimeout() {
timeout := rq.h.config.ClientResponseTimeout
if timeout == 0 {
return
}
from := rq.sentAt()
if from.IsZero() {
from = time.Now()
}
_ = rq.rc.SetWriteDeadline(from.Add(timeout))
}
// sentAt is when the app had been sent the whole request, or zero.
func (rq *request) sentAt() time.Time {
rq.mu.Lock()
defer rq.mu.Unlock()
return rq.requestSent
}
// stopTimers stops the request's timeouts and keeps any from starting
// later: the app's answer is complete, the connection upgraded, or the
// request handled.
func (rq *request) stopTimers() {
rq.mu.Lock()
defer rq.mu.Unlock()
rq.timersStopped = true
stopTimer(rq.clientRequestTimer)
stopTimer(rq.upstreamRequestTimer)
stopTimer(rq.upstreamResponseTimer)
}
// stopTimer stops t, which is nil when its timeout is off.
func stopTimer(t *time.Timer) {
if t != nil {
t.Stop()
}
}
+258
View File
@@ -0,0 +1,258 @@
package proxy_test
import (
"errors"
"io"
"net"
"net/http"
"strconv"
"sync"
"testing"
"time"
"sneak.berlin/go/smallwebwaf/internal/requestlog"
)
// largeBodySize is more than the connections between the client,
// smallwebwaf and the app can hold while nobody reads, so that a sender
// soon waits.
const largeBodySize = 64 << 20
// writeSize is how much a test sender writes at a time.
const writeSize = 32 << 10
func TestRequestTimeouts(t *testing.T) {
t.Parallel()
for _, tc := range []struct {
name string
env map[string]string
// appTakesNothing has the app never read, while the client sends
// as fast as it can; otherwise the app reads, and the client
// stops sending halfway.
appTakesNothing bool
want int
}{
{
name: "client request timeout, waiting on the client",
env: map[string]string{clientRequestTimeout: shortTimeoutSetting},
want: http.StatusRequestTimeout,
},
{
name: "upstream request timeout, waiting on the client",
env: map[string]string{
upstreamRequestTimeout: shortTimeoutSetting,
clientRequestTimeout: longTimeoutSetting,
},
want: http.StatusRequestTimeout,
},
{
name: "upstream request timeout, waiting on the app",
env: map[string]string{upstreamRequestTimeout: shortTimeoutSetting},
appTakesNothing: true,
want: http.StatusGatewayTimeout,
},
{
name: "client request timeout, waiting on the app",
env: map[string]string{
clientRequestTimeout: shortTimeoutSetting,
upstreamRequestTimeout: longTimeoutSetting,
},
appTakesNothing: true,
want: http.StatusGatewayTimeout,
},
} {
t.Run(tc.name, func(t *testing.T) {
t.Parallel()
var (
appURL string
sendRequest func(*testing.T, string) net.Conn
)
if tc.appTakesNothing {
appURL, sendRequest = startAppThatTakesNothing(t), sendLargeBody
} else {
appURL, sendRequest = startApp(t, readBody).URL, sendPartOfBody
}
addr, out := startProxy(t, appURL, tc.env)
start := time.Now()
conn := sendRequest(t, addr)
wantStatus(t, readResponse(t, conn), tc.want)
wantTimedOut(t, start)
wantLine(t, out.requestLine(t), tc.want, requestlog.ActionTimedOut)
})
}
}
// readBody is an app that reads the request body, then answers.
func readBody(_ http.ResponseWriter, r *http.Request) {
_, _ = io.Copy(io.Discard, r.Body)
}
// startAppThatTakesNothing starts an app that accepts connections and
// never reads from them, and returns its URL.
func startAppThatTakesNothing(t *testing.T) string {
t.Helper()
listener, err := (&net.ListenConfig{}).Listen(t.Context(), "tcp", localhost+":0")
if err != nil {
t.Fatalf("listen: %v", err)
}
var (
mu sync.Mutex
held []net.Conn
)
hold := func(conn net.Conn) {
mu.Lock()
defer mu.Unlock()
held = append(held, conn)
}
go func() {
for {
conn, err := listener.Accept()
if err != nil {
return
}
hold(conn)
}
}()
t.Cleanup(func() {
_ = listener.Close()
mu.Lock()
defer mu.Unlock()
for _, conn := range held {
_ = conn.Close()
}
})
return "http://" + listener.Addr().String()
}
// sendPartOfBody sends a request that announces a large body, and only
// the first bytes of it.
func sendPartOfBody(t *testing.T, addr string) net.Conn {
t.Helper()
conn := dial(t, addr)
send(t, conn, "POST /upload HTTP/1.1\r\nHost: app\r\nContent-Length: "+
strconv.Itoa(largeBodySize)+"\r\n\r\nthe first bytes")
return conn
}
// sendLargeBody sends a request with a large body, as fast as smallwebwaf
// takes it, from a goroutine of its own.
func sendLargeBody(t *testing.T, addr string) net.Conn {
t.Helper()
conn := dial(t, addr)
send(t, conn, "POST /upload HTTP/1.1\r\nHost: app\r\nContent-Length: "+
strconv.Itoa(largeBodySize)+"\r\n\r\n")
go func() {
chunk := make([]byte, writeSize)
for range largeBodySize / writeSize {
_, err := conn.Write(chunk)
if err != nil {
return
}
}
}()
return conn
}
func TestAppTooSlowToAnswer(t *testing.T) {
t.Parallel()
app := startApp(t, func(_ http.ResponseWriter, r *http.Request) {
<-r.Context().Done()
})
addr, out := startProxy(t, app.URL, map[string]string{
upstreamResponseTimeout: shortTimeoutSetting,
})
start := time.Now()
wantStatus(t, get(t, addr, "/slow"), http.StatusGatewayTimeout)
wantTimedOut(t, start)
line := out.requestLine(t)
wantLine(t, line, http.StatusGatewayTimeout, requestlog.ActionTimedOut)
_, answered := line.fields["upstream_status"]
if answered {
t.Errorf("log line has upstream_status %v for an app that never answered",
line.fields["upstream_status"])
}
}
func TestAppTooSlowToFinishItsAnswer(t *testing.T) {
t.Parallel()
app := startApp(t, func(w http.ResponseWriter, r *http.Request) {
_, _ = io.WriteString(w, "the first part")
_ = http.NewResponseController(w).Flush()
<-r.Context().Done()
})
addr, out := startProxy(t, app.URL, map[string]string{
upstreamResponseTimeout: shortTimeoutSetting,
})
start := time.Now()
got := get(t, addr, "/slow")
wantStatus(t, got, http.StatusOK)
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)
}
wantTimedOut(t, start)
line := out.requestLine(t)
wantLine(t, line, http.StatusOK, requestlog.ActionTimedOut)
if line.UpstreamStatus != http.StatusOK {
t.Errorf("log line has upstream_status %d, want %d",
line.UpstreamStatus, http.StatusOK)
}
}
func TestClientTooSlowToTakeTheAnswer(t *testing.T) {
t.Parallel()
app := startApp(t, func(w http.ResponseWriter, _ *http.Request) {
chunk := make([]byte, writeSize)
for range largeBodySize / writeSize {
_, err := w.Write(chunk)
if err != nil {
return
}
}
})
addr, out := startProxy(t, app.URL, map[string]string{
clientResponseTimeout: shortTimeoutSetting,
})
start := time.Now()
// The client asks, and never reads the answer.
conn := dial(t, addr)
send(t, conn, "GET /large HTTP/1.1\r\nHost: app\r\n\r\n")
line := out.requestLine(t)
wantTimedOut(t, start)
wantLine(t, line, http.StatusOK, requestlog.ActionTimedOut)
}