diff --git a/.dockerignore b/.dockerignore new file mode 100644 index 0000000..bb5c293 --- /dev/null +++ b/.dockerignore @@ -0,0 +1,75 @@ +# .dockerignore does NOT use .gitignore semantics. Docker matches with +# moby/patternmatcher: filepath.Match plus `**`, so `*` does not cross +# `/` and an unprefixed pattern is anchored at the context root. Every +# depth-independent pattern therefore needs `**/`, or `config/.env` and +# `certs/server.key` still ship while this file reads as solved. Only +# genuinely root-anchored entries go unprefixed. Never transplant these +# into .gitignore, where `**/` is wrong. +# +# Matching is case-sensitive, so secrets use character ranges rather +# than an ALL-CAPS twin, which would still miss `Server.Key`. +# +# Extend with this repo's own host-built artifacts, written anchored: +# `/myapp`, never `**/myapp`, which also matches `cmd/myapp/` and +# deletes the package directory from the context. + +# .git is sent without its config. Without a VERSION build argument the +# stage that compiles runs `git describe --tags --always` on .git, which +# does not need .git/config; that file can hold a credential, such as a +# password in a remote URL or the token the CI checkout step stores there. +# Each submodule keeps a config with the same exposure in its git directory +# under .git/modules/, nested again for a submodule's own submodules, or in +# its own .git directory when it keeps one. +# KNOWN GAP: a submodule whose name has a `config` segment (`config`, +# `deploy/config`, `config/lib`) loses its whole git directory, because +# `**/.git/modules/**/config` also matches that segment's directory +# under .git/modules/. Go's version stamping then fails the build; +# nothing leaks. Name such a submodule without that segment: +# `git submodule add --name`. +**/.git/config +**/.git/modules/**/config + +# Agent scratch: one full checkout of the repo per in-flight agent. +# Anchored because it occurs once where agents run at the repo root. +# KNOWN GAP: a repo running agents in subdirectories still ships +# `services/api/.claude/` and must add its own anchored entry. +.claude + +# Environment files. `*.env` covers bare `.env` and the `prod.env` +# convention. Re-include a committed template with a negation if the +# build needs one: `!docs/example.env`. +**/*.[eE][nN][vV] +**/.[eE][nN][vV].* +**/.[eE][nN][vV][rR][cC] + +# Private keys and the bundles carrying them. Public certificates +# (*.crt, *.cer) are deliberately absent: they are legitimate inputs. +**/*.[pP][eE][mM] +**/*.[kK][eE][yY] +**/*.[pP]12 +**/*.[pP][fF][xX] +**/[iI][dD]_[rR][sS][aA] +**/[iI][dD]_[dD][sS][aA] +**/[iI][dD]_[eE][cC][dD][sS][aA] +**/[iI][dD]_[eE][cC][dD][sS][aA]_[sS][kK] +**/[iI][dD]_[eE][dD]25519 +**/[iI][dD]_[eE][dD]25519_[sS][kK] + +# Dependencies: restored inside the image, never copied in. +**/node_modules + +# OS metadata. +**/.DS_Store +**/Thumbs.db + +# Editor state: never a build input, and it churns COPY. +**/*.swp +**/*.swo +**/*~ +**/*.bak +**/.idea +**/.vscode +**/*.sublime-* + +# The binary `make build` writes on the host; the image builds its own. +/bin diff --git a/.editorconfig b/.editorconfig new file mode 100644 index 0000000..92ec261 --- /dev/null +++ b/.editorconfig @@ -0,0 +1,15 @@ +root = true + +[*] +indent_style = space +indent_size = 4 +end_of_line = lf +charset = utf-8 +trim_trailing_whitespace = true +insert_final_newline = true + +[Makefile] +indent_style = tab + +[*.go] +indent_style = tab diff --git a/.gitea/workflows/check.yml b/.gitea/workflows/check.yml new file mode 100644 index 0000000..ee73864 --- /dev/null +++ b/.gitea/workflows/check.yml @@ -0,0 +1,9 @@ +name: check +on: [push] +jobs: + check: + runs-on: ubuntu-latest + steps: + # actions/checkout v4.2.2, 2026-02-22 + - uses: actions/checkout@11bd71901bbe5b1630ceea73d27597364c9af683 + - run: script/cibuild diff --git a/.gitignore b/.gitignore new file mode 100644 index 0000000..4d6a4b5 --- /dev/null +++ b/.gitignore @@ -0,0 +1,53 @@ +# OS +.DS_Store +Thumbs.db + +# Editors +*.swp +*.swo +*~ +*.bak +.idea/ +.vscode/ +*.sublime-* + +# Agent scratch (worktrees of this repo, created and destroyed by +# in-flight tooling). Unanchored: .gitignore patterns already match at +# every depth, so no prefix is wanted here. This is not a .dockerignore +# entry and must not be given a `**/` prefix on the way into one. +.claude/ + +# Node +node_modules/ + +# Secrets. Unanchored like every entry above, so each matches at every +# depth. Matching is case-sensitive on Linux, so names use character +# ranges rather than a lowercase form that misses `Server.Key`. + +# Environment files. `*.env` covers bare `.env` and the `prod.env` +# convention. Only the templates `example.env` and `sample.env` are +# re-included below. A repository that commits any other template adds +# its own negation after these lines, for example `!.env.example`. +*.[eE][nN][vV] +.[eE][nN][vV].* +.[eE][nN][vV][rR][cC] +!example.env +!sample.env + +# Private keys and the bundles carrying them. +*.[pP][eE][mM] +*.[kK][eE][yY] +*.[pP]12 +*.[pP][fF][xX] +[iI][dD]_[rR][sS][aA] +[iI][dD]_[dD][sS][aA] +[iI][dD]_[eE][cC][dD][sS][aA] +[iI][dD]_[eE][cC][dD][sS][aA]_[sS][kK] +[iI][dD]_[eE][dD]25519 +[iI][dD]_[eE][dD]25519_[sS][kK] + +# Go: the binary `make build` writes, test binaries, profiles and logs +/bin/ +*.test +*.out +*.log diff --git a/.golangci.yml b/.golangci.yml new file mode 100644 index 0000000..1b73eb9 --- /dev/null +++ b/.golangci.yml @@ -0,0 +1,99 @@ +version: "2" + +# Config schema uses the golangci-lint v2 layout (settings live under +# linters.settings, not top-level linters-settings) so that the +# thresholds below are actually applied by golangci-lint >= v2. + +run: + timeout: 5m + modules-download-mode: readonly + +linters: + default: all + enable: + # Successor to the deprecated gomodguard. Named explicitly, rather than + # left to `default: all`, because it carries the module policy below. + - gomodguard_v2 + disable: + # Genuinely incompatible with project patterns + - exhaustruct # Requires all struct fields + - exhaustruct_v5 # Requires all struct fields (successor to exhaustruct) + - godot # Requires comments to end with periods + - wrapcheck # Too verbose for internal packages + - varnamelen # Short names like db, id are idiomatic Go + # Deprecated: the warning is attached to the old name, so it is + # silenced by disabling that name, not by enabling the successor. + - wsl # Deprecated, replaced by wsl_v5 + - gomodguard # Deprecated, replaced by gomodguard_v2 + settings: + lll: + line-length: 88 + funlen: + lines: 80 + statements: 50 + cyclop: + max-complexity: 15 + dupl: + threshold: 100 + depguard: + # Test-support code must not be compiled into the shipped binary. A + # test-support package exists to hand a test privileges the program + # itself must never have, so a file that is not a test must not import + # one. Test files, and the files inside a package whose directory name + # ends in `test`, are where that code belongs, and are exempt. + # + # The deny list below is the one part of this file a repository is + # expected to extend, and the only part it may. depguard matches an + # import path against a list of prefixes, so it cannot be told "any path + # whose last segment ends in test"; a repository's own test-support + # packages have to be named here one at a time, by full import path, + # under a module path that differs from repository to repository. Add + # them; change nothing else. + rules: + test-support: + list-mode: lax + files: + - "$all" + - "!$test" + - "!**/*test/**" + deny: + - pkg: net/http/httptest + desc: >- + Test-support code belongs in test files and in packages whose + directory name ends in test, not in the shipped binary. + # Only decisions already recorded in the Go package defaults are + # listed here. Every entry matches the module path exactly. + gomodguard_v2: + blocked: + - module: github.com/rs/zerolog + recommendations: + - log/slog + reason: "Structured logging is stdlib log/slog." + # One entry per pre-fork module path, because the later releases + # are separate paths. A prefix match would be shorter but would + # also reach github.com/go-redis/redismock, the test double for + # the successor these entries recommend. + - module: github.com/go-redis/redis + recommendations: + - github.com/redis/go-redis/v9 + reason: "Pre-fork module; use the maintained go-redis v9." + - module: github.com/go-redis/redis/v7 + recommendations: + - github.com/redis/go-redis/v9 + reason: "Pre-fork module; use the maintained go-redis v9." + - module: github.com/go-redis/redis/v8 + recommendations: + - github.com/redis/go-redis/v9 + reason: "Pre-fork module; use the maintained go-redis v9." + - module: github.com/sergi/go-diff + recommendations: + - github.com/aymanbagabas/go-udiff + reason: "No unified diff output; use go-udiff." + - module: github.com/hexops/gotextdiff + recommendations: + - github.com/aymanbagabas/go-udiff + reason: "Unmaintained fork; use go-udiff." + +issues: + max-issues-per-linter: 0 + max-same-issues: 0 diff --git a/.prettierignore b/.prettierignore new file mode 100644 index 0000000..23d67fc --- /dev/null +++ b/.prettierignore @@ -0,0 +1,2 @@ +node_modules/ +yarn.lock diff --git a/.prettierrc b/.prettierrc new file mode 100644 index 0000000..8af31cd --- /dev/null +++ b/.prettierrc @@ -0,0 +1,4 @@ +{ + "tabWidth": 4, + "proseWrap": "always" +} diff --git a/Dockerfile b/Dockerfile new file mode 100644 index 0000000..5267b36 --- /dev/null +++ b/Dockerfile @@ -0,0 +1,182 @@ +# Lint phase. The linter is invoked directly rather than through `make +# lint` or `script/lint`, which are themselves a docker build and would +# recurse into a daemon that does not exist in a build step. +# +# golangci/golangci-lint v2.14.0 (built with go1.27.0), 2026-09-24 +FROM golangci/golangci-lint@sha256:ad862ba6b3798cbe0fd9fd7408d498fd74fbd2623a92406b2fd3898faf0bf98f AS lint + +WORKDIR /src + +COPY go.mod go.sum ./ +RUN go mod download + +COPY . . + +RUN golangci-lint run --config .golangci.yml ./... + +# Test phase, same shape and for the same reason. The go directive in +# go.mod is a minimum, so this Go may be newer than the linter's. The +# Debian image rather than the Alpine one, because the race detector +# needs the C compiler it carries. +# +# golang 1.27.1-trixie, 2026-09-19 +FROM golang@sha256:3b77fc618ec235a1ab412de7737f120dd507c57e8d87de4cbb7994fb94275ed5 AS test + +WORKDIR /src + +COPY go.mod go.sum ./ +RUN go mod download + +COPY . . + +# Go's build cache is kept on a tmpfs, out of the image: nothing uses it +# after this step, and writing it into the image takes seconds. +RUN --mount=type=tmpfs,target=/root/.cache/go-build \ + go test -timeout 90s -race -cover ./... || \ + { echo "--- Rerunning with -v for details ---"; \ + go test -timeout 90s -race -v ./...; exit 1; } + +# Build stage. Nothing is wanted from the two phases above; the copies +# are what make BuildKit build them first, so the image, which needs this +# stage, cannot be produced unless lint and test passed. +# +# golang 1.27.1-trixie, 2026-09-19 +FROM golang@sha256:3b77fc618ec235a1ab412de7737f120dd507c57e8d87de4cbb7994fb94275ed5 AS builder + +COPY --from=lint /src/go.sum /dev/null +COPY --from=test /src/go.sum /dev/null + +# This image has git. A tar-stream context keeps the sender's file +# owners, which git refuses. +RUN git config --system --add safe.directory /src + +WORKDIR /src + +COPY go.mod go.sum ./ +RUN go mod download + +COPY . . + +# The VERSION build arg when one is given, otherwise +# `git describe --tags --always` on the .git in the build context. With +# .git present, a version that is still empty, dev or unknown fails the +# build: git is missing or could not read the checkout. +ARG VERSION +RUN VERSION="${VERSION:-$(git describe --tags --always)}"; \ + if [ -e .git ]; then \ + case "$VERSION" in ""|dev|unknown) \ + echo "version is '$VERSION' although .git is present" >&2; \ + exit 1 ;; \ + esac; \ + fi; \ + CGO_ENABLED=0 go build -trimpath \ + -ldflags="-s -w -X main.Version=${VERSION}" \ + -o /usr/local/bin/smallwebwaf ./cmd/smallwebwaf + +# runsvinit, the image's entrypoint, built at the last commit of its +# archived repository. It has no go.mod, and `go build` of its directory +# needs one; it uses only the standard library, so the one written here +# names nothing else. +# +# golang 1.27.1-trixie, 2026-09-19 +FROM golang@sha256:3b77fc618ec235a1ab412de7737f120dd507c57e8d87de4cbb7994fb94275ed5 AS runsvinit + +RUN git clone --quiet https://github.com/peterbourgon/runsvinit /src +WORKDIR /src +# runsvinit v2.0.0-8-gb4b2c78, 2015-10-07 +RUN git checkout --quiet --detach b4b2c785308b1ce785b6155c7fe5f16879080193 \ + && go mod init github.com/peterbourgon/runsvinit \ + && CGO_ENABLED=0 go build -trimpath -ldflags="-s -w" \ + -o /usr/local/bin/runsvinit . + +# The image an app's Dockerfile builds FROM, described under "Deployment" +# in SPEC.md. It is the last stage, so a plain `docker build .` builds it. +# +# ubuntu 26.04, 2026-09-27 +FROM ubuntu@sha256:f144425ff09be612d6d9ad965196e9cdc23dae1f42110a8a11a3e9a8198759f7 + +# runit's install creates its _runit-log user with minsysusers, which +# reads this file in place of runit's /usr/lib/sysusers.d/runit.conf. +# runit's line leaves out the shell, and minsysusers prints a Perl +# warning for that; this copy of it names /sbin/nologin, the shell +# minsysusers gives when none is named. +RUN mkdir /etc/sysusers.d \ + && echo 'u _runit-log - "runit svlogd user" /nonexistent /sbin/nologin' \ + > /etc/sysusers.d/runit.conf + +# ca-certificates, nix-bin and runit, from Ubuntu's archive as it was at +# the snapshot moment, which is never earlier than the Ubuntu image above. +# apt checks every package against the snapshot's InRelease files, and +# this step checks those against the hashes named here, which are those +# of the amd64 archive: other architectures use Ubuntu's ports archive. +# apt also fetches the live archive's InRelease files, which change daily +# and which the install does not use. The snapshot service is HTTPS only +# and this image has no CA certificates yet, so this step uses the Go +# image's. +RUN --mount=type=bind,from=builder,source=/etc/ssl/certs/ca-certificates.crt,target=/tmp/go-image-ca.crt \ + apt-get update --snapshot 20261001T000000Z \ + -o Acquire::https::CaInfo=/tmp/go-image-ca.crt \ + && printf '%s\n' \ + '45f95ce276cdba3e41870516a130e03c58b8b7a79e9546b0efe9e526d255740c snapshot.ubuntu.com_ubuntu_20261001T000000Z_dists_resolute_InRelease' \ + '802e675dd9de4c7f3916434a95e7c1d8eec0e82886622d7805ab19a2c6fe0365 snapshot.ubuntu.com_ubuntu_20261001T000000Z_dists_resolute-updates_InRelease' \ + '64b3353f0bd4970b4f7271962245bcea9ff24d4cc7bea16b433f8a60e42ca3dd snapshot.ubuntu.com_ubuntu_20261001T000000Z_dists_resolute-backports_InRelease' \ + '1d5041572116a8b23aabf79ac7439ad8af83d57ad3fb0f9aa0d4523ec10c5908 snapshot.ubuntu.com_ubuntu_20261001T000000Z_dists_resolute-security_InRelease' \ + | (cd /var/lib/apt/lists && sha256sum --check --strict) \ + && DEBIAN_FRONTEND=noninteractive apt-get install --yes --no-install-recommends \ + --snapshot 20261001T000000Z \ + -o Acquire::https::CaInfo=/tmp/go-image-ca.crt \ + ca-certificates nix-bin runit \ + && rm -rf /var/lib/apt/lists/* + +# Nix run by root expects a group of build users, which nix-bin does not +# create; with the setting empty, root's builds run without them. +RUN mkdir /etc/nix && echo 'build-users-group =' > /etc/nix/nix.conf + +# nixpkgs, from its release file, checked by SHA-256, and set up for root +# as `nixpkgs`, so that an app's Dockerfile installs a package with +# `nix-env -iA nixpkgs.`. curl and xz come with nix-bin. +# +# nixpkgs nixos-26.05.11045.774debe7a0d1, 2026-10-02 +RUN curl -fsSL -o /tmp/nixexprs.tar.xz \ + https://releases.nixos.org/nixos/26.05/nixos-26.05.11045.774debe7a0d1/nixexprs.tar.xz \ + && echo 'b2994104605601690023a5a6a3bb5a07b2bd1716b4e3b208cba1056dacd2ab08 /tmp/nixexprs.tar.xz' \ + | sha256sum --check --strict \ + && mkdir -p /root/.nix-defexpr/nixpkgs \ + && tar -xJf /tmp/nixexprs.tar.xz -C /root/.nix-defexpr/nixpkgs --strip-components=1 \ + && rm /tmp/nixexprs.tar.xz + +# What root installs with nix-env lands in root's profile. This path to +# it works for every user, unlike /root/.nix-profile: only root can +# enter /root. It comes last, so that no package shadows the image's +# own tools: busybox, for one, brings an sv that looks for services +# elsewhere. +ENV PATH=${PATH}:/nix/var/nix/profiles/default/bin + +COPY --from=runsvinit /usr/local/bin/runsvinit /usr/local/bin/runsvinit +COPY --from=builder /usr/local/bin/smallwebwaf /usr/local/bin/smallwebwaf + +# 65532 is above the uids Ubuntu keeps for system users, which end at +# 999; useradd warns about it unless --key raises that end for this call. +RUN groupadd --system --gid 65532 smallwebwaf \ + && useradd --system --key SYS_UID_MAX=65532 --uid 65532 \ + --gid smallwebwaf --no-create-home --shell /usr/sbin/nologin \ + smallwebwaf + +# The state files' directory, SWWAF_STATE_DIR by default, where a volume +# is mounted to keep them across deploys. The run script gives it to the +# smallwebwaf user at each start. +RUN mkdir /var/lib/smallwebwaf + +# runsvinit starts runit's runsvdir on /etc/service, where Ubuntu's sv +# looks too. +COPY --chmod=755 share/smallwebwaf.run /etc/service/smallwebwaf/run + +EXPOSE 8080 + +# traefik sends a container no requests until it is healthy, so the +# check runs every second from the start until it first passes, for up +# to a minute, and every 30 seconds after that. +HEALTHCHECK --start-period=1m --start-interval=1s \ + CMD ["/usr/local/bin/smallwebwaf", "healthcheck"] + +ENTRYPOINT ["/usr/local/bin/runsvinit"] diff --git a/EVALUATION.md b/EVALUATION.md index f671f13..4ea3fc9 100644 --- a/EVALUATION.md +++ b/EVALUATION.md @@ -17,8 +17,8 @@ configured only by environment variables, that provides: - R5: temporary blocks for abusers, permanent bans for repeat offenders - R6: one or more RBLs or IP reputation APIs - R7: AS number lookup -- R8: thresholds biased by AS number or country (for example, listed AS - numbers get 50 percent of the normal limit) +- R8: thresholds biased by AS number or country (for example, listed AS numbers + get 50 percent of the normal limit) - R9: WAF-style attack detection and prevention - R10: runs as a plain env-var-configured sidecar between traefik and one app @@ -37,19 +37,19 @@ Closest existing options, and why each still falls short: - BunkerWeb is the closest single product. - Meets: R1 (rates in requests per second, minute, hour or day), R2 - (`LIMIT_IGNORE_IP`, also by AS number and reverse DNS), R3 partly (webhook, - Slack, Discord, Matrix plugins, fired only on denied requests; ntfy only - through the generic webhook, payload format unverified), R5 partly - (`BAD_BEHAVIOR_BAN_TIME`, `0` means permanent; no escalation for repeat - offenders, the ban length is one fixed value), R6 (DNSBL plugin, external - blacklist URLs, optional CrowdSec), R7 partly (AS number used for + (`LIMIT_IGNORE_IP`, also by AS number and reverse DNS), R3 partly + (webhook, Slack, Discord, Matrix plugins, fired only on denied requests; + ntfy only through the generic webhook, payload format unverified), R5 + partly (`BAD_BEHAVIOR_BAN_TIME`, `0` means permanent; no escalation for + repeat offenders, the ban length is one fixed value), R6 (DNSBL plugin, + external blacklist URLs, optional CrowdSec), R7 partly (AS number used for blacklist and whitelist decisions), R9 (ModSecurity with the Core Rule Set, or Coraza plugin), env-var settings. - Fails: R8 (AS number and country can only allow or deny, never scale a limit), R4 (no volume or byte threshold alerts in the free edition; - reporting is a paid feature), R5 escalation, and R10 in spirit: since - 1.6 it needs a `bunkerweb` container plus a `bw-scheduler` container and - a database, or the all-in-one image that bundles nginx, scheduler, UI and + reporting is a paid feature), R5 escalation, and R10 in spirit: since 1.6 + it needs a `bunkerweb` container plus a `bw-scheduler` container and a + database, or the all-in-one image that bundles nginx, scheduler, UI and Redis in one container. It is designed to be the front door for many sites, not a per-app sidecar. AGPL-3.0. - Unverified: whether several rates (minute, hour, day) can be stacked on @@ -61,15 +61,14 @@ Closest existing options, and why each still falls short: (community blocklist, further blocklists, reputation API), R7 (alerts are enriched with AS number and country), R9 (AppSec component with virtual patching and ModSecurity-syntax rules), R2 (allowlists). - - Partly: R1 and R8. Detection is by leaky-bucket scenarios over logs, and - a scenario can filter on AS number or country, so a stricter bucket for + - Partly: R1 and R8. Detection is by leaky-bucket scenarios over logs, and a + scenario can filter on AS number or country, so a stricter bucket for listed AS numbers is possible, but each is a hand-written YAML scenario, it reacts after the fact by banning, and it is not an inline limiter that answers 429. - - Fails: R10 (needs the security engine container with persistent state, - log acquisition from traefik, a bouncer such as the traefik plugin, and - YAML for acquisition, profiles, scenarios and notifications), bytes half - of R4. + - Fails: R10 (needs the security engine container with persistent state, log + acquisition from traefik, a bouncer such as the traefik plugin, and YAML + for acquisition, profiles, scenarios and notifications), bytes half of R4. - CrowdSec plus traefik's own `rateLimit` middleware is the best combination with no new code. It gives inline limiting (one window per middleware, keyed by IP, with `sourceCriterion` exclusions), bans with escalation, alerts and @@ -89,19 +88,19 @@ Closest existing options, and why each still falls short: central API, or AppSec only. Supports captcha remediation. - Covers: R5, R6, R9, R3 and R7 through the engine (see Verdict). - Misses: the plugin itself does no rate limiting; R8; R4 bytes. - - Configuration: traefik static config to load the plugin, dynamic config - or labels for the middleware, CrowdSec YAML for everything else. + - Configuration: traefik static config to load the plugin, dynamic config or + labels for the middleware, CrowdSec YAML for everything else. - Sidecar fit: no. It lives inside traefik, plus a separate engine container. - Maturity: widely used (about 900 stars, listed in the traefik plugin catalog, documented by CrowdSec itself), actively maintained. Traefik - plugins run in an interpreter inside traefik, which costs some - per-request time. + plugins run in an interpreter inside traefik, which costs some per-request + time. - CrowdSec generic bouncers (nginx, Caddy, firewall) - Same engine, different enforcement point. The firewall bouncer blocks at - nftables level on the host, which is cheap and covers every service on - the host at once; worth considering fleet-wide regardless of this - project. Not a sidecar, same misses as above. + nftables level on the host, which is cheap and covers every service on the + host at once; worth considering fleet-wide regardless of this project. Not + a sidecar, same misses as above. - BunkerWeb 1.6.14 (`bunkerity/bunkerweb`) - What it is: nginx with Lua plugins, ModSecurity and the Core Rule Set, configured by settings that are passed as env vars to its scheduler @@ -123,23 +122,22 @@ Closest existing options, and why each still falls short: edition is capped at 10 applications. - Configuration: web UI backed by PostgreSQL. No env-var configuration. - Sidecar fit: no. Seven containers (postgres, management, detector, - tengine, and three helpers); the proxy container uses host networking. - The detection engine is closed source. + tengine, and three helpers); the proxy container uses host networking. The + detection engine is closed source. - Maturity: very active, vendor-driven. - Coraza (`corazawaf/coraza`) and Coraza-based proxies - - What it is: a Go library that implements the ModSecurity rule language - and runs the OWASP Core Rule Set. OWASP project, actively maintained, the + - What it is: a Go library that implements the ModSecurity rule language and + runs the OWASP Core Rule Set. OWASP project, actively maintained, the successor path now that ModSecurity is in maintenance only. - Packagings: `coraza-caddy` (Caddy module), `coraza-spoa` (HAProxy), `coraza-proxy-wasm` (Envoy), a traefik WASM plugin, and - `coreruleset/coraza-crs-docker` (Caddy plus Coraza plus the Core Rule - Set, with env vars for backend address, engine mode and rule-set - tuning). + `coreruleset/coraza-crs-docker` (Caddy plus Coraza plus the Core Rule Set, + with env vars for backend address, engine mode and rule-set tuning). - Covers: R9 only. `coraza-crs-docker` fits R10 well: one container, env vars, backend address. - - Misses: R1 to R8. ModSecurity-language rules can count requests per IP - in a persistent collection, but Coraza's support for persistent - collections is limited and this is not a practical rate limiter. + - Misses: R1 to R8. ModSecurity-language rules can count requests per IP in + a persistent collection, but Coraza's support for persistent collections + is limited and this is not a practical rate limiter. - Value here: the right library to embed for R9 in a purpose-built sidecar. - ModSecurity Core Rule Set containers (`owasp/modsecurity-crs`) - What it is: official images of Apache or nginx with ModSecurity and the @@ -149,21 +147,21 @@ Closest existing options, and why each still falls short: - Covers: R9, and R10 (single container, env vars, one backend). - Misses: R1 to R8. - Maturity: rule set is very actively maintained; the ModSecurity engine - itself is in maintenance under OWASP after Trustwave ended support in - 2024. + itself is in maintenance under OWASP after Trustwave ended support + in 2024. - Anubis (`TecharoHQ/anubis`), 1.27 current - What it is: a single-binary reverse proxy that makes browsers solve a proof-of-work challenge before passing them to `TARGET`. Aimed at scrapers, which is likely a large share of unwanted traffic on a public gitea. - Covers: R10 well (one container, env vars for listener, target, - difficulty, cookies). Policy rules can match on path, user agent, - headers, IP ranges, and with the vendor's hosted data service also AS - number and country, and can weigh a request toward a harder challenge. + difficulty, cookies). Policy rules can match on path, user agent, headers, + IP ranges, and with the vendor's hosted data service also AS number and + country, and can weigh a request toward a harder challenge. - Misses: R1, R3, R4, R5, R6, R9. Bot policy needs a YAML file, not env vars. AS number and country matching depend on the vendor's hosted - service. Breaks non-browser clients unless paths are exempted; for - gitea, git-over-HTTP and API paths must be allowed through by rule. + service. Breaks non-browser clients unless paths are exempted; for gitea, + git-over-HTTP and API paths must be allowed through by rule. - Maturity: very active, widely deployed on code forges since 2025. - Value here: complementary. It can be chained (traefik, then the sidecar, then Anubis, then the app) if challenge pages are wanted. @@ -187,29 +185,28 @@ Closest existing options, and why each still falls short: - Sidecar fit: no; they live inside traefik. - Traefik built-in middlewares - `rateLimit` (one average-and-burst window per middleware, optional Redis - in traefik 3.x), `inFlightReq`, `ipAllowList`. No bans, alerts, - reputation or AS number awareness. + in traefik 3.x), `inFlightReq`, `ipAllowList`. No bans, alerts, reputation + or AS number awareness. - caddy-waf (`fabriziosalmi/caddy-waf`), 0.4.x - What it is: a Caddy module with regex rules and anomaly scoring, per-IP and per-path rate limiting with one configurable window, IP and DNS - blacklists, Tor exit list fetch, country and AS number allow or deny - from MaxMind databases. + blacklists, Tor exit list fetch, country and AS number allow or deny from + MaxMind databases. - Misses: R8 (allow or deny only), R3, R4, R5 (no documented ban state or alerting), R1's three windows. Caddyfile configuration. One maintainer, pre-1.0, AGPL-3.0. - open-appsec (Check Point) - - Machine-learning WAF agent attached to nginx, Kong, Envoy or similar, - with a declarative policy file or the vendor's cloud console. Covers R9 - only; rate limiting and richer features are in paid tiers. Not a sidecar - in the required sense. + - Machine-learning WAF agent attached to nginx, Kong, Envoy or similar, with + a declarative policy file or the vendor's cloud console. Covers R9 only; + rate limiting and richer features are in paid tiers. Not a sidecar in the + required sense. - iocaine - - Serves generated garbage pages to clients the fronting proxy classifies - as scrapers. Not a limiter, WAF or ban tool; out of scope except as a + - Serves generated garbage pages to clients the fronting proxy classifies as + scrapers. Not a limiter, WAF or ban tool; out of scope except as a curiosity for scraper traffic. - Pangolin - - A tunnelled access platform that bundles traefik and optionally - CrowdSec. Replaces the ingress rather than adding a sidecar; out of - scope. + - A tunnelled access platform that bundles traefik and optionally CrowdSec. + Replaces the ingress rather than adding a sidecar; out of scope. ## Requirement by requirement, across the field diff --git a/LICENSE b/LICENSE new file mode 100644 index 0000000..34edefe --- /dev/null +++ b/LICENSE @@ -0,0 +1,21 @@ +MIT License + +Copyright (c) 2026 sneak + +Permission is hereby granted, free of charge, to any person obtaining a copy +of this software and associated documentation files (the "Software"), to deal +in the Software without restriction, including without limitation the rights +to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +copies of the Software, and to permit persons to whom the Software is +furnished to do so, subject to the following conditions: + +The above copyright notice and this permission notice shall be included in all +copies or substantial portions of the Software. + +THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE +SOFTWARE. diff --git a/Makefile b/Makefile new file mode 100644 index 0000000..36cf289 --- /dev/null +++ b/Makefile @@ -0,0 +1,43 @@ +.PHONY: bootstrap setup test lint fmt fmt-check check docker hooks build run \ + example-app + +# Makefile targets are thin shims; the implementations live in script/ +# per the scripts-to-rule-them-all pattern (see the Entrypoints section +# of README.md). build and run are for working on the code by hand; +# example-app checks the image with an app built on it. + +bootstrap: + @script/bootstrap + +setup: + @script/setup + +test: + @script/test + +lint: + @script/lint + +fmt: + @script/fmt + +fmt-check: + @script/fmt-check + +check: + @script/check + +docker: + @script/docker + +hooks: + @script/install-precommit + +build: + @script/build + +run: + @script/run + +example-app: + @script/example-app diff --git a/README.md b/README.md index 60f7bc6..d09e8b6 100644 --- a/README.md +++ b/README.md @@ -1,18 +1,481 @@ # smallwebwaf -`smallwebwaf` is a simple, fast, logging web application firewall for people who -host their own services. It runs inside the container of the one application it -protects, between your reverse proxy (traefik) and the app: the app's Dockerfile -builds `FROM` the `smallwebwaf` image, traefik sends the app's requests to +`smallwebwaf` is a simple, fast, logging web application firewall, MIT-licensed +and written in Go by [@sneak](https://sneak.berlin), for people who host their +own services. It runs inside the container of the one application it protects, +between your reverse proxy (traefik) and the app: the app's Dockerfile builds +`FROM` the `smallwebwaf` image, traefik sends the app's requests to `smallwebwaf` on port 8080, and `smallwebwaf` passes them on to the app on `127.0.0.1:8081`. It needs no setting, and protects the app from the first request with defaults chosen for a service on the open internet. It keeps its state in memory and in JSON files you can read and edit, and writes a detailed JSON log line for every request. -Status: design stage. This repository currently holds the documents only; no -code has been written. The design is in [`SPEC.md`](SPEC.md), and the survey of -existing tools that led to it is in [`EVALUATION.md`](EVALUATION.md). +Status: the first two milestones are built +(https://git.eeqj.de/sneak/smallwebwaf/issues/13 and +https://git.eeqj.de/sneak/smallwebwaf/issues/14), and so are seven parts of +milestone 3: the static lists, the bans that broken rate limits lead to and the +JSON state files with your edits taken in while it runs, which come next in the +build order, `observe` mode and the rest of the request log's fields, which come +a little later, and the metrics endpoint and the header size and the idle time +as settings, which come last in it. `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, bans a client that sends too many requests, refuses a +client that comes from a country you refuse or from a network you refuse, lets +the networks you choose through, keeps its bans, each client's counters and +history, and GeoJS's answers in JSON files across restarts, takes in your edits +of those files while it runs, writes a JSON log line for every request, serves +Prometheus metrics to a scraper that holds the metrics token, and in `observe` +mode passes on the requests it would refuse, logging what it would have done +with them. It comes as the image the app's own image is built on. The rest of +the design comes 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 + +Build the `smallwebwaf` image from a clone: + +```sh +git clone https://git.eeqj.de/sneak/smallwebwaf.git +cd smallwebwaf +make docker +``` + +`make docker` runs the tests and the linter, then builds the image, tagged +`smallwebwaf`, for amd64 and only on an amd64 host: the hashes the `Dockerfile` +checks Ubuntu's package lists against are those of Ubuntu's amd64 archive. Push +it to a registry your hosts pull from, and build each app's image on it, pinned +by digest, as "How it works, in short" below shows. `make example-app` builds a +small app on the image, the one in `deploy/example-app`, and checks that it +works. + +To work on the code, `make build` builds the binary alone, with Go installed, +and `make run` builds and runs it, listening on port 8080 in front of an app at +`SWWAF_UPSTREAM_URL`, by default `http://127.0.0.1:8081`, with its state files +in `bin/state` unless `SWWAF_STATE_DIR` is set. + +## What it does so far + +- Passes each request to the app and the app's answer back unchanged: method, + path, query, headers, body and status. Bodies stream through in both + directions and are never held whole in memory. A WebSocket, or any other + upgraded connection, passes through, and the timeouts do not cut it. +- Works out the client's address. A TCP peer outside `SWWAF_TRUSTED_PROXIES` is + the client, and the forwarded headers it sends are replaced, not passed on. + For a peer inside it, `X-Forwarded-For` is read from the right, and the first + address outside `SWWAF_TRUSTED_PROXIES` is the client; if every address in it + is inside, the leftmost is, and with no header the peer is. The app sees what + it would see from traefik directly: the same `Host`, the same + `X-Forwarded-Proto`, and `X-Forwarded-For` with the peer added at the end. It + also gets the request's id in `X-Request-ID`, the same id as in the request's + log line (see `request_id` in "Request log" below). +- Enforces the timeouts and the size limits below. A limit passed before the + response has started gets `smallwebwaf`'s own answer: `408` for a client too + slow to send its request, `413` for a request body that is too large, `504` + for an app too slow to answer, and `502` for a response that is too large or + an app that cannot be reached. A request that announces a body over the limit + is refused before anything reaches the app. While a request body is still on + its way, a request timeout that runs out answers `408` if `smallwebwaf` was + waiting for the client to send more, and `504` if it was waiting for the app + to take what it had. Once the response has started, a limit can only cut the + connection. +- Counts each client's requests over a minute, an hour and a day. A request that + takes the client over one of the rate limits below is refused with + `SWWAF_BAN_RESPONSE`, `403` by default, before anything reaches the app, and + bans the client. A client is one IPv4 address, or one IPv6 /64, since one + abuser usually holds a whole /64. Each window is counted in two fixed buckets, + the earlier one weighted by how much of it the window still covers. At most + 20,000 clients are kept, the least recently seen dropped first, with their + history, and a restart gives no client a fresh allowance (see "State files" + below). +- Bans a client that breaks a rate limit, as "Bans" in [`SPEC.md`](SPEC.md) + describes: the first ban lasts an hour, and a limit broken again within a day + of a ban ending bans for three times as long as that ban, so 1, 3, 9, 27 and + 81 hours; a ban that would last longer than seven days is permanent instead. A + ban covers the client's netblock: its IPv4 address, or the netblock around it + that `SWWAF_BAN_SCOPE_V4_PREFIX` sets, or its IPv6 /64. While it lasts, every + request from the netblock is refused with `SWWAF_BAN_RESPONSE` after the + static lists and before the country lists, so the client is not looked up, and + is not counted for the rate limits. A ban sets the client's counters back to + zero. Each ban carries notes for deciding whether to lift it: the limit, its + window and the requests counted in it, the request that broke it, the client's + country when it was looked up, the netblock's requests since it was first + seen, how many of them the ban has refused, and how many bans the netblock had + before. At most `SWWAF_MAX_BANS` bans are kept, past, active and permanent; + past that, the earliest ban of the netblock that has gone longest without a + request is dropped first. `bans.json` shows the bans and their notes, a + restart lifts none, and you add or lift a ban by editing it (see "State files" + below). +- Refuses a request from a country you refuse with `SWWAF_BAN_RESPONSE`, 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 is + never looked up: `SWWAF_EXCLUSIVELY_ALLOWED_COUNTRIES` refuses it unless it is + in `SWWAF_ALLOW_NETS`, and `SWWAF_DENIED_COUNTRIES` does not refuse it. +- Checks the client's own address against the static lists, the three netblock + settings below, before anything else, its country included. A client in + `SWWAF_ALLOW_NETS` skips bans, the country lists and the rate limits, and is + not looked up; the timeouts and size limits still apply. A client in + `SWWAF_DENY_NETS` is refused with `SWWAF_BAN_RESPONSE` before its body is + read, and the request is not counted for the rate limits; an address in + `SWWAF_ALLOW_NETS` too is let through. A client in + `SWWAF_RATE_LIMIT_EXEMPT_NETS` is neither counted nor refused by the rate + limits; the country lists and bans still apply to it. +- In `observe` mode, with `SWWAF_MODE=observe`, refuses none of the requests + that `SWWAF_DENY_NETS`, a ban, the country lists or a rate limit would refuse: + it passes them to the app, and their log lines name what `enforce` mode would + have done (see `would_action` in "Request log" below). The checks run, and + requests are counted, as in `enforce` mode, but a broken rate limit makes no + ban and does not set the client's counters back to zero, so each request over + the limit is logged as one that would be refused. The bans in `bans.json` are + kept, and refuse requests again when `smallwebwaf` next runs in `enforce` + mode, as long as they last. The timeouts and size limits still apply, since + they protect `smallwebwaf` and the app themselves, and a request for the + metrics without the token is still answered `401`. It is for trying a + configuration before enforcing it. +- Answers `GET /_smallwebwaf/healthz` itself with `200` and `ok`, before any + check and without asking the app, for the image's health check. +- Answers `GET /_smallwebwaf/metrics` with its metrics (see "Metrics" below) for + a request that carries `SWWAF_METRICS_TOKEN` as + `Authorization: Bearer `, and with `401` for one that does not. While + the token is unset the metrics answer `404`, as does any other request under + `/_smallwebwaf/`. Unlike the health check, such a request goes through every + check any other request goes through, and is answered where another would be + passed to the app: a banned client stays refused, and each counts toward the + client's rate limits. None of them reaches the app. +- Writes a line in the request log for each request (see "Request log" below). + +## Settings + +Each setting is an environment variable, and each has a default, so none has to +be set. A setting that is set but invalid stops the start with a message naming +it, and the effective settings are logged at start. + +- `SWWAF_LISTEN_ADDR` (default `:8080`): where `smallwebwaf` listens. +- `SWWAF_UPSTREAM_URL` (default `http://127.0.0.1:8081`): the app, as `http` or + `https`, a host and an optional port, and nothing more. +- `SWWAF_INSTANCE_NAME` (default: the host's name, which docker sets to the + first 12 characters of the container's id unless the deployment names one): + the name each request log line gives as `instance`. Set it, for example to + `fsn1app1/gitea`, for a name that stays the same when a deploy replaces the + container, and that tells instances apart when several log to one place. +- `SWWAF_MODE` (default `enforce`): `enforce`, or `observe` to pass on the + requests `smallwebwaf` would refuse and log what it would have done (see "What + it does so far" above). +- `SWWAF_TRUSTED_PROXIES` (default `10.0.0.0/8,172.16.0.0/12,192.168.0.0/16`, + the private address ranges): the netblocks whose `X-Forwarded-For` is + believed. A list given replaces the default; set but empty, it trusts nothing. +- `SWWAF_CLIENT_REQUEST_TIMEOUT` (default `60s`): how long a client may take to + send its request line and headers, and then, from the end of the headers, its + body. +- `SWWAF_CLIENT_REQUEST_HEADER_MAX_BYTES` (default `32K`): the largest request + line and headers a client may send. Over it, the answer is `431` and nothing + reaches the app. It must be more than `4K`, and cannot be `off`: Go's HTTP + server always has such a limit, and reads 4 KiB past the one it is given + before it refuses. +- `SWWAF_CLIENT_IDLE_TIMEOUT` (default `120s`): how long a kept-open connection + may wait for its next request before `smallwebwaf` closes it. The default 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. +- `SWWAF_CLIENT_RESPONSE_TIMEOUT` (default `30m`): how long the response may + take to reach the client, from the end of the request to the last byte. +- `SWWAF_UPSTREAM_REQUEST_TIMEOUT` (default `60s`): how long connecting to the + app and sending it the whole request may take. +- `SWWAF_UPSTREAM_RESPONSE_TIMEOUT` (default `30m`): how long the app may take + to send its whole answer, from the end of the request to the last byte. +- `SWWAF_REQUEST_MAX_BYTES` (default `100M`): the largest request body. +- `SWWAF_RESPONSE_MAX_BYTES` (default `5G`): the largest response body. +- `SWWAF_ALLOW_NETS` (default empty): netblocks whose clients skip bans, the + country lists and the rate limits, such as your monitoring or your own + networks. +- `SWWAF_RATE_LIMIT_EXEMPT_NETS` (default empty): netblocks whose clients the + rate limits do not apply to, such as a machine that talks to the app all day. +- `SWWAF_DENY_NETS` (default empty): netblocks whose clients are always refused. +- `SWWAF_RATE_LIMIT_PER_MINUTE` (default `1000`), `SWWAF_RATE_LIMIT_PER_HOUR` + (default `10000`) and `SWWAF_RATE_LIMIT_PER_DAY` (default `50000`): the most + 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. +- `SWWAF_BAN_RESPONSE` (default `403`): how a refused client is answered, one + that is banned, breaks a rate limit, is in `SWWAF_DENY_NETS` or comes from a + refused country: `403`, `429`, or `close` to close the connection without an + answer. Behind traefik, `close` does not leave the client unanswered: traefik + answers `502`, as it does whenever its backend drops a connection. +- `SWWAF_LIMIT_BAN_DURATION` (default `1h`): the ban for a first broken rate + limit. +- `SWWAF_LIMIT_BAN_REPEAT_WINDOW` (default `24h`): a rate limit broken again + within this time after a ban ended bans for three times as long as that ban. +- `SWWAF_MAX_BAN_DURATION` (default `7d`): a ban that would be longer is + permanent instead. +- `SWWAF_MAX_BANS` (default `5000`): the most bans kept, past, active and + permanent. +- `SWWAF_BAN_SCOPE_V4_PREFIX` (default `32`): the length of the netblock around + an IPv4 client that a ban covers, such as `24` to ban the surrounding /24. An + IPv6 ban covers the client's /64. +- `SWWAF_STATE_DIR` (default `/var/lib/smallwebwaf`): the directory of the state + files, an absolute path. A directory `smallwebwaf` cannot write stops the + start. +- `SWWAF_STATE_WRITE_DELAY` (default `10s`): how long after a ban is made + `bans.json` is written, with every ban made in between. +- `SWWAF_STATE_COUNTER_INTERVAL` (default `15m`): how often every state file is + written. +- `SWWAF_LOG_REQUEST_HEADERS` (default + `accept,accept-language,accept-encoding,content-type,origin,range`): the + request headers whose values the request log gives, in either case. + `Authorization`, `Cookie` and `Set-Cookie` are never logged, even when listed + (see "Request log" below). An entry naming `Host` or `Transfer-Encoding` stops + the start, since Go's HTTP server takes both out of the request; the request's + host is the field `host`. +- `SWWAF_METRICS_TOKEN` (default unset): the token a scraper sends for the + metrics, a long random value. While it is unset the metrics are off; one + shorter than 32 characters stops the start. The settings logged at start show + `********` in its place. +- `SWWAF_METRICS_TOP_N` (default `50`): how many countries get series of their + own in the metrics by country; the others are counted as `other`. + +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. 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; +`SWWAF_CLIENT_REQUEST_HEADER_MAX_BYTES`, the ban settings, the state settings +and `SWWAF_METRICS_TOP_N` cannot be off. + +Several limits are fixed rather than settings. At most 20,000 clients are kept, +with their counters and history, and an IPv6 client 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 + +`smallwebwaf` writes one JSON object per line on stdout for every request, +refused ones included: + +``` +{"type":"request","time":"2026-10-03T12:00:00.123Z","instance":"fsn1app1/gitea","client_ip":"203.0.113.9","method":"GET","scheme":"https","host":"app.example","path":"/","query":"","protocol":"HTTP/1.1","status":200,"request_bytes":0,"response_bytes":5120,"referer":"","user_agent":"curl/8.9.1","request_id":"7Q2NHZ4KJ3VXW5YB6R3MEFTD2A","peer_ip":"172.18.0.2","forwarded_for":"203.0.113.9","client_group":"203.0.113.9/32","country":"DE","request_headers":{"accept":"*/*"},"response_content_type":"text/html; charset=utf-8","upstream_status":200,"action":"forward","counts":{"minute":1,"hour":12,"day":40},"duration_total":3.217,"duration_checks":0.041,"duration_upstream_connect":0.052,"duration_upstream_first_byte":2.874,"duration_upstream_total":3.104} +``` + +A field that does not apply to a request is left out of its line, apart from +`type`, the fields from `time` to `user_agent`, `request_id`, `peer_ip`, +`client_group`, `country`, `action` and `duration_total`, which every line has. + +- `time` is when the request arrived, in UTC. `instance` is + `SWWAF_INSTANCE_NAME`. `scheme` is the `X-Forwarded-Proto` a trusted proxy + sent, and otherwise `http`. `path` and `query` are as the client sent them. +- `request_id` is the `X-Request-ID` a trusted proxy sent, or a new random one + of 26 letters and digits when it sent none, or when the peer is not a trusted + proxy. A request passed to the app takes it there in `X-Request-ID`. +- `peer_ip` is the TCP peer, normally traefik. `forwarded_for` is the + `X-Forwarded-For` header as received, several lines of it joined with `, `. + `client_group` is the client as the rate limits count it: its IPv4 address as + a /32, or the /64 of its IPv6 address. +- `country` is the client's country as GeoJS places it. It is empty with neither + country list set, for a client in `SWWAF_ALLOW_NETS` or `SWWAF_DENY_NETS`, for + a client on a private, loopback or link-local address, when GeoJS cannot place + the client or has not answered in time, and for a request whose client a ban + covers, even when the client's country is known. +- `content_type` is the request's `Content-Type`, and `content_length` the + length the request announced for its body, which is left out for none or zero. +- `request_headers` are the request's headers that `SWWAF_LOG_REQUEST_HEADERS` + names, by name in lower case, several lines of one joined with `, `. + `Authorization`, `Cookie` and `Set-Cookie` are never among them, whatever the + setting says: `has_authorization` and `has_cookie` are there instead, and + true, when the request has an `Authorization` or a `Cookie` header. +- `websocket` is there, and true, when the app switched the connection to + another protocol, as it does for a WebSocket. +- `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. +- `response_content_type`, `cache_control` and `location` are the + `Content-Type`, `Cache-Control` and `Location` headers of the answer: the + app's, as passed on, or those of `smallwebwaf`'s own answer. +- `request_bytes` and `response_bytes` count body bytes. +- `action` is `forward` for a request passed to the app, `denied` for one + refused because its client is in `SWWAF_DENY_NETS`, `banned` for one refused + because a ban covers its client, `country_denied` for one refused for its + client's country, `rate_limited` for one that broke a rate limit and banned + its client, `too_large` for a request or response over its size limit, + `timed_out` for one that ran out of time, `upstream_error` when the app could + not be reached or its answer broke off, and `admin` for one `smallwebwaf` + answered at its own endpoint. +- `would_action` is there in `observe` mode for a request that + `SWWAF_DENY_NETS`, a ban, the country lists or a rate limit would have refused + in `enforce` mode, and names the action that refusal would have had: `denied`, + `banned`, `country_denied` or `rate_limited`. `action` then names what was + done: `forward` for a request passed to the app, and another action, such as + `too_large`, for one a size or time limit refused. +- `counts` gives the client's requests in the minute, the hour and the day as + the rate limits count them, this request included: in each window, those in + the bucket under way and a share of those in the bucket before, so a count can + have a fraction. For a request that broke a limit, they are the counts that + broke it. It is left out for a request the rate limits do not count: the + health check, one from a client in `SWWAF_ALLOW_NETS` or + `SWWAF_RATE_LIMIT_EXEMPT_NETS`, and one that `SWWAF_DENY_NETS`, a ban or the + country lists refuse, or would refuse in `observe` mode. The byte totals come + with the byte limits. +- `limit_hit` is there for a request that broke a rate limit, and names the + window whose limit it went over: `minute`, `hour` or `day`, the shortest if it + went over several. `offence` is then `limit`. +- `ban_expires` is there for a request that made a ban or was refused under one, + or in `observe` mode would have been refused under one, and gives when the ban + ends, in the same form as `time`, or `permanent`. +- `aborted` is there, and true, when the client went away early. +- The timings are in milliseconds, to the microsecond. `duration_total` runs + from when the request's headers had been read to when its line is written, and + `duration_checks` over the same start to when the checks were done; the health + check runs none, and its line has no `duration_checks`. + `duration_upstream_connect`, `duration_upstream_first_byte` and + `duration_upstream_total` are there for a request passed to the app, and run + from when it was handed to the app: until there was a connection to it, new or + kept open from an earlier request, until the first byte of its answer arrived, + and until the end. The first two are left out when that never happened, as for + an app that cannot be reached. + +No body is logged, and no header but those above. `smallwebwaf`'s own messages +(start, the settings, stop, errors) share the stream as JSON lines marked +`"type":"process"`. + +Go's HTTP server, on which `smallwebwaf` is built, reads a request's line and +headers before `smallwebwaf` sees the request, and some requests end there, +without a line in the log: headers over `SWWAF_CLIENT_REQUEST_HEADER_MAX_BYTES`, +which it answers `431`, headers slower than `SWWAF_CLIENT_REQUEST_TIMEOUT`, +whose connection it closes without an answer, and requests it cannot read at +all, which it answers itself, mostly with `400`. + +## State files + +`smallwebwaf` keeps its state in memory and a copy of it in three JSON files in +`SWWAF_STATE_DIR`, `/var/lib/smallwebwaf` by default, as "Persistent state" in +[`SPEC.md`](SPEC.md) describes. Each has a top-level `version`, 1, and lists its +entries by client address, with times in UTC. + +- `bans.json`: every ban with its notes, indented to be read; a permanent ban's + `expires` is `null`. +- `clients.json`: each client's two buckets in the minute, the hour and the day, + and its history: when it was first and last seen, its country as last looked + up and when, its requests, how many were forwarded and how many refused (one + `smallwebwaf` answered at its own endpoints is neither, unless it was refused + with `401` for a missing or wrong token), the body bytes in each direction, + its responses by status class and its offences by kind. Each client is on a + line of its own, so `grep` shows everything about one. +- `lookups.json`: GeoJS's answers, one to a line, with when GeoJS gave each and + when it was last used. + +`bans.json` is written `SWWAF_STATE_WRITE_DELAY` after a ban is made, with every +ban made in between, and every file every `SWWAF_STATE_COUNTER_INTERVAL` and +when `smallwebwaf` stops. Each write goes to a temporary file in the same +directory, which then replaces the file, so a crash leaves the old file or the +new one, whole. A write that fails is logged, and tried again at the next write. +A hard kill loses what changed since the last write. + +At start the files are read back: each client keeps its counts, so a restart +gives it no fresh allowance, and each ban keeps refusing every client in its +netblock until it ends, even after `SWWAF_BAN_SCOPE_V4_PREFIX` has changed. A +netblock whose address has bits past its length, such as `203.0.113.9/24`, is +read as the netblock it is in, `203.0.113.0/24`. Buckets and answers whose time +has passed are dropped. A missing file is empty state, as on a first start. A +file that does not parse, or has another `version`, stops the start with a +message naming the file, and the line and column where Go's JSON decoder gives +them; so does a state directory `smallwebwaf` cannot write. So does an entry +without a field it needs, named with the entry's place in the file: a ban's +`netblock`, `start` or `expires`, which is `null` for a permanent ban; a +client's `client`, or the `start` of a window in which it has requests; an +answer's `client`, `country`, which is `""` for a client GeoJS cannot place, or +`answered`. The AS number and AS name come with their lookup. + +While it runs, `smallwebwaf` watches `SWWAF_STATE_DIR` and takes in your edit of +a state file as soon as you save it: what the file then holds replaces what +`smallwebwaf` held for it, as if read at start. It tells its own writes from +yours by comparing the file with what it last read or wrote, and before it +writes a file it takes in any edit made since, so your edit is not overwritten; +a change `smallwebwaf` made after you opened the file, such as a new ban, is +lost when you save over it. An edit that would stop the start, because it does +not parse, has another `version` or leaves out a field an entry needs, does not +stop the running `smallwebwaf`: it keeps what it holds, and at the file's next +write renames your file to `.bad`, such as `bans.json.bad`, writes the +file again from memory, and logs the file and where the error is. It waits for +that write because an editor's file can be read before the editor has finished +writing it. Mend the `.bad` file and move it back. A file you remove is written +again at its next write. + +To ban a netblock, add an entry to `bans.json` with its `netblock`, its `start` +and its `expires`, `null` for a ban that never ends; its `notes` may be left +out. This `bans.json` bans `203.0.113.0/24` for good: + +```json +{ + "version": 1, + "bans": [ + { + "netblock": "203.0.113.0/24", + "start": "2026-10-06T12:00:00Z", + "expires": null + } + ] +} +``` + +To lift a ban, delete its entry. `smallwebwaf` then forgets the ban, so it does +not make the netblock's next ban longer. + +## Metrics + +`GET /_smallwebwaf/metrics` answers with the metrics in the Prometheus text +format, for a scraper that sends `SWWAF_METRICS_TOKEN`, through traefik like any +other request. No metric carries a client's address. + +- `smallwebwaf_requests_total`, `smallwebwaf_request_bytes_total` and + `smallwebwaf_response_bytes_total`: requests, and their body bytes each way, + by `status_class`, such as `2xx`, or `none` when nothing was sent, and by + `action`, as the request log names it. +- `smallwebwaf_request_duration_seconds`: how long requests took, and + `smallwebwaf_upstream_duration_seconds`: how long those passed to the app took + from then on, as histograms; `smallwebwaf_requests_in_flight`: the requests + under way. +- `smallwebwaf_rate_limit_hits_total` by `window`, + `smallwebwaf_size_and_time_limit_hits_total` by `limit`, the setting whose + limit was passed, `smallwebwaf_offences_total` by `kind`, and + `smallwebwaf_bans_made_total` by `cause`; `smallwebwaf_active_bans` and + `smallwebwaf_permanent_bans`. +- `smallwebwaf_country_requests_total`, + `smallwebwaf_country_request_bytes_total`, + `smallwebwaf_country_response_bytes_total`, and + `smallwebwaf_country_list_refusals_total`, the requests the country lists + refused, by `country`, for the requests whose client's country is known. The + `SWWAF_METRICS_TOP_N` countries with the most requests since the start have + series of their own, and the others are counted as `other`. A country that + drops out of them loses its series, and its later requests count as `other`; + one that comes into them gets a series that counts from then on. +- `smallwebwaf_geojs_requests_total`: the requests to GeoJS; + `smallwebwaf_geojs_failures_total`: those that failed, an answer that leaves + out an address asked about included; and `smallwebwaf_geojs_unanswered_total`: + the requests whose client counted as coming from an unknown country because + GeoJS had not answered about it in time. +- `smallwebwaf_tracked_clients`: the clients in the table of clients. +- `smallwebwaf_state_file_writes_total`, + `smallwebwaf_state_file_write_failures_total`, + `smallwebwaf_state_file_last_write_timestamp_seconds` and + `smallwebwaf_state_file_size_bytes`, by `file`; and, by `file` too, + `smallwebwaf_state_file_edits_taken_in_total`: your edits taken in, and + `smallwebwaf_state_file_edits_set_aside_total`: those renamed to `.bad` + because they would stop the start. +- Go's own `go_` metrics and the process's `process_` metrics. + +The requests Go's HTTP server ends before `smallwebwaf` sees them (see "Request +log") are not counted. The metrics of the features still to come, such as the +rule files, come with them. ## Why @@ -114,7 +577,10 @@ goes through the candidates one by one. answers, the reputation cache, the alerting state) held in memory and kept in readable JSON files, written regularly and at every stop, so a restart loses nothing. Edit a file, or add a rule file, and the running `smallwebwaf` picks - up the change. Nothing is read from disk while serving a request. + up the change. Nothing is read from disk while serving a request. The files + for the bans, the clients and the GeoJS answers are built, with an edit taken + in while running (see "State files" above); the others come with their + features. - Health checks, the metrics, and listing, adding and lifting bans or asking why a given address was refused, all on the one port every request uses: under `/_smallwebwaf/` on the app's own address, through traefik like any other @@ -175,19 +641,24 @@ stand for the app's own options: ```bash #!/usr/bin/env bash set -euo pipefail -sleep 1 -exec chpst -u app:app /usr/local/bin/app \ - --listen 127.0.0.1:8081 \ - --trusted-proxies 10.0.0.0/8,172.16.0.0/12,192.168.0.0/16,127.0.0.1/32,::1/128 + +main() { + sleep 1 + exec chpst -u app:app /usr/local/bin/app \ + --listen 127.0.0.1:8081 \ + --trusted-proxies 10.0.0.0/8,172.16.0.0/12,192.168.0.0/16,127.0.0.1/32,::1/128 +} + +main "$@" ``` - The image's entrypoint, `runsvinit`, has runit start `smallwebwaf` and the app side by side, each as its own user, and start either again a second after it exits. Leave out `ENTRYPOINT` and `USER` from the app's Dockerfile. - `nix-env -iA nixpkgs.` installs a package from the nixpkgs in the image, - and the app finds it on its `PATH`. That nixpkgs is fixed at one commit, so - the same `smallwebwaf` image always gives the app the same packages; newer - ones come with a newer `smallwebwaf` image. + and the app finds it on its `PATH`, after Ubuntu's own commands. That nixpkgs + is fixed at one commit, so the same `smallwebwaf` image always gives the app + the same packages; newer ones come with a newer `smallwebwaf` image. - Deploy it as you deploy any app, with traefik's labels on this one container pointing at port 8080. upaas needs no change for this. - The app has to trust `127.0.0.1` and `::1` for forwarded headers, besides the @@ -197,10 +668,19 @@ exec chpst -u app:app /usr/local/bin/app \ - Port 8080 is the only one the app must leave free: the health check, the metrics and ban management are all on it, under `/_smallwebwaf/`. The image's health check passes while `smallwebwaf` answers and the app accepts - connections. + connections. `SWWAF_LISTEN_ADDR` can move `smallwebwaf` to another port, which + the app then leaves free instead; the health check follows it, and traefik's + labels must point at it. The address part of `SWWAF_LISTEN_ADDR` stays empty + (for example `:9000`, never `127.0.0.1:9000`), so `smallwebwaf` keeps + listening on every address: traefik reaches it on the container's address, and + the health check on `127.0.0.1`. - `smallwebwaf` keeps its state files in `/var/lib/smallwebwaf`. Mount a volume there to keep bans and client history when a deploy replaces the container; - without one, it still starts. + without one, it still starts. At each start the `run` script of `smallwebwaf` + gives that directory and every file in it to the `smallwebwaf` user, so a host + directory mounted there needs no change of owner. +- `docker stop` has runit stop both processes. `smallwebwaf` then stops taking + requests and gives those in progress five seconds to finish. A rule file is one rule per line: a name, what to match against, what to do, and a regex. @@ -216,41 +696,154 @@ 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. 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. +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, except a visitor in +`SWWAF_ALLOW_NETS` or `SWWAF_DENY_NETS` and one whose netblock a ban covers, 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 for seven days, in memory and in `lookups.json`, so that it survives a +restart, and many addresses are asked about in one request. 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 country: `SWWAF_EXCLUSIVELY_ALLOWED_COUNTRIES` refuses it unless you list it in -`SWWAF_ALLOW_NETS`. Such addresses are never sent to GeoJS. +`SWWAF_ALLOW_NETS`, and `SWWAF_DENIED_COUNTRIES` does not refuse it. Such +addresses are never sent to GeoJS. + +## How the code is laid out + +- `cmd/smallwebwaf`: the binary, which only calls `internal/smallwebwaf`. +- `internal/smallwebwaf`: the process: it reads the settings and the state + files, listens, serves requests until `SIGTERM` or `SIGINT`, and stops, + writing the state files. Run as `smallwebwaf healthcheck`, it is the image's + health check instead. +- `internal/config`: reads the settings, the one place they are read. +- `internal/proxy`: what happens to each request: it works out the client, runs + 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 `SWWAF_DENY_NETS`, for + a ban, for the country lists, for a rate limit, which bans the client, and for + an announced body over the size limit; in `observe` mode, only for the size + limit, with what it would have refused for noted in the log line. A request + under `/_smallwebwaf/` that `check` lets through is answered by `answerAdmin` + instead of reaching the app. +- `internal/metrics`: the metrics, counted as the other parts tell it what + happened, and served in the Prometheus text format. +- `internal/bans`: the ban ledger: each netblock's bans with their notes, how + long a new ban lasts, and which ban is dropped when `SWWAF_MAX_BANS` are held. +- `internal/lookup`: looks up each client's country through GeoJS, and keeps the + answers. +- `internal/ratelimit`: the table of clients: counts each client's requests, + tells when one takes it over a rate limit, and keeps each client's history. +- `internal/state`: reads the state files at start, takes in an admin's edit of + one while running, and writes them when they are due and at the stop. +- `internal/requestlog`: the lines on stdout: the request log line and the + process's own messages. +- `Dockerfile`: the lint and test phases, then the image, whose last stage + installs Ubuntu's packages, nixpkgs, `runsvinit` and `smallwebwaf`, with + `share/smallwebwaf.run` as runit's `run` script for `smallwebwaf`. +- `deploy/example-app`: an app built on the image, which `script/example-app` + checks. + +Besides the Go standard library, `github.com/hashicorp/golang-lru/v2` keeps the +table of clients to 20,000, the GeoJS answers to 100,000 and the banned +netblocks to `SWWAF_MAX_BANS`, dropping the least recently seen, and +`github.com/prometheus/client_golang` keeps the metrics and serves them, and +`github.com/fsnotify/fsnotify` tells `smallwebwaf` when a state file is saved. +The country codes are the list in `internal/config/config.go`. + +## Entrypoints + +This repository adheres to the +[Scripts to Rule Them All](https://github.com/github/scripts-to-rule-them-all) +standard: the scripts in `script/` are the entrypoints for working on it, and +the `Makefile` targets are thin shims that call them. The scripts are POSIX sh, +so that they run in minimal containers. + +- `script/bootstrap`: installs what the other scripts need on the host: `make`, + `git`, `curl`, Go for `gofmt`, node and yarn, and prettier. +- `script/setup`: readies a fresh clone: runs `script/bootstrap`, then + `script/install-precommit`. +- `script/projectname`: prints the project's name, `smallwebwaf`, which + `script/docker` and the others tag their images with. +- `script/test`: runs the tests, as the `test` phase of the `Dockerfile`. +- `script/lint`: runs golangci-lint, as the `lint` phase of the `Dockerfile`. +- `script/fmt`: formats the Go code with `gofmt` and the Markdown with prettier. +- `script/fmt-check`: checks the formatting, and changes nothing. +- `script/check`: runs `script/test`, `script/lint` and `script/fmt-check`. +- `script/docker`: builds the image, whose build runs the tests and the linter + first. +- `script/cibuild`: what CI runs: `script/bootstrap`, `script/check`, then the + image build. +- `script/precommit`: run by the git pre-commit hook; runs `script/check`. +- `script/install-precommit`: installs that hook; `make hooks` runs it. +- `script/build`: builds `bin/smallwebwaf` on the host, with Go installed, for + working on the code by hand; `make build` runs it. +- `script/run`: builds `bin/smallwebwaf` with `script/build` and runs it, with + its state files in `bin/state` unless `SWWAF_STATE_DIR` is set; `make run` + runs it. +- `script/example-app`: builds the image and, on it, the example app in + `deploy/example-app`, runs it with a volume for the state files, and checks + that the health check passes, that a request reaches the app through + `smallwebwaf`, that a second request in a minute bans the client, that + `sv stop` and `docker stop` stop it in order, and that a new container on the + same volume still refuses the banned client; then removes the containers, the + volume and both images. It needs network access, for nixpkgs' binary cache, + and `script/check` does not run it; `make example-app` does. + +## TODO + +- The rest of milestone 3: exemptions; then the rest of the design, in the order + of the build order in [`SPEC.md`](SPEC.md). ## Documents - [`SPEC.md`](SPEC.md): the design. - [`EVALUATION.md`](EVALUATION.md): what already exists, what each tool covers and misses, and why none was adopted. +- [`REPO_POLICIES.md`](REPO_POLICIES.md): the policies this repository follows. + +## License + +MIT. See [`LICENSE`](LICENSE). + +## Author + +[@sneak](https://sneak.berlin) diff --git a/REPO_POLICIES.md b/REPO_POLICIES.md new file mode 100644 index 0000000..20382d1 --- /dev/null +++ b/REPO_POLICIES.md @@ -0,0 +1,679 @@ +--- +title: Repository Policies +last_modified: 2026-10-04 +--- + +This document covers repository structure, tooling, and workflow standards. Code +style conventions are in separate documents: + +- [Code Styleguide](https://git.eeqj.de/sneak/prompts/raw/branch/main/prompts/CODE_STYLEGUIDE.md) + (general, bash, Docker) +- [Go](https://git.eeqj.de/sneak/prompts/raw/branch/main/prompts/CODE_STYLEGUIDE_GO.md) +- [JavaScript](https://git.eeqj.de/sneak/prompts/raw/branch/main/prompts/CODE_STYLEGUIDE_JS.md) +- [Python](https://git.eeqj.de/sneak/prompts/raw/branch/main/prompts/CODE_STYLEGUIDE_PYTHON.md) +- [Go HTTP Server Conventions](https://git.eeqj.de/sneak/prompts/raw/branch/main/prompts/GO_HTTP_SERVER_CONVENTIONS.md) + +--- + +- Cross-project documentation (such as this file) must include + `last_modified: YYYY-MM-DD` in the YAML front matter so it can be kept in sync + with the authoritative source as policies evolve. + +- **ALL external references must be pinned by cryptographic hash.** This + includes Docker base images, Go modules, npm packages, GitHub Actions, and + anything else fetched from a remote source. Version tags (`@v4`, `@latest`, + `:3.21`, etc.) are server-mutable and therefore remote code execution + vulnerabilities. The ONLY acceptable way to reference an external dependency + is by its content hash (Docker `@sha256:...`, Go module hash in `go.sum`, npm + integrity hash in lockfile, GitHub Actions `@`). No exceptions. + This also means never `curl | bash` to install tools like pyenv, nvm, rustup, + etc. Instead, download a specific release archive from GitHub, verify its hash + (hardcoded in the Dockerfile or script), and only then install. Unverified + install scripts are arbitrary remote code execution. This is the single most + important rule in this document. Double-check every external reference in + every file before committing. There are zero exceptions to this rule. + +- Every repo with software must have a root `Makefile` with these targets: + `make bootstrap`, `make setup`, `make test`, `make lint`, `make fmt` (writes), + `make fmt-check` (read-only), `make check` (runs `test`, `lint`, `fmt-check`), + `make docker`, and `make hooks` (installs pre-commit hook). A model Makefile + is at `https://git.eeqj.de/sneak/prompts/raw/branch/main/Makefile`. + +- Repos follow the + [Scripts to Rule Them All](https://github.com/github/scripts-to-rule-them-all) + pattern: the implementation of each Makefile target lives in an executable + script in `script/` (`script/bootstrap`, `script/setup`, `script/test`, + `script/lint`, `script/fmt`, `script/fmt-check`, `script/check`, + `script/docker`), and the Makefile targets are thin shims that call them. The + scripts must be POSIX sh (`#!/bin/sh`, `set -eu`, no bashisms) so they run in + minimal containers (e.g. alpine images have no bash); locate the repo root + with `$(cd "$(dirname "$0")/.." && pwd -P)` and `cd` there before acting. From + the standard's canonical set we use `bootstrap`, `setup` (make the repo ready + for development after a fresh clone: runs `bootstrap`, then + `install-precommit`, plus any repo-specific initialization), `test`, and + `cibuild`. `script/bootstrap` installs all dependencies idempotently and + assumes nothing is present: base tools come from nix, apt, brew, or apk + (detected in that order; apt runs noninteractive). For node it uses the + installed node if present; otherwise it installs a PINNED node version via + nvm, first installing nvm itself if missing — from a hash-verified GitHub + release archive (never `curl | sh`), with bash installed as an explicit + prerequisite since nvm requires bash. yarn is then pinned via + `corepack prepare yarn@ --activate`. Never install "latest" or "lts"; + always exact versions. `script/cibuild` runs the CI build: it changes to the + repo root, runs `script/bootstrap`, runs `script/check`, and builds the image + with the version; the Gitea workflow calls it. **`script/cibuild` runs + `script/bootstrap` first**, because the workflow checks out the repo and runs + nothing else, while `script/fmt-check` runs the formatter on the host: on a + pristine checkout with nothing installed the run dies there, after the + containerised gates have passed. **The bootstrap alone is not enough**: + `script/bootstrap` installs node and yarn under nvm and leaves neither on the + `PATH` of the shell that called it, so a bare `yarn` still exits 127. The host + entrypoints that need yarn — `script/fmt` and `script/fmt-check` — therefore + source nvm for the pinned node version before invoking it, exactly as + `script/bootstrap`'s own install step does. A runner carrying nothing but + docker and git then gets through `script/check`. Four further scripts are our + own extensions to the standard: `script/check` runs `script/test`, + `script/lint` and `script/fmt-check`; `script/precommit` is what the git + pre-commit hook runs, and it calls `script/check`; `script/install-precommit` + installs the git pre-commit hook (the `make hooks` target shims to it); and + `script/projectname` (literally that filename) simply outputs the project's + name. Scripts that need the name call `script/projectname` — e.g. + `script/docker` assembles its image tag from it — so those scripts stay + byte-identical across all repos. Repo-type-specific pre-commit extras (e.g. + `go mod tidy` verification in Go repos) belong in `script/precommit`, not in + the hook itself. Model scripts are at + `https://git.eeqj.de/sneak/prompts/raw/branch/main/script/`. The README + must document the provided scripts in an **Entrypoints** section (see the + README requirements below). + +- Always use Makefile targets (`make fmt`, `make test`, `make lint`, etc.) + instead of invoking the underlying tools directly. The Makefile is the single + source of truth for how these operations are run. + +- The Makefile is authoritative documentation for how the repo is used. Beyond + the required targets above, it should have targets for every common operation: + running a local development server (`make run`, `make dev`), re-initializing + or migrating the database (`make db-reset`, `make migrate`), building + artifacts (`make build`), generating code, seeding data, or anything else a + developer would do regularly. If someone checks out the repo and types + `make`, they should see every meaningful operation available. A new + contributor should be able to understand the entire development workflow by + reading the Makefile. + +- Every repo should have a `Dockerfile`, and it carries the repo's gates: a + `lint` phase and a `test` phase, with the final stage depending on both so the + image cannot be built unless they pass. For non-server repos the final stage + brings up a development environment; for server repos it is the runtime image. + The gate phases and the build stage start from their pinned base images and + install what those images lack either inline, as the canonical Go `Dockerfile` + below does for `git`, or by running `script/bootstrap`, as the `prompts` + repo's own `Dockerfile` does for its yarn packages. The development + environment stage installs development prerequisites by running + `script/bootstrap` rather than duplicating its installs inline. A stage that + runs `script/bootstrap` COPYs `script/` and the dependency manifests + (`package.json` + `yarn.lock`, `go.mod` + `go.sum`, etc.) before running it. + +- **Linting and testing run in Docker, as phases of the `Dockerfile`.** There is + no separate lint file. `script/lint` and `script/test` each build one phase + and nothing else: + + ```sh + docker build --no-cache --target lint -t "$(script/projectname)-lint" . + docker build --no-cache --target test -t "$(script/projectname)-test" . + ``` + + **A stage that is not the last one in the file is built only when the final + stage's chain depends on it, or when `--target` names it.** That is why the + two gates are always invoked by name here, and why the final stage carries a + `COPY --from=` of a harmless file from each of them: without that edge a + plain `docker build .` builds the last stage alone and exits 0 having linted + and tested nothing. + + **Every `docker build` in `script/` is tagged**, here and in + `script/cibuild` and `script/docker`. An untagged build leaves a dangling + image behind on every invocation, on every developer host and every CI + runner; a tagged one replaces the previous image. + + Inside a phase the tool is invoked directly — `golangci-lint`, `go test`, + `eslint`, `prettier` — never through `make lint` or `script/test`, which are + themselves a `docker build` and would recurse into a daemon that does not + exist in a build step. Formatting is the exception and stays on the host: + `script/fmt` writes the working tree, and `script/fmt-check` is its + read-only twin. + + **No lint verdict may come from a host invocation of the linter.** On a + shared host golangci-lint reads a result cache keyed on file content rather + than location, so a second checkout of the same content is served the first + one's findings, and a host-global lock in `$TMPDIR` makes concurrent runs + exit non-zero with `parallel golangci-lint is running` — a status a caller + cannot tell from real findings. Both have produced wrong verdicts in this + org, in both directions. A container has its own cache, its own `TMPDIR` and + a digest-pinned binary, so neither is reachable. + +- **Any build that runs checks is built with `--no-cache`.** Docker invalidates + a `COPY` layer only when the copied content changes, so on an unchanged tree + the check `RUN` is served from cache, nothing executes, and the build still + exits 0. Every `docker build` in `script/` therefore passes `--no-cache`: + `script/lint`, `script/test`, `script/cibuild` and `script/docker` are the + four, and there is no fifth — `script/check` runs the two gate phases and + `script/fmt-check`, and builds no image of its own. A bare `docker build .` is + not evidence that anything ran: a sub-second build reporting success is a + cache hit, not a result. Never invalidate by pruning — `docker builder prune` + and friends destroy a build cache shared with every other build on the host. + When a check is added or changed, prove it works by planting a defect it must + catch and watching the run fail on it, then revert the defect. A green run + alone shows neither that the check ran nor that it covers what it should. + +- **The gate phases are separate stages, and the build stage depends on both.** + The lint phase is based on the `golangci/golangci-lint` image (pinned by + hash), so lint failures surface in seconds rather than after a full compile, + and the test phase is based on the Debian Go image. The canonical Go repo + `Dockerfile`: + + ```dockerfile + # Lint phase + # golangci/golangci-lint:v2.x.x, YYYY-MM-DD + FROM golangci/golangci-lint@sha256:... AS lint + WORKDIR /src + COPY go.mod go.sum ./ + RUN go mod download + COPY . . + RUN golangci-lint run --config .golangci.yml ./... + + # Test phase. -race needs cgo and so a C compiler, which the Debian Go + # image ships and the alpine one does not. + # golang:1.x, YYYY-MM-DD + FROM golang@sha256:... AS test + WORKDIR /src + COPY go.mod go.sum ./ + RUN go mod download + COPY . . + RUN go test -timeout 90s -race -cover ./... || \ + { echo "--- Rerunning with -v for details ---"; \ + go test -timeout 90s -race -v ./...; exit 1; } + + # Build stage. Nothing is wanted from either phase above; the copies + # are what make BuildKit build them first, so this stage cannot run + # unless lint and test passed. + # golang:1.x-alpine, YYYY-MM-DD + FROM golang@sha256:... AS builder + COPY --from=lint /src/go.sum /dev/null + COPY --from=test /src/go.sum /dev/null + RUN apk add --no-cache git + # A tar-stream context keeps the sender's file owners, which git refuses. + RUN git config --system --add safe.directory /src + WORKDIR /src + COPY go.mod go.sum ./ + RUN go mod download + COPY . . + + # The VERSION build arg when one is given, otherwise + # `git describe --tags --always` on the .git in the build context. With + # .git present, a version that is still empty, dev or unknown fails the + # build: git is missing or could not read the checkout. + ARG VERSION + RUN VERSION="${VERSION:-$(git describe --tags --always)}"; \ + if [ -e .git ]; then \ + case "$VERSION" in ""|dev|unknown) \ + echo "version is '$VERSION' although .git is present" >&2; \ + exit 1 ;; \ + esac; \ + fi; \ + CGO_ENABLED=0 go build -trimpath \ + -ldflags="-s -w -X main.Version=${VERSION}" \ + -o /app ./cmd/app/ + + # Runtime stage, and the last one + FROM alpine@sha256:... + COPY --from=builder /app /usr/local/bin/app + ENTRYPOINT ["app"] + ``` + + Key points: + - The lint phase uses the `golangci/golangci-lint` image directly (it has + both Go and the linter), so nothing needs installing. + - `COPY --from= /src/go.sum /dev/null` is a no-op copy whose only + purpose is the ordering edge. BuildKit runs stages in parallel by default, + and a stage nothing depends on is not built at all, so without these two + lines a red gate would not fail the build. + - Keep the runtime stage last, and if you add a stage after it, give it the + same two copies. A plain `docker build .` builds the last stage's chain + and nothing else. + - If the project uses `//go:embed` directives that reference build artifacts + (e.g. a web frontend compiled in a separate stage), the lint phase must + create placeholder files so the embed directives resolve. Example: + `RUN mkdir -p web/dist && touch web/dist/index.html web/dist/style.css`. + - If the project requires CGO or system libraries for linting, install them + in the lint phase. The `golangci/golangci-lint` image is Debian-based and + has no `apk`, so install with `apt-get` under the Debian package name + (`libvips-dev`, where alpine says `vips-dev`), and delete the package + lists in the same `RUN`, so the layer does not keep them: + + ```dockerfile + RUN apt-get update \ + && apt-get install -y --no-install-recommends libvips-dev \ + && rm -rf /var/lib/apt/lists/* + ``` + + - `.dockerignore` lets `.git` into the build context. It keeps out every git + `config` at any depth (`**/.git/config`, `**/.git/modules/**/config`): the + repository's own, each submodule's under `.git/modules/`, and that of a + submodule keeping its own `.git` directory. `git describe` does not need + them, and each can hold a credential: a password in a remote URL, or the + token the CI checkout step stores there. A submodule whose name has a + `config` segment (`config`, `deploy/config`, `config/lib`) loses its whole + git directory to `**/.git/modules/**/config`, and Go's version stamping + then fails the build: give it a name without that segment + (`git submodule add --name`). The stage that compiles has `git` (the + Debian Go image has it; an alpine one needs `apk add --no-cache git`) and + takes the version from the `VERSION` build argument when one is given, + otherwise from `git describe --tags --always`. That gives the tag on a + tagged commit; on a later commit, the tag, the number of commits since it + and the short commit (`v1.2.3-4-gabc1234`); and the short commit when no + tag is reachable. The stage that compiles also marks its working directory + safe for git (`git config --system --add safe.directory /src`): a context + sent as a tar stream keeps the sender's file owners, and git refuses a + checkout owned by another user, so the version would come out empty. + `ARG VERSION` has no default, and the build fails if the context carries + `.git` and the version still comes out empty, `dev` or `unknown`. A plain + `docker build .` with no build arguments must succeed; a Dockerfile that + refuses an empty build argument drops that refusal and keeps the argument. + +- Every repo should have a Gitea Actions workflow (`.gitea/workflows/`) that + runs `script/cibuild` on push, and checks out the repo as its only other step. + That script bootstraps, runs the gate phases, and then builds the image, so a + successful run means every check passed; a bare `docker build .` does not + carry the same guarantee, because its gate phases may come from the cache. The + image build is uncached and so runs the gate phases a second time. That is the + price of the rule above, and it is worth paying: the image that ships is built + from a run of its own gates rather than from a cache entry. A separate + workflow limited to `main` by a `branches` list under `on: push` cannot be + checked by review: to try a change to it, add the feature branch to that list + and push, then remove the branch from the list again before merging. Keep any + job in it that publishes behind `if: github.ref_name == 'main'`, so the run + from the feature branch publishes nothing. + +- Use platform-standard formatters: `black` for Python, `prettier` for + JS/CSS/Markdown/HTML, `go fmt` for Go. Always use default configuration with + two exceptions: four-space indents (except Go), and `proseWrap: always` for + Markdown (hard-wrap at 80 columns). Documentation and writing repos (Markdown, + HTML, CSS) should also have `.prettierrc` and `.prettierignore`. + +- Pre-commit hook: runs `script/precommit`, which calls `script/check`. If local + testing is not possible in the repo, `script/precommit` may skip `script/test` + and run only `script/lint` and `script/fmt-check`. The hook is installed by + `script/install-precommit`; the Makefile must provide a `make hooks` target + that shims to it. + +- All repos with software must have tests that run via the platform-standard + test framework (`go test`, `pytest`, `jest`/`vitest`, etc.). If no meaningful + tests exist yet, add the most minimal test possible — e.g. importing the + module under test to verify it compiles/parses. There is no excuse for + `make test` to be a no-op. + +- `make test` must complete in under 60 seconds. That is the hard cap, and a + suite that exceeds it fails. Under 20 seconds is the target. A suite between + 20 and 60 seconds is still green, but the overage must be filed as an + improvement bug against that repo. Add a 90-second timeout to the test + invocation (`go test -timeout 90s`). The backstop deliberately sits above the + hard cap so that it catches a genuinely hung test rather than a merely slow + one. + +- **The test command should use the conditional verbose rerun pattern.** Run + tests without `-v` (verbose) first. If tests fail, automatically rerun with + `-v` to show full output. This keeps CI logs and `docker build` output clean + on success (just package/suite summaries) while providing full diagnostic + detail on failure (every test case, every assertion). The command lives in the + `test` phase of the `Dockerfile`, since `script/test` builds that phase; the + Makefile form below is the same pattern for any repo-local invocation: + + ```makefile + test: + @ || \ + { echo "--- Rerunning with -v for details ---"; \ + ; exit 1; } + ``` + + Go example: + + ```makefile + test: + @go test -count=1 -timeout 90s -race -cover ./... || \ + { echo "--- Rerunning with -v for details ---"; \ + go test -count=1 -timeout 90s -race -v ./...; exit 1; } + ``` + + `-count=1` is required on both invocations: it defeats Go's test _result_ + cache, so neither run can report a stored pass in place of running the + tests. It leaves the build cache alone, so it costs the runtime of the suite + and no recompilation. + + That cache is Go's own, separate from Docker's layer cache. Go stores a + passing result in its cache directory (`GOCACHE`), and when the same tests + run again on unchanged code it prints that result, marked `(cached)`, + without running them. That matters on a developer's machine, where this + target runs and the directory lasts from one run to the next. The `test` + phase of the `Dockerfile` needs no `-count=1`: its base image holds no + result for this repo's tests and nothing before its `go test` step runs a + test, so there is nothing to replay. `--no-cache` (above) is what makes that + step run on an unchanged tree. + + Python example: + + ```makefile + test: + @python -m pytest || \ + { echo "--- Rerunning with -v for details ---"; \ + python -m pytest -v; exit 1; } + ``` + + The `exit 1` ensures the target always fails after a rerun — the first run + already proved the tests are broken, so the build must not pass even if a + flaky test happens to succeed on the second attempt. The rerun exists solely + for diagnostic output. + +- Docker builds must complete in under 5 minutes. + +- `make check` must not modify any files in the repo. Tests may use temporary + directories. + +- `main` must always pass `make check`, no exceptions. + +- Never commit secrets. `.env` files, credentials, API keys, and private keys + must be in `.gitignore`. No exceptions. + +- `.gitignore` should be comprehensive from the start: OS files (`.DS_Store`), + editor files (`.swp`, `*~`), in-repo agent scratch directories (`.claude/`), + language build artifacts, and `node_modules/`. Fetch the standard `.gitignore` + from `https://git.eeqj.de/sneak/prompts/raw/branch/main/.gitignore` when + setting up a new repo. These patterns are written to `.gitignore`'s own + semantics, in which an unanchored pattern already matches at every depth; they + are not a `.dockerignore` and must not be transplanted into one unmodified. + +- **`.dockerignore` does not use `.gitignore` semantics, and copying patterns + across unmodified leaves secrets in the build context.** Docker matches with + `moby/patternmatcher`: `filepath.Match` semantics plus a `**` extension, so + `*` does not cross `/` and a pattern without a leading `**/` is anchored at + the build-context root. A `.dockerignore` listing `.env`, `*.pem` and `*.key` + therefore excludes only the copies at the repository root, while `config/.env` + and `certs/server.key` still reach the context and can land in an image layer + — which is more dangerous than a short file with no secret patterns at all, + because it reads as solved and stops anyone looking. Give every + depth-independent pattern the `**/` prefix and leave only genuinely + root-anchored entries unprefixed: `.claude`, and the repo's own host-built + binary, written `/myapp` and never `**/myapp`, which would also match + `cmd/myapp/` and delete the package directory from the context. Matching is + case-sensitive, and an ALL-CAPS twin per pattern still misses `Server.Key`, so + secret names use character ranges — `**/*.[kK][eE][yY]`, `**/*.[pP][eE][mM]`, + and likewise for `.envrc` and the extensionless SSH keys. Where such a pattern + also catches something the build needs, re-include it with a negation + (`!docs/example.env`); deleting the pattern reopens the exposure for every + other file it covers. Fetch the standard `.dockerignore` from + `https://git.eeqj.de/sneak/prompts/raw/branch/main/.dockerignore` and extend + it with the repo's own artifacts. + +- **In-repo agent scratch belongs in both files, written to each file's own + semantics.** `.claude/` holds one worktree per in-flight agent — an entire + additional checkout of the repo — so under `COPY . .` the build context + inflates by a multiple of the repo and another session's unreviewed work can + be copied into an image layer. In `.gitignore` the entry is `.claude/`, + unanchored. In `.dockerignore` it is `.claude`, anchored and with **no** `**/` + prefix, because the prefixed form would also delete any nested directory of + that name from the build. Anchoring carries a known gap that the canonical + `.dockerignore` states in its own comment, since consuming repos receive the + file and not the tracker: the directory is created in the agent's working + directory, so a repo running agents in subdirectories still ships + `services/api/.claude/` and must add its own anchored entry there. + +- **A plain `docker build .` of a clone stamps the version that + `git describe --tags --always` gives**, derived from the `.git` in the build + context as the canonical `Dockerfile` above shows. Without its failure check, + a missing `git` or an unreadable checkout would leave `-X main.Version=` empty + and the build would still exit 0. `script/docker` and `script/cibuild` pass + the version they compute on the host; it takes precedence. They do this + byte-identically across repos: + + ```sh + # Own line: a failing command substitution inside an argument does not + # trip `set -e`, so the inline form degrades to an empty constant. + version="$(git describe --tags --always --dirty 2>/dev/null || true)" + [ -n "$version" ] || version="unknown" + docker build --no-cache \ + --build-arg VERSION="$version" \ + -t "$(script/projectname)" . + ``` + + `--always` makes an untagged repo yield an abbreviated commit hash rather + than failing, and the `[ -n "$version" ]` line is the single place the + fallback is applied — a live check that fires on a build from an export with + no `.git` and on a repository with no commits yet. Do not fold it into the + substitution as `|| echo unknown`, which makes the guard unreachable. The + Dockerfile's side is `ARG VERSION` in the stage that compiles, declared + there because `ARG` is stage-scoped; passing `VERSION` to a repo whose + Dockerfile declares no such `ARG` is ignored and costs nothing, which is why + the scripts stay byte-identical. One consequence for CI: the standard + checkout action clones shallow and fetches no tags, so a repo that embeds a + tag-derived version must set `fetch-depth: 0` on its checkout step. + +- **Verify `.dockerignore` by enumerating the image, not by reading the + patterns.** Plant files at the root _and_ at least two directories deep, build + a probe image that does `COPY . .`, and list what actually landed + (`docker run --rm --entrypoint find IMAGE /app`). The `transferring context` + size is not a substitute: a nested secret is a few bytes, and BuildKit + transfers only the delta from the previous build. + +- **No build artifacts in version control.** Code-derived data (compiled + bundles, minified output, generated assets) must never be committed to the + repository if it can be avoided. The build process (e.g. Dockerfile, Makefile) + should generate these at build time. Notable exception: Go protobuf generated + files (`.pb.go`) ARE committed because repos need to work with `go get`, which + downloads code but does not execute code generation. + +- Never use `git add -A` or `git add .`. Always stage files explicitly by name. + +- Never force-push to `main`. + +- Make all changes on a feature branch. You can do whatever you want on a + feature branch. + +- `.golangci.yml` is standardized. The vendored copy in a consuming repo must + _NEVER_ be modified by an agent: fetch it from + `https://git.eeqj.de/sneak/prompts/raw/branch/main/.golangci.yml` and keep it + byte-identical, so that no repo can quietly loosen its own linting. Linter + configuration changes are made to the canonical copy in the `prompts` repo and + reach consuming repos by re-vendoring; an agent may open a PR against + canonical, which only the user merges. One list is exempt from byte-identity, + because it cannot be written once for every repo: the `deny` list of the + `test-support` depguard rule, where a repo names its own test-support packages + by full import path. A repo adds entries there and changes nothing else, and a + re-vendor carries its entries forward. The canonical golangci-lint version is + v2.14.0 (released 2026-09-24), pinned as the digest of the lint phase's base + image + (`golangci/golangci-lint@sha256:ad862ba6b3798cbe0fd9fd7408d498fd74fbd2623a92406b2fd3898faf0bf98f`, + which reports `2.14.0 built with go1.27.0 from 114493f9`). A module's `go` + directive must not name a newer Go minor version than the one golangci-lint + was built with, or golangci-lint refuses to lint it: this release lints + `go 1.27.1` but not `go 1.28`. That digest is the only pin, since no repo + installs golangci-lint on the host. A repo sets the lint phase digest to the + one named here and re-vendors `.golangci.yml` in the same commit, whichever of + the two prompted the change: the canonical copy can name linters that an older + golangci-lint rejects, and a newer golangci-lint can add linters that + `default: all` switches on until the canonical copy disables them. + +- **`script/bootstrap` installs a pinned tool by comparing versions, never by + testing presence.** An `if ! command -v ; then install; fi` guard tests + `PATH` only, so on an already-provisioned machine the pin is inert and a + version bump is a silent no-op — while the Dockerfile, installing into a clean + image, gets the pinned version, so a local `make check` and `make docker` can + disagree about what the tool even is. The canonical form: + - compares the installed version against the pin over the **whole** version + token; a parser that stops at the first `-` reports `2.12.2` for a host + running `2.12.2-rc1` and skips the install; + - treats absent, non-zero, empty or unrecognised `--version` output as a + mismatch, so the failure direction is a redundant install and never a + skipped one; + - after installing, re-resolves the binary the way callers do — `hash -r`, + then through `PATH`, not through the directory the installer wrote to — + and fails naming the resolved path, since an install that a shadowing + binary hides succeeds while changing nothing any caller sees; + - is actually called, and prints the version on both success paths: a + function defined and never invoked has the same exit status and the same + empty output as one that worked. + + Keep it POSIX sh: no arrays, no `[[`, no `grep -P`. + + A Go tool a repo needs on the host is installed with `go install` pinned to + a commit hash (`go install @`). It is never tracked as + a `go.mod` tool dependency or through a `tools.go` file, either of which + pulls the tool's own dependencies into the repo's `go.mod` and `go.sum`. + +- When pinning images or packages by hash, add a comment above the reference + with the version and date (YYYY-MM-DD). + +- Use `yarn`, not `npm`. + +- Write all dates as YYYY-MM-DD (ISO 8601). + +- Simple projects should be configured with environment variables. + +- Dockerized web services listen on port 8080 by default, overridable with + `PORT`. + +- **HTTP/web services must be hardened for production internet exposure before + tagging 1.0.** This means full compliance with security best practices + including, without limitation, all of the following: + - **Security headers** on every response: + - `Strict-Transport-Security` (HSTS) with `max-age` of at least one year + and `includeSubDomains`. + - `Content-Security-Policy` (CSP) with a restrictive default policy + (`default-src 'self'` as a baseline, tightened per-resource as + needed). Never use `unsafe-inline` or `unsafe-eval` unless + unavoidable, and document the reason. + - `X-Frame-Options: DENY` (or `SAMEORIGIN` if framing is required). + Prefer the `frame-ancestors` CSP directive as the primary control. + - `X-Content-Type-Options: nosniff`. + - `Referrer-Policy: strict-origin-when-cross-origin` (or stricter). + - `Permissions-Policy` restricting access to browser features the + application does not use (camera, microphone, geolocation, etc.). + - **Request and response limits:** + - Maximum request body size enforced on all endpoints (e.g. Go + `http.MaxBytesReader`). Choose a sane default per-route; never accept + unbounded input. + - Maximum response body size where applicable (e.g. paginated APIs). + - `ReadTimeout` and `ReadHeaderTimeout` on the `http.Server` to defend + against slowloris attacks. + - `WriteTimeout` on the `http.Server`. + - `IdleTimeout` on the `http.Server`. + - Per-handler execution time limits via `context.WithTimeout` or + chi/stdlib `middleware.Timeout`. + - **Authentication and session security:** + - Rate limiting on password-based authentication endpoints. API keys are + high-entropy and not susceptible to brute force, so they are exempt. + - CSRF tokens on all state-mutating HTML forms. API endpoints + authenticated via `Authorization` header (Bearer token, API key) are + exempt because the browser does not attach these automatically. + - Passwords stored using bcrypt, scrypt, or argon2 — never plain-text, + MD5, or SHA. + - Session cookies set with `HttpOnly`, `Secure`, and `SameSite=Lax` (or + `Strict`) attributes. + - **Reverse proxy awareness:** + - True client IP detection when behind a reverse proxy + (`X-Forwarded-For`, `X-Real-IP`). The application must accept + forwarded headers only from a configured set of trusted proxy + addresses — never trust `X-Forwarded-For` unconditionally. + - **CORS:** + - Authenticated endpoints must restrict `Access-Control-Allow-Origin` to + an explicit allowlist of known origins. Wildcard (`*`) is acceptable + only for public, unauthenticated read-only APIs. + - **Error handling:** + - Internal errors must never leak stack traces, SQL queries, file paths, + or other implementation details to the client. Return generic error + messages in production; detailed errors only when `DEBUG` is enabled. + - **TLS:** + - Services never terminate TLS directly. They are always deployed behind + a TLS-terminating reverse proxy. The service itself listens on plain + HTTP. However, HSTS headers and `Secure` cookie flags must still be + set by the application so that the browser enforces HTTPS end-to-end. + + This list is non-exhaustive. Apply defense-in-depth: if a standard security + hardening measure exists for HTTP services and is not listed here, it is + still expected. When in doubt, harden. + +- `README.md` is the primary documentation. Required sections: + - **Description**: First line must include the project name, purpose, + category (web server, SPA, CLI tool, etc.), license, and author. Example: + "µPaaS is an MIT-licensed Go web application by @sneak that receives + git-frontend webhooks and deploys applications via Docker in realtime." + - **Getting Started**: Copy-pasteable install/usage code block. + - **Entrypoints**: Opens by stating that the repo adheres to the + [Scripts to Rule Them All](https://github.com/github/scripts-to-rule-them-all) + standard (with that link), then documents each provided `script/` + entrypoint and its purpose. + - **Rationale**: Why does this exist? + - **Design**: How is the program structured? + - **TODO**: Update meticulously, even between commits. When planning, put + the todo list in the README so a new agent can pick up where the last one + left off. + - **License**: MIT, GPL, or WTFPL. Ask the user for new projects. Include a + `LICENSE` file in the repo root and a License section in the README. + - **Author**: [@sneak](https://sneak.berlin). + +- First commit of a new repo should contain only `README.md`. + +- Go module root: `sneak.berlin/go/`. Always run `go mod tidy` before + committing. + +- Use SemVer. + +- Database migrations live in `internal/db/migrations/` and must be embedded in + the binary. + - `000_migration.sql` — contains ONLY the creation of the migrations + tracking table itself. Nothing else. + - `001_schema.sql` — the full application schema. + - **Pre-1.0.0:** never add additional migration files (002, 003, etc.). + There is no installed base to migrate. Edit `001_schema.sql` directly. + - **Post-1.0.0:** add new numbered migration files for each schema change. + Never edit existing migrations after release. + +- All repos should have an `.editorconfig` enforcing the project's indentation + settings. + +- Avoid putting files in the repo root unless necessary. Root should contain + only project-level config files (`README.md`, `AGENTS.md`, `Makefile`, + `Dockerfile`, `LICENSE`, `.gitignore`, `.editorconfig`, `REPO_POLICIES.md`, + and language-specific config). Everything else goes in a subdirectory. + Canonical subdirectory names: + - `bin/` — executable scripts and tools + - `cmd/` — Go command entrypoints; thin only: one `main.go` per binary whose + body is a single call into `internal/` or `pkg/`, no project logic in + `cmd/` + - `configs/` — configuration templates and examples + - `deploy/` — deployment manifests (k8s, compose, terraform) + - `docs/` — documentation and markdown (README.md stays in root) + - `internal/` — Go internal packages + - `internal/db/migrations/` — database migrations + - `pkg/` — Go library packages + - `share/` — systemd units, data files + - `static/` — static assets (images, fonts, etc.) + - `web/` — web frontend source + +- When setting up a new repo, files from the `prompts` repo may be used as + templates. Fetch them from + `https://git.eeqj.de/sneak/prompts/raw/branch/main/`. + +- New repos must contain at minimum: + - `README.md`, `.git`, `.gitignore`, `.editorconfig` + - `LICENSE`, `REPO_POLICIES.md` (copy from the `prompts` repo) + - `Makefile` + - `script/` entrypoints (`bootstrap`, `setup`, `projectname`, `test`, + `lint`, `fmt`, `fmt-check`, `check`, `docker`, `cibuild`, `precommit`, + `install-precommit`) + - `Dockerfile`, `.dockerignore` + - `.gitea/workflows/check.yml` + - Go: `go.mod`, `go.sum`, `.golangci.yml` + - JS: `package.json`, `yarn.lock`, `.prettierrc`, `.prettierignore` + - Python: `pyproject.toml` + +- Guidance for coding agents lives in one `AGENTS.md` at the repository root. It + is never committed under a file or directory named after one agent tool, such + as `CLAUDE.md` or `.claude/`, and never split into separate memory files. diff --git a/SPEC.md b/SPEC.md index 08a00de..9071387 100644 --- a/SPEC.md +++ b/SPEC.md @@ -1,8 +1,8 @@ # smallwebwaf SPEC (draft): protective reverse proxy for one app -Status: fourth draft, with the owner's rulings to date applied. Nothing has been -built yet. `EVALUATION.md` beside this file explains why no existing tool was -chosen. +Status: fourth draft, with the owner's rulings to date applied. Milestone 1 of +the build order is built. `EVALUATION.md` beside this file explains why no +existing tool was chosen. ## Purpose @@ -293,7 +293,8 @@ it. needs: an alert destination, an account key, a token. - Every setting's name starts with `SWWAF_`, since `smallwebwaf` shares its container, and so its environment variables, with the app it protects. -- Any limit or threshold can be switched off with the value `off`. +- Any limit or threshold can be switched off with the value `off`, except + `SWWAF_CLIENT_REQUEST_HEADER_MAX_BYTES`. - A list set to an empty value is an empty list, and replaces the default. - Every setting may instead be given as a file holding the value, named by the setting's name with `_FILE` added, such as `SWWAF_ADMIN_TOKEN_FILE`, for @@ -409,10 +410,13 @@ The settings, by group: Bodies stream straight through, so a request body reaches the app while the client is still sending it. - `SWWAF_CLIENT_REQUEST_TIMEOUT` (default `60s`): how long a client may take - to send its whole request, headers and body. + to send its request line and headers, and then, from the end of the + headers, its body. - `SWWAF_CLIENT_REQUEST_HEADER_MAX_BYTES` (default `32K`): the largest request line and headers a client may send. Over it, `smallwebwaf` answers - `431` and closes the connection, and nothing reaches the app. + `431` and closes the connection, and nothing reaches the app. It must be + more than `4K`, and cannot be `off`: Go's HTTP server always has such a + limit, and reads 4 KiB past the one it is given before it refuses. - `SWWAF_CLIENT_IDLE_TIMEOUT` (default `120s`): how long a kept-open connection may wait for its next request before `smallwebwaf` closes it. It is longer than the 90 seconds after which traefik, by default, closes a @@ -436,6 +440,19 @@ The settings, by group: an app that is too slow. A request that announces a body larger than its limit is refused before anything reaches the app. Once the response has started it can only be cut off, and the connection is closed. + - Since a request body streams through, each side can hold up the other: a + slow client slows the send to the app, and an app slow to take the body + slows the client's send. So while a request body is still on its way, a + request timeout that runs out, `SWWAF_CLIENT_REQUEST_TIMEOUT` or + `SWWAF_UPSTREAM_REQUEST_TIMEOUT`, answers `408` if `smallwebwaf` was + waiting for the client to send more at that moment, and `504` if it was + waiting for the app to take what it had. + - Go's HTTP server, on which `smallwebwaf` is built, reads a request's line + and headers before `smallwebwaf` sees the request. A client that takes + longer than `SWWAF_CLIENT_REQUEST_TIMEOUT` to send them gets no answer: + the server closes its connection. Headers over + `SWWAF_CLIENT_REQUEST_HEADER_MAX_BYTES` are answered `431` by the server + itself. Neither request gets a line in the request log. - A WebSocket connection leaves these limits behind once it is upgraded: it stays open until either side closes it. - Lookup of AS number and country (R7). On by default through GeoJS, which needs @@ -931,9 +948,9 @@ and the running `smallwebwaf` takes the edit in. - what was broken: the rule ids and target that matched, or the limit, its window, the count reached and the client's limit percentage with what set it; and any reputation sources that listed the client; - - the requests that caused the ban, up to the last ten: time, method, host, - path with its query string, status and user agent, each text cut to 256 - bytes; + - the request that caused the ban, the one that broke the limit or carried + the clear sign of attack: time, method, host, path with its query string, + status and user agent, each text cut to 256 bytes; - how many requests counted toward the ban, and the time span over which they came; - the netblock's total requests since it was first seen, and the requests @@ -944,13 +961,13 @@ and the running `smallwebwaf` takes the edit in. the table is full, so on a public service the file grows to the default `SWWAF_MAX_TRACKED_CLIENTS` of 20,000, about 20 MiB. Written every 15 minutes, that is under 2 GiB of disk writes a day. - - `bans.json` takes about 2 KiB per ban and at most about 8 KiB, since the - texts in the notes are cut short. At the default `SWWAF_MAX_BANS` of 5,000 - it is about 10 MiB, and never more than about 40 MiB, plus whatever bans - an admin made. It is written when a ban is made, lifted or made permanent, - at most once every 10 seconds, and otherwise with the 15-minute write, so - its writes follow the bans made: with a full file, a hundred new bans a - day come to about 1 GiB of disk writes. + - `bans.json` takes about 1.2 KiB per ban and at most about 2.5 KiB, since + the notes hold one request and their texts are cut short. At the default + `SWWAF_MAX_BANS` of 5,000 it is about 6 MiB, and never more than about 12 + MiB, plus whatever bans an admin made. It is written when a ban is made, + lifted or made permanent, at most once every 10 seconds, and otherwise + with the 15-minute write, so its writes follow the bans made: with a full + file, a hundred new bans a day come to about 600 MiB of disk writes. - `lookups.json` takes about 150 bytes per answer, about 15 MiB when full. Written every 15 minutes, that is under 1.5 GiB of disk writes a day. - `reputation.json` and `alerts.json` are usually a few MiB or less. @@ -1012,10 +1029,12 @@ and the running `smallwebwaf` takes the edit in. ## Request log -One JSON object per line on stdout for every request, including refused ones. -stdout is always on. When `SWWAF_LOG_REMOTE_URL` is set the same lines are also -sent to the remote endpoint, so a deployment can stop depending on docker's log -handling while `docker logs` keeps working. +One JSON object per line on stdout for every request, including refused ones, +apart from those Go's HTTP server ends before `smallwebwaf` sees them (see +"Configuration surface", size and time limits). stdout is always on. When +`SWWAF_LOG_REMOTE_URL` is set the same lines are also sent to the remote +endpoint, so a deployment can stop depending on docker's log handling while +`docker logs` keeps working. - Standard web log fields: `time` (RFC 3339 with milliseconds), `instance`, `client_ip`, `method`, `scheme`, `host`, `path`, `query`, `protocol`, @@ -1146,31 +1165,46 @@ The image holds: - Ubuntu 26.04 LTS, the newest long-term support release of Ubuntu, pinned by digest. The image moves to the next LTS release when that ships. +- Three packages from Ubuntu, `ca-certificates`, `nix-bin` and `runit`, + installed from a dated snapshot of Ubuntu's archive and checked by hash (see + "Packages from Ubuntu" below). The Ubuntu image has no CA certificates, and + they come with `nix-bin` only because a library it uses recommends them, so + `ca-certificates` is installed by name. Without it Nix cannot download + packages, and `smallwebwaf`, unable to reach GeoJS, would count every visitor + as coming from an unknown country. - Nix, the package manager, from Ubuntu's own `nix-bin` package, and nixpkgs, the collection of packages Nix installs from, fixed at one commit (see "Packages from nixpkgs" below). Root uses Nix directly, and no Nix daemon runs - in the container. -- runit, from Ubuntu's own `runit` package, and `runsvinit` as the entrypoint. - Neither Ubuntu nor nixpkgs packages `runsvinit`, so the image builds it from - its source (`github.com/peterbourgon/runsvinit`) at a version fixed by hash. - `runsvinit` starts runit's `runsvdir`, which starts a `runsv` for each - directory under `/etc/service`; each `runsv` runs the `run` script in its - directory, and runs it again whenever it exits. Ubuntu's runit looks for + in the container. Nix run by root expects a group of build users, `nixbld`, + which `nix-bin` does not create, so the image writes `build-users-group =` to + `/etc/nix/nix.conf`, and root's builds then run without build users. +- runit, from Ubuntu's own `runit` package, and `runsvinit` as the entrypoint, + as the owner's code style guide asks for service containers. Neither Ubuntu + nor nixpkgs packages `runsvinit`, so the image builds it from its source + (`github.com/peterbourgon/runsvinit`) at a fixed commit hash. Its repository + is archived and has not changed since 2015, and its last tag is `v2.0.0`. It + has no `go.mod`, and `go build` of its directory needs one, so the build + writes one; since `runsvinit` uses only Go's standard library, that file names + nothing else. `runsvinit` starts runit's `runsvdir`, which starts a `runsv` + for each directory under `/etc/service`; each `runsv` runs the `run` script in + its directory, and runs it again whenever it exits. Ubuntu's runit looks for services in `/etc/service` too, so when `docker stop` has `runsvinit` stop each service with runit's `sv`, `sv` finds it. - The `smallwebwaf` binary, and a user of its own, `smallwebwaf` (uid and gid 65532). - The service directory `/etc/service/smallwebwaf`, whose `run` script waits one - second (`sleep 1`), makes `SWWAF_STATE_DIR` belong to the `smallwebwaf` user, - and starts `smallwebwaf` as that user with runit's `chpst`. + second (`sleep 1`), makes `SWWAF_STATE_DIR` and every file in it belong to the + `smallwebwaf` user, and starts `smallwebwaf` as that user with runit's + `chpst`. - The directory `/var/lib/smallwebwaf` for the state files, and `/etc/smallwebwaf/rules.d` with the default rule file (see "Rule files"). - Port 8080 declared (`EXPOSE 8080`), and the health check described below (`HEALTHCHECK`). Every `run` script, the app's included, is a bash script that starts with -`#!/usr/bin/env bash` and `set -euo pipefail`, as the owner's code style guide -asks; Ubuntu ships bash. +`#!/usr/bin/env bash` and `set -euo pipefail` and puts its code in a `main` +function, called on its last line, as the owner's code style guide asks; Ubuntu +ships bash. Beyond its `FROM` line, the app's Dockerfile adds: @@ -1209,23 +1243,64 @@ with `app.run` beside the Dockerfile: ```bash #!/usr/bin/env bash set -euo pipefail -sleep 1 -exec chpst -u app:app /usr/local/bin/app \ - --listen 127.0.0.1:8081 \ - --trusted-proxies 10.0.0.0/8,172.16.0.0/12,192.168.0.0/16,127.0.0.1/32,::1/128 + +main() { + sleep 1 + exec chpst -u app:app /usr/local/bin/app \ + --listen 127.0.0.1:8081 \ + --trusted-proxies 10.0.0.0/8,172.16.0.0/12,192.168.0.0/16,127.0.0.1/32,::1/128 +} + +main "$@" ``` +Packages from Ubuntu: the image installs `ca-certificates`, `nix-bin` and +`runit` from Ubuntu's snapshot service, which serves the archive as it was at a +given moment, rather than from the archive itself, whose packages change with +every update. The image's Dockerfile names that moment, in apt's +`--snapshot 20261001T000000Z` form, on both `apt-get update` and +`apt-get install`. That moment is never earlier than the date of the pinned +Ubuntu image, since packages from an older snapshot can need older versions of +packages the Ubuntu image already holds, and it moves forward whenever that +image's digest does. `apt-get update` keeps the snapshot's `InRelease` files, +which apt checks against the archive's signature, in `/var/lib/apt/lists/`; each +lists the SHA-256 hash of the package lists it covers, and each package list the +hash of every package in it. The Dockerfile also names the SHA-256 hash of each +of the snapshot's `InRelease` files, and the build checks them after +`apt-get update` and before `apt-get install`, so every package apt installs is +checked, through those files, against hashes the Dockerfile names. +`apt-get update` also fetches the live archive's `InRelease` files into the same +directory; their hashes change whenever the archive does, and the install does +not use them, so the check leaves them out. The hashes are those of the archive +for amd64, and so the image is built for amd64: other architectures use Ubuntu's +ports archive, whose `InRelease` files differ. The snapshot service is reached +over HTTPS, and the Ubuntu image has no CA certificates of its own, so this one +install uses those of the Go image that `smallwebwaf` is built in, which is +pinned by digest too: apt's `Acquire::https::CaInfo` option names that image's +CA certificate file, `/etc/ssl/certs/ca-certificates.crt`, mounted for that one +step. + Packages from nixpkgs: nixpkgs is fixed at one commit of its newest release -branch, `nixos-26.05` today. The image's Dockerfile names the commit and the -hash of its contents, and the build checks that hash. nixpkgs is set up for root -under the name `nixpkgs`, so the app's Dockerfile installs a package with -`nix-env -iA nixpkgs.`, and whatever it installs is on the `PATH` of every -service. Because nixpkgs stays at that commit, an app built on the same -`smallwebwaf` image gets the same packages each time it is built. A newer commit -of the branch, with its security fixes, comes with a newer `smallwebwaf` image, -as do Ubuntu's own fixes; an app takes them by changing the digest in its `FROM` -line. When nixpkgs makes its next release, every six months, the image moves to -that release's branch. +branch, `nixos-26.05` today. For each commit of the branch that has passed its +tests, the Nix project publishes a release on `releases.nixos.org`, such as +`nixos-26.05.11045.774debe7a0d1`, and the image takes nixpkgs from that +release's file `nixexprs.tar.xz`, not from a GitHub archive of the commit, whose +bytes can change. The image's Dockerfile names the release and the SHA-256 hash +of that file, which the release's page lists, and the build checks the hash +before unpacking it. nixpkgs is set up for root under the name `nixpkgs`, so the +app's Dockerfile installs a package with `nix-env -iA nixpkgs.`, and +whatever it installs is on the `PATH` of every service: the image adds root's +Nix profile, `/nix/var/nix/profiles/default/bin`, at the end of the `PATH`, +after Ubuntu's own directories, so that no package hides the image's own +commands. busybox, for one, brings its own `sv`, which looks for services +elsewhere. Because nixpkgs stays at that commit, an app built on the same +`smallwebwaf` image gets the same packages each time it is built. Unpacked, +nixpkgs takes about 500 MiB of disk, more on some filesystems such as ZFS, and +each package an app installs from it adds its own size, with everything it +depends on. A newer commit of the branch, with its security fixes, comes with a +newer `smallwebwaf` image, as do Ubuntu's own fixes; an app takes them by +changing the digest in its `FROM` line. When nixpkgs makes its next release, +every six months, the image moves to that release's branch. The two processes: @@ -1251,19 +1326,28 @@ The two processes: - The container's root filesystem stays writable: runit writes each service's status into its directory under `/etc/service`. -The health check: the image's `HEALTHCHECK` passes while `smallwebwaf` answers -`GET /_smallwebwaf/healthz` on `127.0.0.1:8080` and the app accepts connections +The health check: the image's `HEALTHCHECK` runs `smallwebwaf healthcheck`, +which passes while `smallwebwaf` answers `GET /_smallwebwaf/healthz` on +`127.0.0.1`, at the port in `SWWAF_LISTEN_ADDR`, and the app accepts connections at the address in `SWWAF_UPSTREAM_URL`, and fails when either does not. The -container therefore shows as healthy only while both processes are up. An app -with a health check of its own can replace the image's `HEALTHCHECK` with one -that checks both. +container therefore shows as healthy only while both processes are up. traefik +sends a container no requests until it shows as healthy, so the check runs every +second from the container's start until it first passes, for up to a minute, and +every 30 seconds after that. An app with a health check of its own can replace +the image's `HEALTHCHECK` with one that checks both. Ports: `smallwebwaf` listens on port 8080 on every address and on no other port; its health check, metrics and ban management are all on that listener, under -`/_smallwebwaf/` (see "Admin endpoints"). The app must leave port 8080 free. It -listens on `127.0.0.1:8081` only, so that nothing outside the container reaches -it except through `smallwebwaf`: an app that listens on every address can be -reached around `smallwebwaf` by anything that reaches the container. +`/_smallwebwaf/` (see "Admin endpoints"). `SWWAF_LISTEN_ADDR` may set another +port: the image's health check takes its port from that setting, and traefik's +port label (`traefik.http.services..loadbalancer.server.port`) must name +the same port, and the app must leave that port free. The address part of +`SWWAF_LISTEN_ADDR` stays empty (for example `:9000`, never `127.0.0.1:9000`), +so `smallwebwaf` keeps listening on every address: traefik reaches it on the +container's address, and the health check on `127.0.0.1`. The app listens on +`127.0.0.1:8081` only, so that nothing outside the container reaches it except +through `smallwebwaf`: an app that listens on every address can be reached +around `smallwebwaf` by anything that reaches the container. State: `smallwebwaf` keeps its state files in `/var/lib/smallwebwaf` (`SWWAF_STATE_DIR`), a directory of its own beside the app's data, which the app @@ -1271,13 +1355,21 @@ keeps in directories of its own, such as `/var/lib/app`. Without a volume there, the files live in the container: they survive a restart of the container and are lost when a deploy replaces it. A volume mounted at `/var/lib/smallwebwaf`, named or a host directory, keeps them across deploys; the `run` script of -`smallwebwaf` makes it belong to the `smallwebwaf` user, so a host directory -mounted there needs no change of owner. It is a volume of its own, separate from -the app's, holds a few tens of MiB at most with the defaults (see "Persistent -state"), and needs no backup beyond whatever the host already does. The image -declares no volume, since every app image built on it would inherit it. -Milestone 2 (https://git.eeqj.de/sneak/smallwebwaf/issues/14) writes no state -files and needs no volume. +`smallwebwaf` makes it and every file in it belong to the `smallwebwaf` user, so +a host directory mounted there needs no change of owner, and files an earlier +owner left in it can be read and replaced. It is a volume of its own, separate +from the app's, holds a few tens of MiB at most with the defaults (see +"Persistent state"), and needs no backup beyond whatever the host already does. +The image declares no volume, since every app image built on it would inherit +it. Milestone 2 (https://git.eeqj.de/sneak/smallwebwaf/issues/14) writes no +state files and needs no volume. + +Tokens: a token given as a file (`SWWAF_ADMIN_TOKEN_FILE`, +`SWWAF_METRICS_TOKEN_FILE`) is out of the app's reach only while the +`smallwebwaf` user alone can read the file. The operator makes the file on the +host, owned by uid 65532, the `smallwebwaf` user, with mode `0400`, and mounts +the directory that holds it into the container read-only; the container sees the +same owner and mode. Forwarded headers: the app's TCP peer is `smallwebwaf` on `127.0.0.1`, and the `X-Forwarded-For` the app receives ends with traefik's address, which @@ -1296,7 +1388,8 @@ and runs the one container as it runs any app, with the app's traefik labels, environment variables and volumes. The labels route to port 8080 (`traefik.http.services..loadbalancer.server.port=8080`), any `SWWAF_` settings go with the app's environment variables, and the volume for -`/var/lib/smallwebwaf` goes beside the app's own. +`/var/lib/smallwebwaf` goes beside the app's own, as does the directory that +holds any token file. - The endpoints of `smallwebwaf` are reached through traefik like any other request, for example `https://app.example.invalid/_smallwebwaf/metrics` for a @@ -1320,12 +1413,13 @@ settings go with the app's environment variables, and the volume for taking more than 60 seconds is cut off, whether it is a git push, an LFS object, a container image layer, a package file or a release attachment. The client is answered `413` for a body that is too large, before anything - reaches gitea when the request announces its size, or `408` for one that - is too slow; the upload fails, and no one is banned for it. A gitea that - takes large uploads needs `SWWAF_REQUEST_MAX_BYTES`, - `SWWAF_CLIENT_REQUEST_TIMEOUT` and `SWWAF_UPSTREAM_REQUEST_TIMEOUT` raised - to fit. The Core Rule Set does not read an upload's body, which streams - through without being held in memory. + reaches gitea when the request announces its size, or `408` for one the + client sends too slowly (`504` if gitea is too slow to take it); the + upload fails, and no one is banned for it. A gitea that takes large + uploads needs `SWWAF_REQUEST_MAX_BYTES`, `SWWAF_CLIENT_REQUEST_TIMEOUT` + and `SWWAF_UPSTREAM_REQUEST_TIMEOUT` raised to fit. The Core Rule Set does + not read an upload's body, which streams through without being held in + memory. - At the defaults (see "Configuration surface", attack detection), the Core Rule Set lets gitea's ordinary use through, apart from the refusals in the next note: browsing and views of files in a repository, with their @@ -1495,7 +1589,7 @@ settings go with the app's environment variables, and the volume for stops the start. The app starts with the same environment variables as `smallwebwaf`, so it can read a token given as one; a token given as a file that only the `smallwebwaf` user can read (`SWWAF_ADMIN_TOKEN_FILE`, - `SWWAF_METRICS_TOKEN_FILE`) is out of the app's reach. + `SWWAF_METRICS_TOKEN_FILE`) is out of the app's reach (see "Deployment"). - GeoJS, the default lookup source: every new visitor's address goes to a third party, and a swarm of fresh addresses, when lookups peak, is when GeoJS may slow down or block `smallwebwaf`. Keeping answers for 7 days and asking about @@ -1524,16 +1618,23 @@ settings go with the app's environment variables, and the volume for only while a list is set. A client on a private, loopback or link-local address has no country, and neither list checks it. - Like milestone 1, it writes nothing to disk: the GeoJS answers and the - rate counters are kept in memory only, and a restart loses them. + rate counters are kept in memory only, and a restart loses them. The + header size and the idle time stay fixed at their defaults. - The container image described under "Deployment", with runit and the container's health check. The health check calls `/_smallwebwaf/healthz`, so milestone 2 answers that path, although the other admin endpoints come - later. -- After milestone 2, the rest of the design, in this order: + later. The image's `/var/lib/smallwebwaf`, which the `run` script gives to + the `smallwebwaf` user, and `/etc/smallwebwaf/rules.d` come with the state + files and the rule files. +- Milestone 3 and later: the rest of the design, in this order: - static lists, the bans that broken request limits lead to, the ban ledger and the JSON state files with edits taken in while running, exemptions, `observe` mode, the rest of the request log's fields, the metrics - endpoint; + endpoint, and the header size and the idle time as settings + (`SWWAF_CLIENT_REQUEST_HEADER_MAX_BYTES`, `SWWAF_CLIENT_IDLE_TIMEOUT`). + With the static lists comes `SWWAF_ALLOW_NETS`, and from then on + `SWWAF_EXCLUSIVELY_ALLOWED_COUNTRIES` refuses a client on a private, + loopback or link-local address unless `SWWAF_ALLOW_NETS` lists it; - rule files, the other admin endpoints, alerting to all three destinations, remote log sending; - AS number and country lookup for every client, from the file or GeoJS, diff --git a/cmd/smallwebwaf/main.go b/cmd/smallwebwaf/main.go new file mode 100644 index 0000000..108247c --- /dev/null +++ b/cmd/smallwebwaf/main.go @@ -0,0 +1,18 @@ +// Command smallwebwaf is a web application firewall for one app: it runs +// between traefik and the app, and passes requests through within its +// limits. Everything it does is in internal/smallwebwaf. +package main + +import ( + "os" + + "sneak.berlin/go/smallwebwaf/internal/smallwebwaf" +) + +// Version is the version of the binary, set when it is built with +// -ldflags "-X main.Version=...". +var Version = "dev" //nolint:gochecknoglobals // the linker sets it + +func main() { + os.Exit(smallwebwaf.Main(Version)) +} diff --git a/deploy/example-app/Dockerfile b/deploy/example-app/Dockerfile new file mode 100644 index 0000000..a228c17 --- /dev/null +++ b/deploy/example-app/Dockerfile @@ -0,0 +1,20 @@ +# An app built on the smallwebwaf image, as under "Deployment" in +# SPEC.md, which script/example-app builds and checks. The app is +# busybox's web server, from the nixpkgs in the image, serving one page; +# a real app copies in its own binary instead. +# +# A real app names the smallwebwaf image by digest. This one takes the +# image script/example-app has just built, or else the one `make docker` +# builds. +ARG SMALLWEBWAF_IMAGE=smallwebwaf +FROM ${SMALLWEBWAF_IMAGE} + +# Packages the app needs, from the nixpkgs in the image. +RUN nix-env -iA nixpkgs.busybox + +# The app's page, and a user of its own to run it. +RUN mkdir /var/www && echo 'hello from the example app' > /var/www/index.html +RUN useradd --system --no-create-home --shell /usr/sbin/nologin app + +# The app's runit service. +COPY --chmod=755 app.run /etc/service/app/run diff --git a/deploy/example-app/app.run b/deploy/example-app/app.run new file mode 100755 index 0000000..81782d8 --- /dev/null +++ b/deploy/example-app/app.run @@ -0,0 +1,9 @@ +#!/usr/bin/env bash +set -euo pipefail + +main() { + sleep 1 + exec chpst -u app:app busybox httpd -f -p 127.0.0.1:8081 -h /var/www +} + +main "$@" diff --git a/go.mod b/go.mod new file mode 100644 index 0000000..caf899d --- /dev/null +++ b/go.mod @@ -0,0 +1,21 @@ +module sneak.berlin/go/smallwebwaf + +go 1.26.0 + +require ( + github.com/fsnotify/fsnotify v1.10.1 + github.com/hashicorp/golang-lru/v2 v2.0.7 + github.com/prometheus/client_golang v1.24.1 +) + +require ( + github.com/beorn7/perks v1.0.1 // indirect + github.com/cespare/xxhash/v2 v2.3.0 // indirect + github.com/kylelemons/godebug v1.1.0 // indirect + github.com/munnerz/goautoneg v0.0.0-20191010083416-a7dc8b61c822 // indirect + github.com/prometheus/client_model v0.6.2 // indirect + github.com/prometheus/common v0.70.1 // indirect + github.com/prometheus/procfs v0.21.1 // indirect + golang.org/x/sys v0.47.0 // indirect + google.golang.org/protobuf v1.36.11 // indirect +) diff --git a/go.sum b/go.sum new file mode 100644 index 0000000..0e9f7bb --- /dev/null +++ b/go.sum @@ -0,0 +1,40 @@ +github.com/beorn7/perks v1.0.1 h1:VlbKKnNfV8bJzeqoa4cOKqO6bYr3WgKZxO8Z16+hsOM= +github.com/beorn7/perks v1.0.1/go.mod h1:G2ZrVWU2WbWT9wwq4/hrbKbnv/1ERSJQ0ibhJ6rlkpw= +github.com/cespare/xxhash/v2 v2.3.0 h1:UL815xU9SqsFlibzuggzjXhog7bL6oX9BbNZnL2UFvs= +github.com/cespare/xxhash/v2 v2.3.0/go.mod h1:VGX0DQ3Q6kWi7AoAeZDth3/j3BFtOZR5XLFGgcrjCOs= +github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c= +github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= +github.com/fsnotify/fsnotify v1.10.1 h1:b0/UzAf9yR5rhf3RPm9gf3ehBPpf0oZKIjtpKrx59Ho= +github.com/fsnotify/fsnotify v1.10.1/go.mod h1:TLheqan6HD6GBK6PrDWyDPBaEV8LspOxvPSjC+bVfgo= +github.com/google/go-cmp v0.7.0 h1:wk8382ETsv4JYUZwIsn6YpYiWiBsYLSJiTsyBybVuN8= +github.com/google/go-cmp v0.7.0/go.mod h1:pXiqmnSA92OHEEa9HXL2W4E7lf9JzCmGVUdgjX3N/iU= +github.com/hashicorp/golang-lru/v2 v2.0.7 h1:a+bsQ5rvGLjzHuww6tVxozPZFVghXaHOwFs4luLUK2k= +github.com/hashicorp/golang-lru/v2 v2.0.7/go.mod h1:QeFd9opnmA6QUJc5vARoKUSoFhyfM2/ZepoAG6RGpeM= +github.com/klauspost/compress v1.19.1 h1:VsB4HPswih7mmZ8WleSFQ75c/Ui1M4trX5oAsJnhSlk= +github.com/klauspost/compress v1.19.1/go.mod h1:cwPg85FWrGar70rWktvGQj8/hthj3wpl0PGDogxkrSQ= +github.com/kylelemons/godebug v1.1.0 h1:RPNrshWIDI6G2gRW9EHilWtl7Z6Sb1BR0xunSBf0SNc= +github.com/kylelemons/godebug v1.1.0/go.mod h1:9/0rRGxNHcop5bhtWyNeEfOS8JIWk580+fNqagV/RAw= +github.com/munnerz/goautoneg v0.0.0-20191010083416-a7dc8b61c822 h1:C3w9PqII01/Oq1c1nUAm88MOHcQC9l5mIlSMApZMrHA= +github.com/munnerz/goautoneg v0.0.0-20191010083416-a7dc8b61c822/go.mod h1:+n7T8mK8HuQTcFwEeznm/DIxMOiR9yIdICNftLE1DvQ= +github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM= +github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4= +github.com/prometheus/client_golang v1.24.1 h1:JnJkREXzWxUdCuPFpIWZiPispT9xVV59uiuyR2bPlnU= +github.com/prometheus/client_golang v1.24.1/go.mod h1:F+oSRECHg4sse5ucfYpYDeIv/hu68Zo0uoHKetWnzcE= +github.com/prometheus/client_model v0.6.2 h1:oBsgwpGs7iVziMvrGhE53c/GrLUsZdHnqNwqPLxwZyk= +github.com/prometheus/client_model v0.6.2/go.mod h1:y3m2F6Gdpfy6Ut/GBsUqTWZqCUvMVzSfMLjcu6wAwpE= +github.com/prometheus/common v0.70.1 h1:1HvjP4D5oL3t8RsPlwxA9onvvStjtIHYE5XuuwOi/PY= +github.com/prometheus/common v0.70.1/go.mod h1:VdFUQDMZK3VLkurFUVhia6uys/0suUp86TJz5qbJRhc= +github.com/prometheus/procfs v0.21.1 h1:GljZCt+zSTS+NZq88cyQ1LjZ+RCHp3uVuabBWA5+OJI= +github.com/prometheus/procfs v0.21.1/go.mod h1:aB55Cww9pdSJVHk0hUf0inxWyyjPogFIjmHKYgMKmtY= +github.com/stretchr/testify v1.11.1 h1:7s2iGBzp5EwR7/aIZr8ao5+dra3wiQyKjjFuvgVKu7U= +github.com/stretchr/testify v1.11.1/go.mod h1:wZwfW3scLgRK+23gO65QZefKpKQRnfz6sD981Nm4B6U= +go.uber.org/goleak v1.3.0 h1:2K3zAYmnTNqV73imy9J1T3WC+gmCePx2hEGkimedGto= +go.uber.org/goleak v1.3.0/go.mod h1:CoHD4mav9JJNrW/WLlf7HGZPjdw8EucARQHekz1X6bE= +go.yaml.in/yaml/v2 v2.4.4 h1:tuyd0P+2Ont/d6e2rl3be67goVK4R6deVxCUX5vyPaQ= +go.yaml.in/yaml/v2 v2.4.4/go.mod h1:gMZqIpDtDqOfM0uNfy0SkpRhvUryYH0Z6wdMYcacYXQ= +golang.org/x/sys v0.47.0 h1:o7XGOvZQCADBQQ4Y7VNq2dRWQR7JmOUW8Kxx4ZsNgWs= +golang.org/x/sys v0.47.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw= +google.golang.org/protobuf v1.36.11 h1:fV6ZwhNocDyBLK0dj+fg8ektcVegBBuEolpbTQyBNVE= +google.golang.org/protobuf v1.36.11/go.mod h1:HTf+CrKn2C3g5S8VImy6tdcUvCska2kB7j23XfzDpco= +gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA= +gopkg.in/yaml.v3 v3.0.1/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM= diff --git a/internal/bans/bans.go b/internal/bans/bans.go new file mode 100644 index 0000000..19b7c19 --- /dev/null +++ b/internal/bans/bans.go @@ -0,0 +1,437 @@ +// Package bans is the ban ledger: the bans smallwebwaf makes on the +// netblocks of clients that break a rate limit, with their notes, as the +// "Bans" section of SPEC.md describes. The bans are kept in memory, and +// written to bans.json and read from it by the state package. +package bans + +import ( + "net/netip" + "slices" + "strings" + "sync" + "time" + + "github.com/hashicorp/golang-lru/v2/simplelru" +) + +// repeatFactor is how many times as long as the netblock's last ban a ban +// for a limit broken again within the repeat window lasts. +const repeatFactor = 3 + +// maxTextBytes is how much of each text in a ban's notes is kept. +const maxTextBytes = 256 + +// Rules are how long a ban for a broken limit lasts, and how many bans +// are held. +type Rules struct { + // LimitBanDuration is how long a first ban lasts. + LimitBanDuration time.Duration + // LimitBanRepeatWindow is how soon after the end of the netblock's + // ban that ended last a broken limit counts as a repeat, which bans + // for repeatFactor times as long as that ban. + LimitBanRepeatWindow time.Duration + // MaxBanDuration is the longest ban; a ban that would be longer is + // permanent instead. + MaxBanDuration time.Duration + // MaxBans is the most bans held, at least one. Past it, the earliest + // ban of the netblock that has gone longest without a request is + // dropped. + MaxBans int +} + +// Ban is a ban on a netblock for a broken limit, the only kind of ban +// smallwebwaf makes so far. +type Ban struct { + Netblock netip.Prefix + Start time.Time + // Expires is when the ban ends, zero for a permanent ban. + Expires time.Time + Notes Notes +} + +// Permanent reports whether the ban never runs out. +func (b Ban) Permanent() bool { + return b.Expires.IsZero() +} + +// ActiveAt reports whether the ban refuses requests at now. +func (b Ban) ActiveAt(now time.Time) bool { + return b.Permanent() || now.Before(b.Expires) +} + +// Notes are what an admin needs to decide whether to lift a ban. The +// JSON names are those of bans.json. +// +//nolint:tagliatelle // the state files use snake_case, as the request log does +type Notes struct { + // Country is the client's country, when it was looked up. + Country string `json:"country"` + // Limit, Window and Count are the limit that was broken, its window, + // "minute", "hour" or "day", and the count reached: the client's + // requests in the window, the one that broke the limit included. + // These are the requests that counted toward the ban, and the window + // is the time over which they came. + Limit int64 `json:"limit"` + Window string `json:"window"` + Count float64 `json:"count"` + // Request is the request that broke the limit. + Request Request `json:"request"` + // Requests is how many requests the netblock has sent since it was + // first seen, and Refused how many of them the ban has refused so + // far. Both go up with each request the ban refuses. + Requests int64 `json:"requests"` + Refused int64 `json:"refused"` + // EarlierBans is how many bans the netblock had before this one. + EarlierBans int `json:"earlier_bans"` +} + +// Request is a request in a ban's notes. Each text is cut to 256 bytes. +// +//nolint:tagliatelle // the state files use snake_case, as the request log does +type Request struct { + Time time.Time `json:"time"` + Method string `json:"method"` + Host string `json:"host"` + // Path is the path with its query string. + Path string `json:"path"` + // Status is what the client was sent, 0 if nothing was. + Status int `json:"status"` + UserAgent string `json:"user_agent"` +} + +// Ledger holds the bans. It is safe for concurrent use. +type Ledger struct { + rules Rules + // changed receives a value when a ban is made, unless one is waiting + // already. + changed chan struct{} + + mu sync.Mutex + // netblocks holds each banned netblock's bans, oldest first. Check and + // Find make each netblock they find the most recently seen. + netblocks *simplelru.LRU[netip.Prefix, *[]Ban] + // held is how many bans netblocks holds, at most rules.MaxBans. + held int + // made is how many bans BanForLimit has made since the start. + made int + // v4Lengths and v6Lengths are the lengths of the IPv4 and IPv6 + // netblocks that have been banned. Check looks for a ban at each of + // them, so that a ban read from bans.json refuses every client in its + // netblock even when it was made with another SWWAF_BAN_SCOPE_V4_PREFIX, + // or another length of an IPv6 client's netblock. + v4Lengths, v6Lengths []int +} + +// New returns a Ledger with no ban yet. +func New(rules Rules) *Ledger { + // Every netblock held has a ban, so there are never more netblocks + // than rules.MaxBans, and the LRU never drops one itself. + netblocks, err := simplelru.NewLRU[netip.Prefix, *[]Ban](rules.MaxBans, nil) + if err != nil { + panic(err) // NewLRU fails only for a size below one + } + + return &Ledger{ + rules: rules, + changed: make(chan struct{}, 1), + netblocks: netblocks, + } +} + +// Changed receives a value after a ban is made, so that bans.json can be +// written. Several bans made before it is read leave one value. +func (l *Ledger) Changed() <-chan struct{} { + return l.changed +} + +// Check is called for a request from client, at now. It reports whether +// a ban on a netblock client is in is active, and returns that ban, with +// the request counted among those it refused. +func (l *Ledger) Check(client netip.Addr, now time.Time) (Ban, bool) { + l.mu.Lock() + defer l.mu.Unlock() + + ban := l.active(client, now) + if ban == nil { + return Ban{}, false + } + + ban.Notes.Requests++ + ban.Notes.Refused++ + + return *ban, true +} + +// Find is Check without counting the request among those the ban +// refused: in observe mode a ban refuses nothing. +func (l *Ledger) Find(client netip.Addr, now time.Time) (Ban, bool) { + l.mu.Lock() + defer l.mu.Unlock() + + ban := l.active(client, now) + if ban == nil { + return Ban{}, false + } + + return *ban, true +} + +// activeBan returns the ban in bans, a netblock's bans oldest first, that +// is active at now, or nil when none is. If several are, it returns the +// one that started last. Every ban is looked at, since a ban an admin adds +// to bans.json can start before the netblock's others and outlast them. +func activeBan(bans []Ban, now time.Time) *Ban { + for i := len(bans) - 1; i >= 0; i-- { + if bans[i].ActiveAt(now) { + return &bans[i] + } + } + + return nil +} + +// BanForLimit bans netblock at now for a broken limit, with notes, and +// returns the ban. A first ban lasts LimitBanDuration. A ban made within +// LimitBanRepeatWindow after the netblock's ban that ended last lasts +// repeatFactor times as long as that one. A ban that would be longer +// than MaxBanDuration is permanent instead. If a ban on netblock is still +// active, as when two of its requests break a limit at once, that ban is +// returned and no other is made. The ledger fills in the notes' Refused +// and EarlierBans itself. +func (l *Ledger) BanForLimit(netblock netip.Prefix, now time.Time, notes Notes) Ban { + l.mu.Lock() + defer l.mu.Unlock() + + var last *Ban + + bans, found := l.netblocks.Get(netblock) + if found { + active := activeBan(*bans, now) + if active != nil { + return *active + } + + // No ban is active, so each has an end. A ban an admin adds to + // bans.json can start after another and end before it, so the + // ban that ended last is looked for among them all. + ended := slices.MaxFunc(*bans, func(a, b Ban) int { + return a.Expires.Compare(b.Expires) + }) + last = &ended + + // The netblock's first ban held counts the bans it had before that + // one, since dropped to make room, and each ban held adds one. + notes.EarlierBans = (*bans)[0].Notes.EarlierBans + len(*bans) + } + + notes.Request = notes.Request.cut() + ban := Ban{ + Netblock: netblock, + Start: now, + Expires: l.expiry(last, now), + Notes: notes, + } + l.add(ban) + l.made++ + + select { + case l.changed <- struct{}{}: + default: // a value is waiting already + } + + return ban +} + +// Bans returns the bans held on netblock, oldest first. It is not a +// request from netblock, and leaves when it was last seen unchanged. +func (l *Ledger) Bans(netblock netip.Prefix) []Ban { + l.mu.Lock() + defer l.mu.Unlock() + + bans, found := l.netblocks.Peek(netblock) + if !found { + return nil + } + + return slices.Clone(*bans) +} + +// Made returns how many bans the ledger has made since the start; bans +// read from bans.json are not among them. +func (l *Ledger) Made() int { + l.mu.Lock() + defer l.mu.Unlock() + + return l.made +} + +// Count returns how many of the bans held are active at now, and how many +// are permanent. +func (l *Ledger) Count(now time.Time) (int, int) { + l.mu.Lock() + defer l.mu.Unlock() + + active, permanent := 0, 0 + + for _, bans := range l.netblocks.Values() { + for _, ban := range *bans { + if ban.ActiveAt(now) { + active++ + } + + if ban.Permanent() { + permanent++ + } + } + } + + return active, permanent +} + +// Snapshot returns every ban held, sorted by netblock, and each +// netblock's bans oldest first, as bans.json lists them. +func (l *Ledger) Snapshot() []Ban { + l.mu.Lock() + defer l.mu.Unlock() + + held := make([]Ban, 0, l.held) + for _, bans := range l.netblocks.Values() { + held = append(held, *bans...) + } + + slices.SortStableFunc(held, func(a, b Ban) int { + return a.Netblock.Compare(b.Netblock) + }) + + return held +} + +// Load puts bans read from bans.json into the ledger, in place of the +// bans it holds, in the order they started, so that a netblock whose last +// ban started latest counts as the most recently seen. Each netblock is +// masked to its length, so that 203.0.113.9/24 is 203.0.113.0/24, and +// each text in the notes is cut to 256 bytes. Past MaxBans the earliest +// bans are dropped, as when they are made. +func (l *Ledger) Load(bans []Ban) { + bans = slices.Clone(bans) + slices.SortStableFunc(bans, func(a, b Ban) int { + return a.Start.Compare(b.Start) + }) + + l.mu.Lock() + defer l.mu.Unlock() + + l.netblocks.Purge() + l.held = 0 + l.v4Lengths, l.v6Lengths = nil, nil + + for _, ban := range bans { + ban.Netblock = ban.Netblock.Masked() + ban.Notes.Request = ban.Notes.Request.cut() + l.add(ban) + } +} + +// active returns the ban active at now on a netblock client is in, or +// nil. +func (l *Ledger) active(client netip.Addr, now time.Time) *Ban { + lengths := l.v6Lengths + if client.Is4() { + lengths = l.v4Lengths + } + + for _, length := range lengths { + bans, found := l.netblocks.Get(netip.PrefixFrom(client, length).Masked()) + if !found { + continue + } + + ban := activeBan(*bans, now) + if ban != nil { + return ban + } + } + + return nil +} + +// add adds ban to its netblock's bans, after the last, and makes its +// netblock the most recently seen. With MaxBans held, it drops one first. +func (l *Ledger) add(ban Ban) { + if l.held == l.rules.MaxBans { + l.dropOne() + } + + // dropOne can have dropped the netblock's last ban, and the netblock + // with it. + bans, found := l.netblocks.Get(ban.Netblock) + if !found { + bans = &[]Ban{} + l.netblocks.Add(ban.Netblock, bans) + } + + *bans = append(*bans, ban) + l.held++ + + lengths := &l.v6Lengths + if ban.Netblock.Addr().Is4() { + lengths = &l.v4Lengths + } + + if !slices.Contains(*lengths, ban.Netblock.Bits()) { + *lengths = append(*lengths, ban.Netblock.Bits()) + } +} + +// expiry returns when a ban for a broken limit made at now ends, or zero +// when it is permanent. last is the netblock's ban that ended last, or nil +// when it has none. +func (l *Ledger) expiry(last *Ban, now time.Time) time.Time { + length := l.rules.LimitBanDuration + + if last != nil && now.Sub(last.Expires) <= l.rules.LimitBanRepeatWindow { + lastLength := last.Expires.Sub(last.Start) + // This is repeatFactor * lastLength > MaxBanDuration, written so + // that it cannot overflow. + if lastLength > l.rules.MaxBanDuration/repeatFactor { + return time.Time{} + } + + length = repeatFactor * lastLength + } + + if length > l.rules.MaxBanDuration { + return time.Time{} + } + + return now.Add(length) +} + +// dropOne drops the earliest ban of the netblock that has gone longest +// without a request, and the netblock with it if that was its only ban. +func (l *Ledger) dropOne() { + netblock, bans, _ := l.netblocks.GetOldest() + if len(*bans) == 1 { + l.netblocks.Remove(netblock) + } else { + *bans = slices.Delete(*bans, 0, 1) + } + + l.held-- +} + +// cut returns r with each text cut to maxTextBytes and copied, so that +// the notes do not keep the rest of the request in memory. +func (r Request) cut() Request { + r.Method = cutText(r.Method) + r.Host = cutText(r.Host) + r.Path = cutText(r.Path) + r.UserAgent = cutText(r.UserAgent) + + return r +} + +// cutText returns a copy of the first maxTextBytes of text. +func cutText(text string) string { + return strings.Clone(text[:min(len(text), maxTextBytes)]) +} diff --git a/internal/bans/bans_test.go b/internal/bans/bans_test.go new file mode 100644 index 0000000..0bf38f1 --- /dev/null +++ b/internal/bans/bans_test.go @@ -0,0 +1,291 @@ +package bans_test + +import ( + "net/netip" + "strings" + "testing" + "time" + + "sneak.berlin/go/smallwebwaf/internal/bans" +) + +const day = 24 * time.Hour + +func TestRepeatsTripleUntilPermanent(t *testing.T) { + t.Parallel() + + ledger := bans.New(defaultRules()) + netblock := netip.MustParsePrefix("203.0.113.9/32") + now := midnight() + + // Each ban is followed by another as soon as it ends: 1, 3, 9, 27 and + // 81 hours. + for i, hours := range []int{1, 3, 9, 27, 81} { + ban := ledger.BanForLimit(netblock, now, bans.Notes{}) + + length := time.Duration(hours) * time.Hour + if !ban.Expires.Equal(now.Add(length)) || ban.Notes.EarlierBans != i { + t.Fatalf("ban %d lasts %s with %d earlier bans, want %d hours and %d", + i+1, ban.Expires.Sub(now), ban.Notes.EarlierBans, hours, i) + } + + now = ban.Expires + } + + // The sixth would last 243 hours, more than seven days: it is + // permanent, and never ends. + ban := ledger.BanForLimit(netblock, now, bans.Notes{}) + if !ban.Permanent() { + t.Fatalf("sixth ban ends at %s, want a permanent one", ban.Expires) + } + + _, banned := ledger.Check(netblock.Addr(), now.Add(100*365*day)) + if !banned { + t.Error("a permanent ban ended") + } +} + +func TestRepeatWindowRunsOut(t *testing.T) { + t.Parallel() + + for _, tc := range []struct { + name string + // gap is the time between the end of the first ban and the second. + gap time.Duration + want time.Duration + }{ + {"broken again as the window ends", day, 3 * time.Hour}, + {"broken again after the window", day + time.Nanosecond, time.Hour}, + } { + t.Run(tc.name, func(t *testing.T) { + t.Parallel() + + ledger := bans.New(defaultRules()) + netblock := netip.MustParsePrefix("203.0.113.9/32") + + first := ledger.BanForLimit(netblock, midnight(), bans.Notes{}) + second := ledger.BanForLimit(netblock, first.Expires.Add(tc.gap), bans.Notes{}) + + if second.Expires.Sub(second.Start) != tc.want || second.Notes.EarlierBans != 1 { + t.Errorf("second ban lasts %s with %d earlier bans, want %s and 1", + second.Expires.Sub(second.Start), second.Notes.EarlierBans, tc.want) + } + }) + } +} + +func TestFirstBanLongerThanTheMaximumIsPermanent(t *testing.T) { + t.Parallel() + + rules := defaultRules() + rules.LimitBanDuration = rules.MaxBanDuration + time.Hour + ledger := bans.New(rules) + + ban := ledger.BanForLimit(netip.MustParsePrefix("203.0.113.9/32"), midnight(), + bans.Notes{}) + if !ban.Permanent() { + t.Errorf("first ban ends at %s, want a permanent one", ban.Expires) + } +} + +func TestLongestBanSetFarOffDoesNotOverflow(t *testing.T) { + t.Parallel() + + // With bans of up to 100,000 days, the 14th ban in a row, of 3^13 + // hours, is within the maximum, and three times as long would not fit + // in a time.Duration. The 15th is permanent. + rules := defaultRules() + rules.MaxBanDuration = 100000 * day + ledger := bans.New(rules) + netblock := netip.MustParsePrefix("203.0.113.9/32") + now := midnight() + + for i := range 14 { + ban := ledger.BanForLimit(netblock, now, bans.Notes{}) + if !ban.Expires.After(ban.Start) { + t.Fatalf("ban %d starts at %s and ends at %s", i+1, ban.Start, ban.Expires) + } + + now = ban.Expires + } + + ban := ledger.BanForLimit(netblock, now, bans.Notes{}) + if !ban.Permanent() { + t.Errorf("15th ban ends at %s, want a permanent one", ban.Expires) + } +} + +func TestBrokenLimitDuringABanMakesNoOther(t *testing.T) { + t.Parallel() + + ledger := bans.New(defaultRules()) + netblock := netip.MustParsePrefix("203.0.113.9/32") + + first := ledger.BanForLimit(netblock, midnight(), bans.Notes{}) + again := ledger.BanForLimit(netblock, midnight().Add(time.Minute), bans.Notes{}) + + if again != first || len(ledger.Bans(netblock)) != 1 { + t.Errorf("a limit broken during a ban gave %+v and %d bans, want %+v and 1", + again, len(ledger.Bans(netblock)), first) + } +} + +func TestCheckRefusesWhileTheBanLastsAndCountsTheRefusals(t *testing.T) { + t.Parallel() + + ledger := bans.New(defaultRules()) + netblock := netip.MustParsePrefix("203.0.113.9/32") + ban := ledger.BanForLimit(netblock, midnight(), bans.Notes{Requests: 5}) + + for range 3 { + got, banned := ledger.Check(netblock.Addr(), ban.Expires.Add(-time.Nanosecond)) + if !banned || got.Start != ban.Start { + t.Fatalf("check during the ban gives %+v and %t", got, banned) + } + } + + _, banned := ledger.Check(netip.MustParseAddr("203.0.113.10"), midnight()) + if banned { + t.Error("another netblock is banned") + } + + _, banned = ledger.Check(netblock.Addr(), ban.Expires) + if banned { + t.Error("the ban did not end") + } + + // The netblock's requests went from 5 to 8 with the three refused. + notes := ledger.Bans(netblock)[0].Notes + if notes.Refused != 3 || notes.Requests != 8 { + t.Errorf("the notes count %d refused requests of %d, want 3 of 8", + notes.Refused, notes.Requests) + } +} + +func TestFindCountsNothing(t *testing.T) { + t.Parallel() + + ledger := bans.New(defaultRules()) + netblock := netip.MustParsePrefix("203.0.113.9/32") + ban := ledger.BanForLimit(netblock, midnight(), bans.Notes{Requests: 5}) + + got, banned := ledger.Find(netblock.Addr(), ban.Expires.Add(-time.Nanosecond)) + if !banned || got != ban { + t.Errorf("find during the ban gives %+v and %t, want %+v", got, banned, ban) + } + + _, banned = ledger.Find(netblock.Addr(), ban.Expires) + if banned { + t.Error("the ban did not end") + } + + if notes := ledger.Bans(netblock)[0].Notes; notes != ban.Notes { + t.Errorf("the notes are %+v, want them unchanged, %+v", notes, ban.Notes) + } +} + +func TestMaxBansDropsTheEarliestBanOfTheNetblockSeenLongestAgo(t *testing.T) { + t.Parallel() + + rules := defaultRules() + rules.MaxBans = 3 + ledger := bans.New(rules) + a := netip.MustParsePrefix("203.0.113.1/32") + b := netip.MustParsePrefix("203.0.113.2/32") + c := netip.MustParsePrefix("203.0.113.3/32") + d := netip.MustParsePrefix("2001:db8::/64") + now := midnight() + + first := ledger.BanForLimit(a, now, bans.Notes{}) + ledger.BanForLimit(b, now, bans.Notes{}) + ledger.BanForLimit(c, now, bans.Notes{}) + + // A request from a makes b the netblock seen longest ago, and its ban + // goes to make room for d's. + ledger.Check(a.Addr(), now) + ledger.BanForLimit(d, now, bans.Notes{}) + wantBans(t, ledger, map[netip.Prefix]int{a: 1, b: 0, c: 1, d: 1}) + + // a is banned again once its ban has ended; c, seen longest ago, goes. + ledger.BanForLimit(a, first.Expires, bans.Notes{}) + wantBans(t, ledger, map[netip.Prefix]int{a: 2, c: 0, d: 1}) + + // With d seen since, a is seen longest ago, and its earlier ban goes + // first. + ledger.Check(d.Addr(), first.Expires) + ledger.BanForLimit(b, first.Expires, bans.Notes{}) + wantBans(t, ledger, map[netip.Prefix]int{a: 1, b: 1, d: 1}) + + if !ledger.Bans(a)[0].Start.Equal(first.Expires) { + t.Errorf("a kept its ban of %s, want the later one", ledger.Bans(a)[0].Start) + } +} + +func TestFullLedgerDropsTheEarlierBanOfTheNetblockBannedAgain(t *testing.T) { + t.Parallel() + + // With room for one ban, the netblock's ended ban goes to make room for + // its new one, whose notes still count it. + rules := defaultRules() + rules.MaxBans = 1 + ledger := bans.New(rules) + netblock := netip.MustParsePrefix("203.0.113.9/32") + + first := ledger.BanForLimit(netblock, midnight(), bans.Notes{}) + second := ledger.BanForLimit(netblock, first.Expires, bans.Notes{}) + + held := ledger.Bans(netblock) + if len(held) != 1 || held[0] != second || held[0].Notes.EarlierBans != 1 { + t.Errorf("the ledger holds %+v, want only the second ban, with 1 earlier ban", + held) + } +} + +func TestRequestTextsAreCutTo256Bytes(t *testing.T) { + t.Parallel() + + ledger := bans.New(defaultRules()) + netblock := netip.MustParsePrefix("203.0.113.9/32") + long := strings.Repeat("a", 300) + request := bans.Request{ + Time: midnight(), Method: long, Host: long, Path: long, Status: 403, UserAgent: long, + } + + ban := ledger.BanForLimit(netblock, midnight(), bans.Notes{Request: request}) + + cut := long[:256] + want := bans.Request{ + Time: midnight(), Method: cut, Host: cut, Path: cut, Status: 403, UserAgent: cut, + } + + if ban.Notes.Request != want || ledger.Bans(netblock)[0].Notes.Request != want { + t.Errorf("the notes keep %+v, want each text cut to 256 bytes", ban.Notes.Request) + } +} + +// defaultRules are the rules at the settings' defaults. +func defaultRules() bans.Rules { + return bans.Rules{ + LimitBanDuration: time.Hour, + LimitBanRepeatWindow: day, + MaxBanDuration: 7 * day, + MaxBans: 5000, + } +} + +// midnight is when the tests' first bans are made. +func midnight() time.Time { + return time.Date(2026, 10, 6, 0, 0, 0, 0, time.UTC) +} + +// wantBans checks how many bans the ledger holds on each netblock. +func wantBans(t *testing.T, ledger *bans.Ledger, want map[netip.Prefix]int) { + t.Helper() + + for netblock, count := range want { + got := len(ledger.Bans(netblock)) + if got != count { + t.Errorf("%s has %d bans, want %d", netblock, got, count) + } + } +} diff --git a/internal/bans/snapshot_test.go b/internal/bans/snapshot_test.go new file mode 100644 index 0000000..3a1a651 --- /dev/null +++ b/internal/bans/snapshot_test.go @@ -0,0 +1,297 @@ +package bans_test + +import ( + "net/netip" + "slices" + "strings" + "testing" + "time" + + "sneak.berlin/go/smallwebwaf/internal/bans" +) + +func TestChangedAfterABanIsMade(t *testing.T) { + t.Parallel() + + ledger := bans.New(defaultRules()) + netblock := netip.MustParsePrefix("203.0.113.9/32") + + wantChanged(t, ledger, false) + + ledger.BanForLimit(netblock, midnight(), bans.Notes{}) + wantChanged(t, ledger, true) + + // A limit broken during the ban makes no other, and a refusal changes + // only the counts in the notes, which wait for the interval's write. + ledger.BanForLimit(netblock, midnight().Add(time.Minute), bans.Notes{}) + ledger.Check(netblock.Addr(), midnight().Add(time.Minute)) + wantChanged(t, ledger, false) + + // Two bans before the value is read leave one. + ledger.BanForLimit(netip.MustParsePrefix("203.0.113.10/32"), midnight(), bans.Notes{}) + ledger.BanForLimit(netip.MustParsePrefix("203.0.113.11/32"), midnight(), bans.Notes{}) + wantChanged(t, ledger, true) + wantChanged(t, ledger, false) +} + +func TestSnapshotListsEveryBanByNetblock(t *testing.T) { + t.Parallel() + + ledger := bans.New(defaultRules()) + v6 := netip.MustParsePrefix("2001:db8::/64") + high := netip.MustParsePrefix("203.0.113.10/32") + low := netip.MustParsePrefix("203.0.113.9/32") + + first := ledger.BanForLimit(v6, midnight(), bans.Notes{}) + ledger.BanForLimit(high, midnight(), bans.Notes{}) + ledger.BanForLimit(low, midnight(), bans.Notes{}) + ledger.BanForLimit(v6, first.Expires, bans.Notes{}) + + snapshot := ledger.Snapshot() + + got := make([]string, 0, len(snapshot)) + for _, ban := range snapshot { + got = append(got, ban.Netblock.String()+" "+ban.Start.Format(time.Kitchen)) + } + + want := []string{ + "203.0.113.9/32 12:00AM", "203.0.113.10/32 12:00AM", + "2001:db8::/64 12:00AM", "2001:db8::/64 1:00AM", + } + if !slices.Equal(got, want) { + t.Errorf("snapshot %v, want %v", got, want) + } +} + +func TestLoadedBansCarryOn(t *testing.T) { + t.Parallel() + + before := bans.New(defaultRules()) + netblock := netip.MustParsePrefix("203.0.113.9/32") + ban := before.BanForLimit(netblock, midnight(), bans.Notes{Limit: 1}) + + // Loaded into a new ledger, as across a restart, the ban still refuses + // while it lasts, and once it has ended a broken limit bans for three + // times as long, with the loaded ban counted among the earlier ones. + after := bans.New(defaultRules()) + after.Load(before.Snapshot()) + + _, banned := after.Check(netblock.Addr(), ban.Expires.Add(-time.Second)) + if !banned { + t.Error("the loaded ban does not refuse") + } + + again := after.BanForLimit(netblock, ban.Expires, bans.Notes{}) + if again.Expires.Sub(again.Start) != 3*time.Hour || again.Notes.EarlierBans != 1 { + t.Errorf("the next ban lasts %s with %d earlier bans, want 3h and 1", + again.Expires.Sub(again.Start), again.Notes.EarlierBans) + } +} + +func TestLoadedBanRefusesEveryClientInItsNetblock(t *testing.T) { + t.Parallel() + + // Two entries as an admin might write them, with addresses not masked + // to their lengths, the IPv6 one shorter than the /64 an IPv6 client's + // ban covers, beside a ban the ledger makes on one IPv4 address. + ledger := bans.New(defaultRules()) + ledger.Load([]bans.Ban{ + {Netblock: netip.MustParsePrefix("203.0.113.9/24"), Start: midnight()}, + {Netblock: netip.MustParsePrefix("2001:db8::1/48"), Start: midnight()}, + }) + ledger.BanForLimit(netip.MustParsePrefix("198.51.100.7/32"), midnight(), bans.Notes{}) + + for client, want := range map[string]bool{ + "203.0.113.0": true, + "203.0.113.200": true, + "203.0.114.1": false, + "2001:db8:0:5::1": true, + "2001:db8:1::1": false, + "198.51.100.7": true, + "198.51.100.8": false, + } { + _, banned := ledger.Check(netip.MustParseAddr(client), midnight()) + if banned != want { + t.Errorf("%s is refused: %t, want %t", client, banned, want) + } + } + + // The loaded netblocks are written back masked. + snapshot := ledger.Snapshot() + + got := make([]string, 0, len(snapshot)) + for _, ban := range snapshot { + got = append(got, ban.Netblock.String()) + } + + want := []string{"198.51.100.7/32", "203.0.113.0/24", "2001:db8::/48"} + if !slices.Equal(got, want) { + t.Errorf("the ledger holds bans on %v, want %v", got, want) + } +} + +func TestPermanentBanStartedBeforeAnEndedOneRefuses(t *testing.T) { + t.Parallel() + + // As when an admin adds a permanent ban to bans.json with a start + // before that of the netblock's ban that has ended. + netblock := netip.MustParsePrefix("203.0.113.0/24") + permanent := bans.Ban{Netblock: netblock, Start: midnight().Add(-time.Hour)} + ended := bans.Ban{ + Netblock: netblock, + Start: midnight(), + Expires: midnight().Add(time.Hour), + } + + ledger := bans.New(defaultRules()) + ledger.Load([]bans.Ban{permanent, ended}) + + now := midnight().Add(2 * time.Hour) + client := netip.MustParseAddr("203.0.113.9") + + ban, banned := ledger.Find(client, now) + if !banned || !ban.Permanent() { + t.Errorf("find gives %+v and %t, want the permanent ban", ban, banned) + } + + ban, banned = ledger.Check(client, now) + if !banned || !ban.Permanent() { + t.Errorf("the client is refused: %t, under %+v, want under the permanent ban", + banned, ban) + } + + // A limit broken now makes no shorter ban over the permanent one. + ban = ledger.BanForLimit(netblock, now, bans.Notes{}) + if !ban.Permanent() || len(ledger.Bans(netblock)) != 2 { + t.Errorf("a broken limit returned %+v and left the netblock %d bans, "+ + "want the permanent ban and 2", ban, len(ledger.Bans(netblock))) + } +} + +func TestNextBanWorkedOutFromTheBanThatEndedLast(t *testing.T) { + t.Parallel() + + // A 9-hour ban smallwebwaf made, the third in a row, and an admin's + // 1-hour ban added to bans.json over it, with no notes. + netblock := netip.MustParsePrefix("203.0.113.9/32") + nineHours := bans.Ban{ + Netblock: netblock, + Start: midnight(), + Expires: midnight().Add(9 * time.Hour), + Notes: bans.Notes{EarlierBans: 2}, + } + admins := bans.Ban{ + Netblock: netblock, + Start: midnight().Add(time.Hour), + Expires: midnight().Add(2 * time.Hour), + } + + ledger := bans.New(defaultRules()) + ledger.Load([]bans.Ban{nineHours, admins}) + + // Once both have ended, a limit broken within the repeat window bans + // for three times the 9 hours, and the notes count the two bans + // before the 9-hour one, it, and the admin's. + ban := ledger.BanForLimit(netblock, nineHours.Expires.Add(time.Hour), bans.Notes{}) + if ban.Expires.Sub(ban.Start) != 27*time.Hour || ban.Notes.EarlierBans != 4 { + t.Errorf("the next ban lasts %s with %d earlier bans, want 27h and 4", + ban.Expires.Sub(ban.Start), ban.Notes.EarlierBans) + } +} + +func TestLoadKeepsAtMostMaxBansDroppingTheEarliest(t *testing.T) { + t.Parallel() + + // bans.json lists the bans by netblock, not in the order they began. + later := bans.Ban{Netblock: netip.MustParsePrefix("203.0.113.1/32"), Start: midnight()} + earlier := bans.Ban{ + Netblock: netip.MustParsePrefix("203.0.113.2/32"), + Start: midnight().Add(-time.Hour), + } + + rules := defaultRules() + rules.MaxBans = 1 + ledger := bans.New(rules) + ledger.Load([]bans.Ban{later, earlier}) + + held := ledger.Snapshot() + if len(held) != 1 || held[0] != later { + t.Errorf("the ledger holds %+v, want only the ban that began later", held) + } +} + +func TestLoadReplacesTheBansHeld(t *testing.T) { + t.Parallel() + + // Room for three bans, so that the second load, were it added to the + // two bans held, would drop none of them to make room. + rules := defaultRules() + rules.MaxBans = 3 + ledger := bans.New(rules) + kept := bans.Ban{Netblock: netip.MustParsePrefix("2001:db8::/64"), Start: midnight()} + ledger.Load([]bans.Ban{ + {Netblock: netip.MustParsePrefix("203.0.113.0/24"), Start: midnight()}, + kept, + }) + + // Loaded again without the first ban, as when an admin's edit of + // bans.json is taken in, that ban is lifted. + ledger.Load([]bans.Ban{kept}) + + _, banned := ledger.Check(netip.MustParseAddr("203.0.113.9"), midnight()) + if banned { + t.Error("a ban left out of the second load still refuses") + } + + // The ledger holds one ban, so it makes two more without dropping any. + first := ledger.BanForLimit(netip.MustParsePrefix("198.51.100.7/32"), midnight(), + bans.Notes{}) + second := ledger.BanForLimit(netip.MustParsePrefix("198.51.100.8/32"), midnight(), + bans.Notes{}) + + want := []bans.Ban{first, second, kept} + if got := ledger.Snapshot(); !slices.Equal(got, want) { + t.Errorf("the ledger holds %+v, want %+v", got, want) + } +} + +func TestLoadCutsTheTextsTo256Bytes(t *testing.T) { + t.Parallel() + + long := strings.Repeat("a", 300) + ban := bans.Ban{ + Netblock: netip.MustParsePrefix("203.0.113.9/32"), + Start: midnight(), + Notes: bans.Notes{Request: bans.Request{ + Method: long, Host: long, Path: long, UserAgent: long, + }}, + } + + ledger := bans.New(defaultRules()) + ledger.Load([]bans.Ban{ban}) + + cut := long[:256] + want := bans.Request{Method: cut, Host: cut, Path: cut, UserAgent: cut} + + got := ledger.Snapshot()[0].Notes.Request + if got != want { + t.Errorf("the notes keep %+v, want each text cut to 256 bytes", got) + } +} + +// wantChanged checks whether the ledger's Changed has a value to read. +func wantChanged(t *testing.T, ledger *bans.Ledger, want bool) { + t.Helper() + + got := false + + select { + case <-ledger.Changed(): + got = true + default: + } + + if got != want { + t.Errorf("Changed has a value: %t, want %t", got, want) + } +} diff --git a/internal/config/config.go b/internal/config/config.go new file mode 100644 index 0000000..1c62687 --- /dev/null +++ b/internal/config/config.go @@ -0,0 +1,775 @@ +// 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" + "os" + "path/filepath" + "slices" + "strconv" + "strings" + "time" + "unicode/utf8" +) + +// 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 + // InstanceName is the name each request log line gives as instance + // (SWWAF_INSTANCE_NAME), by default the host's name, which docker sets + // to the first 12 characters of the container's id. + InstanceName string + // Observe is true in observe mode, when SWWAF_MODE is observe rather + // than enforce: a request that SWWAF_DENY_NETS, a ban, the country + // lists or a rate limit would refuse is passed to the app instead, and + // no ban is made. + Observe bool + // 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 + // LogRequestHeaders are the request headers whose values the request + // log gives, in lower case (SWWAF_LOG_REQUEST_HEADERS). + LogRequestHeaders []string + // MetricsToken is the bearer token a scraper sends for the metrics + // (SWWAF_METRICS_TOKEN), "" while it is unset and the metrics are off. + // MetricsTopN is how many countries get series of their own in the + // metrics (SWWAF_METRICS_TOP_N). + MetricsToken string + MetricsTopN int + + // 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 + // minTokenLength is the fewest characters a token may have. + minTokenLength = 32 + // masked is what the log shows for a token that is set. + masked = "********" +) + +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") + errNotHeaderName = errors.New( + "is not a header name such as accept-language") + errHeaderTakenOut = errors.New( + "is taken out of every request by Go's HTTP server, so it can never " + + "be logged") + 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") + errShortToken = errors.New("is shorter than 32 characters") + errNotMode = errors.New("is not enforce or observe") +) + +// 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} + hostname, _ := os.Hostname() // "" when the host has no name to give + cfg := &Config{ + ListenAddr: env.address("SWWAF_LISTEN_ADDR", ":8080"), + UpstreamURL: env.appURL("SWWAF_UPSTREAM_URL", "http://127.0.0.1:8081"), + InstanceName: env.value("SWWAF_INSTANCE_NAME", hostname), + Observe: env.observe("SWWAF_MODE", "enforce"), + 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"), + LogRequestHeaders: env.headerNames("SWWAF_LOG_REQUEST_HEADERS", + "accept,accept-language,accept-encoding,content-type,origin,range"), + MetricsToken: env.token("SWWAF_METRICS_TOKEN"), + MetricsTopN: env.numberNotOff("SWWAF_METRICS_TOP_N", "50"), + } + + 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 +} + +// observe reads the setting that is the mode, enforce or observe, and +// reports whether it is observe. +func (e *environment) observe(name, defaultValue string) bool { + mode := e.value(name, defaultValue) + if mode != "enforce" && mode != "observe" { + e.check(name, fmt.Errorf("%q %w", mode, errNotMode)) + } + + return mode == "observe" +} + +// 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 +} + +// headerNames reads a setting that is a list of header names, and +// returns them in lower case. +func (e *environment) headerNames(name, defaultValue string) []string { + headers, err := parseHeaderNames(e.value(name, defaultValue)) + e.check(name, err) + + return headers +} + +// 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 +} + +// token reads a setting that is a bearer token. Unset, it is "", which +// switches off what it guards; set, it must be at least minTokenLength +// characters. Neither the log nor an error shows its value. +func (e *environment) token(name string) string { + value, set := e.lookupEnv(name) + if !set { + e.settings = append(e.settings, slog.String(name, "")) + + return "" + } + + e.settings = append(e.settings, slog.String(name, masked)) + + if utf8.RuneCountInString(value) < minTokenLength { + e.check(name, errShortToken) + } + + return value +} + +// 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 +} + +// headerNameChars are the characters RFC 9110 allows in a header name: +// letters, digits and these marks. +const headerNameChars = "ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz" + + "0123456789!#$%&'*+-.^_`|~" + +// parseHeaderNames reads a comma-separated list of header names in either +// case, and returns them in lower case. Host and Transfer-Encoding are +// refused: Go's HTTP server takes them out of the request's headers. +func parseHeaderNames(value string) ([]string, error) { + items, err := parseList(value) + if err != nil { + return nil, err + } + + headers := make([]string, 0, len(items)) + + for _, item := range items { + for _, char := range item { + if !strings.ContainsRune(headerNameChars, char) { + return nil, fmt.Errorf("%q %w", item, errNotHeaderName) + } + } + + header := strings.ToLower(item) + switch header { + case "host": + return nil, fmt.Errorf("%q %w; the request's host is the field host", + item, errHeaderTakenOut) + case "transfer-encoding": + return nil, fmt.Errorf("%q %w", item, errHeaderTakenOut) + } + + headers = append(headers, header) + } + + return headers, 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 +} diff --git a/internal/config/config_test.go b/internal/config/config_test.go new file mode 100644 index 0000000..ccfe4c7 --- /dev/null +++ b/internal/config/config_test.go @@ -0,0 +1,599 @@ +package config_test + +import ( + "bytes" + "encoding/json" + "log/slog" + "maps" + "net/netip" + "os" + "slices" + "strings" + "testing" + "time" + + "sneak.berlin/go/smallwebwaf/internal/config" +) + +// The settings, by name. +const ( + listenAddr = "SWWAF_LISTEN_ADDR" + upstreamURL = "SWWAF_UPSTREAM_URL" + mode = "SWWAF_MODE" + trustedProxies = "SWWAF_TRUSTED_PROXIES" + clientRequestTimeout = "SWWAF_CLIENT_REQUEST_TIMEOUT" + clientHeaderMaxBytes = "SWWAF_CLIENT_REQUEST_HEADER_MAX_BYTES" + clientIdleTimeout = "SWWAF_CLIENT_IDLE_TIMEOUT" + clientResponseTimeout = "SWWAF_CLIENT_RESPONSE_TIMEOUT" + upstreamRequestTimeout = "SWWAF_UPSTREAM_REQUEST_TIMEOUT" + upstreamResponseTimeout = "SWWAF_UPSTREAM_RESPONSE_TIMEOUT" + requestMaxBytes = "SWWAF_REQUEST_MAX_BYTES" + responseMaxBytes = "SWWAF_RESPONSE_MAX_BYTES" + allowNets = "SWWAF_ALLOW_NETS" + rateLimitExemptNets = "SWWAF_RATE_LIMIT_EXEMPT_NETS" + denyNets = "SWWAF_DENY_NETS" + 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" + banResponse = "SWWAF_BAN_RESPONSE" + limitBanDuration = "SWWAF_LIMIT_BAN_DURATION" + limitBanRepeatWindow = "SWWAF_LIMIT_BAN_REPEAT_WINDOW" + maxBanDuration = "SWWAF_MAX_BAN_DURATION" + maxBans = "SWWAF_MAX_BANS" + banScopeV4Prefix = "SWWAF_BAN_SCOPE_V4_PREFIX" + stateDir = "SWWAF_STATE_DIR" + stateWriteDelay = "SWWAF_STATE_WRITE_DELAY" + stateCounterInterval = "SWWAF_STATE_COUNTER_INTERVAL" + metricsToken = "SWWAF_METRICS_TOKEN" //nolint:gosec // the setting's name + metricsTopN = "SWWAF_METRICS_TOP_N" + instanceName = "SWWAF_INSTANCE_NAME" + logRequestHeaders = "SWWAF_LOG_REQUEST_HEADERS" +) + +// defaultLogRequestHeaders is the default of SWWAF_LOG_REQUEST_HEADERS. +const defaultLogRequestHeaders = "accept,accept-language,accept-encoding," + + "content-type,origin,range" + +// token is a token of 32 characters, the shortest allowed. +const token = "0123456789abcdef0123456789abcdef" + +// off switches a timeout, a size limit or a rate limit off. +const off = "off" + +// environment is a set of environment variables, for FromEnvironment. +type environment map[string]string + +// lookupEnv reads one of the variables, as os.LookupEnv does. +func (e environment) lookupEnv(name string) (string, bool) { + value, ok := e[name] + + return value, ok +} + +// fromEnvironment reads the settings from env, which must be valid. +func fromEnvironment(t *testing.T, env environment) *config.Config { + t.Helper() + + cfg, err := config.FromEnvironment(env.lookupEnv) + if err != nil { + t.Fatalf("settings %v: %v", env, err) + } + + return cfg +} + +func TestDefaults(t *testing.T) { + t.Parallel() + + cfg := fromEnvironment(t, environment{}) + + wantSettings(t, cfg, config.Config{ + ListenAddr: ":8080", + Observe: false, + ClientRequestTimeout: time.Minute, + ClientRequestHeaderMaxBytes: 32 << 10, + ClientIdleTimeout: 2 * time.Minute, + ClientResponseTimeout: 30 * time.Minute, + UpstreamRequestTimeout: time.Minute, + UpstreamResponseTimeout: 30 * time.Minute, + RequestMaxBytes: 100 << 20, + ResponseMaxBytes: 5 << 30, + RateLimitPerMinute: 1000, + RateLimitPerHour: 10000, + RateLimitPerDay: 50000, + BanResponse: 403, + LimitBanDuration: time.Hour, + LimitBanRepeatWindow: 24 * time.Hour, + MaxBanDuration: 7 * 24 * time.Hour, + MaxBans: 5000, + BanScopeV4Prefix: 32, + StateDir: "/var/lib/smallwebwaf", + StateWriteDelay: 10 * time.Second, + StateCounterInterval: 15 * time.Minute, + MetricsToken: "", + MetricsTopN: 50, + }) + + if cfg.UpstreamURL.String() != "http://127.0.0.1:8081" { + t.Errorf("%s is %s", upstreamURL, cfg.UpstreamURL) + } + + wantNetblocks(t, cfg.TrustedProxies, + "10.0.0.0/8", "172.16.0.0/12", "192.168.0.0/16") + wantNetblocks(t, cfg.AllowNets) + wantNetblocks(t, cfg.RateLimitExemptNets) + wantNetblocks(t, cfg.DenyNets) + wantCountries(t, deniedCountries, cfg.DeniedCountries) + wantCountries(t, allowedCountries, cfg.ExclusivelyAllowedCountries) + + hostname, err := os.Hostname() + if err != nil || hostname == "" || cfg.InstanceName != hostname { + t.Errorf("%s is %q, want the host's name %q (%v)", instanceName, + cfg.InstanceName, hostname, err) + } + + wantHeaders := strings.Split(defaultLogRequestHeaders, ",") + if !slices.Equal(cfg.LogRequestHeaders, wantHeaders) { + t.Errorf("%s gave %v, want %v", logRequestHeaders, cfg.LogRequestHeaders, + wantHeaders) + } +} + +func TestValuesAsSet(t *testing.T) { + t.Parallel() + + cfg := fromEnvironment(t, environment{ + listenAddr: "127.0.0.1:9000", + upstreamURL: "https://app.internal:8443/", + mode: "observe", + trustedProxies: " 192.0.2.1, 10.1.2.3/8 ,2001:db8::/32", + clientRequestTimeout: "90s", + clientHeaderMaxBytes: "8K", + clientIdleTimeout: "5m", + clientResponseTimeout: "7d", + upstreamRequestTimeout: "1h30m", + upstreamResponseTimeout: off, + requestMaxBytes: "512K", + responseMaxBytes: "1234", + allowNets: "192.0.2.7", + rateLimitExemptNets: "2001:db8::/48, 10.9.8.7", + denyNets: "198.51.100.0/24", + rateLimitPerMinute: "60", + rateLimitPerHour: "600", + rateLimitPerDay: "6000", + deniedCountries: "cn, RU,kp,Xk", + allowedCountries: "de", + banResponse: "429", + limitBanDuration: "15m", + limitBanRepeatWindow: "2d", + maxBanDuration: "30d", + maxBans: "100", + banScopeV4Prefix: "24", + stateDir: "/srv/waf-state", + stateWriteDelay: "500ms", + stateCounterInterval: "1h", + metricsToken: token, + metricsTopN: "10", + }) + + wantSettings(t, cfg, config.Config{ + ListenAddr: "127.0.0.1:9000", + Observe: true, + ClientRequestTimeout: 90 * time.Second, + ClientRequestHeaderMaxBytes: 8 << 10, + ClientIdleTimeout: 5 * time.Minute, + ClientResponseTimeout: 7 * 24 * time.Hour, + UpstreamRequestTimeout: 90 * time.Minute, + UpstreamResponseTimeout: 0, + RequestMaxBytes: 512 << 10, + ResponseMaxBytes: 1234, + RateLimitPerMinute: 60, + RateLimitPerHour: 600, + RateLimitPerDay: 6000, + BanResponse: 429, + LimitBanDuration: 15 * time.Minute, + LimitBanRepeatWindow: 48 * time.Hour, + MaxBanDuration: 30 * 24 * time.Hour, + MaxBans: 100, + BanScopeV4Prefix: 24, + StateDir: "/srv/waf-state", + StateWriteDelay: 500 * time.Millisecond, + StateCounterInterval: time.Hour, + MetricsToken: token, + MetricsTopN: 10, + }) + + if cfg.UpstreamURL.String() != "https://app.internal:8443/" { + t.Errorf("%s is %s", upstreamURL, cfg.UpstreamURL) + } + + wantNetblocks(t, cfg.TrustedProxies, "192.0.2.1/32", "10.0.0.0/8", "2001:db8::/32") + wantNetblocks(t, cfg.AllowNets, "192.0.2.7/32") + wantNetblocks(t, cfg.RateLimitExemptNets, "2001:db8::/48", "10.9.8.7/32") + wantNetblocks(t, cfg.DenyNets, "198.51.100.0/24") + wantCountries(t, deniedCountries, cfg.DeniedCountries, "CN", "RU", "KP", "XK") + wantCountries(t, allowedCountries, cfg.ExclusivelyAllowedCountries, "DE") +} + +func TestInstanceNameAndLoggedHeadersAsSet(t *testing.T) { + t.Parallel() + + cfg := fromEnvironment(t, environment{ + instanceName: "fsn1app1/gitea", + logRequestHeaders: " Accept , X-Custom", + }) + + if cfg.InstanceName != "fsn1app1/gitea" || + !slices.Equal(cfg.LogRequestHeaders, []string{"accept", "x-custom"}) { + t.Errorf("%s is %q and %s gives %v", instanceName, cfg.InstanceName, + logRequestHeaders, cfg.LogRequestHeaders) + } +} + +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) { + t.Parallel() + + cfg := fromEnvironment(t, environment{ + requestMaxBytes: "3G", + responseMaxBytes: off, + clientRequestTimeout: off, + clientIdleTimeout: off, + }) + + if cfg.RequestMaxBytes != 3<<30 || cfg.ResponseMaxBytes != 0 || + cfg.ClientRequestTimeout != 0 || cfg.ClientIdleTimeout != 0 { + t.Errorf("3G, off, off and off read as %d, %d, %s and %s", + cfg.RequestMaxBytes, cfg.ResponseMaxBytes, cfg.ClientRequestTimeout, + cfg.ClientIdleTimeout) + } +} + +func TestRequestHeaderMaxBytesJustOver4K(t *testing.T) { + t.Parallel() + + cfg := fromEnvironment(t, environment{clientHeaderMaxBytes: "4097"}) + if cfg.ClientRequestHeaderMaxBytes != 4097 { + t.Errorf("4097 read as %d", cfg.ClientRequestHeaderMaxBytes) + } +} + +func TestRequestHeaderMaxBytesRefusalNeverOffersOff(t *testing.T) { + t.Parallel() + + for _, value := range []string{"32KB", "0", "4K", off} { + t.Run(value, func(t *testing.T) { + t.Parallel() + + _, err := config.FromEnvironment( + environment{clientHeaderMaxBytes: value}.lookupEnv) + + want := clientHeaderMaxBytes + `: "` + value + + `" is not a size of more than 4K, such as 32K` + if err == nil || err.Error() != want { + t.Errorf("error %v, want %s", err, want) + } + }) + } +} + +func TestRateLimitsOff(t *testing.T) { + t.Parallel() + + cfg := fromEnvironment(t, environment{ + rateLimitPerMinute: off, + rateLimitPerHour: off, + rateLimitPerDay: off, + }) + + if cfg.RateLimitPerMinute != 0 || cfg.RateLimitPerHour != 0 || + cfg.RateLimitPerDay != 0 { + t.Errorf("off read as %d, %d and %d", + cfg.RateLimitPerMinute, cfg.RateLimitPerHour, cfg.RateLimitPerDay) + } +} + +func TestBanResponseCloseIsZero(t *testing.T) { + t.Parallel() + + cfg := fromEnvironment(t, environment{banResponse: "close"}) + if cfg.BanResponse != 0 { + t.Errorf("close read as %d, want 0", cfg.BanResponse) + } +} + +func TestTrustedProxiesSetButEmptyTrustNothing(t *testing.T) { + t.Parallel() + + cfg := fromEnvironment(t, environment{trustedProxies: ""}) + if len(cfg.TrustedProxies) != 0 { + t.Errorf("trusted proxies %v, want none", cfg.TrustedProxies) + } +} + +func TestInvalidValueStopsTheStart(t *testing.T) { + t.Parallel() + + for _, tc := range []struct{ name, value string }{ + {listenAddr, "8080"}, {listenAddr, ":http"}, {listenAddr, ":65536"}, + {upstreamURL, "127.0.0.1:8081"}, + {upstreamURL, "ftp://127.0.0.1:8081"}, + {upstreamURL, "http://"}, + {upstreamURL, "http://:8081"}, + {upstreamURL, "http://127.0.0.1:0"}, + {upstreamURL, "http://127.0.0.1:99999"}, + {upstreamURL, "http://127.0.0.1:8081/app"}, + {upstreamURL, "http://127.0.0.1:8081/?a=1"}, + {upstreamURL, "http://user:secret@127.0.0.1:8081"}, + {mode, "Observe"}, {mode, "block"}, {mode, ""}, + {trustedProxies, "10.0.0.0/33"}, + {trustedProxies, "traefik"}, + {trustedProxies, "10.0.0.0/8,,192.168.0.0/16"}, + {trustedProxies, "fe80::1%eth0"}, + {allowNets, "192.0.2.0/24,monitoring"}, + {rateLimitExemptNets, "2001:db8::/129"}, + {denyNets, "198.51.100.0/24,"}, + {clientRequestTimeout, "60"}, {clientRequestTimeout, ""}, + {clientIdleTimeout, "0s"}, + {clientIdleTimeout, "2 minutes"}, + {clientResponseTimeout, "1y"}, + {upstreamRequestTimeout, "-1s"}, + {upstreamResponseTimeout, "0s"}, + {upstreamResponseTimeout, "1.5d"}, + {requestMaxBytes, "100MB"}, + {requestMaxBytes, "100m"}, + {requestMaxBytes, "1.5M"}, + {responseMaxBytes, "0"}, + {responseMaxBytes, "-5"}, + {responseMaxBytes, "99999999999G"}, + {rateLimitPerMinute, ""}, + {rateLimitPerMinute, "1K"}, + {rateLimitPerHour, "0"}, + {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"}, + {banResponse, "404"}, {banResponse, "drop"}, {banResponse, ""}, + {limitBanDuration, off}, {limitBanDuration, "0s"}, {limitBanDuration, "1"}, + {limitBanRepeatWindow, off}, {limitBanRepeatWindow, "-1h"}, + {maxBanDuration, off}, {maxBanDuration, "1w"}, + {maxBans, off}, {maxBans, "0"}, {maxBans, "5K"}, + {banScopeV4Prefix, "33"}, {banScopeV4Prefix, "-1"}, {banScopeV4Prefix, "/24"}, + {stateDir, ""}, {stateDir, "state"}, {stateDir, "./var/lib/smallwebwaf"}, + {stateWriteDelay, off}, {stateWriteDelay, "0s"}, + {stateCounterInterval, off}, {stateCounterInterval, "15"}, + {metricsTopN, off}, {metricsTopN, "0"}, {metricsTopN, "-1"}, + {logRequestHeaders, "accept,,origin"}, {logRequestHeaders, "accept;origin"}, + {logRequestHeaders, "accept language"}, {logRequestHeaders, "x-foo:"}, + {logRequestHeaders, "host"}, {logRequestHeaders, "accept,Host"}, + {logRequestHeaders, "transfer-encoding"}, {logRequestHeaders, "TRANSFER-ENCODING"}, + } { + t.Run(tc.name+"="+tc.value, func(t *testing.T) { + t.Parallel() + + _, err := config.FromEnvironment(environment{tc.name: tc.value}.lookupEnv) + if err == nil { + t.Fatalf("%s=%q was accepted", tc.name, tc.value) + } + + if !strings.HasPrefix(err.Error(), tc.name+": ") { + t.Errorf("error %q does not name %s", err, tc.name) + } + }) + } +} + +func TestHostOrTransferEncodingStopsTheStart(t *testing.T) { + t.Parallel() + + // Only Host's message points to the field host. + for value, want := range map[string]string{ + "Host": `"Host" is taken out of every request by Go's HTTP server, ` + + "so it can never be logged; the request's host is the field host", + "transfer-encoding": `"transfer-encoding" is taken out of every ` + + "request by Go's HTTP server, so it can never be logged", + } { + t.Run(value, func(t *testing.T) { + t.Parallel() + + _, err := config.FromEnvironment(environment{logRequestHeaders: value}.lookupEnv) + if err == nil || err.Error() != logRequestHeaders+": "+want { + t.Errorf("error %v, want %s: %s", err, logRequestHeaders, want) + } + }) + } +} + +func TestShortTokenStopsTheStartWithoutShowingIt(t *testing.T) { + t.Parallel() + + // Characters are counted, not bytes: each é takes two. + for _, value := range []string{"", token[1:], strings.Repeat("é", 31)} { + t.Run(value, func(t *testing.T) { + t.Parallel() + + _, err := config.FromEnvironment(environment{metricsToken: value}.lookupEnv) + + want := metricsToken + ": is shorter than 32 characters" + if err == nil || err.Error() != want { + t.Errorf("error %v, want %s", err, want) + } + }) + } +} + +func TestTokenIsLoggedMasked(t *testing.T) { + t.Parallel() + + cfg := fromEnvironment(t, environment{metricsToken: token}) + + var out bytes.Buffer + + slog.New(slog.NewJSONHandler(&out, nil)).Info("starting", "settings", cfg) + + if strings.Contains(out.String(), token) || + !strings.Contains(out.String(), `"`+metricsToken+`":"********"`) { + t.Errorf("the token is not logged masked: %s", out.String()) + } +} + +func TestLogsEachSettingWithItsValue(t *testing.T) { + t.Parallel() + + cfg := fromEnvironment(t, environment{clientRequestTimeout: "45s"}) + + var out bytes.Buffer + + slog.New(slog.NewJSONHandler(&out, nil)).Info("starting", "settings", cfg) + + var line struct { + Settings map[string]string `json:"settings"` + } + + err := json.Unmarshal(out.Bytes(), &line) + if err != nil { + t.Fatalf("decode %s: %v", out.Bytes(), err) + } + + hostname, _ := os.Hostname() + + want := map[string]string{ + listenAddr: ":8080", + upstreamURL: "http://127.0.0.1:8081", + mode: "enforce", + trustedProxies: "10.0.0.0/8,172.16.0.0/12,192.168.0.0/16", + clientRequestTimeout: "45s", + clientHeaderMaxBytes: "32K", + clientIdleTimeout: "120s", + clientResponseTimeout: "30m", + upstreamRequestTimeout: "60s", + upstreamResponseTimeout: "30m", + requestMaxBytes: "100M", + responseMaxBytes: "5G", + allowNets: "", + rateLimitExemptNets: "", + denyNets: "", + rateLimitPerMinute: "1000", + rateLimitPerHour: "10000", + rateLimitPerDay: "50000", + deniedCountries: "", + allowedCountries: "", + banResponse: "403", + limitBanDuration: "1h", + limitBanRepeatWindow: "24h", + maxBanDuration: "7d", + maxBans: "5000", + banScopeV4Prefix: "32", + stateDir: "/var/lib/smallwebwaf", + stateWriteDelay: "10s", + stateCounterInterval: "15m", + metricsToken: "", + metricsTopN: "50", + instanceName: hostname, + logRequestHeaders: defaultLogRequestHeaders, + } + if !maps.Equal(line.Settings, want) { + t.Errorf("logged settings\n%v\nwant\n%v", line.Settings, want) + } +} + +// wantSettings checks the settings that are plain values. +func wantSettings(t *testing.T, got *config.Config, want config.Config) { + t.Helper() + + if got.ListenAddr != want.ListenAddr || + got.Observe != want.Observe || + got.ClientRequestTimeout != want.ClientRequestTimeout || + got.ClientRequestHeaderMaxBytes != want.ClientRequestHeaderMaxBytes || + got.ClientIdleTimeout != want.ClientIdleTimeout || + got.ClientResponseTimeout != want.ClientResponseTimeout || + got.UpstreamRequestTimeout != want.UpstreamRequestTimeout || + got.UpstreamResponseTimeout != want.UpstreamResponseTimeout || + got.RequestMaxBytes != want.RequestMaxBytes || + got.ResponseMaxBytes != want.ResponseMaxBytes || + got.RateLimitPerMinute != want.RateLimitPerMinute || + got.RateLimitPerHour != want.RateLimitPerHour || + got.RateLimitPerDay != want.RateLimitPerDay { + t.Errorf("settings\n%+v\nwant\n%+v", got, want) + } + + wantBanSettings(t, got, want) +} + +// wantBanSettings checks the settings for bans, the state files and the +// metrics. +func wantBanSettings(t *testing.T, got *config.Config, want config.Config) { + t.Helper() + + if got.BanResponse != want.BanResponse || + got.LimitBanDuration != want.LimitBanDuration || + got.LimitBanRepeatWindow != want.LimitBanRepeatWindow || + got.MaxBanDuration != want.MaxBanDuration || + got.MaxBans != want.MaxBans || + got.BanScopeV4Prefix != want.BanScopeV4Prefix { + t.Errorf("ban settings\n%+v\nwant\n%+v", got, want) + } + + if got.StateDir != want.StateDir || + got.StateWriteDelay != want.StateWriteDelay || + got.StateCounterInterval != want.StateCounterInterval { + t.Errorf("state settings\n%+v\nwant\n%+v", got, want) + } + + if got.MetricsToken != want.MetricsToken || got.MetricsTopN != want.MetricsTopN { + t.Errorf("metrics token %q and top %d, want %q and %d", + got.MetricsToken, got.MetricsTopN, want.MetricsToken, want.MetricsTopN) + } +} + +// wantNetblocks checks a list of netblocks. +func wantNetblocks(t *testing.T, got []netip.Prefix, want ...string) { + t.Helper() + + gotText := make([]string, 0, len(got)) + for _, netblock := range got { + gotText = append(gotText, netblock.String()) + } + + if !slices.Equal(gotText, want) { + t.Errorf("netblocks %v, want %v", gotText, want) + } +} + +// 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/export_test.go b/internal/lookup/export_test.go new file mode 100644 index 0000000..3276a6c --- /dev/null +++ b/internal/lookup/export_test.go @@ -0,0 +1,9 @@ +package lookup + +import "net/http" + +// SetTransport has g's requests to GeoJS go through transport instead of +// the network. +func (g *GeoJS) SetTransport(transport http.RoundTripper) { + g.httpClient.Transport = transport +} diff --git a/internal/lookup/lookup.go b/internal/lookup/lookup.go new file mode 100644 index 0000000..a2b7139 --- /dev/null +++ b/internal/lookup/lookup.go @@ -0,0 +1,456 @@ +// 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. The answers are written to lookups.json and read +// from it by the state package. +package lookup + +import ( + "context" + "encoding/json" + "errors" + "fmt" + "io" + "log/slog" + "net/http" + "net/netip" + "slices" + "strings" + "sync" + "time" + + "github.com/hashicorp/golang-lru/v2/simplelru" + "sneak.berlin/go/smallwebwaf/internal/metrics" +) + +// 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 + // Metrics count the requests to GeoJS, those that failed, and the + // clients that go without an answer. + Metrics *metrics.Metrics +} + +// 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 + metrics *metrics.Metrics + // 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, as lookups.json holds it: its +// country, "" when GeoJS cannot place it, when GeoJS said so, and when +// the answer was last used. +type Answer struct { + Client netip.Prefix `json:"client"` + Country string `json:"country"` + Answered time.Time `json:"answered"` + Used time.Time `json:"used"` +} + +// 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, + metrics: params.Metrics, + 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) + if !found { + g.metrics.GeoJSUnanswered.Inc() + } + + w, waiting := g.waiting[client] + if !found && waiting { + w.late = true + } + + return country +} + +// Snapshot returns every answer kept, sorted by client, as lookups.json +// lists them. +func (g *GeoJS) Snapshot() []Answer { + g.mu.Lock() + + answers := make([]Answer, 0, g.answers.Len()) + for _, kept := range g.answers.Values() { + answers = append(answers, *kept) + } + + g.mu.Unlock() + + slices.SortFunc(answers, func(a, b Answer) int { + return a.Client.Compare(b.Client) + }) + + return answers +} + +// Load keeps answers read from lookups.json, in place of the answers it +// keeps, in the order they were last used, so that the one used longest +// ago is dropped first. Answers GeoJS gave keepFor ago or more are +// dropped. +func (g *GeoJS) Load(answers []Answer) { + answers = slices.Clone(answers) + slices.SortStableFunc(answers, func(a, b Answer) int { + return a.Used.Compare(b.Used) + }) + + g.mu.Lock() + defer g.mu.Unlock() + + g.answers.Purge() + + now := g.now() + + for _, answer := range answers { + if now.Sub(answer.Answered) < keepFor { + g.answers.Add(answer.Client, &answer) + } + } +} + +// 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 { + g.metrics.GeoJSUnanswered.Inc() + + 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 { + g.metrics.GeoJSUnanswered.Inc() + + return "", nil + } + + return "", w.asked +} + +// kept returns client's answer, if GeoJS gave it less than keepFor ago, +// and notes that it was used. +func (g *GeoJS) kept(client netip.Prefix) (string, bool) { + now := g.now() + + kept, found := g.answers.Get(client) + if !found || now.Sub(kept.Answered) >= keepFor { + return "", false + } + + kept.Used = now + + 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{ + Client: client, Country: country, Answered: now, Used: 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.metrics.GeoJSFailures.Inc() + + 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, ",") + + g.metrics.GeoJSRequests.Inc() + + 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..4d718fc --- /dev/null +++ b/internal/lookup/lookup_test.go @@ -0,0 +1,631 @@ +package lookup_test + +import ( + "encoding/json" + "log/slog" + "net/http" + "net/http/httptest" + "net/netip" + "slices" + "strings" + "sync" + "testing" + "testing/synctest" + "time" + + "github.com/prometheus/client_golang/prometheus/testutil" + "sneak.berlin/go/smallwebwaf/internal/lookup" + "sneak.berlin/go/smallwebwaf/internal/metrics" +) + +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 + // week is how long an answer is kept. + week = 7 * 24 * time.Hour +) + +// The tests that have GeoJS asked run in a synctest bubble, where the time +// package runs on a clock of the test's own: a wait lasts exactly as long +// as it should, however slowly the test process runs, and synctest.Wait +// returns once g has done all it can before time passes. The stand-in for +// GeoJS answers without the network, since a request waiting on the +// network would keep that clock from moving on. + +func TestKeptAnswerIsUsedFor7DaysThenAskedAgain(t *testing.T) { + t.Parallel() + + synctest.Test(t, func(t *testing.T) { + geojs, clock, g := start() + 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() + + synctest.Test(t, func(t *testing.T) { + geojs, clock, g := start() + 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 { + 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 != 0 { + 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() + + synctest.Test(t, func(t *testing.T) { + geojs, clock, g := start() + 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() + + synctest.Test(t, func(t *testing.T) { + geojs, _, g := start() + geojs.set(redirecting) + + wantCountry(t, g, netip.MustParsePrefix("203.0.113.9/32"), "") + wantRequests(t, geojs, 1) + }) +} + +func TestCountryIsKeptInCapitals(t *testing.T) { + t.Parallel() + + synctest.Test(t, func(t *testing.T) { + geojs, _, g := start() + geojs.set(answeringInLowerCase) + + wantCountry(t, g, netip.MustParsePrefix("203.0.113.9/32"), germany) + }) +} + +func TestFailureIsLoggedWithoutTheAddressesAskedAbout(t *testing.T) { + t.Parallel() + + synctest.Test(t, func(t *testing.T) { + var log strings.Builder + + // GeoJS does not answer, so the request to it is abandoned, and fails. + geojs := &standIn{answers: hanging} + g := lookup.New(lookup.Params{ + URL: lookup.URL, + Now: time.Now, + ProcessLog: slog.New(slog.NewTextHandler(&log, nil)), + Metrics: metrics.New(1), + }) + g.SetTransport(geojs) + + wantCountry(t, g, netip.MustParsePrefix("203.0.113.9/32"), "") + synctest.Wait() + + 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() + + synctest.Test(t, func(t *testing.T) { + geojs, clock, g := start() + + // 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() + + synctest.Test(t, func(t *testing.T) { + geojs, clock, g := start() + 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() + + synctest.Test(t, func(t *testing.T) { + geojs, clock, g := start() + 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() + + synctest.Test(t, func(t *testing.T) { + geojs, clock, g := start() + 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) + }) +} + +func TestClientsWithoutAnAnswerAreCounted(t *testing.T) { + t.Parallel() + + synctest.Test(t, func(t *testing.T) { + m := metrics.New(1) + g := lookup.New(lookup.Params{ + URL: lookup.URL, + Now: time.Now, + ProcessLog: slog.New(slog.DiscardHandler), + Metrics: m, + }) + g.SetTransport(&standIn{answers: failing}) + + clients := newClients() + + // GeoJS fails, so the first client goes without an answer, and GeoJS + // is left alone for a second, which does not pass in this test. + wantCountry(t, g, clients(), "") + wantUnanswered(t, m, 1) + + // Meanwhile each new client goes without one at once, while there is + // room for it among the 10,000 that may wait. + for range 9999 { + wantCountry(t, g, clients(), "") + } + + wantUnanswered(t, m, 10000) + + // One more, for which there is no room, goes without one too. + wantCountry(t, g, clients(), "") + wantUnanswered(t, m, 10001) + }) +} + +// 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 { + mu sync.Mutex + answers int + requests [][]string +} + +// RoundTrip has the stand-in answer req, in place of the network. A request +// abandoned before the stand-in answers fails, as over the network. +func (s *standIn) RoundTrip(req *http.Request) (*http.Response, error) { + answer := httptest.NewRecorder() + s.ServeHTTP(answer, req) + + err := req.Context().Err() + if err != nil { + return nil, err + } + + return answer.Result(), nil +} + +// 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. GeoJS tells the time by it, while +// waits run on the bubble's clock. +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 returns a stand-in for GeoJS that answers, a clock, and a GeoJS +// asking the stand-in by that clock. +func start() (*standIn, *testClock, *lookup.GeoJS) { + geojs := &standIn{} + clock := &testClock{now: time.Date(2026, 10, 4, 0, 0, 0, 0, time.UTC)} + g := lookup.New(lookup.Params{ + URL: lookup.URL, + Now: clock.Now, + ProcessLog: slog.New(slog.DiscardHandler), + Metrics: metrics.New(1), + }) + g.SetTransport(geojs) + + 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) + } +} + +// wantUnanswered checks how many requests m counts as having gone without +// an answer from GeoJS. +func wantUnanswered(t *testing.T, m *metrics.Metrics, want float64) { + t.Helper() + + got := testutil.ToFloat64(m.GeoJSUnanswered) + if got != want { + t.Errorf("%v requests went without an answer, want %v", got, want) + } +} + +// waitForRequests waits until g has done all it can before time passes, +// checks that GeoJS has had count requests, and returns the addresses each +// asked about. +func waitForRequests(t *testing.T, geojs *standIn, count int) [][]string { + t.Helper() + + synctest.Wait() + + asked := geojs.asked() + if len(asked) != count { + t.Fatalf("GeoJS had %d requests, want %d", len(asked), count) + } + + return asked +} + +// waitForCountry lets a request to GeoJS under way be abandoned, and moves +// the clock on a minute, so that GeoJS may be asked again after a failure. +// It then checks that client's next request does not wait but has it asked +// about again in the background, after which g gives it the country want. +func waitForCountry( + t *testing.T, g *lookup.GeoJS, clock *testClock, client netip.Prefix, want string, +) { + t.Helper() + + time.Sleep(timeout) + clock.advance(time.Minute) + wantCountry(t, g, client, "") + synctest.Wait() + wantCountry(t, g, client, want) +} diff --git a/internal/lookup/snapshot_test.go b/internal/lookup/snapshot_test.go new file mode 100644 index 0000000..eb3b2df --- /dev/null +++ b/internal/lookup/snapshot_test.go @@ -0,0 +1,99 @@ +package lookup_test + +import ( + "net/netip" + "slices" + "testing" + "testing/synctest" + "time" + + "sneak.berlin/go/smallwebwaf/internal/lookup" +) + +func TestSnapshotHoldsEachAnswerAndWhenItWasLastUsed(t *testing.T) { + t.Parallel() + + synctest.Test(t, func(t *testing.T) { + _, clock, g := start() + placed := netip.MustParsePrefix("203.0.113.9/32") + notPlaced := netip.MustParsePrefix(unplaced + "/32") + asked := clock.Now() + + wantCountry(t, g, placed, germany) + wantCountry(t, g, notPlaced, "") + + clock.advance(time.Hour) + wantCountry(t, g, placed, germany) + + want := []lookup.Answer{ + {Client: notPlaced, Country: "", Answered: asked, Used: asked}, + {Client: placed, Country: germany, Answered: asked, Used: asked.Add(time.Hour)}, + } + if got := g.Snapshot(); !slices.Equal(got, want) { + t.Errorf("snapshot\n%+v\nwant\n%+v", got, want) + } + }) +} + +func TestLoadedAnswersAreKeptFor7DaysFromWhenGeoJSGaveThem(t *testing.T) { + t.Parallel() + + geojs, clock, g := start() + now := clock.Now() + kept := lookup.Answer{ + Client: netip.MustParsePrefix("203.0.113.9/32"), + Country: "FR", + Answered: now.Add(-week + time.Second), + Used: now.Add(-time.Hour), + } + stale := lookup.Answer{ + Client: netip.MustParsePrefix("203.0.113.10/32"), + Country: "FR", + Answered: now.Add(-week), + Used: now.Add(-time.Hour), + } + + g.Load([]lookup.Answer{kept, stale}) + + if got := g.Snapshot(); !slices.Equal(got, []lookup.Answer{kept}) { + t.Errorf("kept %+v, want only the answer GeoJS gave less than 7 days ago", got) + } + + wantCountry(t, g, kept.Client, "FR") + wantRequests(t, geojs, 0) +} + +func TestLoadDropsTheAnswerUsedLongestAgoFirst(t *testing.T) { + t.Parallel() + + const maxAnswers = 100000 + + _, clock, g := start() + now := clock.Now() + + // lookups.json lists the answers by client. Here each was last used a + // second before the one listed before it, so the last listed is the + // one used longest ago, and the one dropped. + answers := make([]lookup.Answer, maxAnswers+1) + addr := netip.MustParseAddr("10.0.0.0") + + for i := range answers { + answers[i] = lookup.Answer{ + Client: netip.PrefixFrom(addr, addr.BitLen()), + Country: germany, + Answered: now, + Used: now.Add(-time.Duration(i) * time.Second), + } + addr = addr.Next() + } + + g.Load(answers) + + got := g.Snapshot() + if len(got) != maxAnswers || got[0] != answers[0] || + got[maxAnswers-1] != answers[maxAnswers-1] { + t.Errorf("%d answers kept, from %s to %s; want %d, from %s to %s", + len(got), got[0].Client, got[len(got)-1].Client, maxAnswers, + answers[0].Client, answers[maxAnswers-1].Client) + } +} diff --git a/internal/metrics/countries.go b/internal/metrics/countries.go new file mode 100644 index 0000000..27cf010 --- /dev/null +++ b/internal/metrics/countries.go @@ -0,0 +1,116 @@ +package metrics + +import ( + "sync" + + "github.com/prometheus/client_golang/prometheus" + "sneak.berlin/go/smallwebwaf/internal/requestlog" +) + +// other is the label under which the countries outside the busiest are +// counted. +const other = "other" + +// countries are the metrics by the client's country, for requests whose +// client's country is known. The topN busiest countries, by their requests +// since the start, have series of their own, and the others are counted +// under other, so that there are never more than topN + 1 series. A +// country that drops out of the busiest loses its series, and its next +// requests are counted under other; one that becomes one of them gets a +// series that counts from then on. Each series therefore only ever goes +// up. +type countries struct { + topN int + + requests *prometheus.CounterVec + requestBytes *prometheus.CounterVec + responseBytes *prometheus.CounterVec + // refused are the requests the country lists refused. + refused *prometheus.CounterVec + + mu sync.Mutex + // seen is each country's requests since the start, by which the + // countries are ranked. GeoJS gives two-letter codes, so it holds at + // most a few hundred. + seen map[string]int64 + // top are the countries with series of their own. + top map[string]bool +} + +// newCountries returns the metrics by country, with series of their own +// for the topN busiest countries. +func newCountries(topN int) *countries { + byCountry := []string{"country"} + + return &countries{ + topN: topN, + requests: counterVec("smallwebwaf_country_requests_total", + "Requests, by the client's country.", byCountry), + requestBytes: counterVec("smallwebwaf_country_request_bytes_total", + "Request body bytes, by the client's country.", byCountry), + responseBytes: counterVec("smallwebwaf_country_response_bytes_total", + "Response body bytes, by the client's country.", byCountry), + refused: counterVec("smallwebwaf_country_list_refusals_total", + "Requests the country lists refused, by the client's country.", + byCountry), + seen: map[string]int64{}, + top: map[string]bool{}, + } +} + +// add counts a request from its log line, whose country is known. +func (c *countries) add(line *requestlog.Line) { + c.mu.Lock() + defer c.mu.Unlock() + + c.seen[line.Country]++ + + label := c.label(line.Country) + c.requests.WithLabelValues(label).Inc() + c.requestBytes.WithLabelValues(label).Add(float64(line.RequestBytes)) + c.responseBytes.WithLabelValues(label).Add(float64(line.ResponseBytes)) + + if line.Action == requestlog.ActionCountryDenied { + c.refused.WithLabelValues(label).Inc() + } +} + +// label returns the label a request from country is counted under: the +// country while it is one of the busiest, other while it is not. A +// country busier than the least busy of them takes its place, and that +// country's series are dropped. +func (c *countries) label(country string) string { + if c.top[country] { + return country + } + + if len(c.top) < c.topN { + c.top[country] = true + + return country + } + + least := "" + + for top := range c.top { + if least == "" || c.seen[top] < c.seen[least] { + least = top + } + } + + if c.seen[country] <= c.seen[least] { + return other + } + + delete(c.top, least) + + for _, vec := range []*prometheus.CounterVec{ + c.requests, c.requestBytes, c.responseBytes, c.refused, + } { + vec.DeleteLabelValues(least) + } + + c.top[country] = true + + return country +} diff --git a/internal/metrics/metrics.go b/internal/metrics/metrics.go new file mode 100644 index 0000000..f08a714 --- /dev/null +++ b/internal/metrics/metrics.go @@ -0,0 +1,277 @@ +// Package metrics keeps smallwebwaf's Prometheus metrics, as the "Metrics +// endpoint" section of SPEC.md lists them, and serves them in the +// Prometheus text format. No metric carries a client's address. +package metrics + +import ( + "net/http" + "strconv" + "time" + + "github.com/prometheus/client_golang/prometheus" + "github.com/prometheus/client_golang/prometheus/collectors" + "github.com/prometheus/client_golang/prometheus/promhttp" + "sneak.berlin/go/smallwebwaf/internal/bans" + "sneak.berlin/go/smallwebwaf/internal/ratelimit" + "sneak.berlin/go/smallwebwaf/internal/requestlog" +) + +// Metrics are smallwebwaf's metrics. They are safe for concurrent use. +type Metrics struct { + registry *prometheus.Registry + handler http.Handler + + inFlight prometheus.Gauge + requests *prometheus.CounterVec + requestBytes *prometheus.CounterVec + responseBytes *prometheus.CounterVec + requestDuration prometheus.Histogram + upstreamDuration prometheus.Histogram + rateLimitHits *prometheus.CounterVec + sizeAndTimeLimitHits *prometheus.CounterVec + offences *prometheus.CounterVec + countries *countries + + // GeoJSRequests are the requests to GeoJS, and GeoJSFailures those + // that failed. GeoJSUnanswered are the requests whose client counted + // as coming from an unknown country because GeoJS had not answered + // about it in time. + GeoJSRequests prometheus.Counter + GeoJSFailures prometheus.Counter + GeoJSUnanswered prometheus.Counter + + stateFileWrites *prometheus.CounterVec + stateFileWriteFailures *prometheus.CounterVec + stateFileLastWrite *prometheus.GaugeVec + stateFileSize *prometheus.GaugeVec + stateFileEditsTakenIn *prometheus.CounterVec + stateFileEditsSetAside *prometheus.CounterVec +} + +// New returns the metrics, with the Go runtime's and the process's own. +// topN is how many countries get series of their own +// (SWWAF_METRICS_TOP_N). +func New(topN int) *Metrics { + byStatus := []string{"status_class", "action"} + byFile := []string{"file"} + + m := &Metrics{ + registry: prometheus.NewRegistry(), + inFlight: prometheus.NewGauge(prometheus.GaugeOpts{ + Name: "smallwebwaf_requests_in_flight", + Help: "Requests under way.", + }), + requests: counterVec("smallwebwaf_requests_total", + "Requests, by the class of their status and their action.", byStatus), + requestBytes: counterVec("smallwebwaf_request_bytes_total", + "Request body bytes, by the class of the status and the action.", + byStatus), + responseBytes: counterVec("smallwebwaf_response_bytes_total", + "Response body bytes, by the class of the status and the action.", + byStatus), + requestDuration: prometheus.NewHistogram(prometheus.HistogramOpts{ + Name: "smallwebwaf_request_duration_seconds", + Help: "How long requests took, from their arrival to their end.", + }), + upstreamDuration: prometheus.NewHistogram(prometheus.HistogramOpts{ + Name: "smallwebwaf_upstream_duration_seconds", + Help: "How long requests passed to the app took, from then to their end.", + }), + rateLimitHits: counterVec("smallwebwaf_rate_limit_hits_total", + "Requests that broke a rate limit, by its window.", + []string{"window"}), + sizeAndTimeLimitHits: counterVec("smallwebwaf_size_and_time_limit_hits_total", + "Requests that passed a size or time limit, by its setting.", + []string{"limit"}), + offences: counterVec("smallwebwaf_offences_total", + "Offences, by kind.", []string{"kind"}), + countries: newCountries(topN), + GeoJSRequests: prometheus.NewCounter(prometheus.CounterOpts{ + Name: "smallwebwaf_geojs_requests_total", + Help: "Requests to GeoJS.", + }), + GeoJSFailures: prometheus.NewCounter(prometheus.CounterOpts{ + Name: "smallwebwaf_geojs_failures_total", + Help: "Requests to GeoJS that failed.", + }), + GeoJSUnanswered: prometheus.NewCounter(prometheus.CounterOpts{ + Name: "smallwebwaf_geojs_unanswered_total", + Help: "Requests whose client counted as coming from an unknown " + + "country because GeoJS had not answered about it in time.", + }), + stateFileWrites: counterVec("smallwebwaf_state_file_writes_total", + "Writes of each state file.", byFile), + stateFileWriteFailures: counterVec("smallwebwaf_state_file_write_failures_total", + "Writes of each state file that failed.", byFile), + stateFileLastWrite: gaugeVec("smallwebwaf_state_file_last_write_timestamp_seconds", + "When each state file was last written, in seconds since 1970.", byFile), + stateFileSize: gaugeVec("smallwebwaf_state_file_size_bytes", + "The size of each state file, as it was last written.", byFile), + stateFileEditsTakenIn: counterVec("smallwebwaf_state_file_edits_taken_in_total", + "Edits of each state file taken in while running.", byFile), + stateFileEditsSetAside: counterVec("smallwebwaf_state_file_edits_set_aside_total", + "Edits of each state file renamed to .bad because they did not parse.", + byFile), + } + + m.handler = promhttp.HandlerFor(m.registry, promhttp.HandlerOpts{}) + + m.registry.MustRegister( + collectors.NewGoCollector(), + collectors.NewProcessCollector(collectors.ProcessCollectorOpts{}), + m.inFlight, m.requests, m.requestBytes, m.responseBytes, + m.requestDuration, m.upstreamDuration, + m.rateLimitHits, m.sizeAndTimeLimitHits, m.offences, + m.countries.requests, m.countries.requestBytes, m.countries.responseBytes, + m.countries.refused, + m.GeoJSRequests, m.GeoJSFailures, m.GeoJSUnanswered, + m.stateFileWrites, m.stateFileWriteFailures, + m.stateFileLastWrite, m.stateFileSize, + m.stateFileEditsTakenIn, m.stateFileEditsSetAside, + ) + + return m +} + +// AddBansAndClients adds the metrics read from the ledger and the table +// of clients as the metrics are asked for: the bans made since the start, +// the bans active and permanent at now, and the clients in the table. +func (m *Metrics) AddBansAndClients( + ledger *bans.Ledger, limiter *ratelimit.Limiter, now func() time.Time, +) { + m.registry.MustRegister( + // Every ban smallwebwaf makes so far is for a broken limit. + prometheus.NewCounterFunc(prometheus.CounterOpts{ + Name: "smallwebwaf_bans_made_total", + Help: "Bans made, by cause.", + ConstLabels: prometheus.Labels{"cause": "limit"}, + }, func() float64 { + return float64(ledger.Made()) + }), + prometheus.NewGaugeFunc(prometheus.GaugeOpts{ + Name: "smallwebwaf_active_bans", + Help: "Bans active now, the permanent ones included.", + }, func() float64 { + active, _ := ledger.Count(now()) + + return float64(active) + }), + prometheus.NewGaugeFunc(prometheus.GaugeOpts{ + Name: "smallwebwaf_permanent_bans", + Help: "Permanent bans.", + }, func() float64 { + _, permanent := ledger.Count(now()) + + return float64(permanent) + }), + prometheus.NewGaugeFunc(prometheus.GaugeOpts{ + Name: "smallwebwaf_tracked_clients", + Help: "Clients in the table of clients.", + }, func() float64 { + return float64(limiter.Len()) + }), + ) +} + +// ServeHTTP answers with the metrics in the Prometheus text format. +func (m *Metrics) ServeHTTP(w http.ResponseWriter, r *http.Request) { + m.handler.ServeHTTP(w, r) +} + +// RequestStarted counts a request as under way. +func (m *Metrics) RequestStarted() { + m.inFlight.Inc() +} + +// RequestEnded counts a request that has ended, from its log line. limit +// is the setting whose size or time limit the request passed, "" if none. +// duration is how long the request took, and upstreamDuration how long it +// took from when it was passed to the app, zero if it was not. +func (m *Metrics) RequestEnded( + line *requestlog.Line, limit string, duration, upstreamDuration time.Duration, +) { + m.inFlight.Dec() + + class := statusClass(line.Status) + m.requests.WithLabelValues(class, line.Action).Inc() + m.requestBytes.WithLabelValues(class, line.Action).Add(float64(line.RequestBytes)) + m.responseBytes.WithLabelValues(class, line.Action).Add(float64(line.ResponseBytes)) + m.requestDuration.Observe(duration.Seconds()) + + if upstreamDuration > 0 { + m.upstreamDuration.Observe(upstreamDuration.Seconds()) + } + + if line.LimitHit != "" { + m.rateLimitHits.WithLabelValues(line.LimitHit).Inc() + } + + if limit != "" { + m.sizeAndTimeLimitHits.WithLabelValues(limit).Inc() + } + + if line.Offence != "" { + m.offences.WithLabelValues(line.Offence).Inc() + } + + if line.Country != "" { + m.countries.add(line) + } +} + +// StateFileWritten counts a write of the state file name, of size bytes, +// that ended with err. +func (m *Metrics) StateFileWritten(name string, size int, err error) { + m.stateFileWrites.WithLabelValues(name).Inc() + + // The series of failures is there from the first write, at zero until + // one fails. + failures := m.stateFileWriteFailures.WithLabelValues(name) + + if err != nil { + failures.Inc() + + return + } + + m.stateFileLastWrite.WithLabelValues(name).SetToCurrentTime() + m.stateFileSize.WithLabelValues(name).Set(float64(size)) +} + +// StateFileEditTakenIn counts an admin's edit of the state file name +// taken in while smallwebwaf runs. +func (m *Metrics) StateFileEditTakenIn(name string) { + m.stateFileEditsTakenIn.WithLabelValues(name).Inc() +} + +// StateFileEditSetAside counts an admin's edit of the state file name +// renamed to name.bad because it did not parse. +func (m *Metrics) StateFileEditSetAside(name string) { + m.stateFileEditsSetAside.WithLabelValues(name).Inc() +} + +// statusClass returns the class of status, such as 2xx, or none when no +// status was sent. +func statusClass(status int) string { + if status == 0 { + return "none" + } + + // A status's class is its hundreds: 404 is in 4xx. + const hundred = 100 + + return strconv.Itoa(status/hundred) + "xx" +} + +// counterVec returns a counter named name, described by help, with a +// series for each set of values of labels. +func counterVec(name, help string, labels []string) *prometheus.CounterVec { + return prometheus.NewCounterVec(prometheus.CounterOpts{Name: name, Help: help}, + labels) +} + +// gaugeVec returns a gauge named name, described by help, with a series +// for each set of values of labels. +func gaugeVec(name, help string, labels []string) *prometheus.GaugeVec { + return prometheus.NewGaugeVec(prometheus.GaugeOpts{Name: name, Help: help}, labels) +} diff --git a/internal/proxy/admin.go b/internal/proxy/admin.go new file mode 100644 index 0000000..7075ffc --- /dev/null +++ b/internal/proxy/admin.go @@ -0,0 +1,43 @@ +package proxy + +import ( + "crypto/subtle" + "net/http" + "strings" + + "sneak.berlin/go/smallwebwaf/internal/requestlog" +) + +// answerAdmin answers a request for smallwebwaf itself, under +// /_smallwebwaf/, once it has passed the checks: GET MetricsPath with +// SWWAF_METRICS_TOKEN gets the metrics, and without it is refused with +// 401. Any other request gets 404, as the metrics do while +// SWWAF_METRICS_TOKEN is unset. +func (rq *request) answerAdmin() { + rq.line.Action = requestlog.ActionAdmin + rq.startClientResponseTimeout() + + token := rq.h.config.MetricsToken + + switch { + case token == "" || rq.in.Method != http.MethodGet || rq.in.URL.Path != MetricsPath: + http.Error(rq.out, http.StatusText(http.StatusNotFound), http.StatusNotFound) + case !hasToken(rq.in, token): + rq.out.Header().Set("WWW-Authenticate", "Bearer") + rq.answer(refusal{ + status: http.StatusUnauthorized, + action: requestlog.ActionAdmin, + }) + default: + rq.h.metrics.ServeHTTP(rq.out, rq.in) + } +} + +// hasToken reports whether r carries token, as Authorization: Bearer +// . +func hasToken(r *http.Request, token string) bool { + scheme, sent, _ := strings.Cut(r.Header.Get("Authorization"), " ") + + return strings.EqualFold(scheme, "Bearer") && + subtle.ConstantTimeCompare([]byte(sent), []byte(token)) == 1 +} diff --git a/internal/proxy/bans.go b/internal/proxy/bans.go new file mode 100644 index 0000000..96abb66 --- /dev/null +++ b/internal/proxy/bans.go @@ -0,0 +1,98 @@ +package proxy + +import ( + "net/netip" + "time" + + "sneak.berlin/go/smallwebwaf/internal/bans" + "sneak.berlin/go/smallwebwaf/internal/requestlog" +) + +// banResponse is a refusal answered with SWWAF_BAN_RESPONSE, and logged +// with action. +func (rq *request) banResponse(action string) *refusal { + return &refusal{status: rq.h.config.BanResponse, action: action} +} + +// banned reports whether a ban on a netblock the client is in covers the +// request at now, and notes for the log line when that ban ends. +func (rq *request) banned(now time.Time) bool { + check := rq.h.ledger.Check + if rq.h.config.Observe { + check = rq.h.ledger.Find // in observe mode the ban refuses nothing + } + + ban, banned := check(rq.client, now) + if banned { + rq.line.BanExpires = banExpires(ban) + } + + return banned +} + +// limitBroken counts the request for the rate limits at now, notes the +// client's counts for the log line, and reports whether the request takes +// the client over a limit. In enforce mode such a request bans the +// client's netblock, and sets the client's counters back to zero; in +// observe mode it does neither. +func (rq *request) limitBroken(now time.Time) bool { + group := clientGroup(rq.client) + + counts, hit, over := rq.h.limiter.Count(group, now) + rq.line.Counts = counts + + if !over { + return false + } + + rq.line.LimitHit = hit.Window + rq.line.Offence = requestlog.OffenceLimit + + if rq.h.config.Observe { + return true + } + + netblock := rq.netblock() + ban := rq.h.ledger.BanForLimit(netblock, now, bans.Notes{ + Country: rq.line.Country, + Limit: hit.Limit, + Window: hit.Window, + Count: hit.Requests, + Request: bans.Request{ + Time: now, + Method: rq.in.Method, + Host: rq.in.Host, + Path: rq.in.URL.RequestURI(), + Status: rq.h.config.BanResponse, + UserAgent: rq.in.UserAgent(), + }, + // The histories count this request only once it has ended. + Requests: rq.h.limiter.Requests(netblock) + 1, + }) + rq.h.limiter.Reset(group) + rq.line.BanExpires = banExpires(ban) + + return true +} + +// netblock is the netblock a ban on the client covers: its IPv4 address, +// widened to SWWAF_BAN_SCOPE_V4_PREFIX, or the IPv6 group clientGroup +// counts it in. +func (rq *request) netblock() netip.Prefix { + addr := rq.client.Unmap() + if addr.Is4() { + return netip.PrefixFrom(addr, rq.h.config.BanScopeV4Prefix).Masked() + } + + return clientGroup(addr) +} + +// banExpires is when ban ends, as the log line gives it: a time, or +// permanent. +func banExpires(ban bans.Ban) string { + if ban.Permanent() { + return "permanent" + } + + return requestlog.FormatTime(ban.Expires) +} diff --git a/internal/proxy/bans_test.go b/internal/proxy/bans_test.go new file mode 100644 index 0000000..5683aa2 --- /dev/null +++ b/internal/proxy/bans_test.go @@ -0,0 +1,455 @@ +package proxy_test + +import ( + "bufio" + "errors" + "io" + "maps" + "net/http" + "net/netip" + "slices" + "sync" + "testing" + "time" + + "sneak.berlin/go/smallwebwaf/internal/bans" + "sneak.berlin/go/smallwebwaf/internal/proxy" + "sneak.berlin/go/smallwebwaf/internal/requestlog" +) + +const ( + // otherClient is a client next to client. + otherClient = "203.0.113.10" + // userAgent is the user agent of every request a sender sends. + userAgent = "ban-test/1.0" + // permanent is the log line's ban_expires for a permanent ban. + permanent = "permanent" +) + +func TestBrokenLimitBansTheClient(t *testing.T) { + t.Parallel() + + s, clk, _ := startWithClock(t, "", map[string]string{rateLimitPerMinute: "1"}) + expires := requestlog.FormatTime(clk.Now().Add(time.Hour)) + + // The request over the limit of one a minute is refused, and bans the + // client for an hour, the default. + s.get(client, http.StatusOK, requestlog.ActionForward) + + line := s.get(client, http.StatusForbidden, requestlog.ActionRateLimited) + if line.LimitHit != minute || line.Offence != requestlog.OffenceLimit || + line.BanExpires != expires { + t.Errorf("log line has limit_hit %q, offence %q and ban_expires %q, "+ + "want minute, limit and %s", line.LimitHit, line.Offence, line.BanExpires, + expires) + } + + // Every request while the ban lasts is refused. + clk.advance(time.Hour - time.Second) + + line = s.get(client, http.StatusForbidden, requestlog.ActionBanned) + if line.BanExpires != expires || line.Offence != "" || line.LimitHit != "" { + t.Errorf("log line has ban_expires %q, offence %q and limit_hit %q, "+ + "want %s and neither of the others", line.BanExpires, line.Offence, + line.LimitHit, expires) + } + + // Once it ends, the client is let through. + clk.advance(time.Second) + s.get(client, http.StatusOK, requestlog.ActionForward) +} + +func TestBanLengthsFollowTheSettings(t *testing.T) { + t.Parallel() + + s, clk, _ := startWithClock(t, "", map[string]string{ + rateLimitPerMinute: "1", + limitBanDuration: "10m", + limitBanRepeatWindow: "1h", + maxBanDuration: "1h", + }) + + // breakLimit has client go over the limit of one a minute, and + // returns when the ban that makes ends. + breakLimit := func() string { + s.get(client, http.StatusOK, requestlog.ActionForward) + + return s.get(client, http.StatusForbidden, requestlog.ActionRateLimited).BanExpires + } + wantExpires := func(got string, length time.Duration) { + t.Helper() + + want := requestlog.FormatTime(clk.Now().Add(length)) + if got != want { + t.Errorf("ban ends at %s, want %s", got, want) + } + } + + // A first ban lasts SWWAF_LIMIT_BAN_DURATION; one within + // SWWAF_LIMIT_BAN_REPEAT_WINDOW after it ended, three times as long. + wantExpires(breakLimit(), 10*time.Minute) + clk.advance(10*time.Minute + time.Hour) + wantExpires(breakLimit(), 30*time.Minute) + + // Later than that, SWWAF_LIMIT_BAN_DURATION again. + clk.advance(30*time.Minute + time.Hour + time.Second) + wantExpires(breakLimit(), 10*time.Minute) + clk.advance(10 * time.Minute) + wantExpires(breakLimit(), 30*time.Minute) + + // 90 minutes would be longer than SWWAF_MAX_BAN_DURATION: the ban is + // permanent. + clk.advance(30 * time.Minute) + + got := breakLimit() + if got != permanent { + t.Errorf("ban ends at %s, want a permanent one", got) + } + + clk.advance(365 * 24 * time.Hour) + s.get(client, http.StatusForbidden, requestlog.ActionBanned) +} + +func TestBanIsNotCountedAndResetsTheCounters(t *testing.T) { + t.Parallel() + + s, clk, server := startWithClock(t, "", map[string]string{rateLimitPerDay: "2"}) + + // The third request in a day is over the limit of two, and bans the + // client for an hour. + s.get(client, http.StatusOK, requestlog.ActionForward) + s.get(client, http.StatusOK, requestlog.ActionForward) + s.get(client, http.StatusForbidden, requestlog.ActionRateLimited) + + for range 3 { + s.get(client, http.StatusForbidden, requestlog.ActionBanned) + } + + // Later the same day the client has its whole allowance again: the + // ban set its counters back to zero, and the requests it refused were + // not counted for the rate limits, only in its notes. + clk.advance(time.Hour) + s.get(client, http.StatusOK, requestlog.ActionForward) + s.get(client, http.StatusOK, requestlog.ActionForward) + s.get(client, http.StatusForbidden, requestlog.ActionRateLimited) + + banned := server.Ledger.Bans(netip.MustParsePrefix(client + "/32")) + if len(banned) != 2 || banned[0].Notes.Refused != 3 { + t.Errorf("bans %+v, want two, the first with 3 requests refused", banned) + } +} + +func TestBanCoversTheClientsNetblock(t *testing.T) { + t.Parallel() + + // In the IPv4 cases, client breaks the limit; these two are next to it. + const ( + allowed = "203.0.113.60" // in SWWAF_ALLOW_NETS + exempt = "203.0.113.50" // in SWWAF_RATE_LIMIT_EXEMPT_NETS + ) + + for _, tc := range []struct { + name string + env map[string]string + breaker string // the client that breaks the limit + refused []string + let []string // let through + }{ + { + "an IPv4 address, by default", nil, client, + nil, []string{otherClient, exempt}, + }, + { + "the IPv4 netblock SWWAF_BAN_SCOPE_V4_PREFIX sets", + map[string]string{banScopeV4Prefix: "24"}, client, + []string{otherClient, exempt}, []string{"203.0.112.9", allowed}, + }, + { + "an IPv6 /64", nil, "2001:db8:5::1", + []string{"2001:db8:5::ffff:1"}, []string{"2001:db8:5:1::1"}, + }, + } { + t.Run(tc.name, func(t *testing.T) { + t.Parallel() + + env := map[string]string{ + rateLimitPerMinute: "1", + allowNets: allowed, + rateLimitExemptNets: exempt, + } + maps.Copy(env, tc.env) + s, _, _ := startWithClock(t, "", env) + + s.get(tc.breaker, http.StatusOK, requestlog.ActionForward) + s.get(tc.breaker, http.StatusForbidden, requestlog.ActionRateLimited) + + for _, sent := range tc.refused { + s.get(sent, http.StatusForbidden, requestlog.ActionBanned) + } + + for _, sent := range tc.let { + s.get(sent, http.StatusOK, requestlog.ActionForward) + } + }) + } +} + +func TestBannedClientIsRefusedBeforeItsCountryIsLookedUp(t *testing.T) { + t.Parallel() + + geojsURL, asked := startGeoJS(t) + s, _, _ := startWithClock(t, geojsURL, map[string]string{ + rateLimitPerMinute: "1", + banScopeV4Prefix: "24", + deniedCountries: "kp", + }) + + // fromDE's ban covers otherClient, which is refused unasked about. + s.get(fromDE, http.StatusOK, requestlog.ActionForward) + s.get(fromDE, http.StatusForbidden, requestlog.ActionRateLimited) + + line := s.get(otherClient, http.StatusForbidden, requestlog.ActionBanned) + if line.Country != "" { + t.Errorf("log line has country %q, want none", line.Country) + } + + if !slices.Equal(asked(), []string{fromDE}) { + t.Errorf("GeoJS was asked about %v, want %s alone", asked(), fromDE) + } +} + +func TestBanResponseAnswersEveryRefusalButTheSizeLimits(t *testing.T) { + t.Parallel() + + const denied = "192.0.2.50" // in SWWAF_DENY_NETS + + for _, tc := range []struct { + setting string // "" leaves SWWAF_BAN_RESPONSE at its default + status int // 0 is the connection closed without an answer + }{ + {"", http.StatusForbidden}, + {"403", http.StatusForbidden}, + {"429", http.StatusTooManyRequests}, + {"close", 0}, + } { + t.Run(banResponse+"="+tc.setting, func(t *testing.T) { + t.Parallel() + + geojsURL, _ := startGeoJS(t) + env := map[string]string{ + rateLimitPerMinute: "1", + denyNets: denied, + deniedCountries: "kp", + } + + if tc.setting != "" { + env[banResponse] = tc.setting + } + + s, _, _ := startWithClock(t, geojsURL, env) + + s.get(denied, tc.status, requestlog.ActionDenied) + s.get(fromKP, tc.status, requestlog.ActionCountryDenied) + s.get(fromDE, http.StatusOK, requestlog.ActionForward) + s.get(fromDE, tc.status, requestlog.ActionRateLimited) + s.get(fromDE, tc.status, requestlog.ActionBanned) + }) + } +} + +func TestBanNotes(t *testing.T) { + t.Parallel() + + geojsURL, _ := startGeoJS(t) + s, clk, server := startWithClock(t, geojsURL, map[string]string{ + rateLimitPerMinute: "1", + deniedCountries: "kp", + }) + start := clk.Now() + + s.get(fromDE, http.StatusOK, requestlog.ActionForward) + s.request(fromDE, "/repo/commits?page=2", + http.StatusForbidden, requestlog.ActionRateLimited) + s.get(fromDE, http.StatusForbidden, requestlog.ActionBanned) + s.get(fromDE, http.StatusForbidden, requestlog.ActionBanned) + + netblock := netip.MustParsePrefix(fromDE + "/32") + want := bans.Ban{ + Netblock: netblock, + Start: start, + Expires: start.Add(time.Hour), + Notes: bans.Notes{ + Country: "DE", + Limit: 1, + Window: minute, + Count: 2, + Request: bans.Request{ + Time: start, + Method: http.MethodGet, + Host: appHost, + Path: "/repo/commits?page=2", + Status: http.StatusForbidden, + UserAgent: userAgent, + }, + // The one let through, the one that broke the limit and the two + // refused under the ban. + Requests: 4, + Refused: 2, + EarlierBans: 0, + }, + } + + ledger := server.Ledger + + got := ledger.Bans(netblock) + if len(got) != 1 || got[0] != want { + t.Fatalf("bans\n%+v\nwant\n%+v", got, want) + } + + // The next ban counts this one among the earlier. + clk.advance(time.Hour) + s.get(fromDE, http.StatusOK, requestlog.ActionForward) + s.get(fromDE, http.StatusForbidden, requestlog.ActionRateLimited) + + got = ledger.Bans(netblock) + if len(got) != 2 || got[1].Notes.EarlierBans != 1 { + t.Errorf("bans %+v, want two, the second with one earlier ban", got) + } +} + +func TestMaxBansDropsTheBanOfTheNetblockSeenLongestAgo(t *testing.T) { + t.Parallel() + + s, _, _ := startWithClock(t, "", map[string]string{ + rateLimitPerMinute: "1", + maxBans: "1", + }) + + // One ban is held, so otherClient's ban drops client's. + s.get(client, http.StatusOK, requestlog.ActionForward) + s.get(client, http.StatusForbidden, requestlog.ActionRateLimited) + s.get(otherClient, http.StatusOK, requestlog.ActionForward) + s.get(otherClient, http.StatusForbidden, requestlog.ActionRateLimited) + + s.get(client, http.StatusOK, requestlog.ActionForward) + s.get(otherClient, http.StatusForbidden, requestlog.ActionBanned) +} + +// clock is the time a test sets, by which smallwebwaf counts requests and +// makes bans. +type clock struct { + mu sync.Mutex + now time.Time +} + +// Now tells the time. +func (c *clock) Now() time.Time { + c.mu.Lock() + defer c.mu.Unlock() + + return c.now +} + +// advance moves the clock on by d. +func (c *clock) advance(d time.Duration) { + c.mu.Lock() + defer c.mu.Unlock() + + c.now = c.now.Add(d) +} + +// startWithClock starts smallwebwaf in front of an app that answers 200, +// with the settings in env on top of trusting localhost's +// X-Forwarded-For, clients' countries looked up at geojsURL, and a clock +// set to midnight, the start of a bucket in every window. +func startWithClock( + t *testing.T, geojsURL string, env map[string]string, +) (*sender, *clock, *proxy.Server) { + t.Helper() + + app := startApp(t, func(http.ResponseWriter, *http.Request) {}) + clk := &clock{now: time.Date(2026, 10, 6, 0, 0, 0, 0, time.UTC)} + settings := map[string]string{trustedProxies: trustLocalhost} + maps.Copy(settings, env) + + addr, out, server := startProxyWithClock(t, app.URL, geojsURL, clk.Now, settings) + + return &sender{t: t, addr: addr, out: out}, clk, server +} + +// sender sends requests to smallwebwaf one after another, each on a +// connection of its own, and checks each one's answer and log line. They +// must be the only requests smallwebwaf is sent, since the log lines are +// matched to them in order. +type sender struct { + t *testing.T + addr string + out *output + sent int +} + +// get sends a GET request for / from the client at from. +func (s *sender) get(from string, status int, action string) logLine { + s.t.Helper() + + return s.request(from, "/", status, action) +} + +// request sends a GET request for path from the client at from, as +// X-Forwarded-For names it, and checks that its answer and its log line +// have status, 0 for the connection closed without an answer, and that +// the line has action. It returns the log line. +func (s *sender) request(from, path string, status int, action string) logLine { + s.t.Helper() + + line, _ := s.requestWithHeader(from, path, "", status, action) + + return line +} + +// requestWithHeader is request with header, such as "Authorization: +// Bearer x", added to the request unless it is "". It returns the body of +// the answer too. +func (s *sender) requestWithHeader( + from, path, header string, status int, action string, +) (logLine, string) { + s.t.Helper() + + if header != "" { + header += "\r\n" + } + + conn := dial(s.t, s.addr) + send(s.t, conn, "GET "+path+" HTTP/1.1\r\nHost: "+appHost+ + "\r\nUser-Agent: "+userAgent+"\r\n"+forwardedFor+": "+from+"\r\n"+ + header+"\r\n") + + err := conn.SetReadDeadline(time.Now().Add(waitLimit)) + if err != nil { + s.t.Fatalf("set read deadline: %v", err) + } + + var got answer + + res, err := http.ReadResponse(bufio.NewReader(conn), nil) + + switch { + case err == nil: + got = readAnswer(res) + case !errors.Is(err, io.ErrUnexpectedEOF): + s.t.Fatalf("read response: %v", err) + } + + _ = conn.Close() + + if got.status != status { + s.t.Errorf("request %d, from %s: status %d, want %d", s.sent+1, from, + got.status, status) + } + + line := s.out.requestLines(s.t, s.sent+1)[s.sent] + s.sent++ + wantLine(s.t, line, status, action) + + return line, string(got.body) +} diff --git a/internal/proxy/bodies.go b/internal/proxy/bodies.go new file mode 100644 index 0000000..d4c23ae --- /dev/null +++ b/internal/proxy/bodies.go @@ -0,0 +1,173 @@ +package proxy + +import ( + "errors" + "io" + "net/http" + "sync/atomic" + + "sneak.berlin/go/smallwebwaf/internal/requestlog" +) + +// errResponseTooLarge ends an app's response body that is longer than +// SWWAF_RESPONSE_MAX_BYTES. +var errResponseTooLarge = errors.New( + "the response body is over SWWAF_RESPONSE_MAX_BYTES") + +// requestBody is the client's request body on its way to the app. The +// transport reads it on a goroutine of its own. +type requestBody struct { + // body is the client's body, ending in an *http.MaxBytesError past + // SWWAF_REQUEST_MAX_BYTES. + body io.ReadCloser + rq *request + // waiting is true while a Read waits for the client to send more. + waiting atomic.Bool + // received is true once the client has sent the whole body. + received atomic.Bool + // bytes is how much of the body has been read. + bytes atomic.Int64 +} + +// Read reads from the client's body. +func (b *requestBody) Read(p []byte) (int, error) { + b.waiting.Store(true) + n, err := b.body.Read(p) + b.waiting.Store(false) + b.bytes.Add(int64(n)) + + var tooLarge *http.MaxBytesError + + switch { + case errors.Is(err, io.EOF): + b.received.Store(true) + b.rq.bodyReceived() + case errors.As(err, &tooLarge): + b.rq.refuse(refusal{ + status: http.StatusRequestEntityTooLarge, + action: requestlog.ActionTooLarge, + limit: "SWWAF_REQUEST_MAX_BYTES", + }) + } + + return n, err +} + +// Close closes the client's body. +func (b *requestBody) Close() error { + return b.body.Close() +} + +// responseBody is the app's response body on its way to the client. +type responseBody struct { + // body is the app's body, ending in an *http.MaxBytesError past + // SWWAF_RESPONSE_MAX_BYTES. + body io.ReadCloser + rq *request +} + +// Read reads from the app's body. +func (b *responseBody) Read(p []byte) (int, error) { + n, err := b.body.Read(p) + if err == nil { + return n, nil + } + + var tooLarge *http.MaxBytesError + + switch { + case errors.Is(err, io.EOF): + b.rq.responseReceived() + case errors.As(err, &tooLarge): + b.rq.refuse(refusal{ + status: http.StatusBadGateway, + action: requestlog.ActionTooLarge, + limit: "SWWAF_RESPONSE_MAX_BYTES", + }) + + return n, errResponseTooLarge + case b.rq.in.Context().Err() == nil: + // The answer broke off, not because the client went away. If a + // timeout cut it, that refusal came first and is the one kept. + b.rq.refuse(refusal{ + status: http.StatusBadGateway, + action: requestlog.ActionUpstreamError, + }) + } + + return n, err +} + +// Close closes the app's body. +func (b *responseBody) Close() error { + return b.body.Close() +} + +// limitBody returns body, cut off with an *http.MaxBytesError after +// maxBytes, or unchanged if maxBytes is zero, which is off. +func limitBody(body io.ReadCloser, maxBytes int64) io.ReadCloser { + if maxBytes == 0 { + return body + } + + // Without a ResponseWriter, MaxBytesReader only counts and cuts off. + return http.MaxBytesReader(nil, body, maxBytes) +} + +// responseWriter is the response to the client. It notes the status and +// size for the log line, and the first error writing to the client. +type responseWriter struct { + http.ResponseWriter + + // status is the final status sent, or zero before one is. + status int + bytes int64 + err error +} + +// WriteHeader sends the status and headers. An informational 1xx status +// is passed on and the final status still comes later. +func (w *responseWriter) WriteHeader(status int) { + if status >= http.StatusOK && w.status == 0 { + w.status = status + } + + w.ResponseWriter.WriteHeader(status) +} + +// Write sends part of the body. +func (w *responseWriter) Write(p []byte) (int, error) { + if w.status == 0 { + w.status = http.StatusOK + } + + n, err := w.ResponseWriter.Write(p) + w.bytes += int64(n) + w.noteError(err) + + return n, err +} + +// FlushError sends what has been written so far. +// http.ResponseController calls it, as ReverseProxy does after each +// write. +func (w *responseWriter) FlushError() error { + err := http.NewResponseController(w.ResponseWriter).Flush() + w.noteError(err) + + return err +} + +// Unwrap lets http.ResponseController reach net/http's own +// ResponseWriter, which is how ReverseProxy takes over the connection of +// an upgraded request. +func (w *responseWriter) Unwrap() http.ResponseWriter { + return w.ResponseWriter +} + +// noteError keeps the first error writing to the client. +func (w *responseWriter) noteError(err error) { + if w.err == nil { + w.err = err + } +} diff --git a/internal/proxy/client.go b/internal/proxy/client.go new file mode 100644 index 0000000..11058b9 --- /dev/null +++ b/internal/proxy/client.go @@ -0,0 +1,144 @@ +package proxy + +import ( + "crypto/rand" + "net/http" + "net/netip" + "slices" + "strings" +) + +// peerAddress is the address of the request's TCP peer, normally traefik. +func peerAddress(r *http.Request) netip.Addr { + addrPort, err := netip.ParseAddrPort(r.RemoteAddr) + if err != nil { + return netip.Addr{} + } + + return addrPort.Addr().Unmap() +} + +// clientAddress works out who the client is. A peer outside the trusted +// proxies is the client, and what it says in X-Forwarded-For is ignored. +// For a peer inside them, X-Forwarded-For is read from the right, and the +// first address outside them is the client; if every address in it is +// inside, the leftmost is, and with no header, the peer. An entry that is +// not an address ends the reading, since nothing to its left can be +// believed. +func clientAddress( + peer netip.Addr, forwardedFor []string, trusted []netip.Prefix, +) netip.Addr { + client := peer + if !isInside(peer, trusted) { + return client + } + + entries := strings.Split(strings.Join(forwardedFor, ","), ",") + for _, entry := range slices.Backward(entries) { + addr, err := netip.ParseAddr(strings.TrimSpace(entry)) + if err != nil { + break + } + + client = addr.Unmap() + if !isInside(client, trusted) { + break + } + } + + return client +} + +// requestIDHeader carries the request's id, from traefik and to the app. +const requestIDHeader = "X-Request-ID" + +// requestID is the request's id: the one a trusted proxy sent, or a new +// random one. A peer outside the trusted proxies did not come through +// traefik, so the id it sends is its own claim, and is replaced. +func requestID(r *http.Request, peerTrusted bool) string { + id := r.Header.Get(requestIDHeader) + if !peerTrusted || id == "" { + id = rand.Text() + } + + return id +} + +// scheme is how the client reached traefik, as a trusted proxy says in +// X-Forwarded-Proto, or otherwise http, the only scheme smallwebwaf +// serves. +func scheme(r *http.Request, peerTrusted bool) string { + proto := r.Header.Get("X-Forwarded-Proto") + if !peerTrusted || proto == "" { + return "http" + } + + return proto +} + +// ipv6GroupPrefix is the length of the IPv6 netblock that is one client. +const ipv6GroupPrefix = 64 + +// clientGroup is the client a request is counted toward: its IPv4 +// address, or the /64 its IPv6 address is in, since one abuser usually +// holds a whole /64. An IPv4 address in IPv6 form counts as IPv4. +func clientGroup(addr netip.Addr) netip.Prefix { + addr = addr.Unmap() + if addr.Is6() { + return netip.PrefixFrom(addr, ipv6GroupPrefix).Masked() + } + + return netip.PrefixFrom(addr, addr.BitLen()) +} + +// isInside reports whether addr is in one of the netblocks. +func isInside(addr netip.Addr, netblocks []netip.Prefix) bool { + return slices.ContainsFunc(netblocks, func(netblock netip.Prefix) bool { + return netblock.Contains(addr) + }) +} + +// setForwardedHeaders sets the headers in which the app learns about the +// client, so that it sees what it would see from traefik directly. A +// trusted proxy's forwarded headers pass on, with the proxy's own address +// added to X-Forwarded-For. Those of any other peer are its own claims and +// are replaced: X-Forwarded-For names the peer, X-Forwarded-Host the host +// it asked for, and X-Forwarded-Proto plain http, which is how it reached +// smallwebwaf. +func setForwardedHeaders(in, out *http.Request, peer netip.Addr, trusted bool) { + forwardedFor := peer.String() + + if trusted { + // ReverseProxy removes these from out before Rewrite. + for _, name := range []string{"Forwarded", "X-Forwarded-Host", "X-Forwarded-Proto"} { + values, ok := in.Header[name] + if ok { + out.Header[name] = values + } + } + + prior := in.Header.Values("X-Forwarded-For") + if len(prior) > 0 { + forwardedFor = strings.Join(prior, ", ") + ", " + forwardedFor + } + + out.Header.Set("X-Forwarded-For", forwardedFor) + + return + } + + // ReverseProxy has removed Forwarded and the three set below; these + // are the other headers in which traefik tells the app about the + // client and its request. + for _, name := range []string{ + "X-Forwarded-Port", "X-Forwarded-Server", "X-Forwarded-Uri", + "X-Forwarded-Method", "X-Forwarded-Prefix", "X-Forwarded-Tls-Client-Cert", + "X-Forwarded-Tls-Client-Cert-Info", "X-Real-Ip", + } { + out.Header.Del(name) + } + + out.Header.Set("X-Forwarded-For", forwardedFor) + out.Header.Set("X-Forwarded-Host", in.Host) + out.Header.Set("X-Forwarded-Proto", "http") +} diff --git a/internal/proxy/client_test.go b/internal/proxy/client_test.go new file mode 100644 index 0000000..cf4ece5 --- /dev/null +++ b/internal/proxy/client_test.go @@ -0,0 +1,165 @@ +package proxy_test + +import ( + "encoding/json" + "net/http" + "testing" +) + +const ( + // trustLocalhost trusts the address every test connects from, and a + // network for proxies in front of it. + trustLocalhost = localhost + "/32,10.0.0.0/8" + // appHost is the host every test asks for. + appHost = "app.example" + // client is the client's address, as a proxy names it. + client = "203.0.113.9" + // forwardedFor is the header that lists the client and its proxies, + // and forwardedProto the one that gives the scheme the client used. + forwardedFor = "X-Forwarded-For" + forwardedProto = "X-Forwarded-Proto" + // secure is the scheme a client reached traefik with, and plain the + // one smallwebwaf serves. + secure = "https" + plain = "http" +) + +// appHeaders is what the app tells about the headers it received. +type appHeaders struct { + Host string `json:"host"` + ForwardedFor string `json:"forwardedFor"` + ForwardedHost string `json:"forwardedHost"` + ForwardedProto string `json:"forwardedProto"` + RealIP string `json:"realIp"` +} + +// clientAddressCase is a request and what smallwebwaf makes of it. +type clientAddressCase struct { + name string + env map[string]string + header http.Header + wantClient string + wantApp appHeaders +} + +func TestClientAddressAndForwardedHeaders(t *testing.T) { + t.Parallel() + + for _, tc := range clientAddressCases() { + t.Run(tc.name, func(t *testing.T) { + t.Parallel() + + got, line := requestWithHeaders(t, tc.env, tc.header) + + tc.wantApp.Host = appHost + if got != tc.wantApp { + t.Errorf("app received %+v, want %+v", got, tc.wantApp) + } + + if line.ClientIP != tc.wantClient || line.PeerIP != localhost { + t.Errorf("log line has client_ip %q and peer_ip %q, want %q and %q", + line.ClientIP, line.PeerIP, tc.wantClient, localhost) + } + }) + } +} + +// clientAddressCases are the requests TestClientAddressAndForwardedHeaders +// sends, from 127.0.0.1, which the default trusted proxies leave out. +func clientAddressCases() []clientAddressCase { + trusted := map[string]string{trustedProxies: trustLocalhost} + forged := http.Header{ + forwardedFor: {client}, + "X-Forwarded-Host": {"forged.example"}, + forwardedProto: {secure}, + "X-Real-Ip": {client}, + } + replaced := appHeaders{ + ForwardedFor: localhost, ForwardedHost: appHost, ForwardedProto: plain, + } + + return []clientAddressCase{{ + name: "a peer outside the trusted proxies is the client, " + + "and its forwarded headers are replaced", + header: forged, wantClient: localhost, wantApp: replaced, + }, { + name: "set but empty, the trusted proxies trust nothing", + env: map[string]string{trustedProxies: ""}, + header: forged, wantClient: localhost, wantApp: replaced, + }, { + name: "behind a trusted peer, the client is the first address " + + "outside the trusted proxies from the right", + env: trusted, + header: http.Header{ + forwardedFor: {"198.51.100.7, " + client + ", 10.0.0.2"}, + "X-Forwarded-Host": {appHost}, + forwardedProto: {secure}, + "X-Real-Ip": {client}, + }, + wantClient: client, + wantApp: appHeaders{ + ForwardedFor: "198.51.100.7, " + client + ", 10.0.0.2, " + localhost, + ForwardedHost: appHost, ForwardedProto: secure, RealIP: client, + }, + }, { + name: "when every address is a trusted proxy, the leftmost is the client", + env: trusted, + header: http.Header{forwardedFor: {"10.0.0.5, 10.0.0.2"}}, + wantClient: "10.0.0.5", + wantApp: appHeaders{ForwardedFor: "10.0.0.5, 10.0.0.2, " + localhost}, + }, { + name: "with no header, a trusted peer is the client", + env: trusted, + wantClient: localhost, + wantApp: appHeaders{ForwardedFor: localhost}, + }, { + name: "an entry that is not an address ends the reading", + env: trusted, + header: http.Header{forwardedFor: {client + ", unknown, 10.0.0.2"}}, + wantClient: "10.0.0.2", + wantApp: appHeaders{ + ForwardedFor: client + ", unknown, 10.0.0.2, " + localhost, + }, + }, { + name: "several header lines are read as one list", + env: trusted, + header: http.Header{forwardedFor: {"2001:db8::7", "10.0.0.2"}}, + wantClient: "2001:db8::7", + wantApp: appHeaders{ForwardedFor: "2001:db8::7, 10.0.0.2, " + localhost}, + }} +} + +// requestWithHeaders sends a request for appHost with header through +// smallwebwaf, with the settings in env, and returns the headers the app +// received and the request's log line. +func requestWithHeaders( + t *testing.T, env map[string]string, header http.Header, +) (appHeaders, logLine) { + t.Helper() + + app := startApp(t, func(w http.ResponseWriter, r *http.Request) { + _ = json.NewEncoder(w).Encode(appHeaders{ + Host: r.Host, + ForwardedFor: r.Header.Get(forwardedFor), + ForwardedHost: r.Header.Get("X-Forwarded-Host"), + ForwardedProto: r.Header.Get(forwardedProto), + RealIP: r.Header.Get("X-Real-IP"), + }) + }) + addr, out := startProxy(t, app.URL, env) + + req := newRequest(t, http.MethodGet, addr, "/", http.NoBody) + req.Host = appHost + req.Header = header.Clone() + + answered := do(t, req) + + var got appHeaders + + err := json.Unmarshal(answered.body, &got) + if err != nil { + t.Fatalf("decode the app's answer %q: %v", answered.body, err) + } + + return got, out.requestLine(t) +} diff --git a/internal/proxy/countries.go b/internal/proxy/countries.go new file mode 100644 index 0000000..2b6ccab --- /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. A client without a country, or 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 + } + + var country string + if hasCountry(rq.client) { + 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..9afeb37 --- /dev/null +++ b/internal/proxy/countries_test.go @@ -0,0 +1,308 @@ +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{deniedCountries: "kp"}, + []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()) + } + }) + } +} + +func TestExclusiveListRefusesAPrivateAddressUnlessAllowed(t *testing.T) { + t.Parallel() + + for _, tc := range []struct { + name string + allowNets string + status int + action string + }{ + { + "not in SWWAF_ALLOW_NETS", "", + http.StatusForbidden, requestlog.ActionCountryDenied, + }, + { + "in SWWAF_ALLOW_NETS", "10.0.0.7,fd00::/8", + http.StatusOK, requestlog.ActionForward, + }, + } { + t.Run(tc.name, func(t *testing.T) { + t.Parallel() + + app := startApp(t, func(http.ResponseWriter, *http.Request) {}) + geojsURL, asked := startGeoJS(t) + addr, out := startProxyWithGeoJS(t, app.URL, geojsURL, map[string]string{ + trustedProxies: trustLocalhost, + allowedCountries: "de", + allowNets: tc.allowNets, + }) + + wantAnswers(t, addr, out, []sentRequest{ + {"10.0.0.7", tc.status, tc.action}, + {"fd00::5", tc.status, tc.action}, + }) + + 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/health_test.go b/internal/proxy/health_test.go new file mode 100644 index 0000000..e069d98 --- /dev/null +++ b/internal/proxy/health_test.go @@ -0,0 +1,56 @@ +package proxy_test + +import ( + "net/http" + "sync/atomic" + "testing" + + "sneak.berlin/go/smallwebwaf/internal/proxy" + "sneak.berlin/go/smallwebwaf/internal/requestlog" +) + +func TestHealthEndpointIsAnsweredBeforeAnyCheck(t *testing.T) { + t.Parallel() + + var calls atomic.Int32 + + app := startApp(t, func(http.ResponseWriter, *http.Request) { + calls.Add(1) + }) + // With a limit of one request a minute, any request counted before + // the last one would have it refused. + addr, out := startProxy(t, app.URL, map[string]string{rateLimitPerMinute: "1"}) + + const ( + healthChecks = 3 + contentType = "text/plain; charset=utf-8" + ) + + for range healthChecks { + got := get(t, addr, proxy.HealthPath) + wantStatus(t, got, http.StatusOK) + + if string(got.body) != "ok\n" || got.header.Get("Content-Type") != contentType { + t.Errorf("health endpoint answered %q with Content-Type %q, want ok "+ + "with %q", got.body, got.header.Get("Content-Type"), contentType) + } + } + + wantStatus(t, get(t, addr, "/"), http.StatusOK) + + lines := out.requestLines(t, healthChecks+1) + for _, line := range lines[:healthChecks] { + wantLine(t, line, http.StatusOK, requestlog.ActionAdmin) + + if line.ResponseContentType != contentType { + t.Errorf("health check's log line has response_content_type %q, "+ + "want %q", line.ResponseContentType, contentType) + } + } + + wantLine(t, lines[healthChecks], http.StatusOK, requestlog.ActionForward) + + if calls.Load() != 1 { + t.Errorf("the app was called %d times, want once", calls.Load()) + } +} diff --git a/internal/proxy/history_test.go b/internal/proxy/history_test.go new file mode 100644 index 0000000..29550e6 --- /dev/null +++ b/internal/proxy/history_test.go @@ -0,0 +1,124 @@ +package proxy_test + +import ( + "io" + "net/http" + "net/netip" + "strings" + "testing" + "time" + + "sneak.berlin/go/smallwebwaf/internal/proxy" + "sneak.berlin/go/smallwebwaf/internal/ratelimit" + "sneak.berlin/go/smallwebwaf/internal/requestlog" +) + +func TestHistoryKeepsEachRequestOfTheClient(t *testing.T) { + t.Parallel() + + geojsURL, _ := startGeoJS(t) + s, clk, server := startWithClock(t, geojsURL, map[string]string{ + rateLimitPerMinute: "2", + deniedCountries: "kp", + }) + start := clk.Now() + + // Two let through, one over the limit, which bans the client, and one + // refused under that ban, for which the country is not looked up. + s.get(fromDE, http.StatusOK, requestlog.ActionForward) + clk.advance(time.Second) + s.get(fromDE, http.StatusOK, requestlog.ActionForward) + s.get(fromDE, http.StatusForbidden, requestlog.ActionRateLimited) + clk.advance(time.Second) + s.get(fromDE, http.StatusForbidden, requestlog.ActionBanned) + + want := ratelimit.History{ + FirstSeen: start, + LastSeen: start.Add(2 * time.Second), + Country: "DE", + LookedUp: start.Add(time.Second), + Requests: 4, + Forwarded: 2, + Refused: 2, + // The app answers with no body, smallwebwaf with its status text. + ResponseBytes: 2 * int64(len("Forbidden\n")), + Responses: ratelimit.Responses{Status2xx: 2, Status4xx: 2}, + Offences: ratelimit.Offences{Limit: 1}, + } + + got := historyOf(t, server, fromDE) + if got != want { + t.Errorf("history\n%+v\nwant\n%+v", got, want) + } +} + +func TestHistoryCountsTheBodiesEachWay(t *testing.T) { + t.Parallel() + + app := startApp(t, func(w http.ResponseWriter, r *http.Request) { + _, _ = io.Copy(io.Discard, r.Body) + _, _ = io.WriteString(w, "hello") + }) + addr, out, server := startProxyWithClock(t, app.URL, "", time.Now, nil) + + got := do(t, newRequest(t, http.MethodPost, addr, "/", strings.NewReader("abc"))) + wantStatus(t, got, http.StatusOK) + out.requestLine(t) + + history := historyOf(t, server, localhost) + if history.RequestBytes != 3 || history.ResponseBytes != 5 { + t.Errorf("history counts %d bytes in and %d out, want 3 and 5", + history.RequestBytes, history.ResponseBytes) + } +} + +func TestHealthEndpointIsNotInTheHistory(t *testing.T) { + t.Parallel() + + app := startApp(t, func(http.ResponseWriter, *http.Request) {}) + addr, out, server := startProxyWithClock(t, app.URL, "", time.Now, nil) + + wantStatus(t, get(t, addr, proxy.HealthPath), http.StatusOK) + out.requestLine(t) + + if clients := server.Limiter.Snapshot(); len(clients) != 0 { + t.Errorf("the table holds %+v, want no client", clients) + } +} + +func TestRequestForSmallwebwafIsRefusedOnlyWithoutTheToken(t *testing.T) { + t.Parallel() + + app := startApp(t, func(http.ResponseWriter, *http.Request) {}) + addr, out, server := startProxyWithClock(t, app.URL, "", time.Now, + map[string]string{metricsToken: token}) + + // The metrics and the 404 are neither forwarded nor refused; the 401 + // is refused. + scrape(t, addr) + wantStatus(t, get(t, addr, "/_smallwebwaf/nothing"), http.StatusNotFound) + wantStatus(t, get(t, addr, proxy.MetricsPath), http.StatusUnauthorized) + out.requestLines(t, 3) + + history := historyOf(t, server, localhost) + if history.Requests != 3 || history.Forwarded != 0 || history.Refused != 1 { + t.Errorf("history counts %d requests, %d forwarded and %d refused, "+ + "want 3, 0 and 1", history.Requests, history.Forwarded, history.Refused) + } +} + +// historyOf returns the history of the client at addr. +func historyOf(t *testing.T, server *proxy.Server, addr string) ratelimit.History { + t.Helper() + + client := netip.MustParsePrefix(addr + "/32") + for _, c := range server.Limiter.Snapshot() { + if c.Client == client { + return c.History + } + } + + t.Fatalf("%s is not in the table", client) + + return ratelimit.History{} +} diff --git a/internal/proxy/limits_test.go b/internal/proxy/limits_test.go new file mode 100644 index 0000000..1f7efe3 --- /dev/null +++ b/internal/proxy/limits_test.go @@ -0,0 +1,162 @@ +package proxy_test + +import ( + "bytes" + "errors" + "io" + "net/http" + "strconv" + "sync/atomic" + "testing" + + "sneak.berlin/go/smallwebwaf/internal/requestlog" +) + +// sizeLimit is the size limit the tests set, 1K as a setting. +const ( + sizeLimit = 1 << 10 + sizeLimitSetting = "1K" +) + +func TestRequestBodyLimit(t *testing.T) { + t.Parallel() + + for _, tc := range []struct { + name string + size int + // announced sends the size in Content-Length; otherwise the body + // is sent in chunks with no length given. + announced bool + want int + action string + // refusedBeforeApp is a refusal before anything reaches the app. + // A body over the limit with no length given has already partly + // reached the app when it is refused. + refusedBeforeApp bool + }{ + {"announced, over the limit", 2 * sizeLimit, true, + http.StatusRequestEntityTooLarge, requestlog.ActionTooLarge, true}, + {"announced, at the limit", sizeLimit, true, + http.StatusOK, requestlog.ActionForward, false}, + {"not announced, over the limit", 4 * sizeLimit, false, + http.StatusRequestEntityTooLarge, requestlog.ActionTooLarge, false}, + {"not announced, at the limit", sizeLimit, false, + http.StatusOK, requestlog.ActionForward, false}, + } { + t.Run(tc.name, func(t *testing.T) { + t.Parallel() + + var calls atomic.Int32 + + app := startApp(t, func(_ http.ResponseWriter, r *http.Request) { + calls.Add(1) + + _, _ = io.Copy(io.Discard, r.Body) + }) + addr, out := startProxy(t, app.URL, map[string]string{ + requestMaxBytes: sizeLimitSetting, + metricsToken: token, + }) + + var body io.Reader = bytes.NewReader(make([]byte, tc.size)) + if !tc.announced { + body = io.MultiReader(body) // hides the length + } + + wantStatus(t, do(t, newRequest(t, http.MethodPost, addr, "/upload", body)), + tc.want) + wantLine(t, out.requestLine(t), tc.want, tc.action) + + hits := 0 + if tc.action == requestlog.ActionTooLarge { + hits = 1 + } + + wantLimitHits(t, addr, requestMaxBytes, hits) + + if tc.refusedBeforeApp && calls.Load() != 0 { + t.Errorf("the app was called %d times, want never", calls.Load()) + } + }) + } +} + +func TestResponseBodyLimit(t *testing.T) { + t.Parallel() + + for _, tc := range []struct { + name string + size int + // announced sends the size in Content-Length; otherwise the body + // is sent in chunks with no length given. + announced bool + want int + action string + // received is how much of a body the client gets, and cutOff + // whether the connection is then cut. + received int + cutOff bool + }{ + {"announced, over the limit", 2 * sizeLimit, true, http.StatusBadGateway, + requestlog.ActionTooLarge, len("Bad Gateway\n"), false}, + {"announced, at the limit", sizeLimit, true, http.StatusOK, + requestlog.ActionForward, sizeLimit, false}, + {"not announced, over the limit", 4 * sizeLimit, false, http.StatusOK, + requestlog.ActionTooLarge, sizeLimit, true}, + {"not announced, at the limit", sizeLimit, false, http.StatusOK, + requestlog.ActionForward, sizeLimit, false}, + } { + t.Run(tc.name, func(t *testing.T) { + t.Parallel() + + app := startApp(t, func(w http.ResponseWriter, _ *http.Request) { + answerWithSize(w, tc.size, tc.announced) + }) + addr, out := startProxy(t, app.URL, map[string]string{ + responseMaxBytes: sizeLimitSetting, + metricsToken: token, + }) + + got := get(t, addr, "/download") + wantStatus(t, got, tc.want) + + if len(got.body) != tc.received || + errors.Is(got.err, io.ErrUnexpectedEOF) != tc.cutOff { + t.Errorf("client got %d bytes (%v), want %d", + len(got.body), got.err, tc.received) + } + + line := out.requestLine(t) + wantLine(t, line, tc.want, tc.action) + + if line.UpstreamStatus != http.StatusOK { + t.Errorf("log line has upstream_status %d", line.UpstreamStatus) + } + + hits := 0 + if tc.action == requestlog.ActionTooLarge { + hits = 1 + } + + wantLimitHits(t, addr, responseMaxBytes, hits) + }) + } +} + +// answerWithSize answers with a body of size bytes, announced in +// Content-Length or sent in chunks with no length given. +func answerWithSize(w http.ResponseWriter, size int, announced bool) { + body := make([]byte, size) + if announced { + w.Header().Set("Content-Length", strconv.Itoa(size)) + _, _ = w.Write(body) + + return + } + + // Sending part of it before the end keeps Go's server from working + // out the length. + _, _ = w.Write(body[:size/2]) + _ = http.NewResponseController(w).Flush() + _, _ = w.Write(body[size/2:]) +} diff --git a/internal/proxy/metrics_test.go b/internal/proxy/metrics_test.go new file mode 100644 index 0000000..6765f2e --- /dev/null +++ b/internal/proxy/metrics_test.go @@ -0,0 +1,455 @@ +package proxy_test + +import ( + "io" + "net/http" + "net/http/httptest" + "net/netip" + "strconv" + "strings" + "sync" + "sync/atomic" + "testing" + "time" + + "sneak.berlin/go/smallwebwaf/internal/lookup" + "sneak.berlin/go/smallwebwaf/internal/proxy" + "sneak.berlin/go/smallwebwaf/internal/requestlog" +) + +const ( + metricsToken = "SWWAF_METRICS_TOKEN" //nolint:gosec // the setting's name + metricsTopN = "SWWAF_METRICS_TOP_N" + // token is the SWWAF_METRICS_TOKEN the tests set, and bearer how a + // request carries it. + token = "0123456789abcdef0123456789abcdef" + bearer = "Bearer " + token +) + +func TestMetricsAreOffWhileTheTokenIsUnset(t *testing.T) { + t.Parallel() + + var calls atomic.Int32 + + app := startApp(t, func(http.ResponseWriter, *http.Request) { + calls.Add(1) + }) + addr, out := startProxy(t, app.URL, nil) + + // An empty token does not match the unset one either. + for i, authorization := range []string{bearer, "Bearer ", ""} { + req := newRequest(t, http.MethodGet, addr, proxy.MetricsPath, http.NoBody) + if authorization != "" { + req.Header.Set("Authorization", authorization) + } + + wantStatus(t, do(t, req), http.StatusNotFound) + wantLine(t, out.requestLines(t, i+1)[i], http.StatusNotFound, + requestlog.ActionAdmin) + } + + if calls.Load() != 0 { + t.Errorf("the app was called %d times, want never", calls.Load()) + } +} + +func TestMetricsNeedTheTokenAndOtherPathsAreNotFound(t *testing.T) { + t.Parallel() + + var calls atomic.Int32 + + app := startApp(t, func(http.ResponseWriter, *http.Request) { + calls.Add(1) + }) + addr, out := startProxy(t, app.URL, map[string]string{metricsToken: token}) + + for i, tc := range []struct { + method, path, authorization string + status int + }{ + {http.MethodGet, proxy.MetricsPath, "", http.StatusUnauthorized}, + { + http.MethodGet, proxy.MetricsPath, "Bearer " + strings.ToUpper(token), + http.StatusUnauthorized, + }, + {http.MethodGet, proxy.MetricsPath, "Basic " + token, http.StatusUnauthorized}, + {http.MethodGet, proxy.MetricsPath, bearer, http.StatusOK}, + {http.MethodGet, proxy.MetricsPath, "bearer " + token, http.StatusOK}, + {http.MethodPost, proxy.MetricsPath, bearer, http.StatusNotFound}, + {http.MethodGet, proxy.MetricsPath + "/", bearer, http.StatusNotFound}, + {http.MethodGet, "/_smallwebwaf/bans", bearer, http.StatusNotFound}, + {http.MethodPost, proxy.HealthPath, "", http.StatusNotFound}, + } { + req := newRequest(t, tc.method, addr, tc.path, http.NoBody) + if tc.authorization != "" { + req.Header.Set("Authorization", tc.authorization) + } + + got := do(t, req) + wantStatus(t, got, tc.status) + wantLine(t, out.requestLines(t, i+1)[i], tc.status, requestlog.ActionAdmin) + + if tc.status == http.StatusUnauthorized && + got.header.Get("WWW-Authenticate") != "Bearer" { + t.Errorf("%q was answered without WWW-Authenticate: Bearer", + tc.authorization) + } + + if tc.status == http.StatusOK && + !strings.Contains(string(got.body), "# TYPE smallwebwaf_requests_total counter") { + t.Errorf("the metrics are\n%s", got.body) + } + } + + if calls.Load() != 0 { + t.Errorf("the app was called %d times, want never", calls.Load()) + } +} + +func TestMetricsAreAskedForThroughTheChecks(t *testing.T) { + t.Parallel() + + s, _, _ := startWithClock(t, "", map[string]string{ + metricsToken: token, + rateLimitPerMinute: "1", + }) + + // Asking for the metrics counts toward the client's limit of one + // request a minute, so its next request breaks it, and bans it. A + // banned client is refused the metrics too. + s.scrape(client) + s.get(client, http.StatusForbidden, requestlog.ActionRateLimited) + s.requestWithHeader(client, proxy.MetricsPath, "Authorization: "+bearer, + http.StatusForbidden, requestlog.ActionBanned) +} + +func TestMetricsCountTheTraffic(t *testing.T) { + t.Parallel() + + arrived, release := make(chan struct{}), make(chan struct{}) + app := startApp(t, func(w http.ResponseWriter, r *http.Request) { + _, _ = io.Copy(io.Discard, r.Body) + + if r.URL.Path == "/held" { + close(arrived) + <-release + } + + _, _ = io.WriteString(w, "hello") + }) + releaseApp := sync.OnceFunc(func() { close(release) }) + t.Cleanup(releaseApp) + + addr, out := startProxy(t, app.URL, map[string]string{metricsToken: token}) + + got := do(t, newRequest(t, http.MethodPost, addr, "/", strings.NewReader("abc"))) + wantStatus(t, got, http.StatusOK) + wantStatus(t, get(t, addr, "/_smallwebwaf/nothing"), http.StatusNotFound) + out.requestLines(t, 2) + + forward := `{action="forward",status_class="2xx"}` + notFound := `{action="admin",status_class="4xx"}` + + // The request for the metrics is itself under way. + metrics := scrape(t, addr) + wantMetric(t, metrics, "smallwebwaf_requests_total"+forward, 1) + wantMetric(t, metrics, "smallwebwaf_requests_total"+notFound, 1) + wantMetric(t, metrics, "smallwebwaf_request_bytes_total"+forward, 3) + wantMetric(t, metrics, "smallwebwaf_response_bytes_total"+forward, 5) + wantMetric(t, metrics, "smallwebwaf_response_bytes_total"+notFound, + float64(len("Not Found\n"))) + wantMetric(t, metrics, "smallwebwaf_request_duration_seconds_count", 2) + wantMetric(t, metrics, "smallwebwaf_upstream_duration_seconds_count", 1) + wantMetric(t, metrics, "smallwebwaf_requests_in_flight", 1) + metric(t, metrics, "go_goroutines") + metric(t, metrics, "process_start_time_seconds") + + // A request the app holds is under way until it ends. + httpClient := newClient(t) + held := newRequest(t, http.MethodGet, addr, "/held", http.NoBody) + ended := make(chan error, 1) + + go func() { + res, err := httpClient.Do(held) + if err == nil { + err = readAnswer(res).err + } + + ended <- err + }() + + <-arrived + wantMetric(t, scrape(t, addr), "smallwebwaf_requests_in_flight", 2) + releaseApp() + + err := <-ended + if err != nil { + t.Fatalf("held request: %v", err) + } + + out.requestLines(t, 5) + wantMetric(t, scrape(t, addr), "smallwebwaf_requests_in_flight", 1) +} + +func TestMetricsCountLimitsAndBans(t *testing.T) { + t.Parallel() + + const ( + scraper = "192.0.2.200" // in SWWAF_RATE_LIMIT_EXEMPT_NETS + denied = "192.0.2.50" // in SWWAF_DENY_NETS + ) + + s, clk, _ := startWithClock(t, "", map[string]string{ + metricsToken: token, + rateLimitPerMinute: "1", + rateLimitExemptNets: scraper, + denyNets: denied, + banResponse: "close", + limitBanDuration: "1h", + maxBanDuration: "2h", + }) + + // SWWAF_BAN_RESPONSE=close sends no status at all. + s.get(denied, 0, requestlog.ActionDenied) + + // A first broken limit bans for an hour. + s.get(client, http.StatusOK, requestlog.ActionForward) + s.get(client, 0, requestlog.ActionRateLimited) + + metrics := s.scrape(scraper) + wantMetric(t, metrics, + `smallwebwaf_requests_total{action="denied",status_class="none"}`, 1) + wantMetric(t, metrics, `smallwebwaf_rate_limit_hits_total{window="minute"}`, 1) + wantMetric(t, metrics, `smallwebwaf_offences_total{kind="limit"}`, 1) + wantMetric(t, metrics, `smallwebwaf_bans_made_total{cause="limit"}`, 1) + wantMetric(t, metrics, "smallwebwaf_active_bans", 1) + wantMetric(t, metrics, "smallwebwaf_permanent_bans", 0) + + clk.advance(time.Hour) + wantMetric(t, s.scrape(scraper), "smallwebwaf_active_bans", 0) + + // A limit broken again right after would ban for three hours, longer + // than SWWAF_MAX_BAN_DURATION, so the ban is permanent. + s.get(client, http.StatusOK, requestlog.ActionForward) + s.get(client, 0, requestlog.ActionRateLimited) + + metrics = s.scrape(scraper) + wantMetric(t, metrics, `smallwebwaf_rate_limit_hits_total{window="minute"}`, 2) + wantMetric(t, metrics, `smallwebwaf_offences_total{kind="limit"}`, 2) + wantMetric(t, metrics, `smallwebwaf_bans_made_total{cause="limit"}`, 2) + wantMetric(t, metrics, "smallwebwaf_active_bans", 1) + wantMetric(t, metrics, "smallwebwaf_permanent_bans", 1) + // denied, client, and the scraper as of its earlier requests. + wantMetric(t, metrics, "smallwebwaf_tracked_clients", 3) +} + +func TestMetricsByCountryKeepTheBusiestAndCountTheRestAsOther(t *testing.T) { + t.Parallel() + + const fromFR = "198.51.100.20" + + app := startApp(t, func(w http.ResponseWriter, r *http.Request) { + _, _ = io.Copy(io.Discard, r.Body) + _, _ = io.WriteString(w, "hello") + }) + env := map[string]string{ + trustedProxies: trustLocalhost, + metricsToken: token, + metricsTopN: "2", + deniedCountries: "kp", + } + addr, out, server := startProxyWithClock(t, app.URL, "", time.Now, env) + + // The answers are kept before the requests, so that none waits for + // GeoJS. + server.GeoJS.Load([]lookup.Answer{ + keptAnswer(fromKP, "KP"), keptAnswer(fromDE, "DE"), keptAnswer(fromFR, "FR"), + }) + + lines := 0 + send := func(from string, times, status int) { + t.Helper() + + for range times { + req := newRequest(t, http.MethodPost, addr, "/", strings.NewReader("abc")) + req.Header.Set(forwardedFor, from) + wantStatus(t, do(t, req), status) + + // Each is counted before the next is sent, so that the + // countries are ranked in the order sent. + lines++ + out.requestLines(t, lines) + } + } + + // With two countries of their own, the third is counted as other. + send(fromKP, 3, http.StatusForbidden) + send(fromDE, 2, http.StatusOK) + send(fromFR, 1, http.StatusOK) + + metrics := scrape(t, addr) + lines++ + + wantMetric(t, metrics, `smallwebwaf_country_requests_total{country="KP"}`, 3) + wantMetric(t, metrics, `smallwebwaf_country_requests_total{country="DE"}`, 2) + wantMetric(t, metrics, `smallwebwaf_country_requests_total{country="other"}`, 1) + wantMetric(t, metrics, `smallwebwaf_country_list_refusals_total{country="KP"}`, 3) + wantMetric(t, metrics, `smallwebwaf_country_request_bytes_total{country="KP"}`, 0) + wantMetric(t, metrics, `smallwebwaf_country_request_bytes_total{country="DE"}`, 6) + wantMetric(t, metrics, `smallwebwaf_country_response_bytes_total{country="KP"}`, + float64(3*len("Forbidden\n"))) + wantMetric(t, metrics, `smallwebwaf_country_response_bytes_total{country="other"}`, + float64(len("hello"))) + wantNoSeries(t, metrics, `smallwebwaf_country_requests_total{country="FR"}`) + + // Once FR is busier than DE, it takes DE's place: its series counts + // from then on, and DE's is gone. + send(fromFR, 3, http.StatusOK) + + metrics = scrape(t, addr) + wantMetric(t, metrics, `smallwebwaf_country_requests_total{country="KP"}`, 3) + wantMetric(t, metrics, `smallwebwaf_country_requests_total{country="FR"}`, 2) + wantMetric(t, metrics, `smallwebwaf_country_requests_total{country="other"}`, 2) + wantNoSeries(t, metrics, `smallwebwaf_country_requests_total{country="DE"}`) + wantNoSeries(t, metrics, `smallwebwaf_country_request_bytes_total{country="DE"}`) +} + +func TestMetricsCountGeoJSRequestsAndFailures(t *testing.T) { + t.Parallel() + + geojs := httptest.NewServer(http.HandlerFunc( + func(w http.ResponseWriter, _ *http.Request) { + w.WriteHeader(http.StatusServiceUnavailable) + })) + t.Cleanup(geojs.Close) + + app := startApp(t, func(http.ResponseWriter, *http.Request) {}) + addr, _ := startProxyWithGeoJS(t, app.URL, geojs.URL, map[string]string{ + trustedProxies: trustLocalhost, + metricsToken: token, + deniedCountries: "kp", + }) + + // GeoJS fails, so the client counts as coming from an unknown country, + // which SWWAF_DENIED_COUNTRIES does not refuse. + req := newRequest(t, http.MethodGet, addr, "/", http.NoBody) + req.Header.Set(forwardedFor, fromDE) + wantStatus(t, do(t, req), http.StatusOK) + + // The client stops waiting for GeoJS after a second, so GeoJS's + // failure can come after its request has ended. + deadline := time.Now().Add(waitLimit) + metrics := scrape(t, addr) + + for metric(t, metrics, "smallwebwaf_geojs_failures_total") == 0 && + time.Now().Before(deadline) { + time.Sleep(pollInterval) + + metrics = scrape(t, addr) + } + + wantMetric(t, metrics, "smallwebwaf_geojs_requests_total", 1) + wantMetric(t, metrics, "smallwebwaf_geojs_failures_total", 1) + wantMetric(t, metrics, "smallwebwaf_geojs_unanswered_total", 1) +} + +// keptAnswer returns GeoJS's answer that the client at addr is in +// country, given now. +func keptAnswer(addr, country string) lookup.Answer { + now := time.Now() + + return lookup.Answer{ + Client: netip.MustParsePrefix(addr + "/32"), Country: country, + Answered: now, Used: now, + } +} + +// scrape asks smallwebwaf at addr for the metrics, with the token, and +// returns them. +func scrape(t *testing.T, addr string) string { + t.Helper() + + req := newRequest(t, http.MethodGet, addr, proxy.MetricsPath, http.NoBody) + req.Header.Set("Authorization", bearer) + + got := do(t, req) + if got.status != http.StatusOK { + t.Fatalf("the metrics were answered %d", got.status) + } + + return string(got.body) +} + +// scrape asks for the metrics, with the token, from the client at from, +// and returns them. +func (s *sender) scrape(from string) string { + s.t.Helper() + + _, metrics := s.requestWithHeader(from, proxy.MetricsPath, "Authorization: "+bearer, + http.StatusOK, requestlog.ActionAdmin) + + return metrics +} + +// metric returns the value of series in metrics, which are in the +// Prometheus text format. series is a name and its labels in the order of +// their names, such as smallwebwaf_offences_total{kind="limit"}. It fails +// the test if there is no such series. +func metric(t *testing.T, metrics, series string) float64 { + t.Helper() + + for line := range strings.Lines(metrics) { + value, found := strings.CutPrefix(strings.TrimSuffix(line, "\n"), series+" ") + if !found { + continue + } + + number, err := strconv.ParseFloat(value, 64) + if err != nil { + t.Fatalf("%s has the value %q", series, value) + } + + return number + } + + t.Fatalf("no series %s in the metrics:\n%s", series, metrics) + + return 0 +} + +// wantMetric checks the value of series in metrics, as metric reads it. +func wantMetric(t *testing.T, metrics, series string, want float64) { + t.Helper() + + got := metric(t, metrics, series) + if got != want { + t.Errorf("%s is %v, want %v", series, got, want) + } +} + +// wantNoSeries checks that metrics have no series series. +func wantNoSeries(t *testing.T, metrics, series string) { + t.Helper() + + if strings.Contains(metrics, "\n"+series+" ") { + t.Errorf("there is a series %s", series) + } +} + +// wantLimitHits checks that the metrics of smallwebwaf at addr count hits +// requests that passed the size or time limit of the setting limit, with +// no series for it when hits is 0. +func wantLimitHits(t *testing.T, addr, limit string, hits int) { + t.Helper() + + series := `smallwebwaf_size_and_time_limit_hits_total{limit="` + limit + `"}` + metrics := scrape(t, addr) + + if hits == 0 { + wantNoSeries(t, metrics, series) + + return + } + + wantMetric(t, metrics, series, float64(hits)) +} diff --git a/internal/proxy/observe_test.go b/internal/proxy/observe_test.go new file mode 100644 index 0000000..7a87b7b --- /dev/null +++ b/internal/proxy/observe_test.go @@ -0,0 +1,202 @@ +package proxy_test + +import ( + "bytes" + "io" + "net/http" + "net/netip" + "sync/atomic" + "testing" + "time" + + "sneak.berlin/go/smallwebwaf/internal/bans" + "sneak.berlin/go/smallwebwaf/internal/proxy" + "sneak.berlin/go/smallwebwaf/internal/requestlog" +) + +// observe is the value of SWWAF_MODE for observe mode. +const observe = "observe" + +func TestObserveModeForwardsWhatEnforceModeRefuses(t *testing.T) { + t.Parallel() + + const ( + denied = "192.0.2.50" // in SWWAF_DENY_NETS + banned = otherClient // under a ban read from bans.json + ) + + for _, tc := range []struct { + setting string // "" leaves SWWAF_MODE at its default + observe bool + }{ + {"", false}, + {"enforce", false}, + {observe, true}, + } { + t.Run(mode+"="+tc.setting, func(t *testing.T) { + t.Parallel() + + geojsURL, _ := startGeoJS(t) + env := map[string]string{ + rateLimitPerMinute: "1", + denyNets: denied, + deniedCountries: "kp", + } + + if tc.setting != "" { + env[mode] = tc.setting + } + + s, clk, server := startWithClock(t, geojsURL, env) + server.Ledger.Load([]bans.Ban{{ + Netblock: netip.MustParsePrefix(banned + "/32"), + Start: clk.Now(), + Expires: clk.Now().Add(time.Hour), + }}) + + // fromDE's first request is within the limit of one a minute, + // and its second breaks it. + s.get(fromDE, http.StatusOK, requestlog.ActionForward) + + for _, sent := range []struct{ from, refusal string }{ + {denied, requestlog.ActionDenied}, + {banned, requestlog.ActionBanned}, + {fromKP, requestlog.ActionCountryDenied}, + {fromDE, requestlog.ActionRateLimited}, + } { + if !tc.observe { + line := s.get(sent.from, http.StatusForbidden, sent.refusal) + wantWouldAction(t, line, "") + + continue + } + + // Passed to the app, which answered it. + line := s.get(sent.from, http.StatusOK, requestlog.ActionForward) + wantWouldAction(t, line, sent.refusal) + + if line.UpstreamStatus != http.StatusOK { + t.Errorf("log line has upstream_status %d, want 200", + line.UpstreamStatus) + } + } + }) + } +} + +func TestObserveModeMakesNoBanAndKeepsTheBansItHas(t *testing.T) { + t.Parallel() + + s, clk, server := startWithClock(t, "", map[string]string{ + mode: observe, + rateLimitPerMinute: "1", + }) + kept := bans.Ban{ + Netblock: netip.MustParsePrefix(otherClient + "/32"), + Start: clk.Now(), + Expires: clk.Now().Add(time.Hour), + } + server.Ledger.Load([]bans.Ban{kept}) + + // No ban sets client's counters back to zero, so each request after + // the first breaks the limit of one a minute. + s.get(client, http.StatusOK, requestlog.ActionForward) + + for range 2 { + line := s.get(client, http.StatusOK, requestlog.ActionForward) + wantWouldAction(t, line, requestlog.ActionRateLimited) + + if line.LimitHit != minute || line.Offence != requestlog.OffenceLimit || + line.BanExpires != "" { + t.Errorf("log line has limit_hit %q, offence %q and ban_expires %q, "+ + "want minute, limit and none", line.LimitHit, line.Offence, + line.BanExpires) + } + } + + // The ban read from bans.json refuses nothing, and so counts no + // refusal in its notes, but is kept. + line := s.get(otherClient, http.StatusOK, requestlog.ActionForward) + wantWouldAction(t, line, requestlog.ActionBanned) + + if line.BanExpires != requestlog.FormatTime(kept.Expires) { + t.Errorf("log line has ban_expires %q, want %s", line.BanExpires, + requestlog.FormatTime(kept.Expires)) + } + + got := server.Ledger.Snapshot() + if len(got) != 1 || got[0] != kept { + t.Errorf("bans\n%+v\nwant only\n%+v", got, kept) + } +} + +func TestObserveModeKeepsTheSizeLimitsAndTheToken(t *testing.T) { + t.Parallel() + + const denied = "192.0.2.50" // in SWWAF_DENY_NETS + + var calls atomic.Int32 + + app := startApp(t, func(w http.ResponseWriter, _ *http.Request) { + calls.Add(1) + answerWithSize(w, 2*sizeLimit, true) + }) + addr, out := startProxy(t, app.URL, map[string]string{ + mode: observe, + trustedProxies: trustLocalhost, + denyNets: denied, + requestMaxBytes: sizeLimitSetting, + responseMaxBytes: sizeLimitSetting, + metricsToken: token, + }) + + // SWWAF_DENY_NETS would refuse each request; instead a size limit or + // the missing token does. + for i, tc := range []struct { + method, path string + body io.Reader + status int + action string + }{ + { + http.MethodPost, "/upload", bytes.NewReader(make([]byte, 2*sizeLimit)), + http.StatusRequestEntityTooLarge, requestlog.ActionTooLarge, + }, + { + http.MethodGet, "/download", http.NoBody, + http.StatusBadGateway, requestlog.ActionTooLarge, + }, + { + http.MethodGet, proxy.MetricsPath, http.NoBody, + http.StatusUnauthorized, requestlog.ActionAdmin, + }, + } { + req := newRequest(t, tc.method, addr, tc.path, tc.body) + req.Header.Set(forwardedFor, denied) + wantStatus(t, do(t, req), tc.status) + + line := out.requestLines(t, i+1)[i] + wantLine(t, line, tc.status, tc.action) + wantWouldAction(t, line, requestlog.ActionDenied) + } + + // The upload was refused before it reached the app. + if calls.Load() != 1 { + t.Errorf("the app was called %d times, want once", calls.Load()) + } +} + +// wantWouldAction checks the request log line's would_action, and that a +// line that should have none has no such field. +func wantWouldAction(t *testing.T, line logLine, want string) { + t.Helper() + + got, present := line.fields["would_action"] + + switch { + case want == "" && present: + t.Errorf("log line has would_action %v, want none", got) + case want != "" && got != want: + t.Errorf("log line has would_action %v, want %s", got, want) + } +} diff --git a/internal/proxy/passthrough_test.go b/internal/proxy/passthrough_test.go new file mode 100644 index 0000000..5a7f015 --- /dev/null +++ b/internal/proxy/passthrough_test.go @@ -0,0 +1,447 @@ +package proxy_test + +import ( + "bufio" + "bytes" + "errors" + "io" + "net/http" + "os" + "reflect" + "slices" + "strings" + "sync/atomic" + "testing" + "time" + + "sneak.berlin/go/smallwebwaf/internal/config" + "sneak.berlin/go/smallwebwaf/internal/proxy" + "sneak.berlin/go/smallwebwaf/internal/ratelimit" + "sneak.berlin/go/smallwebwaf/internal/requestlog" +) + +// A request target with an escaped slash and space in its path, and a +// query with a parameter ReverseProxy cannot parse. +const ( + rawPath = "/some%2Fpath/with%20space" + rawQuery = "b=2&a=1&bad=%zz;x" +) + +// chunkSize is the size of each part of a body a test sends in parts. +const chunkSize = 1 << 10 + +var errNotStreamed = errors.New("the first part never reached the app") + +// appSaw is what the app received. +type appSaw struct { + method string + target string + header http.Header + body []byte +} + +func TestPassesRequestAndAnswerUnchanged(t *testing.T) { + t.Parallel() + + requestBody := bytes.Repeat([]byte("request body "), 8000) + answerBody := bytes.Repeat([]byte("answer body "), 8000) + saw := make(chan appSaw, 1) + + app := startApp(t, func(w http.ResponseWriter, r *http.Request) { + body, _ := io.ReadAll(r.Body) + saw <- appSaw{r.Method, r.RequestURI, r.Header.Clone(), body} + + w.Header().Set("X-App", "yes") + w.Header().Add("Set-Cookie", "a=1") + w.Header().Add("Set-Cookie", "b=2") + w.WriteHeader(http.StatusTeapot) + _, _ = w.Write(answerBody) + }) + addr, out := startProxy(t, app.URL, nil) + + req := newRequest(t, http.MethodPatch, addr, rawPath+"?"+rawQuery, + bytes.NewReader(requestBody)) + req.Header.Add("X-Test", "one") + req.Header.Add("X-Test", "two") + req.Header.Set("User-Agent", "test-agent") + + got := do(t, req) + + wantAppSaw(t, <-saw, requestBody) + wantAnswer(t, got, answerBody) + + line := out.requestLine(t) + wantLine(t, line, http.StatusTeapot, requestlog.ActionForward) + wantRequestFields(t, line, addr, len(requestBody), len(answerBody)) +} + +// wantAppSaw checks that the app received the test's request unchanged. +func wantAppSaw(t *testing.T, saw appSaw, body []byte) { + t.Helper() + + if saw.method != http.MethodPatch || saw.target != rawPath+"?"+rawQuery { + t.Errorf("app saw %s %s, want %s %s", saw.method, saw.target, + http.MethodPatch, rawPath+"?"+rawQuery) + } + + if !slices.Equal(saw.header.Values("X-Test"), []string{"one", "two"}) { + t.Errorf("app saw X-Test %q", saw.header.Values("X-Test")) + } + + if saw.header.Get("User-Agent") != "test-agent" { + t.Errorf("app saw User-Agent %q", saw.header.Get("User-Agent")) + } + + if !bytes.Equal(saw.body, body) { + t.Errorf("app saw a body of %d bytes, want the %d sent", + len(saw.body), len(body)) + } +} + +// wantAnswer checks that the client received the app's answer unchanged. +func wantAnswer(t *testing.T, got answer, body []byte) { + t.Helper() + + wantStatus(t, got, http.StatusTeapot) + + if got.header.Get("X-App") != "yes" { + t.Errorf("client got X-App %q", got.header.Get("X-App")) + } + + if !slices.Equal(got.header.Values("Set-Cookie"), []string{"a=1", "b=2"}) { + t.Errorf("client got Set-Cookie %q", got.header.Values("Set-Cookie")) + } + + if got.err != nil || !bytes.Equal(got.body, body) { + t.Errorf("client got %d bytes (%v), want the %d the app sent", + len(got.body), got.err, len(body)) + } +} + +// wantRequestFields checks the log line's fields about the request. Its +// time, its id and its timings are checked only for being there. +func wantRequestFields(t *testing.T, line logLine, host string, sent, received int) { + t.Helper() + + hostname, _ := os.Hostname() + + want := withTimings(line, requestlog.Line{ + Type: requestType, Time: line.Time, Instance: hostname, + ClientIP: localhost, Method: http.MethodPatch, Scheme: plain, Host: host, + Path: rawPath, Query: rawQuery, Protocol: protocol, + Status: http.StatusTeapot, RequestBytes: int64(sent), + ResponseBytes: int64(received), UserAgent: "test-agent", + RequestID: line.RequestID, PeerIP: localhost, ClientGroup: localhost + "/32", + ContentLength: int64(sent), ResponseContentType: "text/plain; charset=utf-8", + UpstreamStatus: http.StatusTeapot, Action: requestlog.ActionForward, + Counts: ratelimit.Counts{Minute: 1, Hour: 1, Day: 1}, + }) + if !reflect.DeepEqual(line.Line, want) { + t.Errorf("log line\n%+v\nwant\n%+v", line.Line, want) + } + + _, err := time.Parse(time.RFC3339, line.Time) + if err != nil || line.RequestID == "" || line.DurationTotal <= 0 || + line.DurationUpstreamTotal == nil || *line.DurationUpstreamTotal <= 0 { + t.Errorf("log line has time %q, request_id %q and durations %v and %v", + line.Time, line.RequestID, line.DurationTotal, + line.fields["duration_upstream_total"]) + } +} + +func TestStreamsTheRequestBody(t *testing.T) { + t.Parallel() + + chunk := bytes.Repeat([]byte("x"), chunkSize) + firstArrived := make(chan struct{}) + + app := startApp(t, func(w http.ResponseWriter, r *http.Request) { + first := make([]byte, len(chunk)) + + _, err := io.ReadFull(r.Body, first) + if err != nil { + return + } + + close(firstArrived) + + rest, _ := io.ReadAll(r.Body) + _, _ = w.Write(rest) + }) + addr, _ := startProxy(t, app.URL, nil) + + body, writer := io.Pipe() + + go func() { + _, _ = writer.Write(chunk) + + select { + case <-firstArrived: + _, _ = writer.Write(chunk) + _ = writer.Close() + case <-time.After(waitLimit): + _ = writer.CloseWithError(errNotStreamed) + } + }() + + got := do(t, newRequest(t, http.MethodPost, addr, "/upload", body)) + if got.err != nil || !bytes.Equal(got.body, chunk) { + t.Errorf("app read %d bytes after the first part (%v), want %d", + len(got.body), got.err, len(chunk)) + } +} + +func TestStreamsTheAnswerBody(t *testing.T) { + t.Parallel() + + chunk := bytes.Repeat([]byte("y"), chunkSize) + firstArrived := make(chan struct{}) + + app := startApp(t, func(w http.ResponseWriter, _ *http.Request) { + _, _ = w.Write(chunk) + _ = http.NewResponseController(w).Flush() + + select { + case <-firstArrived: + _, _ = w.Write(chunk) + case <-time.After(waitLimit): + } + }) + addr, _ := startProxy(t, app.URL, nil) + + req := newRequest(t, http.MethodGet, addr, "/download", http.NoBody) + + res, err := newClient(t).Do(req) + if err != nil { + t.Fatalf("request: %v", err) + } + + first := make([]byte, len(chunk)) + _, err = io.ReadFull(res.Body, first) + + close(firstArrived) + + got := readAnswer(res) + if err != nil || got.err != nil || !bytes.Equal(got.body, chunk) { + t.Errorf("client read %d bytes after the first part (%v, %v), want %d", + len(got.body), err, got.err, len(chunk)) + } +} + +func TestUpgradedConnectionOutlastsTheTimeouts(t *testing.T) { + t.Parallel() + + app := startApp(t, echoAfterUpgrade) + addr, out := startProxy(t, app.URL, map[string]string{ + clientRequestTimeout: shortTimeoutSetting, + clientResponseTimeout: shortTimeoutSetting, + upstreamRequestTimeout: shortTimeoutSetting, + upstreamResponseTimeout: shortTimeoutSetting, + }) + + conn := dial(t, addr) + send(t, conn, "GET /socket HTTP/1.1\r\nHost: app\r\n"+ + "Connection: Upgrade\r\nUpgrade: websocket\r\n\r\n") + + reader := bufio.NewReader(conn) + + res, err := http.ReadResponse(reader, nil) + if err != nil { + t.Fatalf("read the answer to the upgrade: %v", err) + } + + answered := time.Now() + _ = res.Body.Close() + + if res.StatusCode != http.StatusSwitchingProtocols { + t.Fatalf("status %d, want %d", res.StatusCode, http.StatusSwitchingProtocols) + } + + // Every timeout started before the upgrade was answered, the response + // timeouts last, at the end of the request: wait until just past + // shortTimeout after the answer was read, then use the connection. + time.Sleep(time.Until(answered.Add(shortTimeout + 100*time.Millisecond))) + send(t, conn, "still here\n") + + echoed, err := reader.ReadString('\n') + if err != nil || echoed != "still here\n" { + t.Errorf("echo %q (%v), want %q", echoed, err, "still here\n") + } + + _ = conn.Close() + + wantLine(t, out.requestLine(t), http.StatusSwitchingProtocols, + requestlog.ActionForward) +} + +// echoAfterUpgrade is an app that switches protocols on request, and then +// sends back each line it receives. +func echoAfterUpgrade(w http.ResponseWriter, r *http.Request) { + if r.Header.Get("Upgrade") != "websocket" { + http.Error(w, "not an upgrade", http.StatusBadRequest) + + return + } + + conn, buffered, err := http.NewResponseController(w).Hijack() + if err != nil { + return + } + + defer func() { + _ = conn.Close() + }() + + _, _ = buffered.WriteString("HTTP/1.1 101 Switching Protocols\r\n" + + "Connection: Upgrade\r\nUpgrade: websocket\r\n\r\n") + _ = buffered.Flush() + + for { + line, err := buffered.ReadString('\n') + if err != nil { + return + } + + _, _ = buffered.WriteString(line) + _ = buffered.Flush() + } +} + +func TestServerHasTheDefaultLimits(t *testing.T) { + t.Parallel() + + cfg, err := config.FromEnvironment(func(string) (string, bool) { return "", false }) + if err != nil { + t.Fatalf("default settings: %v", err) + } + + server := proxy.New(proxy.Params{ + Config: cfg, + RequestLog: io.Discard, + ProcessLog: requestlog.NewProcessLogger(io.Discard), + }) + + if server.Addr != ":8080" || server.MaxHeaderBytes != 28<<10 || + server.IdleTimeout != 2*time.Minute || server.ReadHeaderTimeout != time.Minute { + t.Errorf("server listens on %q with header limit %d, idle time %s and "+ + "header timeout %s", server.Addr, server.MaxHeaderBytes, + server.IdleTimeout, server.ReadHeaderTimeout) + } +} + +func TestRefusesHeadersOverTheLimit(t *testing.T) { + t.Parallel() + + for _, tc := range []struct { + name string + env map[string]string + limit int + }{ + {"by default", nil, 32 << 10}, + {"as set", map[string]string{clientHeaderMaxBytes: "8K"}, 8 << 10}, + } { + 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) + }) + addr, _ := startProxy(t, app.URL, tc.env) + + // size counts every byte of the request: the request line, + // the headers and the blank line that ends them. + const ( + start = "GET / HTTP/1.1\r\nHost: app\r\nX-Large: " + end = "\r\n\r\n" + ) + + for _, sent := range []struct{ size, want int }{ + {tc.limit, http.StatusOK}, + {tc.limit + 1, http.StatusRequestHeaderFieldsTooLarge}, + } { + conn := dial(t, addr) + send(t, conn, + start+strings.Repeat("a", sent.size-len(start)-len(end))+end) + wantStatus(t, readResponse(t, conn), sent.want) + } + + if calls.Load() != 1 { + t.Errorf("the app was called %d times, want once", calls.Load()) + } + }) + } +} + +func TestAnswers502WhenTheAppCannotBeReached(t *testing.T) { + t.Parallel() + + // No test can listen on port 1: listening on port 0 gets one from 32768 up. + addr, out := startProxy(t, "http://"+localhost+":1", nil) + + wantStatus(t, get(t, addr, "/"), http.StatusBadGateway) + + line := out.requestLine(t) + wantLine(t, line, http.StatusBadGateway, requestlog.ActionUpstreamError) + + // There never was a connection to the app, nor an answer from it. + wantTimings(t, line, "duration_total", "duration_checks", + "duration_upstream_total") + + logged := slices.ContainsFunc(out.lines(t), func(line map[string]any) bool { + return line["type"] == "process" && line["msg"] == "request to the app failed" + }) + if !logged { + t.Errorf("no process line says the request to the app failed") + } +} + +func TestLogsAnAnswerThatBrokeOff(t *testing.T) { + t.Parallel() + + app := startApp(t, func(w http.ResponseWriter, _ *http.Request) { + _, _ = io.WriteString(w, "the first part") + _ = http.NewResponseController(w).Flush() + + panic(http.ErrAbortHandler) // drops the connection mid-answer + }) + addr, out := startProxy(t, app.URL, nil) + + got := get(t, addr, "/") + if string(got.body) != "the first part" || !errors.Is(got.err, io.ErrUnexpectedEOF) { + t.Errorf("client read %q (%v), want the first part cut off", got.body, got.err) + } + + wantLine(t, out.requestLine(t), http.StatusOK, requestlog.ActionUpstreamError) +} + +func TestLogsAClientThatWentAway(t *testing.T) { + t.Parallel() + + arrived := make(chan struct{}) + + app := startApp(t, func(_ http.ResponseWriter, r *http.Request) { + close(arrived) + <-r.Context().Done() + }) + addr, out := startProxy(t, app.URL, nil) + + conn := dial(t, addr) + send(t, conn, "GET /slow HTTP/1.1\r\nHost: app\r\n\r\n") + + select { + case <-arrived: + case <-time.After(waitLimit): + t.Fatal("the request never reached the app") + } + + _ = conn.Close() + + line := out.requestLine(t) + if !line.Aborted || line.Status != 0 || line.Action != requestlog.ActionForward { + t.Errorf("log line has aborted %v, status %d and action %q, "+ + "want true, 0 and %q", + line.Aborted, line.Status, line.Action, requestlog.ActionForward) + } +} diff --git a/internal/proxy/proxy.go b/internal/proxy/proxy.go new file mode 100644 index 0000000..a651f4d --- /dev/null +++ b/internal/proxy/proxy.go @@ -0,0 +1,192 @@ +// Package proxy passes each request to the app and the app's answer back, +// unchanged, within the size and time limits, and writes one request log +// line for each request. +package proxy + +import ( + "io" + "log" + "log/slog" + "net/http" + "strings" + "time" + + "sneak.berlin/go/smallwebwaf/internal/bans" + "sneak.berlin/go/smallwebwaf/internal/config" + "sneak.berlin/go/smallwebwaf/internal/lookup" + "sneak.berlin/go/smallwebwaf/internal/metrics" + "sneak.berlin/go/smallwebwaf/internal/ratelimit" + "sneak.berlin/go/smallwebwaf/internal/requestlog" +) + +// How smallwebwaf keeps connections to the app open between requests. +const ( + appIdleConns = 100 + appIdleConnTimeout = 90 * time.Second +) + +// adminPrefix starts the path of every request for smallwebwaf itself, +// which never reaches the app. +const adminPrefix = "/_smallwebwaf/" + +// HealthPath is smallwebwaf's health endpoint, which the container's +// health check asks. +const HealthPath = "/_smallwebwaf/healthz" + +// MetricsPath is where the metrics are, for a request that carries +// SWWAF_METRICS_TOKEN. +const MetricsPath = "/_smallwebwaf/metrics" + +// Params are what New needs. +type Params struct { + Config *config.Config + // RequestLog receives one JSON line per request. + RequestLog io.Writer + // ProcessLog receives the process's own messages. + ProcessLog *slog.Logger + // GeoJSURL is where clients' countries are looked up, normally + // lookup.URL. GeoJS is asked only while a country list is set. + GeoJSURL string + // Now tells the time by which requests are counted for the rate + // limits, bans are made and run out, and GeoJS's answers are kept, + // normally time.Now in UTC, the time the state files give. + Now func() time.Time +} + +// Server is the server smallwebwaf runs, with the parts of the proxy +// whose state the state files keep, and the metrics. +type Server struct { + *http.Server + + Ledger *bans.Ledger + Limiter *ratelimit.Limiter + GeoJS *lookup.GeoJS + Metrics *metrics.Metrics +} + +// New returns the server smallwebwaf runs: each request it reads passes +// through the proxy. Go's server itself refuses a request line and +// headers over SWWAF_CLIENT_REQUEST_HEADER_MAX_BYTES, with 431, closes a +// connection idle for SWWAF_CLIENT_IDLE_TIMEOUT, and applies +// SWWAF_CLIENT_REQUEST_TIMEOUT while the headers arrive; the proxy +// applies the timeouts and size limits from then on. +func New(params Params) *Server { + errorLog := slog.NewLogLogger(params.ProcessLog.Handler(), slog.LevelWarn) + m := metrics.New(params.Config.MetricsTopN) + h := &handler{ + config: params.Config, + requestLog: params.RequestLog, + processLog: params.ProcessLog, + errorLog: errorLog, + transport: newTransport(), + now: params.Now, + metrics: m, + limiter: ratelimit.New(ratelimit.Limits{ + PerMinute: params.Config.RateLimitPerMinute, + PerHour: params.Config.RateLimitPerHour, + PerDay: params.Config.RateLimitPerDay, + }), + ledger: bans.New(bans.Rules{ + LimitBanDuration: params.Config.LimitBanDuration, + LimitBanRepeatWindow: params.Config.LimitBanRepeatWindow, + MaxBanDuration: params.Config.MaxBanDuration, + MaxBans: params.Config.MaxBans, + }), + geojs: lookup.New(lookup.Params{ + URL: params.GeoJSURL, + Now: params.Now, + ProcessLog: params.ProcessLog, + Metrics: m, + }), + } + m.AddBansAndClients(h.ledger, h.limiter, params.Now) + + return &Server{ + Server: &http.Server{ + Addr: params.Config.ListenAddr, + Handler: h, + ReadHeaderTimeout: params.Config.ClientRequestTimeout, + // Off is an IdleTimeout of 0, which Go's server replaces with + // ReadTimeout: no limit, as long as ReadTimeout stays unset. + IdleTimeout: params.Config.ClientIdleTimeout, + // Go's server reads 4 KiB past MaxHeaderBytes before it + // refuses, so the limit a client meets is the setting. + MaxHeaderBytes: int(params.Config.ClientRequestHeaderMaxBytes - 4<<10), + ErrorLog: errorLog, + }, + Ledger: h.ledger, + Limiter: h.limiter, + GeoJS: h.geojs, + Metrics: m, + } +} + +// handler is the proxy. It holds what every request shares; what belongs +// to one request is in a request. +type handler struct { + config *config.Config + requestLog io.Writer + processLog *slog.Logger + errorLog *log.Logger + transport http.RoundTripper + now func() time.Time + metrics *metrics.Metrics + limiter *ratelimit.Limiter + ledger *bans.Ledger + geojs *lookup.GeoJS +} + +// newTransport returns what carries requests to the app. It never goes +// through a proxy named in the environment, and leaves the app's answers +// compressed or not as the app sent them. +func newTransport() *http.Transport { + return &http.Transport{ + MaxIdleConns: appIdleConns, + MaxIdleConnsPerHost: appIdleConns, + IdleConnTimeout: appIdleConnTimeout, + DisableCompression: true, + } +} + +// ServeHTTP handles one request: it works out the client, runs the +// checks, passes the request to the app and the answer back within the +// limits, or answers it itself if it is for smallwebwaf, and writes the +// request's log line. +func (h *handler) ServeHTTP(w http.ResponseWriter, r *http.Request) { + rq := h.newRequest(w, r) + defer rq.finish() + + // The health endpoint is answered at once, before any check, so that + // a health checker is never refused. It does not ask the app. + if r.Method == http.MethodGet && r.URL.Path == HealthPath { + rq.line.Action = requestlog.ActionAdmin + // Set here rather than left to Go's server, which would set it only + // after the log line has taken the response's headers. + rq.out.Header().Set("Content-Type", "text/plain; charset=utf-8") + _, _ = io.WriteString(rq.out, "ok\n") + + return + } + + // Once the request has ended, before its log line is written. + defer rq.addToHistory() + + refused := rq.check(r.Context()) + rq.checked = time.Now() + + if refused != nil { + rq.answer(*refused) + + return + } + + // A request for smallwebwaf itself is answered where another would be + // passed to the app, so that it goes through every check first. + if strings.HasPrefix(r.URL.Path, adminPrefix) { + rq.answerAdmin() + + return + } + + rq.forward(r.Context()) +} diff --git a/internal/proxy/proxy_test.go b/internal/proxy/proxy_test.go new file mode 100644 index 0000000..ccb3bed --- /dev/null +++ b/internal/proxy/proxy_test.go @@ -0,0 +1,390 @@ +package proxy_test + +import ( + "bufio" + "bytes" + "encoding/json" + "io" + "maps" + "net" + "net/http" + "net/http/httptest" + "strings" + "sync" + "testing" + "time" + + "sneak.berlin/go/smallwebwaf/internal/config" + "sneak.berlin/go/smallwebwaf/internal/proxy" + "sneak.berlin/go/smallwebwaf/internal/requestlog" +) + +const ( + // shortTimeout is what a test sets a timeout to, to see it run out. + // It starts before the test has set up its case, such as an upgrade + // or the app's buffers filling, so it is as long as the hold-up of the + // test process that wantTimedOut allows: a shorter one can run out + // first on a busy host. + shortTimeout = waitLimit / 2 + // longTimeoutSetting is a timeout that does not run out in a test. + longTimeoutSetting = "1m" + // waitLimit bounds how long a test waits for what should happen. + waitLimit = 10 * time.Second + // pollInterval is how often a test looks for a log line. + pollInterval = 10 * time.Millisecond + // localhost is where every test server listens, and so the address + // smallwebwaf sees each test's requests come from. + localhost = "127.0.0.1" + // requestType is the type that marks a request log line. + requestType = "request" + // protocol is the protocol of every test's requests. + protocol = "HTTP/1.1" +) + +// shortTimeoutSetting is shortTimeout as a setting's value. +// +//nolint:gochecknoglobals // a constant cannot call String +var shortTimeoutSetting = shortTimeout.String() + +// The settings the tests set. +const ( + clientRequestTimeout = "SWWAF_CLIENT_REQUEST_TIMEOUT" + clientHeaderMaxBytes = "SWWAF_CLIENT_REQUEST_HEADER_MAX_BYTES" + clientIdleTimeout = "SWWAF_CLIENT_IDLE_TIMEOUT" + clientResponseTimeout = "SWWAF_CLIENT_RESPONSE_TIMEOUT" + upstreamRequestTimeout = "SWWAF_UPSTREAM_REQUEST_TIMEOUT" + upstreamResponseTimeout = "SWWAF_UPSTREAM_RESPONSE_TIMEOUT" + mode = "SWWAF_MODE" + requestMaxBytes = "SWWAF_REQUEST_MAX_BYTES" + responseMaxBytes = "SWWAF_RESPONSE_MAX_BYTES" + trustedProxies = "SWWAF_TRUSTED_PROXIES" + allowNets = "SWWAF_ALLOW_NETS" + rateLimitExemptNets = "SWWAF_RATE_LIMIT_EXEMPT_NETS" + denyNets = "SWWAF_DENY_NETS" + rateLimitPerMinute = "SWWAF_RATE_LIMIT_PER_MINUTE" + rateLimitPerDay = "SWWAF_RATE_LIMIT_PER_DAY" + deniedCountries = "SWWAF_DENIED_COUNTRIES" + allowedCountries = "SWWAF_EXCLUSIVELY_ALLOWED_COUNTRIES" + banResponse = "SWWAF_BAN_RESPONSE" + limitBanDuration = "SWWAF_LIMIT_BAN_DURATION" + limitBanRepeatWindow = "SWWAF_LIMIT_BAN_REPEAT_WINDOW" + maxBanDuration = "SWWAF_MAX_BAN_DURATION" + maxBans = "SWWAF_MAX_BANS" + banScopeV4Prefix = "SWWAF_BAN_SCOPE_V4_PREFIX" + instanceName = "SWWAF_INSTANCE_NAME" + logRequestHeaders = "SWWAF_LOG_REQUEST_HEADERS" +) + +// output collects what smallwebwaf writes on stdout. +type output struct { + mu sync.Mutex + buf bytes.Buffer +} + +// Write adds lines smallwebwaf writes. +func (o *output) Write(p []byte) (int, error) { + o.mu.Lock() + defer o.mu.Unlock() + + return o.buf.Write(p) +} + +// text returns everything written so far. +func (o *output) text() string { + o.mu.Lock() + defer o.mu.Unlock() + + return o.buf.String() +} + +// lines returns every line written so far, decoded. +func (o *output) lines(t *testing.T) []map[string]any { + t.Helper() + o.mu.Lock() + defer o.mu.Unlock() + + var lines []map[string]any + + for text := range strings.Lines(o.buf.String()) { + var line map[string]any + + err := json.Unmarshal([]byte(text), &line) + if err != nil { + t.Fatalf("output line %q is not JSON: %v", text, err) + } + + lines = append(lines, line) + } + + return lines +} + +// logLine is a request log line, as typed fields and as the JSON object +// it was written as. +type logLine struct { + requestlog.Line + + fields map[string]any +} + +// requestLines waits for count request log lines and returns them. +func (o *output) requestLines(t *testing.T, count int) []logLine { + t.Helper() + + deadline := time.Now().Add(waitLimit) + for time.Now().Before(deadline) { + var found []logLine + + for _, fields := range o.lines(t) { + if fields["type"] == requestType { + found = append(found, decodeLine(t, fields)) + } + } + + if len(found) >= count { + return found + } + + time.Sleep(pollInterval) + } + + t.Fatalf("fewer than %d request log lines after %s", count, waitLimit) + + return nil +} + +// requestLine waits for the request log line of a test's one request. +func (o *output) requestLine(t *testing.T) logLine { + t.Helper() + + return o.requestLines(t, 1)[0] +} + +// decodeLine reads a request log line's fields into a logLine. +func decodeLine(t *testing.T, fields map[string]any) logLine { + t.Helper() + + encoded, err := json.Marshal(fields) + if err != nil { + t.Fatalf("encode %v: %v", fields, err) + } + + line := logLine{fields: fields} + + err = json.Unmarshal(encoded, &line.Line) + if err != nil { + t.Fatalf("decode %s: %v", encoded, err) + } + + return line +} + +// startApp starts app as the app smallwebwaf passes requests to. +func startApp(t *testing.T, app http.HandlerFunc) *httptest.Server { + t.Helper() + + server := httptest.NewServer(app) + t.Cleanup(server.Close) + + return server +} + +// startProxy starts smallwebwaf in front of the app at appURL, with the +// settings in env on top of the defaults, and returns where it listens and +// what it writes. +func startProxy(t *testing.T, appURL string, env map[string]string) (string, *output) { + t.Helper() + + 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() + + addr, out, _ := startProxyWithClock(t, appURL, geojsURL, time.Now, env) + + return addr, out +} + +// startProxyWithClock is startProxyWithGeoJS with requests counted and +// bans made by the time now tells, and returns the server as well. +func startProxyWithClock( + t *testing.T, appURL, geojsURL string, now func() time.Time, + env map[string]string, +) (string, *output, *proxy.Server) { + t.Helper() + + settings := map[string]string{"SWWAF_UPSTREAM_URL": appURL} + maps.Copy(settings, env) + + cfg, err := config.FromEnvironment(func(name string) (string, bool) { + value, ok := settings[name] + + return value, ok + }) + if err != nil { + t.Fatalf("settings %v: %v", settings, err) + } + + out := &output{} + server := proxy.New(proxy.Params{ + Config: cfg, + RequestLog: out, + ProcessLog: requestlog.NewProcessLogger(out), + GeoJSURL: geojsURL, + Now: now, + }) + + listener, err := (&net.ListenConfig{}).Listen(t.Context(), "tcp", localhost+":0") + if err != nil { + t.Fatalf("listen: %v", err) + } + + go func() { + _ = server.Serve(listener) + }() + + t.Cleanup(func() { + _ = server.Close() + }) + + return listener.Addr().String(), out, server +} + +// newClient returns an HTTP client that sends requests as they are made, +// with no compression of its own. +func newClient(t *testing.T) *http.Client { + t.Helper() + + transport := &http.Transport{DisableCompression: true} + t.Cleanup(transport.CloseIdleConnections) + + return &http.Client{Transport: transport} +} + +// answer is a response as a test reads it: the status, the headers, as +// much of the body as arrived, and the error that ended the reading, nil +// when the whole body arrived. +type answer struct { + status int + header http.Header + body []byte + err error +} + +// readAnswer reads all of res, and closes its body. +func readAnswer(res *http.Response) answer { + body, err := io.ReadAll(res.Body) + _ = res.Body.Close() + + return answer{status: res.StatusCode, header: res.Header, body: body, err: err} +} + +// newRequest makes a request for path to smallwebwaf at addr. +func newRequest(t *testing.T, method, addr, path string, body io.Reader) *http.Request { + t.Helper() + + req, err := http.NewRequestWithContext(t.Context(), method, "http://"+addr+path, body) + if err != nil { + t.Fatalf("new request: %v", err) + } + + return req +} + +// do sends req and reads the answer. +func do(t *testing.T, req *http.Request) answer { + t.Helper() + + res, err := newClient(t).Do(req) + if err != nil { + t.Fatalf("%s %s: %v", req.Method, req.URL.Path, err) + } + + return readAnswer(res) +} + +// get sends a GET request for path to smallwebwaf at addr. +func get(t *testing.T, addr, path string) answer { + t.Helper() + + return do(t, newRequest(t, http.MethodGet, addr, path, http.NoBody)) +} + +// dial opens a connection to smallwebwaf at addr, for requests the HTTP +// client cannot make, such as one that stops sending halfway. +func dial(t *testing.T, addr string) net.Conn { + t.Helper() + + conn, err := (&net.Dialer{}).DialContext(t.Context(), "tcp", addr) + if err != nil { + t.Fatalf("dial %s: %v", addr, err) + } + + t.Cleanup(func() { + _ = conn.Close() + }) + + return conn +} + +// send writes text to conn. +func send(t *testing.T, conn net.Conn, text string) { + t.Helper() + + _, err := io.WriteString(conn, text) + if err != nil { + t.Fatalf("send: %v", err) + } +} + +// readResponse reads the answer to a request sent on conn. +func readResponse(t *testing.T, conn net.Conn) answer { + t.Helper() + + err := conn.SetReadDeadline(time.Now().Add(waitLimit)) + if err != nil { + t.Fatalf("set read deadline: %v", err) + } + + res, err := http.ReadResponse(bufio.NewReader(conn), nil) + if err != nil { + t.Fatalf("read response: %v", err) + } + + return readAnswer(res) +} + +// wantLine checks the request log line's status and action. +func wantLine(t *testing.T, line logLine, status int, action string) { + t.Helper() + + if line.Status != status || line.Action != action { + t.Errorf("log line has status %d and action %q, want %d and %q", + line.Status, line.Action, status, action) + } +} + +// wantStatus checks an answer's status. +func wantStatus(t *testing.T, got answer, status int) { + t.Helper() + + if got.status != status { + t.Errorf("status %d, want %d", got.status, status) + } +} + +// wantTimedOut checks that what began at start ended once shortTimeout +// had run out, and not much later. +func wantTimedOut(t *testing.T, start time.Time) { + t.Helper() + + took := time.Since(start) + if took < shortTimeout || took > shortTimeout+waitLimit/2 { + t.Errorf("took %s, want %s", took, shortTimeout) + } +} diff --git a/internal/proxy/ratelimits_test.go b/internal/proxy/ratelimits_test.go new file mode 100644 index 0000000..8d57142 --- /dev/null +++ b/internal/proxy/ratelimits_test.go @@ -0,0 +1,71 @@ +package proxy_test + +import ( + "net/http" + "sync/atomic" + "testing" + + "sneak.berlin/go/smallwebwaf/internal/requestlog" +) + +// minute is the window of SWWAF_RATE_LIMIT_PER_MINUTE, as a log line's +// limit_hit names it. +const minute = "minute" + +func TestRateLimitRefusesBeforeTheApp(t *testing.T) { + t.Parallel() + + var calls atomic.Int32 + + app := startApp(t, func(http.ResponseWriter, *http.Request) { + calls.Add(1) + }) + addr, out := startProxy(t, app.URL, map[string]string{ + trustedProxies: trustLocalhost, + rateLimitPerMinute: "1", + }) + + const otherClient = "203.0.113.10" + + // With a limit of one request a minute, a client's second request is + // refused, with 403 by default. A client is one IPv4 address, or one + // IPv6 /64; an IPv4 address in IPv6 form is that IPv4 address. + requests := []struct { + client string // as X-Forwarded-For names it + logged string // as the log line's client_ip names it + want int + }{ + {client, client, http.StatusOK}, + {client, client, http.StatusForbidden}, + {otherClient, otherClient, http.StatusOK}, + {"::ffff:" + otherClient, otherClient, http.StatusForbidden}, + {"2001:db8::1", "2001:db8::1", http.StatusOK}, + {"2001:db8::8000:0:0:1", "2001:db8::8000:0:0:1", http.StatusForbidden}, + {"2001:db8:0:1::1", "2001:db8:0:1::1", http.StatusOK}, + } + + for i, sent := range requests { + req := newRequest(t, http.MethodGet, addr, "/", http.NoBody) + req.Header.Set(forwardedFor, sent.client) + wantStatus(t, do(t, req), sent.want) + + line := out.requestLines(t, i+1)[i] + if line.ClientIP != sent.logged { + t.Errorf("log line has client_ip %q, want %q", line.ClientIP, sent.logged) + } + + if sent.want == http.StatusOK { + wantLine(t, line, http.StatusOK, requestlog.ActionForward) + } else { + wantLine(t, line, http.StatusForbidden, requestlog.ActionRateLimited) + + if line.LimitHit != minute { + t.Errorf("log line has limit_hit %q, want minute", line.LimitHit) + } + } + } + + if calls.Load() != 4 { + t.Errorf("the app was called %d times, want 4", calls.Load()) + } +} diff --git a/internal/proxy/request.go b/internal/proxy/request.go new file mode 100644 index 0000000..09dfef1 --- /dev/null +++ b/internal/proxy/request.go @@ -0,0 +1,656 @@ +package proxy + +import ( + "context" + "errors" + "net/http" + "net/http/httptrace" + "net/http/httputil" + "net/netip" + "os" + "strings" + "sync" + "sync/atomic" + "time" + + "sneak.berlin/go/smallwebwaf/internal/ratelimit" + "sneak.berlin/go/smallwebwaf/internal/requestlog" +) + +// flushAfterEachWrite has ReverseProxy pass on each part of the app's +// answer as soon as it arrives. +const flushAfterEachWrite time.Duration = -1 + +// refusal is smallwebwaf refusing a request, or refusing to go on with it: +// the status the client is answered if the response has not started yet, +// 0 to close the connection without an answer, the action the log line +// names, and the setting whose size or time limit the request passed, if +// that is why. +type refusal struct { + status int + action string + limit string +} + +// request is one request on its way through smallwebwaf, from the moment +// its headers have been read to its log line. +type request struct { + h *handler + in *http.Request + // rc sets the deadlines of the connection to the client. + rc *http.ResponseController + out *responseWriter + body *requestBody // nil for a request without a body + line requestlog.Line + + client netip.Addr + peer netip.Addr + peerTrusted bool + start time.Time + // checked is when the checks were done, and upstreamStart when the + // request was handed to the app. + checked time.Time + upstreamStart time.Time + // cancel ends the request to the app. + cancel context.CancelFunc + // refused is the first refusal, from whichever goroutine meets it. + refused atomic.Pointer[refusal] + // complete is true once the app's whole answer has been passed on. + complete bool + + // mu guards what follows. The timeouts run on goroutines of their + // own, and the transport starts and stops them, and notes the times + // below, from its own; once timersStopped is set, none of the timeouts + // acts any more. + mu sync.Mutex + timersStopped bool + clientRequestTimer *time.Timer + upstreamRequestTimer *time.Timer + upstreamResponseTimer *time.Timer + // connected is when there was a connection to the app, requestSent + // when the app had been sent the whole request, and answerStarted + // when the first byte of its answer arrived. + connected time.Time + requestSent time.Time + answerStarted time.Time +} + +// newRequest starts handling r: it notes the time, counts the request as +// under way, works out the client, and starts the log line with what is +// known of the request. +func (h *handler) newRequest(w http.ResponseWriter, r *http.Request) *request { + h.metrics.RequestStarted() + + start := time.Now() + peer := peerAddress(r) + trusted := h.config.TrustedProxies + peerTrusted := isInside(peer, trusted) + forwardedFor := r.Header.Values("X-Forwarded-For") + client := clientAddress(peer, forwardedFor, trusted) + + rq := &request{ + h: h, + in: r, + rc: http.NewResponseController(w), + out: &responseWriter{ResponseWriter: w}, + client: client, + peer: peer, + peerTrusted: peerTrusted, + start: start, + line: requestlog.Line{ + Time: requestlog.FormatTime(start), + Instance: h.config.InstanceName, + ClientIP: client.String(), + Method: r.Method, + Scheme: scheme(r, peerTrusted), + Host: r.Host, + Path: r.URL.EscapedPath(), + Query: r.URL.RawQuery, + Protocol: r.Proto, + Referer: r.Referer(), + UserAgent: r.UserAgent(), + RequestID: requestID(r, peerTrusted), + PeerIP: peer.String(), + ForwardedFor: strings.Join(forwardedFor, ", "), + ClientGroup: clientGroup(client).String(), + ContentType: r.Header.Get("Content-Type"), + RequestHeaders: requestHeaders(r, h.config.LogRequestHeaders), + HasAuthorization: len(r.Header.Values("Authorization")) > 0, + HasCookie: len(r.Header.Values("Cookie")) > 0, + Action: requestlog.ActionForward, + }, + } + + // A length of -1 is a body whose length was not announced. + if r.ContentLength > 0 { + rq.line.ContentLength = r.ContentLength + } + + if r.Body != http.NoBody { + rq.body = &requestBody{body: limitBody(r.Body, h.config.RequestMaxBytes), rq: rq} + } + + return rq +} + +// requestHeaders returns the headers of r that names lists, by name in +// lower case, each with its values joined by ", ". Authorization, Cookie +// and Set-Cookie are never among them, whatever names says. +func requestHeaders(r *http.Request, names []string) map[string]string { + headers := map[string]string{} + + for _, name := range names { + switch name { + case "authorization", "cookie", "set-cookie": + continue + } + + values := r.Header.Values(name) + if len(values) > 0 { + headers[name] = strings.Join(values, ", ") + } + } + + return headers +} + +// 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 checks of checkClient come +// first, answered with SWWAF_BAN_RESPONSE, and then the size limit, so +// that a request the rate limits count is counted even when it is +// refused for its size. In observe mode a request checkClient refuses +// goes on to the size limit like any other. ctx is the request's own +// context. +func (rq *request) check(ctx context.Context) *refusal { + action := rq.checkClient(ctx) + if action != "" { + if !rq.h.config.Observe { + return rq.banResponse(action) + } + + // The log line names what enforce mode would have done. + rq.line.WouldAction = action + } + + maxBytes := rq.h.config.RequestMaxBytes + if maxBytes > 0 && rq.in.ContentLength > maxBytes { + return &refusal{ + status: http.StatusRequestEntityTooLarge, + action: requestlog.ActionTooLarge, + limit: "SWWAF_REQUEST_MAX_BYTES", + } + } + + return nil +} + +// checkClient runs the checks on the request's client, and returns the +// action of the first that refuses the request, or "" when none does. A +// client in SWWAF_ALLOW_NETS skips them. For any other client, +// SWWAF_DENY_NETS comes first, then a ban on its netblock, so that a +// client either refuses is not looked up, and then the country lists; a +// request any of them refuses is not counted for the rate limits. Then +// come the rate limits, unless the client is in +// SWWAF_RATE_LIMIT_EXEMPT_NETS, so that every other request is counted. +// ctx is the request's own context. +func (rq *request) checkClient(ctx context.Context) string { + cfg := rq.h.config + if isInside(rq.client, cfg.AllowNets) { + return "" + } + + now := rq.h.now() + + if isInside(rq.client, cfg.DenyNets) { + return requestlog.ActionDenied + } + + if rq.banned(now) { + return requestlog.ActionBanned + } + + if rq.countryDenied(ctx) { + return requestlog.ActionCountryDenied + } + + if !isInside(rq.client, cfg.RateLimitExemptNets) && rq.limitBroken(now) { + return requestlog.ActionRateLimited + } + + return "" +} + +// forward passes the request to the app and the app's answer back. ctx +// is the request's own context. +func (rq *request) forward(ctx context.Context) { + ctx, cancel := context.WithCancel(ctx) + defer cancel() + + rq.cancel = cancel + ctx = httptrace.WithClientTrace(ctx, &httptrace.ClientTrace{ + GotConn: rq.gotConn, + WroteRequest: rq.wroteRequest, + GotFirstResponseByte: rq.gotFirstResponseByte, + }) + + out := rq.in.WithContext(ctx) + if rq.body != nil { + out.Body = rq.body + } + + reverseProxy := &httputil.ReverseProxy{ + Rewrite: rq.rewrite, + Transport: rq.h.transport, + FlushInterval: flushAfterEachWrite, + ErrorLog: rq.h.errorLog, + ModifyResponse: rq.modifyResponse, + ErrorHandler: rq.answerError, + } + + rq.startRequestTimers() + rq.upstreamStart = time.Now() + reverseProxy.ServeHTTP(rq.out, out) +} + +// rewrite makes the request the app receives: the client's request, +// unchanged, sent to SWWAF_UPSTREAM_URL, with the forwarded headers and +// the request's id set. +func (rq *request) rewrite(pr *httputil.ProxyRequest) { + upstream := rq.h.config.UpstreamURL + pr.Out.URL.Scheme = upstream.Scheme + pr.Out.URL.Host = upstream.Host + // ReverseProxy drops query parameters it cannot parse; the app gets + // the query as the client sent it. + pr.Out.URL.RawQuery = pr.In.URL.RawQuery + setForwardedHeaders(pr.In, pr.Out, rq.peer, rq.peerTrusted) + pr.Out.Header.Set(requestIDHeader, rq.line.RequestID) +} + +// modifyResponse looks at the app's answer before ReverseProxy passes it +// on. +func (rq *request) modifyResponse(res *http.Response) error { + rq.line.UpstreamStatus = res.StatusCode + + if res.StatusCode == http.StatusSwitchingProtocols { + // An upgraded connection, such as a WebSocket, is not cut by the + // timeouts. ReverseProxy writes this answer straight to the + // connection it takes over, not through rq.out. + rq.stopTimers() + rq.out.status = res.StatusCode + rq.line.Websocket = true + + return nil + } + + maxBytes := rq.h.config.ResponseMaxBytes + if maxBytes > 0 && res.Body != http.NoBody && res.ContentLength > maxBytes { + rq.refuse(refusal{ + status: http.StatusBadGateway, + action: requestlog.ActionTooLarge, + limit: "SWWAF_RESPONSE_MAX_BYTES", + }) + + return errResponseTooLarge + } + + res.Body = &responseBody{body: limitBody(res.Body, maxBytes), rq: rq} + rq.startClientResponseTimeout() + + return nil +} + +// answerError is ReverseProxy's ErrorHandler: the request could not be +// passed to the app, or the app's answer cannot be passed on. +func (rq *request) answerError(_ http.ResponseWriter, _ *http.Request, err error) { + refused := rq.refused.Load() + if refused == nil { + if rq.in.Context().Err() != nil { + return // the client has gone, and there is no one to answer + } + + rq.h.processLog.Warn("request to the app failed", "error", err.Error()) + + refused = &refusal{ + status: http.StatusBadGateway, + action: requestlog.ActionUpstreamError, + } + } + + rq.answer(*refused) +} + +// answer sends smallwebwaf's own answer, unless the response has already +// started, and records the refusal for the log line. +func (rq *request) answer(r refusal) { + rq.refused.CompareAndSwap(nil, &r) + + if rq.out.status != 0 { + return // too late to answer: the connection can only be cut + } + + if r.status == 0 { + // SWWAF_BAN_RESPONSE is close. This panic has Go's server close + // the connection without an answer, and log nothing; the log line + // is still written as the handler returns. + panic(http.ErrAbortHandler) + } + + // A client found too slow is read no more; any other may go on + // sending until its time is up, so that Go's server can read the + // rest of the body and end the request cleanly. + deadline := rq.clientRequestDeadline() + if r.status == http.StatusRequestTimeout { + deadline = time.Now() + } + + rq.stopReadingBody(deadline) + + timeout := rq.h.config.ClientResponseTimeout + if timeout > 0 { + _ = rq.rc.SetWriteDeadline(time.Now().Add(timeout)) + } + + http.Error(rq.out, http.StatusText(r.status), r.status) +} + +// refuse records r, unless an earlier refusal was, and ends the request +// to the app. +func (rq *request) refuse(r refusal) { + rq.refused.CompareAndSwap(nil, &r) + rq.cancel() +} + +// finish ends the request's timeouts, counts it in the metrics and writes +// its log line. +func (rq *request) finish() { + rq.stopTimers() + + refused := rq.refused.Load() + if refused == nil { + rq.stopReadingBody(rq.clientRequestDeadline()) + } + + line := &rq.line + line.Status = rq.out.status + line.ResponseBytes = rq.out.bytes + header := rq.out.Header() + line.ResponseContentType = header.Get("Content-Type") + line.CacheControl = header.Get("Cache-Control") + line.Location = header.Get("Location") + + if rq.body != nil { + line.RequestBytes = rq.body.bytes.Load() + } + + // limit is the setting whose size or time limit the request passed. + var limit string + + switch { + case refused != nil: + line.Action = refused.action + limit = refused.limit + case errors.Is(rq.out.err, os.ErrDeadlineExceeded): + // The client took longer than SWWAF_CLIENT_RESPONSE_TIMEOUT to + // take the response. + line.Action = requestlog.ActionTimedOut + limit = "SWWAF_CLIENT_RESPONSE_TIMEOUT" + case !rq.complete && (rq.out.err != nil || rq.in.Context().Err() != nil): + line.Aborted = true + } + + now := time.Now() + duration := now.Sub(rq.start) + line.DurationTotal = requestlog.Milliseconds(duration) + line.DurationChecks = timing(rq.start, rq.checked) + + var upstreamDuration time.Duration + + if !rq.upstreamStart.IsZero() { + upstreamDuration = now.Sub(rq.upstreamStart) + line.DurationUpstreamTotal = new(requestlog.Milliseconds(upstreamDuration)) + + rq.mu.Lock() + line.DurationUpstreamConnect = timing(rq.upstreamStart, rq.connected) + line.DurationUpstreamFirstByte = timing(rq.upstreamStart, rq.answerStarted) + rq.mu.Unlock() + } + + // Counted before the log line is written, so that the metrics count + // every request whose line is out. + rq.h.metrics.RequestEnded(line, limit, duration, upstreamDuration) + + err := requestlog.Write(rq.h.requestLog, line) + if err != nil { + rq.h.processLog.Error("writing the request log failed", "error", err.Error()) + } +} + +// timing is the time from start to end in milliseconds, for one of the +// log line's timings, or nil when end is zero: what it times never +// happened. +func timing(start, end time.Time) *float64 { + if end.IsZero() { + return nil + } + + return new(requestlog.Milliseconds(end.Sub(start))) +} + +// addToHistory adds the request, which has ended, to its client's +// history. +func (rq *request) addToHistory() { + var requestBytes int64 + if rq.body != nil { + requestBytes = rq.body.bytes.Load() + } + + forwarded := !rq.upstreamStart.IsZero() + + rq.h.limiter.AddToHistory(clientGroup(rq.client), rq.h.now(), ratelimit.Request{ + Country: rq.line.Country, + Forwarded: forwarded, + Refused: !forwarded && rq.refused.Load() != nil, + Status: rq.out.status, + RequestBytes: requestBytes, + ResponseBytes: rq.out.bytes, + BrokeLimit: rq.line.Offence == requestlog.OffenceLimit, + }) +} + +// clientRequestDeadline is when the client must have sent its whole +// request, or zero when SWWAF_CLIENT_REQUEST_TIMEOUT is off. +func (rq *request) clientRequestDeadline() time.Time { + timeout := rq.h.config.ClientRequestTimeout + if timeout == 0 { + return time.Time{} + } + + return rq.start.Add(timeout) +} + +// stopReadingBody ends, at deadline, the reading of a client body that has +// not arrived whole: Go's server then reads no more of it, and closes the +// connection after the answer. +func (rq *request) stopReadingBody(deadline time.Time) { + if rq.body == nil || rq.body.received.Load() { + return + } + + _ = rq.rc.SetReadDeadline(deadline) +} + +// startRequestTimers starts the timeouts that run while the request goes +// to the app: SWWAF_CLIENT_REQUEST_TIMEOUT until the client has sent its +// whole body, and SWWAF_UPSTREAM_REQUEST_TIMEOUT until the app has been +// sent the whole request. +func (rq *request) startRequestTimers() { + rq.mu.Lock() + defer rq.mu.Unlock() + + if rq.body != nil && rq.h.config.ClientRequestTimeout > 0 { + rq.clientRequestTimer = time.AfterFunc( + time.Until(rq.clientRequestDeadline()), func() { + rq.requestTimedOut("SWWAF_CLIENT_REQUEST_TIMEOUT") + }) + } + + timeout := rq.h.config.UpstreamRequestTimeout + if timeout > 0 { + rq.upstreamRequestTimer = time.AfterFunc(timeout, func() { + rq.requestTimedOut("SWWAF_UPSTREAM_REQUEST_TIMEOUT") + }) + } +} + +// requestTimedOut is called when limit, SWWAF_CLIENT_REQUEST_TIMEOUT or +// SWWAF_UPSTREAM_REQUEST_TIMEOUT, runs out while the request is still on +// its way to the app. The answer names the side smallwebwaf was waiting +// on at that moment: 408 when it was waiting for the client to send more +// of its body, 504 when it was waiting for the app to be reached or to +// take what it had. +func (rq *request) requestTimedOut(limit string) { + rq.mu.Lock() + defer rq.mu.Unlock() + + if rq.timersStopped { + return + } + + if rq.body == nil || !rq.body.waiting.Load() { + rq.refuse(refusal{ + status: http.StatusGatewayTimeout, + action: requestlog.ActionTimedOut, + limit: limit, + }) + + return + } + + rq.refuse(refusal{ + status: http.StatusRequestTimeout, + action: requestlog.ActionTimedOut, + limit: limit, + }) + // The transport gives up on the app only once its Read of the + // client's body returns, so that Read is ended now. The lock keeps + // this from reaching the connection after the request is handled. + _ = rq.rc.SetReadDeadline(time.Now()) +} + +// bodyReceived is called once the client has sent its whole body. +func (rq *request) bodyReceived() { + rq.mu.Lock() + defer rq.mu.Unlock() + + stopTimer(rq.clientRequestTimer) +} + +// gotConn is called once there is a connection to the app, a new one or +// one kept open from an earlier request. +func (rq *request) gotConn(httptrace.GotConnInfo) { + rq.mu.Lock() + defer rq.mu.Unlock() + + rq.connected = time.Now() +} + +// gotFirstResponseByte is called once the first byte of the app's answer +// has arrived. +func (rq *request) gotFirstResponseByte() { + rq.mu.Lock() + defer rq.mu.Unlock() + + rq.answerStarted = time.Now() +} + +// wroteRequest is called once the app has been sent the whole request: +// the request timeouts end and SWWAF_UPSTREAM_RESPONSE_TIMEOUT starts. +func (rq *request) wroteRequest(info httptrace.WroteRequestInfo) { + if info.Err != nil { + return // the transport gives up, or tries again + } + + rq.mu.Lock() + defer rq.mu.Unlock() + + if rq.timersStopped { + return + } + + stopTimer(rq.clientRequestTimer) + stopTimer(rq.upstreamRequestTimer) + rq.requestSent = time.Now() + + timeout := rq.h.config.UpstreamResponseTimeout + if timeout > 0 { + rq.upstreamResponseTimer = time.AfterFunc(timeout, rq.responseTimedOut) + } +} + +// responseTimedOut is called when SWWAF_UPSTREAM_RESPONSE_TIMEOUT runs out +// before the app has sent its whole answer. +func (rq *request) responseTimedOut() { + rq.mu.Lock() + defer rq.mu.Unlock() + + if !rq.timersStopped { + rq.refuse(refusal{ + status: http.StatusGatewayTimeout, + action: requestlog.ActionTimedOut, + limit: "SWWAF_UPSTREAM_RESPONSE_TIMEOUT", + }) + } +} + +// responseReceived is called once the app has sent its whole answer. +func (rq *request) responseReceived() { + rq.complete = true + rq.stopTimers() +} + +// startClientResponseTimeout sets SWWAF_CLIENT_RESPONSE_TIMEOUT on the +// connection to the client: the response must reach the client within it +// of the end of the request, or of now if the app answers before it has +// the whole request. +func (rq *request) startClientResponseTimeout() { + timeout := rq.h.config.ClientResponseTimeout + if timeout == 0 { + return + } + + from := rq.sentAt() + if from.IsZero() { + from = time.Now() + } + + _ = rq.rc.SetWriteDeadline(from.Add(timeout)) +} + +// sentAt is when the app had been sent the whole request, or zero. +func (rq *request) sentAt() time.Time { + rq.mu.Lock() + defer rq.mu.Unlock() + + return rq.requestSent +} + +// stopTimers stops the request's timeouts and keeps any from starting +// later: the app's answer is complete, the connection upgraded, or the +// request handled. +func (rq *request) stopTimers() { + rq.mu.Lock() + defer rq.mu.Unlock() + + rq.timersStopped = true + stopTimer(rq.clientRequestTimer) + stopTimer(rq.upstreamRequestTimer) + stopTimer(rq.upstreamResponseTimer) +} + +// stopTimer stops t, which is nil when its timeout is off. +func stopTimer(t *time.Timer) { + if t != nil { + t.Stop() + } +} diff --git a/internal/proxy/requestlog_test.go b/internal/proxy/requestlog_test.go new file mode 100644 index 0000000..98d2104 --- /dev/null +++ b/internal/proxy/requestlog_test.go @@ -0,0 +1,368 @@ +package proxy_test + +import ( + "io" + "maps" + "math" + "net/http" + "reflect" + "slices" + "strings" + "testing" + "time" + + "sneak.berlin/go/smallwebwaf/internal/proxy" + "sneak.berlin/go/smallwebwaf/internal/ratelimit" + "sneak.berlin/go/smallwebwaf/internal/requestlog" +) + +const ( + // requestIDHeader carries the request's id. + requestIDHeader = "X-Request-ID" + // instance is the SWWAF_INSTANCE_NAME a test sets. + instance = "fsn1app1/gitea" + // ipv6Client is a client on IPv6, and ipv6Group the netblock the rate + // limits count it as. + ipv6Client = "2001:db8::7" + ipv6Group = "2001:db8::/64" +) + +func TestLogLineHasEachFieldWhereItApplies(t *testing.T) { + t.Parallel() + + received := make(chan string, 2) // the request ids the app received + app := startApp(t, func(w http.ResponseWriter, r *http.Request) { + received <- r.Header.Get(requestIDHeader) + + _, _ = io.Copy(io.Discard, r.Body) + + if r.URL.Path != "/full" { + w.WriteHeader(http.StatusNoContent) + + return + } + + w.Header().Set("Content-Type", "text/html") + w.Header().Set("Cache-Control", "no-store") + w.Header().Set("Location", "/elsewhere") + w.WriteHeader(http.StatusFound) + _, _ = io.WriteString(w, "moved") + }) + addr, out := startProxy(t, app.URL, map[string]string{ + trustedProxies: trustLocalhost, + rateLimitExemptNets: localhost, + instanceName: instance, + logRequestHeaders: "Accept,x-custom,Authorization,cookie,SET-COOKIE", + }) + + // This request comes from ipv6Client through a trusted proxy, with a + // body and each header the log line looks at, and is answered with a + // redirect. + conn := dial(t, addr) + send(t, conn, "POST /full HTTP/1.1\r\nHost: "+appHost+"\r\n"+ + forwardedFor+": 198.51.100.7, "+ipv6Client+"\r\n"+ + forwardedProto+": "+secure+"\r\n"+requestIDHeader+": from-traefik\r\n"+ + "Content-Type: application/x-www-form-urlencoded\r\nContent-Length: 3\r\n"+ + "Accept: text/html\r\nX-Custom: one\r\nX-Custom: two\r\n"+ + "Authorization: Bearer secret-token\r\nCookie: session=secret-cookie\r\n"+ + "Set-Cookie: secret-set-cookie\r\n\r\na=b") + wantStatus(t, readResponse(t, conn), http.StatusFound) + + // A request's log line can come after its answer: each is waited for + // before the next request, so that the lines are in order. + full := out.requestLines(t, 1)[0] + + // This one comes from 127.0.0.1, which the rate limits do not count, + // with a body of 4 bytes whose length it does not announce, so that its + // request_bytes is not its content_length, and no header the log line + // looks at, and is answered with 204 and no header. + conn = dial(t, addr) + send(t, conn, "POST /bare HTTP/1.1\r\nHost: "+appHost+"\r\n"+ + "Transfer-Encoding: chunked\r\n\r\n4\r\nbody\r\n0\r\n\r\n") + wantStatus(t, readResponse(t, conn), http.StatusNoContent) + + bare := out.requestLines(t, 2)[1] + + wantFullLine(t, full) + wantBareLine(t, bare) + + for _, line := range []logLine{full, bare} { + got := <-received + if got != line.RequestID { + t.Errorf("the app received request id %q, the log line has %q", + got, line.RequestID) + } + } + + if strings.Contains(out.text(), "secret") { + t.Errorf("a value of Authorization, Cookie or Set-Cookie is logged:\n%s", + out.text()) + } +} + +// wantFullLine checks the log line of the request with every header the +// line looks at. Its timings are checked by TestTimingsAreInOrder. +func wantFullLine(t *testing.T, line logLine) { + t.Helper() + + headers := map[string]string{"accept": "text/html", "x-custom": "one, two"} + + want := withTimings(line, requestlog.Line{ + Type: requestType, Time: line.Time, Instance: instance, + ClientIP: ipv6Client, Method: http.MethodPost, Scheme: secure, + Host: appHost, Path: "/full", Protocol: protocol, + Status: http.StatusFound, RequestBytes: 3, ResponseBytes: 5, + RequestID: "from-traefik", PeerIP: localhost, + ForwardedFor: "198.51.100.7, " + ipv6Client, ClientGroup: ipv6Group, + ContentType: "application/x-www-form-urlencoded", ContentLength: 3, + RequestHeaders: headers, HasAuthorization: true, HasCookie: true, + ResponseContentType: "text/html", UpstreamStatus: http.StatusFound, + CacheControl: "no-store", Location: "/elsewhere", + Action: requestlog.ActionForward, + Counts: ratelimit.Counts{Minute: 1, Hour: 1, Day: 1}, + }) + if !reflect.DeepEqual(line.Line, want) { + t.Errorf("log line\n%+v\nwant\n%+v", line.Line, want) + } +} + +// wantBareLine checks the log line of the request with none of them, and +// that the fields that do not apply to it are left out. +func wantBareLine(t *testing.T, line logLine) { + t.Helper() + + want := withTimings(line, requestlog.Line{ + Type: requestType, Time: line.Time, Instance: instance, + ClientIP: localhost, Method: http.MethodPost, Scheme: plain, + Host: appHost, Path: "/bare", Protocol: protocol, + Status: http.StatusNoContent, RequestBytes: 4, RequestID: line.RequestID, + PeerIP: localhost, ClientGroup: localhost + "/32", + UpstreamStatus: http.StatusNoContent, Action: requestlog.ActionForward, + }) + if !reflect.DeepEqual(line.Line, want) || line.RequestID == "" { + t.Errorf("log line\n%+v\nwant\n%+v, with a request id", line.Line, want) + } + + for _, name := range []string{ + "forwarded_for", "content_type", "content_length", "request_headers", + "has_authorization", "has_cookie", "websocket", "response_content_type", + "cache_control", "location", "counts", + } { + _, present := line.fields[name] + if present { + t.Errorf("log line has %s, which does not apply", name) + } + } +} + +// withTimings returns want with the timings of line. +func withTimings(line logLine, want requestlog.Line) requestlog.Line { + want.DurationTotal = line.DurationTotal + want.DurationChecks = line.DurationChecks + want.DurationUpstreamConnect = line.DurationUpstreamConnect + want.DurationUpstreamFirstByte = line.DurationUpstreamFirstByte + want.DurationUpstreamTotal = line.DurationUpstreamTotal + + return want +} + +func TestHasAuthorizationAndHasCookieEachComeFromTheirOwnHeader(t *testing.T) { + t.Parallel() + + const hasAuthorization, hasCookie = "has_authorization", "has_cookie" + + for _, tc := range []struct{ header, field, other string }{ + {"Authorization", hasAuthorization, hasCookie}, + {"Cookie", hasCookie, hasAuthorization}, + } { + t.Run("only "+tc.header, func(t *testing.T) { + t.Parallel() + + app := startApp(t, func(http.ResponseWriter, *http.Request) {}) + addr, out := startProxy(t, app.URL, nil) + + req := newRequest(t, http.MethodGet, addr, "/", http.NoBody) + req.Header.Set(tc.header, "secret") + wantStatus(t, do(t, req), http.StatusOK) + + line := out.requestLine(t) + + _, otherPresent := line.fields[tc.other] + if line.fields[tc.field] != true || otherPresent { + t.Errorf("log line has %s %v and %s %v, want true and none", + tc.field, line.fields[tc.field], tc.other, line.fields[tc.other]) + } + }) + } +} + +func TestRequestIDAndSchemeComeOnlyFromATrustedProxy(t *testing.T) { + t.Parallel() + + const sentID = "from-traefik" + + sent := http.Header{requestIDHeader: {sentID}, forwardedProto: {secure}} + trusted := map[string]string{trustedProxies: trustLocalhost} + + for _, tc := range []struct { + name string + env map[string]string + header http.Header + // wantID is the request id logged, "" for a new one. + wantID, wantScheme string + }{ + {"a trusted proxy's are kept", trusted, sent, sentID, secure}, + {"without them, the id is new and the scheme http", trusted, nil, "", plain}, + {"another peer's are replaced", nil, sent, "", plain}, + } { + t.Run(tc.name, func(t *testing.T) { + t.Parallel() + + received := make(chan string, 2) + app := startApp(t, func(_ http.ResponseWriter, r *http.Request) { + received <- r.Header.Get(requestIDHeader) + }) + addr, out := startProxy(t, app.URL, tc.env) + + // Two requests, so that two new ids can be told apart. + ids := make([]string, 0, 2) + + for i := range 2 { + req := newRequest(t, http.MethodGet, addr, "/", http.NoBody) + maps.Copy(req.Header, tc.header) + wantStatus(t, do(t, req), http.StatusOK) + + line := out.requestLines(t, i+1)[i] + ids = append(ids, line.RequestID) + + got := <-received + if line.RequestID != got || line.Scheme != tc.wantScheme { + t.Errorf("log line has request_id %q and scheme %q, and the "+ + "app received id %q; want the same id and scheme %q", + line.RequestID, line.Scheme, got, tc.wantScheme) + } + } + + switch { + case tc.wantID != "" && (ids[0] != tc.wantID || ids[1] != tc.wantID): + t.Errorf("request ids %q, want %q", ids, tc.wantID) + case tc.wantID == "" && (slices.Contains(ids, sentID) || + slices.Contains(ids, "") || ids[0] == ids[1]): + t.Errorf("request ids %q, want two new ones", ids) + } + }) + } +} + +func TestTimingsAreInOrder(t *testing.T) { + t.Parallel() + + const denied = "192.0.2.50" // in SWWAF_DENY_NETS + + app := startApp(t, func(w http.ResponseWriter, _ *http.Request) { + // The pauses set the times apart; a hold-up of the test only + // lengthens them. + time.Sleep(time.Millisecond) + w.WriteHeader(http.StatusOK) + _ = http.NewResponseController(w).Flush() + + time.Sleep(time.Millisecond) + + _, _ = io.WriteString(w, "done") + }) + addr, out := startProxy(t, app.URL, map[string]string{ + trustedProxies: trustLocalhost, + denyNets: denied, + }) + + // Each log line is waited for before the next request, so that the + // lines are in order. + wantStatus(t, get(t, addr, "/"), http.StatusOK) + forwarded := out.requestLines(t, 1)[0] + + req := newRequest(t, http.MethodGet, addr, "/", http.NoBody) + req.Header.Set(forwardedFor, denied) + wantStatus(t, do(t, req), http.StatusForbidden) + refused := out.requestLines(t, 2)[1] + + wantStatus(t, get(t, addr, proxy.HealthPath), http.StatusOK) + health := out.requestLines(t, 3)[2] + + // A request passed to the app has every timing; one refused, none of + // the app's; the health check, which runs no check, only the total. + wantTimings(t, forwarded, "duration_total", "duration_checks", + "duration_upstream_connect", "duration_upstream_first_byte", + "duration_upstream_total") + wantTimings(t, refused, "duration_total", "duration_checks") + wantTimings(t, health, "duration_total") + + if t.Failed() { + return + } + + // In whole microseconds, as they are logged, so that the sum below is + // exact. + total := microseconds(forwarded.DurationTotal) + checks := microseconds(*forwarded.DurationChecks) + connect := microseconds(*forwarded.DurationUpstreamConnect) + firstByte := microseconds(*forwarded.DurationUpstreamFirstByte) + upstream := microseconds(*forwarded.DurationUpstreamTotal) + + // The checks end before the request is handed to the app, and the + // connection comes before the answer, which the app ends after a + // pause. + if checks+upstream > total || connect >= firstByte || firstByte >= upstream { + t.Errorf("timings in microseconds: total %d, checks %d, connect %d, "+ + "first byte %d, upstream total %d", total, checks, connect, firstByte, + upstream) + } + + if *refused.DurationChecks > refused.DurationTotal { + t.Errorf("refused request's checks took %v of %v milliseconds", + *refused.DurationChecks, refused.DurationTotal) + } +} + +// wantTimings checks that the timings named are the only ones line has. +func wantTimings(t *testing.T, line logLine, want ...string) { + t.Helper() + + var got []string + + for name := range line.fields { + if strings.HasPrefix(name, "duration_") { + got = append(got, name) + } + } + + slices.Sort(got) + slices.Sort(want) + + if !slices.Equal(got, want) { + t.Errorf("log line of %s has timings %v, want %v", line.Path, got, want) + } +} + +// microseconds is a timing in whole microseconds. +func microseconds(milliseconds float64) int64 { + return int64(math.Round(milliseconds * 1000)) +} + +func TestLogsAnUpgradedConnection(t *testing.T) { + t.Parallel() + + app := startApp(t, echoAfterUpgrade) + addr, out := startProxy(t, app.URL, nil) + + conn := dial(t, addr) + send(t, conn, "GET /socket HTTP/1.1\r\nHost: app\r\n"+ + "Connection: Upgrade\r\nUpgrade: websocket\r\n\r\n") + wantStatus(t, readResponse(t, conn), http.StatusSwitchingProtocols) + + _ = conn.Close() + + line := out.requestLine(t) + if line.fields["websocket"] != true { + t.Errorf("log line has websocket %v, want true", line.fields["websocket"]) + } +} diff --git a/internal/proxy/staticlists_test.go b/internal/proxy/staticlists_test.go new file mode 100644 index 0000000..d757f19 --- /dev/null +++ b/internal/proxy/staticlists_test.go @@ -0,0 +1,184 @@ +package proxy_test + +import ( + "net/http" + "strings" + "sync/atomic" + "testing" + + "sneak.berlin/go/smallwebwaf/internal/requestlog" +) + +// The rate limits count an IPv6 client by its /64, so these two addresses +// are one client for them. The static lists match each address on its own, +// and the tests list listedAddr alone. +const ( + listedAddr = "2001:db8::1" + unlistedAddr = "2001:db8::2" +) + +func TestAllowNetsSkipEveryCheckButTheSizeLimit(t *testing.T) { + t.Parallel() + + var calls atomic.Int32 + + app := startApp(t, func(http.ResponseWriter, *http.Request) { + calls.Add(1) + }) + geojsURL, asked := startGeoJS(t) + // fromKP is in SWWAF_ALLOW_NETS, and in SWWAF_DENY_NETS too, which + // comes after it. + addr, out := startProxyWithGeoJS(t, app.URL, geojsURL, map[string]string{ + trustedProxies: trustLocalhost, + allowNets: "198.51.100.0/24", + denyNets: fromKP, + deniedCountries: "kp", + rateLimitPerMinute: "1", + requestMaxBytes: "1K", + }) + + // Neither SWWAF_DENY_NETS, the country lists nor the limit of one + // request a minute refuses the client, and its country is not looked + // up. + wantAnswers(t, addr, out, []sentRequest{ + {fromKP, http.StatusOK, requestlog.ActionForward}, + {fromKP, http.StatusOK, requestlog.ActionForward}, + }) + + if len(asked()) != 0 { + t.Errorf("GeoJS was asked about %v, want nothing", asked()) + } + + // The size limit still applies. + body := strings.NewReader(strings.Repeat("a", 2<<10)) + req := newRequest(t, http.MethodPost, addr, "/", body) + req.Header.Set(forwardedFor, fromKP) + wantStatus(t, do(t, req), http.StatusRequestEntityTooLarge) + wantLine(t, out.requestLines(t, 3)[2], + http.StatusRequestEntityTooLarge, requestlog.ActionTooLarge) + + if calls.Load() != 2 { + t.Errorf("the app was called %d times, want 2", calls.Load()) + } +} + +func TestRequestFromAllowNetsIsNotCounted(t *testing.T) { + t.Parallel() + + app := startApp(t, func(http.ResponseWriter, *http.Request) {}) + addr, out := startProxy(t, app.URL, map[string]string{ + trustedProxies: trustLocalhost, + allowNets: listedAddr, + rateLimitPerMinute: "1", + }) + + // listedAddr's requests are not counted, so the first request from + // unlistedAddr is within the limit of one a minute. + wantAnswers(t, addr, out, []sentRequest{ + {listedAddr, http.StatusOK, requestlog.ActionForward}, + {listedAddr, http.StatusOK, requestlog.ActionForward}, + {unlistedAddr, http.StatusOK, requestlog.ActionForward}, + {unlistedAddr, http.StatusForbidden, requestlog.ActionRateLimited}, + }) +} + +func TestDenyNetsRefuseBeforeTheLookupAndTheBody(t *testing.T) { + t.Parallel() + + var calls atomic.Int32 + + app := startApp(t, func(http.ResponseWriter, *http.Request) { + calls.Add(1) + }) + geojsURL, asked := startGeoJS(t) + addr, out := startProxyWithGeoJS(t, app.URL, geojsURL, map[string]string{ + trustedProxies: trustLocalhost, + denyNets: "203.0.113.0/24", + deniedCountries: "kp", + }) + + req := newRequest(t, http.MethodPost, addr, "/", strings.NewReader("a body")) + req.Header.Set(forwardedFor, fromDE) + wantStatus(t, do(t, req), http.StatusForbidden) + + line := out.requestLine(t) + wantLine(t, line, http.StatusForbidden, requestlog.ActionDenied) + + if line.RequestBytes != 0 { + t.Errorf("log line has request_bytes %d, want 0", line.RequestBytes) + } + + if len(asked()) != 0 { + t.Errorf("GeoJS was asked about %v, want nothing", asked()) + } + + if calls.Load() != 0 { + t.Errorf("the app was called %d times, want none", calls.Load()) + } +} + +func TestRequestRefusedByDenyNetsIsNotCounted(t *testing.T) { + t.Parallel() + + app := startApp(t, func(http.ResponseWriter, *http.Request) {}) + addr, out := startProxy(t, app.URL, map[string]string{ + trustedProxies: trustLocalhost, + denyNets: listedAddr, + rateLimitPerMinute: "1", + }) + + // listedAddr's refused requests are not counted, so the first request + // from unlistedAddr is within the limit of one a minute. + wantAnswers(t, addr, out, []sentRequest{ + {listedAddr, http.StatusForbidden, requestlog.ActionDenied}, + {listedAddr, http.StatusForbidden, requestlog.ActionDenied}, + {unlistedAddr, http.StatusOK, requestlog.ActionForward}, + {unlistedAddr, http.StatusForbidden, requestlog.ActionRateLimited}, + }) +} + +func TestRateLimitExemptNetsAreNeitherCountedNorRefused(t *testing.T) { + t.Parallel() + + app := startApp(t, func(http.ResponseWriter, *http.Request) {}) + geojsURL, _ := startGeoJS(t) + addr, out := startProxyWithGeoJS(t, app.URL, geojsURL, map[string]string{ + trustedProxies: trustLocalhost, + rateLimitExemptNets: listedAddr + "," + fromKP, + deniedCountries: "kp", + rateLimitPerMinute: "1", + }) + + // listedAddr's requests are neither refused nor counted, so the first + // request from unlistedAddr is within the limit of one a minute. The + // country lists still refuse an exempt client. + wantAnswers(t, addr, out, []sentRequest{ + {listedAddr, http.StatusOK, requestlog.ActionForward}, + {listedAddr, http.StatusOK, requestlog.ActionForward}, + {unlistedAddr, http.StatusOK, requestlog.ActionForward}, + {unlistedAddr, http.StatusForbidden, requestlog.ActionRateLimited}, + {fromKP, http.StatusForbidden, requestlog.ActionCountryDenied}, + }) +} + +// sentRequest is a GET request from client, as X-Forwarded-For names it, +// and the status and log line action it should get. +type sentRequest struct { + client string + status int + action string +} + +// wantAnswers sends requests to smallwebwaf at addr one after another and +// checks each one's answer and log line. They must be the first requests +// smallwebwaf is sent, since the log lines are matched to them in order. +func wantAnswers(t *testing.T, addr string, out *output, requests []sentRequest) { + t.Helper() + + for i, sent := range requests { + req := newRequest(t, http.MethodGet, addr, "/", http.NoBody) + req.Header.Set(forwardedFor, sent.client) + wantStatus(t, do(t, req), sent.status) + wantLine(t, out.requestLines(t, i+1)[i], sent.status, sent.action) + } +} diff --git a/internal/proxy/timeouts_test.go b/internal/proxy/timeouts_test.go new file mode 100644 index 0000000..dab7791 --- /dev/null +++ b/internal/proxy/timeouts_test.go @@ -0,0 +1,307 @@ +package proxy_test + +import ( + "errors" + "io" + "net" + "net/http" + "net/http/httptest" + "strconv" + "sync" + "sync/atomic" + "testing" + "time" + + "sneak.berlin/go/smallwebwaf/internal/requestlog" +) + +// largeBodySize is more than the connections between the client, +// smallwebwaf and the app can hold while nobody reads, so that a sender +// soon waits. +const largeBodySize = 64 << 20 + +// writeSize is how much a test sender writes at a time. +const writeSize = 32 << 10 + +func TestRequestTimeouts(t *testing.T) { + t.Parallel() + + for _, tc := range []struct { + name string + // limit is the setting set to shortTimeout, which runs out; long + // is one set to longTimeoutSetting, which does not, or "". + limit, long string + // appTakesNothing has the app never read, while the client sends + // as fast as it can; otherwise the app reads, and the client + // stops sending halfway. + appTakesNothing bool + want int + }{ + { + name: "client request timeout, waiting on the client", + limit: clientRequestTimeout, + want: http.StatusRequestTimeout, + }, + { + name: "upstream request timeout, waiting on the client", + limit: upstreamRequestTimeout, + long: clientRequestTimeout, + want: http.StatusRequestTimeout, + }, + { + name: "upstream request timeout, waiting on the app", + limit: upstreamRequestTimeout, + appTakesNothing: true, + want: http.StatusGatewayTimeout, + }, + { + name: "client request timeout, waiting on the app", + limit: clientRequestTimeout, + long: upstreamRequestTimeout, + appTakesNothing: true, + want: http.StatusGatewayTimeout, + }, + } { + t.Run(tc.name, func(t *testing.T) { + t.Parallel() + + var ( + app *httptest.Server + appURL string + sendRequest func(*testing.T, string) net.Conn + appGotBody atomic.Bool + ) + + if tc.appTakesNothing { + appURL, sendRequest = startAppThatTakesNothing(t), sendLargeBody + } else { + app = startApp(t, func(_ http.ResponseWriter, r *http.Request) { + n, _ := io.Copy(io.Discard, r.Body) + appGotBody.Store(n > 0) + }) + appURL, sendRequest = app.URL, sendPartOfBody + } + + env := map[string]string{tc.limit: shortTimeoutSetting, metricsToken: token} + if tc.long != "" { + env[tc.long] = longTimeoutSetting + } + + addr, out := startProxy(t, appURL, env) + start := time.Now() + got := readResponse(t, sendRequest(t, addr)) + wantTimedOut(t, start) + + want := tc.want + + if app != nil { + // Close returns once the app has finished with the request. + app.Close() + + // Until some of the body has reached the app, smallwebwaf + // waits on the app, and SPEC.md asks for 504; the timeout + // runs out then only if the test process is held up. + if !appGotBody.Load() { + want = http.StatusGatewayTimeout + } + } + + wantStatus(t, got, want) + wantLine(t, out.requestLine(t), want, requestlog.ActionTimedOut) + wantLimitHits(t, addr, tc.limit, 1) + }) + } +} + +// startAppThatTakesNothing starts an app that accepts connections and +// never reads from them, and returns its URL. +func startAppThatTakesNothing(t *testing.T) string { + t.Helper() + + listener, err := (&net.ListenConfig{}).Listen(t.Context(), "tcp", localhost+":0") + if err != nil { + t.Fatalf("listen: %v", err) + } + + var ( + mu sync.Mutex + held []net.Conn + ) + + hold := func(conn net.Conn) { + mu.Lock() + defer mu.Unlock() + + held = append(held, conn) + } + + go func() { + for { + conn, err := listener.Accept() + if err != nil { + return + } + + hold(conn) + } + }() + + t.Cleanup(func() { + _ = listener.Close() + + mu.Lock() + defer mu.Unlock() + + for _, conn := range held { + _ = conn.Close() + } + }) + + return "http://" + listener.Addr().String() +} + +// sendPartOfBody sends a request that announces a large body, and only +// the first bytes of it. +func sendPartOfBody(t *testing.T, addr string) net.Conn { + t.Helper() + + conn := dial(t, addr) + send(t, conn, "POST /upload HTTP/1.1\r\nHost: app\r\nContent-Length: "+ + strconv.Itoa(largeBodySize)+"\r\n\r\nthe first bytes") + + return conn +} + +// sendLargeBody sends a request with a large body, as fast as smallwebwaf +// takes it, from a goroutine of its own. +func sendLargeBody(t *testing.T, addr string) net.Conn { + t.Helper() + + conn := dial(t, addr) + send(t, conn, "POST /upload HTTP/1.1\r\nHost: app\r\nContent-Length: "+ + strconv.Itoa(largeBodySize)+"\r\n\r\n") + + go func() { + chunk := make([]byte, writeSize) + for range largeBodySize / writeSize { + _, err := conn.Write(chunk) + if err != nil { + return + } + } + }() + + return conn +} + +func TestAppTooSlowToAnswer(t *testing.T) { + t.Parallel() + + app := startApp(t, func(_ http.ResponseWriter, r *http.Request) { + <-r.Context().Done() + }) + addr, out := startProxy(t, app.URL, map[string]string{ + upstreamResponseTimeout: shortTimeoutSetting, + metricsToken: token, + }) + + start := time.Now() + + wantStatus(t, get(t, addr, "/slow"), http.StatusGatewayTimeout) + wantTimedOut(t, start) + + line := out.requestLine(t) + wantLine(t, line, http.StatusGatewayTimeout, requestlog.ActionTimedOut) + + _, answered := line.fields["upstream_status"] + if answered { + t.Errorf("log line has upstream_status %v for an app that never answered", + line.fields["upstream_status"]) + } + + wantLimitHits(t, addr, upstreamResponseTimeout, 1) +} + +func TestAppTooSlowToFinishItsAnswer(t *testing.T) { + t.Parallel() + + app := startApp(t, func(w http.ResponseWriter, r *http.Request) { + _, _ = io.WriteString(w, "the first part") + _ = http.NewResponseController(w).Flush() + + <-r.Context().Done() + }) + addr, out := startProxy(t, app.URL, map[string]string{ + upstreamResponseTimeout: shortTimeoutSetting, + }) + + start := time.Now() + got := get(t, addr, "/slow") + wantStatus(t, got, http.StatusOK) + + if string(got.body) != "the first part" || !errors.Is(got.err, io.ErrUnexpectedEOF) { + t.Errorf("client read %q (%v), want the first part cut off", got.body, got.err) + } + + wantTimedOut(t, start) + + line := out.requestLine(t) + wantLine(t, line, http.StatusOK, requestlog.ActionTimedOut) + + if line.UpstreamStatus != http.StatusOK { + t.Errorf("log line has upstream_status %d, want %d", + line.UpstreamStatus, http.StatusOK) + } +} + +func TestClientTooSlowToTakeTheAnswer(t *testing.T) { + t.Parallel() + + app := startApp(t, func(w http.ResponseWriter, _ *http.Request) { + chunk := make([]byte, writeSize) + for range largeBodySize / writeSize { + _, err := w.Write(chunk) + if err != nil { + return + } + } + }) + addr, out := startProxy(t, app.URL, map[string]string{ + clientResponseTimeout: shortTimeoutSetting, + metricsToken: token, + }) + + start := time.Now() + + // The client asks, and never reads the answer. + conn := dial(t, addr) + send(t, conn, "GET /large HTTP/1.1\r\nHost: app\r\n\r\n") + + line := out.requestLine(t) + wantTimedOut(t, start) + wantLine(t, line, http.StatusOK, requestlog.ActionTimedOut) + wantLimitHits(t, addr, clientResponseTimeout, 1) +} + +func TestClosesAnIdleConnection(t *testing.T) { + t.Parallel() + + app := startApp(t, func(http.ResponseWriter, *http.Request) {}) + addr, _ := startProxy(t, app.URL, map[string]string{ + clientIdleTimeout: shortTimeoutSetting, + }) + + // The idle time starts once the answer is sent, so after start. + start := time.Now() + conn := dial(t, addr) + send(t, conn, "GET / HTTP/1.1\r\nHost: app\r\n\r\n") + wantStatus(t, readResponse(t, conn), http.StatusOK) + + // The read deadline readResponse set still bounds this read. + _, err := conn.Read(make([]byte, 1)) + if !errors.Is(err, io.EOF) { + t.Fatalf("read on the idle connection: %v, want it closed", err) + } + + wantTimedOut(t, start) +} diff --git a/internal/ratelimit/history_test.go b/internal/ratelimit/history_test.go new file mode 100644 index 0000000..55ff667 --- /dev/null +++ b/internal/ratelimit/history_test.go @@ -0,0 +1,119 @@ +package ratelimit_test + +import ( + "net/netip" + "testing" + "time" + + "sneak.berlin/go/smallwebwaf/internal/ratelimit" +) + +func TestHistoryKeepsEveryRequest(t *testing.T) { + t.Parallel() + + limiter := ratelimit.New(ratelimit.Limits{}) + client := netip.MustParsePrefix("203.0.113.9/32") + start := midnight() + + for i, r := range []ratelimit.Request{ + {Country: "DE", Forwarded: true, Status: 200, RequestBytes: 10, ResponseBytes: 100}, + {Forwarded: true, Status: 101}, + {Forwarded: true, Status: 304, RequestBytes: 5}, + {Country: "FR", Refused: true, Status: 403, ResponseBytes: 10, BrokeLimit: true}, + {Forwarded: true, Status: 502, ResponseBytes: 12}, + // Closed without an answer: refused, and no response. + {Refused: true, Status: 0}, + // Answered 404 at smallwebwaf's own endpoints: neither forwarded + // nor refused. + {Status: 404}, + } { + limiter.AddToHistory(client, start.Add(time.Duration(i)*time.Minute), r) + } + + want := ratelimit.History{ + FirstSeen: start, + LastSeen: start.Add(6 * time.Minute), + Country: "FR", + LookedUp: start.Add(3 * time.Minute), + Requests: 7, + Forwarded: 4, + Refused: 2, + RequestBytes: 15, + ResponseBytes: 122, + Responses: ratelimit.Responses{ + Status1xx: 1, Status2xx: 1, Status3xx: 1, Status4xx: 2, Status5xx: 1, + }, + Offences: ratelimit.Offences{Limit: 1}, + } + + got := historyOf(t, limiter, client) + if got != want { + t.Errorf("history\n%+v\nwant\n%+v", got, want) + } +} + +func TestResetKeepsTheHistory(t *testing.T) { + t.Parallel() + + limiter := ratelimit.New(ratelimit.Limits{PerMinute: limit}) + client := netip.MustParsePrefix("203.0.113.9/32") + start := midnight() + + for range limit { + wantCount(t, limiter, client, start, "") + limiter.AddToHistory(client, start, ratelimit.Request{Forwarded: true}) + } + + limiter.Reset(client) + + if got := historyOf(t, limiter, client).Requests; got != limit { + t.Errorf("the history counts %d requests, want %d", got, limit) + } +} + +func TestRequestsAddsUpTheClientsInsideTheNetblock(t *testing.T) { + t.Parallel() + + limiter := ratelimit.New(ratelimit.Limits{}) + + for client, requests := range map[string]int{ + "198.51.100.9/32": 2, + "198.51.100.10/32": 3, + "192.0.2.1/32": 5, + "2001:db8:5::/64": 7, + } { + for range requests { + limiter.AddToHistory(netip.MustParsePrefix(client), midnight(), + ratelimit.Request{}) + } + } + + for netblock, want := range map[string]int64{ + "198.51.100.9/32": 2, + "198.51.100.0/24": 5, + "2001:db8:5::/64": 7, + "203.0.113.0/24": 0, + } { + got := limiter.Requests(netip.MustParsePrefix(netblock)) + if got != want { + t.Errorf("%s has sent %d requests, want %d", netblock, got, want) + } + } +} + +// historyOf returns client's history. +func historyOf( + t *testing.T, limiter *ratelimit.Limiter, client netip.Prefix, +) ratelimit.History { + t.Helper() + + for _, c := range limiter.Snapshot() { + if c.Client == client { + return c.History + } + } + + t.Fatalf("%s is not in the table", client) + + return ratelimit.History{} +} diff --git a/internal/ratelimit/ratelimit.go b/internal/ratelimit/ratelimit.go new file mode 100644 index 0000000..9a7b68b --- /dev/null +++ b/internal/ratelimit/ratelimit.go @@ -0,0 +1,389 @@ +// Package ratelimit keeps the table of clients: each client's requests +// counted over a minute, an hour and a day, as the "Counting method" +// section of SPEC.md describes, which tell when a request takes the client +// over a rate limit, and each client's history since it was first seen. +// At most 20,000 clients are kept, in memory, and written to clients.json +// and read from it by the state package. +package ratelimit + +import ( + "net/http" + "net/netip" + "slices" + "sync" + "time" + + "github.com/hashicorp/golang-lru/v2/simplelru" +) + +// maxClients is how many clients are kept. Past it, the least recently +// seen client is dropped, with its history, and starts afresh if it comes +// back. +const maxClients = 20000 + +const day = 24 * time.Hour + +// Limits are the most requests a client may make in a minute, an hour and +// a day. Zero is no limit. +type Limits struct { + PerMinute int64 + PerHour int64 + PerDay int64 +} + +// Limiter counts each client's requests against the limits, and keeps +// its history. It is safe for concurrent use. +type Limiter struct { + // windows are the minute, the hour and the day, in the order of + // Client.buckets. + windows [3]window + + mu sync.Mutex + clients *simplelru.LRU[netip.Prefix, *Client] +} + +// Client is a client in the table, as clients.json holds it: its buckets +// in each window, and its history. +type Client struct { + Client netip.Prefix `json:"client"` + Minute Buckets `json:"minute"` + Hour Buckets `json:"hour"` + Day Buckets `json:"day"` + History History `json:"history"` +} + +// Buckets are a client's two buckets in one window: the requests in the +// bucket under way, which began at Start, and in the bucket before it. +type Buckets struct { + Start time.Time `json:"start"` + Current int64 `json:"current"` + Previous int64 `json:"previous"` +} + +// History is what is known of a client since it was first seen. +// +//nolint:tagliatelle // the state files use snake_case, as the request log does +type History struct { + FirstSeen time.Time `json:"first_seen"` + LastSeen time.Time `json:"last_seen"` + // Country is the client's country as it was last looked up, and + // LookedUp when that was; both are empty while it never was. + Country string `json:"country,omitempty"` + LookedUp time.Time `json:"looked_up,omitzero"` + // Requests are all the client's requests: Forwarded those passed to + // the app, Refused those refused before anything reached it, a 401 at + // smallwebwaf's own endpoints included, and neither the others + // smallwebwaf answered there. + Requests int64 `json:"requests"` + Forwarded int64 `json:"forwarded"` + Refused int64 `json:"refused"` + // RequestBytes and ResponseBytes are the body bytes of its requests + // and of the responses it was sent. + RequestBytes int64 `json:"request_bytes"` + ResponseBytes int64 `json:"response_bytes"` + Responses Responses `json:"responses,omitzero"` + Offences Offences `json:"offences,omitzero"` +} + +// Responses are the responses a client was sent, by status class; +// Status5xx counts every status from 500 up. +type Responses struct { + Status1xx int64 `json:"1xx,omitempty"` + Status2xx int64 `json:"2xx,omitempty"` + Status3xx int64 `json:"3xx,omitempty"` + Status4xx int64 `json:"4xx,omitempty"` + Status5xx int64 `json:"5xx,omitempty"` +} + +// Offences are a client's offences, by kind. +type Offences struct { + // Limit is its requests that broke a rate limit. + Limit int64 `json:"limit"` +} + +// Request is what a client's history keeps of one of its requests. +type Request struct { + // Country is the client's country, when the request looked it up. + Country string + // Forwarded is true for a request passed to the app, Refused for one + // refused before anything reached it, a 401 at smallwebwaf's own + // endpoints included. Both are false for any other request smallwebwaf + // answered there. + Forwarded bool + Refused bool + // Status is what the client was sent, 0 if nothing was. + Status int + // RequestBytes and ResponseBytes are the body bytes of the request + // and of its response. + RequestBytes int64 + ResponseBytes int64 + // BrokeLimit is true for a request that broke a rate limit. + BrokeLimit bool +} + +// New returns a Limiter for limits, with no client counted yet. +func New(limits Limits) *Limiter { + clients, err := simplelru.NewLRU[netip.Prefix, *Client](maxClients, nil) + if err != nil { + panic(err) // NewLRU fails only for a size below one + } + + return &Limiter{ + windows: [3]window{ + {name: "minute", length: time.Minute, limit: limits.PerMinute}, + {name: "hour", length: time.Hour, limit: limits.PerHour}, + {name: "day", length: day, limit: limits.PerDay}, + }, + clients: clients, + } +} + +// Hit is a request that takes a client over a rate limit. +type Hit struct { + // Window is "minute", "hour" or "day". + Window string + // Limit is the window's limit. + Limit int64 + // Requests is the client's requests counted in the window, this one + // included. + Requests float64 +} + +// Counts are a client's requests in the minute, the hour and the day that +// end at a request, that request included. +type Counts struct { + Minute float64 `json:"minute"` + Hour float64 `json:"hour"` + Day float64 `json:"day"` +} + +// Count counts a request from client at now, in every window, whether or +// not it is refused, and returns the client's requests in each window. It +// reports whether the request takes the client over a limit, and the +// window whose limit it goes over, the shortest if it is over several. +func (l *Limiter) Count(client netip.Prefix, now time.Time) (Counts, Hit, bool) { + l.mu.Lock() + defer l.mu.Unlock() + + var ( + requests [3]float64 + hit Hit + ) + + for i, b := range l.get(client).buckets() { + w := l.windows[i] + + requests[i] = b.add(now, w.length) + if hit.Window == "" && w.limit > 0 && requests[i] > float64(w.limit) { + hit = Hit{Window: w.name, Limit: w.limit, Requests: requests[i]} + } + } + + counts := Counts{Minute: requests[0], Hour: requests[1], Day: requests[2]} + + return counts, hit, hit.Window != "" +} + +// Reset sets client's counts in every window back to zero. Its history +// keeps its totals. +func (l *Limiter) Reset(client netip.Prefix) { + l.mu.Lock() + defer l.mu.Unlock() + + c, seen := l.clients.Peek(client) + if seen { + c.Minute, c.Hour, c.Day = Buckets{}, Buckets{}, Buckets{} + } +} + +// AddToHistory adds r, a request from client at now, to the client's +// history. +func (l *Limiter) AddToHistory(client netip.Prefix, now time.Time, r Request) { + l.mu.Lock() + defer l.mu.Unlock() + + h := &l.get(client).History + if h.FirstSeen.IsZero() { + h.FirstSeen = now + } + + h.LastSeen = now + + if r.Country != "" { + h.Country = r.Country + h.LookedUp = now + } + + h.Requests++ + if r.Forwarded { + h.Forwarded++ + } + + if r.Refused { + h.Refused++ + } + + h.RequestBytes += r.RequestBytes + h.ResponseBytes += r.ResponseBytes + h.Responses.add(r.Status) + + if r.BrokeLimit { + h.Offences.Limit++ + } +} + +// Requests returns how many requests the clients inside netblock have +// sent, as their histories count them. +func (l *Limiter) Requests(netblock netip.Prefix) int64 { + l.mu.Lock() + defer l.mu.Unlock() + + // Most often the netblock is one client. + c, seen := l.clients.Peek(netblock) + if seen { + return c.History.Requests + } + + var requests int64 + + for _, c := range l.clients.Values() { + if netblock.Overlaps(c.Client) { + requests += c.History.Requests + } + } + + return requests +} + +// Len returns how many clients are in the table. +func (l *Limiter) Len() int { + l.mu.Lock() + defer l.mu.Unlock() + + return l.clients.Len() +} + +// Snapshot returns every client in the table, sorted by address, as +// clients.json lists them. +func (l *Limiter) Snapshot() []Client { + l.mu.Lock() + + clients := make([]Client, 0, l.clients.Len()) + for _, c := range l.clients.Values() { + clients = append(clients, *c) + } + + l.mu.Unlock() + + slices.SortFunc(clients, func(a, b Client) int { + return a.Client.Compare(b.Client) + }) + + return clients +} + +// Load puts clients read from clients.json into the table, in place of +// the clients it holds, in the order they were last seen, so that the +// least recently seen is dropped first. Buckets whose time has passed at +// now are emptied. +func (l *Limiter) Load(clients []Client, now time.Time) { + clients = slices.Clone(clients) + slices.SortStableFunc(clients, func(a, b Client) int { + return a.History.LastSeen.Compare(b.History.LastSeen) + }) + + l.mu.Lock() + defer l.mu.Unlock() + + l.clients.Purge() + + for _, c := range clients { + for i, b := range c.buckets() { + // The window that ends at now covers neither bucket once it + // begins after the bucket under way has ended. + length := l.windows[i].length + if !now.Add(-length).Before(b.Start.Add(length)) { + *b = Buckets{} + } + } + + l.clients.Add(c.Client, &c) + } +} + +// get returns client's entry in the table, a new one if it has none, and +// makes it the most recently seen. +func (l *Limiter) get(client netip.Prefix) *Client { + c, seen := l.clients.Get(client) + if !seen { + c = &Client{Client: client} + l.clients.Add(client, c) + } + + return c +} + +// buckets returns c's buckets in the minute, the hour and the day. +func (c *Client) buckets() [3]*Buckets { + return [3]*Buckets{&c.Minute, &c.Hour, &c.Day} +} + +// window is a length of time over which requests are counted, and the +// most requests a client may make in it. +type window struct { + name string + length time.Duration + limit int64 +} + +// add counts a request at now in a window of length, and returns the +// client's requests in the window that ends at now: those in the bucket +// under way, and those in the bucket before it weighted by how much of +// that bucket the window still covers. +// +// Concurrent requests can be counted out of order, so now can be a moment +// before the bucket under way began; such a request is counted in that +// bucket. A request dated more than a second before it means the clock +// was set back, and the buckets start afresh: otherwise the bucket before +// would keep its full weight until the clock caught up. +func (b *Buckets) add(now time.Time, length time.Duration) float64 { + if now.Before(b.Start.Add(-time.Second)) { + *b = Buckets{} + } + + start := now.Truncate(length) + if start.After(b.Start) { + if start.Equal(b.Start.Add(length)) { + b.Previous = b.Current + } else { + b.Previous = 0 + } + + b.Start = start + b.Current = 0 + } + + b.Current++ + + elapsed := max(now.Sub(b.Start), 0) + covered := 1 - float64(elapsed)/float64(length) + + return float64(b.Previous)*covered + float64(b.Current) +} + +// add counts a response with status in its class. A status of 0, for +// nothing sent, is not a response. +func (r *Responses) add(status int) { + switch { + case status >= http.StatusInternalServerError: + r.Status5xx++ + case status >= http.StatusBadRequest: + r.Status4xx++ + case status >= http.StatusMultipleChoices: + r.Status3xx++ + case status >= http.StatusOK: + r.Status2xx++ + case status >= http.StatusContinue: + r.Status1xx++ + } +} diff --git a/internal/ratelimit/ratelimit_test.go b/internal/ratelimit/ratelimit_test.go new file mode 100644 index 0000000..49bb3ec --- /dev/null +++ b/internal/ratelimit/ratelimit_test.go @@ -0,0 +1,269 @@ +package ratelimit_test + +import ( + "net/netip" + "testing" + "time" + + "sneak.berlin/go/smallwebwaf/internal/ratelimit" +) + +// limit is the limit the tests set. +const limit = 3 + +// The windows, as Count names them. +const ( + minute = "minute" + hour = "hour" +) + +func TestEachWindowRefusesAtItsLimitAndLetsTheClientBack(t *testing.T) { + t.Parallel() + + for _, tc := range []struct { + window string + limits ratelimit.Limits + length time.Duration + }{ + {minute, ratelimit.Limits{PerMinute: limit}, time.Minute}, + {hour, ratelimit.Limits{PerHour: limit}, time.Hour}, + {"day", ratelimit.Limits{PerDay: limit}, 24 * time.Hour}, + } { + t.Run(tc.window, func(t *testing.T) { + t.Parallel() + + limiter := ratelimit.New(tc.limits) + client := netip.MustParsePrefix("203.0.113.9/32") + start := midnight() + quarter := tc.length / 4 + + for range limit { + wantCount(t, limiter, client, start, "") + } + + wantCount(t, limiter, client, start, tc.window) + + // A quarter into the next bucket, the window still covers three + // quarters of the bucket before, with its four requests: 3 + 1 + // is over the limit. + wantCount(t, limiter, client, start.Add(tc.length+quarter), tc.window) + + // Three quarters into it, a quarter: 1 + 2 is within. + wantCount(t, limiter, client, start.Add(tc.length+3*quarter), "") + }) + } +} + +func TestHitGivesTheLimitAndTheRequestsCounted(t *testing.T) { + t.Parallel() + + limiter := ratelimit.New(ratelimit.Limits{PerMinute: limit, PerHour: limit}) + client := netip.MustParsePrefix("203.0.113.9/32") + start := midnight() + + for range limit { + _, _, over := limiter.Count(client, start) + if over { + t.Fatal("a request within the limit is over it") + } + } + + // Over both limits; the minute's is named, with the four requests. + _, hit, over := limiter.Count(client, start) + + want := ratelimit.Hit{Window: minute, Limit: limit, Requests: limit + 1} + if !over || hit != want { + t.Errorf("request over the limit gives %+v and %t, want %+v and true", + hit, over, want) + } +} + +func TestCountGivesTheRequestsInEachWindow(t *testing.T) { + t.Parallel() + + limiter := ratelimit.New(ratelimit.Limits{}) + client := netip.MustParsePrefix("203.0.113.9/32") + start := midnight() + + for range 3 { + limiter.Count(client, start) + } + + // A quarter into the next hour, the minute has only this request. The + // hour still covers three quarters of the bucket before, with its three + // requests, which count 2.25, and this one: 3.25. The day covers all + // four. + counts, _, _ := limiter.Count(client, start.Add(time.Hour+time.Hour/4)) + + want := ratelimit.Counts{Minute: 1, Hour: 3.25, Day: 4} + if counts != want { + t.Errorf("counts %+v, want %+v", counts, want) + } +} + +func TestResetSetsTheCountsBackToZero(t *testing.T) { + t.Parallel() + + limiter := ratelimit.New(ratelimit.Limits{PerMinute: limit, PerDay: limit}) + client := netip.MustParsePrefix("203.0.113.9/32") + start := midnight() + + for range limit { + wantCount(t, limiter, client, start, "") + } + + wantCount(t, limiter, client, start, minute) + limiter.Reset(client) + + // At the same moment, the client has its whole allowance again. + for range limit { + wantCount(t, limiter, client, start, "") + } + + wantCount(t, limiter, client, start, minute) +} + +func TestClientBackAfterAWholeBucketIsWithinTheLimitAtOnce(t *testing.T) { + t.Parallel() + + limiter := ratelimit.New(ratelimit.Limits{PerHour: limit}) + client := netip.MustParsePrefix("203.0.113.9/32") + start := midnight() + + for range limit { + wantCount(t, limiter, client, start, "") + } + + wantCount(t, limiter, client, start, hour) + + // No request in the whole next bucket, so a quarter into the one after + // it the window covers none of the four requests: 1 is within the + // limit. Were they counted as the bucket before, 3 + 1 would be over. + wantCount(t, limiter, client, start.Add(2*time.Hour+time.Hour/4), "") +} + +func TestRefusedRequestsCount(t *testing.T) { + t.Parallel() + + limiter := ratelimit.New(ratelimit.Limits{PerMinute: limit, PerHour: 2 * limit}) + refused := netip.MustParsePrefix("203.0.113.9/32") + within := netip.MustParsePrefix("203.0.113.10/32") + start := midnight() + + for range limit { + wantCount(t, limiter, refused, start, "") + wantCount(t, limiter, within, start, "") + } + + for range limit { + wantCount(t, limiter, refused, start, minute) + } + + // Half a minute into the next bucket the window covers half of the + // bucket before: 3 + 1 is over the minute's limit for the client + // whose three refused requests count, and 1.5 + 1 within it for the + // other. The first is over the hour's limit too, and the shorter + // window is named. + halfway := start.Add(time.Minute + time.Minute/2) + wantCount(t, limiter, refused, halfway, minute) + wantCount(t, limiter, within, halfway, "") + + // The refused requests count in the hour as well: 6 + 1 + 1 is over + // its limit, and 3 + 1 + 1 within it. + later := start.Add(10 * time.Minute) + wantCount(t, limiter, refused, later, hour) + wantCount(t, limiter, within, later, "") +} + +func TestRequestCountedLateGoesInTheBucketUnderWay(t *testing.T) { + t.Parallel() + + limiter := ratelimit.New(ratelimit.Limits{PerMinute: limit}) + client := netip.MustParsePrefix("203.0.113.9/32") + start := midnight() + + for range limit { + wantCount(t, limiter, client, start, "") + } + + // A concurrent request dated a moment before the bucket under way, but + // counted after it began, is counted in it: 3 + 1 is over the limit. + wantCount(t, limiter, client, start.Add(-time.Millisecond), minute) +} + +func TestClockSetBackStartsTheBucketsAfresh(t *testing.T) { + t.Parallel() + + limiter := ratelimit.New(ratelimit.Limits{PerHour: limit}) + client := netip.MustParsePrefix("203.0.113.9/32") + start := midnight() + + for range limit { + wantCount(t, limiter, client, start, "") + } + + // Half an hour into the next bucket: 3 / 2 + 1 is within the limit. + wantCount(t, limiter, client, start.Add(time.Hour+time.Hour/2), "") + + // The clock is set back an hour. Counted in the bucket under way, the + // next request would find the bucket before it at full weight, 3 + 2, + // over the limit until the clock caught up. The buckets start afresh + // instead, and the client is refused only past the limit again. + setBack := start.Add(time.Hour / 2) + for range limit { + wantCount(t, limiter, client, setBack, "") + } + + wantCount(t, limiter, client, setBack, hour) +} + +func TestKeepsAtMost20000ClientsDroppingTheLeastRecentlySeen(t *testing.T) { + t.Parallel() + + const maxClients = 20000 + + limiter := ratelimit.New(ratelimit.Limits{PerMinute: 1}) + now := midnight() + + clients := make([]netip.Prefix, maxClients+1) + addr := netip.MustParseAddr("10.0.0.0") + + for i := range clients { + clients[i] = netip.PrefixFrom(addr, addr.BitLen()) + addr = addr.Next() + } + + for _, client := range clients[:maxClients] { + wantCount(t, limiter, client, now, "") + } + + // The first client is seen again: its second request is over the + // limit of one, so it is still counted. + wantCount(t, limiter, clients[0], now, minute) + + // One client more drops the least recently seen, the second, which + // starts afresh, while the first is kept. + wantCount(t, limiter, clients[maxClients], now, "") + wantCount(t, limiter, clients[1], now, "") + wantCount(t, limiter, clients[0], now, minute) +} + +// midnight is the start of a bucket in every window. +func midnight() time.Time { + return time.Date(2026, 10, 4, 0, 0, 0, 0, time.UTC) +} + +// wantCount counts a request from client at now, and checks the window +// whose limit it goes over, "" for none. +func wantCount( + t *testing.T, limiter *ratelimit.Limiter, client netip.Prefix, now time.Time, + want string, +) { + t.Helper() + + _, hit, _ := limiter.Count(client, now) + if hit.Window != want { + t.Errorf("request from %s at %s is over %q, want %q", + client, now.Format(time.RFC3339), hit.Window, want) + } +} diff --git a/internal/ratelimit/snapshot_test.go b/internal/ratelimit/snapshot_test.go new file mode 100644 index 0000000..251ae4d --- /dev/null +++ b/internal/ratelimit/snapshot_test.go @@ -0,0 +1,122 @@ +package ratelimit_test + +import ( + "net/netip" + "slices" + "testing" + "time" + + "sneak.berlin/go/smallwebwaf/internal/ratelimit" +) + +func TestSnapshotListsTheClientsByAddress(t *testing.T) { + t.Parallel() + + want := []string{"192.0.2.1/32", "203.0.113.9/32", "203.0.113.10/32", "2001:db8::/64"} + + limiter := ratelimit.New(ratelimit.Limits{}) + for _, i := range []int{2, 3, 0, 1} { + limiter.Count(netip.MustParsePrefix(want[i]), midnight()) + } + + snapshot := limiter.Snapshot() + + got := make([]string, 0, len(snapshot)) + for _, c := range snapshot { + got = append(got, c.Client.String()) + } + + if !slices.Equal(got, want) { + t.Errorf("snapshot %v, want %v", got, want) + } + + counted := ratelimit.Buckets{Start: midnight(), Current: 1} + if snapshot[0].Minute != counted || snapshot[0].Day != counted { + t.Errorf("buckets %+v and %+v, want %+v", snapshot[0].Minute, snapshot[0].Day, + counted) + } +} + +func TestLoadedCountsCarryOn(t *testing.T) { + t.Parallel() + + client := netip.MustParsePrefix("203.0.113.9/32") + start := midnight() + + before := ratelimit.New(ratelimit.Limits{PerHour: limit}) + for range limit { + wantCount(t, before, client, start, "") + } + + // Loaded into a new limiter, as across a restart, the client has no + // fresh allowance. + later := start.Add(time.Minute) + after := ratelimit.New(ratelimit.Limits{PerHour: limit}) + after.Load(before.Snapshot(), later) + wantCount(t, after, client, later, hour) +} + +func TestLoadEmptiesBucketsWhoseTimeHasPassed(t *testing.T) { + t.Parallel() + + client := netip.MustParsePrefix("203.0.113.9/32") + start := midnight() + + limiter := ratelimit.New(ratelimit.Limits{}) + limiter.Count(client, start) + limiter.AddToHistory(client, start, ratelimit.Request{Forwarded: true}) + + loaded := func(now time.Time) ratelimit.Client { + t.Helper() + + after := ratelimit.New(ratelimit.Limits{}) + after.Load(limiter.Snapshot(), now) + + return after.Snapshot()[0] + } + + // Two minutes on, the window that ends then covers neither of the + // minute's buckets, which are emptied; the hour's and the day's stay, + // and so does the history. + got := loaded(start.Add(2 * time.Minute)) + if got.Minute != (ratelimit.Buckets{}) || got.Hour.Current != 1 || + got.Day.Current != 1 || got.History.Requests != 1 { + t.Errorf("loaded two minutes on as %+v", got) + } + + // A moment before, the window still covers some of the earlier one. + got = loaded(start.Add(2*time.Minute - time.Nanosecond)) + if got.Minute.Current != 1 { + t.Errorf("loaded just under two minutes on with minute buckets %+v", + got.Minute) + } +} + +func TestLoadDropsTheLeastRecentlySeenFirst(t *testing.T) { + t.Parallel() + + const maxClients = 20000 + + // clients.json lists the clients by address. Here each was last seen + // a second before the one listed before it, so the last listed is the + // one seen longest ago, and the one dropped. + clients := make([]ratelimit.Client, maxClients+1) + addr := netip.MustParseAddr("10.0.0.0") + + for i := range clients { + clients[i].Client = netip.PrefixFrom(addr, addr.BitLen()) + clients[i].History.LastSeen = midnight().Add(-time.Duration(i) * time.Second) + addr = addr.Next() + } + + limiter := ratelimit.New(ratelimit.Limits{}) + limiter.Load(clients, midnight()) + + got := limiter.Snapshot() + if len(got) != maxClients || got[0].Client != clients[0].Client || + got[maxClients-1].Client != clients[maxClients-1].Client { + t.Errorf("%d clients kept, from %s to %s; want %d, from %s to %s", + len(got), got[0].Client, got[len(got)-1].Client, maxClients, + clients[0].Client, clients[maxClients-1].Client) + } +} diff --git a/internal/requestlog/requestlog.go b/internal/requestlog/requestlog.go new file mode 100644 index 0000000..a36003a --- /dev/null +++ b/internal/requestlog/requestlog.go @@ -0,0 +1,180 @@ +// Package requestlog writes the lines smallwebwaf prints on stdout: one +// JSON object per request, marked "type":"request", and the process's own +// messages as JSON lines marked "type":"process". +package requestlog + +import ( + "encoding/json" + "fmt" + "io" + "log/slog" + "time" + + "sneak.berlin/go/smallwebwaf/internal/ratelimit" +) + +// The action a request line names: what smallwebwaf did with the +// request. +const ( + // ActionForward is a request passed to the app. + ActionForward = "forward" + // ActionTooLarge is a request or response over its size limit. + ActionTooLarge = "too_large" + // ActionTimedOut is a request or response that ran out of time. + ActionTimedOut = "timed_out" + // ActionUpstreamError is a request the app could not be reached + // for, or whose answer could not be passed on. + ActionUpstreamError = "upstream_error" + // ActionRateLimited is a request refused because it took its client + // over a rate limit, which bans the client. + ActionRateLimited = "rate_limited" + // ActionBanned is a request refused because a ban covers its client. + ActionBanned = "banned" + // ActionDenied is a request refused because its client is in + // SWWAF_DENY_NETS. + ActionDenied = "denied" + // ActionCountryDenied is a request refused for its client's country. + ActionCountryDenied = "country_denied" + // ActionAdmin is a request smallwebwaf answered at one of its own + // endpoints, under /_smallwebwaf/. + ActionAdmin = "admin" +) + +// OffenceLimit is the offence a request line names for a request that +// broke a rate limit. +const OffenceLimit = "limit" + +// timeLayout is RFC 3339 with milliseconds. +const timeLayout = "2006-01-02T15:04:05.000Z07:00" + +// Line is one request's line in the request log. The field names, and +// their order, are those of the "Request log" section of SPEC.md. A field +// that may not apply to a request is left out of its line when it does +// not. +// +//nolint:tagliatelle // SPEC.md's request log names its fields in snake_case +type Line struct { + Type string `json:"type"` + + // The standard web log fields. Scheme is how the client reached + // smallwebwaf, or the trusted proxy in front of it. + Time string `json:"time"` + Instance string `json:"instance"` + ClientIP string `json:"client_ip"` + Method string `json:"method"` + Scheme string `json:"scheme"` + Host string `json:"host"` + Path string `json:"path"` + Query string `json:"query"` + Protocol string `json:"protocol"` + Status int `json:"status"` + RequestBytes int64 `json:"request_bytes"` + ResponseBytes int64 `json:"response_bytes"` + Referer string `json:"referer"` + UserAgent string `json:"user_agent"` + + // Request detail. RequestID is the X-Request-ID a trusted proxy sent, + // or a new one, and is sent on to the app. ForwardedFor is the + // X-Forwarded-For header as received. ClientGroup is the netblock the + // client is counted as. + RequestID string `json:"request_id"` + PeerIP string `json:"peer_ip"` + ForwardedFor string `json:"forwarded_for,omitempty"` + ClientGroup string `json:"client_group"` + Country string `json:"country"` + ContentType string `json:"content_type,omitempty"` + // ContentLength is the length of its body the request announced. + ContentLength int64 `json:"content_length,omitempty"` + // RequestHeaders are the headers SWWAF_LOG_REQUEST_HEADERS names that + // the request carried, by name in lower case. + RequestHeaders map[string]string `json:"request_headers,omitempty"` + HasAuthorization bool `json:"has_authorization,omitempty"` + HasCookie bool `json:"has_cookie,omitempty"` + // Websocket is true when the connection was upgraded, as for a + // WebSocket. + Websocket bool `json:"websocket,omitempty"` + + // Response detail, from the headers of the answer: the app's, as + // passed on, or those of smallwebwaf's own. Aborted is true when the + // client went away early. + ResponseContentType string `json:"response_content_type,omitempty"` + UpstreamStatus int `json:"upstream_status,omitempty"` + CacheControl string `json:"cache_control,omitempty"` + Location string `json:"location,omitempty"` + Aborted bool `json:"aborted,omitempty"` + + // The decision. + Action string `json:"action"` + // WouldAction is, in observe mode, the action enforce mode would have + // taken with a request it would have refused: ActionDenied, + // ActionBanned, ActionCountryDenied or ActionRateLimited. + WouldAction string `json:"would_action,omitempty"` + // Counts are the client's requests as the rate limits counted them + // with this one, for a request they counted. + Counts ratelimit.Counts `json:"counts,omitzero"` + // LimitHit is the window whose rate limit the request went over: + // minute, hour or day. + LimitHit string `json:"limit_hit,omitempty"` + // Offence is the offence the request was held as, OffenceLimit. + Offence string `json:"offence,omitempty"` + // BanExpires is when the ban the request made, or was refused under, + // ends: a time, or "permanent". + BanExpires string `json:"ban_expires,omitempty"` + + // The timings, in milliseconds. DurationChecks is the time until the + // checks were done. DurationUpstreamConnect, DurationUpstreamFirstByte + // and DurationUpstreamTotal run from when the request was handed to the + // app: until there was a connection to it, until the first byte of its + // answer arrived, and until the end. Each but DurationTotal is nil for + // a request that did not get that far. + DurationTotal float64 `json:"duration_total"` + DurationChecks *float64 `json:"duration_checks,omitempty"` + DurationUpstreamConnect *float64 `json:"duration_upstream_connect,omitempty"` + DurationUpstreamFirstByte *float64 `json:"duration_upstream_first_byte,omitempty"` + DurationUpstreamTotal *float64 `json:"duration_upstream_total,omitempty"` +} + +// Write writes line to w as one JSON line marked "type":"request". +func Write(w io.Writer, line *Line) error { + line.Type = "request" + + encoded, err := json.Marshal(line) + if err != nil { + return fmt.Errorf("encode the request log line: %w", err) + } + + _, err = w.Write(append(encoded, '\n')) + if err != nil { + return fmt.Errorf("write the request log line: %w", err) + } + + return nil +} + +// FormatTime formats t for a line's time field: RFC 3339 in UTC, with +// milliseconds. +func FormatTime(t time.Time) string { + return t.UTC().Format(timeLayout) +} + +// Milliseconds is d in milliseconds, to the microsecond. +func Milliseconds(d time.Duration) float64 { + return float64(d.Microseconds()) / float64(time.Millisecond/time.Microsecond) +} + +// NewProcessLogger returns the logger for the process's own messages: +// JSON lines on w, marked "type":"process", with the time in the same form +// as a request line's. +func NewProcessLogger(w io.Writer) *slog.Logger { + handler := slog.NewJSONHandler(w, &slog.HandlerOptions{ + ReplaceAttr: func(groups []string, attr slog.Attr) slog.Attr { + if attr.Key == slog.TimeKey && len(groups) == 0 { + return slog.String(slog.TimeKey, FormatTime(attr.Value.Time())) + } + + return attr + }, + }) + + return slog.New(handler).With("type", "process") +} diff --git a/internal/requestlog/requestlog_test.go b/internal/requestlog/requestlog_test.go new file mode 100644 index 0000000..38303b6 --- /dev/null +++ b/internal/requestlog/requestlog_test.go @@ -0,0 +1,95 @@ +package requestlog_test + +import ( + "bytes" + "encoding/json" + "strings" + "testing" + "time" + + "sneak.berlin/go/smallwebwaf/internal/requestlog" +) + +func TestWriteWritesOneJSONLineMarkedRequest(t *testing.T) { + t.Parallel() + + var out bytes.Buffer + + err := requestlog.Write(&out, &requestlog.Line{ + Time: requestlog.FormatTime(time.Date(2026, 10, 3, 12, 0, 0, 0, time.UTC)), + ClientIP: "203.0.113.9", + Status: 200, + Action: requestlog.ActionForward, + DurationTotal: requestlog.Milliseconds(1500 * time.Microsecond), + }) + if err != nil { + t.Fatalf("write: %v", err) + } + + text := out.String() + if strings.Count(text, "\n") != 1 || !strings.HasSuffix(text, "\n") { + t.Fatalf("wrote %q, want one line", text) + } + + var fields map[string]any + + err = json.Unmarshal(out.Bytes(), &fields) + if err != nil { + t.Fatalf("decode %q: %v", text, err) + } + + want := map[string]any{ + "type": "request", "time": "2026-10-03T12:00:00.000Z", + "client_ip": "203.0.113.9", "status": 200.0, "action": "forward", + "duration_total": 1.5, + } + for name, value := range want { + if fields[name] != value { + t.Errorf("%s is %v, want %v", name, fields[name], value) + } + } + + unset := []string{ + "forwarded_for", "content_type", "content_length", "request_headers", + "has_authorization", "has_cookie", "websocket", "response_content_type", + "upstream_status", "cache_control", "location", "aborted", "counts", + "limit_hit", "offence", "ban_expires", "duration_checks", + "duration_upstream_connect", "duration_upstream_first_byte", + "duration_upstream_total", + } + for _, name := range unset { + _, present := fields[name] + if present { + t.Errorf("%s is there with no value to give", name) + } + } +} + +func TestProcessLinesAreMarkedProcess(t *testing.T) { + t.Parallel() + + var out bytes.Buffer + + requestlog.NewProcessLogger(&out).Info("starting", "version", "v1") + + var fields map[string]any + + err := json.Unmarshal(out.Bytes(), &fields) + if err != nil { + t.Fatalf("decode %q: %v", out.String(), err) + } + + if fields["type"] != "process" || fields["msg"] != "starting" || + fields["level"] != "INFO" || fields["version"] != "v1" { + t.Errorf("process line %v", fields) + } + + timeText, _ := fields["time"].(string) + + logged, err := time.Parse(time.RFC3339, timeText) + if err != nil || !strings.HasSuffix(timeText, "Z") || + len(timeText) != len("2006-01-02T15:04:05.000Z") || + time.Since(logged) > time.Minute { + t.Errorf("process line time %q, want now in UTC with milliseconds", timeText) + } +} diff --git a/internal/smallwebwaf/healthcheck.go b/internal/smallwebwaf/healthcheck.go new file mode 100644 index 0000000..9269f3e --- /dev/null +++ b/internal/smallwebwaf/healthcheck.go @@ -0,0 +1,100 @@ +package smallwebwaf + +import ( + "context" + "errors" + "fmt" + "io" + "net" + "net/http" + "net/url" + "time" + + "sneak.berlin/go/smallwebwaf/internal/config" + "sneak.berlin/go/smallwebwaf/internal/proxy" +) + +// healthCheckTimeout bounds the whole health check. +const healthCheckTimeout = 5 * time.Second + +var errHealthEndpoint = errors.New("smallwebwaf's health endpoint answered") + +// HealthCheck is the container's health check. It returns 0 while +// smallwebwaf answers its health endpoint on 127.0.0.1, at the port in +// SWWAF_LISTEN_ADDR, and the app accepts connections at the address in +// SWWAF_UPSTREAM_URL. Otherwise it writes why to stderr and returns 1. +// args are the arguments after `healthcheck`; it takes none, and given +// one it names it on stderr and returns 1 without checking anything. +func HealthCheck( + ctx context.Context, args []string, lookupEnv func(string) (string, bool), + stderr io.Writer, +) int { + if len(args) > 0 { + _, _ = fmt.Fprintf(stderr, + "smallwebwaf healthcheck: unexpected argument %q\n", args[0]) + + return 1 + } + + err := healthCheck(ctx, lookupEnv) + if err != nil { + _, _ = fmt.Fprintln(stderr, "unhealthy:", err) + + return 1 + } + + return 0 +} + +func healthCheck(ctx context.Context, lookupEnv func(string) (string, bool)) error { + ctx, cancel := context.WithTimeout(ctx, healthCheckTimeout) + defer cancel() + + cfg, err := config.FromEnvironment(lookupEnv) + if err != nil { + return fmt.Errorf("invalid setting: %w", err) + } + + // The settings have checked that the address has a port. + _, port, _ := net.SplitHostPort(cfg.ListenAddr) + health := "http://" + net.JoinHostPort("127.0.0.1", port) + proxy.HealthPath + + req, err := http.NewRequestWithContext(ctx, http.MethodGet, health, http.NoBody) + if err != nil { + return fmt.Errorf("make the request: %w", err) + } + + res, err := http.DefaultClient.Do(req) + if err != nil { + return fmt.Errorf("ask smallwebwaf: %w", err) + } + + _ = res.Body.Close() + + if res.StatusCode != http.StatusOK { + return fmt.Errorf("%w %s", errHealthEndpoint, res.Status) + } + + conn, err := (&net.Dialer{}).DialContext(ctx, "tcp", appAddress(cfg.UpstreamURL)) + if err != nil { + return fmt.Errorf("connect to the app: %w", err) + } + + _ = conn.Close() + + return nil +} + +// appAddress is the host and port of the app's URL, the port being the +// scheme's own when the URL names none. +func appAddress(app *url.URL) string { + port := app.Port() + if port == "" { + port = "80" + if app.Scheme == "https" { + port = "443" + } + } + + return net.JoinHostPort(app.Hostname(), port) +} diff --git a/internal/smallwebwaf/healthcheck_internal_test.go b/internal/smallwebwaf/healthcheck_internal_test.go new file mode 100644 index 0000000..7aec7f9 --- /dev/null +++ b/internal/smallwebwaf/healthcheck_internal_test.go @@ -0,0 +1,27 @@ +package smallwebwaf + +import ( + "net/url" + "testing" +) + +func TestAppAddress(t *testing.T) { + t.Parallel() + + for app, want := range map[string]string{ + "http://127.0.0.1:8081": "127.0.0.1:8081", + "http://app": "app:80", + "https://app/": "app:443", + "https://[::1]": "[::1]:443", + } { + parsed, err := url.Parse(app) + if err != nil { + t.Fatalf("parse %q: %v", app, err) + } + + got := appAddress(parsed) + if got != want { + t.Errorf("appAddress(%q) is %q, want %q", app, got, want) + } + } +} diff --git a/internal/smallwebwaf/healthcheck_test.go b/internal/smallwebwaf/healthcheck_test.go new file mode 100644 index 0000000..ae88403 --- /dev/null +++ b/internal/smallwebwaf/healthcheck_test.go @@ -0,0 +1,99 @@ +package smallwebwaf_test + +import ( + "bytes" + "context" + "net" + "net/http" + "net/http/httptest" + "strings" + "testing" + "time" + + "sneak.berlin/go/smallwebwaf/internal/smallwebwaf" +) + +func TestHealthCheck(t *testing.T) { + t.Parallel() + + app := httptest.NewServer(http.NotFoundHandler()) + defer app.Close() + + ctx, stop := context.WithCancel(t.Context()) + defer stop() + + out := &output{} + exited := make(chan int, 1) + settings := map[string]string{ + listenAddr: localhost + ":0", + upstreamURL: app.URL, + stateDir: t.TempDir(), + } + + go func() { + exited <- run(ctx, settings, out) + }() + + addr, _ := out.line(t, "msg", "starting")["address"].(string) + _, port, _ := net.SplitHostPort(addr) + // The container's settings: an address to listen on with an empty + // host part, which the health check asks at 127.0.0.1. + env := map[string]string{listenAddr: ":" + port, upstreamURL: app.URL} + + wantHealthCheck(t, env, 0, "") + + app.Close() + wantHealthCheck(t, env, 1, "unhealthy: connect to the app: ") + + stop() + + select { + case <-exited: + case <-time.After(waitLimit): + t.Fatal("still running after being told to stop") + } + + wantHealthCheck(t, env, 1, "unhealthy: ask smallwebwaf: ") + wantHealthCheck(t, map[string]string{listenAddr: "8080"}, 1, + "unhealthy: invalid setting: SWWAF_LISTEN_ADDR: ") +} + +func TestHealthCheckRefusesAnArgument(t *testing.T) { + t.Parallel() + + var stderr bytes.Buffer + + noSettings := func(string) (string, bool) { + return "", false + } + + got := smallwebwaf.HealthCheck(t.Context(), []string{"now"}, noSettings, &stderr) + + want := "smallwebwaf healthcheck: unexpected argument \"now\"\n" + if got != 1 || stderr.String() != want { + t.Errorf("health check returned %d and wrote %q, want 1 and %q", + got, stderr.String(), want) + } +} + +// wantHealthCheck runs the health check with the settings in env, and +// checks its exit status and the start of what it writes to stderr, +// which is nothing when message is empty. +func wantHealthCheck(t *testing.T, env map[string]string, status int, message string) { + t.Helper() + + var stderr bytes.Buffer + + got := smallwebwaf.HealthCheck(t.Context(), nil, func(name string) (string, bool) { + value, ok := env[name] + + return value, ok + }, &stderr) + + wrote := stderr.String() + if got != status || !strings.HasPrefix(wrote, message) || + (message == "" && wrote != "") { + t.Errorf("health check returned %d and wrote %q, want %d and %q", + got, wrote, status, message) + } +} diff --git a/internal/smallwebwaf/smallwebwaf.go b/internal/smallwebwaf/smallwebwaf.go new file mode 100644 index 0000000..06b446b --- /dev/null +++ b/internal/smallwebwaf/smallwebwaf.go @@ -0,0 +1,195 @@ +// Package smallwebwaf runs the smallwebwaf process: it reads the settings +// and the state files, serves requests until it is told to stop, and then +// stops in an orderly way, writing the state files. +package smallwebwaf + +import ( + "context" + "errors" + "io" + "log/slog" + "net" + "net/http" + "os" + "os/signal" + "syscall" + "time" + + "sneak.berlin/go/smallwebwaf/internal/config" + "sneak.berlin/go/smallwebwaf/internal/lookup" + "sneak.berlin/go/smallwebwaf/internal/proxy" + "sneak.berlin/go/smallwebwaf/internal/requestlog" + "sneak.berlin/go/smallwebwaf/internal/state" +) + +// shutdownTimeout is how long requests in progress may take to finish +// once smallwebwaf is told to stop, before their connections are closed. +// runit and docker wait a little longer before they kill the process. +const shutdownTimeout = 5 * time.Second + +// Params are what Run needs from the process. +type Params struct { + // Version is the version of the binary, set when it is built. + Version string + // LookupEnv reads an environment variable, normally os.LookupEnv. + LookupEnv func(string) (string, bool) + // Stdout receives the request log and the process's own messages. + Stdout io.Writer +} + +// Main runs smallwebwaf until SIGTERM or SIGINT, and returns the +// process's exit status. Run as `smallwebwaf healthcheck`, it is the +// container's health check instead. +func Main(version string) int { + if len(os.Args) > 1 && os.Args[1] == "healthcheck" { + return HealthCheck(context.Background(), os.Args[2:], os.LookupEnv, os.Stderr) + } + + ctx, stop := signal.NotifyContext(context.Background(), + syscall.SIGTERM, os.Interrupt) + defer stop() + + return Run(ctx, Params{ + Version: version, + LookupEnv: os.LookupEnv, + Stdout: os.Stdout, + }) +} + +// Run reads the settings and the state files, then serves requests until +// ctx is done. It returns the process's exit status, 1 when smallwebwaf +// cannot start. +func Run(ctx context.Context, params Params) int { + processLog := requestlog.NewProcessLogger(params.Stdout) + + cfg, err := config.FromEnvironment(params.LookupEnv) + if err != nil { + processLog.Error("invalid setting", "error", err.Error()) + + return 1 + } + + // The state files give times in UTC. + now := func() time.Time { return time.Now().UTC() } + + server := proxy.New(proxy.Params{ + Config: cfg, + RequestLog: params.Stdout, + ProcessLog: processLog, + GeoJSURL: lookup.URL, + Now: now, + }) + + files, err := state.Load(state.Params{ + Dir: cfg.StateDir, + WriteDelay: cfg.StateWriteDelay, + CounterInterval: cfg.StateCounterInterval, + Ledger: server.Ledger, + Limiter: server.Limiter, + GeoJS: server.GeoJS, + Now: now, + ProcessLog: processLog, + Metrics: server.Metrics, + }) + if err != nil { + processLog.Error("cannot use the state files", "error", err.Error()) + + return 1 + } + + listener, err := (&net.ListenConfig{}).Listen(ctx, "tcp", cfg.ListenAddr) + if err != nil { + processLog.Error("cannot listen on SWWAF_LISTEN_ADDR", + "error", err.Error()) + + return 1 + } + + processLog.Info("starting", + "version", params.Version, + "address", listener.Addr().String(), + "settings", cfg) + + return serve(ctx, server.Server, listener, files, processLog) +} + +// serve serves requests on listener, writes the state files as they are +// due, and takes in an admin's edits of them, until ctx is done. Then it +// gives the requests in progress shutdownTimeout to finish, and writes +// every state file. +func serve( + ctx context.Context, server *http.Server, listener net.Listener, + files *state.Files, processLog *slog.Logger, +) int { + served := make(chan error, 1) + + go func() { + served <- server.Serve(listener) + }() + + writing, stopWriting := context.WithCancel(ctx) + defer stopWriting() + + written := make(chan struct{}) + watched := make(chan struct{}) + + go func() { + files.Run(writing) + close(written) + }() + + go func() { + files.Watch(writing) + close(watched) + }() + + select { + case err := <-served: + processLog.Error("serving failed", "error", err.Error()) + + return 1 + case <-ctx.Done(): + } + + processLog.Info("stopping") + + shutdownCtx, cancel := context.WithTimeout(context.WithoutCancel(ctx), + shutdownTimeout) + defer cancel() + + err := server.Shutdown(shutdownCtx) + if err != nil { + processLog.Warn("requests still in progress were cut off", + "error", err.Error()) + + _ = server.Close() + } + + err = <-served + if !errors.Is(err, http.ErrServerClosed) { + processLog.Error("serving failed", "error", err.Error()) + + return 1 + } + + // Run and Watch have ended, so nothing else reads or writes the + // files. Every request has ended too, but for two kinds + // that Go's server does not wait for: one cut off because Shutdown + // timed out, and one whose connection switched protocols, such as a + // WebSocket. Such a request adds to its client's history only as it + // ends, which can be after this write, and then that request is + // missing from clients.json. + <-written + <-watched + + err = files.WriteAll() + if err != nil { + processLog.Error("writing the state files failed", "error", err.Error()) + + return 1 + } + + processLog.Info("stopped") + + return 0 +} diff --git a/internal/smallwebwaf/smallwebwaf_test.go b/internal/smallwebwaf/smallwebwaf_test.go new file mode 100644 index 0000000..befbab9 --- /dev/null +++ b/internal/smallwebwaf/smallwebwaf_test.go @@ -0,0 +1,570 @@ +package smallwebwaf_test + +import ( + "bytes" + "context" + "encoding/json" + "io" + "net" + "net/http" + "net/http/httptest" + "os" + "path/filepath" + "strings" + "sync" + "testing" + "time" + + "sneak.berlin/go/smallwebwaf/internal/smallwebwaf" +) + +const ( + // waitLimit bounds how long a test waits for what should happen. + waitLimit = 10 * time.Second + // pollInterval is how often a test looks for a line. + pollInterval = 10 * time.Millisecond + // testVersion is the version the tests give smallwebwaf. + testVersion = "test" + // localhost is where the tests listen. + localhost = "127.0.0.1" + listenAddr = "SWWAF_LISTEN_ADDR" + upstreamURL = "SWWAF_UPSTREAM_URL" + trustedProxies = "SWWAF_TRUSTED_PROXIES" + stateDir = "SWWAF_STATE_DIR" + stateWriteDelay = "SWWAF_STATE_WRITE_DELAY" + stateCounterInterval = "SWWAF_STATE_COUNTER_INTERVAL" + rateLimitPerDay = "SWWAF_RATE_LIMIT_PER_DAY" + // greeting is what the tests' app answers. + greeting = "hello from the app" +) + +// output collects what smallwebwaf writes on stdout. +type output struct { + mu sync.Mutex + buf bytes.Buffer +} + +// Write adds lines smallwebwaf writes. +func (o *output) Write(p []byte) (int, error) { + o.mu.Lock() + defer o.mu.Unlock() + + return o.buf.Write(p) +} + +// line returns the first line whose field key is value, waiting for it. +func (o *output) line(t *testing.T, key, value string) map[string]any { + t.Helper() + + deadline := time.Now().Add(waitLimit) + for time.Now().Before(deadline) { + o.mu.Lock() + text := o.buf.String() + o.mu.Unlock() + + for line := range strings.Lines(text) { + var fields map[string]any + + err := json.Unmarshal([]byte(line), &fields) + if err != nil { + t.Fatalf("output line %q is not JSON: %v", line, err) + } + + if fields[key] == value { + return fields + } + } + + time.Sleep(pollInterval) + } + + t.Fatalf("no line with %s %q in the output:\n%s", key, value, o.text()) + + return nil +} + +// text returns everything written so far. +func (o *output) text() string { + o.mu.Lock() + defer o.mu.Unlock() + + return o.buf.String() +} + +// run runs smallwebwaf with the settings in env until ctx is done, and +// returns its exit status. +func run(ctx context.Context, env map[string]string, out *output) int { + return smallwebwaf.Run(ctx, smallwebwaf.Params{ + Version: testVersion, + LookupEnv: func(name string) (string, bool) { + value, ok := env[name] + + return value, ok + }, + Stdout: out, + }) +} + +func TestInvalidSettingStopsTheStart(t *testing.T) { + t.Parallel() + + out := &output{} + + status := run(t.Context(), map[string]string{"SWWAF_REQUEST_MAX_BYTES": "lots"}, out) + if status != 1 { + t.Errorf("exit status %d, want 1", status) + } + + line := out.line(t, "msg", "invalid setting") + message, _ := line["error"].(string) + + if line["type"] != "process" || line["level"] != "ERROR" || + !strings.HasPrefix(message, "SWWAF_REQUEST_MAX_BYTES: ") { + t.Errorf("start refused with %v", line) + } +} + +func TestShortMetricsTokenStopsTheStartUnshown(t *testing.T) { + t.Parallel() + + const token = "a-token-of-31-characters-at-all" //nolint:gosec // too short to use + + out := &output{} + + status := run(t.Context(), map[string]string{"SWWAF_METRICS_TOKEN": token}, out) + if status != 1 { + t.Errorf("exit status %d, want 1", status) + } + + line := out.line(t, "msg", "invalid setting") + if line["error"] != "SWWAF_METRICS_TOKEN: is shorter than 32 characters" { + t.Errorf("start refused with %v", line) + } + + if strings.Contains(out.text(), token) { + t.Errorf("the output shows the token:\n%s", out.text()) + } +} + +func TestAddressInUseStopsTheStart(t *testing.T) { + t.Parallel() + + taken, err := (&net.ListenConfig{}).Listen(t.Context(), "tcp", localhost+":0") + if err != nil { + t.Fatalf("listen: %v", err) + } + + defer func() { + _ = taken.Close() + }() + + out := &output{} + + status := run(t.Context(), map[string]string{ + listenAddr: taken.Addr().String(), + stateDir: t.TempDir(), + }, out) + if status != 1 { + t.Errorf("exit status %d, want 1", status) + } + + out.line(t, "msg", "cannot listen on SWWAF_LISTEN_ADDR") +} + +func TestServesUntilToldToStop(t *testing.T) { + t.Parallel() + + appURL := startApp(t) + dir := t.TempDir() + + ctx, stop := context.WithCancel(t.Context()) + out := &output{} + exited := make(chan int, 1) + + go func() { + exited <- run(ctx, map[string]string{ + listenAddr: localhost + ":0", + upstreamURL: appURL, + stateDir: dir, + }, out) + }() + + starting := out.line(t, "msg", "starting") + wantStartingLine(t, starting, appURL, dir) + + addr, _ := starting["address"].(string) + wantGreeting(t, "http://"+addr+"/") + out.line(t, "type", "request") + + stop() + + select { + case status := <-exited: + if status != 0 { + t.Errorf("exit status %d, want 0", status) + } + case <-time.After(waitLimit): + t.Fatal("still running after being told to stop") + } + + out.line(t, "msg", "stopped") +} + +func TestStateKeptAcrossRestarts(t *testing.T) { + t.Parallel() + + env := map[string]string{ + listenAddr: localhost + ":0", + upstreamURL: startApp(t), + stateDir: t.TempDir(), + rateLimitPerDay: "2", + // Neither comes due in the test: the files are written as + // smallwebwaf stops. + stateWriteDelay: "1h", + stateCounterInterval: "1h", + } + + // The two requests a day allows, and a stop. + runUntilStopped(t, env, func(url string) { + wantGreeting(t, url) + wantGreeting(t, url) + }) + + // After a restart the client has no fresh allowance: its third + // request breaks the day limit, and bans it. + out := runUntilStopped(t, env, func(url string) { + wantRefused(t, url) + }) + out.line(t, "action", "rate_limited") + + // After another, the ban still refuses it. + out = runUntilStopped(t, env, func(url string) { + wantRefused(t, url) + }) + out.line(t, "action", "banned") +} + +func TestBanRefusesItsNetblockAfterARestartWithAnotherScope(t *testing.T) { + t.Parallel() + + const scope = "SWWAF_BAN_SCOPE_V4_PREFIX" + + env := map[string]string{ + listenAddr: localhost + ":0", + upstreamURL: startApp(t), + stateDir: t.TempDir(), + trustedProxies: localhost + "/32", + rateLimitPerDay: "1", + scope: "24", + } + + // 203.0.113.9's second request breaks the day limit, and bans + // 203.0.113.0/24. + runUntilStopped(t, env, func(url string) { + wantStatus(t, url, "203.0.113.9", http.StatusOK) + wantStatus(t, url, "203.0.113.9", http.StatusForbidden) + }) + + // With each address a netblock of its own after a restart, that ban + // still refuses all of 203.0.113.0/24. 198.51.100.7 is banned alone. + env[scope] = "32" + runUntilStopped(t, env, func(url string) { + wantStatus(t, url, "203.0.113.200", http.StatusForbidden) + wantStatus(t, url, "203.0.114.1", http.StatusOK) + wantStatus(t, url, "198.51.100.7", http.StatusOK) + wantStatus(t, url, "198.51.100.7", http.StatusForbidden) + }) + + // With /24 netblocks again, that ban still refuses 198.51.100.7, and + // no other address. + env[scope] = "24" + runUntilStopped(t, env, func(url string) { + wantStatus(t, url, "198.51.100.7", http.StatusForbidden) + wantStatus(t, url, "198.51.100.8", http.StatusOK) + }) +} + +func TestBanAddedAndLiftedByEditingBansJSON(t *testing.T) { + t.Parallel() + + const ( + // bans.json as an admin writes it with a ban, permanent, on + // 203.0.113.0/24, and with none. + oneBan = `{"version": 1, "bans": [{"netblock": "203.0.113.0/24", ` + + `"start": "2026-10-06T00:00:00Z", "expires": null}]}` + noBan = `{"version": 1, "bans": []}` + ) + + dir := t.TempDir() + env := map[string]string{ + listenAddr: localhost + ":0", + upstreamURL: startApp(t), + stateDir: dir, + trustedProxies: localhost + "/32", + // No write comes due in the test, so only the watch on the + // directory can take the edits in. + stateWriteDelay: "1h", + stateCounterInterval: "1h", + } + + runUntilStopped(t, env, func(url string) { + path := filepath.Join(dir, "bans.json") + + saveUntilAnswered(t, path, oneBan, url, "203.0.113.9", http.StatusForbidden) + wantStatus(t, url, "198.51.100.7", http.StatusOK) + saveUntilAnswered(t, path, noBan, url, "203.0.113.9", http.StatusOK) + }) +} + +func TestStateFileThatDoesNotParseStopsTheStart(t *testing.T) { + t.Parallel() + + dir := t.TempDir() + + err := os.WriteFile(filepath.Join(dir, "bans.json"), []byte("{\n"), 0o600) + if err != nil { + t.Fatalf("write bans.json: %v", err) + } + + // The file ends at the newline that is the second byte of its first + // line. + wantStartRefused(t, dir, filepath.Join(dir, "bans.json")+", line 1, column 2: ") +} + +func TestUnwritableStateDirStopsTheStart(t *testing.T) { + t.Parallel() + + wantStartRefused(t, filepath.Join(t.TempDir(), "missing"), + "SWWAF_STATE_DIR cannot be written: ") +} + +// wantStartRefused runs smallwebwaf with its state files in dir, and +// checks that it stops at start, with an error that starts with want. If +// it starts instead, it is stopped after waitLimit. +func wantStartRefused(t *testing.T, dir, want string) { + t.Helper() + + ctx, stop := context.WithTimeout(t.Context(), waitLimit) + defer stop() + + out := &output{} + + status := run(ctx, map[string]string{listenAddr: localhost + ":0", stateDir: dir}, out) + if status != 1 { + t.Fatalf("exit status %d, want 1", status) + } + + line := out.line(t, "msg", "cannot use the state files") + message, _ := line["error"].(string) + + if !strings.HasPrefix(message, want) { + t.Errorf("start refused with %q, want an error starting %q", message, want) + } +} + +// startApp starts an app that answers every request with greeting, and +// returns its URL. +func startApp(t *testing.T) string { + t.Helper() + + app := httptest.NewServer(http.HandlerFunc( + func(w http.ResponseWriter, _ *http.Request) { + _, _ = io.WriteString(w, greeting) + })) + t.Cleanup(app.Close) + + return app.URL +} + +// runUntilStopped runs smallwebwaf with the settings in env, has use send +// it requests at url, then stops it as SIGTERM does, checks that it +// stopped in order, and returns its output. +func runUntilStopped( + t *testing.T, env map[string]string, use func(url string), +) *output { + t.Helper() + + ctx, stop := context.WithCancel(t.Context()) + out := &output{} + exited := make(chan int, 1) + + go func() { + exited <- run(ctx, env, out) + }() + + addr, _ := out.line(t, "msg", "starting")["address"].(string) + use("http://" + addr + "/") + stop() + + select { + case status := <-exited: + if status != 0 { + t.Fatalf("exit status %d, want 0; output:\n%s", status, out.text()) + } + case <-time.After(waitLimit): + t.Fatal("still running after being told to stop") + } + + return out +} + +// wantStartingLine checks that the line at start gives the version and +// every setting's value. +func wantStartingLine(t *testing.T, line map[string]any, appURL, dir string) { + t.Helper() + + settings, _ := line["settings"].(map[string]any) + want := map[string]any{ + listenAddr: localhost + ":0", + upstreamURL: appURL, + stateDir: dir, + "SWWAF_MODE": "enforce", + stateWriteDelay: "10s", + stateCounterInterval: "15m", + trustedProxies: "10.0.0.0/8,172.16.0.0/12,192.168.0.0/16", + "SWWAF_CLIENT_REQUEST_TIMEOUT": "60s", + "SWWAF_CLIENT_REQUEST_HEADER_MAX_BYTES": "32K", + "SWWAF_CLIENT_IDLE_TIMEOUT": "120s", + "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_ALLOW_NETS": "", + "SWWAF_RATE_LIMIT_EXEMPT_NETS": "", + "SWWAF_DENY_NETS": "", + "SWWAF_RATE_LIMIT_PER_MINUTE": "1000", + "SWWAF_RATE_LIMIT_PER_HOUR": "10000", + rateLimitPerDay: "50000", + "SWWAF_DENIED_COUNTRIES": "", + "SWWAF_EXCLUSIVELY_ALLOWED_COUNTRIES": "", + "SWWAF_BAN_RESPONSE": "403", + "SWWAF_LIMIT_BAN_DURATION": "1h", + "SWWAF_LIMIT_BAN_REPEAT_WINDOW": "24h", + "SWWAF_MAX_BAN_DURATION": "7d", + "SWWAF_MAX_BANS": "5000", + "SWWAF_BAN_SCOPE_V4_PREFIX": "32", + } + + for name, value := range want { + if settings[name] != value { + t.Errorf("starting line gives %s=%v, want %v", name, settings[name], value) + } + } + + if line["version"] != testVersion || line["type"] != "process" { + t.Errorf("starting line %v", line) + } +} + +// wantGreeting checks that a request to url gets the app's answer. +func wantGreeting(t *testing.T, url string) { + t.Helper() + + req, err := http.NewRequestWithContext(t.Context(), http.MethodGet, url, + http.NoBody) + if err != nil { + t.Fatalf("new request: %v", err) + } + + transport := &http.Transport{} + defer transport.CloseIdleConnections() + + res, err := (&http.Client{Transport: transport}).Do(req) + if err != nil { + t.Fatalf("request: %v", err) + } + + body, err := io.ReadAll(res.Body) + _ = res.Body.Close() + + if err != nil || string(body) != greeting { + t.Errorf("got %q (%v), want the app's answer", body, err) + } +} + +// wantRefused checks that a request to url is refused with 403, the +// default SWWAF_BAN_RESPONSE. +func wantRefused(t *testing.T, url string) { + t.Helper() + + req, err := http.NewRequestWithContext(t.Context(), http.MethodGet, url, + http.NoBody) + if err != nil { + t.Fatalf("new request: %v", err) + } + + transport := &http.Transport{} + defer transport.CloseIdleConnections() + + res, err := (&http.Client{Transport: transport}).Do(req) + if err != nil { + t.Fatalf("request: %v", err) + } + + _ = res.Body.Close() + + if res.StatusCode != http.StatusForbidden { + t.Errorf("status %d, want %d", res.StatusCode, http.StatusForbidden) + } +} + +// wantStatus checks that a request to url from the client at from, as +// X-Forwarded-For names it, is answered with status. +func wantStatus(t *testing.T, url, from string, status int) { + t.Helper() + + got := statusFrom(t, url, from) + if got != status { + t.Errorf("request from %s: status %d, want %d", from, got, status) + } +} + +// saveUntilAnswered writes content to the state file at path, as an +// admin saves an edit of it, until a request to url from the client at +// from is answered with status. The file is written again before each +// request, since smallwebwaf may not watch its directory yet when it is +// first written. It waits as long as that takes, so that a slow test +// process cannot fail the test. +func saveUntilAnswered(t *testing.T, path, content, url, from string, status int) { + t.Helper() + + for { + err := os.WriteFile(path, []byte(content), 0o600) + if err != nil { + t.Fatalf("write %s: %v", path, err) + } + + if statusFrom(t, url, from) == status { + return + } + + time.Sleep(pollInterval) + } +} + +// statusFrom returns the status a request to url from the client at +// from, as X-Forwarded-For names it, is answered with. +func statusFrom(t *testing.T, url, from string) int { + t.Helper() + + req, err := http.NewRequestWithContext(t.Context(), http.MethodGet, url, + http.NoBody) + if err != nil { + t.Fatalf("new request: %v", err) + } + + req.Header.Set("X-Forwarded-For", from) + + transport := &http.Transport{} + defer transport.CloseIdleConnections() + + res, err := (&http.Client{Transport: transport}).Do(req) + if err != nil { + t.Fatalf("request: %v", err) + } + + _ = res.Body.Close() + + return res.StatusCode +} diff --git a/internal/state/state.go b/internal/state/state.go new file mode 100644 index 0000000..35551cb --- /dev/null +++ b/internal/state/state.go @@ -0,0 +1,698 @@ +// Package state keeps smallwebwaf's state in JSON files in +// SWWAF_STATE_DIR, as the "Persistent state" section of SPEC.md describes: +// bans.json holds the bans, clients.json each client's counters and +// history, and lookups.json GeoJS's answers. Load reads them at start, +// Watch takes in an admin's edit of one while smallwebwaf runs, and Run +// and WriteAll write them. The disk is read and written outside the +// parts' locks, which are held only to take a snapshot or to put in what +// a file holds, so that no request waits on the disk. +package state + +import ( + "bytes" + "context" + "crypto/sha256" + "encoding/json" + "errors" + "fmt" + "io/fs" + "log/slog" + "net/netip" + "os" + "path/filepath" + "sync" + "time" + + "github.com/fsnotify/fsnotify" + "sneak.berlin/go/smallwebwaf/internal/bans" + "sneak.berlin/go/smallwebwaf/internal/lookup" + "sneak.berlin/go/smallwebwaf/internal/metrics" + "sneak.berlin/go/smallwebwaf/internal/ratelimit" +) + +// version is the version of the files' format, the only one read. +const version = 1 + +// fileMode lets the smallwebwaf user alone read and write the files, which +// hold visitors' addresses. +const fileMode = 0o600 + +// The state files' names. +const ( + bansJSON = "bans.json" + clientsJSON = "clients.json" + lookupsJSON = "lookups.json" +) + +var ( + errVersion = errors.New("unknown version") + // errMissing is for an entry without a field it needs. + errMissing = errors.New("has no") +) + +// Params are what Load needs. +type Params struct { + // Dir is the directory of the state files (SWWAF_STATE_DIR). + Dir string + // WriteDelay is how long after a ban is made bans.json is written + // (SWWAF_STATE_WRITE_DELAY), and CounterInterval how often every file + // is (SWWAF_STATE_COUNTER_INTERVAL). + WriteDelay time.Duration + CounterInterval time.Duration + // Ledger, Limiter and GeoJS hold the state. + Ledger *bans.Ledger + Limiter *ratelimit.Limiter + GeoJS *lookup.GeoJS + // Now tells the time by which the counters' buckets run out, normally + // time.Now in UTC. + Now func() time.Time + // ProcessLog receives what was read and taken in, the edits set aside, + // and the writes that fail. + ProcessLog *slog.Logger + // Metrics count each file's writes, and the edits taken in and set + // aside. + Metrics *metrics.Metrics +} + +// Files are the state files of a running smallwebwaf. +type Files struct { + params Params + + // mu is held while a file is read for an edit, and while it is + // written, so that Watch and the writes take turns. No request takes + // it. + mu sync.Mutex + // sums are the SHA-256 sums of what each file held, by name, when + // smallwebwaf last read or wrote it. A file that holds anything else + // has been edited since. + sums map[string][sha256.Size]byte +} + +// bansFile is bans.json, indented for an admin to read and edit. +type bansFile struct { + Version int `json:"version"` + Bans []banEntry `json:"bans"` +} + +// banEntry is a ban as bans.json holds it: a permanent ban's expires is +// null. +type banEntry struct { + Netblock netip.Prefix `json:"netblock"` + Start time.Time `json:"start"` + Expires *time.Time `json:"expires"` + Notes bans.Notes `json:"notes"` +} + +// clientsFile is clients.json, with each client on a line of its own. +type clientsFile struct { + Version int `json:"version"` + Clients []ratelimit.Client `json:"clients"` +} + +// lookupsFile is lookups.json, with each answer on a line of its own. +type lookupsFile struct { + Version int `json:"version"` + Lookups []lookup.Answer `json:"lookups"` +} + +// stateFile is the struct of a state file. Once the file is decoded, its +// check refuses the first entry without a field it needs, which would +// otherwise be read as something the entry does not say. data is the +// file, for a field that may be null or "" but not left out, which the +// struct cannot tell apart. +type stateFile interface { + check(data []byte) error +} + +// Load checks that files can be written in Dir, and reads the state files +// in it into the ledger, the limiter and GeoJS. A missing file is empty +// state, as on a first start. A file that does not parse, has an unknown +// version, or has an entry without a field it needs, is an error that +// names the file and, where the JSON decoder tells it, the line and +// column, or else the entry. +func Load(params Params) (*Files, error) { + err := checkWritable(params.Dir) + if err != nil { + return nil, fmt.Errorf("SWWAF_STATE_DIR cannot be written: %w", err) + } + + f := &Files{params: params, sums: map[string][sha256.Size]byte{}} + + bansRead, bansErr := f.read(bansJSON) + clientsRead, clientsErr := f.read(clientsJSON) + lookupsRead, lookupsErr := f.read(lookupsJSON) + + err = errors.Join(bansErr, clientsErr, lookupsErr) + if err != nil { + return nil, err + } + + params.ProcessLog.Info("read the state files", "directory", params.Dir, + "bans", bansRead, "clients", clientsRead, "lookups", lookupsRead) + + return f, nil +} + +// Run writes bans.json WriteDelay after a ban is made, with every ban +// made in between, and every file every CounterInterval, until ctx is +// done. A write that fails is logged, and the file is written again at +// its next write. Each write takes in an admin's edit of its file first, +// as writeFile describes. +func (f *Files) Run(ctx context.Context) { + interval := time.NewTicker(f.params.CounterInterval) + defer interval.Stop() + + var bansDue <-chan time.Time // nil while no ban waits to be written + + for { + select { + case <-ctx.Done(): + return + case <-f.params.Ledger.Changed(): + if bansDue == nil { + bansDue = time.After(f.params.WriteDelay) + } + case <-bansDue: + bansDue = nil + + f.logFailure(f.writeFile(bansJSON)) + case <-interval.C: + f.logFailure(f.WriteAll()) + } + } +} + +// WriteAll writes every state file, as smallwebwaf stops. A file that +// fails does not keep the others from being written. +func (f *Files) WriteAll() error { + return errors.Join(f.writeFile(bansJSON), f.writeFile(clientsJSON), + f.writeFile(lookupsJSON)) +} + +// Watch watches Dir until ctx is done, and takes in an admin's edit of a +// state file as soon as it is saved: what the file holds replaces what +// smallwebwaf held for it. An edit that does not parse is left for the +// file's next write, which sets it aside, since a file can be read while +// an editor is still writing it. If Dir cannot be watched, that is +// logged, and an edit is taken in only before its file is written. +func (f *Files) Watch(ctx context.Context) { + watcher, err := fsnotify.NewWatcher() + if err == nil { + defer func() { + _ = watcher.Close() + }() + + err = watcher.Add(f.params.Dir) + } + + if err != nil { + f.params.ProcessLog.Error("cannot watch the state files for edits", + "error", err.Error()) + + return + } + + f.params.ProcessLog.Info("watching the state files for edits", + "directory", f.params.Dir) + + for { + select { + case <-ctx.Done(): + return + case event := <-watcher.Events: + switch name := filepath.Base(event.Name); name { + case bansJSON, clientsJSON, lookupsJSON: + f.fileChanged(name) + } + case err = <-watcher.Errors: + f.params.ProcessLog.Warn("watching the state files failed", + "error", err.Error()) + } + } +} + +// logFailure logs a write that failed. +func (f *Files) logFailure(err error) { + if err != nil { + f.params.ProcessLog.Error("writing the state files failed", + "error", err.Error()) + } +} + +// fileChanged takes in what the state file name holds, as Watch sees it +// change, if that is an edit made since smallwebwaf last read or wrote +// the file. A file that cannot be read or does not parse is left for its +// next write. +func (f *Files) fileChanged(name string) { + f.mu.Lock() + defer f.mu.Unlock() + + data, changed, err := f.readChanged(name) + if err != nil || !changed { + return + } + + _ = f.takeInEdit(name, data) +} + +// takeInEdit takes in data, an edit of the state file name, as takeIn +// does, and counts and logs it. Every edit taken in while smallwebwaf +// runs, by Watch or by a write, is taken in here. An edit that does not +// parse is neither counted nor logged, and takeIn's error returned. +func (f *Files) takeInEdit(name string, data []byte) error { + _, err := f.takeIn(name, data) + if err != nil { + return err + } + + // Counted before it is logged, so that the count is there once the + // log line is. + f.params.Metrics.StateFileEditTakenIn(name) + f.params.ProcessLog.Info("took in an edit of a state file", + "file", filepath.Join(f.params.Dir, name)) + + return nil +} + +// read takes in the state file name at start, and returns how many +// entries it holds. A missing file holds none. +func (f *Files) read(name string) (int, error) { + data, changed, err := f.readChanged(name) + if err != nil || !changed { + return 0, err + } + + return f.takeIn(name, data) +} + +// readChanged returns what the state file name holds, and whether that +// has changed since smallwebwaf last read or wrote the file, as it has +// for a file smallwebwaf never read or wrote. A missing file has not +// changed: it is written again at its next write. +func (f *Files) readChanged(name string) ([]byte, bool, error) { + path := filepath.Join(f.params.Dir, name) + + data, err := os.ReadFile(path) //nolint:gosec // a state file, in SWWAF_STATE_DIR + if errors.Is(err, fs.ErrNotExist) { + return nil, false, nil + } + + if err != nil { + return nil, false, err + } + + return data, sha256.Sum256(data) != f.sums[name], nil +} + +// takeIn parses data, what the state file name holds, puts it into the +// part that keeps that state, in place of what the part held, and returns +// how many entries the file holds. An error names the file and, where the +// JSON decoder tells it, the line and column, or else the entry. +func (f *Files) takeIn(name string, data []byte) (int, error) { + path := filepath.Join(f.params.Dir, name) + + var entries int + + switch name { + case bansJSON: + var file bansFile + + err := parse(path, data, &file) + if err != nil { + return 0, err + } + + held := make([]bans.Ban, 0, len(file.Bans)) + for _, entry := range file.Bans { + held = append(held, entry.ban()) + } + + f.params.Ledger.Load(held) + entries = len(held) + case clientsJSON: + var file clientsFile + + err := parse(path, data, &file) + if err != nil { + return 0, err + } + + f.params.Limiter.Load(file.Clients, f.params.Now()) + entries = len(file.Clients) + case lookupsJSON: + var file lookupsFile + + err := parse(path, data, &file) + if err != nil { + return 0, err + } + + f.params.GeoJS.Load(file.Lookups) + entries = len(file.Lookups) + } + + f.sums[name] = sha256.Sum256(data) + + return entries, nil +} + +// writeFile writes the state file name from what smallwebwaf holds. An +// edit made since smallwebwaf last read or wrote the file is taken in +// first, so that it is not overwritten, or set aside if it does not +// parse. A file that cannot be read, or an edit that cannot be set +// aside, is left as it is, and the write given up. Every write is counted +// in the metrics, and one that fails or is given up as a failure. +func (f *Files) writeFile(name string) error { + f.mu.Lock() + defer f.mu.Unlock() + + data, changed, err := f.readChanged(name) + if err == nil && changed { + err = f.takeInEdit(name, data) + if err != nil { + err = f.setAside(name, err) + } + } + + if err == nil { + data, err = f.encode(name) + if err != nil { + err = fmt.Errorf("encode %s: %w", name, err) + } + } + + if err == nil { + err = write(f.params.Dir, name, data) + } + + if err == nil { + // The file holds data from here on, even if the directory sync + // fails, so that its next read does not take it for an admin's + // edit. + f.sums[name] = sha256.Sum256(data) + err = syncDirectory(f.params.Dir) + } + + f.params.Metrics.StateFileWritten(name, len(data), err) + + return err +} + +// setAside renames the state file name, an edit that does not parse with +// parseErr, to name.bad, for the admin to mend, and logs it with where in +// the file the error is. If the rename fails, the edit is left as it is, +// and the error returned is parseErr joined with the rename's. +func (f *Files) setAside(name string, parseErr error) error { + path := filepath.Join(f.params.Dir, name) + + err := os.Rename(path, path+".bad") + if err != nil { + return errors.Join(parseErr, err) + } + + f.params.ProcessLog.Error("set aside an edit of a state file that does not parse", + "file", path+".bad", "error", parseErr.Error()) + f.params.Metrics.StateFileEditSetAside(name) + + return nil +} + +// encode returns the state file name as smallwebwaf writes it, from a +// snapshot of the part that keeps that state. +func (f *Files) encode(name string) ([]byte, error) { + switch name { + case bansJSON: + held := f.params.Ledger.Snapshot() + + file := bansFile{Version: version, Bans: make([]banEntry, 0, len(held))} + for _, ban := range held { + file.Bans = append(file.Bans, newBanEntry(ban)) + } + + data, err := json.MarshalIndent(file, "", " ") + if err != nil { + return nil, err + } + + return append(data, '\n'), nil + case clientsJSON: + return encodeOnePerLine("clients", f.params.Limiter.Snapshot()) + default: // lookups.json + return encodeOnePerLine("lookups", f.params.GeoJS.Snapshot()) + } +} + +// newBanEntry returns ban as bans.json holds it. +func newBanEntry(ban bans.Ban) banEntry { + entry := banEntry{Netblock: ban.Netblock, Start: ban.Start, Notes: ban.Notes} + if !ban.Permanent() { + entry.Expires = &ban.Expires + } + + return entry +} + +// ban returns the ban an entry of bans.json holds. +func (e banEntry) ban() bans.Ban { + ban := bans.Ban{Netblock: e.Netblock, Start: e.Start, Notes: e.Notes} + if e.Expires != nil { + ban.Expires = *e.Expires + } + + return ban +} + +// check refuses a ban without a netblock, which would refuse every IPv6 +// client, a start, from which the length of the netblock's next ban is +// worked out, or an expires, which would make it permanent. A permanent +// ban's expires is null, which Bans cannot tell from a missing one, so +// each expires is read again as written. +func (f *bansFile) check(data []byte) error { + var written struct { + Bans []struct { + Expires json.RawMessage `json:"expires"` + } `json:"bans"` + } + + err := json.Unmarshal(data, &written) + if err != nil { + return err + } + + for i, entry := range f.Bans { + switch { + case !entry.Netblock.IsValid(): + return missing(i, "netblock") + case entry.Start.IsZero(): + return missing(i, "start") + case written.Bans[i].Expires == nil: + return missing(i, "expires") + } + } + + return nil +} + +// check refuses a client without its address, which would count nobody's +// requests, or with requests in a window but no start, which would drop +// them and give the client a fresh allowance. +func (f *clientsFile) check([]byte) error { + for i, client := range f.Clients { + switch { + case !client.Client.IsValid(): + return missing(i, "client") + case countsWithoutStart(client.Minute): + return missing(i, "minute.start") + case countsWithoutStart(client.Hour): + return missing(i, "hour.start") + case countsWithoutStart(client.Day): + return missing(i, "day.start") + } + } + + return nil +} + +// check refuses an answer without a client, which would answer for +// nobody, a country, which would place the client nowhere, or the time +// GeoJS gave it, which would drop it. "" is the country of a client +// GeoJS cannot place, which Lookups cannot tell from a missing one, so +// each country is read again as written. +func (f *lookupsFile) check(data []byte) error { + var written struct { + Lookups []struct { + Country *string `json:"country"` + } `json:"lookups"` + } + + err := json.Unmarshal(data, &written) + if err != nil { + return err + } + + for i, answer := range f.Lookups { + switch { + case !answer.Client.IsValid(): + return missing(i, "client") + case written.Lookups[i].Country == nil: + return missing(i, "country") + case answer.Answered.IsZero(): + return missing(i, "answered") + } + } + + return nil +} + +// countsWithoutStart reports whether b holds requests but no start, which +// places them in time. +func countsWithoutStart(b ratelimit.Buckets) bool { + return b.Start.IsZero() && (b.Current != 0 || b.Previous != 0) +} + +// missing returns the error for entry i, counted from 0, of a state file, +// which has no field. +func missing(i int, field string) error { + return fmt.Errorf("entry %d %w %q", i+1, errMissing, field) +} + +// encodeOnePerLine encodes a state file whose entries, under key, are one +// to a line, so that grep shows everything about one client. +func encodeOnePerLine[E any](key string, entries []E) ([]byte, error) { + var b bytes.Buffer + + fmt.Fprintf(&b, "{\n \"version\": %d,\n %q: [", version, key) + + for i, entry := range entries { + line, err := json.Marshal(entry) + if err != nil { + return nil, err + } + + if i > 0 { + b.WriteString(",") + } + + b.WriteString("\n ") + b.Write(line) + } + + b.WriteString("\n ]\n}\n") + + return b.Bytes(), nil +} + +// checkWritable makes a file in dir and removes it again. +func checkWritable(dir string) error { + file, err := os.CreateTemp(dir, "write-check-*") + if err != nil { + return err + } + + return errors.Join(file.Close(), os.Remove(file.Name())) +} + +// parse reads data, what the state file at path holds, into file, a +// pointer to that file's struct, and checks its entries. +func parse(path string, data []byte, file stateFile) error { + // The version is read first, so that a file of another version is + // refused for that, and not for an entry this version cannot read. + var header struct { + Version int `json:"version"` + } + + err := json.Unmarshal(data, &header) + if err == nil && header.Version != version { + err = fmt.Errorf("%w %d, where this smallwebwaf reads version %d", + errVersion, header.Version, version) + } + + if err == nil { + decoder := json.NewDecoder(bytes.NewReader(data)) + // A field this version does not know is most likely misspelt, and + // its value would be lost without a word. + decoder.DisallowUnknownFields() + err = decoder.Decode(file) + } + + if err == nil { + err = file.check(data) + } + + if err != nil { + return fmt.Errorf("%s%s: %w", path, position(data, err), err) + } + + return nil +} + +// position returns where in data err was found, as ", line L, column C" +// of the last byte the JSON decoder read, or "" when err does not tell. +func position(data []byte, err error) string { + var ( + syntaxErr *json.SyntaxError + typeErr *json.UnmarshalTypeError + read int64 + ) + + switch { + case errors.As(err, &syntaxErr): + read = syntaxErr.Offset + case errors.As(err, &typeErr): + read = typeErr.Offset + default: + return "" + } + + before := data[:max(min(read, int64(len(data)))-1, 0)] + line := bytes.Count(before, []byte("\n")) + 1 + column := len(before) - bytes.LastIndexByte(before, '\n') + + return fmt.Sprintf(", line %d, column %d", line, column) +} + +// write writes data to the file name in dir so that a crash at any +// moment leaves either the old file or the new one, whole: data goes to a +// temporary file in the same directory, which is synced and renamed over +// name. syncDirectory must follow, so that the rename lasts. +func write(dir, name string, data []byte) error { + path := filepath.Join(dir, name) + temporary := path + ".tmp" + + err := writeSynced(temporary, data) + if err == nil { + err = os.Rename(temporary, path) + } + + if err != nil { + _ = os.Remove(temporary) + } + + return err +} + +// syncDirectory syncs dir to the disk, so that a rename in it lasts. +func syncDirectory(dir string) error { + directory, err := os.Open(dir) //nolint:gosec // SWWAF_STATE_DIR itself + if err != nil { + return err + } + + return errors.Join(directory.Sync(), directory.Close()) +} + +// writeSynced writes data to the file at path, and syncs it to the disk. +func writeSynced(path string, data []byte) error { + //nolint:gosec // a state file's temporary file, in SWWAF_STATE_DIR + file, err := os.OpenFile(path, os.O_WRONLY|os.O_CREATE|os.O_TRUNC, fileMode) + if err != nil { + return err + } + + _, err = file.Write(data) + if err == nil { + err = file.Sync() + } + + return errors.Join(err, file.Close()) +} diff --git a/internal/state/state_test.go b/internal/state/state_test.go new file mode 100644 index 0000000..b508b96 --- /dev/null +++ b/internal/state/state_test.go @@ -0,0 +1,1217 @@ +package state_test + +import ( + "context" + "encoding/json" + "io/fs" + "log/slog" + "maps" + "net" + "net/http" + "net/http/httptest" + "net/netip" + "os" + "path/filepath" + "slices" + "strconv" + "strings" + "testing" + "testing/synctest" + "time" + + "sneak.berlin/go/smallwebwaf/internal/bans" + "sneak.berlin/go/smallwebwaf/internal/lookup" + "sneak.berlin/go/smallwebwaf/internal/metrics" + "sneak.berlin/go/smallwebwaf/internal/ratelimit" + "sneak.berlin/go/smallwebwaf/internal/state" +) + +const ( + // The state files. + bansJSON = "bans.json" + clientsJSON = "clients.json" + lookupsJSON = "lookups.json" + // What the process log says once Watch watches the directory, and as + // it takes in an edit. + watching = "watching the state files for edits" + tookIn = "took in an edit of a state file" + // maxLogLines is how many lines of the process log wait for a test to + // read them. + maxLogLines = 64 +) + +// permanentBansJSON is bans.json holding permanentBan. +const permanentBansJSON = `{ + "version": 1, + "bans": [ + { + "netblock": "2001:db8::/64", + "start": "2026-10-06T00:00:00Z", + "expires": null, + "notes": { + "country": "DE", + "limit": 1000, + "window": "minute", + "count": 1000.5, + "request": { + "time": "2026-10-06T00:00:00Z", + "method": "GET", + "host": "app.example", + "path": "/repo?page=2", + "status": 403, + "user_agent": "scraper/1.0" + }, + "requests": 1500, + "refused": 3, + "earlier_bans": 5 + } + } + ] +} +` + +func TestFilesWrittenAndReadBack(t *testing.T) { + t.Parallel() + + dir := t.TempDir() + before := newParams(dir) + fill(before) + + files, err := state.Load(before) + if err != nil { + t.Fatalf("load: %v", err) + } + + err = files.WriteAll() + if err != nil { + t.Fatalf("write: %v", err) + } + + // Read into new parts, as at the next start, the files give back what + // was written. + after := newParams(dir) + load(t, after) + + wantEqual(t, bansJSON, after.Ledger.Snapshot(), before.Ledger.Snapshot()) + wantEqual(t, clientsJSON, after.Limiter.Snapshot(), before.Limiter.Snapshot()) + wantEqual(t, lookupsJSON, after.GeoJS.Snapshot(), before.GeoJS.Snapshot()) + + // Each one-per-line file lists its entries by client, and nothing + // but the three files is left in the directory. + wantEntries(t, filepath.Join(dir, clientsJSON), "clients", + "192.0.2.1/32", "203.0.113.9/32", "2001:db8::/64") + wantEntries(t, filepath.Join(dir, lookupsJSON), "lookups", + "192.0.2.1/32", "203.0.113.9/32") + wantFiles(t, dir, bansJSON, clientsJSON, lookupsJSON) +} + +func TestBansJSONIsIndentedWithNullForAPermanentBan(t *testing.T) { + t.Parallel() + + dir := t.TempDir() + params := newParams(dir) + params.Ledger.Load([]bans.Ban{permanentBan()}) + + files := load(t, params) + + err := files.WriteAll() + if err != nil { + t.Fatalf("write: %v", err) + } + + got := readFile(t, filepath.Join(dir, bansJSON)) + if got != permanentBansJSON { + t.Errorf("bans.json\n%s\nwant\n%s", got, permanentBansJSON) + } +} + +func TestMissingFilesAreEmptyState(t *testing.T) { + t.Parallel() + + params := newParams(t.TempDir()) + load(t, params) + + if len(params.Ledger.Snapshot()) != 0 || len(params.Limiter.Snapshot()) != 0 || + len(params.GeoJS.Snapshot()) != 0 { + t.Error("state from no files") + } +} + +func TestFileThatDoesNotParseStopsTheStart(t *testing.T) { + t.Parallel() + + for _, tc := range []struct { + name, file, content string + // want is what the error says after the file's path. + want string + }{ + { + "a syntax error", bansJSON, + "{\n \"version\": 1,\n \"bans\": [\n" + + " {\"netblock\": \"203.0.113.9/32\",}\n ]\n}\n", + ", line 4, column 39: invalid character '}'", + }, + { + "a value of the wrong kind", clientsJSON, + "{\n \"version\": 1,\n \"clients\": [\n" + + " {\"client\":\"203.0.113.9/32\",\"history\":{\"requests\":\"many\"}}\n" + + " ]\n}\n", + ", line 4, column ", + }, + { + // Found at the newline that ends the file. + "a cut-off file", lookupsJSON, + "{\n \"version\": 1,\n \"lookups\": [\n", + ", line 3, column 17: unexpected end of JSON input", + }, + { + "an unknown field", lookupsJSON, + `{"version": 1, "lookups": [{"client": "203.0.113.9/32", "contry": "DE"}]}`, + `: json: unknown field "contry"`, + }, + { + "a netblock that does not read", bansJSON, + `{"version": 1, "bans": [{"netblock": "203.0.113.300/32"}]}`, + `: netip.ParsePrefix("203.0.113.300/32")`, + }, + } { + t.Run(tc.name, func(t *testing.T) { + t.Parallel() + + wantRefused(t, tc.file, tc.content, tc.want) + }) + } +} + +func TestEntryWithoutAFieldItNeedsStopsTheStart(t *testing.T) { + t.Parallel() + + const ( + // The other fields each entry needs. + ban = `"start": "2026-10-06T00:00:00Z", "expires": null` + answer = `"answered": "2026-10-06T00:00:00Z"` + + noNetblock = `: entry 1 has no "netblock"` + ) + + for _, tc := range []struct { + name, file, content string + // want is what the error says after the file's path. + want string + }{ + { + "a ban without a netblock", bansJSON, + `{"version": 1, "bans": [{` + ban + `}]}`, + noNetblock, + }, + { + "a ban whose netblock is null", bansJSON, + `{"version": 1, "bans": [{"netblock": null, ` + ban + `}]}`, + noNetblock, + }, + { + "a ban whose netblock is empty", bansJSON, + `{"version": 1, "bans": [{"netblock": "", ` + ban + `}]}`, + noNetblock, + }, + { + "a ban without a start", bansJSON, + `{"version": 1, "bans": [{"netblock": "203.0.113.9/32", "expires": null}]}`, + `: entry 1 has no "start"`, + }, + { + // The first ban's expires is null, as a permanent ban's is. + "a ban without an expires", bansJSON, + `{"version": 1, "bans": [{"netblock": "203.0.113.9/32", ` + ban + `}, ` + + `{"netblock": "203.0.113.10/32", "start": "2026-10-06T00:00:00Z"}]}`, + `: entry 2 has no "expires"`, + }, + { + "a client without its address", clientsJSON, + `{"version": 1, "clients": [{"history": {"requests": 3}}]}`, + `: entry 1 has no "client"`, + }, + { + "a client with requests in a window without its start", clientsJSON, + `{"version": 1, "clients": [{"client": "203.0.113.9/32", ` + + `"hour": {"current": 3}}]}`, + `: entry 1 has no "hour.start"`, + }, + { + "an answer without a client", lookupsJSON, + `{"version": 1, "lookups": [{"country": "DE", ` + answer + `}]}`, + `: entry 1 has no "client"`, + }, + { + // A country of "" is a client GeoJS cannot place. + "an answer without a country", lookupsJSON, + `{"version": 1, "lookups": [{"client": "192.0.2.1/32", "country": "", ` + + answer + `}, {"client": "203.0.113.9/32", ` + answer + `}]}`, + `: entry 2 has no "country"`, + }, + { + "an answer without the time GeoJS gave it", lookupsJSON, + `{"version": 1, "lookups": [{"client": "203.0.113.9/32", "country": "DE"}]}`, + `: entry 1 has no "answered"`, + }, + } { + t.Run(tc.name, func(t *testing.T) { + t.Parallel() + + wantRefused(t, tc.file, tc.content, tc.want) + }) + } +} + +func TestUnknownVersionStopsTheStart(t *testing.T) { + t.Parallel() + + for _, file := range []string{bansJSON, clientsJSON, lookupsJSON} { + for _, content := range []string{`{"version": 2}`, `{}`} { + t.Run(file+" "+content, func(t *testing.T) { + t.Parallel() + + wantRefused(t, file, content, ": unknown version ") + }) + } + } +} + +func TestUnwritableDirectoryStopsTheStart(t *testing.T) { + t.Parallel() + + notADirectory := filepath.Join(t.TempDir(), "file") + + err := os.WriteFile(notADirectory, nil, 0o600) + if err != nil { + t.Fatalf("write: %v", err) + } + + for _, dir := range []string{ + filepath.Join(t.TempDir(), "missing"), + notADirectory, + } { + const want = "SWWAF_STATE_DIR cannot be written: " + + _, err := state.Load(newParams(dir)) + if err == nil || !strings.HasPrefix(err.Error(), want) { + t.Errorf("state directory %s: error %v, want one starting %s", dir, err, want) + } + } +} + +// The three tests below run Run in a synctest bubble, where time is a clock +// of the test's own: time.Sleep moves it on at once, and synctest.Wait +// returns once Run waits for its next write, so that every write due by +// then is on disk. + +func TestBansWrittenOnceWriteDelayAfterABan(t *testing.T) { + t.Parallel() + + synctest.Test(t, func(t *testing.T) { + dir := t.TempDir() + params := newParams(dir) + params.WriteDelay = 10 * time.Second + run(t, load(t, params).Run) + + // A second ban, made while the first waits to be written, puts the + // write off no further, and is written with it. + first := params.Ledger.BanForLimit(netip.MustParsePrefix("203.0.113.9/32"), + midnight(), bans.Notes{}) + + time.Sleep(5 * time.Second) + + second := params.Ledger.BanForLimit(netip.MustParsePrefix("203.0.113.10/32"), + midnight(), bans.Notes{}) + + time.Sleep(5*time.Second - time.Nanosecond) + synctest.Wait() + wantFiles(t, dir) + + time.Sleep(time.Nanosecond) + synctest.Wait() + wantFiles(t, dir, bansJSON) + + read := newParams(dir) + load(t, read) + + want := []bans.Ban{first, second} + if got := read.Ledger.Snapshot(); !slices.Equal(got, want) { + t.Errorf("bans.json holds %+v, want %+v", got, want) + } + + // That write was the only one: bans.json is not written again for + // the second ban. The other files wait for the interval, an hour + // away. + removeFiles(t, dir, bansJSON) + time.Sleep(params.WriteDelay) + synctest.Wait() + wantFiles(t, dir) + }) +} + +func TestEveryFileWrittenEveryCounterInterval(t *testing.T) { + t.Parallel() + + synctest.Test(t, func(t *testing.T) { + dir := t.TempDir() + params := newParams(dir) + params.CounterInterval = time.Minute + run(t, load(t, params).Run) + + // The files are removed once written, so that each interval shows + // them written again. + for range 3 { + time.Sleep(time.Minute - time.Nanosecond) + synctest.Wait() + wantFiles(t, dir) + + time.Sleep(time.Nanosecond) + synctest.Wait() + wantFiles(t, dir, bansJSON, clientsJSON, lookupsJSON) + removeFiles(t, dir, bansJSON, clientsJSON, lookupsJSON) + } + }) +} + +func TestEditJustBeforeAScheduledWriteSurvivesIt(t *testing.T) { + t.Parallel() + + synctest.Test(t, func(t *testing.T) { + dir := t.TempDir() + params := newParams(dir) + params.WriteDelay = 10 * time.Second + run(t, load(t, params).Run) + + // A ban, and bans.json written with it. + params.Ledger.BanForLimit(netip.MustParsePrefix("203.0.113.9/32"), + midnight(), bans.Notes{}) + time.Sleep(params.WriteDelay) + synctest.Wait() + + // A second ban is to be written WriteDelay later. Just before + // then, an admin saves bans.json with the first ban lifted and + // another added. + params.Ledger.BanForLimit(netip.MustParsePrefix("203.0.113.10/32"), + midnight(), bans.Notes{}) + time.Sleep(params.WriteDelay - time.Nanosecond) + synctest.Wait() + edit(t, dir, bansJSON, permanentBansJSON) + + // The write takes the edit in first, and writes it back. The second + // ban, made after the admin opened the file, is lost, as "Edits + // while running" in SPEC.md says. + time.Sleep(time.Nanosecond) + synctest.Wait() + + if got := readFile(t, filepath.Join(dir, bansJSON)); got != permanentBansJSON { + t.Errorf("bans.json holds\n%s\nwant the edit", got) + } + + wantEqual(t, bansJSON, params.Ledger.Snapshot(), []bans.Ban{permanentBan()}) + }) +} + +func TestFailedWriteLeavesTheFileAsItWas(t *testing.T) { + t.Parallel() + + dir := t.TempDir() + params := newParams(dir) + params.Ledger.Load([]bans.Ban{permanentBan()}) + files := load(t, params) + + err := files.WriteAll() + if err != nil { + t.Fatalf("write: %v", err) + } + + // A directory in the way of bans.json's temporary file fails its + // next write, but not the others'. + err = os.Mkdir(filepath.Join(dir, bansJSON+".tmp"), 0o700) + if err != nil { + t.Fatalf("mkdir: %v", err) + } + + params.Ledger.BanForLimit(netip.MustParsePrefix("203.0.113.9/32"), midnight(), + bans.Notes{}) + params.Limiter.Count(netip.MustParsePrefix("203.0.113.9/32"), midnight()) + + err = files.WriteAll() + if err == nil || !strings.Contains(err.Error(), bansJSON+".tmp") { + t.Errorf("error %v, want one naming bans.json's temporary file", err) + } + + got := readFile(t, filepath.Join(dir, bansJSON)) + if got != permanentBansJSON { + t.Errorf("bans.json is now\n%s\nwant it as it was", got) + } + + read := newParams(dir) + load(t, read) + + if len(read.Limiter.Snapshot()) != 1 { + t.Error("clients.json was not written") + } +} + +func TestWritesAreCountedInTheMetrics(t *testing.T) { + t.Parallel() + + dir := t.TempDir() + params := newParams(dir) + params.Ledger.Load([]bans.Ban{permanentBan()}) + files := load(t, params) + + err := files.WriteAll() + if err != nil { + t.Fatalf("write: %v", err) + } + + const ( + ofBans = `{file="bans.json"}` + ofClients = `{file="clients.json"}` + ) + + got := scrape(t, params) + wantMetric(t, got, "smallwebwaf_state_file_writes_total"+ofBans, 1) + wantMetric(t, got, "smallwebwaf_state_file_writes_total"+ofClients, 1) + wantMetric(t, got, "smallwebwaf_state_file_write_failures_total"+ofBans, 0) + wantMetric(t, got, "smallwebwaf_state_file_size_bytes"+ofBans, + float64(len(permanentBansJSON))) + + written := metric(t, got, "smallwebwaf_state_file_last_write_timestamp_seconds"+ofBans) + if written < float64(time.Now().Add(-time.Hour).Unix()) { + t.Errorf("bans.json was last written at %v, not by that write", written) + } + + // A directory in the way of bans.json's temporary file fails its next + // write, which leaves its size as it was, although it has a ban more. + err = os.Mkdir(filepath.Join(dir, bansJSON+".tmp"), 0o700) + if err != nil { + t.Fatalf("mkdir: %v", err) + } + + params.Ledger.BanForLimit(netip.MustParsePrefix("203.0.113.9/32"), midnight(), + bans.Notes{}) + + err = files.WriteAll() + if err == nil { + t.Fatal("the write did not fail") + } + + got = scrape(t, params) + wantMetric(t, got, "smallwebwaf_state_file_writes_total"+ofBans, 2) + wantMetric(t, got, "smallwebwaf_state_file_write_failures_total"+ofBans, 1) + wantMetric(t, got, "smallwebwaf_state_file_write_failures_total"+ofClients, 0) + wantMetric(t, got, "smallwebwaf_state_file_size_bytes"+ofBans, + float64(len(permanentBansJSON))) +} + +func TestFileThatCannotBeReadIsNotWrittenOver(t *testing.T) { + t.Parallel() + + dir := t.TempDir() + path := filepath.Join(dir, bansJSON) + params := newParams(dir) + files := load(t, params) + + // bans.json is a socket, which cannot be opened as a file, even by + // root, as the tests run in Docker, but which a rename could replace. + // Whether it holds an edit cannot be told, so it is left as it is. + socket, err := (&net.ListenConfig{}).Listen(t.Context(), "unix", path) + if err != nil { + t.Fatalf("listen: %v", err) + } + + defer func() { + _ = socket.Close() + }() + + err = files.WriteAll() + if err == nil { + t.Error("writing with bans.json unreadable did not fail") + } + + info, err := os.Lstat(path) + if err != nil || info.Mode().Type() != fs.ModeSocket { + t.Errorf("bans.json is now %v (%v), want the socket", info, err) + } + + wantFiles(t, dir, bansJSON, clientsJSON, lookupsJSON) + wantWriteFailed(t, params, bansJSON) +} + +func TestBrokenEditThatCannotBeSetAsideIsNotWrittenOver(t *testing.T) { + t.Parallel() + + const broken = `{"version": 1, "bans": [` + + dir := t.TempDir() + path := filepath.Join(dir, bansJSON) + params := newParams(dir) + files := load(t, params) + + // A directory named bans.json.bad cannot be renamed over, so the + // broken edit cannot be set aside, and is left as it is. + edit(t, dir, bansJSON, broken) + + err := os.Mkdir(path+".bad", 0o700) + if err != nil { + t.Fatalf("mkdir: %v", err) + } + + err = files.WriteAll() + if err == nil { + t.Error("writing with bans.json.bad in the way did not fail") + } + + if got := readFile(t, path); got != broken { + t.Errorf("bans.json holds\n%s\nwant the edit", got) + } + + wantWriteFailed(t, params, bansJSON) +} + +func TestEditOfEachFileTakenIn(t *testing.T) { + t.Parallel() + + dir := t.TempDir() + params := newParams(dir) + lines := logInto(¶ms) + fill(params) + files := load(t, params) + + err := files.WriteAll() + if err != nil { + t.Fatalf("write: %v", err) + } + + watch(t, files, lines) + + // Each edit holds one entry, for a client the parts did not hold, and + // takes the place of everything the part held. + client := netip.MustParsePrefix("198.51.100.7/32") + + edit(t, dir, bansJSON, `{"version": 1, "bans": [{"netblock": "198.51.100.7/32", `+ + `"start": "2026-10-06T00:00:00Z", "expires": null}]}`) + wantTakenIn(t, lines, dir, bansJSON) + wantEqual(t, bansJSON, params.Ledger.Snapshot(), + []bans.Ban{{Netblock: client, Start: midnight()}}) + + edit(t, dir, clientsJSON, `{"version": 1, "clients": [`+ + `{"client": "198.51.100.7/32", "history": {"requests": 7}}]}`) + wantTakenIn(t, lines, dir, clientsJSON) + wantEqual(t, clientsJSON, params.Limiter.Snapshot(), + []ratelimit.Client{{Client: client, History: ratelimit.History{Requests: 7}}}) + + edit(t, dir, lookupsJSON, `{"version": 1, "lookups": [{"client": "198.51.100.7/32", `+ + `"country": "FR", "answered": "2026-10-06T00:00:00Z"}]}`) + wantTakenIn(t, lines, dir, lookupsJSON) + wantEqual(t, lookupsJSON, params.GeoJS.Snapshot(), + []lookup.Answer{{Client: client, Country: "FR", Answered: midnight()}}) +} + +func TestOwnWritesAreNotTakenIn(t *testing.T) { + t.Parallel() + + dir := t.TempDir() + params := newParams(dir) + lines := logInto(¶ms) + fill(params) + files := load(t, params) + watch(t, files, lines) + + // Every file is written while watched, and then lookups.json edited: + // the first edit taken in is that one. + err := files.WriteAll() + if err != nil { + t.Fatalf("write: %v", err) + } + + edit(t, dir, lookupsJSON, `{"version": 1, "lookups": []}`) + wantTakenIn(t, lines, dir, lookupsJSON) +} + +func TestFileRenamedOverAStateFileTakenIn(t *testing.T) { + t.Parallel() + + dir := t.TempDir() + path := filepath.Join(dir, bansJSON) + params := newParams(dir) + lines := logInto(¶ms) + files := load(t, params) + + err := files.WriteAll() + if err != nil { + t.Fatalf("write: %v", err) + } + + watch(t, files, lines) + + // The admin mends bans.json.bad and moves it back, as editors that + // save by renaming do with a file of their own: nothing is written + // into bans.json itself. An edit of clients.json after it must be + // taken in second. + edit(t, dir, bansJSON+".bad", permanentBansJSON) + + err = os.Rename(path+".bad", path) + if err != nil { + t.Fatalf("rename: %v", err) + } + + edit(t, dir, clientsJSON, `{"version": 1, "clients": []}`) + + wantTakenIn(t, lines, dir, bansJSON) + wantEqual(t, bansJSON, params.Ledger.Snapshot(), []bans.Ban{permanentBan()}) +} + +func TestBanAddedAndLiftedThroughBansJSON(t *testing.T) { + t.Parallel() + + dir := t.TempDir() + params := newParams(dir) + lines := logInto(¶ms) + watch(t, load(t, params), lines) + + client := netip.MustParseAddr("203.0.113.9") + + // An entry added, as an admin writes it, bans its netblock. + edit(t, dir, bansJSON, `{"version": 1, "bans": [{"netblock": "203.0.113.0/24", `+ + `"start": "2026-10-06T00:00:00Z", "expires": null}]}`) + wantTakenIn(t, lines, dir, bansJSON) + + _, banned := params.Ledger.Check(client, midnight()) + if !banned { + t.Error("the ban added to bans.json does not refuse") + } + + // The entry removed lifts the ban. + edit(t, dir, bansJSON, `{"version": 1, "bans": []}`) + wantTakenIn(t, lines, dir, bansJSON) + + _, banned = params.Ledger.Check(client, midnight()) + if banned { + t.Error("the ban removed from bans.json still refuses") + } +} + +func TestBrokenEditSetAsideAtTheNextWrite(t *testing.T) { + t.Parallel() + + // It ends a ban's entry with a comma. + const broken = "{\n \"version\": 1,\n \"bans\": [\n" + + " {\"netblock\": \"203.0.113.9/32\",}\n ]\n}\n" + + dir := t.TempDir() + path := filepath.Join(dir, bansJSON) + params := newParams(dir) + lines := logInto(¶ms) + params.Ledger.Load([]bans.Ban{permanentBan()}) + files := load(t, params) + + err := files.WriteAll() + if err != nil { + t.Fatalf("write: %v", err) + } + + watch(t, files, lines) + + // While smallwebwaf runs, the broken edit is left as it is: an edit + // of clients.json, made after it and taken in, shows that it has been + // seen. + edit(t, dir, bansJSON, broken) + edit(t, dir, clientsJSON, `{"version": 1, "clients": []}`) + wantTakenIn(t, lines, dir, clientsJSON) + wantFiles(t, dir, bansJSON, clientsJSON, lookupsJSON) + + // The next write sets it aside, logged with where the error is, and + // writes bans.json again from what smallwebwaf still holds. + err = files.WriteAll() + if err != nil { + t.Fatalf("write: %v", err) + } + + line := lines.waitFor(t, "set aside an edit of a state file that does not parse") + message, _ := line["error"].(string) + + if line["file"] != path+".bad" || + !strings.HasPrefix(message, path+", line 4, column 39: ") { + t.Errorf("set aside with %v", line) + } + + wantFiles(t, dir, bansJSON, bansJSON+".bad", clientsJSON, lookupsJSON) + + if got := readFile(t, path+".bad"); got != broken { + t.Errorf("bans.json.bad holds\n%s\nwant the edit", got) + } + + if got := readFile(t, path); got != permanentBansJSON { + t.Errorf("bans.json holds\n%s\nwant\n%s", got, permanentBansJSON) + } +} + +func TestEditsTakenInAreCountedInTheMetrics(t *testing.T) { + t.Parallel() + + dir := t.TempDir() + params := newParams(dir) + lines := logInto(¶ms) + files := load(t, params) + + // One edit is taken in by the write of its file, before Watch runs, + // and one by Watch. + edit(t, dir, bansJSON, `{"version": 1, "bans": []}`) + + err := files.WriteAll() + if err != nil { + t.Fatalf("write: %v", err) + } + + watch(t, files, lines) + edit(t, dir, bansJSON, permanentBansJSON) + wantTakenIn(t, lines, dir, bansJSON) + + wantMetric(t, scrape(t, params), + `smallwebwaf_state_file_edits_taken_in_total{file="bans.json"}`, 2) +} + +func TestEditTakenInByAWriteIsLoggedAsWatchLogsIt(t *testing.T) { + t.Parallel() + + dir := t.TempDir() + params := newParams(dir) + lines := logInto(¶ms) + files := load(t, params) + + // An edit taken in by Watch, which is then stopped. + ctx, stop := context.WithCancel(t.Context()) + stopped := make(chan struct{}) + + go func() { + files.Watch(ctx) + close(stopped) + }() + + lines.waitFor(t, watching) + edit(t, dir, bansJSON, `{"version": 1, "bans": []}`) + byWatch := lines.waitFor(t, tookIn) + + stop() + <-stopped + + // An edit taken in by the write of its file. Nothing logs after the + // write, so the log is closed, and a write that does not log the edit + // fails the test at once instead of waiting for the line. + edit(t, dir, bansJSON, permanentBansJSON) + + err := files.WriteAll() + if err != nil { + t.Fatalf("write: %v", err) + } + + close(lines) + + byWrite := lines.waitFor(t, tookIn) + + // The two lines differ only in their time. + delete(byWatch, "time") + delete(byWrite, "time") + + if !maps.Equal(byWrite, byWatch) { + t.Errorf("the write logged %v, where Watch logged %v", byWrite, byWatch) + } +} + +func TestEditsSetAsideAreCountedInTheMetrics(t *testing.T) { + t.Parallel() + + dir := t.TempDir() + params := newParams(dir) + files := load(t, params) + + edit(t, dir, bansJSON, `{"version": 1, "bans": [`) + + err := files.WriteAll() + if err != nil { + t.Fatalf("write: %v", err) + } + + wantMetric(t, scrape(t, params), + `smallwebwaf_state_file_edits_set_aside_total{file="bans.json"}`, 1) +} + +func TestDirectoryThatCannotBeWatchedIsLogged(t *testing.T) { + t.Parallel() + + dir := t.TempDir() + params := newParams(dir) + lines := logInto(¶ms) + files := load(t, params) + + err := os.Remove(dir) + if err != nil { + t.Fatalf("remove: %v", err) + } + + // Watch returns at once. + files.Watch(t.Context()) + + line := lines.waitFor(t, "cannot watch the state files for edits") + if line["level"] != "ERROR" { + t.Errorf("logged as %v", line) + } +} + +// midnight is the time of the tests' clock. +func midnight() time.Time { + return time.Date(2026, 10, 6, 0, 0, 0, 0, time.UTC) +} + +// newParams returns Params for the state files in dir, with parts that +// hold nothing yet. GeoJS is never asked. +func newParams(dir string) state.Params { + discard := slog.New(slog.DiscardHandler) + m := metrics.New(1) + + return state.Params{ + Dir: dir, + WriteDelay: time.Hour, + CounterInterval: time.Hour, + Ledger: bans.New(bans.Rules{ + LimitBanDuration: time.Hour, + LimitBanRepeatWindow: 24 * time.Hour, + MaxBanDuration: 7 * 24 * time.Hour, + MaxBans: 5000, + }), + Limiter: ratelimit.New(ratelimit.Limits{}), + GeoJS: lookup.New(lookup.Params{ + Now: midnight, ProcessLog: discard, Metrics: m, + }), + Now: midnight, + ProcessLog: discard, + Metrics: m, + } +} + +// fill puts a ban that ends and one that does not, clients with counts +// and histories, and GeoJS answers into the parts of params. +func fill(params state.Params) { + now := midnight() + client := netip.MustParsePrefix("203.0.113.9/32") + + params.Ledger.Load([]bans.Ban{permanentBan()}) + params.Ledger.BanForLimit(client, now, bans.Notes{Country: "DE", Limit: 1}) + + for _, c := range []string{"2001:db8::/64", "203.0.113.9/32", "192.0.2.1/32"} { + params.Limiter.Count(netip.MustParsePrefix(c), now) + } + + params.Limiter.AddToHistory(client, now, ratelimit.Request{ + Country: "DE", Forwarded: true, Status: 200, RequestBytes: 3, ResponseBytes: 5, + }) + + params.GeoJS.Load([]lookup.Answer{ + {Client: client, Country: "DE", Answered: now.Add(-time.Hour), Used: now}, + { + Client: netip.MustParsePrefix("192.0.2.1/32"), + Answered: now.Add(-time.Hour), Used: now.Add(-time.Minute), + }, + }) +} + +// permanentBan is the ban permanentBansJSON holds. +func permanentBan() bans.Ban { + return bans.Ban{ + Netblock: netip.MustParsePrefix("2001:db8::/64"), + Start: midnight(), + Notes: bans.Notes{ + Country: "DE", + Limit: 1000, + Window: "minute", + Count: 1000.5, + Request: bans.Request{ + Time: midnight(), + Method: "GET", + Host: "app.example", + Path: "/repo?page=2", + Status: 403, + UserAgent: "scraper/1.0", + }, + Requests: 1500, + Refused: 3, + EarlierBans: 5, + }, + } +} + +// load reads the state files into the parts of params. +func load(t *testing.T, params state.Params) *state.Files { + t.Helper() + + files, err := state.Load(params) + if err != nil { + t.Fatalf("load: %v", err) + } + + return files +} + +// run runs task, the Run or the Watch of state files, until the test +// ends. +func run(t *testing.T, task func(context.Context)) { + t.Helper() + + ctx, stop := context.WithCancel(t.Context()) + stopped := make(chan struct{}) + + go func() { + task(ctx) + close(stopped) + }() + + t.Cleanup(func() { + stop() + <-stopped + }) +} + +// watch runs files' Watch until the test ends, and waits until it +// watches the directory. +func watch(t *testing.T, files *state.Files, lines processLog) { + t.Helper() + + run(t, files.Watch) + lines.waitFor(t, watching) +} + +// processLog receives the lines of a process log, each a JSON object, for +// a test to wait for. +type processLog chan string + +// logInto has params' process log write its lines into a new processLog, +// and returns that. +func logInto(params *state.Params) processLog { + lines := make(processLog, maxLogLines) + params.ProcessLog = slog.New(slog.NewJSONHandler(lines, nil)) + + return lines +} + +// Write receives a line of the process log. +func (l processLog) Write(line []byte) (int, error) { + l <- string(line) + + return len(line), nil +} + +// waitFor returns the next line of the process log whose message is msg, +// passing over the lines before it, or nil if the log is closed first. It +// waits as long as that takes, so that a slow test process cannot fail +// the test. +func (l processLog) waitFor(t *testing.T, msg string) map[string]any { + t.Helper() + + for line := range l { + var fields map[string]any + + err := json.Unmarshal([]byte(line), &fields) + if err != nil { + t.Fatalf("process log line %q is not JSON: %v", line, err) + } + + if fields["msg"] == msg { + return fields + } + } + + return nil +} + +// wantTakenIn waits for the next edit taken in, and checks that it is of +// the state file name in dir. +func wantTakenIn(t *testing.T, lines processLog, dir, name string) { + t.Helper() + + line := lines.waitFor(t, tookIn) + if line["file"] != filepath.Join(dir, name) { + t.Fatalf("took in %v, want an edit of %s", line, name) + } +} + +// edit writes content to the state file name in dir, as an admin saves an +// edit of it. +func edit(t *testing.T, dir, name, content string) { + t.Helper() + + err := os.WriteFile(filepath.Join(dir, name), []byte(content), 0o600) + if err != nil { + t.Fatalf("write %s: %v", name, err) + } +} + +// wantEqual checks that the entries read back from file are those +// written. +func wantEqual[E comparable](t *testing.T, file string, got, want []E) { + t.Helper() + + if !slices.Equal(got, want) { + t.Errorf("%s read back\n%+v\nwant\n%+v", file, got, want) + } +} + +// readFile returns what the file at path holds. +func readFile(t *testing.T, path string) string { + t.Helper() + + data, err := os.ReadFile(path) //nolint:gosec // a file the test wrote + if err != nil { + t.Fatalf("read: %v", err) + } + + return string(data) +} + +// wantRefused writes content to the state file named file in a new +// directory, and checks that Load refuses it with an error that is the +// file's path and then starts with want. +func wantRefused(t *testing.T, file, content, want string) { + t.Helper() + + dir := t.TempDir() + path := filepath.Join(dir, file) + + err := os.WriteFile(path, []byte(content), 0o600) + if err != nil { + t.Fatalf("write %s: %v", file, err) + } + + _, err = state.Load(newParams(dir)) + if err == nil || !strings.HasPrefix(err.Error(), path+want) { + t.Errorf("error %v, want one starting %s%s", err, path, want) + } +} + +// removeFiles removes the named files from dir. +func removeFiles(t *testing.T, dir string, names ...string) { + t.Helper() + + for _, name := range names { + err := os.Remove(filepath.Join(dir, name)) + if err != nil { + t.Fatalf("remove: %v", err) + } + } +} + +// wantFiles checks the names of the files in dir. +func wantFiles(t *testing.T, dir string, want ...string) { + t.Helper() + + entries, err := os.ReadDir(dir) + if err != nil { + t.Fatalf("read %s: %v", dir, err) + } + + got := make([]string, 0, len(entries)) + for _, entry := range entries { + got = append(got, entry.Name()) + } + + if !slices.Equal(got, want) { + t.Errorf("%s holds %v, want %v", dir, got, want) + } +} + +// wantEntries checks that the file at path has its version, then its +// entries under key, each on a line of its own, for the clients want +// names in that order. +func wantEntries(t *testing.T, path, key string, want ...string) { + t.Helper() + + data := readFile(t, path) + lines := strings.Split(strings.TrimSuffix(data, "\n"), "\n") + head := []string{"{", ` "version": 1,`, ` "` + key + `": [`} + tail := []string{" ]", "}"} + + if len(lines) != len(head)+len(want)+len(tail) || + !slices.Equal(lines[:len(head)], head) || + !slices.Equal(lines[len(lines)-len(tail):], tail) { + t.Fatalf("%s is\n%s", path, data) + } + + for i, client := range want { + line := strings.TrimSuffix(lines[len(head)+i], ",") + + var entry struct { + Client string `json:"client"` + } + + err := json.Unmarshal([]byte(line), &entry) + if err != nil || entry.Client != client { + t.Errorf("entry %d of %s is %s (%v), want %s's", i, path, line, err, client) + } + } +} + +// scrape returns the metrics of params, in the Prometheus text format. +func scrape(t *testing.T, params state.Params) string { + t.Helper() + + recorder := httptest.NewRecorder() + params.Metrics.ServeHTTP(recorder, + httptest.NewRequestWithContext(t.Context(), http.MethodGet, "/", http.NoBody)) + + if recorder.Code != http.StatusOK { + t.Fatalf("the metrics were answered %d", recorder.Code) + } + + return recorder.Body.String() +} + +// metric returns the value of series in text, the metrics, such as +// smallwebwaf_state_file_writes_total{file="bans.json"}, or fails the test +// if there is no such series. +func metric(t *testing.T, text, series string) float64 { + t.Helper() + + for line := range strings.Lines(text) { + value, found := strings.CutPrefix(strings.TrimSuffix(line, "\n"), series+" ") + if !found { + continue + } + + number, err := strconv.ParseFloat(value, 64) + if err != nil { + t.Fatalf("%s has the value %q", series, value) + } + + return number + } + + t.Fatalf("no series %s in the metrics:\n%s", series, text) + + return 0 +} + +// wantWriteFailed checks that the metrics of params count one write of the +// state file name, and that it failed. +func wantWriteFailed(t *testing.T, params state.Params, name string) { + t.Helper() + + got := scrape(t, params) + file := `{file="` + name + `"}` + + wantMetric(t, got, "smallwebwaf_state_file_writes_total"+file, 1) + wantMetric(t, got, "smallwebwaf_state_file_write_failures_total"+file, 1) +} + +// wantMetric checks the value of series in text, the metrics, as metric +// reads it. +func wantMetric(t *testing.T, text, series string, want float64) { + t.Helper() + + got := metric(t, text, series) + if got != want { + t.Errorf("%s is %v, want %v", series, got, want) + } +} diff --git a/internal/state/write_internal_test.go b/internal/state/write_internal_test.go new file mode 100644 index 0000000..698fb3b --- /dev/null +++ b/internal/state/write_internal_test.go @@ -0,0 +1,35 @@ +package state + +import ( + "os" + "path/filepath" + "testing" +) + +// The test is on write itself: a state file is read before it is +// written, and a directory in its place fails that read first. +func TestFailedRenameLeavesNoTemporaryFile(t *testing.T) { + t.Parallel() + + dir := t.TempDir() + + // A directory named bans.json cannot be renamed over. + err := os.Mkdir(filepath.Join(dir, bansJSON), 0o700) + if err != nil { + t.Fatalf("mkdir: %v", err) + } + + err = write(dir, bansJSON, []byte("{}\n")) + if err == nil { + t.Error("writing over a directory did not fail") + } + + entries, err := os.ReadDir(dir) + if err != nil { + t.Fatalf("read %s: %v", dir, err) + } + + if len(entries) != 1 || entries[0].Name() != bansJSON { + t.Errorf("%s holds %v, want only bans.json", dir, entries) + } +} diff --git a/package.json b/package.json new file mode 100644 index 0000000..d846c06 --- /dev/null +++ b/package.json @@ -0,0 +1,6 @@ +{ + "license": "MIT", + "devDependencies": { + "prettier": "3.8.1" + } +} diff --git a/script/bootstrap b/script/bootstrap new file mode 100755 index 0000000..967bf12 --- /dev/null +++ b/script/bootstrap @@ -0,0 +1,144 @@ +#!/bin/sh +# script/bootstrap: install all dependencies needed to build and develop +# this repo. Idempotent: every install is guarded by a check so already +# installed tools are skipped. Base tooling comes from nix, apt, brew, +# or apk (detected in that order); assumes nothing is present. Node is +# used directly if installed; otherwise it is installed at a pinned +# version via nvm (installing nvm itself first, from a hash-verified +# release archive, never curl | sh). Go comes from the package manager, +# for gofmt in script/fmt and script/fmt-check: the tests and the linter +# run in docker and need no Go on the host. +set -eu + +ROOT="$(cd "$(dirname "$0")/.." && pwd -P)" + +# Pinned versions, 2026-07-06 +NODE_VERSION="22.17.0" +NVM_VERSION="0.40.3" +# sha256 of https://github.com/nvm-sh/nvm/archive/refs/tags/v0.40.3.tar.gz +NVM_SHA256="5f4d6aaa04a177dc93c985e31dbc411ab6b8c6e1e21d8015dbc1372625fcd1d0" +YARN_VERSION="1.22.22" + +PKGMGR="" +SUDO="" + +detect_pkgmgr() { + [ -n "$PKGMGR" ] && return 0 + if command -v nix-env >/dev/null 2>&1; then + PKGMGR="nix" + elif command -v apt-get >/dev/null 2>&1; then + PKGMGR="apt" + elif command -v brew >/dev/null 2>&1; then + PKGMGR="brew" + elif command -v apk >/dev/null 2>&1; then + PKGMGR="apk" + else + echo "bootstrap: no supported package manager (nix, apt, brew, apk)" >&2 + exit 1 + fi + if [ "$PKGMGR" = "apt" ]; then + export DEBIAN_FRONTEND=noninteractive + if [ "$(id -u)" != "0" ]; then + SUDO="sudo" + fi + fi +} + +# pkg_install +pkg_install() { + detect_pkgmgr + case "$PKGMGR" in + nix) nix-env -iA "nixpkgs.$1" ;; + apt) + # A fresh CI runner may carry no package lists at all. + $SUDO env DEBIAN_FRONTEND=noninteractive apt-get update -q + $SUDO env DEBIAN_FRONTEND=noninteractive apt-get install -y "$2" + ;; + brew) brew install "$3" ;; + apk) apk add --no-cache "$4" ;; + esac +} + +missing() { + ! command -v "$1" >/dev/null 2>&1 +} + +# verify_sha256 +verify_sha256() { + if command -v sha256sum >/dev/null 2>&1; then + actual="$(sha256sum "$1" | cut -d' ' -f1)" + else + actual="$(shasum -a 256 "$1" | cut -d' ' -f1)" + fi + if [ "$actual" != "$2" ]; then + echo "bootstrap: sha256 mismatch for $1" >&2 + echo " expected: $2" >&2 + echo " actual: $actual" >&2 + exit 1 + fi +} + +# nvm is a bash script; run a command in a bash with nvm loaded +nvm_sh() { + bash -c ". \"\$HOME/.nvm/nvm.sh\" && $*" +} + +ensure_nvm() { + [ -s "$HOME/.nvm/nvm.sh" ] && return 0 + # nvm prerequisites; nvm itself requires bash + if missing bash; then pkg_install bash bash bash bash; fi + if missing curl; then pkg_install curl curl curl curl; fi + if missing git; then pkg_install git git git git; fi + tmp="$(mktemp -d)" + curl -fsSL -o "$tmp/nvm.tar.gz" \ + "https://github.com/nvm-sh/nvm/archive/refs/tags/v${NVM_VERSION}.tar.gz" + verify_sha256 "$tmp/nvm.tar.gz" "$NVM_SHA256" + mkdir -p "$HOME/.nvm" + tar -xzf "$tmp/nvm.tar.gz" -C "$HOME/.nvm" --strip-components=1 + rm -rf "$tmp" +} + +ensure_node() { + if ! missing node; then return 0; fi + ensure_nvm + nvm_sh "nvm install $NODE_VERSION" +} + +ensure_yarn() { + if ! missing yarn; then return 0; fi + if ! missing corepack; then + corepack enable + corepack prepare "yarn@$YARN_VERSION" --activate + elif [ -s "$HOME/.nvm/nvm.sh" ]; then + nvm_sh "nvm use $NODE_VERSION >/dev/null && corepack enable && \ + corepack prepare yarn@$YARN_VERSION --activate" + else + npm install -g "yarn@$YARN_VERSION" + fi +} + +install_js_deps() { + if missing yarn && [ -s "$HOME/.nvm/nvm.sh" ]; then + nvm_sh "nvm use $NODE_VERSION >/dev/null && cd \"$ROOT\" && \ + yarn install --frozen-lockfile" + else + yarn install --frozen-lockfile + fi +} + +main() { + cd "$ROOT" + + if missing make; then pkg_install gnumake make make make; fi + if missing git; then pkg_install git git git git; fi + if missing curl; then pkg_install curl curl curl curl; fi + if missing gofmt; then pkg_install go golang go go; fi + + ensure_node + ensure_yarn + install_js_deps + + echo "bootstrap complete" +} + +main "$@" diff --git a/script/build b/script/build new file mode 100755 index 0000000..94bd69c --- /dev/null +++ b/script/build @@ -0,0 +1,18 @@ +#!/bin/sh +# script/build: build bin/smallwebwaf on the host, with Go installed, for +# working on the code by hand. The version it reports comes from git, as +# in script/docker. +set -eu + +SCRIPT_DIR="$(cd "$(dirname "$0")" && pwd -P)" +ROOT="$(cd "$SCRIPT_DIR/.." && pwd -P)" + +main() { + cd "$ROOT" + version="$(git describe --tags --always --dirty 2>/dev/null || true)" + [ -n "$version" ] || version="unknown" + go build -trimpath -ldflags "-X main.Version=$version" \ + -o bin/smallwebwaf ./cmd/smallwebwaf +} + +main "$@" diff --git a/script/check b/script/check new file mode 100755 index 0000000..92875f7 --- /dev/null +++ b/script/check @@ -0,0 +1,16 @@ +#!/bin/sh +# script/check: run all checks (test, lint, fmt-check). Our own +# extension to scripts-to-rule-them-all. test and lint are Docker +# phases; fmt-check is native, because a formatter writes the working +# tree. Must not modify any files. +set -eu + +SCRIPT_DIR="$(cd "$(dirname "$0")" && pwd -P)" + +main() { + "$SCRIPT_DIR/test" + "$SCRIPT_DIR/lint" + "$SCRIPT_DIR/fmt-check" +} + +main "$@" diff --git a/script/cibuild b/script/cibuild new file mode 100755 index 0000000..d8d3200 --- /dev/null +++ b/script/cibuild @@ -0,0 +1,28 @@ +#!/bin/sh +# script/cibuild: run the CI build. It bootstraps first: a CI runner +# checks out and runs this and nothing else, and script/fmt-check runs +# the formatter on the host, which a pristine checkout cannot do. +# --no-cache for the same reason as script/docker: the gate phases the +# final stage depends on are RUN steps, and a cached one is a check that +# did not run. +set -eu + +SCRIPT_DIR="$(cd "$(dirname "$0")" && pwd -P)" +ROOT="$(cd "$SCRIPT_DIR/.." && pwd -P)" + +main() { + cd "$ROOT" + "$SCRIPT_DIR/bootstrap" + "$SCRIPT_DIR/check" + # Own line: a failing command substitution inside an argument does + # not trip `set -e`, so the inline form degrades silently to an + # empty constant. The VERSION build argument takes precedence over + # the version a build stage derives from the .git in the context. + version="$(git describe --tags --always --dirty 2>/dev/null || true)" + [ -n "$version" ] || version="unknown" + docker build --no-cache \ + --build-arg VERSION="$version" \ + -t "$("$SCRIPT_DIR/projectname")" . +} + +main "$@" diff --git a/script/docker b/script/docker new file mode 100755 index 0000000..07b626c --- /dev/null +++ b/script/docker @@ -0,0 +1,24 @@ +#!/bin/sh +# script/docker: build the Docker image tagged with the project name. +# Identical in all repos; the tag comes from script/projectname. +# --no-cache because the gate phases the final stage depends on are RUN +# steps, and a cached one is a check that did not run. +set -eu + +SCRIPT_DIR="$(cd "$(dirname "$0")" && pwd -P)" +ROOT="$(cd "$SCRIPT_DIR/.." && pwd -P)" + +main() { + cd "$ROOT" + # Own line: a failing command substitution inside an argument does + # not trip `set -e`, so the inline form degrades silently to an + # empty constant. The VERSION build argument takes precedence over + # the version a build stage derives from the .git in the context. + version="$(git describe --tags --always --dirty 2>/dev/null || true)" + [ -n "$version" ] || version="unknown" + docker build --no-cache \ + --build-arg VERSION="$version" \ + -t "$("$SCRIPT_DIR/projectname")" . +} + +main "$@" diff --git a/script/example-app b/script/example-app new file mode 100755 index 0000000..b07236b --- /dev/null +++ b/script/example-app @@ -0,0 +1,121 @@ +#!/bin/sh +# script/example-app: build the image, and on it the example app in +# deploy/example-app, then run the app's container with a volume for the +# state files and check that the health check passes, that a request is +# served through smallwebwaf, that a second one in a minute bans the +# client, that `sv stop` stops smallwebwaf in order, that `docker stop` +# stops the container without having to kill it, and that a new +# container on the same volume still refuses the banned client. The +# containers, the volume and both images are removed however the script +# ends. Building the app needs network access, for nixpkgs' binary cache. +# script/check does not run this. +set -eu + +SCRIPT_DIR="$(cd "$(dirname "$0")" && pwd -P)" +ROOT="$(cd "$SCRIPT_DIR/.." && pwd -P)" + +# Named after this run, so that runs in other clones on the same host +# never touch each other's. +NAME="$("$SCRIPT_DIR/projectname")-example-$$" +IMAGE="$NAME-base" +APP_IMAGE="$NAME-app" +CONTAINER="$NAME" +VOLUME="$NAME-state" + +cleanup() { + docker rm --force "$CONTAINER" >/dev/null 2>&1 || true + docker volume rm --force "$VOLUME" >/dev/null 2>&1 || true + docker rmi --force "$APP_IMAGE" "$IMAGE" >/dev/null 2>&1 || true +} + +fail() { + echo "example-app: $*; the container's output:" >&2 + docker logs "$CONTAINER" >&2 || true + exit 1 +} + +# wait_for ...: run the command every second until +# it succeeds, for at most a minute. +wait_for() { + failure="$1" + shift + tries=0 + until "$@"; do + tries=$((tries + 1)) + [ "$tries" -lt 60 ] || fail "$failure" + sleep 1 + done +} + +healthy() { + status="$(docker inspect --format '{{.State.Health.Status}}' "$CONTAINER")" + [ "$status" = healthy ] +} + +# logged : the container's output holds text. +logged() { + docker logs "$CONTAINER" 2>&1 | grep -qF "$1" +} + +# start_container: run the app's container, with the state files on the +# volume and a rate limit of one request a minute, and wait until it is +# healthy. +start_container() { + docker run --detach --name "$CONTAINER" --publish 127.0.0.1::8080 \ + --volume "$VOLUME:/var/lib/smallwebwaf" \ + --env SWWAF_RATE_LIMIT_PER_MINUTE=1 \ + "$APP_IMAGE" >/dev/null + wait_for "the health check did not pass" healthy + address="$(docker port "$CONTAINER" 8080/tcp)" +} + +# refused: a request to the container gets 403, SWWAF_BAN_RESPONSE's +# default. +refused() { + code="$(curl --silent --output /dev/null --write-out '%{http_code}' \ + --max-time 10 "http://$address/")" || true + [ "$code" = 403 ] +} + +main() { + cd "$ROOT" + trap cleanup EXIT + trap 'exit 1' HUP INT TERM + + docker build --no-cache -t "$IMAGE" . + docker build --no-cache --build-arg SMALLWEBWAF_IMAGE="$IMAGE" \ + -t "$APP_IMAGE" deploy/example-app + + docker volume create "$VOLUME" >/dev/null + start_container + echo "example-app: the health check passes" + + page="$(curl --fail --silent --show-error --max-time 10 "http://$address/")" || + fail "no answer on port 8080" + [ "$page" = "hello from the example app" ] || fail "port 8080 answered $page" + wait_for "smallwebwaf logged no request it forwarded" logged '"action":"forward"' + echo "example-app: smallwebwaf passes a request to the app and its answer back" + + refused || fail "a second request in a minute was not refused" + wait_for "smallwebwaf logged no ban" logged '"action":"rate_limited"' + echo "example-app: a second request in a minute bans the client" + + docker exec "$CONTAINER" sv stop smallwebwaf >/dev/null || + fail "sv stop smallwebwaf failed" + wait_for "smallwebwaf did not stop in order" logged '"msg":"stopped"' + echo "example-app: sv stop stops smallwebwaf in order" + + docker stop "$CONTAINER" >/dev/null + status="$(docker inspect --format '{{.State.ExitCode}}' "$CONTAINER")" + [ "$status" = 0 ] || fail "docker stop left exit status $status" + echo "example-app: docker stop stops the container in order" + + docker rm "$CONTAINER" >/dev/null + start_container + refused || fail "the new container let the banned client through" + wait_for "smallwebwaf logged no request refused under the ban" \ + logged '"action":"banned"' + echo "example-app: a new container on the same volume keeps the ban" +} + +main "$@" diff --git a/script/fmt b/script/fmt new file mode 100755 index 0000000..9febb30 --- /dev/null +++ b/script/fmt @@ -0,0 +1,33 @@ +#!/bin/sh +# script/fmt: format all files (writes): the Go code with gofmt, the +# Markdown with prettier. +set -eu + +ROOT="$(cd "$(dirname "$0")/.." && pwd -P)" + +# Must match the pin in script/bootstrap. +NODE_VERSION="22.17.0" + +# script/bootstrap installs node and yarn under nvm and leaves neither +# on the PATH of the shell that called it, so resolve the pinned +# toolchain here the way bootstrap's own install step does. nvm is a +# bash script, hence the subshell. +run_yarn() { + if command -v yarn >/dev/null 2>&1; then + exec yarn "$@" + fi + if [ ! -s "$HOME/.nvm/nvm.sh" ]; then + echo "fmt: no yarn; run script/bootstrap first" >&2 + exit 1 + fi + exec bash -c '. "$HOME/.nvm/nvm.sh" && nvm use "$1" >/dev/null && + shift && exec yarn "$@"' bash "$NODE_VERSION" "$@" +} + +main() { + cd "$ROOT" + gofmt -w cmd internal + run_yarn run prettier --write '**/*.md' --tab-width 4 --prose-wrap always +} + +main "$@" diff --git a/script/fmt-check b/script/fmt-check new file mode 100755 index 0000000..52ec452 --- /dev/null +++ b/script/fmt-check @@ -0,0 +1,38 @@ +#!/bin/sh +# script/fmt-check: check formatting (read-only): the Go code with +# gofmt, the Markdown with prettier. +set -eu + +ROOT="$(cd "$(dirname "$0")/.." && pwd -P)" + +# Must match the pin in script/bootstrap. +NODE_VERSION="22.17.0" + +# script/bootstrap installs node and yarn under nvm and leaves neither +# on the PATH of the shell that called it, so resolve the pinned +# toolchain here the way bootstrap's own install step does. nvm is a +# bash script, hence the subshell. +run_yarn() { + if command -v yarn >/dev/null 2>&1; then + exec yarn "$@" + fi + if [ ! -s "$HOME/.nvm/nvm.sh" ]; then + echo "fmt-check: no yarn; run script/bootstrap first" >&2 + exit 1 + fi + exec bash -c '. "$HOME/.nvm/nvm.sh" && nvm use "$1" >/dev/null && + shift && exec yarn "$@"' bash "$NODE_VERSION" "$@" +} + +main() { + cd "$ROOT" + unformatted="$(gofmt -l cmd internal)" + if [ -n "$unformatted" ]; then + echo "fmt-check: gofmt would change:" >&2 + echo "$unformatted" >&2 + exit 1 + fi + run_yarn run prettier --check '**/*.md' --tab-width 4 --prose-wrap always +} + +main "$@" diff --git a/script/install-precommit b/script/install-precommit new file mode 100755 index 0000000..bef6406 --- /dev/null +++ b/script/install-precommit @@ -0,0 +1,16 @@ +#!/bin/sh +# script/install-precommit: install the git pre-commit hook that runs +# script/precommit. Our own extension to scripts-to-rule-them-all. +set -eu + +ROOT="$(cd "$(dirname "$0")/.." && pwd -P)" + +main() { + cd "$ROOT" + hook=".git/hooks/pre-commit" + printf '#!/bin/sh\nset -e\nscript/precommit\n' > .git/hooks/pre-commit + chmod +x .git/hooks/pre-commit + echo "pre-commit hook installed: runs script/precommit" +} + +main "$@" diff --git a/script/lint b/script/lint new file mode 100755 index 0000000..2d8b075 --- /dev/null +++ b/script/lint @@ -0,0 +1,23 @@ +#!/bin/sh +# script/lint: run the linter. Linting is a phase of the Dockerfile and +# this builds that phase alone; the linter is never installed or run on +# a developer host, where a shared result cache and a host-global lock +# make its answer untrustworthy. +# +# The phase is not the last stage in the file, so it is built only when +# --target names it. --no-cache because a cached lint layer is a lint +# that did not run. The tag makes each build replace the previous image +# instead of leaving a dangling one behind. +set -eu + +SCRIPT_DIR="$(cd "$(dirname "$0")" && pwd -P)" +ROOT="$(cd "$SCRIPT_DIR/.." && pwd -P)" + +main() { + cd "$ROOT" + docker build --no-cache \ + --target lint \ + -t "$("$SCRIPT_DIR/projectname")-lint" . +} + +main "$@" diff --git a/script/precommit b/script/precommit new file mode 100755 index 0000000..c0a7867 --- /dev/null +++ b/script/precommit @@ -0,0 +1,12 @@ +#!/bin/sh +# script/precommit: run by the git pre-commit hook; fails the commit if +# checks fail. Our own extension to scripts-to-rule-them-all. +set -eu + +SCRIPT_DIR="$(cd "$(dirname "$0")" && pwd -P)" + +main() { + "$SCRIPT_DIR/check" +} + +main "$@" diff --git a/script/projectname b/script/projectname new file mode 100755 index 0000000..a822b81 --- /dev/null +++ b/script/projectname @@ -0,0 +1,12 @@ +#!/bin/sh +# script/projectname: output the name of this project. Our own +# extension to scripts-to-rule-them-all. Other scripts that need the +# name (e.g. script/docker) call this, so they can stay identical +# across all repos. +set -eu + +main() { + echo "smallwebwaf" +} + +main "$@" diff --git a/script/run b/script/run new file mode 100755 index 0000000..b0cd063 --- /dev/null +++ b/script/run @@ -0,0 +1,20 @@ +#!/bin/sh +# script/run: build bin/smallwebwaf with script/build and run it, with +# the settings in the environment. Unless SWWAF_STATE_DIR is set, the +# state files go in bin/state, beside the binary. +set -eu + +SCRIPT_DIR="$(cd "$(dirname "$0")" && pwd -P)" +ROOT="$(cd "$SCRIPT_DIR/.." && pwd -P)" + +main() { + "$SCRIPT_DIR/build" + if [ -z "${SWWAF_STATE_DIR+set}" ]; then + SWWAF_STATE_DIR="$ROOT/bin/state" + export SWWAF_STATE_DIR + mkdir -p "$SWWAF_STATE_DIR" + fi + exec "$ROOT/bin/smallwebwaf" +} + +main "$@" diff --git a/script/setup b/script/setup new file mode 100755 index 0000000..4cc5b6b --- /dev/null +++ b/script/setup @@ -0,0 +1,13 @@ +#!/bin/sh +# script/setup: set up the repo for development after a fresh clone: +# installs dependencies and the git pre-commit hook. +set -eu + +SCRIPT_DIR="$(cd "$(dirname "$0")" && pwd -P)" + +main() { + "$SCRIPT_DIR/bootstrap" + "$SCRIPT_DIR/install-precommit" +} + +main "$@" diff --git a/script/test b/script/test new file mode 100755 index 0000000..cd239f2 --- /dev/null +++ b/script/test @@ -0,0 +1,19 @@ +#!/bin/sh +# script/test: run the test suite. Testing is a phase of the Dockerfile +# and this builds that phase alone, on the same terms as script/lint: +# --target because a phase that is not the last stage is built only when +# named, --no-cache because a cached test layer is a test that did not +# run, and a tag so each build replaces the previous image. +set -eu + +SCRIPT_DIR="$(cd "$(dirname "$0")" && pwd -P)" +ROOT="$(cd "$SCRIPT_DIR/.." && pwd -P)" + +main() { + cd "$ROOT" + docker build --no-cache \ + --target test \ + -t "$("$SCRIPT_DIR/projectname")-test" . +} + +main "$@" diff --git a/share/smallwebwaf.run b/share/smallwebwaf.run new file mode 100755 index 0000000..a13a17b --- /dev/null +++ b/share/smallwebwaf.run @@ -0,0 +1,16 @@ +#!/usr/bin/env bash +set -euo pipefail + +# runit's run script for smallwebwaf, run again whenever smallwebwaf +# exits; the wait spaces out the restarts. The state directory and every +# file in it are given to the smallwebwaf user, so that a volume mounted +# there needs no change of owner; chown -R changes a symbolic link itself, +# never what it points to. exec, so that the signal `sv stop` sends +# reaches smallwebwaf itself. +main() { + sleep 1 + chown -R smallwebwaf:smallwebwaf "${SWWAF_STATE_DIR:-/var/lib/smallwebwaf}" + exec chpst -u smallwebwaf:smallwebwaf /usr/local/bin/smallwebwaf +} + +main "$@" diff --git a/yarn.lock b/yarn.lock new file mode 100644 index 0000000..d846639 --- /dev/null +++ b/yarn.lock @@ -0,0 +1,8 @@ +# THIS IS AN AUTOGENERATED FILE. DO NOT EDIT THIS FILE DIRECTLY. +# yarn lockfile v1 + + +prettier@3.8.1: + version "3.8.1" + resolved "https://registry.yarnpkg.com/prettier/-/prettier-3.8.1.tgz#edf48977cf991558f4fcbd8a3ba6015ba2a3a173" + integrity sha512-UOnG6LftzbdaHZcKoPFtOcCKztrQ57WkHDeRD9t/PTQtmT0NHSeWWepj6pS0z/N7+08BHFDQVUrfmfMRcZwbMg==