Pass-through proxy with timeouts, size limits and a request log (closes #13)
check / check (push) Successful in 1m29s
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:
@@ -0,0 +1,352 @@
|
||||
// Package config reads smallwebwaf's settings. Every setting is an
|
||||
// environment variable whose name starts with SWWAF_, every setting has a
|
||||
// default, and this package is the one place they are read.
|
||||
package config
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"math"
|
||||
"net"
|
||||
"net/netip"
|
||||
"net/url"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
)
|
||||
|
||||
// Config is smallwebwaf's settings. A timeout or size of zero is off.
|
||||
type Config struct {
|
||||
// ListenAddr is where smallwebwaf listens (SWWAF_LISTEN_ADDR).
|
||||
ListenAddr string
|
||||
// UpstreamURL is the app (SWWAF_UPSTREAM_URL).
|
||||
UpstreamURL *url.URL
|
||||
// TrustedProxies are the netblocks whose X-Forwarded-For is
|
||||
// believed (SWWAF_TRUSTED_PROXIES).
|
||||
TrustedProxies []netip.Prefix
|
||||
// ClientRequestTimeout bounds reading the whole request from the
|
||||
// client (SWWAF_CLIENT_REQUEST_TIMEOUT).
|
||||
ClientRequestTimeout time.Duration
|
||||
// ClientResponseTimeout bounds writing the whole response to the
|
||||
// client (SWWAF_CLIENT_RESPONSE_TIMEOUT).
|
||||
ClientResponseTimeout time.Duration
|
||||
// UpstreamRequestTimeout bounds connecting to the app and writing
|
||||
// the whole request to it (SWWAF_UPSTREAM_REQUEST_TIMEOUT).
|
||||
UpstreamRequestTimeout time.Duration
|
||||
// UpstreamResponseTimeout bounds reading the whole response from
|
||||
// the app (SWWAF_UPSTREAM_RESPONSE_TIMEOUT).
|
||||
UpstreamResponseTimeout time.Duration
|
||||
// RequestMaxBytes is the largest request body
|
||||
// (SWWAF_REQUEST_MAX_BYTES).
|
||||
RequestMaxBytes int64
|
||||
// ResponseMaxBytes is the largest response body
|
||||
// (SWWAF_RESPONSE_MAX_BYTES).
|
||||
ResponseMaxBytes int64
|
||||
|
||||
// settings are the values read, as given or by default, for the
|
||||
// log line at start.
|
||||
settings []slog.Attr
|
||||
}
|
||||
|
||||
// off is the value that switches a timeout or a size limit off.
|
||||
const off = "off"
|
||||
|
||||
const (
|
||||
day = 24 * time.Hour
|
||||
kibibyte = 1 << 10
|
||||
mebibyte = 1 << 20
|
||||
gibibyte = 1 << 30
|
||||
)
|
||||
|
||||
var (
|
||||
errNotDuration = errors.New(
|
||||
"is not a duration such as 90s, 15m or 7d, or off")
|
||||
errNotSize = errors.New(
|
||||
"is not a size such as 512K, 100M or 5G, or off")
|
||||
errNotPositive = errors.New("must be more than zero, or off")
|
||||
errEmptyItem = errors.New("has an empty item in its list")
|
||||
errNotNetblock = errors.New(
|
||||
"is not a netblock such as 10.0.0.0/8, or an address")
|
||||
errNotListenAddr = errors.New(
|
||||
"is not an address to listen on, such as :8080")
|
||||
errNotUpstreamURL = errors.New(
|
||||
"is not a URL with only a scheme, a host and an optional port, " +
|
||||
"such as http://127.0.0.1:8081")
|
||||
)
|
||||
|
||||
// FromEnvironment reads the settings with lookupEnv, normally
|
||||
// os.LookupEnv. A setting that is not set takes its default. A setting
|
||||
// that is set but invalid is an error that names it.
|
||||
func FromEnvironment(lookupEnv func(string) (string, bool)) (*Config, error) {
|
||||
env := &environment{lookupEnv: lookupEnv}
|
||||
cfg := &Config{
|
||||
ListenAddr: env.address("SWWAF_LISTEN_ADDR", ":8080"),
|
||||
UpstreamURL: env.appURL("SWWAF_UPSTREAM_URL", "http://127.0.0.1:8081"),
|
||||
TrustedProxies: env.netblocks("SWWAF_TRUSTED_PROXIES", privateRanges),
|
||||
ClientRequestTimeout: env.duration("SWWAF_CLIENT_REQUEST_TIMEOUT", "60s"),
|
||||
ClientResponseTimeout: env.duration("SWWAF_CLIENT_RESPONSE_TIMEOUT", "30m"),
|
||||
UpstreamRequestTimeout: env.duration("SWWAF_UPSTREAM_REQUEST_TIMEOUT", "60s"),
|
||||
UpstreamResponseTimeout: env.duration("SWWAF_UPSTREAM_RESPONSE_TIMEOUT", "30m"),
|
||||
RequestMaxBytes: env.size("SWWAF_REQUEST_MAX_BYTES", "100M"),
|
||||
ResponseMaxBytes: env.size("SWWAF_RESPONSE_MAX_BYTES", "5G"),
|
||||
}
|
||||
|
||||
if env.err != nil {
|
||||
return nil, env.err
|
||||
}
|
||||
|
||||
cfg.settings = env.settings
|
||||
|
||||
return cfg, nil
|
||||
}
|
||||
|
||||
// privateRanges are the private address ranges, the default trusted
|
||||
// proxies.
|
||||
const privateRanges = "10.0.0.0/8,172.16.0.0/12,192.168.0.0/16"
|
||||
|
||||
// LogValue makes a Config log as each setting's name with the value it was
|
||||
// given, or its default.
|
||||
func (c *Config) LogValue() slog.Value {
|
||||
return slog.GroupValue(c.settings...)
|
||||
}
|
||||
|
||||
// environment is where FromEnvironment reads the settings: it notes each
|
||||
// value for the log, and keeps the first error.
|
||||
type environment struct {
|
||||
lookupEnv func(string) (string, bool)
|
||||
settings []slog.Attr
|
||||
err error
|
||||
}
|
||||
|
||||
// value returns a setting's value, or its default when it is not set,
|
||||
// and notes it for the log.
|
||||
func (e *environment) value(name, defaultValue string) string {
|
||||
value, ok := e.lookupEnv(name)
|
||||
if !ok {
|
||||
value = defaultValue
|
||||
}
|
||||
|
||||
e.settings = append(e.settings, slog.String(name, value))
|
||||
|
||||
return value
|
||||
}
|
||||
|
||||
// check keeps the first error, naming the setting it is about.
|
||||
func (e *environment) check(name string, err error) {
|
||||
if err != nil && e.err == nil {
|
||||
e.err = fmt.Errorf("%s: %w", name, err)
|
||||
}
|
||||
}
|
||||
|
||||
// address reads a setting that is an address to listen on.
|
||||
func (e *environment) address(name, defaultValue string) string {
|
||||
address, err := parseListenAddr(e.value(name, defaultValue))
|
||||
e.check(name, err)
|
||||
|
||||
return address
|
||||
}
|
||||
|
||||
// appURL reads a setting that is the app's URL.
|
||||
func (e *environment) appURL(name, defaultValue string) *url.URL {
|
||||
upstream, err := parseUpstreamURL(e.value(name, defaultValue))
|
||||
e.check(name, err)
|
||||
|
||||
return upstream
|
||||
}
|
||||
|
||||
// netblocks reads a setting that is a list of netblocks.
|
||||
func (e *environment) netblocks(name, defaultValue string) []netip.Prefix {
|
||||
netblocks, err := parseNetblocks(e.value(name, defaultValue))
|
||||
e.check(name, err)
|
||||
|
||||
return netblocks
|
||||
}
|
||||
|
||||
// duration reads a setting that is a duration.
|
||||
func (e *environment) duration(name, defaultValue string) time.Duration {
|
||||
duration, err := parseDuration(e.value(name, defaultValue))
|
||||
e.check(name, err)
|
||||
|
||||
return duration
|
||||
}
|
||||
|
||||
// size reads a setting that is a number of bytes.
|
||||
func (e *environment) size(name, defaultValue string) int64 {
|
||||
size, err := parseSize(e.value(name, defaultValue))
|
||||
e.check(name, err)
|
||||
|
||||
return size
|
||||
}
|
||||
|
||||
// parseDuration reads a duration in Go's syntax, such as 90s or 15m, a
|
||||
// whole number of days such as 7d, or off.
|
||||
func parseDuration(value string) (time.Duration, error) {
|
||||
if value == off {
|
||||
return 0, nil
|
||||
}
|
||||
|
||||
duration, err := durationOrDays(value)
|
||||
if err != nil {
|
||||
return 0, fmt.Errorf("%q %w", value, errNotDuration)
|
||||
}
|
||||
|
||||
if duration <= 0 {
|
||||
return 0, fmt.Errorf("%q %w", value, errNotPositive)
|
||||
}
|
||||
|
||||
return duration, nil
|
||||
}
|
||||
|
||||
// durationOrDays reads Go's duration syntax, or a whole number of days.
|
||||
func durationOrDays(value string) (time.Duration, error) {
|
||||
days, isDays := strings.CutSuffix(value, "d")
|
||||
if !isDays {
|
||||
return time.ParseDuration(value)
|
||||
}
|
||||
|
||||
n, err := strconv.ParseInt(days, 10, 64)
|
||||
if err != nil || n < 0 || n > math.MaxInt64/int64(day) {
|
||||
return 0, errNotDuration
|
||||
}
|
||||
|
||||
return time.Duration(n) * day, nil
|
||||
}
|
||||
|
||||
// parseSize reads a number of bytes with an optional K, M or G suffix, in
|
||||
// powers of 1024 (1K is 1024 bytes), or off.
|
||||
func parseSize(value string) (int64, error) {
|
||||
if value == off {
|
||||
return 0, nil
|
||||
}
|
||||
|
||||
number, unit := splitUnit(value)
|
||||
|
||||
n, err := strconv.ParseInt(number, 10, 64)
|
||||
if err != nil || n > math.MaxInt64/unit {
|
||||
return 0, fmt.Errorf("%q %w", value, errNotSize)
|
||||
}
|
||||
|
||||
if n <= 0 {
|
||||
return 0, fmt.Errorf("%q %w", value, errNotPositive)
|
||||
}
|
||||
|
||||
return n * unit, nil
|
||||
}
|
||||
|
||||
// splitUnit splits a size into its number and the bytes its suffix
|
||||
// stands for.
|
||||
func splitUnit(value string) (string, int64) {
|
||||
switch {
|
||||
case strings.HasSuffix(value, "K"):
|
||||
return strings.TrimSuffix(value, "K"), kibibyte
|
||||
case strings.HasSuffix(value, "M"):
|
||||
return strings.TrimSuffix(value, "M"), mebibyte
|
||||
case strings.HasSuffix(value, "G"):
|
||||
return strings.TrimSuffix(value, "G"), gibibyte
|
||||
default:
|
||||
return value, 1
|
||||
}
|
||||
}
|
||||
|
||||
// parseList splits a comma-separated list and trims the spaces around
|
||||
// each item. An empty value is an empty list.
|
||||
func parseList(value string) ([]string, error) {
|
||||
if strings.TrimSpace(value) == "" {
|
||||
return []string{}, nil
|
||||
}
|
||||
|
||||
items := strings.Split(value, ",")
|
||||
for i, item := range items {
|
||||
items[i] = strings.TrimSpace(item)
|
||||
if items[i] == "" {
|
||||
return nil, fmt.Errorf("%q %w", value, errEmptyItem)
|
||||
}
|
||||
}
|
||||
|
||||
return items, nil
|
||||
}
|
||||
|
||||
// parseNetblocks reads a comma-separated list of netblocks.
|
||||
func parseNetblocks(value string) ([]netip.Prefix, error) {
|
||||
items, err := parseList(value)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
netblocks := make([]netip.Prefix, 0, len(items))
|
||||
|
||||
for _, item := range items {
|
||||
netblock, err := parseNetblock(item)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
netblocks = append(netblocks, netblock)
|
||||
}
|
||||
|
||||
return netblocks, nil
|
||||
}
|
||||
|
||||
// parseNetblock reads a netblock in CIDR form, such as 10.0.0.0/8. A bare
|
||||
// address is a netblock of that address alone, a /32 or a /128.
|
||||
func parseNetblock(value string) (netip.Prefix, error) {
|
||||
if strings.Contains(value, "/") {
|
||||
netblock, err := netip.ParsePrefix(value)
|
||||
if err != nil {
|
||||
return netip.Prefix{}, fmt.Errorf("%q %w", value, errNotNetblock)
|
||||
}
|
||||
|
||||
return netblock.Masked(), nil
|
||||
}
|
||||
|
||||
addr, err := netip.ParseAddr(value)
|
||||
if err != nil || addr.Zone() != "" {
|
||||
return netip.Prefix{}, fmt.Errorf("%q %w", value, errNotNetblock)
|
||||
}
|
||||
|
||||
return netip.PrefixFrom(addr, addr.BitLen()), nil
|
||||
}
|
||||
|
||||
// parseListenAddr checks an address to listen on: an optional host and a
|
||||
// port number.
|
||||
func parseListenAddr(value string) (string, error) {
|
||||
_, port, err := net.SplitHostPort(value)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("%q %w", value, errNotListenAddr)
|
||||
}
|
||||
|
||||
_, err = strconv.ParseUint(port, 10, 16)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("%q %w", value, errNotListenAddr)
|
||||
}
|
||||
|
||||
return value, nil
|
||||
}
|
||||
|
||||
// parseUpstreamURL reads the app's URL: http or https, a host and an
|
||||
// optional port from 1 to 65535, and nothing else, since the request's
|
||||
// own path and query go to the app unchanged.
|
||||
func parseUpstreamURL(value string) (*url.URL, error) {
|
||||
upstream, err := url.Parse(value)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("%q %w", value, errNotUpstreamURL)
|
||||
}
|
||||
|
||||
onlySchemeAndHost := (upstream.Scheme == "http" || upstream.Scheme == "https") &&
|
||||
upstream.Hostname() != "" && upstream.User == nil && upstream.Opaque == "" &&
|
||||
(upstream.Path == "" || upstream.Path == "/") &&
|
||||
upstream.RawQuery == "" && upstream.Fragment == ""
|
||||
if !onlySchemeAndHost {
|
||||
return nil, fmt.Errorf("%q %w", value, errNotUpstreamURL)
|
||||
}
|
||||
|
||||
if upstream.Port() != "" {
|
||||
port, err := strconv.ParseUint(upstream.Port(), 10, 16)
|
||||
if err != nil || port == 0 {
|
||||
return nil, fmt.Errorf("%q %w", value, errNotUpstreamURL)
|
||||
}
|
||||
}
|
||||
|
||||
return upstream, nil
|
||||
}
|
||||
@@ -0,0 +1,244 @@
|
||||
package config_test
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"log/slog"
|
||||
"maps"
|
||||
"net/netip"
|
||||
"slices"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"sneak.berlin/go/smallwebwaf/internal/config"
|
||||
)
|
||||
|
||||
// The settings, by name.
|
||||
const (
|
||||
listenAddr = "SWWAF_LISTEN_ADDR"
|
||||
upstreamURL = "SWWAF_UPSTREAM_URL"
|
||||
trustedProxies = "SWWAF_TRUSTED_PROXIES"
|
||||
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"
|
||||
)
|
||||
|
||||
// off switches a timeout or a size limit off.
|
||||
const off = "off"
|
||||
|
||||
// environment is a set of environment variables, for FromEnvironment.
|
||||
type environment map[string]string
|
||||
|
||||
// lookupEnv reads one of the variables, as os.LookupEnv does.
|
||||
func (e environment) lookupEnv(name string) (string, bool) {
|
||||
value, ok := e[name]
|
||||
|
||||
return value, ok
|
||||
}
|
||||
|
||||
// fromEnvironment reads the settings from env, which must be valid.
|
||||
func fromEnvironment(t *testing.T, env environment) *config.Config {
|
||||
t.Helper()
|
||||
|
||||
cfg, err := config.FromEnvironment(env.lookupEnv)
|
||||
if err != nil {
|
||||
t.Fatalf("settings %v: %v", env, err)
|
||||
}
|
||||
|
||||
return cfg
|
||||
}
|
||||
|
||||
func TestDefaults(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
cfg := fromEnvironment(t, environment{})
|
||||
|
||||
wantSettings(t, cfg, config.Config{
|
||||
ListenAddr: ":8080",
|
||||
ClientRequestTimeout: time.Minute,
|
||||
ClientResponseTimeout: 30 * time.Minute,
|
||||
UpstreamRequestTimeout: time.Minute,
|
||||
UpstreamResponseTimeout: 30 * time.Minute,
|
||||
RequestMaxBytes: 100 << 20,
|
||||
ResponseMaxBytes: 5 << 30,
|
||||
})
|
||||
|
||||
if cfg.UpstreamURL.String() != "http://127.0.0.1:8081" {
|
||||
t.Errorf("%s is %s", upstreamURL, cfg.UpstreamURL)
|
||||
}
|
||||
|
||||
wantNetblocks(t, cfg.TrustedProxies,
|
||||
"10.0.0.0/8", "172.16.0.0/12", "192.168.0.0/16")
|
||||
}
|
||||
|
||||
func TestValuesAsSet(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
cfg := fromEnvironment(t, environment{
|
||||
listenAddr: "127.0.0.1:9000",
|
||||
upstreamURL: "https://app.internal:8443/",
|
||||
trustedProxies: " 192.0.2.1, 10.1.2.3/8 ,2001:db8::/32",
|
||||
clientRequestTimeout: "90s",
|
||||
clientResponseTimeout: "7d",
|
||||
upstreamRequestTimeout: "1h30m",
|
||||
upstreamResponseTimeout: off,
|
||||
requestMaxBytes: "512K",
|
||||
responseMaxBytes: "1234",
|
||||
})
|
||||
|
||||
wantSettings(t, cfg, config.Config{
|
||||
ListenAddr: "127.0.0.1:9000",
|
||||
ClientRequestTimeout: 90 * time.Second,
|
||||
ClientResponseTimeout: 7 * 24 * time.Hour,
|
||||
UpstreamRequestTimeout: 90 * time.Minute,
|
||||
UpstreamResponseTimeout: 0,
|
||||
RequestMaxBytes: 512 << 10,
|
||||
ResponseMaxBytes: 1234,
|
||||
})
|
||||
|
||||
if cfg.UpstreamURL.String() != "https://app.internal:8443/" {
|
||||
t.Errorf("%s is %s", upstreamURL, cfg.UpstreamURL)
|
||||
}
|
||||
|
||||
wantNetblocks(t, cfg.TrustedProxies, "192.0.2.1/32", "10.0.0.0/8", "2001:db8::/32")
|
||||
}
|
||||
|
||||
func TestSizesAndOff(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
cfg := fromEnvironment(t, environment{
|
||||
requestMaxBytes: "3G",
|
||||
responseMaxBytes: off,
|
||||
clientRequestTimeout: off,
|
||||
})
|
||||
|
||||
if cfg.RequestMaxBytes != 3<<30 || cfg.ResponseMaxBytes != 0 ||
|
||||
cfg.ClientRequestTimeout != 0 {
|
||||
t.Errorf("3G, off and off read as %d, %d and %s",
|
||||
cfg.RequestMaxBytes, cfg.ResponseMaxBytes, cfg.ClientRequestTimeout)
|
||||
}
|
||||
}
|
||||
|
||||
func TestTrustedProxiesSetButEmptyTrustNothing(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
cfg := fromEnvironment(t, environment{trustedProxies: ""})
|
||||
if len(cfg.TrustedProxies) != 0 {
|
||||
t.Errorf("trusted proxies %v, want none", cfg.TrustedProxies)
|
||||
}
|
||||
}
|
||||
|
||||
func TestInvalidValueStopsTheStart(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
for _, tc := range []struct{ name, value string }{
|
||||
{listenAddr, "8080"},
|
||||
{listenAddr, ":http"},
|
||||
{listenAddr, ":65536"},
|
||||
{upstreamURL, "127.0.0.1:8081"},
|
||||
{upstreamURL, "ftp://127.0.0.1:8081"},
|
||||
{upstreamURL, "http://"},
|
||||
{upstreamURL, "http://:8081"},
|
||||
{upstreamURL, "http://127.0.0.1:0"},
|
||||
{upstreamURL, "http://127.0.0.1:99999"},
|
||||
{upstreamURL, "http://127.0.0.1:8081/app"},
|
||||
{upstreamURL, "http://127.0.0.1:8081/?a=1"},
|
||||
{upstreamURL, "http://user:secret@127.0.0.1:8081"},
|
||||
{trustedProxies, "10.0.0.0/33"},
|
||||
{trustedProxies, "traefik"},
|
||||
{trustedProxies, "10.0.0.0/8,,192.168.0.0/16"},
|
||||
{trustedProxies, "fe80::1%eth0"},
|
||||
{clientRequestTimeout, "60"},
|
||||
{clientRequestTimeout, ""},
|
||||
{clientResponseTimeout, "1y"},
|
||||
{upstreamRequestTimeout, "-1s"},
|
||||
{upstreamResponseTimeout, "0s"},
|
||||
{upstreamResponseTimeout, "1.5d"},
|
||||
{requestMaxBytes, "100MB"},
|
||||
{requestMaxBytes, "100m"},
|
||||
{requestMaxBytes, "1.5M"},
|
||||
{responseMaxBytes, "0"},
|
||||
{responseMaxBytes, "-5"},
|
||||
{responseMaxBytes, "99999999999G"},
|
||||
} {
|
||||
t.Run(tc.name+"="+tc.value, func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
_, err := config.FromEnvironment(environment{tc.name: tc.value}.lookupEnv)
|
||||
if err == nil {
|
||||
t.Fatalf("%s=%q was accepted", tc.name, tc.value)
|
||||
}
|
||||
|
||||
if !strings.HasPrefix(err.Error(), tc.name+": ") {
|
||||
t.Errorf("error %q does not name %s", err, tc.name)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestLogsEachSettingWithItsValue(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
cfg := fromEnvironment(t, environment{clientRequestTimeout: "45s"})
|
||||
|
||||
var out bytes.Buffer
|
||||
|
||||
slog.New(slog.NewJSONHandler(&out, nil)).Info("starting", "settings", cfg)
|
||||
|
||||
var line struct {
|
||||
Settings map[string]string `json:"settings"`
|
||||
}
|
||||
|
||||
err := json.Unmarshal(out.Bytes(), &line)
|
||||
if err != nil {
|
||||
t.Fatalf("decode %s: %v", out.Bytes(), err)
|
||||
}
|
||||
|
||||
want := map[string]string{
|
||||
listenAddr: ":8080",
|
||||
upstreamURL: "http://127.0.0.1:8081",
|
||||
trustedProxies: "10.0.0.0/8,172.16.0.0/12,192.168.0.0/16",
|
||||
clientRequestTimeout: "45s",
|
||||
clientResponseTimeout: "30m",
|
||||
upstreamRequestTimeout: "60s",
|
||||
upstreamResponseTimeout: "30m",
|
||||
requestMaxBytes: "100M",
|
||||
responseMaxBytes: "5G",
|
||||
}
|
||||
if !maps.Equal(line.Settings, want) {
|
||||
t.Errorf("logged settings\n%v\nwant\n%v", line.Settings, want)
|
||||
}
|
||||
}
|
||||
|
||||
// wantSettings checks the settings that are plain values.
|
||||
func wantSettings(t *testing.T, got *config.Config, want config.Config) {
|
||||
t.Helper()
|
||||
|
||||
if got.ListenAddr != want.ListenAddr ||
|
||||
got.ClientRequestTimeout != want.ClientRequestTimeout ||
|
||||
got.ClientResponseTimeout != want.ClientResponseTimeout ||
|
||||
got.UpstreamRequestTimeout != want.UpstreamRequestTimeout ||
|
||||
got.UpstreamResponseTimeout != want.UpstreamResponseTimeout ||
|
||||
got.RequestMaxBytes != want.RequestMaxBytes ||
|
||||
got.ResponseMaxBytes != want.ResponseMaxBytes {
|
||||
t.Errorf("settings\n%+v\nwant\n%+v", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
// wantNetblocks checks a list of netblocks.
|
||||
func wantNetblocks(t *testing.T, got []netip.Prefix, want ...string) {
|
||||
t.Helper()
|
||||
|
||||
gotText := make([]string, 0, len(got))
|
||||
for _, netblock := range got {
|
||||
gotText = append(gotText, netblock.String())
|
||||
}
|
||||
|
||||
if !slices.Equal(gotText, want) {
|
||||
t.Errorf("netblocks %v, want %v", gotText, want)
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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:])
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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())
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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()
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -0,0 +1,102 @@
|
||||
// Package requestlog writes the lines smallwebwaf prints on stdout: one
|
||||
// JSON object per request, marked "type":"request", and the process's own
|
||||
// messages as JSON lines marked "type":"process".
|
||||
package requestlog
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"log/slog"
|
||||
"time"
|
||||
)
|
||||
|
||||
// The action a request line names: what smallwebwaf did with the
|
||||
// request.
|
||||
const (
|
||||
// ActionForward is a request passed to the app.
|
||||
ActionForward = "forward"
|
||||
// ActionTooLarge is a request or response over its size limit.
|
||||
ActionTooLarge = "too_large"
|
||||
// ActionTimedOut is a request or response that ran out of time.
|
||||
ActionTimedOut = "timed_out"
|
||||
// ActionUpstreamError is a request the app could not be reached
|
||||
// for, or whose answer could not be passed on.
|
||||
ActionUpstreamError = "upstream_error"
|
||||
)
|
||||
|
||||
// timeLayout is RFC 3339 with milliseconds.
|
||||
const timeLayout = "2006-01-02T15:04:05.000Z07:00"
|
||||
|
||||
// Line is one request's line in the request log. The field names are
|
||||
// those of the "Request log" section of SPEC.md.
|
||||
//
|
||||
//nolint:tagliatelle // SPEC.md's request log names its fields in snake_case
|
||||
type Line struct {
|
||||
Type string `json:"type"`
|
||||
Time string `json:"time"`
|
||||
ClientIP string `json:"client_ip"`
|
||||
PeerIP string `json:"peer_ip"`
|
||||
Method string `json:"method"`
|
||||
Host string `json:"host"`
|
||||
Path string `json:"path"`
|
||||
Query string `json:"query"`
|
||||
Protocol string `json:"protocol"`
|
||||
Status int `json:"status"`
|
||||
UpstreamStatus int `json:"upstream_status,omitempty"`
|
||||
RequestBytes int64 `json:"request_bytes"`
|
||||
ResponseBytes int64 `json:"response_bytes"`
|
||||
Referer string `json:"referer"`
|
||||
UserAgent string `json:"user_agent"`
|
||||
Action string `json:"action"`
|
||||
// Aborted is true when the client went away early.
|
||||
Aborted bool `json:"aborted,omitempty"`
|
||||
// DurationTotal and DurationUpstreamTotal are in milliseconds.
|
||||
DurationTotal float64 `json:"duration_total"`
|
||||
DurationUpstreamTotal float64 `json:"duration_upstream_total,omitempty"`
|
||||
}
|
||||
|
||||
// Write writes line to w as one JSON line marked "type":"request".
|
||||
func Write(w io.Writer, line *Line) error {
|
||||
line.Type = "request"
|
||||
|
||||
encoded, err := json.Marshal(line)
|
||||
if err != nil {
|
||||
return fmt.Errorf("encode the request log line: %w", err)
|
||||
}
|
||||
|
||||
_, err = w.Write(append(encoded, '\n'))
|
||||
if err != nil {
|
||||
return fmt.Errorf("write the request log line: %w", err)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// FormatTime formats t for a line's time field: RFC 3339 in UTC, with
|
||||
// milliseconds.
|
||||
func FormatTime(t time.Time) string {
|
||||
return t.UTC().Format(timeLayout)
|
||||
}
|
||||
|
||||
// Milliseconds is d in milliseconds, to the microsecond.
|
||||
func Milliseconds(d time.Duration) float64 {
|
||||
return float64(d.Microseconds()) / float64(time.Millisecond/time.Microsecond)
|
||||
}
|
||||
|
||||
// NewProcessLogger returns the logger for the process's own messages:
|
||||
// JSON lines on w, marked "type":"process", with the time in the same form
|
||||
// as a request line's.
|
||||
func NewProcessLogger(w io.Writer) *slog.Logger {
|
||||
handler := slog.NewJSONHandler(w, &slog.HandlerOptions{
|
||||
ReplaceAttr: func(groups []string, attr slog.Attr) slog.Attr {
|
||||
if attr.Key == slog.TimeKey && len(groups) == 0 {
|
||||
return slog.String(slog.TimeKey, FormatTime(attr.Value.Time()))
|
||||
}
|
||||
|
||||
return attr
|
||||
},
|
||||
})
|
||||
|
||||
return slog.New(handler).With("type", "process")
|
||||
}
|
||||
@@ -0,0 +1,88 @@
|
||||
package requestlog_test
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"sneak.berlin/go/smallwebwaf/internal/requestlog"
|
||||
)
|
||||
|
||||
func TestWriteWritesOneJSONLineMarkedRequest(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
var out bytes.Buffer
|
||||
|
||||
err := requestlog.Write(&out, &requestlog.Line{
|
||||
Time: requestlog.FormatTime(time.Date(2026, 10, 3, 12, 0, 0, 0, time.UTC)),
|
||||
ClientIP: "203.0.113.9",
|
||||
Status: 200,
|
||||
Action: requestlog.ActionForward,
|
||||
DurationTotal: requestlog.Milliseconds(1500 * time.Microsecond),
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("write: %v", err)
|
||||
}
|
||||
|
||||
text := out.String()
|
||||
if strings.Count(text, "\n") != 1 || !strings.HasSuffix(text, "\n") {
|
||||
t.Fatalf("wrote %q, want one line", text)
|
||||
}
|
||||
|
||||
var fields map[string]any
|
||||
|
||||
err = json.Unmarshal(out.Bytes(), &fields)
|
||||
if err != nil {
|
||||
t.Fatalf("decode %q: %v", text, err)
|
||||
}
|
||||
|
||||
want := map[string]any{
|
||||
"type": "request", "time": "2026-10-03T12:00:00.000Z",
|
||||
"client_ip": "203.0.113.9", "status": 200.0, "action": "forward",
|
||||
"duration_total": 1.5,
|
||||
}
|
||||
for name, value := range want {
|
||||
if fields[name] != value {
|
||||
t.Errorf("%s is %v, want %v", name, fields[name], value)
|
||||
}
|
||||
}
|
||||
|
||||
unset := []string{"upstream_status", "aborted", "duration_upstream_total"}
|
||||
for _, name := range unset {
|
||||
_, present := fields[name]
|
||||
if present {
|
||||
t.Errorf("%s is there with no value to give", name)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestProcessLinesAreMarkedProcess(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
var out bytes.Buffer
|
||||
|
||||
requestlog.NewProcessLogger(&out).Info("starting", "version", "v1")
|
||||
|
||||
var fields map[string]any
|
||||
|
||||
err := json.Unmarshal(out.Bytes(), &fields)
|
||||
if err != nil {
|
||||
t.Fatalf("decode %q: %v", out.String(), err)
|
||||
}
|
||||
|
||||
if fields["type"] != "process" || fields["msg"] != "starting" ||
|
||||
fields["level"] != "INFO" || fields["version"] != "v1" {
|
||||
t.Errorf("process line %v", fields)
|
||||
}
|
||||
|
||||
timeText, _ := fields["time"].(string)
|
||||
|
||||
logged, err := time.Parse(time.RFC3339, timeText)
|
||||
if err != nil || !strings.HasSuffix(timeText, "Z") ||
|
||||
len(timeText) != len("2006-01-02T15:04:05.000Z") ||
|
||||
time.Since(logged) > time.Minute {
|
||||
t.Errorf("process line time %q, want now in UTC with milliseconds", timeText)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,130 @@
|
||||
// Package smallwebwaf runs the smallwebwaf process: it reads the settings,
|
||||
// serves requests until it is told to stop, and then stops in an orderly
|
||||
// way.
|
||||
package smallwebwaf
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"io"
|
||||
"log/slog"
|
||||
"net"
|
||||
"net/http"
|
||||
"os"
|
||||
"os/signal"
|
||||
"syscall"
|
||||
"time"
|
||||
|
||||
"sneak.berlin/go/smallwebwaf/internal/config"
|
||||
"sneak.berlin/go/smallwebwaf/internal/proxy"
|
||||
"sneak.berlin/go/smallwebwaf/internal/requestlog"
|
||||
)
|
||||
|
||||
// shutdownTimeout is how long requests in progress may take to finish
|
||||
// once smallwebwaf is told to stop, before their connections are closed.
|
||||
// runit and docker wait a little longer before they kill the process.
|
||||
const shutdownTimeout = 5 * time.Second
|
||||
|
||||
// Params are what Run needs from the process.
|
||||
type Params struct {
|
||||
// Version is the version of the binary, set when it is built.
|
||||
Version string
|
||||
// LookupEnv reads an environment variable, normally os.LookupEnv.
|
||||
LookupEnv func(string) (string, bool)
|
||||
// Stdout receives the request log and the process's own messages.
|
||||
Stdout io.Writer
|
||||
}
|
||||
|
||||
// Main runs smallwebwaf until SIGTERM or SIGINT, and returns the
|
||||
// process's exit status.
|
||||
func Main(version string) int {
|
||||
ctx, stop := signal.NotifyContext(context.Background(),
|
||||
syscall.SIGTERM, os.Interrupt)
|
||||
defer stop()
|
||||
|
||||
return Run(ctx, Params{
|
||||
Version: version,
|
||||
LookupEnv: os.LookupEnv,
|
||||
Stdout: os.Stdout,
|
||||
})
|
||||
}
|
||||
|
||||
// Run reads the settings, then serves requests until ctx is done. It
|
||||
// returns the process's exit status, 1 when smallwebwaf cannot start.
|
||||
func Run(ctx context.Context, params Params) int {
|
||||
processLog := requestlog.NewProcessLogger(params.Stdout)
|
||||
|
||||
cfg, err := config.FromEnvironment(params.LookupEnv)
|
||||
if err != nil {
|
||||
processLog.Error("invalid setting", "error", err.Error())
|
||||
|
||||
return 1
|
||||
}
|
||||
|
||||
listener, err := (&net.ListenConfig{}).Listen(ctx, "tcp", cfg.ListenAddr)
|
||||
if err != nil {
|
||||
processLog.Error("cannot listen on SWWAF_LISTEN_ADDR",
|
||||
"error", err.Error())
|
||||
|
||||
return 1
|
||||
}
|
||||
|
||||
server := proxy.New(proxy.Params{
|
||||
Config: cfg,
|
||||
RequestLog: params.Stdout,
|
||||
ProcessLog: processLog,
|
||||
})
|
||||
|
||||
processLog.Info("starting",
|
||||
"version", params.Version,
|
||||
"address", listener.Addr().String(),
|
||||
"settings", cfg)
|
||||
|
||||
return serve(ctx, server, listener, processLog)
|
||||
}
|
||||
|
||||
// serve serves requests on listener until ctx is done, then gives the
|
||||
// requests in progress shutdownTimeout to finish.
|
||||
func serve(
|
||||
ctx context.Context, server *http.Server, listener net.Listener,
|
||||
processLog *slog.Logger,
|
||||
) int {
|
||||
served := make(chan error, 1)
|
||||
|
||||
go func() {
|
||||
served <- server.Serve(listener)
|
||||
}()
|
||||
|
||||
select {
|
||||
case err := <-served:
|
||||
processLog.Error("serving failed", "error", err.Error())
|
||||
|
||||
return 1
|
||||
case <-ctx.Done():
|
||||
}
|
||||
|
||||
processLog.Info("stopping")
|
||||
|
||||
shutdownCtx, cancel := context.WithTimeout(context.WithoutCancel(ctx),
|
||||
shutdownTimeout)
|
||||
defer cancel()
|
||||
|
||||
err := server.Shutdown(shutdownCtx)
|
||||
if err != nil {
|
||||
processLog.Warn("requests still in progress were cut off",
|
||||
"error", err.Error())
|
||||
|
||||
_ = server.Close()
|
||||
}
|
||||
|
||||
err = <-served
|
||||
if !errors.Is(err, http.ErrServerClosed) {
|
||||
processLog.Error("serving failed", "error", err.Error())
|
||||
|
||||
return 1
|
||||
}
|
||||
|
||||
processLog.Info("stopped")
|
||||
|
||||
return 0
|
||||
}
|
||||
@@ -0,0 +1,225 @@
|
||||
package smallwebwaf_test
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"io"
|
||||
"net"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"sneak.berlin/go/smallwebwaf/internal/smallwebwaf"
|
||||
)
|
||||
|
||||
const (
|
||||
// waitLimit bounds how long a test waits for what should happen.
|
||||
waitLimit = 10 * time.Second
|
||||
// pollInterval is how often a test looks for a line.
|
||||
pollInterval = 10 * time.Millisecond
|
||||
// testVersion is the version the tests give smallwebwaf.
|
||||
testVersion = "test"
|
||||
// localhost is where the tests listen.
|
||||
localhost = "127.0.0.1"
|
||||
listenAddr = "SWWAF_LISTEN_ADDR"
|
||||
)
|
||||
|
||||
// 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)
|
||||
}
|
||||
|
||||
// line returns the first line whose field key is value, waiting for it.
|
||||
func (o *output) line(t *testing.T, key, value string) map[string]any {
|
||||
t.Helper()
|
||||
|
||||
deadline := time.Now().Add(waitLimit)
|
||||
for time.Now().Before(deadline) {
|
||||
o.mu.Lock()
|
||||
text := o.buf.String()
|
||||
o.mu.Unlock()
|
||||
|
||||
for line := range strings.Lines(text) {
|
||||
var fields map[string]any
|
||||
|
||||
err := json.Unmarshal([]byte(line), &fields)
|
||||
if err != nil {
|
||||
t.Fatalf("output line %q is not JSON: %v", line, err)
|
||||
}
|
||||
|
||||
if fields[key] == value {
|
||||
return fields
|
||||
}
|
||||
}
|
||||
|
||||
time.Sleep(pollInterval)
|
||||
}
|
||||
|
||||
t.Fatalf("no line with %s %q in the output:\n%s", key, value, o.buf.String())
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// run runs smallwebwaf with the settings in env until ctx is done, and
|
||||
// returns its exit status.
|
||||
func run(ctx context.Context, env map[string]string, out *output) int {
|
||||
return smallwebwaf.Run(ctx, smallwebwaf.Params{
|
||||
Version: testVersion,
|
||||
LookupEnv: func(name string) (string, bool) {
|
||||
value, ok := env[name]
|
||||
|
||||
return value, ok
|
||||
},
|
||||
Stdout: out,
|
||||
})
|
||||
}
|
||||
|
||||
func TestInvalidSettingStopsTheStart(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
out := &output{}
|
||||
|
||||
status := run(t.Context(), map[string]string{"SWWAF_REQUEST_MAX_BYTES": "lots"}, out)
|
||||
if status != 1 {
|
||||
t.Errorf("exit status %d, want 1", status)
|
||||
}
|
||||
|
||||
line := out.line(t, "msg", "invalid setting")
|
||||
message, _ := line["error"].(string)
|
||||
|
||||
if line["type"] != "process" || line["level"] != "ERROR" ||
|
||||
!strings.HasPrefix(message, "SWWAF_REQUEST_MAX_BYTES: ") {
|
||||
t.Errorf("start refused with %v", line)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAddressInUseStopsTheStart(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
taken, err := (&net.ListenConfig{}).Listen(t.Context(), "tcp", localhost+":0")
|
||||
if err != nil {
|
||||
t.Fatalf("listen: %v", err)
|
||||
}
|
||||
|
||||
defer func() {
|
||||
_ = taken.Close()
|
||||
}()
|
||||
|
||||
out := &output{}
|
||||
|
||||
status := run(t.Context(), map[string]string{listenAddr: taken.Addr().String()}, out)
|
||||
if status != 1 {
|
||||
t.Errorf("exit status %d, want 1", status)
|
||||
}
|
||||
|
||||
out.line(t, "msg", "cannot listen on SWWAF_LISTEN_ADDR")
|
||||
}
|
||||
|
||||
func TestServesUntilToldToStop(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
app := httptest.NewServer(http.HandlerFunc(
|
||||
func(w http.ResponseWriter, _ *http.Request) {
|
||||
_, _ = io.WriteString(w, "hello from the app")
|
||||
}))
|
||||
defer app.Close()
|
||||
|
||||
ctx, stop := context.WithCancel(t.Context())
|
||||
out := &output{}
|
||||
exited := make(chan int, 1)
|
||||
|
||||
go func() {
|
||||
exited <- run(ctx, map[string]string{
|
||||
listenAddr: localhost + ":0",
|
||||
"SWWAF_UPSTREAM_URL": app.URL,
|
||||
}, out)
|
||||
}()
|
||||
|
||||
starting := out.line(t, "msg", "starting")
|
||||
wantStartingLine(t, starting, app.URL)
|
||||
|
||||
addr, _ := starting["address"].(string)
|
||||
wantGreeting(t, "http://"+addr+"/")
|
||||
out.line(t, "type", "request")
|
||||
|
||||
stop()
|
||||
|
||||
select {
|
||||
case status := <-exited:
|
||||
if status != 0 {
|
||||
t.Errorf("exit status %d, want 0", status)
|
||||
}
|
||||
case <-time.After(waitLimit):
|
||||
t.Fatal("still running after being told to stop")
|
||||
}
|
||||
|
||||
out.line(t, "msg", "stopped")
|
||||
}
|
||||
|
||||
// wantStartingLine checks that the line at start gives the version and
|
||||
// every setting's value.
|
||||
func wantStartingLine(t *testing.T, line map[string]any, appURL string) {
|
||||
t.Helper()
|
||||
|
||||
settings, _ := line["settings"].(map[string]any)
|
||||
want := map[string]any{
|
||||
listenAddr: localhost + ":0",
|
||||
"SWWAF_UPSTREAM_URL": appURL,
|
||||
"SWWAF_TRUSTED_PROXIES": "10.0.0.0/8,172.16.0.0/12,192.168.0.0/16",
|
||||
"SWWAF_CLIENT_REQUEST_TIMEOUT": "60s",
|
||||
"SWWAF_CLIENT_RESPONSE_TIMEOUT": "30m",
|
||||
"SWWAF_UPSTREAM_REQUEST_TIMEOUT": "60s",
|
||||
"SWWAF_UPSTREAM_RESPONSE_TIMEOUT": "30m",
|
||||
"SWWAF_REQUEST_MAX_BYTES": "100M",
|
||||
"SWWAF_RESPONSE_MAX_BYTES": "5G",
|
||||
}
|
||||
|
||||
for name, value := range want {
|
||||
if settings[name] != value {
|
||||
t.Errorf("starting line gives %s=%v, want %v", name, settings[name], value)
|
||||
}
|
||||
}
|
||||
|
||||
if line["version"] != testVersion || line["type"] != "process" {
|
||||
t.Errorf("starting line %v", line)
|
||||
}
|
||||
}
|
||||
|
||||
// wantGreeting checks that a request to url gets the app's answer.
|
||||
func wantGreeting(t *testing.T, url string) {
|
||||
t.Helper()
|
||||
|
||||
req, err := http.NewRequestWithContext(t.Context(), http.MethodGet, url,
|
||||
http.NoBody)
|
||||
if err != nil {
|
||||
t.Fatalf("new request: %v", err)
|
||||
}
|
||||
|
||||
transport := &http.Transport{}
|
||||
defer transport.CloseIdleConnections()
|
||||
|
||||
res, err := (&http.Client{Transport: transport}).Do(req)
|
||||
if err != nil {
|
||||
t.Fatalf("request: %v", err)
|
||||
}
|
||||
|
||||
body, err := io.ReadAll(res.Body)
|
||||
_ = res.Body.Close()
|
||||
|
||||
if err != nil || string(body) != "hello from the app" {
|
||||
t.Errorf("got %q (%v), want the app's answer", body, err)
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user