From 0750879e582193e6922604f58c92c9cea09f91eb Mon Sep 17 00:00:00 2001 From: clawbot <35+clawbot@noreply.example.org> Date: Sun, 4 Oct 2026 08:29:41 +0200 Subject: [PATCH] Country allow and deny lists, looked up through GeoJS (closes #44) SWWAF_DENIED_COUNTRIES and SWWAF_EXCLUSIVELY_ALLOWED_COUNTRIES refuse a request with 403 before its body is read or rate-limited, logged as country_denied. internal/lookup asks GeoJS only while a list is set, 200 clients per request, one at a time, keeping answers 7 days. Failures, a redirect or an answer leaving an address out included, are logged without addresses; GeoJS is then left alone a second, doubling to five minutes. Private, loopback and link-local clients are never sent. Deviation, per the issue: no SWWAF_LOOKUP_SOURCE or SWWAF_LOOKUP_TIMEOUT; 403, not SWWAF_BAN_RESPONSE. Deviation: GeoJS's country endpoint, not geo.json. Judgement call: an IPv6 /64 is asked about by its first address; at most 10,000 clients wait. Judgement call: config.go lists the ISO 3166-1 codes; no widely used library holds them. Model: opus-5-5 --- README.md | 134 ++++-- internal/config/config.go | 85 ++++ internal/config/config_test.go | 48 ++ internal/lookup/lookup.go | 385 ++++++++++++++++ internal/lookup/lookup_test.go | 553 +++++++++++++++++++++++ internal/proxy/countries.go | 41 ++ internal/proxy/countries_test.go | 267 +++++++++++ internal/proxy/proxy.go | 12 +- internal/proxy/proxy_test.go | 13 + internal/proxy/request.go | 15 +- internal/requestlog/requestlog.go | 3 + internal/smallwebwaf/smallwebwaf.go | 2 + internal/smallwebwaf/smallwebwaf_test.go | 26 +- 13 files changed, 1521 insertions(+), 63 deletions(-) create mode 100644 internal/lookup/lookup.go create mode 100644 internal/lookup/lookup_test.go create mode 100644 internal/proxy/countries.go create mode 100644 internal/proxy/countries_test.go diff --git a/README.md b/README.md index 593e3f5..0b41fd2 100644 --- a/README.md +++ b/README.md @@ -12,15 +12,15 @@ state in memory and in JSON files you can read and edit, and writes a detailed JSON log line for every request. Status: the first milestone is built -(https://git.eeqj.de/sneak/smallwebwaf/issues/13), and the rate limits of the -second (https://git.eeqj.de/sneak/smallwebwaf/issues/14). `smallwebwaf` passes -each request to the app and the app's answer back, unchanged, within its -timeouts and size limits, works out each client's address, refuses a client that -sends too many requests, and writes a JSON log line for every request. The -country lists and the image an app builds on come with the rest of milestone 2, -and the rest of the design after that, in the order of the build order in -[`SPEC.md`](SPEC.md). The survey of existing tools that led to the design is in -[`EVALUATION.md`](EVALUATION.md). +(https://git.eeqj.de/sneak/smallwebwaf/issues/13), and the rate limits and +country lists of the second (https://git.eeqj.de/sneak/smallwebwaf/issues/14). +`smallwebwaf` passes each request to the app and the app's answer back, +unchanged, within its timeouts and size limits, works out each client's address, +refuses a client that sends too many requests or comes from a country you +refuse, and writes a JSON log line for every request. The image an app builds on +comes with the rest of milestone 2, and the rest of the design after that, in +the order of the build order in [`SPEC.md`](SPEC.md). The survey of existing +tools that led to the design is in [`EVALUATION.md`](EVALUATION.md). ## Getting started @@ -72,6 +72,13 @@ gives those in progress five seconds to finish. much of it the window still covers. At most 20,000 clients are kept, the least recently seen dropped first, and only in memory: a restart starts every client afresh. +- Refuses a request from a country you refuse with `403`, as soon as the + client's country is known and before its body is read; such a request is not + counted for the rate limits. While one of the country lists below is set, each + client's country is looked up through GeoJS (see "Country and AS number + lookup" below); with neither set, no visitor's address leaves the host. A + client on a private, loopback or link-local address has no country, and + neither list checks it. - Writes a line in the request log for each request (see "Request log" below). ## Settings @@ -102,20 +109,30 @@ it, and the effective settings are logged at start. requests a client may make in a minute, an hour and a day. The defaults are several times what one busy person produces, since a browser loading a heavy page makes a few hundred requests and several people often share one address. +- `SWWAF_DENIED_COUNTRIES` (default empty): countries whose clients are refused, + for example `cn,ru,kp`. +- `SWWAF_EXCLUSIVELY_ALLOWED_COUNTRIES` (default empty): when set, the only + countries whose clients get through, for example `us,de`. A client whose + country cannot be found is refused too, so that new clients are not let in + whenever GeoJS stops answering. Durations are in Go's syntax, with `d` for days (`90s`, `15m`, `7d`). Sizes are bytes, with an optional `K`, `M` or `G`, which are powers of 1024 (`1K` is 1024 bytes). Rate limits are whole numbers of requests. Netblocks are in CIDR form, -and a bare address stands for itself alone. `off` switches a timeout, a size -limit or a rate limit off. +and a bare address stands for itself alone. Countries are the two-letter codes +ISO 3166-1 assigns today, and `xk` for Kosovo, in either case (`de` and `DE` are +the same); any other code, such as `nk` (North Korea is `kp`) or the withdrawn +`su`, stops the start, and so does a code on both country lists. `off` switches +a timeout, a size limit or a rate limit off. -Four limits are fixed rather than settings. The request line and headers may +Several limits are fixed rather than settings. The request line and headers may take up to 32 KiB, above which the answer is `431` and nothing reaches the app. A kept-open connection that sends nothing for 120 seconds is closed. That 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. At most 20,000 clients are kept for the rate limits, and an IPv6 client -is counted by its /64. +is counted by its /64. A new client waits at most a second for its country, and +at most 100,000 answers from GeoJS are kept, for 7 days each. ## Request log @@ -123,18 +140,23 @@ is counted by its /64. refused ones included: ``` -{"type":"request","time":"2026-10-03T12:00:00.123Z","client_ip":"203.0.113.9","peer_ip":"172.18.0.2","method":"GET","host":"app.example","path":"/","query":"","protocol":"HTTP/1.1","status":200,"upstream_status":200,"request_bytes":0,"response_bytes":5120,"referer":"","user_agent":"curl/8.9.1","action":"forward","duration_total":3.217,"duration_upstream_total":3.104} +{"type":"request","time":"2026-10-03T12:00:00.123Z","client_ip":"203.0.113.9","peer_ip":"172.18.0.2","country":"DE","method":"GET","host":"app.example","path":"/","query":"","protocol":"HTTP/1.1","status":200,"upstream_status":200,"request_bytes":0,"response_bytes":5120,"referer":"","user_agent":"curl/8.9.1","action":"forward","duration_total":3.217,"duration_upstream_total":3.104} ``` - `time` is when the request arrived, in UTC. `peer_ip` is the TCP peer, normally traefik. `path` and `query` are as the client sent them. +- `country` is the client's country as GeoJS places it, and empty when it is not + known: with neither country list set, for a client on a private, loopback or + link-local address, and when GeoJS cannot place the client or has not answered + in time. - `status` is what the client was sent, `0` if nothing was; `upstream_status` is what the app answered, and is left out when the app did not answer. - `request_bytes` and `response_bytes` count body bytes. -- `action` is `forward` for a request passed to the app, `rate_limited` for one - refused for a rate limit, `too_large` for a request or response over its size - limit, `timed_out` for one that ran out of time, and `upstream_error` when the - app could not be reached or its answer broke off. +- `action` is `forward` for a request passed to the app, `country_denied` for + one refused for its client's country, `rate_limited` for one refused for a + rate limit, `too_large` for a request or response over its size limit, + `timed_out` for one that ran out of time, and `upstream_error` when the app + could not be reached or its answer broke off. - `limit_hit` is there for a request refused for a rate limit, and names the window whose limit it went over: `minute`, `hour` or `day`, the shortest if it went over several. @@ -367,34 +389,49 @@ the metrics, failure behaviour and the build order. ## Country and AS number lookup -`smallwebwaf` looks up the AS number and country of every client, for the -request log, the metrics and the ban notes, and for the country lists and biased -limits when you set them. It works with no setup: by default it asks the free -GeoJS web service, which needs no account and no file. This means that, by -default, the address of every new visitor is sent to GeoJS. Each answer is kept -in memory for seven days, and many addresses are asked about in one request; -writing the answers to disk, so that they survive a restart, comes in milestone -3 or later (see the build order in [`SPEC.md`](SPEC.md)). GeoJS publishes no -rate limit but may block a caller it thinks asks too much; while it is not -answering, new visitors count as coming from an unknown country, which +So far `smallwebwaf` looks up only the country, only through GeoJS, and only +while `SWWAF_DENIED_COUNTRIES` or `SWWAF_EXCLUSIVELY_ALLOWED_COUNTRIES` is set: +then the address of every new visitor is sent to GeoJS, and with neither set, +none is. An IPv6 visitor is asked about by the first address of its /64. A new +visitor waits at most a second for its answer, and without one counts as coming +from an unknown country until the answer arrives. The addresses waiting are +asked about together, up to 200 in one request, one request at a time; at most +10,000 visitors wait, and one more counts as coming from an unknown country +until there is room. While GeoJS fails, visitors with a kept answer are +unaffected and new ones count as coming from an unknown country. GeoJS is then +left alone for a second, twice as long after each further failure up to five +minutes, and asked again by the next request that needs it. + +In the full design, `smallwebwaf` looks up the AS number and country of every +client, for the request log, the metrics and the ban notes, and for the country +lists and biased limits when you set them. It works with no setup: by default it +asks the free GeoJS web service, which needs no account and no file. This means +that, by default, the address of every new visitor is sent to GeoJS. Each answer +is kept in memory for seven days, and many addresses are asked about in one +request; writing the answers to disk, so that they survive a restart, comes in +milestone 3 or later (see the build order in [`SPEC.md`](SPEC.md)). GeoJS +publishes no rate limit but may block a caller it thinks asks too much; while it +is not answering, new visitors count as coming from an unknown country, which `SWWAF_EXCLUSIVELY_ALLOWED_COUNTRIES` refuses. To keep your visitors' addresses on your own host, set `SWWAF_LOOKUP_SOURCE=off`, or use the database file instead of GeoJS: `SWWAF_LOOKUP_SOURCE=file` reads the free IPinfo Lite database -(`ipinfo_lite.mmdb`). You download it with your own IPinfo account, mount the -directory that holds it into the container, point `SWWAF_LOOKUP_DB_PATH` at the -file and refresh it when you choose; `smallwebwaf` never downloads it itself, -and reads it again when you replace it. It has to be the directory rather than -the file itself: docker does not show a single mounted file being replaced, so a -refresh would go unseen. IPinfo releases it under the Creative Commons -Attribution-ShareAlike 4.0 International License and asks for attribution, in -its own words on https://ipinfo.io/lite: "The attribution requirements can be -met by giving our service credit as your data source. Simply place a link to -IPinfo on the website, application, or social media account that uses our data." -Its example of such a credit is a link mentioning "IP address data is powered by -IPinfo". A service that uses the database through `smallwebwaf` should carry -that link. +(`ipinfo_lite.mmdb`). `SWWAF_LOOKUP_SOURCE` comes in milestone 3 or later (see +the build order in [`SPEC.md`](SPEC.md)); until then GeoJS is asked only while a +country list is set. You download the database with your own IPinfo account, +mount the directory that holds it into the container, point +`SWWAF_LOOKUP_DB_PATH` at the file and refresh it when you choose; `smallwebwaf` +never downloads it itself, and reads it again when you replace it. It has to be +the directory rather than the file itself: docker does not show a single mounted +file being replaced, so a refresh would go unseen. IPinfo releases it under the +Creative Commons Attribution-ShareAlike 4.0 International License and asks for +attribution, in its own words on https://ipinfo.io/lite: "The attribution +requirements can be met by giving our service credit as your data source. Simply +place a link to IPinfo on the website, application, or social media account that +uses our data." Its example of such a credit is a link mentioning "IP address +data is powered by IPinfo". A service that uses the database through +`smallwebwaf` should carry that link. Neither source can place a private address, so a client on one, such as a visitor on your local network, another container or your monitoring, has no @@ -413,16 +450,18 @@ refusal comes with `SWWAF_ALLOW_NETS` in milestone 3 or later. the checks, passes the request to the app and the answer back with the standard library's `httputil.ReverseProxy` within the timeouts and size limits, and writes the request's log line. Its `check` method is where a - request is refused before anything reaches the app: for a rate limit, for an - announced body over the size limit, and, with the rest of milestone 2, for the - country lists. + request is refused before anything reaches the app: for the country lists, for + a rate limit, and for an announced body over the size limit. +- `internal/lookup`: looks up each client's country through GeoJS, and keeps the + answers. - `internal/ratelimit`: counts each client's requests and tells when one takes it over a rate limit. - `internal/requestlog`: the lines on stdout: the request log line and the process's own messages. Besides the Go standard library, `github.com/hashicorp/golang-lru/v2` keeps the -table of clients to 20,000, dropping the least recently seen. +table of clients to 20,000 and the GeoJS answers to 100,000, dropping the least +recently seen. The country codes are the list in `internal/config/config.go`. ## Entrypoints @@ -456,8 +495,9 @@ so that they run in minimal containers. ## TODO -- Milestone 2: the country lists and the image an app builds on - (https://git.eeqj.de/sneak/smallwebwaf/issues/14); its rate limits are built. +- Milestone 2: the image an app builds on + (https://git.eeqj.de/sneak/smallwebwaf/issues/14); its rate limits and country + lists are built. - The rest of the design, in the order of the build order in [`SPEC.md`](SPEC.md). diff --git a/internal/config/config.go b/internal/config/config.go index 703aa18..2dda27e 100644 --- a/internal/config/config.go +++ b/internal/config/config.go @@ -11,6 +11,7 @@ import ( "net" "net/netip" "net/url" + "slices" "strconv" "strings" "time" @@ -51,6 +52,13 @@ type Config struct { 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. @@ -84,6 +92,9 @@ var ( 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 @@ -104,6 +115,16 @@ func FromEnvironment(lookupEnv func(string) (string, bool)) (*Config, error) { 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 { @@ -201,6 +222,14 @@ func (e *environment) count(name, defaultValue string) int64 { 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) { @@ -348,6 +377,62 @@ func parseNetblock(value string) (netip.Prefix, error) { 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) { diff --git a/internal/config/config_test.go b/internal/config/config_test.go index 83a910b..2d129e3 100644 --- a/internal/config/config_test.go +++ b/internal/config/config_test.go @@ -28,6 +28,8 @@ const ( rateLimitPerMinute = "SWWAF_RATE_LIMIT_PER_MINUTE" rateLimitPerHour = "SWWAF_RATE_LIMIT_PER_HOUR" rateLimitPerDay = "SWWAF_RATE_LIMIT_PER_DAY" + deniedCountries = "SWWAF_DENIED_COUNTRIES" + allowedCountries = "SWWAF_EXCLUSIVELY_ALLOWED_COUNTRIES" ) // off switches a timeout, a size limit or a rate limit off. @@ -79,6 +81,8 @@ func TestDefaults(t *testing.T) { wantNetblocks(t, cfg.TrustedProxies, "10.0.0.0/8", "172.16.0.0/12", "192.168.0.0/16") + wantCountries(t, deniedCountries, cfg.DeniedCountries) + wantCountries(t, allowedCountries, cfg.ExclusivelyAllowedCountries) } func TestValuesAsSet(t *testing.T) { @@ -97,6 +101,8 @@ func TestValuesAsSet(t *testing.T) { rateLimitPerMinute: "60", rateLimitPerHour: "600", rateLimitPerDay: "6000", + deniedCountries: "cn, RU,kp,Xk", + allowedCountries: "de", }) wantSettings(t, cfg, config.Config{ @@ -117,6 +123,25 @@ func TestValuesAsSet(t *testing.T) { } wantNetblocks(t, cfg.TrustedProxies, "192.0.2.1/32", "10.0.0.0/8", "2001:db8::/32") + wantCountries(t, deniedCountries, cfg.DeniedCountries, "CN", "RU", "KP", "XK") + wantCountries(t, allowedCountries, cfg.ExclusivelyAllowedCountries, "DE") +} + +func TestCodeOnBothCountryListsStopsTheStart(t *testing.T) { + t.Parallel() + + _, err := config.FromEnvironment(environment{ + deniedCountries: "cn,ru", + allowedCountries: "de,RU", + }.lookupEnv) + if err == nil { + t.Fatal("ru on both country lists was accepted") + } + + if !strings.HasPrefix(err.Error(), allowedCountries+": ") || + !strings.Contains(err.Error(), `"RU"`) { + t.Errorf("error %q does not name %s and RU", err, allowedCountries) + } } func TestSizesAndOff(t *testing.T) { @@ -198,6 +223,18 @@ func TestInvalidValueStopsTheStart(t *testing.T) { {rateLimitPerHour, "1.5"}, {rateLimitPerDay, "-1"}, {rateLimitPerDay, "lots"}, + {deniedCountries, "nk"}, + {deniedCountries, "kp,,ir"}, + {deniedCountries, "prk"}, + {deniedCountries, "408"}, + {deniedCountries, "k"}, + {deniedCountries, "eu"}, + {deniedCountries, "un"}, + {deniedCountries, "su"}, + {allowedCountries, "ac"}, + {allowedCountries, "uk"}, + {allowedCountries, "zz"}, + {allowedCountries, "de,germany"}, } { t.Run(tc.name+"="+tc.value, func(t *testing.T) { t.Parallel() @@ -245,6 +282,8 @@ func TestLogsEachSettingWithItsValue(t *testing.T) { rateLimitPerMinute: "1000", rateLimitPerHour: "10000", rateLimitPerDay: "50000", + deniedCountries: "", + allowedCountries: "", } if !maps.Equal(line.Settings, want) { t.Errorf("logged settings\n%v\nwant\n%v", line.Settings, want) @@ -282,3 +321,12 @@ func wantNetblocks(t *testing.T, got []netip.Prefix, want ...string) { t.Errorf("netblocks %v, want %v", gotText, want) } } + +// wantCountries checks the list of countries the setting name gave. +func wantCountries(t *testing.T, name string, got []string, want ...string) { + t.Helper() + + if !slices.Equal(got, want) { + t.Errorf("%s gave %v, want %v", name, got, want) + } +} diff --git a/internal/lookup/lookup.go b/internal/lookup/lookup.go new file mode 100644 index 0000000..145df17 --- /dev/null +++ b/internal/lookup/lookup.go @@ -0,0 +1,385 @@ +// Package lookup looks up each client's country through the GeoJS web +// service, and keeps the answers in memory, for at most 100,000 clients +// and for 7 days each. +package lookup + +import ( + "context" + "encoding/json" + "errors" + "fmt" + "io" + "log/slog" + "net/http" + "net/netip" + "strings" + "sync" + "time" + + "github.com/hashicorp/golang-lru/v2/simplelru" +) + +// URL is GeoJS's country endpoint. Asked about several addresses at once, +// comma separated in its ip parameter, it answers with a list. +const URL = "https://get.geojs.io/v1/ip/country.json" + +const ( + // keepFor is how long an answer is used instead of asking GeoJS again. + keepFor = 7 * 24 * time.Hour + // maxAnswers is how many answers are kept. Past it, the one used + // longest ago is dropped. + maxAnswers = 100000 + // maxWaiting is how many clients may wait to be asked about. Past it, + // a new client counts as not found and is not asked about until there + // is room, so that a swarm of new addresses while GeoJS is down cannot + // fill the memory. + maxWaiting = 10000 + // maxPerRequest is how many addresses one request to GeoJS asks about. + maxPerRequest = 200 + // timeout is how long a new client waits for its answer, and how long + // a request to GeoJS may take before it is abandoned. + timeout = time.Second + // After a failure GeoJS is not asked again for a second, and for + // retryDelayFactor times as long after each further failure in a row, + // up to five minutes. + firstRetryDelay = time.Second + retryDelayFactor = 2 + maxRetryDelay = 5 * time.Minute + // maxResponseBytes is the most of GeoJS's answer that is read. + maxResponseBytes = 1 << 20 +) + +var ( + errStatus = errors.New("GeoJS answered") + errLeftOut = errors.New("GeoJS's answer left out") +) + +// Params are what New needs. +type Params struct { + // URL is where GeoJS is asked, normally URL. + URL string + // Now tells the time, normally time.Now. + Now func() time.Time + // ProcessLog receives GeoJS's failures. + ProcessLog *slog.Logger +} + +// GeoJS looks up clients' countries through GeoJS. At most one request +// to GeoJS is under way at a time, and it asks about every client waiting, +// up to maxPerRequest. It is safe for concurrent use. +type GeoJS struct { + url string + now func() time.Time + processLog *slog.Logger + // httpClient follows no redirect, so that visitors' addresses go to + // GeoJS alone: a redirect is a failure. + httpClient *http.Client + + mu sync.Mutex + answers *simplelru.LRU[netip.Prefix, answer] + // waiting are the clients without an answer: those to ask GeoJS about, + // and those it is being asked about. + waiting map[netip.Prefix]*wait + // asking is true while a request to GeoJS is under way. + asking bool + // retryDelay is how long GeoJS is left alone after its last failure, + // zero after an answer; retryAt is when it may be asked again. + retryDelay time.Duration + retryAt time.Time +} + +// answer is what GeoJS said about a client: its country, "" when GeoJS +// cannot place it, and when GeoJS said so. +type answer struct { + country string + received time.Time +} + +// wait is a client waiting for its answer. +type wait struct { + // asked is closed when the client gets its answer, and closed and + // replaced each time GeoJS fails before then. + asked chan struct{} + // late is true once the client has gone without an answer, for a + // whole timeout or because GeoJS failed: its requests no longer wait. + late bool +} + +// New returns a GeoJS with no answer kept yet. +func New(params Params) *GeoJS { + answers, err := simplelru.NewLRU[netip.Prefix, answer](maxAnswers, nil) + if err != nil { + panic(err) // NewLRU fails only for a size below one + } + + return &GeoJS{ + url: params.URL, + now: params.Now, + processLog: params.ProcessLog, + httpClient: &http.Client{ + CheckRedirect: func(*http.Request, []*http.Request) error { + return http.ErrUseLastResponse + }, + }, + answers: answers, + waiting: map[netip.Prefix]*wait{}, + } +} + +// Country returns the country GeoJS places client in, as a two-letter +// code in capitals, or "" when the country cannot be found: GeoJS cannot +// place the client, or has not answered in time. An answer is kept for 7 +// days. Without one, a client waits up to timeout for it, unless it has +// gone without one before; until GeoJS answers, the client is asked about +// again in the background. ctx is the context of the client's request, +// and ends the wait when it ends. +// +// GeoJS is asked about the client's first address, which is the client's +// own address for IPv4, and an address in the same place for an IPv6 /64. +func (g *GeoJS) Country(ctx context.Context, client netip.Prefix) string { + country, asked := g.answerOrWait(ctx, client) + if asked == nil { + return country + } + + timer := time.NewTimer(timeout) + defer timer.Stop() + + select { + case <-asked: + case <-timer.C: + case <-ctx.Done(): + } + + g.mu.Lock() + defer g.mu.Unlock() + + country, found := g.kept(client) + + w, waiting := g.waiting[client] + if !found && waiting { + w.late = true + } + + return country +} + +// answerOrWait returns client's kept answer if it has one. Otherwise it +// puts the client among those waiting if there is room, has GeoJS asked +// about them if it can be, and returns what to wait on for the answer, or +// nil when there is nothing to wait for. +func (g *GeoJS) answerOrWait( + ctx context.Context, client netip.Prefix, +) (string, <-chan struct{}) { + g.mu.Lock() + defer g.mu.Unlock() + + country, found := g.kept(client) + if found { + return country, nil + } + + w, waiting := g.waiting[client] + if !waiting && len(g.waiting) < maxWaiting { + w = &wait{asked: make(chan struct{})} + g.waiting[client] = w + } + + g.ask(ctx) + + if w == nil { + return "", nil // too many clients wait already + } + + if !g.asking { + // GeoJS is left alone after a failure, so no answer can come. + w.late = true + } + + if w.late { + return "", nil + } + + return "", w.asked +} + +// kept returns client's answer, if one was received less than keepFor +// ago. +func (g *GeoJS) kept(client netip.Prefix) (string, bool) { + kept, found := g.answers.Get(client) + if !found || g.now().Sub(kept.received) >= keepFor { + return "", false + } + + return kept.country, true +} + +// ask starts asking GeoJS about the waiting clients, unless a request to +// it is under way or it is left alone after a failure. The requests to +// GeoJS are for every client waiting, so they go on when the client's +// request whose ctx is given ends. +func (g *GeoJS) ask(ctx context.Context) { + if g.asking || g.now().Before(g.retryAt) { + return + } + + g.asking = true + + go g.askAboutWaiting(context.WithoutCancel(ctx)) +} + +// askAboutWaiting asks GeoJS about the waiting clients, one request at a +// time, until none is left or GeoJS fails. +func (g *GeoJS) askAboutWaiting(ctx context.Context) { + for { + clients := g.nextClients() + if len(clients) == 0 { + return + } + + countries, err := g.request(ctx, clients) + if !g.keep(clients, countries, err) { + return + } + } +} + +// nextClients returns up to maxPerRequest of the waiting clients. When +// none is waiting, it returns none and notes that no request to GeoJS is +// under way. +func (g *GeoJS) nextClients() []netip.Prefix { + g.mu.Lock() + defer g.mu.Unlock() + + if len(g.waiting) == 0 { + g.asking = false + + return nil + } + + clients := make([]netip.Prefix, 0, min(len(g.waiting), maxPerRequest)) + + for client := range g.waiting { + if len(clients) == maxPerRequest { + break + } + + clients = append(clients, client) + } + + return clients +} + +// keep notes how a request to GeoJS about clients ended, and reports +// whether GeoJS answered about all of them. Each client whose address +// GeoJS's answer names gets its answer, with no country when GeoJS gave +// none. An answer that leaves an address out is a failure. After a +// failure GeoJS is left alone for a while, and every client still waiting +// stops waiting and is asked about once GeoJS is asked again. +func (g *GeoJS) keep( + clients []netip.Prefix, countries map[netip.Addr]string, err error, +) bool { + g.mu.Lock() + defer g.mu.Unlock() + + now := g.now() + leftOut := 0 + + for _, client := range clients { + country, named := countries[client.Addr()] + if !named { + leftOut++ + + continue + } + + g.answers.Add(client, answer{country: country, received: now}) + close(g.waiting[client].asked) + delete(g.waiting, client) + } + + if err == nil && leftOut > 0 { + err = fmt.Errorf("%w %d of %d addresses", errLeftOut, leftOut, len(clients)) + } + + if err != nil { + g.retryDelay = min(max(retryDelayFactor*g.retryDelay, firstRetryDelay), + maxRetryDelay) + g.retryAt = now.Add(g.retryDelay) + g.asking = false + + for _, w := range g.waiting { + close(w.asked) + + w.asked = make(chan struct{}) + w.late = true + } + + g.processLog.Warn("asking GeoJS failed", + "error", err.Error(), "asking_again_in", g.retryDelay.String()) + + return false + } + + g.retryDelay = 0 + + return true +} + +// request asks GeoJS about clients in one request, and returns the +// country it gave, in capitals, for each address its answer names. +func (g *GeoJS) request( + ctx context.Context, clients []netip.Prefix, +) (map[netip.Addr]string, error) { + addrs := make([]string, 0, len(clients)) + + for _, client := range clients { + addrs = append(addrs, client.Addr().String()) + } + + ctx, cancel := context.WithTimeout(ctx, timeout) + defer cancel() + + req, err := http.NewRequestWithContext(ctx, http.MethodGet, g.url, http.NoBody) + if err != nil { + return nil, fmt.Errorf("make the request to GeoJS: %w", err) + } + + req.URL.RawQuery = "ip=" + strings.Join(addrs, ",") + + res, err := g.httpClient.Do(req) + if err != nil { + // Do's error names the URL, and so the visitors' addresses, which + // are not to be logged: only what went wrong is kept. + return nil, fmt.Errorf("ask GeoJS: %w", errors.Unwrap(err)) + } + + defer func() { + _ = res.Body.Close() + }() + + if res.StatusCode != http.StatusOK { + return nil, fmt.Errorf("%w %s", errStatus, res.Status) + } + + var answers []struct { + IP string `json:"ip"` + Country string `json:"country"` + } + + err = json.NewDecoder(io.LimitReader(res.Body, maxResponseBytes)).Decode(&answers) + if err != nil { + return nil, fmt.Errorf("read GeoJS's answer: %w", err) + } + + countries := make(map[netip.Addr]string, len(answers)) + + for _, item := range answers { + addr, err := netip.ParseAddr(item.IP) + if err == nil { + countries[addr] = strings.ToUpper(item.Country) + } + } + + return countries, nil +} diff --git a/internal/lookup/lookup_test.go b/internal/lookup/lookup_test.go new file mode 100644 index 0000000..bcfc3a1 --- /dev/null +++ b/internal/lookup/lookup_test.go @@ -0,0 +1,553 @@ +package lookup_test + +import ( + "encoding/json" + "log/slog" + "net/http" + "net/http/httptest" + "net/netip" + "slices" + "strings" + "sync" + "testing" + "time" + + "sneak.berlin/go/smallwebwaf/internal/lookup" +) + +const ( + // germany is where the stand-in for GeoJS places every address but + // unplaced. + germany = "DE" + // unplaced is the address it cannot place. + unplaced = "192.0.2.1" + // leftOut is the address it leaves out of its answer when + // answeringWithoutLeftOut. + leftOut = "203.0.113.7" + // timeout is how long a new client waits for its answer. + timeout = time.Second + // waitLimit bounds how long a test waits for what should happen. + waitLimit = 10 * time.Second + // pollInterval is how often a test looks again. + pollInterval = 10 * time.Millisecond + // week is how long an answer is kept. + week = 7 * 24 * time.Hour +) + +func TestKeptAnswerIsUsedFor7DaysThenAskedAgain(t *testing.T) { + t.Parallel() + + geojs, clock, g := start(t) + placed := netip.MustParsePrefix("203.0.113.9/32") + notPlaced := netip.MustParsePrefix(unplaced + "/32") + + wantCountry(t, g, placed, germany) + wantCountry(t, g, notPlaced, "") + wantRequests(t, geojs, 2) + + // An answer without a country is kept too. + clock.advance(week - time.Second) + wantCountry(t, g, placed, germany) + wantCountry(t, g, notPlaced, "") + wantRequests(t, geojs, 2) + + clock.advance(time.Second) + wantCountry(t, g, placed, germany) + wantRequests(t, geojs, 3) + wantAsked(t, geojs, 2, "203.0.113.9") +} + +func TestNewClientWaitsAtMostOneSecondThenCountsAsNotFound(t *testing.T) { + t.Parallel() + + geojs, clock, g := start(t) + client := netip.MustParsePrefix("203.0.113.9/32") + + // The client comes while GeoJS is asked about an earlier client, which + // it answers most of a second later. It is then asked about the client + // and does not answer: that request is abandoned a second after it + // began, well after the client's wait is over. + geojs.set(answeringSlowly) + + var earlier sync.WaitGroup + + earlier.Go(func() { g.Country(t.Context(), netip.MustParsePrefix("203.0.113.1/32")) }) + defer earlier.Wait() + + waitForRequests(t, geojs, 1) + geojs.set(hanging) + + began := time.Now() + + wantCountry(t, g, client, "") + + took := time.Since(began) + if took < timeout || took > timeout+timeout/2 { + t.Errorf("waited %s for the answer, want %s", took, timeout) + } + + // Its next request does not wait. + began = time.Now() + + wantCountry(t, g, client, "") + + took = time.Since(began) + if took > timeout/2 { + t.Errorf("waited %s again, want no wait", took) + } + + // Once GeoJS answers, the client is asked about again in the + // background, and has its country. + geojs.set(answering) + waitForCountry(t, g, clock, client, germany) +} + +func TestAddressLeftOutOfAnAnswerIsAskedAboutAgain(t *testing.T) { + t.Parallel() + + for _, tc := range []struct { + name string + answers int + // named is whether the answer names the other client asked about. + named bool + }{ + {"null", answeringNull, false}, + {"empty list", answeringEmptyList, false}, + {"list without " + leftOut, answeringWithoutLeftOut, true}, + } { + t.Run(tc.name, func(t *testing.T) { + t.Parallel() + + geojs, clock, g := start(t) + other := netip.MustParsePrefix("203.0.113.1/32") + client := netip.MustParsePrefix(leftOut + "/32") + + // GeoJS fails, and is left alone for a second while the client + // comes too, so that the next request asks about both. + geojs.set(failing) + wantCountry(t, g, other, "") + wantCountry(t, g, client, "") + + geojs.set(tc.answers) + clock.advance(time.Second) + wantCountry(t, g, other, "") + waitForRequests(t, geojs, 2) + + // The answer counts as a failure, and the client is asked about + // again, with the other client only if the answer left it out too. + geojs.set(answering) + waitForCountry(t, g, clock, client, germany) + wantCountry(t, g, other, germany) + wantRequests(t, geojs, 3) + + if tc.named { + wantAsked(t, geojs, 2, leftOut) + } else { + wantAsked(t, geojs, 2, leftOut, "203.0.113.1") + } + }) + } +} + +func TestRedirectCountsAsFailure(t *testing.T) { + t.Parallel() + + geojs, _, g := start(t) + geojs.set(redirecting) + + wantCountry(t, g, netip.MustParsePrefix("203.0.113.9/32"), "") + wantRequests(t, geojs, 1) +} + +func TestCountryIsKeptInCapitals(t *testing.T) { + t.Parallel() + + geojs, _, g := start(t) + geojs.set(answeringInLowerCase) + + wantCountry(t, g, netip.MustParsePrefix("203.0.113.9/32"), germany) +} + +func TestFailureIsLoggedWithoutTheAddressesAskedAbout(t *testing.T) { + t.Parallel() + + var log strings.Builder + + // Nothing listens on port 1, so asking GeoJS fails. + g := lookup.New(lookup.Params{ + URL: "http://127.0.0.1:1", + Now: time.Now, + ProcessLog: slog.New(slog.NewTextHandler(&log, nil)), + }) + + wantCountry(t, g, netip.MustParsePrefix("203.0.113.9/32"), "") + + logged := log.String() + if !strings.Contains(logged, "asking GeoJS failed") || + strings.Contains(logged, "203.0.113.9") { + t.Errorf("logged %q, want the failure without the address asked about", logged) + } +} + +func TestWaitingClientsAreAskedAboutInOneRequest(t *testing.T) { + t.Parallel() + + geojs, clock, g := start(t) + + // GeoJS fails, and is then left alone for a second, while three more + // clients come. An IPv6 client is a /64, and GeoJS is asked about its + // first address. + geojs.set(failing) + wantCountry(t, g, netip.MustParsePrefix("203.0.113.1/32"), "") + wantCountry(t, g, netip.MustParsePrefix("203.0.113.2/32"), "") + wantCountry(t, g, netip.MustParsePrefix("2001:db8:1:2::/64"), "") + wantRequests(t, geojs, 1) + + geojs.set(answering) + clock.advance(time.Second) + wantCountry(t, g, netip.MustParsePrefix("203.0.113.3/32"), germany) + wantRequests(t, geojs, 2) + wantAsked(t, geojs, 1, "203.0.113.1", "203.0.113.2", "2001:db8:1:2::", "203.0.113.3") +} + +func TestKeptAnswersUnaffectedWhileGeoJSFailsAndAskedAgainWithBackoff(t *testing.T) { + t.Parallel() + + geojs, clock, g := start(t) + clients := newClients() + kept := clients() + + wantCountry(t, g, kept, germany) + + geojs.set(failing) + wantCountry(t, g, kept, germany) + wantRequests(t, geojs, 1) + + // Each failure leaves GeoJS alone twice as long as the one before, up + // to five minutes. New clients meanwhile count as not found, and the + // client with a kept answer still gets its country, without GeoJS being + // asked. + requests := 1 + + for _, delay := range []time.Duration{ + time.Second, 2 * time.Second, 4 * time.Second, 8 * time.Second, + 16 * time.Second, 32 * time.Second, 64 * time.Second, 128 * time.Second, + 256 * time.Second, 5 * time.Minute, 5 * time.Minute, + } { + wantCountry(t, g, clients(), "") + + requests++ + wantRequests(t, geojs, requests) + + clock.advance(delay - time.Millisecond) + wantCountry(t, g, clients(), "") + wantCountry(t, g, kept, germany) + wantRequests(t, geojs, requests) + + clock.advance(time.Millisecond) + } + + // Once GeoJS answers again, it is asked about every client waiting. + geojs.set(answering) + wantCountry(t, g, clients(), germany) + wantRequests(t, geojs, requests+1) + + asked := waitForRequests(t, geojs, requests+1) + if len(asked[requests]) != 23 { + t.Errorf("GeoJS was asked about %d clients, want 23", len(asked[requests])) + } +} + +func TestAtMost200AddressesInOneRequest(t *testing.T) { + t.Parallel() + + geojs, clock, g := start(t) + clients := newClients() + first := clients() + + // 201 clients wait while GeoJS is left alone after a failure. + geojs.set(failing) + wantCountry(t, g, first, "") + + for range 200 { + wantCountry(t, g, clients(), "") + } + + // The first one's next request has GeoJS asked again. + geojs.set(answering) + clock.advance(time.Second) + wantCountry(t, g, first, "") + + asked := waitForRequests(t, geojs, 3) + if len(asked[1]) != 200 || len(asked[2]) != 1 { + t.Errorf("GeoJS was asked about %d and then %d clients, want 200 and 1", + len(asked[1]), len(asked[2])) + } +} + +func TestAtMost10000ClientsWait(t *testing.T) { + t.Parallel() + + geojs, clock, g := start(t) + clients := newClients() + first := clients() + + // 10,000 clients wait while GeoJS is left alone after a failure, and + // one more cannot join them. + geojs.set(failing) + wantCountry(t, g, first, "") + + for range 9999 { + wantCountry(t, g, clients(), "") + } + + extra := clients() + wantCountry(t, g, extra, "") + + // The first one's next request has GeoJS asked about the 10,000, 200 + // at a time, and not about the one more. + geojs.set(answering) + clock.advance(time.Second) + wantCountry(t, g, first, "") + + asked := waitForRequests(t, geojs, 51) + for i, request := range asked { + if slices.Contains(request, extra.Addr().String()) { + t.Errorf("request %d asked about %s", i, extra.Addr()) + } + } + + // With room among those waiting, it is asked about. + wantCountry(t, g, extra, germany) +} + +// How the stand-in for GeoJS answers. +const ( + answering = iota + answeringSlowly // most of a second later + answeringInLowerCase // with each country in lower case + answeringWithoutLeftOut // with a list that leaves leftOut out + answeringEmptyList // with [] + answeringNull // with null + failing // with 503 + hanging // not at all, until the request is abandoned + redirecting // with a redirect to itself +) + +// standIn is a stand-in for GeoJS. It notes the addresses each request +// asks about. +type standIn struct { + server *httptest.Server + + mu sync.Mutex + answers int + requests [][]string +} + +// ServeHTTP answers a request about the addresses in its ip parameter. +func (s *standIn) ServeHTTP(w http.ResponseWriter, r *http.Request) { + addrs := strings.Split(r.URL.Query().Get("ip"), ",") + + s.mu.Lock() + s.requests = append(s.requests, addrs) + answers := s.answers + s.mu.Unlock() + + switch answers { + case failing: + w.WriteHeader(http.StatusServiceUnavailable) + + return + case hanging: + <-r.Context().Done() + + return + case redirecting: + http.Redirect(w, r, "/", http.StatusFound) + + return + case answeringSlowly: + select { + case <-time.After(timeout * 4 / 5): + case <-r.Context().Done(): + return + } + } + + list := make([]map[string]string, 0, len(addrs)) + + for _, addr := range addrs { + country := germany + + switch { + case addr == unplaced: + country = "" + case addr == leftOut && answers == answeringWithoutLeftOut: + continue + case answers == answeringInLowerCase: + country = strings.ToLower(germany) + } + + list = append(list, map[string]string{"ip": addr, "country": country}) + } + + var answer any = list + + switch answers { + case answeringEmptyList: + answer = []string{} + case answeringNull: + answer = nil + } + + err := json.NewEncoder(w).Encode(answer) + if err != nil { + http.Error(w, err.Error(), http.StatusInternalServerError) + } +} + +// set sets how the stand-in answers. +func (s *standIn) set(answers int) { + s.mu.Lock() + defer s.mu.Unlock() + + s.answers = answers +} + +// asked returns the addresses each request has asked about so far. +func (s *standIn) asked() [][]string { + s.mu.Lock() + defer s.mu.Unlock() + + return slices.Clone(s.requests) +} + +// testClock is a clock the test sets. +type testClock struct { + mu sync.Mutex + now time.Time +} + +// Now tells the time. +func (c *testClock) Now() time.Time { + c.mu.Lock() + defer c.mu.Unlock() + + return c.now +} + +// advance moves the clock on by d. +func (c *testClock) advance(d time.Duration) { + c.mu.Lock() + defer c.mu.Unlock() + + c.now = c.now.Add(d) +} + +// start starts a stand-in for GeoJS that answers, and returns it, a +// clock, and a GeoJS asking it by that clock. +func start(t *testing.T) (*standIn, *testClock, *lookup.GeoJS) { + t.Helper() + + geojs := &standIn{} + geojs.server = httptest.NewServer(geojs) + t.Cleanup(geojs.server.Close) + + clock := &testClock{now: time.Date(2026, 10, 4, 0, 0, 0, 0, time.UTC)} + g := lookup.New(lookup.Params{ + URL: geojs.server.URL, + Now: clock.Now, + ProcessLog: slog.New(slog.DiscardHandler), + }) + + return geojs, clock, g +} + +// newClients returns what returns a new IPv4 client each time it is +// called. +func newClients() func() netip.Prefix { + addr := netip.MustParseAddr("10.0.0.0") + + return func() netip.Prefix { + addr = addr.Next() + + return netip.PrefixFrom(addr, addr.BitLen()) + } +} + +// wantCountry checks the country g gives client. +func wantCountry(t *testing.T, g *lookup.GeoJS, client netip.Prefix, want string) { + t.Helper() + + got := g.Country(t.Context(), client) + if got != want { + t.Errorf("%s is in %q, want %q", client, got, want) + } +} + +// wantRequests checks how many requests GeoJS has had. +func wantRequests(t *testing.T, geojs *standIn, want int) { + t.Helper() + + got := len(geojs.asked()) + if got != want { + t.Errorf("GeoJS had %d requests, want %d", got, want) + } +} + +// wantAsked checks the addresses request i asked about, in any order. +func wantAsked(t *testing.T, geojs *standIn, i int, want ...string) { + t.Helper() + + asked := geojs.asked() + if len(asked) <= i { + t.Fatalf("GeoJS had %d requests, want more than %d", len(asked), i) + } + + got := slices.Sorted(slices.Values(asked[i])) + + slices.Sort(want) + + if !slices.Equal(got, want) { + t.Errorf("request %d asked about %v, want %v", i, got, want) + } +} + +// waitForRequests waits for GeoJS to have had count requests, and returns +// the addresses each asked about. +func waitForRequests(t *testing.T, geojs *standIn, count int) [][]string { + t.Helper() + + deadline := time.Now().Add(waitLimit) + for time.Now().Before(deadline) { + asked := geojs.asked() + if len(asked) >= count { + return asked + } + + time.Sleep(pollInterval) + } + + t.Fatalf("fewer than %d requests to GeoJS after %s", count, waitLimit) + + return nil +} + +// waitForCountry waits for g to give client the country want, moving the +// clock on a minute at a time, so that GeoJS is asked again after a +// failure. +func waitForCountry( + t *testing.T, g *lookup.GeoJS, clock *testClock, client netip.Prefix, want string, +) { + t.Helper() + + deadline := time.Now().Add(waitLimit) + for g.Country(t.Context(), client) != want { + if time.Now().After(deadline) { + t.Fatalf("%s is not in %q after %s", client, want, waitLimit) + } + + clock.advance(time.Minute) + time.Sleep(pollInterval) + } +} diff --git a/internal/proxy/countries.go b/internal/proxy/countries.go new file mode 100644 index 0000000..1ab04d3 --- /dev/null +++ b/internal/proxy/countries.go @@ -0,0 +1,41 @@ +package proxy + +import ( + "context" + "net/netip" + "slices" +) + +// countryDenied reports whether the country lists refuse the request. +// The client's country is looked up only while a list is set, and never +// for a client on a private, loopback or link-local address, which has +// no country and which neither list checks. A client whose country +// cannot be found is refused only by SWWAF_EXCLUSIVELY_ALLOWED_COUNTRIES. +// ctx is the request's own context. +func (rq *request) countryDenied(ctx context.Context) bool { + denied := rq.h.config.DeniedCountries + allowed := rq.h.config.ExclusivelyAllowedCountries + + if len(denied) == 0 && len(allowed) == 0 { + return false + } + + if !hasCountry(rq.client) { + return false + } + + country := rq.h.geojs.Country(ctx, clientGroup(rq.client)) + rq.line.Country = country + + if slices.Contains(denied, country) { + return true + } + + return len(allowed) > 0 && !slices.Contains(allowed, country) +} + +// hasCountry reports whether addr can be placed in a country: private, +// loopback and link-local addresses cannot. +func hasCountry(addr netip.Addr) bool { + return !addr.IsPrivate() && !addr.IsLoopback() && !addr.IsLinkLocalUnicast() +} diff --git a/internal/proxy/countries_test.go b/internal/proxy/countries_test.go new file mode 100644 index 0000000..b5b999d --- /dev/null +++ b/internal/proxy/countries_test.go @@ -0,0 +1,267 @@ +package proxy_test + +import ( + "encoding/json" + "maps" + "net/http" + "net/http/httptest" + "slices" + "strings" + "sync" + "sync/atomic" + "testing" + "time" + + "sneak.berlin/go/smallwebwaf/internal/requestlog" +) + +// The clients the stand-in for GeoJS knows about. +const ( + // fromDE is placed in Germany. + fromDE = client + // fromKP is placed in North Korea. + fromKP = "198.51.100.7" + // unplaced cannot be placed in any country. + unplaced = "192.0.2.1" +) + +func TestCountryLists(t *testing.T) { + t.Parallel() + + for _, tc := range []struct { + name string + env map[string]string + refused []string + }{ + {"denied", map[string]string{deniedCountries: "kp"}, []string{fromKP}}, + { + "exclusively allowed", map[string]string{allowedCountries: "DE"}, + []string{fromKP, unplaced}, + }, + { + "both", map[string]string{deniedCountries: "kp", allowedCountries: "de,fr"}, + []string{fromKP, unplaced}, + }, + } { + t.Run(tc.name, func(t *testing.T) { + t.Parallel() + + var calls atomic.Int32 + + app := startApp(t, func(http.ResponseWriter, *http.Request) { + calls.Add(1) + }) + geojsURL, _ := startGeoJS(t) + env := map[string]string{trustedProxies: trustLocalhost} + maps.Copy(env, tc.env) + addr, out := startProxyWithGeoJS(t, app.URL, geojsURL, env) + + for i, sent := range []struct{ client, country string }{ + {fromDE, "DE"}, {fromKP, "KP"}, {unplaced, ""}, + } { + req := newRequest(t, http.MethodGet, addr, "/", http.NoBody) + req.Header.Set(forwardedFor, sent.client) + got := do(t, req) + + line := out.requestLines(t, i+1)[i] + if line.Country != sent.country { + t.Errorf("log line has country %q, want %q", line.Country, sent.country) + } + + if slices.Contains(tc.refused, sent.client) { + wantStatus(t, got, http.StatusForbidden) + wantLine(t, line, http.StatusForbidden, requestlog.ActionCountryDenied) + } else { + wantStatus(t, got, http.StatusOK) + wantLine(t, line, http.StatusOK, requestlog.ActionForward) + } + } + + if int(calls.Load()) != 3-len(tc.refused) { + t.Errorf("the app was called %d times, want %d", + calls.Load(), 3-len(tc.refused)) + } + }) + } +} + +func TestCountryRefusalComesBeforeTheBody(t *testing.T) { + t.Parallel() + + var calls atomic.Int32 + + app := startApp(t, func(http.ResponseWriter, *http.Request) { + calls.Add(1) + }) + geojsURL, _ := startGeoJS(t) + addr, out := startProxyWithGeoJS(t, app.URL, geojsURL, map[string]string{ + trustedProxies: trustLocalhost, + deniedCountries: "kp", + }) + + req := newRequest(t, http.MethodPost, addr, "/", strings.NewReader("a body")) + req.Header.Set(forwardedFor, fromKP) + wantStatus(t, do(t, req), http.StatusForbidden) + + line := out.requestLine(t) + wantLine(t, line, http.StatusForbidden, requestlog.ActionCountryDenied) + + if line.RequestBytes != 0 { + t.Errorf("log line has request_bytes %d, want 0", line.RequestBytes) + } + + if calls.Load() != 0 { + t.Errorf("the app was called %d times, want none", calls.Load()) + } +} + +func TestRequestRefusedByCountryIsNotCounted(t *testing.T) { + t.Parallel() + + // The stand-in for GeoJS fails until placing is set, and then places + // every address in Germany. + var placing atomic.Bool + + geojs := httptest.NewServer(http.HandlerFunc( + func(w http.ResponseWriter, r *http.Request) { + if !placing.Load() { + w.WriteHeader(http.StatusServiceUnavailable) + + return + } + + answer := []map[string]string{{"ip": r.URL.Query().Get("ip"), "country": "DE"}} + + err := json.NewEncoder(w).Encode(answer) + if err != nil { + http.Error(w, err.Error(), http.StatusInternalServerError) + } + })) + t.Cleanup(geojs.Close) + + app := startApp(t, func(http.ResponseWriter, *http.Request) {}) + addr, _ := startProxyWithGeoJS(t, app.URL, geojs.URL, map[string]string{ + trustedProxies: trustLocalhost, + allowedCountries: "de", + rateLimitPerMinute: "1", + }) + + request := func() answer { + req := newRequest(t, http.MethodGet, addr, "/", http.NoBody) + req.Header.Set(forwardedFor, fromDE) + + return do(t, req) + } + + // While GeoJS fails, the client's country cannot be found, and its + // request is refused. + wantStatus(t, request(), http.StatusForbidden) + + // Once GeoJS places it, a second after the failure, its requests are let + // through. No refused one was counted, so the first let through is + // within the limit of one a minute. + placing.Store(true) + + deadline := time.Now().Add(waitLimit) + got := request() + + for got.status == http.StatusForbidden && time.Now().Before(deadline) { + time.Sleep(pollInterval) + + got = request() + } + + wantStatus(t, got, http.StatusOK) +} + +func TestCountryNotLookedUpWithoutAListOrForAPrivateAddress(t *testing.T) { + t.Parallel() + + for _, tc := range []struct { + name string + env map[string]string + clients []string // "" sends no X-Forwarded-For: the client is 127.0.0.1 + }{ + {"no country list is set", nil, []string{fromKP, fromDE}}, + { + "private, loopback and link-local addresses", + map[string]string{allowedCountries: "de"}, + []string{"10.0.0.5", "192.168.1.9", "fd00::5", "", "169.254.0.9", "fe80::9"}, + }, + } { + t.Run(tc.name, func(t *testing.T) { + t.Parallel() + + app := startApp(t, func(http.ResponseWriter, *http.Request) {}) + geojsURL, asked := startGeoJS(t) + env := map[string]string{trustedProxies: trustLocalhost} + maps.Copy(env, tc.env) + addr, out := startProxyWithGeoJS(t, app.URL, geojsURL, env) + + for i, sent := range tc.clients { + req := newRequest(t, http.MethodGet, addr, "/", http.NoBody) + if sent != "" { + req.Header.Set(forwardedFor, sent) + } + + wantStatus(t, do(t, req), http.StatusOK) + + line := out.requestLines(t, i+1)[i] + wantLine(t, line, http.StatusOK, requestlog.ActionForward) + + country, present := line.fields["country"] + if !present || country != "" { + t.Errorf("log line for %q has country %v, want an empty one", + line.ClientIP, country) + } + } + + if len(asked()) != 0 { + t.Errorf("GeoJS was asked about %v, want nothing", asked()) + } + }) + } +} + +// startGeoJS starts a stand-in for GeoJS, which places fromDE and fromKP +// and no other address. It returns its URL, and what returns the +// addresses it has been asked about. +func startGeoJS(t *testing.T) (string, func() []string) { + t.Helper() + + places := map[string]string{fromDE: "DE", fromKP: "KP"} + + var asked struct { + mu sync.Mutex + addrs []string + } + + geojs := httptest.NewServer(http.HandlerFunc( + func(w http.ResponseWriter, r *http.Request) { + addrs := strings.Split(r.URL.Query().Get("ip"), ",") + + asked.mu.Lock() + asked.addrs = append(asked.addrs, addrs...) + asked.mu.Unlock() + + answers := make([]map[string]string, 0, len(addrs)) + for _, addr := range addrs { + answers = append(answers, map[string]string{ + "ip": addr, "country": places[addr], + }) + } + + err := json.NewEncoder(w).Encode(answers) + if err != nil { + http.Error(w, err.Error(), http.StatusInternalServerError) + } + })) + t.Cleanup(geojs.Close) + + return geojs.URL, func() []string { + asked.mu.Lock() + defer asked.mu.Unlock() + + return slices.Clone(asked.addrs) + } +} diff --git a/internal/proxy/proxy.go b/internal/proxy/proxy.go index a650c11..531478b 100644 --- a/internal/proxy/proxy.go +++ b/internal/proxy/proxy.go @@ -11,6 +11,7 @@ import ( "time" "sneak.berlin/go/smallwebwaf/internal/config" + "sneak.berlin/go/smallwebwaf/internal/lookup" "sneak.berlin/go/smallwebwaf/internal/ratelimit" ) @@ -40,6 +41,9 @@ type Params struct { RequestLog io.Writer // ProcessLog receives the process's own messages. ProcessLog *slog.Logger + // GeoJSURL is where clients' countries are looked up, normally + // lookup.URL. GeoJS is asked only while a country list is set. + GeoJSURL string } // New returns the server smallwebwaf runs: each request it reads passes @@ -63,6 +67,11 @@ func New(params Params) *http.Server { PerHour: params.Config.RateLimitPerHour, PerDay: params.Config.RateLimitPerDay, }), + geojs: lookup.New(lookup.Params{ + URL: params.GeoJSURL, + Now: time.Now, + ProcessLog: params.ProcessLog, + }), }, ReadHeaderTimeout: params.Config.ClientRequestTimeout, IdleTimeout: clientIdleTimeout, @@ -80,6 +89,7 @@ type handler struct { errorLog *log.Logger transport http.RoundTripper limiter *ratelimit.Limiter + geojs *lookup.GeoJS } // newTransport returns what carries requests to the app. It never goes @@ -101,7 +111,7 @@ func (h *handler) ServeHTTP(w http.ResponseWriter, r *http.Request) { rq := h.newRequest(w, r) defer rq.finish() - refused := rq.check() + refused := rq.check(r.Context()) if refused != nil { rq.answer(*refused) diff --git a/internal/proxy/proxy_test.go b/internal/proxy/proxy_test.go index c6aefe1..0e7dea7 100644 --- a/internal/proxy/proxy_test.go +++ b/internal/proxy/proxy_test.go @@ -49,6 +49,8 @@ const ( responseMaxBytes = "SWWAF_RESPONSE_MAX_BYTES" trustedProxies = "SWWAF_TRUSTED_PROXIES" rateLimitPerMinute = "SWWAF_RATE_LIMIT_PER_MINUTE" + deniedCountries = "SWWAF_DENIED_COUNTRIES" + allowedCountries = "SWWAF_EXCLUSIVELY_ALLOWED_COUNTRIES" ) // output collects what smallwebwaf writes on stdout. @@ -163,6 +165,16 @@ func startApp(t *testing.T, app http.HandlerFunc) *httptest.Server { func startProxy(t *testing.T, appURL string, env map[string]string) (string, *output) { t.Helper() + return startProxyWithGeoJS(t, appURL, "", env) +} + +// startProxyWithGeoJS is startProxy with clients' countries looked up at +// geojsURL. +func startProxyWithGeoJS( + t *testing.T, appURL, geojsURL string, env map[string]string, +) (string, *output) { + t.Helper() + settings := map[string]string{"SWWAF_UPSTREAM_URL": appURL} maps.Copy(settings, env) @@ -180,6 +192,7 @@ func startProxy(t *testing.T, appURL string, env map[string]string) (string, *ou Config: cfg, RequestLog: out, ProcessLog: requestlog.NewProcessLogger(out), + GeoJSURL: geojsURL, }) listener, err := (&net.ListenConfig{}).Listen(t.Context(), "tcp", localhost+":0") diff --git a/internal/proxy/request.go b/internal/proxy/request.go index 6520d6b..185e0b1 100644 --- a/internal/proxy/request.go +++ b/internal/proxy/request.go @@ -103,9 +103,18 @@ func (h *handler) newRequest(w http.ResponseWriter, r *http.Request) *request { // 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. It -// returns nil to let the request through. The rate limits come first, so -// that every request is counted, one refused for its size too. -func (rq *request) check() *refusal { +// returns nil to let the request through. The country lists come first, +// and a request they refuse is not counted for the rate limits; then the +// rate limits, so that every other request is counted, one refused for +// its size too. ctx is the request's own context. +func (rq *request) check(ctx context.Context) *refusal { + if rq.countryDenied(ctx) { + return &refusal{ + status: http.StatusForbidden, + action: requestlog.ActionCountryDenied, + } + } + limitHit := rq.h.limiter.Count(clientGroup(rq.client), rq.start) if limitHit != "" { rq.line.LimitHit = limitHit diff --git a/internal/requestlog/requestlog.go b/internal/requestlog/requestlog.go index c31b9d9..fddcd19 100644 --- a/internal/requestlog/requestlog.go +++ b/internal/requestlog/requestlog.go @@ -26,6 +26,8 @@ const ( // ActionRateLimited is a request refused because it took its client // over a rate limit, or came while the client was over one. ActionRateLimited = "rate_limited" + // ActionCountryDenied is a request refused for its client's country. + ActionCountryDenied = "country_denied" ) // timeLayout is RFC 3339 with milliseconds. @@ -40,6 +42,7 @@ type Line struct { Time string `json:"time"` ClientIP string `json:"client_ip"` PeerIP string `json:"peer_ip"` + Country string `json:"country"` Method string `json:"method"` Host string `json:"host"` Path string `json:"path"` diff --git a/internal/smallwebwaf/smallwebwaf.go b/internal/smallwebwaf/smallwebwaf.go index 7141b63..97e1047 100644 --- a/internal/smallwebwaf/smallwebwaf.go +++ b/internal/smallwebwaf/smallwebwaf.go @@ -16,6 +16,7 @@ import ( "time" "sneak.berlin/go/smallwebwaf/internal/config" + "sneak.berlin/go/smallwebwaf/internal/lookup" "sneak.berlin/go/smallwebwaf/internal/proxy" "sneak.berlin/go/smallwebwaf/internal/requestlog" ) @@ -73,6 +74,7 @@ func Run(ctx context.Context, params Params) int { Config: cfg, RequestLog: params.Stdout, ProcessLog: processLog, + GeoJSURL: lookup.URL, }) processLog.Info("starting", diff --git a/internal/smallwebwaf/smallwebwaf_test.go b/internal/smallwebwaf/smallwebwaf_test.go index fb7341b..b65fedf 100644 --- a/internal/smallwebwaf/smallwebwaf_test.go +++ b/internal/smallwebwaf/smallwebwaf_test.go @@ -176,18 +176,20 @@ func wantStartingLine(t *testing.T, line map[string]any, appURL string) { 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", - "SWWAF_RATE_LIMIT_PER_MINUTE": "1000", - "SWWAF_RATE_LIMIT_PER_HOUR": "10000", - "SWWAF_RATE_LIMIT_PER_DAY": "50000", + 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", + "SWWAF_RATE_LIMIT_PER_MINUTE": "1000", + "SWWAF_RATE_LIMIT_PER_HOUR": "10000", + "SWWAF_RATE_LIMIT_PER_DAY": "50000", + "SWWAF_DENIED_COUNTRIES": "", + "SWWAF_EXCLUSIVELY_ALLOWED_COUNTRIES": "", } for name, value := range want {