Files
smallwebwaf/internal/config/config.go
T
sneak 545ce67f44
check / check (push) Successful in 2m9s
Pass-through proxy with timeouts, size limits and a request log (closes #13)
The repo's first code, with the layout the prompts policies ask for:
Makefile, script/ entrypoints, a Dockerfile whose lint and test phases
gate the build, the Gitea workflow, the canonical dotfiles and
REPO_POLICIES.md. smallwebwaf passes each request to the app through
httputil.ReverseProxy within the four timeouts and two size limits,
works out the client's address behind trusted proxies, and writes one
JSON line per request. The tests run against real local servers.
SPEC.md now says what Go's HTTP server does before smallwebwaf sees a
request; make fmt only rewraps EVALUATION.md.

Model: opus-5-5
2026-10-03 14:19:01 +00:00

346 lines
9.9 KiB
Go

// 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 a 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, 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.Host != "" && upstream.User == nil && upstream.Opaque == "" &&
(upstream.Path == "" || upstream.Path == "/") &&
upstream.RawQuery == "" && upstream.Fragment == ""
if !onlySchemeAndHost {
return nil, fmt.Errorf("%q %w", value, errNotUpstreamURL)
}
return upstream, nil
}