Files
smallwebwaf/internal/config/config.go
T
clawbot e001f1c830
check / check (push) Successful in 4m27s
Network lists: always allowed, exempt from rate limits, always refused (closes #19)
Adds SWWAF_ALLOW_NETS, SWWAF_RATE_LIMIT_EXEMPT_NETS and SWWAF_DENY_NETS,
read like SWWAF_TRUSTED_PROXIES and empty by default, and checked against
the client's own address before its country is looked up. A client in
SWWAF_ALLOW_NETS skips the country lists and the rate limits and is not
looked up. One in SWWAF_DENY_NETS is refused with 403, logged as denied
and not counted. One in SWWAF_RATE_LIMIT_EXEMPT_NETS is neither counted
nor refused by the rate limits. SWWAF_EXCLUSIVELY_ALLOWED_COUNTRIES now
refuses a private, loopback or link-local client unless SWWAF_ALLOW_NETS
lists it.

Judgement call: an address in both SWWAF_ALLOW_NETS and SWWAF_DENY_NETS is let through.
Judgement call: the size and time limits still apply to SWWAF_ALLOW_NETS.

Model: opus-5-5
2026-10-06 00:06:10 +00:00

489 lines
14 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"
"slices"
"strconv"
"strings"
"time"
)
// Config is smallwebwaf's settings. A timeout, size or rate limit 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
// AllowNets are the netblocks whose clients skip every check
// (SWWAF_ALLOW_NETS). RateLimitExemptNets are those whose clients the
// rate limits neither count nor refuse (SWWAF_RATE_LIMIT_EXEMPT_NETS).
// DenyNets are those whose clients are always refused
// (SWWAF_DENY_NETS).
AllowNets []netip.Prefix
RateLimitExemptNets []netip.Prefix
DenyNets []netip.Prefix
// RateLimitPerMinute, RateLimitPerHour and RateLimitPerDay are the
// most requests a client may make in a minute, an hour and a day
// (SWWAF_RATE_LIMIT_PER_MINUTE, SWWAF_RATE_LIMIT_PER_HOUR and
// SWWAF_RATE_LIMIT_PER_DAY).
RateLimitPerMinute int64
RateLimitPerHour int64
RateLimitPerDay int64
// DeniedCountries are the countries whose clients are refused
// (SWWAF_DENIED_COUNTRIES). ExclusivelyAllowedCountries, when not
// empty, are the only countries whose clients are let through
// (SWWAF_EXCLUSIVELY_ALLOWED_COUNTRIES). Both hold two-letter codes in
// capitals, as GeoJS gives them.
DeniedCountries []string
ExclusivelyAllowedCountries []string
// 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, a size limit or a rate 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")
errNotCount = errors.New(
"is not a whole number of requests such as 1000, 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")
errNotCountry = errors.New(
"is not a two-letter country code such as de or kp")
errOnBothLists = errors.New("is in SWWAF_DENIED_COUNTRIES too")
)
// 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"),
AllowNets: env.netblocks("SWWAF_ALLOW_NETS", ""),
RateLimitExemptNets: env.netblocks("SWWAF_RATE_LIMIT_EXEMPT_NETS", ""),
DenyNets: env.netblocks("SWWAF_DENY_NETS", ""),
RateLimitPerMinute: env.count("SWWAF_RATE_LIMIT_PER_MINUTE", "1000"),
RateLimitPerHour: env.count("SWWAF_RATE_LIMIT_PER_HOUR", "10000"),
RateLimitPerDay: env.count("SWWAF_RATE_LIMIT_PER_DAY", "50000"),
DeniedCountries: env.countries("SWWAF_DENIED_COUNTRIES", ""),
ExclusivelyAllowedCountries: env.countries(
"SWWAF_EXCLUSIVELY_ALLOWED_COUNTRIES", ""),
}
for _, country := range cfg.ExclusivelyAllowedCountries {
if slices.Contains(cfg.DeniedCountries, country) {
env.check("SWWAF_EXCLUSIVELY_ALLOWED_COUNTRIES",
fmt.Errorf("%q %w", country, errOnBothLists))
}
}
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
}
// count reads a setting that is a number of requests.
func (e *environment) count(name, defaultValue string) int64 {
count, err := parseCount(e.value(name, defaultValue))
e.check(name, err)
return count
}
// countries reads a setting that is a list of countries.
func (e *environment) countries(name, defaultValue string) []string {
countries, err := parseCountries(e.value(name, defaultValue))
e.check(name, err)
return countries
}
// 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
}
}
// parseCount reads a whole number of requests, or off.
func parseCount(value string) (int64, error) {
if value == off {
return 0, nil
}
n, err := strconv.ParseInt(value, 10, 64)
if err != nil {
return 0, fmt.Errorf("%q %w", value, errNotCount)
}
if n <= 0 {
return 0, fmt.Errorf("%q %w", value, errNotPositive)
}
return n, nil
}
// 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
}
// countryCodes are the two-letter codes ISO 3166-1 assigns today, and XK,
// the code in common use for Kosovo. golang.org/x/text/language cannot
// check them: it also takes withdrawn codes such as su, and reserved ones
// such as ac, as countries.
const countryCodes = `
AD AE AF AG AI AL AM AO AQ AR AS AT AU AW AX AZ
BA BB BD BE BF BG BH BI BJ BL BM BN BO BQ BR BS BT BV BW BY BZ
CA CC CD CF CG CH CI CK CL CM CN CO CR CU CV CW CX CY CZ
DE DJ DK DM DO DZ
EC EE EG EH ER ES ET
FI FJ FK FM FO FR
GA GB GD GE GF GG GH GI GL GM GN GP GQ GR GS GT GU GW GY
HK HM HN HR HT HU
ID IE IL IM IN IO IQ IR IS IT
JE JM JO JP
KE KG KH KI KM KN KP KR KW KY KZ
LA LB LC LI LK LR LS LT LU LV LY
MA MC MD ME MF MG MH MK ML MM MN MO MP MQ MR MS MT MU MV MW MX MY MZ
NA NC NE NF NG NI NL NO NP NR NU NZ
OM
PA PE PF PG PH PK PL PM PN PR PS PT PW PY
QA
RE RO RS RU RW
SA SB SC SD SE SG SH SI SJ SK SL SM SN SO SR SS ST SV SX SY SZ
TC TD TF TG TH TJ TK TL TM TN TO TR TT TV TW TZ
UA UG UM US UY UZ
VA VC VE VG VI VN VU
WF WS
XK
YE YT
ZA ZM ZW
`
// parseCountries reads a comma-separated list of country codes in either
// case, and returns them in capitals.
func parseCountries(value string) ([]string, error) {
items, err := parseList(value)
if err != nil {
return nil, err
}
known := strings.Fields(countryCodes)
countries := make([]string, 0, len(items))
for _, item := range items {
country := strings.ToUpper(item)
if !slices.Contains(known, country) {
return nil, fmt.Errorf("%q %w", item, errNotCountry)
}
countries = append(countries, country)
}
return countries, 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
}