// 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/http" "net/netip" "net/url" "path/filepath" "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 // ClientRequestHeaderMaxBytes is the largest request line and headers // a client may send (SWWAF_CLIENT_REQUEST_HEADER_MAX_BYTES). It is // never off, and always more than 4K. ClientRequestHeaderMaxBytes int64 // ClientIdleTimeout bounds how long a kept-open client connection // may wait for its next request (SWWAF_CLIENT_IDLE_TIMEOUT). ClientIdleTimeout 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 // BanResponse is the status a refused client is answered with, 403 // or 429, or 0 to close the connection without an answer // (SWWAF_BAN_RESPONSE). It answers a banned client, a request that // breaks a rate limit, SWWAF_DENY_NETS and the country lists. BanResponse int // LimitBanDuration is the ban for a first broken rate limit // (SWWAF_LIMIT_BAN_DURATION). A limit broken again within // LimitBanRepeatWindow after the last ban ended bans for three times // as long as that ban (SWWAF_LIMIT_BAN_REPEAT_WINDOW), and a ban that // would be longer than MaxBanDuration is permanent instead // (SWWAF_MAX_BAN_DURATION). None of them can be off. LimitBanDuration time.Duration LimitBanRepeatWindow time.Duration MaxBanDuration time.Duration // MaxBans is the most bans held (SWWAF_MAX_BANS). MaxBans int // BanScopeV4Prefix is the length of the netblock around an IPv4 // client that a ban covers (SWWAF_BAN_SCOPE_V4_PREFIX). BanScopeV4Prefix int // StateDir is the directory of the state files, an absolute path // (SWWAF_STATE_DIR). bans.json is written StateWriteDelay after a ban // is made (SWWAF_STATE_WRITE_DELAY), and every state file every // StateCounterInterval (SWWAF_STATE_COUNTER_INTERVAL). Neither can be // off. StateDir string StateWriteDelay time.Duration StateCounterInterval time.Duration // 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 ipv4Bits = 32 ) 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") errNotOver4K = errors.New("is not a size of more than 4K, such as 32K") errNotDurationAboveZero = errors.New( "is not a duration above zero, such as 1h or 7d") errNotNumberAboveZero = errors.New( "is not a whole number above zero, such as 5000") errNotBanResponse = errors.New("is not 403, 429 or close") errNotV4Prefix = errors.New( "is not the length of an IPv4 netblock, from 0 to 32, such as 24") errNotAbsolutePath = errors.New( "is not an absolute path, such as /var/lib/smallwebwaf") ) // 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"), ClientRequestHeaderMaxBytes: env.headerSize( "SWWAF_CLIENT_REQUEST_HEADER_MAX_BYTES", "32K"), ClientIdleTimeout: env.duration("SWWAF_CLIENT_IDLE_TIMEOUT", "120s"), 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", ""), BanResponse: env.banResponse("SWWAF_BAN_RESPONSE", "403"), LimitBanDuration: env.durationNotOff("SWWAF_LIMIT_BAN_DURATION", "1h"), LimitBanRepeatWindow: env.durationNotOff("SWWAF_LIMIT_BAN_REPEAT_WINDOW", "24h"), MaxBanDuration: env.durationNotOff("SWWAF_MAX_BAN_DURATION", "7d"), MaxBans: env.numberNotOff("SWWAF_MAX_BANS", "5000"), BanScopeV4Prefix: env.v4Prefix("SWWAF_BAN_SCOPE_V4_PREFIX", "32"), StateDir: env.absolutePath("SWWAF_STATE_DIR", "/var/lib/smallwebwaf"), StateWriteDelay: env.durationNotOff("SWWAF_STATE_WRITE_DELAY", "10s"), StateCounterInterval: env.durationNotOff("SWWAF_STATE_COUNTER_INTERVAL", "15m"), } 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 } // headerSize reads the setting that is the largest request line and // headers. func (e *environment) headerSize(name, defaultValue string) int64 { size, err := parseHeaderSize(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 } // durationNotOff reads a setting that is a duration and, unlike a // timeout, cannot be off. func (e *environment) durationNotOff(name, defaultValue string) time.Duration { duration, err := parseDurationNotOff(e.value(name, defaultValue)) e.check(name, err) return duration } // numberNotOff reads a setting that is a whole number above zero, which // cannot be off. func (e *environment) numberNotOff(name, defaultValue string) int { number, err := parseNumberNotOff(e.value(name, defaultValue)) e.check(name, err) return number } // banResponse reads a setting that is how a refused client is answered. func (e *environment) banResponse(name, defaultValue string) int { status, err := parseBanResponse(e.value(name, defaultValue)) e.check(name, err) return status } // v4Prefix reads a setting that is the length of an IPv4 netblock. func (e *environment) v4Prefix(name, defaultValue string) int { length, err := parseV4Prefix(e.value(name, defaultValue)) e.check(name, err) return length } // absolutePath reads a setting that is an absolute path. func (e *environment) absolutePath(name, defaultValue string) string { path := e.value(name, defaultValue) if !filepath.IsAbs(path) { e.check(name, fmt.Errorf("%q %w", path, errNotAbsolutePath)) } return path } // 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 } // parseHeaderSize reads the largest request line and headers: a size as // parseSize reads it, but more than 4K and never off. Go's server reads 4K // past the limit it is given before it refuses, so proxy.New gives it this // size less 4K, which must leave a limit. func parseHeaderSize(value string) (int64, error) { size, err := parseSize(value) if err != nil || size <= 4*kibibyte { return 0, fmt.Errorf("%q %w", value, errNotOver4K) } return size, 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 } // parseDurationNotOff reads a duration above zero, as parseDuration does, // but not off. func parseDurationNotOff(value string) (time.Duration, error) { duration, err := parseDuration(value) if err != nil || duration == 0 { return 0, fmt.Errorf("%q %w", value, errNotDurationAboveZero) } return duration, nil } // parseNumberNotOff reads a whole number above zero. func parseNumberNotOff(value string) (int, error) { n, err := strconv.Atoi(value) if err != nil || n <= 0 { return 0, fmt.Errorf("%q %w", value, errNotNumberAboveZero) } return n, nil } // parseBanResponse reads how a refused client is answered: 403, 429, or // close, which is 0. func parseBanResponse(value string) (int, error) { switch value { case "403": return http.StatusForbidden, nil case "429": return http.StatusTooManyRequests, nil case "close": return 0, nil default: return 0, fmt.Errorf("%q %w", value, errNotBanResponse) } } // parseV4Prefix reads the length of an IPv4 netblock, from 0 to 32. func parseV4Prefix(value string) (int, error) { n, err := strconv.Atoi(value) if err != nil || n < 0 || n > ipv4Bits { return 0, fmt.Errorf("%q %w", value, errNotV4Prefix) } 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 }