Compare commits
1
Commits
next
..
23df4310f8
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
23df4310f8 |
+6
-20
@@ -13,21 +13,9 @@
|
|||||||
# `/myapp`, never `**/myapp`, which also matches `cmd/myapp/` and
|
# `/myapp`, never `**/myapp`, which also matches `cmd/myapp/` and
|
||||||
# deletes the package directory from the context.
|
# deletes the package directory from the context.
|
||||||
|
|
||||||
# .git is sent without its config. Without a VERSION build argument the
|
# Excluding .git means `git describe` cannot run in any build stage and
|
||||||
# stage that compiles runs `git describe --tags --always` on .git, which
|
# fails quietly there; pass the version in with --build-arg VERSION.
|
||||||
# does not need .git/config; that file can hold a credential, such as a
|
.git
|
||||||
# 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.
|
# Agent scratch: one full checkout of the repo per in-flight agent.
|
||||||
# Anchored because it occurs once where agents run at the repo root.
|
# Anchored because it occurs once where agents run at the repo root.
|
||||||
@@ -51,13 +39,14 @@
|
|||||||
**/[iI][dD]_[rR][sS][aA]
|
**/[iI][dD]_[rR][sS][aA]
|
||||||
**/[iI][dD]_[dD][sS][aA]
|
**/[iI][dD]_[dD][sS][aA]
|
||||||
**/[iI][dD]_[eE][cC][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
|
||||||
**/[iI][dD]_[eE][dD]25519_[sS][kK]
|
|
||||||
|
|
||||||
# Dependencies: restored inside the image, never copied in.
|
# Dependencies: restored inside the image, never copied in.
|
||||||
**/node_modules
|
**/node_modules
|
||||||
|
|
||||||
|
# The binary `make build` writes on the host; the image builds its own.
|
||||||
|
/bin
|
||||||
|
|
||||||
# OS metadata.
|
# OS metadata.
|
||||||
**/.DS_Store
|
**/.DS_Store
|
||||||
**/Thumbs.db
|
**/Thumbs.db
|
||||||
@@ -70,6 +59,3 @@
|
|||||||
**/.idea
|
**/.idea
|
||||||
**/.vscode
|
**/.vscode
|
||||||
**/*.sublime-*
|
**/*.sublime-*
|
||||||
|
|
||||||
# The binary `make build` writes on the host; the image builds its own.
|
|
||||||
/bin
|
|
||||||
|
|||||||
+5
-25
@@ -20,31 +20,11 @@ Thumbs.db
|
|||||||
# Node
|
# Node
|
||||||
node_modules/
|
node_modules/
|
||||||
|
|
||||||
# Secrets. Unanchored like every entry above, so each matches at every
|
# Environment / secrets
|
||||||
# depth. Matching is case-sensitive on Linux, so names use character
|
.env
|
||||||
# ranges rather than a lowercase form that misses `Server.Key`.
|
.env.*
|
||||||
|
*.pem
|
||||||
# Environment files. `*.env` covers bare `.env` and the `prod.env`
|
*.key
|
||||||
# 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
|
# Go: the binary `make build` writes, test binaries, profiles and logs
|
||||||
/bin/
|
/bin/
|
||||||
|
|||||||
@@ -17,7 +17,6 @@ linters:
|
|||||||
disable:
|
disable:
|
||||||
# Genuinely incompatible with project patterns
|
# Genuinely incompatible with project patterns
|
||||||
- exhaustruct # Requires all struct fields
|
- exhaustruct # Requires all struct fields
|
||||||
- exhaustruct_v5 # Requires all struct fields (successor to exhaustruct)
|
|
||||||
- godot # Requires comments to end with periods
|
- godot # Requires comments to end with periods
|
||||||
- wrapcheck # Too verbose for internal packages
|
- wrapcheck # Too verbose for internal packages
|
||||||
- varnamelen # Short names like db, id are idiomatic Go
|
- varnamelen # Short names like db, id are idiomatic Go
|
||||||
@@ -61,10 +60,6 @@ linters:
|
|||||||
desc: >-
|
desc: >-
|
||||||
Test-support code belongs in test files and in packages whose
|
Test-support code belongs in test files and in packages whose
|
||||||
directory name ends in test, not in the shipped binary.
|
directory name ends in test, not in the shipped binary.
|
||||||
- pkg: sneak.berlin/go/smallwebwaf/internal/lookup/lookuptest
|
|
||||||
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
|
# Only decisions already recorded in the Go package defaults are
|
||||||
# listed here. Every entry matches the module path exactly.
|
# listed here. Every entry matches the module path exactly.
|
||||||
gomodguard_v2:
|
gomodguard_v2:
|
||||||
|
|||||||
+19
-140
@@ -2,8 +2,8 @@
|
|||||||
# lint` or `script/lint`, which are themselves a docker build and would
|
# lint` or `script/lint`, which are themselves a docker build and would
|
||||||
# recurse into a daemon that does not exist in a build step.
|
# 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
|
# golangci/golangci-lint v2.12.2 (built with go1.26.2), 2026-05-06
|
||||||
FROM golangci/golangci-lint@sha256:ad862ba6b3798cbe0fd9fd7408d498fd74fbd2623a92406b2fd3898faf0bf98f AS lint
|
FROM golangci/golangci-lint@sha256:5cceeef04e53efe1470638d4b4b4f5ceefd574955ab3941b2d9a68a8c9ad5240 AS lint
|
||||||
|
|
||||||
WORKDIR /src
|
WORKDIR /src
|
||||||
|
|
||||||
@@ -29,27 +29,24 @@ RUN go mod download
|
|||||||
|
|
||||||
COPY . .
|
COPY . .
|
||||||
|
|
||||||
# Go's build cache is kept on a tmpfs, out of the image: nothing uses it
|
RUN go test -count=1 -timeout 90s -race -cover ./... || \
|
||||||
# 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 ---"; \
|
{ echo "--- Rerunning with -v for details ---"; \
|
||||||
go test -timeout 90s -race -v ./...; exit 1; }
|
go test -count=1 -timeout 90s -race -v ./...; exit 1; }
|
||||||
|
|
||||||
# Build stage. Nothing is wanted from the two phases above; the copies
|
# Build stage, and the last one: a plain `docker build .` names no
|
||||||
# are what make BuildKit build them first, so the image, which needs this
|
# target and so builds this one. Nothing is wanted from the two phases
|
||||||
# stage, cannot be produced unless lint and test passed.
|
# above; the copies are what make BuildKit build them first, so this
|
||||||
|
# image cannot be produced unless lint and test passed. The image an
|
||||||
|
# app's Dockerfile builds FROM comes with milestone 2
|
||||||
|
# (https://git.eeqj.de/sneak/smallwebwaf/issues/12); until then this
|
||||||
|
# stage builds the binary and can run it.
|
||||||
#
|
#
|
||||||
# golang 1.27.1-trixie, 2026-09-19
|
# golang 1.27.1-trixie, 2026-09-19
|
||||||
FROM golang@sha256:3b77fc618ec235a1ab412de7737f120dd507c57e8d87de4cbb7994fb94275ed5 AS builder
|
FROM golang@sha256:3b77fc618ec235a1ab412de7737f120dd507c57e8d87de4cbb7994fb94275ed5
|
||||||
|
|
||||||
COPY --from=lint /src/go.sum /dev/null
|
COPY --from=lint /src/go.sum /dev/null
|
||||||
COPY --from=test /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
|
WORKDIR /src
|
||||||
|
|
||||||
COPY go.mod go.sum ./
|
COPY go.mod go.sum ./
|
||||||
@@ -57,130 +54,12 @@ RUN go mod download
|
|||||||
|
|
||||||
COPY . .
|
COPY . .
|
||||||
|
|
||||||
# The VERSION build arg when one is given, otherwise
|
# The version is computed on the host and passed in, because
|
||||||
# `git describe --tags --always` on the .git in the build context. With
|
# .dockerignore excludes .git.
|
||||||
# .git present, a version that is still empty, dev or unknown fails the
|
ARG VERSION=dev
|
||||||
# build: git is missing or could not read the checkout.
|
RUN CGO_ENABLED=0 go build -trimpath \
|
||||||
ARG VERSION
|
-ldflags="-s -w -X main.Version=${VERSION}" \
|
||||||
RUN VERSION="${VERSION:-$(git describe --tags --always)}"; \
|
-o /usr/local/bin/smallwebwaf ./cmd/smallwebwaf
|
||||||
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.<name>`. 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
|
|
||||||
|
|
||||||
# The default rule file, in SWWAF_RULES_DIR by default, where an app's
|
|
||||||
# Dockerfile can copy rule files of its own beside it.
|
|
||||||
COPY share/rules.d/00-default.rules /etc/smallwebwaf/rules.d/00-default.rules
|
|
||||||
|
|
||||||
# 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
|
EXPOSE 8080
|
||||||
|
ENTRYPOINT ["/usr/local/bin/smallwebwaf"]
|
||||||
# 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"]
|
|
||||||
|
|||||||
@@ -1,10 +1,8 @@
|
|||||||
.PHONY: bootstrap setup test lint fmt fmt-check check docker hooks build run \
|
.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/
|
# Makefile targets are thin shims; the implementations live in script/
|
||||||
# per the scripts-to-rule-them-all pattern (see the Entrypoints section
|
# 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;
|
# 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:
|
bootstrap:
|
||||||
@script/bootstrap
|
@script/bootstrap
|
||||||
@@ -38,6 +36,3 @@ build:
|
|||||||
|
|
||||||
run:
|
run:
|
||||||
@script/run
|
@script/run
|
||||||
|
|
||||||
example-app:
|
|
||||||
@script/example-app
|
|
||||||
|
|||||||
+44
-120
@@ -1,6 +1,6 @@
|
|||||||
---
|
---
|
||||||
title: Repository Policies
|
title: Repository Policies
|
||||||
last_modified: 2026-10-04
|
last_modified: 2026-09-08
|
||||||
---
|
---
|
||||||
|
|
||||||
This document covers repository structure, tooling, and workflow standards. Code
|
This document covers repository structure, tooling, and workflow standards. Code
|
||||||
@@ -104,14 +104,10 @@ style conventions are in separate documents:
|
|||||||
`lint` phase and a `test` phase, with the final stage depending on both so the
|
`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
|
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.
|
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
|
Dockerfiles install development prerequisites by running `script/bootstrap`
|
||||||
install what those images lack either inline, as the canonical Go `Dockerfile`
|
rather than duplicating installs inline; COPY `script/` and the dependency
|
||||||
below does for `git`, or by running `script/bootstrap`, as the `prompts`
|
manifests (`package.json` + `yarn.lock`, `go.mod` + `go.sum`, etc.) before
|
||||||
repo's own `Dockerfile` does for its yarn packages. The development
|
running it.
|
||||||
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
|
- **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
|
no separate lint file. `script/lint` and `script/test` each build one phase
|
||||||
@@ -160,14 +156,11 @@ style conventions are in separate documents:
|
|||||||
not evidence that anything ran: a sub-second build reporting success is a
|
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`
|
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.
|
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 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
|
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,
|
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
|
and the test phase is based on the Go image. The canonical Go repo
|
||||||
`Dockerfile`:
|
`Dockerfile`:
|
||||||
|
|
||||||
```dockerfile
|
```dockerfile
|
||||||
@@ -180,9 +173,8 @@ style conventions are in separate documents:
|
|||||||
COPY . .
|
COPY . .
|
||||||
RUN golangci-lint run --config .golangci.yml ./...
|
RUN golangci-lint run --config .golangci.yml ./...
|
||||||
|
|
||||||
# Test phase. -race needs cgo and so a C compiler, which the Debian Go
|
# Test phase
|
||||||
# image ships and the alpine one does not.
|
# golang:1.x-alpine, YYYY-MM-DD
|
||||||
# golang:1.x, YYYY-MM-DD
|
|
||||||
FROM golang@sha256:... AS test
|
FROM golang@sha256:... AS test
|
||||||
WORKDIR /src
|
WORKDIR /src
|
||||||
COPY go.mod go.sum ./
|
COPY go.mod go.sum ./
|
||||||
@@ -199,29 +191,15 @@ style conventions are in separate documents:
|
|||||||
FROM golang@sha256:... AS builder
|
FROM golang@sha256:... AS builder
|
||||||
COPY --from=lint /src/go.sum /dev/null
|
COPY --from=lint /src/go.sum /dev/null
|
||||||
COPY --from=test /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
|
WORKDIR /src
|
||||||
COPY go.mod go.sum ./
|
COPY go.mod go.sum ./
|
||||||
RUN go mod download
|
RUN go mod download
|
||||||
COPY . .
|
COPY . .
|
||||||
|
|
||||||
# The VERSION build arg when one is given, otherwise
|
ARG VERSION=dev
|
||||||
# `git describe --tags --always` on the .git in the build context. With
|
RUN CGO_ENABLED=0 go build -trimpath \
|
||||||
# .git present, a version that is still empty, dev or unknown fails the
|
-ldflags="-s -w -X main.Version=${VERSION}" \
|
||||||
# build: git is missing or could not read the checkout.
|
-o /app ./cmd/app/
|
||||||
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
|
# Runtime stage, and the last one
|
||||||
FROM alpine@sha256:...
|
FROM alpine@sha256:...
|
||||||
@@ -243,41 +221,10 @@ style conventions are in separate documents:
|
|||||||
(e.g. a web frontend compiled in a separate stage), the lint phase must
|
(e.g. a web frontend compiled in a separate stage), the lint phase must
|
||||||
create placeholder files so the embed directives resolve. Example:
|
create placeholder files so the embed directives resolve. Example:
|
||||||
`RUN mkdir -p web/dist && touch web/dist/index.html web/dist/style.css`.
|
`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
|
- If the project requires CGO or system libraries for linting (e.g.
|
||||||
in the lint phase. The `golangci/golangci-lint` image is Debian-based and
|
`vips-dev`), install them in the lint phase with `apk add`.
|
||||||
has no `apk`, so install with `apt-get` under the Debian package name
|
- `ARG VERSION=dev` is declared in the stage that compiles and supplied by
|
||||||
(`libvips-dev`, where alpine says `vips-dev`), and delete the package
|
`script/docker` and `script/cibuild`; no stage may call `git describe`.
|
||||||
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
|
- 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.
|
runs `script/cibuild` on push, and checks out the repo as its only other step.
|
||||||
@@ -286,12 +233,7 @@ style conventions are in separate documents:
|
|||||||
carry the same guarantee, because its gate phases may come from the cache. The
|
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
|
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
|
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
|
from a run of its own gates rather than from a cache entry.
|
||||||
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
|
- Use platform-standard formatters: `black` for Python, `prettier` for
|
||||||
JS/CSS/Markdown/HTML, `go fmt` for Go. Always use default configuration with
|
JS/CSS/Markdown/HTML, `go fmt` for Go. Always use default configuration with
|
||||||
@@ -344,19 +286,17 @@ style conventions are in separate documents:
|
|||||||
```
|
```
|
||||||
|
|
||||||
`-count=1` is required on both invocations: it defeats Go's test _result_
|
`-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
|
cache, so the target cannot report a pass it did not earn, and the rerun
|
||||||
tests. It leaves the build cache alone, so it costs the runtime of the suite
|
reproduces a failure instead of replaying it. It leaves the build cache
|
||||||
and no recompilation.
|
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
|
Note that this is a second, independent cache, stacked below the Docker
|
||||||
passing result in its cache directory (`GOCACHE`), and when the same tests
|
layer cache that [issue #26](https://git.eeqj.de/sneak/prompts/issues/26)
|
||||||
run again on unchanged code it prints that result, marked `(cached)`,
|
addresses. `CHECK_EPOCH` guarantees the `RUN make test` _step_ re-executes;
|
||||||
without running them. That matters on a developer's machine, where this
|
it does not guarantee `go test` inside that step does any work, because the
|
||||||
target runs and the directory lasts from one run to the next. The `test`
|
`GOCACHE` baked into earlier image layers survives into the re-executed
|
||||||
phase of the `Dockerfile` needs no `-count=1`: its base image holds no
|
step. They are two separate defects requiring two separate fixes, and a fix
|
||||||
result for this repo's tests and nothing before its `go test` step runs a
|
for one must not be recorded as covering the other.
|
||||||
test, so there is nothing to replay. `--no-cache` (above) is what makes that
|
|
||||||
step run on an unchanged tree.
|
|
||||||
|
|
||||||
Python example:
|
Python example:
|
||||||
|
|
||||||
@@ -400,7 +340,7 @@ style conventions are in separate documents:
|
|||||||
— which is more dangerous than a short file with no secret patterns at all,
|
— 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
|
because it reads as solved and stops anyone looking. Give every
|
||||||
depth-independent pattern the `**/` prefix and leave only genuinely
|
depth-independent pattern the `**/` prefix and leave only genuinely
|
||||||
root-anchored entries unprefixed: `.claude`, and the repo's own host-built
|
root-anchored entries unprefixed: `.git`, and the repo's own host-built
|
||||||
binary, written `/myapp` and never `**/myapp`, which would also match
|
binary, written `/myapp` and never `**/myapp`, which would also match
|
||||||
`cmd/myapp/` and delete the package directory from the context. Matching is
|
`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
|
case-sensitive, and an ALL-CAPS twin per pattern still misses `Server.Key`, so
|
||||||
@@ -425,13 +365,12 @@ style conventions are in separate documents:
|
|||||||
directory, so a repo running agents in subdirectories still ships
|
directory, so a repo running agents in subdirectories still ships
|
||||||
`services/api/.claude/` and must add its own anchored entry there.
|
`services/api/.claude/` and must add its own anchored entry there.
|
||||||
|
|
||||||
- **A plain `docker build .` of a clone stamps the version that
|
- **Excluding `.git` means `git describe` cannot run inside any build stage, and
|
||||||
`git describe --tags --always` gives**, derived from the `.git` in the build
|
it fails quietly there.** In a build stage there is no repository, so
|
||||||
context as the canonical `Dockerfile` above shows. Without its failure check,
|
`git describe` writes nothing to stdout, `-X main.Version=` comes out empty,
|
||||||
a missing `git` or an unreadable checkout would leave `-X main.Version=` empty
|
the binary reports no version at all, and the build still exits 0. Compute the
|
||||||
and the build would still exit 0. `script/docker` and `script/cibuild` pass
|
version on the host and thread it in as a build arg. `script/docker` and
|
||||||
the version they compute on the host; it takes precedence. They do this
|
`script/cibuild` do this, byte-identically across repos:
|
||||||
byte-identically across repos:
|
|
||||||
|
|
||||||
```sh
|
```sh
|
||||||
# Own line: a failing command substitution inside an argument does not
|
# Own line: a failing command substitution inside an argument does not
|
||||||
@@ -448,7 +387,7 @@ style conventions are in separate documents:
|
|||||||
fallback is applied — a live check that fires on a build from an export with
|
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
|
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
|
substitution as `|| echo unknown`, which makes the guard unreachable. The
|
||||||
Dockerfile's side is `ARG VERSION` in the stage that compiles, declared
|
Dockerfile's side is `ARG VERSION=dev` in the stage that compiles, declared
|
||||||
there because `ARG` is stage-scoped; passing `VERSION` to a repo whose
|
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
|
Dockerfile declares no such `ARG` is ignored and costs nothing, which is why
|
||||||
the scripts stay byte-identical. One consequence for CI: the standard
|
the scripts stay byte-identical. One consequence for CI: the standard
|
||||||
@@ -487,18 +426,12 @@ style conventions are in separate documents:
|
|||||||
`test-support` depguard rule, where a repo names its own test-support packages
|
`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
|
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
|
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
|
v2.12.2 (released 2026-05-06), pinned as the digest of the lint phase's base
|
||||||
image
|
image
|
||||||
(`golangci/golangci-lint@sha256:ad862ba6b3798cbe0fd9fd7408d498fd74fbd2623a92406b2fd3898faf0bf98f`,
|
(`golangci/golangci-lint@sha256:5cceeef04e53efe1470638d4b4b4f5ceefd574955ab3941b2d9a68a8c9ad5240`,
|
||||||
which reports `2.14.0 built with go1.27.0 from 114493f9`). A module's `go`
|
which reports `2.12.2 built with go1.26.2 from c0d3ddc9`). That digest is the
|
||||||
directive must not name a newer Go minor version than the one golangci-lint
|
only pin, since no repo installs golangci-lint on the host: bumping the
|
||||||
was built with, or golangci-lint refuses to lint it: this release lints
|
version means changing it and nothing else.
|
||||||
`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
|
- **`script/bootstrap` installs a pinned tool by comparing versions, never by
|
||||||
testing presence.** An `if ! command -v <tool>; then install; fi` guard tests
|
testing presence.** An `if ! command -v <tool>; then install; fi` guard tests
|
||||||
@@ -522,11 +455,6 @@ style conventions are in separate documents:
|
|||||||
|
|
||||||
Keep it POSIX sh: no arrays, no `[[`, no `grep -P`.
|
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 <package>@<commit hash>`). 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
|
- When pinning images or packages by hash, add a comment above the reference
|
||||||
with the version and date (YYYY-MM-DD).
|
with the version and date (YYYY-MM-DD).
|
||||||
|
|
||||||
@@ -639,10 +567,10 @@ style conventions are in separate documents:
|
|||||||
settings.
|
settings.
|
||||||
|
|
||||||
- Avoid putting files in the repo root unless necessary. Root should contain
|
- Avoid putting files in the repo root unless necessary. Root should contain
|
||||||
only project-level config files (`README.md`, `AGENTS.md`, `Makefile`,
|
only project-level config files (`README.md`, `Makefile`, `Dockerfile`,
|
||||||
`Dockerfile`, `LICENSE`, `.gitignore`, `.editorconfig`, `REPO_POLICIES.md`,
|
`LICENSE`, `.gitignore`, `.editorconfig`, `REPO_POLICIES.md`, and
|
||||||
and language-specific config). Everything else goes in a subdirectory.
|
language-specific config). Everything else goes in a subdirectory. Canonical
|
||||||
Canonical subdirectory names:
|
subdirectory names:
|
||||||
- `bin/` — executable scripts and tools
|
- `bin/` — executable scripts and tools
|
||||||
- `cmd/` — Go command entrypoints; thin only: one `main.go` per binary whose
|
- `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
|
body is a single call into `internal/` or `pkg/`, no project logic in
|
||||||
@@ -673,7 +601,3 @@ style conventions are in separate documents:
|
|||||||
- Go: `go.mod`, `go.sum`, `.golangci.yml`
|
- Go: `go.mod`, `go.sum`, `.golangci.yml`
|
||||||
- JS: `package.json`, `yarn.lock`, `.prettierrc`, `.prettierignore`
|
- JS: `package.json`, `yarn.lock`, `.prettierrc`, `.prettierignore`
|
||||||
- Python: `pyproject.toml`
|
- 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.
|
|
||||||
|
|||||||
@@ -32,7 +32,7 @@ from a directory of hand-editable text files.
|
|||||||
- Defence against traffic floods that saturate the host's network link. That
|
- Defence against traffic floods that saturate the host's network link. That
|
||||||
needs help upstream of the host.
|
needs help upstream of the host.
|
||||||
- A web UI or a configuration file. Settings are environment variables. Apart
|
- A web UI or a configuration file. Settings are environment variables. Apart
|
||||||
from settings given as files (the `_FILE` form of a setting, such as
|
from settings given as files (the `_FILE` form of any setting, such as
|
||||||
`SWWAF_ADMIN_TOKEN_FILE`, and `SWWAF_LOG_REMOTE_TLS_CA_FILE`), its own state
|
`SWWAF_ADMIN_TOKEN_FILE`, and `SWWAF_LOG_REMOTE_TLS_CA_FILE`), its own state
|
||||||
files and the lookup database, the only files read are the rule files, which
|
files and the lookup database, the only files read are the rule files, which
|
||||||
hold one regex per line and nothing more elaborate.
|
hold one regex per line and nothing more elaborate.
|
||||||
@@ -293,13 +293,11 @@ it.
|
|||||||
needs: an alert destination, an account key, a token.
|
needs: an alert destination, an account key, a token.
|
||||||
- Every setting's name starts with `SWWAF_`, since `smallwebwaf` shares its
|
- Every setting's name starts with `SWWAF_`, since `smallwebwaf` shares its
|
||||||
container, and so its environment variables, with the app it protects.
|
container, and so its environment variables, with the app it protects.
|
||||||
- Any limit or threshold can be switched off with the value `off`, except
|
- Any limit or threshold can be switched off with the value `off`.
|
||||||
`SWWAF_CLIENT_REQUEST_HEADER_MAX_BYTES`.
|
|
||||||
- A list set to an empty value is an empty list, and replaces the default.
|
- 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
|
- 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
|
setting's name with `_FILE` added, such as `SWWAF_ADMIN_TOKEN_FILE`, for
|
||||||
secrets and long lists. `SWWAF_LOG_REMOTE_TLS_CA_FILE`, whose value names a
|
secrets and long lists.
|
||||||
file already, has no `_FILE` form.
|
|
||||||
- Settings, including those given as files, are read once at start; changing one
|
- Settings, including those given as files, are read once at start; changing one
|
||||||
means restarting the container. The files `smallwebwaf` watches while it runs
|
means restarting the container. The files `smallwebwaf` watches while it runs
|
||||||
are its state files, its rule files and the lookup database.
|
are its state files, its rule files and the lookup database.
|
||||||
@@ -415,9 +413,7 @@ The settings, by group:
|
|||||||
headers, its body.
|
headers, its body.
|
||||||
- `SWWAF_CLIENT_REQUEST_HEADER_MAX_BYTES` (default `32K`): the largest
|
- `SWWAF_CLIENT_REQUEST_HEADER_MAX_BYTES` (default `32K`): the largest
|
||||||
request line and headers a client may send. Over it, `smallwebwaf` answers
|
request line and headers a client may send. Over it, `smallwebwaf` answers
|
||||||
`431` and closes the connection, and nothing reaches the app. It must be
|
`431` and closes the connection, and nothing reaches the app.
|
||||||
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
|
- `SWWAF_CLIENT_IDLE_TIMEOUT` (default `120s`): how long a kept-open
|
||||||
connection may wait for its next request before `smallwebwaf` closes it.
|
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
|
It is longer than the 90 seconds after which traefik, by default, closes a
|
||||||
@@ -949,9 +945,9 @@ and the running `smallwebwaf` takes the edit in.
|
|||||||
- what was broken: the rule ids and target that matched, or the limit, its
|
- 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
|
window, the count reached and the client's limit percentage with what set
|
||||||
it; and any reputation sources that listed the client;
|
it; and any reputation sources that listed the client;
|
||||||
- the request that caused the ban, the one that broke the limit or carried
|
- the requests that caused the ban, up to the last ten: time, method, host,
|
||||||
the clear sign of attack: time, method, host, path with its query string,
|
path with its query string, status and user agent, each text cut to 256
|
||||||
status and user agent, each text cut to 256 bytes;
|
bytes;
|
||||||
- how many requests counted toward the ban, and the time span over which
|
- how many requests counted toward the ban, and the time span over which
|
||||||
they came;
|
they came;
|
||||||
- the netblock's total requests since it was first seen, and the requests
|
- the netblock's total requests since it was first seen, and the requests
|
||||||
@@ -962,13 +958,13 @@ and the running `smallwebwaf` takes the edit in.
|
|||||||
the table is full, so on a public service the file grows to the default
|
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
|
`SWWAF_MAX_TRACKED_CLIENTS` of 20,000, about 20 MiB. Written every 15
|
||||||
minutes, that is under 2 GiB of disk writes a day.
|
minutes, that is under 2 GiB of disk writes a day.
|
||||||
- `bans.json` takes about 1.2 KiB per ban and at most about 2.5 KiB, since
|
- `bans.json` takes about 2 KiB per ban and at most about 8 KiB, since the
|
||||||
the notes hold one request and their texts are cut short. At the default
|
texts in the notes are cut short. At the default `SWWAF_MAX_BANS` of 5,000
|
||||||
`SWWAF_MAX_BANS` of 5,000 it is about 6 MiB, and never more than about 12
|
it is about 10 MiB, and never more than about 40 MiB, plus whatever bans
|
||||||
MiB, plus whatever bans an admin made. It is written when a ban is made,
|
an admin made. It is written when a ban is made, lifted or made permanent,
|
||||||
lifted or made permanent, at most once every 10 seconds, and otherwise
|
at most once every 10 seconds, and otherwise with the 15-minute write, so
|
||||||
with the 15-minute write, so its writes follow the bans made: with a full
|
its writes follow the bans made: with a full file, a hundred new bans a
|
||||||
file, a hundred new bans a day come to about 600 MiB of disk writes.
|
day come to about 1 GiB of disk writes.
|
||||||
- `lookups.json` takes about 150 bytes per answer, about 15 MiB when full.
|
- `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.
|
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.
|
- `reputation.json` and `alerts.json` are usually a few MiB or less.
|
||||||
@@ -1267,19 +1263,14 @@ 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
|
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
|
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
|
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
|
`InRelease` file apt uses, and the build checks them after `apt-get update` and
|
||||||
`apt-get update` and before `apt-get install`, so every package apt installs is
|
before `apt-get install`, so every package apt installs is checked, through
|
||||||
checked, through those files, against hashes the Dockerfile names.
|
those files, against hashes the Dockerfile names. The snapshot service is
|
||||||
`apt-get update` also fetches the live archive's `InRelease` files into the same
|
reached over HTTPS, and the Ubuntu image has no CA certificates of its own, so
|
||||||
directory; their hashes change whenever the archive does, and the install does
|
this one install uses those of the Go image that `smallwebwaf` is built in,
|
||||||
not use them, so the check leaves them out. The hashes are those of the archive
|
which is pinned by digest too: apt's `Acquire::https::CaInfo` option names that
|
||||||
for amd64, and so the image is built for amd64: other architectures use Ubuntu's
|
image's CA certificate file, `/etc/ssl/certs/ca-certificates.crt`, copied into
|
||||||
ports archive, whose `InRelease` files differ. The snapshot service is reached
|
the build.
|
||||||
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
|
Packages from nixpkgs: nixpkgs is fixed at one commit of its newest release
|
||||||
branch, `nixos-26.05` today. For each commit of the branch that has passed its
|
branch, `nixos-26.05` today. For each commit of the branch that has passed its
|
||||||
@@ -1290,18 +1281,15 @@ 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
|
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
|
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.<name>`, and
|
app's Dockerfile installs a package with `nix-env -iA nixpkgs.<name>`, and
|
||||||
whatever it installs is on the `PATH` of every service: the image adds root's
|
whatever it installs is on the `PATH` of every service. Because nixpkgs stays at
|
||||||
Nix profile, `/nix/var/nix/profiles/default/bin`, at the end of the `PATH`,
|
that commit, an app built on the same `smallwebwaf` image gets the same packages
|
||||||
after Ubuntu's own directories, so that no package hides the image's own
|
each time it is built. Unpacked, nixpkgs takes about 500 MiB of disk, more on
|
||||||
commands. busybox, for one, brings its own `sv`, which looks for services
|
some filesystems such as ZFS, and each package an app installs from it adds its
|
||||||
elsewhere. Because nixpkgs stays at that commit, an app built on the same
|
own size, with everything it depends on. A newer commit of the branch, with its
|
||||||
`smallwebwaf` image gets the same packages each time it is built. Unpacked,
|
security fixes, comes with a newer `smallwebwaf` image, as do Ubuntu's own
|
||||||
nixpkgs takes about 500 MiB of disk, more on some filesystems such as ZFS, and
|
fixes; an app takes them by changing the digest in its `FROM` line. When nixpkgs
|
||||||
each package an app installs from it adds its own size, with everything it
|
makes its next release, every six months, the image moves to that release's
|
||||||
depends on. A newer commit of the branch, with its security fixes, comes with a
|
branch.
|
||||||
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:
|
The two processes:
|
||||||
|
|
||||||
@@ -1327,15 +1315,12 @@ The two processes:
|
|||||||
- The container's root filesystem stays writable: runit writes each service's
|
- The container's root filesystem stays writable: runit writes each service's
|
||||||
status into its directory under `/etc/service`.
|
status into its directory under `/etc/service`.
|
||||||
|
|
||||||
The health check: the image's `HEALTHCHECK` runs `smallwebwaf healthcheck`,
|
The health check: the image's `HEALTHCHECK` passes while `smallwebwaf` answers
|
||||||
which passes while `smallwebwaf` answers `GET /_smallwebwaf/healthz` on
|
`GET /_smallwebwaf/healthz` on `127.0.0.1`, at the port in `SWWAF_LISTEN_ADDR`,
|
||||||
`127.0.0.1`, at the port in `SWWAF_LISTEN_ADDR`, and the app accepts connections
|
and the app accepts connections at the address in `SWWAF_UPSTREAM_URL`, and
|
||||||
at the address in `SWWAF_UPSTREAM_URL`, and fails when either does not. The
|
fails when either does not. The container therefore shows as healthy only while
|
||||||
container therefore shows as healthy only while both processes are up. traefik
|
both processes are up. An app with a health check of its own can replace the
|
||||||
sends a container no requests until it shows as healthy, so the check runs every
|
image's `HEALTHCHECK` with one that checks both.
|
||||||
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;
|
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
|
its health check, metrics and ban management are all on that listener, under
|
||||||
@@ -1624,9 +1609,7 @@ holds any token file.
|
|||||||
- The container image described under "Deployment", with runit and the
|
- The container image described under "Deployment", with runit and the
|
||||||
container's health check. The health check calls `/_smallwebwaf/healthz`,
|
container's health check. The health check calls `/_smallwebwaf/healthz`,
|
||||||
so milestone 2 answers that path, although the other admin endpoints come
|
so milestone 2 answers that path, although the other admin endpoints come
|
||||||
later. The image's `/var/lib/smallwebwaf`, which the `run` script gives to
|
later.
|
||||||
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:
|
- 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
|
- 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,
|
and the JSON state files with edits taken in while running, exemptions,
|
||||||
|
|||||||
@@ -1,20 +0,0 @@
|
|||||||
# 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
|
|
||||||
@@ -1,9 +0,0 @@
|
|||||||
#!/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 "$@"
|
|
||||||
@@ -2,23 +2,4 @@ module sneak.berlin/go/smallwebwaf
|
|||||||
|
|
||||||
go 1.26.0
|
go 1.26.0
|
||||||
|
|
||||||
require (
|
require github.com/hashicorp/golang-lru/v2 v2.0.7
|
||||||
github.com/fsnotify/fsnotify v1.10.1
|
|
||||||
github.com/hashicorp/golang-lru/v2 v2.0.7
|
|
||||||
github.com/maxmind/mmdbwriter v1.2.0
|
|
||||||
github.com/oschwald/maxminddb-golang/v2 v2.7.0
|
|
||||||
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
|
|
||||||
go4.org/netipx v0.0.0-20231129151722-fdeea329fbba // indirect
|
|
||||||
golang.org/x/sys v0.48.0 // indirect
|
|
||||||
google.golang.org/protobuf v1.36.11 // indirect
|
|
||||||
)
|
|
||||||
|
|||||||
@@ -1,42 +1,2 @@
|
|||||||
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/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 h1:a+bsQ5rvGLjzHuww6tVxozPZFVghXaHOwFs4luLUK2k=
|
||||||
github.com/hashicorp/golang-lru/v2 v2.0.7/go.mod h1:QeFd9opnmA6QUJc5vARoKUSoFhyfM2/ZepoAG6RGpeM=
|
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/maxmind/mmdbwriter v1.2.0 h1:hyvDopImmgvle3aR8AaddxXnT0iQH2KWJX3vNfkwzYM=
|
|
||||||
github.com/maxmind/mmdbwriter v1.2.0/go.mod h1:EQmKHhk2y9DRVvyNxwCLKC5FrkXZLx4snc5OlLY5XLE=
|
|
||||||
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/oschwald/maxminddb-golang/v2 v2.7.0 h1:ZcAr3GYc2LYC8aec2mCMX9+QOF0EolH3jDFKRV/Z1+U=
|
|
||||||
github.com/oschwald/maxminddb-golang/v2 v2.7.0/go.mod h1:DuKJLbbug6TXC0yJXgs1MWifvXHmudRWzMobMIUu04g=
|
|
||||||
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.12.1 h1:EuwCh5fleGS7H32xRwO3wRGT7DxrDhLAT6FF8MpWDWE=
|
|
||||||
github.com/stretchr/testify v1.12.1/go.mod h1:MDEgiDPPsNp5cuIrHPPCyornHKgEVbtFUmoNlxoYthg=
|
|
||||||
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=
|
|
||||||
go.yaml.in/yaml/v3 v3.0.5 h1:N6y/pJk8buWs9NY5ERU2HSMfm+IuD/OtfdAnq6kESPw=
|
|
||||||
go.yaml.in/yaml/v3 v3.0.5/go.mod h1:HVTZu1O7/Vkt2N+BFy8Zza+lnLsABggaTM2ZpNIGuKg=
|
|
||||||
go4.org/netipx v0.0.0-20231129151722-fdeea329fbba h1:0b9z3AuHCjxk0x/opv64kcgZLBseWJUpBw5I82+2U4M=
|
|
||||||
go4.org/netipx v0.0.0-20231129151722-fdeea329fbba/go.mod h1:PLyyIXexvUFg3Owu6p/WfdlivPbZJsZdgWZlrGope/Y=
|
|
||||||
golang.org/x/sys v0.48.0 h1:bbX/i/6MgT9BVLM9RT1thmxL04yeTAhbEz4SyadbXoo=
|
|
||||||
golang.org/x/sys v0.48.0/go.mod h1:hNLxWAXmnKAxqDtdwIYC4bM9oQPEecfsnNMuSxOs3og=
|
|
||||||
google.golang.org/protobuf v1.36.11 h1:fV6ZwhNocDyBLK0dj+fg8ektcVegBBuEolpbTQyBNVE=
|
|
||||||
google.golang.org/protobuf v1.36.11/go.mod h1:HTf+CrKn2C3g5S8VImy6tdcUvCska2kB7j23XfzDpco=
|
|
||||||
|
|||||||
@@ -1,884 +0,0 @@
|
|||||||
// Package alerts sends alerts on bans, on a source that fails and on a
|
|
||||||
// file with an error to each destination set: to the webhook
|
|
||||||
// SWWAF_ALERT_WEBHOOK_URL names, each as one JSON object, as the "Alert
|
|
||||||
// webhook schema" section of SPEC.md describes, to the Slack incoming
|
|
||||||
// webhook SWWAF_ALERT_SLACK_WEBHOOK_URL names, as a message, and to the
|
|
||||||
// ntfy topic SWWAF_ALERT_NTFY_URL names. A repeat within
|
|
||||||
// SWWAF_ALERT_COOLDOWN is held back, and so is an alert past
|
|
||||||
// SWWAF_ALERT_MAX_PER_HOUR, for the hour's summary. The others wait in a
|
|
||||||
// bounded queue of each destination's own, so that a destination that is
|
|
||||||
// slow or unreachable holds up neither the others nor any request. The
|
|
||||||
// state is written to alerts.json and read from it by the state package.
|
|
||||||
// Nothing logged names a destination's URL, whose path or query can carry
|
|
||||||
// a secret.
|
|
||||||
package alerts
|
|
||||||
|
|
||||||
import (
|
|
||||||
"bytes"
|
|
||||||
"cmp"
|
|
||||||
"context"
|
|
||||||
"encoding/json"
|
|
||||||
"errors"
|
|
||||||
"fmt"
|
|
||||||
"io"
|
|
||||||
"log/slog"
|
|
||||||
"maps"
|
|
||||||
"net/http"
|
|
||||||
"net/netip"
|
|
||||||
"net/url"
|
|
||||||
"slices"
|
|
||||||
"strings"
|
|
||||||
"sync"
|
|
||||||
"sync/atomic"
|
|
||||||
"time"
|
|
||||||
)
|
|
||||||
|
|
||||||
// The events an alert is for, as SWWAF_ALERT_EVENTS names them.
|
|
||||||
const (
|
|
||||||
// EventBan is a ban smallwebwaf made.
|
|
||||||
EventBan = "ban"
|
|
||||||
// EventPermanentBan is a permanent ban smallwebwaf made, or a ban it
|
|
||||||
// made permanent.
|
|
||||||
EventPermanentBan = "permanent_ban"
|
|
||||||
// EventWAFBlock, EventAnomaly and EventReputationHit come with the
|
|
||||||
// Core Rule Set, the anomaly thresholds and the reputation sources;
|
|
||||||
// nothing raises them yet.
|
|
||||||
EventWAFBlock = "waf_block"
|
|
||||||
EventAnomaly = "anomaly"
|
|
||||||
EventReputationHit = "reputation_hit"
|
|
||||||
// EventSourceFailure is GeoJS failing or refusing smallwebwaf.
|
|
||||||
EventSourceFailure = "source_failure"
|
|
||||||
// EventFileError is a rule file or state file edited while smallwebwaf
|
|
||||||
// runs that does not parse, a replacement of the lookup database that
|
|
||||||
// cannot be read, or a state file that cannot be written.
|
|
||||||
EventFileError = "file_error"
|
|
||||||
// EventSummary is the summary of the alerts an hour held back past
|
|
||||||
// SWWAF_ALERT_MAX_PER_HOUR. SWWAF_ALERT_EVENTS does not name it.
|
|
||||||
EventSummary = "summary"
|
|
||||||
)
|
|
||||||
|
|
||||||
// Events returns every event SWWAF_ALERT_EVENTS can name, which is its
|
|
||||||
// default.
|
|
||||||
func Events() []string {
|
|
||||||
return []string{
|
|
||||||
EventBan, EventPermanentBan, EventWAFBlock, EventAnomaly,
|
|
||||||
EventReputationHit, EventSourceFailure, EventFileError,
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// The destinations alerts are sent to, as the metrics and alerts.json
|
|
||||||
// name them.
|
|
||||||
const (
|
|
||||||
// DestinationWebhook is the webhook SWWAF_ALERT_WEBHOOK_URL names.
|
|
||||||
DestinationWebhook = "webhook"
|
|
||||||
// DestinationSlack is the Slack incoming webhook
|
|
||||||
// SWWAF_ALERT_SLACK_WEBHOOK_URL names.
|
|
||||||
DestinationSlack = "slack"
|
|
||||||
// DestinationNtfy is the ntfy topic SWWAF_ALERT_NTFY_URL names.
|
|
||||||
DestinationNtfy = "ntfy"
|
|
||||||
)
|
|
||||||
|
|
||||||
// Destinations returns every destination alerts can be sent to.
|
|
||||||
func Destinations() []string {
|
|
||||||
return []string{DestinationWebhook, DestinationSlack, DestinationNtfy}
|
|
||||||
}
|
|
||||||
|
|
||||||
const (
|
|
||||||
// queueSize is the most alerts that wait to be sent to a destination.
|
|
||||||
// Past it, the oldest is dropped.
|
|
||||||
queueSize = 1000
|
|
||||||
// sendTimeout bounds one request to a destination.
|
|
||||||
sendTimeout = 10 * time.Second
|
|
||||||
// After a request to a destination fails, the alert is sent again a
|
|
||||||
// second later, and retryDelayFactor times as long after each further
|
|
||||||
// failure in a row, up to a minute.
|
|
||||||
firstRetryDelay = time.Second
|
|
||||||
retryDelayFactor = 2
|
|
||||||
maxRetryDelay = time.Minute
|
|
||||||
// maxAnswerBytes is the most of a destination's answer that is read.
|
|
||||||
maxAnswerBytes = 64 << 10
|
|
||||||
)
|
|
||||||
|
|
||||||
var (
|
|
||||||
errStatus = errors.New("the destination answered")
|
|
||||||
// errRefused is a 4xx answer other than 408 and 429: the destination
|
|
||||||
// refuses the alert itself, and would refuse it again.
|
|
||||||
errRefused = errors.New("the destination refused the alert, answering")
|
|
||||||
)
|
|
||||||
|
|
||||||
// Params are what New needs. With none of WebhookURL, SlackURL and
|
|
||||||
// NtfyURL set, no alert is sent.
|
|
||||||
type Params struct {
|
|
||||||
// WebhookURL is where each alert is posted as JSON
|
|
||||||
// (SWWAF_ALERT_WEBHOOK_URL), nil while it is unset. WebhookHeaders
|
|
||||||
// are sent with each (SWWAF_ALERT_WEBHOOK_HEADERS).
|
|
||||||
WebhookURL *url.URL
|
|
||||||
WebhookHeaders http.Header
|
|
||||||
// SlackURL is the Slack incoming webhook each alert is posted to as a
|
|
||||||
// message (SWWAF_ALERT_SLACK_WEBHOOK_URL), nil while it is unset.
|
|
||||||
SlackURL *url.URL
|
|
||||||
// NtfyURL is the ntfy topic each alert is published to
|
|
||||||
// (SWWAF_ALERT_NTFY_URL), nil while it is unset. NtfyToken, unless
|
|
||||||
// empty, is sent with each as a bearer token (SWWAF_ALERT_NTFY_TOKEN).
|
|
||||||
NtfyURL *url.URL
|
|
||||||
NtfyToken string
|
|
||||||
// Events are the events alerts are sent for (SWWAF_ALERT_EVENTS).
|
|
||||||
Events []string
|
|
||||||
// Cooldown is how long a repeat of an alert is held back
|
|
||||||
// (SWWAF_ALERT_COOLDOWN), 0 for no time. MaxPerHour is the most alerts
|
|
||||||
// sent in an hour (SWWAF_ALERT_MAX_PER_HOUR), 0 for no limit.
|
|
||||||
Cooldown time.Duration
|
|
||||||
MaxPerHour int
|
|
||||||
// Instance is SWWAF_INSTANCE_NAME, which every alert gives.
|
|
||||||
Instance string
|
|
||||||
// Now tells the time of an alert, normally time.Now in UTC.
|
|
||||||
Now func() time.Time
|
|
||||||
// ProcessLog receives the requests to a destination that fail.
|
|
||||||
ProcessLog *slog.Logger
|
|
||||||
}
|
|
||||||
|
|
||||||
// Alert is one alert, as the webhook is sent it and alerts.json holds it,
|
|
||||||
// with the fields of the "Alert webhook schema" section of SPEC.md. ASN,
|
|
||||||
// ASName and Country are, for a ban, the client's as the ban's notes give
|
|
||||||
// them.
|
|
||||||
//
|
|
||||||
//nolint:tagliatelle // SPEC.md's alert webhook schema names its fields in snake_case
|
|
||||||
type Alert struct {
|
|
||||||
Instance string `json:"instance"`
|
|
||||||
Time time.Time `json:"time"`
|
|
||||||
Event string `json:"event"`
|
|
||||||
Client netip.Addr `json:"client"`
|
|
||||||
Netblock netip.Prefix `json:"netblock"`
|
|
||||||
ASN string `json:"asn"`
|
|
||||||
ASName string `json:"as_name"`
|
|
||||||
Country string `json:"country"`
|
|
||||||
// Reason is a short sentence, and Detail what is particular to the
|
|
||||||
// event: for a file_error, its "file", and for a source_failure, its
|
|
||||||
// "source", which the cooldown tells repeats by.
|
|
||||||
Reason string `json:"reason"`
|
|
||||||
Detail map[string]any `json:"detail"`
|
|
||||||
// SuppressedRepeats is how many repeats of the alert the cooldown
|
|
||||||
// held back since the last one let through.
|
|
||||||
SuppressedRepeats int `json:"suppressed_repeats"`
|
|
||||||
}
|
|
||||||
|
|
||||||
// Cooldown is, for an event on a netblock, or about a file or a source,
|
|
||||||
// when the last alert let through was raised, and how many repeats the
|
|
||||||
// cooldown has held back since, as alerts.json holds it.
|
|
||||||
//
|
|
||||||
//nolint:tagliatelle // the state files use snake_case, as the request log does
|
|
||||||
type Cooldown struct {
|
|
||||||
Event string `json:"event"`
|
|
||||||
Netblock netip.Prefix `json:"netblock"`
|
|
||||||
File string `json:"file,omitempty"`
|
|
||||||
Source string `json:"source,omitempty"`
|
|
||||||
Sent time.Time `json:"sent"`
|
|
||||||
SuppressedRepeats int `json:"suppressed_repeats"`
|
|
||||||
}
|
|
||||||
|
|
||||||
// Hour is the hour under way, by the clock, as alerts.json holds it: when
|
|
||||||
// it started, how many alerts were let through in it, and how many were
|
|
||||||
// held back in it past MaxPerHour, by event, for its summary.
|
|
||||||
//
|
|
||||||
//nolint:tagliatelle // the state files use snake_case, as the request log does
|
|
||||||
type Hour struct {
|
|
||||||
Start time.Time `json:"start"`
|
|
||||||
Sent int `json:"sent"`
|
|
||||||
HeldBack map[string]int `json:"held_back"`
|
|
||||||
}
|
|
||||||
|
|
||||||
// State is what alerts.json holds: the cooldowns, the hour under way, and
|
|
||||||
// for each destination set, the alerts waiting to be sent to it, oldest
|
|
||||||
// first.
|
|
||||||
type State struct {
|
|
||||||
Cooldowns []Cooldown `json:"cooldowns"`
|
|
||||||
Hour Hour `json:"hour"`
|
|
||||||
Waiting map[string][]Alert `json:"waiting"`
|
|
||||||
}
|
|
||||||
|
|
||||||
// Counts are, for a destination, how many alerts it took, how many
|
|
||||||
// requests to it failed, and how many alerts were dropped from its full
|
|
||||||
// queue or given up as it refused them.
|
|
||||||
type Counts struct {
|
|
||||||
Sent, Failed, Dropped int64
|
|
||||||
}
|
|
||||||
|
|
||||||
// Queue takes the alerts raised, holds back those it must, and sends the
|
|
||||||
// others to each destination set, from a queue of the destination's own.
|
|
||||||
// It is safe for concurrent use.
|
|
||||||
type Queue struct {
|
|
||||||
params Params
|
|
||||||
// destinations are the destinations set, in the order of
|
|
||||||
// Destinations.
|
|
||||||
destinations []*destination
|
|
||||||
|
|
||||||
mu sync.Mutex
|
|
||||||
// cooldowns are the alerts last let through, by event and netblock,
|
|
||||||
// file or source.
|
|
||||||
cooldowns map[cooldownKey]*Cooldown
|
|
||||||
hour Hour
|
|
||||||
|
|
||||||
suppressed atomic.Int64
|
|
||||||
}
|
|
||||||
|
|
||||||
// destination is a destination set, with the alerts waiting to be sent
|
|
||||||
// to it. Its mu is taken after the Queue's, never before.
|
|
||||||
type destination struct {
|
|
||||||
// name is how the metrics and alerts.json name the destination, and
|
|
||||||
// setting the setting that is its URL, which the log names in place
|
|
||||||
// of the URL.
|
|
||||||
name string
|
|
||||||
setting string
|
|
||||||
url *url.URL
|
|
||||||
// message returns the body an alert is posted with, and the headers
|
|
||||||
// sent with it.
|
|
||||||
message func(alert *Alert) ([]byte, http.Header, error)
|
|
||||||
// httpClient follows no redirect: a redirect is a failure.
|
|
||||||
httpClient *http.Client
|
|
||||||
processLog *slog.Logger
|
|
||||||
// queued receives a value when an alert joins the queue, unless one
|
|
||||||
// waits already, so that run looks at the queue again.
|
|
||||||
queued chan struct{}
|
|
||||||
|
|
||||||
mu sync.Mutex
|
|
||||||
// waiting are the alerts waiting to be sent, oldest first.
|
|
||||||
waiting []*Alert
|
|
||||||
|
|
||||||
sent, failed, dropped atomic.Int64
|
|
||||||
}
|
|
||||||
|
|
||||||
// cooldownKey is what makes an alert a repeat of another: the same event
|
|
||||||
// on the same netblock, and about the same file or source, as its detail
|
|
||||||
// names them. Each is empty for an alert without one.
|
|
||||||
type cooldownKey struct {
|
|
||||||
event string
|
|
||||||
netblock netip.Prefix
|
|
||||||
file string
|
|
||||||
source string
|
|
||||||
}
|
|
||||||
|
|
||||||
// cooldownKeyOf returns what makes another alert a repeat of alert.
|
|
||||||
func cooldownKeyOf(alert *Alert) cooldownKey {
|
|
||||||
file, _ := alert.Detail["file"].(string)
|
|
||||||
source, _ := alert.Detail["source"].(string)
|
|
||||||
|
|
||||||
return cooldownKey{alert.Event, alert.Netblock, file, source}
|
|
||||||
}
|
|
||||||
|
|
||||||
// New returns a Queue with no alert yet.
|
|
||||||
func New(params Params) *Queue {
|
|
||||||
q := &Queue{
|
|
||||||
params: params,
|
|
||||||
cooldowns: map[cooldownKey]*Cooldown{},
|
|
||||||
hour: Hour{HeldBack: map[string]int{}},
|
|
||||||
}
|
|
||||||
|
|
||||||
if params.WebhookURL != nil {
|
|
||||||
q.addDestination(DestinationWebhook, "SWWAF_ALERT_WEBHOOK_URL",
|
|
||||||
params.WebhookURL, q.webhookMessage)
|
|
||||||
}
|
|
||||||
|
|
||||||
if params.SlackURL != nil {
|
|
||||||
q.addDestination(DestinationSlack, "SWWAF_ALERT_SLACK_WEBHOOK_URL",
|
|
||||||
params.SlackURL, slackMessage)
|
|
||||||
}
|
|
||||||
|
|
||||||
if params.NtfyURL != nil {
|
|
||||||
q.addDestination(DestinationNtfy, "SWWAF_ALERT_NTFY_URL",
|
|
||||||
params.NtfyURL, q.ntfyMessage)
|
|
||||||
}
|
|
||||||
|
|
||||||
return q
|
|
||||||
}
|
|
||||||
|
|
||||||
// Raise sends alert, which names its event and what is particular to it,
|
|
||||||
// unless no destination is set or SWWAF_ALERT_EVENTS leaves its event
|
|
||||||
// out. It gives alert the instance and the time. An alert that repeats
|
|
||||||
// the last one let through less than Cooldown before is held back and
|
|
||||||
// counted, and the next one let through gives that count. Past MaxPerHour
|
|
||||||
// alerts let through in the hour under way, by the clock, an alert is
|
|
||||||
// held back for that hour's summary instead, which is sent once the hour
|
|
||||||
// has ended; it starts no cooldown, and the repeats held back before it
|
|
||||||
// are given by the next alert let through. Raise never waits: an alert
|
|
||||||
// let through joins the queue of each destination, from which Run sends
|
|
||||||
// it, and with queueSize alerts waiting for a destination, the oldest is
|
|
||||||
// dropped.
|
|
||||||
func (q *Queue) Raise(alert Alert) {
|
|
||||||
if len(q.destinations) == 0 || !slices.Contains(q.params.Events, alert.Event) {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
q.mu.Lock()
|
|
||||||
defer q.mu.Unlock()
|
|
||||||
|
|
||||||
now := q.params.Now()
|
|
||||||
alert.Instance = q.params.Instance
|
|
||||||
alert.Time = now
|
|
||||||
|
|
||||||
if q.repeat(&alert, now) {
|
|
||||||
q.suppressed.Add(1)
|
|
||||||
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
q.endHour(now)
|
|
||||||
|
|
||||||
if q.params.MaxPerHour > 0 && q.hour.Sent >= q.params.MaxPerHour {
|
|
||||||
q.hour.HeldBack[alert.Event]++
|
|
||||||
q.suppressed.Add(1)
|
|
||||||
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
q.startCooldown(&alert, now)
|
|
||||||
q.hour.Sent++
|
|
||||||
q.queue(&alert)
|
|
||||||
}
|
|
||||||
|
|
||||||
// WouldSend reports whether Raise would let an alert for event on
|
|
||||||
// netblock through now: a destination is set, SWWAF_ALERT_EVENTS chooses
|
|
||||||
// event, no alert for event on netblock was let through less than
|
|
||||||
// Cooldown before, and fewer than MaxPerHour alerts have been let through
|
|
||||||
// in the hour under way. Unlike Raise, it counts nothing.
|
|
||||||
func (q *Queue) WouldSend(event string, netblock netip.Prefix) bool {
|
|
||||||
if len(q.destinations) == 0 || !slices.Contains(q.params.Events, event) {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
|
|
||||||
q.mu.Lock()
|
|
||||||
defer q.mu.Unlock()
|
|
||||||
|
|
||||||
now := q.params.Now()
|
|
||||||
|
|
||||||
last, found := q.cooldowns[cooldownKey{event: event, netblock: netblock}]
|
|
||||||
if q.params.Cooldown > 0 && found && now.Sub(last.Sent) < q.params.Cooldown {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
|
|
||||||
q.endHour(now)
|
|
||||||
|
|
||||||
return q.params.MaxPerHour == 0 || q.hour.Sent < q.params.MaxPerHour
|
|
||||||
}
|
|
||||||
|
|
||||||
// Run sends the alerts waiting to each destination, from its own queue,
|
|
||||||
// as destination.run does, until ctx is done. It also ends each hour as
|
|
||||||
// Raise does, so that the hour's summary is sent as it ends. With no
|
|
||||||
// destination set, it returns at once.
|
|
||||||
func (q *Queue) Run(ctx context.Context) {
|
|
||||||
if len(q.destinations) == 0 {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
var sending sync.WaitGroup
|
|
||||||
|
|
||||||
for _, d := range q.destinations {
|
|
||||||
sending.Go(func() { d.run(ctx) })
|
|
||||||
}
|
|
||||||
|
|
||||||
for {
|
|
||||||
q.mu.Lock()
|
|
||||||
untilHourEnds := q.hour.Start.Add(time.Hour).Sub(q.params.Now())
|
|
||||||
q.mu.Unlock()
|
|
||||||
|
|
||||||
select {
|
|
||||||
case <-ctx.Done():
|
|
||||||
sending.Wait()
|
|
||||||
|
|
||||||
return
|
|
||||||
case <-time.After(untilHourEnds):
|
|
||||||
q.mu.Lock()
|
|
||||||
q.endHour(q.params.Now())
|
|
||||||
q.mu.Unlock()
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// Counts returns the counts of the destination name, all 0 for one not
|
|
||||||
// set.
|
|
||||||
func (q *Queue) Counts(name string) Counts {
|
|
||||||
for _, d := range q.destinations {
|
|
||||||
if d.name == name {
|
|
||||||
return Counts{
|
|
||||||
Sent: d.sent.Load(), Failed: d.failed.Load(), Dropped: d.dropped.Load(),
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
return Counts{}
|
|
||||||
}
|
|
||||||
|
|
||||||
// DestinationsSet returns the destinations set, in the order of
|
|
||||||
// Destinations.
|
|
||||||
func (q *Queue) DestinationsSet() []string {
|
|
||||||
names := make([]string, 0, len(q.destinations))
|
|
||||||
for _, d := range q.destinations {
|
|
||||||
names = append(names, d.name)
|
|
||||||
}
|
|
||||||
|
|
||||||
return names
|
|
||||||
}
|
|
||||||
|
|
||||||
// Suppressed is how many alerts were held back: by the cooldown, and past
|
|
||||||
// MaxPerHour. No destination is sent such an alert.
|
|
||||||
func (q *Queue) Suppressed() int64 {
|
|
||||||
return q.suppressed.Load()
|
|
||||||
}
|
|
||||||
|
|
||||||
// Snapshot returns the queue's state, as alerts.json holds it, with the
|
|
||||||
// cooldowns sorted by netblock, then by event, file and source.
|
|
||||||
func (q *Queue) Snapshot() State {
|
|
||||||
q.mu.Lock()
|
|
||||||
defer q.mu.Unlock()
|
|
||||||
|
|
||||||
state := State{
|
|
||||||
Cooldowns: make([]Cooldown, 0, len(q.cooldowns)),
|
|
||||||
Hour: q.hour,
|
|
||||||
Waiting: map[string][]Alert{},
|
|
||||||
}
|
|
||||||
state.Hour.HeldBack = maps.Clone(q.hour.HeldBack)
|
|
||||||
|
|
||||||
for _, cooldown := range q.cooldowns {
|
|
||||||
state.Cooldowns = append(state.Cooldowns, *cooldown)
|
|
||||||
}
|
|
||||||
|
|
||||||
slices.SortFunc(state.Cooldowns, func(a, b Cooldown) int {
|
|
||||||
return cmp.Or(a.Netblock.Compare(b.Netblock), cmp.Compare(a.Event, b.Event),
|
|
||||||
cmp.Compare(a.File, b.File), cmp.Compare(a.Source, b.Source))
|
|
||||||
})
|
|
||||||
|
|
||||||
for _, d := range q.destinations {
|
|
||||||
state.Waiting[d.name] = d.snapshot()
|
|
||||||
}
|
|
||||||
|
|
||||||
return state
|
|
||||||
}
|
|
||||||
|
|
||||||
// Load puts state, read from alerts.json, in place of the queue's state.
|
|
||||||
// Each cooldown's netblock is masked to its length, so that
|
|
||||||
// 203.0.113.9/24 is 203.0.113.0/24. The alerts waiting for a destination
|
|
||||||
// that is not set are dropped, and so are the oldest past queueSize
|
|
||||||
// alerts waiting for one that is.
|
|
||||||
func (q *Queue) Load(state State) {
|
|
||||||
q.mu.Lock()
|
|
||||||
defer q.mu.Unlock()
|
|
||||||
|
|
||||||
q.cooldowns = map[cooldownKey]*Cooldown{}
|
|
||||||
|
|
||||||
for _, cooldown := range state.Cooldowns {
|
|
||||||
cooldown.Netblock = cooldown.Netblock.Masked()
|
|
||||||
key := cooldownKey{cooldown.Event, cooldown.Netblock, cooldown.File, cooldown.Source}
|
|
||||||
q.cooldowns[key] = &cooldown
|
|
||||||
}
|
|
||||||
|
|
||||||
q.hour = state.Hour
|
|
||||||
q.hour.HeldBack = maps.Clone(state.Hour.HeldBack)
|
|
||||||
|
|
||||||
if q.hour.HeldBack == nil {
|
|
||||||
q.hour.HeldBack = map[string]int{}
|
|
||||||
}
|
|
||||||
|
|
||||||
for _, d := range q.destinations {
|
|
||||||
d.load(state.Waiting[d.name])
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// addDestination adds a destination: its name, the setting that gives
|
|
||||||
// its URL, that URL, target, and message, which makes the messages sent
|
|
||||||
// to it.
|
|
||||||
func (q *Queue) addDestination(
|
|
||||||
name, setting string, target *url.URL,
|
|
||||||
message func(alert *Alert) ([]byte, http.Header, error),
|
|
||||||
) {
|
|
||||||
q.destinations = append(q.destinations, &destination{
|
|
||||||
name: name,
|
|
||||||
setting: setting,
|
|
||||||
url: target,
|
|
||||||
message: message,
|
|
||||||
httpClient: &http.Client{
|
|
||||||
CheckRedirect: func(*http.Request, []*http.Request) error {
|
|
||||||
return http.ErrUseLastResponse
|
|
||||||
},
|
|
||||||
},
|
|
||||||
processLog: q.params.ProcessLog,
|
|
||||||
queued: make(chan struct{}, 1),
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
// repeat reports whether alert, raised at now, repeats the last one let
|
|
||||||
// through less than Cooldown before, and counts it if it does.
|
|
||||||
func (q *Queue) repeat(alert *Alert, now time.Time) bool {
|
|
||||||
if q.params.Cooldown == 0 {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
|
|
||||||
last, found := q.cooldowns[cooldownKeyOf(alert)]
|
|
||||||
if !found || now.Sub(last.Sent) >= q.params.Cooldown {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
|
|
||||||
last.SuppressedRepeats++
|
|
||||||
|
|
||||||
return true
|
|
||||||
}
|
|
||||||
|
|
||||||
// startCooldown gives alert, let through at now, the count of the repeats
|
|
||||||
// held back since the last one let through, and notes alert as the last
|
|
||||||
// one let through.
|
|
||||||
func (q *Queue) startCooldown(alert *Alert, now time.Time) {
|
|
||||||
if q.params.Cooldown == 0 {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
key := cooldownKeyOf(alert)
|
|
||||||
|
|
||||||
last, found := q.cooldowns[key]
|
|
||||||
if found {
|
|
||||||
alert.SuppressedRepeats = last.SuppressedRepeats
|
|
||||||
}
|
|
||||||
|
|
||||||
q.cooldowns[key] = &Cooldown{
|
|
||||||
Event: alert.Event, Netblock: alert.Netblock, File: key.file, Source: key.source,
|
|
||||||
Sent: now,
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// endHour ends the hour under way, if now is past it: it queues that
|
|
||||||
// hour's summary when alerts were held back in it past MaxPerHour, and
|
|
||||||
// forgets the cooldowns that have run out with no repeat held back, which
|
|
||||||
// no alert needs any more.
|
|
||||||
func (q *Queue) endHour(now time.Time) {
|
|
||||||
start := now.Truncate(time.Hour)
|
|
||||||
if !start.After(q.hour.Start) {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
heldBack := 0
|
|
||||||
for _, count := range q.hour.HeldBack {
|
|
||||||
heldBack += count
|
|
||||||
}
|
|
||||||
|
|
||||||
if heldBack > 0 {
|
|
||||||
q.queue(&Alert{
|
|
||||||
Instance: q.params.Instance,
|
|
||||||
Time: now,
|
|
||||||
Event: EventSummary,
|
|
||||||
Reason: fmt.Sprintf("%d alerts held back in the hour from %s, past the %d "+
|
|
||||||
"an hour SWWAF_ALERT_MAX_PER_HOUR allows", heldBack,
|
|
||||||
q.hour.Start.Format(time.RFC3339), q.params.MaxPerHour),
|
|
||||||
Detail: map[string]any{
|
|
||||||
"hour": q.hour.Start, "count": heldBack, "events": q.hour.HeldBack,
|
|
||||||
},
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
q.hour = Hour{Start: start, HeldBack: map[string]int{}}
|
|
||||||
|
|
||||||
for key, cooldown := range q.cooldowns {
|
|
||||||
if now.Sub(cooldown.Sent) >= q.params.Cooldown && cooldown.SuppressedRepeats == 0 {
|
|
||||||
delete(q.cooldowns, key)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// queue adds alert to the alerts waiting for each destination.
|
|
||||||
func (q *Queue) queue(alert *Alert) {
|
|
||||||
for _, d := range q.destinations {
|
|
||||||
d.add(alert)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// run sends the alerts waiting, oldest first, until ctx is done. An alert
|
|
||||||
// stays in the queue until the destination answers it with a 2xx status,
|
|
||||||
// or refuses it with a 4xx status other than 408 and 429: a refused alert
|
|
||||||
// is logged, counted as dropped, and given up, so that the next is sent.
|
|
||||||
// Any other request that fails is logged, and the alert sent again
|
|
||||||
// firstRetryDelay later, retryDelayFactor times as long after each
|
|
||||||
// further failure in a row, up to maxRetryDelay.
|
|
||||||
func (d *destination) run(ctx context.Context) {
|
|
||||||
var (
|
|
||||||
retryDelay time.Duration
|
|
||||||
retryAt time.Time
|
|
||||||
)
|
|
||||||
|
|
||||||
for {
|
|
||||||
alert := d.oldest()
|
|
||||||
|
|
||||||
var due <-chan time.Time // nil while no alert waits
|
|
||||||
if alert != nil {
|
|
||||||
due = time.After(time.Until(retryAt))
|
|
||||||
}
|
|
||||||
|
|
||||||
select {
|
|
||||||
case <-ctx.Done():
|
|
||||||
return
|
|
||||||
case <-d.queued:
|
|
||||||
case <-due:
|
|
||||||
err := d.send(ctx, alert)
|
|
||||||
|
|
||||||
switch {
|
|
||||||
case err == nil:
|
|
||||||
d.remove(alert)
|
|
||||||
d.sent.Add(1)
|
|
||||||
|
|
||||||
retryDelay = 0
|
|
||||||
retryAt = time.Time{}
|
|
||||||
case errors.Is(err, errRefused):
|
|
||||||
d.remove(alert)
|
|
||||||
d.failed.Add(1)
|
|
||||||
d.dropped.Add(1)
|
|
||||||
|
|
||||||
retryDelay = 0
|
|
||||||
retryAt = time.Time{}
|
|
||||||
|
|
||||||
d.processLog.Warn("gave up an alert "+d.setting+" refused",
|
|
||||||
"event", alert.Event, "error", err.Error())
|
|
||||||
case ctx.Err() == nil: // not cut off as smallwebwaf stops
|
|
||||||
d.failed.Add(1)
|
|
||||||
|
|
||||||
retryDelay = min(max(retryDelayFactor*retryDelay, firstRetryDelay),
|
|
||||||
maxRetryDelay)
|
|
||||||
retryAt = time.Now().Add(retryDelay)
|
|
||||||
|
|
||||||
d.processLog.Warn("sending an alert to "+d.setting+" failed",
|
|
||||||
"error", err.Error(), "sending_again_in", retryDelay.String())
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// add adds alert to the alerts waiting, first dropping the oldest while
|
|
||||||
// queueSize wait, and has run look at the queue again.
|
|
||||||
func (d *destination) add(alert *Alert) {
|
|
||||||
d.mu.Lock()
|
|
||||||
defer d.mu.Unlock()
|
|
||||||
|
|
||||||
if len(d.waiting) == queueSize {
|
|
||||||
d.waiting = slices.Delete(d.waiting, 0, 1)
|
|
||||||
d.dropped.Add(1)
|
|
||||||
}
|
|
||||||
|
|
||||||
d.waiting = append(d.waiting, alert)
|
|
||||||
|
|
||||||
select {
|
|
||||||
case d.queued <- struct{}{}:
|
|
||||||
default: // a value waits already
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// load puts waiting, read from alerts.json, in place of the alerts
|
|
||||||
// waiting, as add adds them.
|
|
||||||
func (d *destination) load(waiting []Alert) {
|
|
||||||
d.mu.Lock()
|
|
||||||
d.waiting = nil
|
|
||||||
d.mu.Unlock()
|
|
||||||
|
|
||||||
for _, alert := range waiting {
|
|
||||||
d.add(&alert)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// snapshot returns the alerts waiting, oldest first.
|
|
||||||
func (d *destination) snapshot() []Alert {
|
|
||||||
d.mu.Lock()
|
|
||||||
defer d.mu.Unlock()
|
|
||||||
|
|
||||||
waiting := make([]Alert, 0, len(d.waiting))
|
|
||||||
for _, alert := range d.waiting {
|
|
||||||
waiting = append(waiting, *alert)
|
|
||||||
}
|
|
||||||
|
|
||||||
return waiting
|
|
||||||
}
|
|
||||||
|
|
||||||
// oldest returns the oldest alert waiting, nil when none waits.
|
|
||||||
func (d *destination) oldest() *Alert {
|
|
||||||
d.mu.Lock()
|
|
||||||
defer d.mu.Unlock()
|
|
||||||
|
|
||||||
if len(d.waiting) == 0 {
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
return d.waiting[0]
|
|
||||||
}
|
|
||||||
|
|
||||||
// remove takes alert, which run has sent or given up, out of the queue,
|
|
||||||
// unless it has been dropped from it, or load has replaced the queue,
|
|
||||||
// since run took it. Only the oldest alert is ever dropped, so alert is
|
|
||||||
// the oldest if it is there at all.
|
|
||||||
func (d *destination) remove(alert *Alert) {
|
|
||||||
d.mu.Lock()
|
|
||||||
defer d.mu.Unlock()
|
|
||||||
|
|
||||||
if len(d.waiting) > 0 && d.waiting[0] == alert {
|
|
||||||
d.waiting = slices.Delete(d.waiting, 0, 1)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// send posts alert to the destination, as message makes it, and returns
|
|
||||||
// an error unless the destination answers with a 2xx status: one that
|
|
||||||
// wraps errRefused for a 4xx status other than 408 and 429. No error
|
|
||||||
// names the destination's URL, whose path or query can carry a secret.
|
|
||||||
func (d *destination) send(ctx context.Context, alert *Alert) error {
|
|
||||||
body, header, err := d.message(alert)
|
|
||||||
if err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
|
|
||||||
ctx, cancel := context.WithTimeout(ctx, sendTimeout)
|
|
||||||
defer cancel()
|
|
||||||
|
|
||||||
req, err := http.NewRequestWithContext(ctx, http.MethodPost, d.url.String(),
|
|
||||||
bytes.NewReader(body))
|
|
||||||
if err != nil {
|
|
||||||
return fmt.Errorf("make the request: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
maps.Copy(req.Header, header)
|
|
||||||
|
|
||||||
res, err := d.httpClient.Do(req)
|
|
||||||
if err != nil {
|
|
||||||
// The client's error names the URL: only what went wrong is kept.
|
|
||||||
if urlErr, ok := errors.AsType[*url.Error](err); ok {
|
|
||||||
return urlErr.Err
|
|
||||||
}
|
|
||||||
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
|
|
||||||
defer func() {
|
|
||||||
_ = res.Body.Close()
|
|
||||||
}()
|
|
||||||
|
|
||||||
// Read, so that the connection can be used again.
|
|
||||||
_, _ = io.Copy(io.Discard, io.LimitReader(res.Body, maxAnswerBytes))
|
|
||||||
|
|
||||||
switch status := res.StatusCode; {
|
|
||||||
case status >= http.StatusOK && status < http.StatusMultipleChoices:
|
|
||||||
return nil
|
|
||||||
case status >= http.StatusBadRequest && status < http.StatusInternalServerError &&
|
|
||||||
status != http.StatusRequestTimeout && status != http.StatusTooManyRequests:
|
|
||||||
return fmt.Errorf("%w %s", errRefused, res.Status)
|
|
||||||
default:
|
|
||||||
return fmt.Errorf("%w %s", errStatus, res.Status)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// webhookMessage returns alert as JSON, for the webhook, and the headers
|
|
||||||
// sent with it: WebhookHeaders, and its Content-Type.
|
|
||||||
func (q *Queue) webhookMessage(alert *Alert) ([]byte, http.Header, error) {
|
|
||||||
body, err := json.Marshal(alert)
|
|
||||||
if err != nil {
|
|
||||||
return nil, nil, fmt.Errorf("encode the alert: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
header := http.Header{}
|
|
||||||
maps.Copy(header, q.params.WebhookHeaders)
|
|
||||||
header.Set("Content-Type", "application/json")
|
|
||||||
|
|
||||||
return body, header, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// slackMessage returns alert as a message for a Slack incoming webhook,
|
|
||||||
// in JSON: its title in bold, then its text, with &, < and > escaped, as
|
|
||||||
// Slack asks, so that nothing in them is read as a link or a mention.
|
|
||||||
func slackMessage(alert *Alert) ([]byte, http.Header, error) {
|
|
||||||
escape := strings.NewReplacer("&", "&", "<", "<", ">", ">").Replace
|
|
||||||
|
|
||||||
body, err := json.Marshal(map[string]string{
|
|
||||||
"text": "*" + escape(title(alert)) + "*\n" + escape(text(alert)),
|
|
||||||
})
|
|
||||||
if err != nil {
|
|
||||||
return nil, nil, fmt.Errorf("encode the message: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
return body, http.Header{"Content-Type": {"application/json"}}, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// ntfyMessage returns alert's text, as the message published to ntfy,
|
|
||||||
// and the headers sent with it: its title, the priority and the tag of
|
|
||||||
// its event, and NtfyToken, unless it is empty, as a bearer token.
|
|
||||||
func (q *Queue) ntfyMessage(alert *Alert) ([]byte, http.Header, error) {
|
|
||||||
header := http.Header{
|
|
||||||
"Title": {title(alert)},
|
|
||||||
"Priority": {ntfyPriority(alert.Event)},
|
|
||||||
"Tags": {ntfyTag(alert.Event)},
|
|
||||||
}
|
|
||||||
|
|
||||||
if q.params.NtfyToken != "" {
|
|
||||||
header.Set("Authorization", "Bearer "+q.params.NtfyToken)
|
|
||||||
}
|
|
||||||
|
|
||||||
return []byte(text(alert)), header, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// ntfyPriority returns the priority an alert for event is published to
|
|
||||||
// ntfy with: high for an event the admin needs to look at.
|
|
||||||
func ntfyPriority(event string) string {
|
|
||||||
switch event {
|
|
||||||
case EventPermanentBan, EventAnomaly, EventSourceFailure, EventFileError:
|
|
||||||
return "high"
|
|
||||||
case EventReputationHit:
|
|
||||||
return "low"
|
|
||||||
default: // ban, waf_block and summary
|
|
||||||
return "default"
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// ntfyTag returns the tag an alert for event is published to ntfy with,
|
|
||||||
// which ntfy shows as an emoji.
|
|
||||||
func ntfyTag(event string) string {
|
|
||||||
switch event {
|
|
||||||
case EventBan, EventPermanentBan:
|
|
||||||
return "no_entry"
|
|
||||||
case EventWAFBlock:
|
|
||||||
return "shield"
|
|
||||||
case EventAnomaly:
|
|
||||||
return "chart_with_upwards_trend"
|
|
||||||
case EventReputationHit:
|
|
||||||
return "label"
|
|
||||||
case EventSourceFailure, EventFileError:
|
|
||||||
return "warning"
|
|
||||||
default: // summary
|
|
||||||
return "bar_chart"
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// title returns the title of alert in Slack and ntfy: the instance and
|
|
||||||
// the event.
|
|
||||||
func title(alert *Alert) string {
|
|
||||||
return alert.Instance + ": " + alert.Event
|
|
||||||
}
|
|
||||||
|
|
||||||
// text returns the text of alert in Slack and ntfy: its reason, then a
|
|
||||||
// line for each of its client, netblock and country, the file, source,
|
|
||||||
// error and mode its detail gives, and its suppressed repeats, that it
|
|
||||||
// has.
|
|
||||||
func text(alert *Alert) string {
|
|
||||||
lines := []string{alert.Reason}
|
|
||||||
|
|
||||||
if alert.Client.IsValid() {
|
|
||||||
lines = append(lines, "client: "+alert.Client.String())
|
|
||||||
}
|
|
||||||
|
|
||||||
if alert.Netblock.IsValid() {
|
|
||||||
lines = append(lines, "netblock: "+alert.Netblock.String())
|
|
||||||
}
|
|
||||||
|
|
||||||
if alert.Country != "" {
|
|
||||||
lines = append(lines, "country: "+alert.Country)
|
|
||||||
}
|
|
||||||
|
|
||||||
for _, name := range []string{"file", "source", "error", "mode"} {
|
|
||||||
value, _ := alert.Detail[name].(string)
|
|
||||||
if value != "" {
|
|
||||||
lines = append(lines, name+": "+value)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
if alert.SuppressedRepeats > 0 {
|
|
||||||
lines = append(lines, fmt.Sprintf("suppressed repeats: %d", alert.SuppressedRepeats))
|
|
||||||
}
|
|
||||||
|
|
||||||
return strings.Join(lines, "\n")
|
|
||||||
}
|
|
||||||
File diff suppressed because it is too large
Load Diff
@@ -1,16 +0,0 @@
|
|||||||
package alerts
|
|
||||||
|
|
||||||
import "net/http"
|
|
||||||
|
|
||||||
// QueueSize is the most alerts that wait to be sent to a destination.
|
|
||||||
const QueueSize = queueSize
|
|
||||||
|
|
||||||
// SetTransport has q's requests to the destination name go through
|
|
||||||
// transport instead of the network.
|
|
||||||
func (q *Queue) SetTransport(name string, transport http.RoundTripper) {
|
|
||||||
for _, d := range q.destinations {
|
|
||||||
if d.name == name {
|
|
||||||
d.httpClient.Transport = transport
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@@ -1,286 +0,0 @@
|
|||||||
package bans_test
|
|
||||||
|
|
||||||
import (
|
|
||||||
"net/netip"
|
|
||||||
"testing"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
"sneak.berlin/go/smallwebwaf/internal/bans"
|
|
||||||
)
|
|
||||||
|
|
||||||
func TestBanWithoutACauseIsAnAdmins(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
netblock := netip.MustParsePrefix("203.0.113.0/24")
|
|
||||||
ledger := bans.New(defaultRules())
|
|
||||||
ledger.Load([]bans.Ban{{Netblock: netblock, Start: midnight()}})
|
|
||||||
|
|
||||||
if got := ledger.Bans(netblock)[0].Cause; got != bans.CauseAdmin {
|
|
||||||
t.Errorf("the ban's cause is %q, want admin", got)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestAdminsBansAreNeverDroppedAndDoNotCountTowardMaxBans(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
rules := defaultRules()
|
|
||||||
rules.MaxBans = 1
|
|
||||||
ledger := bans.New(rules)
|
|
||||||
adminsOnly := netip.MustParsePrefix("198.51.100.0/24")
|
|
||||||
both := netip.MustParsePrefix("203.0.113.1/32")
|
|
||||||
second := netip.MustParsePrefix("203.0.113.2/32")
|
|
||||||
third := netip.MustParsePrefix("203.0.113.3/32")
|
|
||||||
|
|
||||||
// Seen longest ago, a netblock with two of an admin's bans alone, and
|
|
||||||
// then one with an admin's ban before a ban smallwebwaf made: the one
|
|
||||||
// ban counted toward MaxBans.
|
|
||||||
ledger.Load([]bans.Ban{
|
|
||||||
{Netblock: adminsOnly, Start: midnight().Add(-3 * time.Hour), Cause: bans.CauseAdmin},
|
|
||||||
{Netblock: adminsOnly, Start: midnight().Add(-2 * time.Hour), Cause: bans.CauseAdmin},
|
|
||||||
{Netblock: both, Start: midnight().Add(-time.Hour), Cause: bans.CauseAdmin},
|
|
||||||
{
|
|
||||||
Netblock: both,
|
|
||||||
Start: midnight(),
|
|
||||||
Expires: midnight().Add(time.Hour),
|
|
||||||
Cause: bans.CauseLimit,
|
|
||||||
},
|
|
||||||
})
|
|
||||||
wantBans(t, ledger, map[netip.Prefix]int{adminsOnly: 2, both: 2})
|
|
||||||
|
|
||||||
// A new ban drops the ban smallwebwaf made, and only that one.
|
|
||||||
ledger.BanForLimit(second, midnight(), bans.Notes{})
|
|
||||||
wantBans(t, ledger, map[netip.Prefix]int{adminsOnly: 2, both: 1, second: 1})
|
|
||||||
|
|
||||||
if ledger.Bans(both)[0].Cause != bans.CauseAdmin {
|
|
||||||
t.Errorf("%s kept %+v, want the admin's ban", both, ledger.Bans(both))
|
|
||||||
}
|
|
||||||
|
|
||||||
// And the next drops that one.
|
|
||||||
ledger.BanForLimit(third, midnight(), bans.Notes{})
|
|
||||||
wantBans(t, ledger, map[netip.Prefix]int{adminsOnly: 2, both: 1, second: 0, third: 1})
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestReasonOfTheBansSmallwebwafMakes(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
ledger := bans.New(defaultRules())
|
|
||||||
|
|
||||||
limit, _ := ledger.BanForLimit(netip.MustParsePrefix("203.0.113.1/32"), midnight(),
|
|
||||||
bans.Notes{Limit: 1000, Window: "minute"})
|
|
||||||
attack, _ := ledger.BanForAttack(netip.MustParsePrefix("203.0.113.2/32"), midnight(),
|
|
||||||
bans.Notes{RuleID: "git-dir", Target: "path"})
|
|
||||||
|
|
||||||
for _, tc := range []struct{ got, want string }{
|
|
||||||
{limit.Reason, "requests per minute over the limit of 1000"},
|
|
||||||
{attack.Reason, "matched the rule git-dir"},
|
|
||||||
} {
|
|
||||||
if tc.got != tc.want {
|
|
||||||
t.Errorf("the reason is %q, want %q", tc.got, tc.want)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestLiftedBanForALimitRefusesNothingAndMakesNoBanLonger(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
// An hour's ban lifted ten minutes after it started.
|
|
||||||
netblock := netip.MustParsePrefix("203.0.113.9/32")
|
|
||||||
lifted := bans.Ban{
|
|
||||||
Netblock: netblock,
|
|
||||||
Start: midnight(),
|
|
||||||
Expires: midnight().Add(time.Hour),
|
|
||||||
Cause: bans.CauseLimit,
|
|
||||||
Lifted: midnight().Add(10 * time.Minute),
|
|
||||||
}
|
|
||||||
|
|
||||||
ledger := bans.New(defaultRules())
|
|
||||||
ledger.Load([]bans.Ban{lifted})
|
|
||||||
|
|
||||||
// While it would still last, it refuses nothing, and a limit broken
|
|
||||||
// bans for an hour, as a first broken limit does; the lifted ban is
|
|
||||||
// kept, and counted among the earlier bans.
|
|
||||||
now := midnight().Add(30 * time.Minute)
|
|
||||||
|
|
||||||
_, banned, _ := ledger.Check(netblock.Addr(), now)
|
|
||||||
if banned {
|
|
||||||
t.Error("the lifted ban refuses")
|
|
||||||
}
|
|
||||||
|
|
||||||
ban, _ := ledger.BanForLimit(netblock, now, bans.Notes{})
|
|
||||||
if ban.Expires.Sub(ban.Start) != time.Hour ||
|
|
||||||
ban.Notes.EarlierBans != (bans.EarlierBans{Limit: 1}) {
|
|
||||||
t.Errorf("the next ban lasts %s with earlier bans %+v, want 1h and 1 for a limit",
|
|
||||||
ban.Expires.Sub(ban.Start), ban.Notes.EarlierBans)
|
|
||||||
}
|
|
||||||
|
|
||||||
held := ledger.Bans(netblock)
|
|
||||||
if len(held) != 2 || held[0] != lifted {
|
|
||||||
t.Errorf("the ledger holds %+v, want the lifted ban and the new one", held)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestLiftedBanForAnAttackRefusesNothingAndMakesNoBanLonger(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
// A permanent ban for a clear sign of attack, lifted.
|
|
||||||
netblock := netip.MustParsePrefix("203.0.113.9/32")
|
|
||||||
ledger := bans.New(defaultRules())
|
|
||||||
ledger.Load([]bans.Ban{{
|
|
||||||
Netblock: netblock,
|
|
||||||
Start: midnight(),
|
|
||||||
Cause: bans.CauseAttack,
|
|
||||||
Lifted: midnight().Add(time.Hour),
|
|
||||||
}})
|
|
||||||
|
|
||||||
now := midnight().Add(2 * time.Hour)
|
|
||||||
|
|
||||||
_, banned, _ := ledger.Find(netblock.Addr(), now)
|
|
||||||
if banned {
|
|
||||||
t.Error("the lifted ban refuses")
|
|
||||||
}
|
|
||||||
|
|
||||||
active, permanent := ledger.Count(now)
|
|
||||||
if active != 0 || permanent != 0 {
|
|
||||||
t.Errorf("%d bans are active and %d permanent, want none", active, permanent)
|
|
||||||
}
|
|
||||||
|
|
||||||
// The next clear sign of attack bans for seven days, as a first does.
|
|
||||||
ban, _ := ledger.BanForAttack(netblock, now, bans.Notes{})
|
|
||||||
if ban.Expires.Sub(ban.Start) != 7*day {
|
|
||||||
t.Errorf("the next ban for an attack ends at %s, want seven days on", ban.Expires)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestLoadEditCountsTheBansAnAdminMade(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
ledger := bans.New(defaultRules())
|
|
||||||
made, _ := ledger.BanForLimit(netip.MustParsePrefix("203.0.113.1/32"), midnight(),
|
|
||||||
bans.Notes{})
|
|
||||||
atStart := bans.Ban{
|
|
||||||
Netblock: netip.MustParsePrefix("203.0.113.2/32"),
|
|
||||||
Start: midnight(),
|
|
||||||
}
|
|
||||||
|
|
||||||
// The bans read at the start were made before it.
|
|
||||||
ledger.Load([]bans.Ban{made, atStart})
|
|
||||||
|
|
||||||
if got := ledger.Made(bans.CauseAdmin); got != 0 {
|
|
||||||
t.Fatalf("%d bans made by an admin after the start's, want none", got)
|
|
||||||
}
|
|
||||||
|
|
||||||
// The admin keeps the ban smallwebwaf made, keeps the one read at the
|
|
||||||
// start, and adds one without a cause: that one alone is made.
|
|
||||||
kept := made
|
|
||||||
kept.Cause = bans.CauseAdmin
|
|
||||||
added := bans.Ban{
|
|
||||||
Netblock: netip.MustParsePrefix("203.0.113.3/32"),
|
|
||||||
Start: midnight(),
|
|
||||||
}
|
|
||||||
ledger.LoadEdit([]bans.Ban{kept, atStart, added})
|
|
||||||
|
|
||||||
if ledger.Made(bans.CauseAdmin) != 1 || ledger.Made(bans.CauseLimit) != 1 {
|
|
||||||
t.Errorf("%d bans made by an admin and %d for a limit, want 1 of each",
|
|
||||||
ledger.Made(bans.CauseAdmin), ledger.Made(bans.CauseLimit))
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestAdminsBanIsMadeWhileAnotherLasts(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
netblock := netip.MustParsePrefix("203.0.113.0/24")
|
|
||||||
ledger := bans.New(defaultRules())
|
|
||||||
|
|
||||||
// An hour's ban for a broken limit.
|
|
||||||
ledger.BanForLimit(netblock, midnight(), bans.Notes{})
|
|
||||||
wantChanged(t, ledger, true)
|
|
||||||
|
|
||||||
// A minute later an admin bans the netblock for good, named by an
|
|
||||||
// address in it: that ban is made, and counts the other among the
|
|
||||||
// earlier bans.
|
|
||||||
now := midnight().Add(time.Minute)
|
|
||||||
want := bans.Ban{
|
|
||||||
Netblock: netblock,
|
|
||||||
Start: now,
|
|
||||||
Cause: bans.CauseAdmin,
|
|
||||||
Reason: "probes for logins",
|
|
||||||
Notes: bans.Notes{EarlierBans: bans.EarlierBans{Limit: 1}},
|
|
||||||
}
|
|
||||||
|
|
||||||
got := ledger.BanForAdmin(netip.MustParsePrefix("203.0.113.9/24"), now, time.Time{},
|
|
||||||
"probes for logins")
|
|
||||||
if got != want {
|
|
||||||
t.Errorf("the admin's ban is\n%+v\nwant\n%+v", got, want)
|
|
||||||
}
|
|
||||||
|
|
||||||
wantChanged(t, ledger, true)
|
|
||||||
|
|
||||||
if made := ledger.Made(bans.CauseAdmin); made != 1 {
|
|
||||||
t.Errorf("%d bans made by an admin, want 1", made)
|
|
||||||
}
|
|
||||||
|
|
||||||
// It refuses once the ban for the limit has ended.
|
|
||||||
ban, banned, _ := ledger.Find(netblock.Addr(), midnight().Add(2*time.Hour))
|
|
||||||
if !banned || ban != want {
|
|
||||||
t.Errorf("after the limit's ban the netblock is under %+v (%t), want %+v",
|
|
||||||
ban, banned, want)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestLiftLiftsEveryActiveBanCoveringTheClient(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
client := netip.MustParseAddr("203.0.113.9")
|
|
||||||
own := netip.MustParsePrefix("203.0.113.9/32")
|
|
||||||
wide := netip.MustParsePrefix("203.0.113.0/24")
|
|
||||||
other := netip.MustParsePrefix("203.0.113.10/32")
|
|
||||||
|
|
||||||
ledger := bans.New(defaultRules())
|
|
||||||
ledger.Load([]bans.Ban{
|
|
||||||
// Ended an hour ago.
|
|
||||||
{
|
|
||||||
Netblock: own, Start: midnight().Add(-2 * time.Hour),
|
|
||||||
Expires: midnight().Add(-time.Hour), Cause: bans.CauseLimit,
|
|
||||||
},
|
|
||||||
// Active, on the client's address and on its /24.
|
|
||||||
{
|
|
||||||
Netblock: own, Start: midnight(), Expires: midnight().Add(time.Hour),
|
|
||||||
Cause: bans.CauseLimit,
|
|
||||||
},
|
|
||||||
{Netblock: wide, Start: midnight(), Cause: bans.CauseAdmin},
|
|
||||||
// Another client's.
|
|
||||||
{Netblock: other, Start: midnight(), Cause: bans.CauseAdmin},
|
|
||||||
})
|
|
||||||
|
|
||||||
now := midnight().Add(time.Minute)
|
|
||||||
|
|
||||||
lifted := ledger.Lift(client, now)
|
|
||||||
if len(lifted) != 2 || lifted[0].Lifted != now || lifted[1].Lifted != now {
|
|
||||||
t.Errorf("lifted %+v, want the two active bans covering the client", lifted)
|
|
||||||
}
|
|
||||||
|
|
||||||
wantChanged(t, ledger, true)
|
|
||||||
|
|
||||||
if _, banned, _ := ledger.Check(client, now); banned {
|
|
||||||
t.Error("the client is still banned")
|
|
||||||
}
|
|
||||||
|
|
||||||
if _, banned, _ := ledger.Check(other.Addr(), now); !banned {
|
|
||||||
t.Error("the other client's ban was lifted")
|
|
||||||
}
|
|
||||||
|
|
||||||
// The lifted bans are kept, and the one that had ended is not lifted.
|
|
||||||
covering := ledger.Covering(client)
|
|
||||||
if len(covering) != 3 || covering[0].Netblock != wide ||
|
|
||||||
!covering[1].Lifted.IsZero() || covering[2].Lifted != now {
|
|
||||||
t.Errorf("the bans covering the client are %+v, want the /24's and both "+
|
|
||||||
"of its own, the earlier not lifted", covering)
|
|
||||||
}
|
|
||||||
|
|
||||||
// With none active, nothing is lifted or changed.
|
|
||||||
if lifted = ledger.Lift(client, now); len(lifted) != 0 {
|
|
||||||
t.Errorf("lifted %+v again", lifted)
|
|
||||||
}
|
|
||||||
|
|
||||||
wantChanged(t, ledger, false)
|
|
||||||
}
|
|
||||||
@@ -1,813 +0,0 @@
|
|||||||
// Package bans is the ban ledger: the bans smallwebwaf makes on the
|
|
||||||
// netblocks of clients that break a rate limit or show a clear sign of
|
|
||||||
// attack, and those an admin makes, 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 (
|
|
||||||
"fmt"
|
|
||||||
"math"
|
|
||||||
"net/netip"
|
|
||||||
"slices"
|
|
||||||
"strings"
|
|
||||||
"sync"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
"github.com/hashicorp/golang-lru/v2/simplelru"
|
|
||||||
)
|
|
||||||
|
|
||||||
// The causes of bans.
|
|
||||||
const (
|
|
||||||
// CauseLimit is a ban smallwebwaf made for a broken limit.
|
|
||||||
CauseLimit = "limit"
|
|
||||||
// CauseAttack is a ban smallwebwaf made for a clear sign of attack.
|
|
||||||
CauseAttack = "attack"
|
|
||||||
// CauseAdmin is a ban an admin made, or one smallwebwaf made that an
|
|
||||||
// admin keeps. It is never dropped.
|
|
||||||
CauseAdmin = "admin"
|
|
||||||
)
|
|
||||||
|
|
||||||
// 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 lasts, and how many bans are held.
|
|
||||||
type Rules struct {
|
|
||||||
// LimitBanDuration is how long a first ban for a broken limit lasts.
|
|
||||||
LimitBanDuration time.Duration
|
|
||||||
// LimitBanRepeatWindow is how soon after the end of the netblock's
|
|
||||||
// ban that ended last, other than one for a clear sign of attack, 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 for a broken limit; one that would
|
|
||||||
// be longer is permanent instead.
|
|
||||||
MaxBanDuration time.Duration
|
|
||||||
// AttackBanDuration is how long a first ban for a clear sign of attack
|
|
||||||
// lasts.
|
|
||||||
AttackBanDuration time.Duration
|
|
||||||
// MaxBans is the most bans held whose cause is not CauseAdmin, at
|
|
||||||
// least one. Past it, the earliest such ban of the netblock that has
|
|
||||||
// gone longest without a request is dropped. Bans whose cause is
|
|
||||||
// CauseAdmin are held besides, and never dropped.
|
|
||||||
MaxBans int
|
|
||||||
}
|
|
||||||
|
|
||||||
// Ban is a ban on a netblock.
|
|
||||||
type Ban struct {
|
|
||||||
Netblock netip.Prefix
|
|
||||||
Start time.Time
|
|
||||||
// Expires is when the ban ends, zero for a permanent ban.
|
|
||||||
Expires time.Time
|
|
||||||
// Cause is CauseLimit, CauseAttack or CauseAdmin.
|
|
||||||
Cause string
|
|
||||||
// Reason is a short text: for a ban smallwebwaf made, the limit broken
|
|
||||||
// or the rule that matched; for an admin's, what the admin wrote.
|
|
||||||
Reason string
|
|
||||||
// Lifted is when an admin lifted the ban, zero while no admin has. A
|
|
||||||
// lifted ban refuses nothing, and does not make the netblock's next
|
|
||||||
// ban longer.
|
|
||||||
Lifted 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: it has not
|
|
||||||
// been lifted, and has not run out.
|
|
||||||
func (b Ban) ActiveAt(now time.Time) bool {
|
|
||||||
return b.Lifted.IsZero() && (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 {
|
|
||||||
// ASN, ASName and Country are the client's AS number, AS name and
|
|
||||||
// country, when they were looked up: when the request that caused the
|
|
||||||
// ban was made, or when GeoJS answered about the client afterwards.
|
|
||||||
ASN string `json:"asn"`
|
|
||||||
ASName string `json:"as_name"`
|
|
||||||
Country string `json:"country"`
|
|
||||||
// Limit, Window and Count are, for a ban for a broken limit, 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,omitempty"`
|
|
||||||
Window string `json:"window,omitempty"`
|
|
||||||
Count float64 `json:"count,omitempty"`
|
|
||||||
// RuleID and Target are, for a ban for a clear sign of attack, the id
|
|
||||||
// of the rule file rule that matched, and its target.
|
|
||||||
RuleID string `json:"rule_id,omitempty"`
|
|
||||||
Target string `json:"target,omitempty"`
|
|
||||||
// Request is the request that broke the limit, or that was the clear
|
|
||||||
// sign of attack.
|
|
||||||
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, by
|
|
||||||
// cause.
|
|
||||||
EarlierBans EarlierBans `json:"earlier_bans"`
|
|
||||||
}
|
|
||||||
|
|
||||||
// EarlierBans counts a netblock's bans before a ban, by cause.
|
|
||||||
type EarlierBans struct {
|
|
||||||
Limit int `json:"limit"`
|
|
||||||
Attack int `json:"attack"`
|
|
||||||
Admin int `json:"admin"`
|
|
||||||
}
|
|
||||||
|
|
||||||
// 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 whose cause is not CauseAdmin,
|
|
||||||
// at most rules.MaxBans.
|
|
||||||
held int
|
|
||||||
// made is how many bans have been made since the start, by cause: by
|
|
||||||
// the ledger, and by an admin, through BanForAdmin or in an edit of
|
|
||||||
// bans.json.
|
|
||||||
made map[string]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 {
|
|
||||||
// The ledger drops bans itself, and never those whose cause is
|
|
||||||
// CauseAdmin, however many there are, so the LRU has no limit of its
|
|
||||||
// own: it keeps the netblocks in the order they were last seen.
|
|
||||||
netblocks, err := simplelru.NewLRU[netip.Prefix, *[]Ban](math.MaxInt, 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,
|
|
||||||
made: map[string]int{},
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// Changed receives a value after a ban is made, lifted or made permanent,
|
|
||||||
// so that bans.json can be written. Several changes 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. A ban for a clear sign of
|
|
||||||
// attack is made permanent by the request: the netblock is malicious.
|
|
||||||
// The last result reports whether the request made the ban permanent.
|
|
||||||
func (l *Ledger) Check(client netip.Addr, now time.Time) (Ban, bool, bool) {
|
|
||||||
l.mu.Lock()
|
|
||||||
defer l.mu.Unlock()
|
|
||||||
|
|
||||||
ban := l.active(client, now)
|
|
||||||
if ban == nil {
|
|
||||||
return Ban{}, false, false
|
|
||||||
}
|
|
||||||
|
|
||||||
ban.Notes.Requests++
|
|
||||||
ban.Notes.Refused++
|
|
||||||
|
|
||||||
madePermanent := ban.Cause == CauseAttack && !ban.Permanent()
|
|
||||||
if madePermanent {
|
|
||||||
ban.Expires = time.Time{}
|
|
||||||
|
|
||||||
l.markChanged()
|
|
||||||
}
|
|
||||||
|
|
||||||
return *ban, true, madePermanent
|
|
||||||
}
|
|
||||||
|
|
||||||
// Find is Check without counting the request among those the ban
|
|
||||||
// refused, and without making the ban permanent: in observe mode a ban
|
|
||||||
// refuses nothing. The last result reports whether Check would have made
|
|
||||||
// the ban permanent.
|
|
||||||
func (l *Ledger) Find(client netip.Addr, now time.Time) (Ban, bool, bool) {
|
|
||||||
l.mu.Lock()
|
|
||||||
defer l.mu.Unlock()
|
|
||||||
|
|
||||||
ban := l.active(client, now)
|
|
||||||
if ban == nil {
|
|
||||||
return Ban{}, false, false
|
|
||||||
}
|
|
||||||
|
|
||||||
return *ban, true, ban.Cause == CauseAttack && !ban.Permanent()
|
|
||||||
}
|
|
||||||
|
|
||||||
// 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, and true. A first ban lasts LimitBanDuration. A ban
|
|
||||||
// made within LimitBanRepeatWindow after the netblock's ban that ended
|
|
||||||
// last, other than one for a clear sign of attack or a lifted one, 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 with false, and no other is made. The ledger fills in the
|
|
||||||
// notes' Refused and EarlierBans itself, and gives the ban the reason
|
|
||||||
// "requests per <Window> over the limit of <Limit>", from the notes.
|
|
||||||
func (l *Ledger) BanForLimit(
|
|
||||||
netblock netip.Prefix, now time.Time, notes Notes,
|
|
||||||
) (Ban, bool) {
|
|
||||||
return l.ban(netblock, now, CauseLimit, limitReason(notes), notes, true)
|
|
||||||
}
|
|
||||||
|
|
||||||
// WouldBanForLimit returns what BanForLimit would, without making the ban:
|
|
||||||
// what observe mode would have done.
|
|
||||||
func (l *Ledger) WouldBanForLimit(
|
|
||||||
netblock netip.Prefix, now time.Time, notes Notes,
|
|
||||||
) (Ban, bool) {
|
|
||||||
return l.ban(netblock, now, CauseLimit, limitReason(notes), notes, false)
|
|
||||||
}
|
|
||||||
|
|
||||||
// BanForAttack bans netblock at now for a clear sign of attack, with
|
|
||||||
// notes, and returns the ban, and whether it made it, as BanForLimit
|
|
||||||
// does. A first ban lasts AttackBanDuration; once the netblock has had
|
|
||||||
// one that was not lifted, the next is permanent. Its reason is "matched
|
|
||||||
// the rule <RuleID>".
|
|
||||||
func (l *Ledger) BanForAttack(
|
|
||||||
netblock netip.Prefix, now time.Time, notes Notes,
|
|
||||||
) (Ban, bool) {
|
|
||||||
return l.ban(netblock, now, CauseAttack, attackReason(notes), notes, true)
|
|
||||||
}
|
|
||||||
|
|
||||||
// WouldBanForAttack returns what BanForAttack would, without making the
|
|
||||||
// ban: what observe mode would have done.
|
|
||||||
func (l *Ledger) WouldBanForAttack(
|
|
||||||
netblock netip.Prefix, now time.Time, notes Notes,
|
|
||||||
) (Ban, bool) {
|
|
||||||
return l.ban(netblock, now, CauseAttack, attackReason(notes), notes, false)
|
|
||||||
}
|
|
||||||
|
|
||||||
// WouldBePermanent reports whether a ban on netblock for cause, CauseLimit
|
|
||||||
// or CauseAttack, made at now would be permanent, as BanForLimit or
|
|
||||||
// BanForAttack would make it. It works out nothing else of the ban.
|
|
||||||
func (l *Ledger) WouldBePermanent(
|
|
||||||
netblock netip.Prefix, now time.Time, cause string,
|
|
||||||
) bool {
|
|
||||||
l.mu.Lock()
|
|
||||||
defer l.mu.Unlock()
|
|
||||||
|
|
||||||
var held []Ban
|
|
||||||
if bans, found := l.netblocks.Peek(netblock); found {
|
|
||||||
held = *bans
|
|
||||||
}
|
|
||||||
|
|
||||||
if cause == CauseAttack {
|
|
||||||
return l.attackExpiry(held, now).IsZero()
|
|
||||||
}
|
|
||||||
|
|
||||||
return l.limitExpiry(held, now).IsZero()
|
|
||||||
}
|
|
||||||
|
|
||||||
// limitReason is the reason of a ban for a broken limit, with notes.
|
|
||||||
func limitReason(notes Notes) string {
|
|
||||||
return fmt.Sprintf("requests per %s over the limit of %d", notes.Window, notes.Limit)
|
|
||||||
}
|
|
||||||
|
|
||||||
// attackReason is the reason of a ban for a clear sign of attack, with
|
|
||||||
// notes.
|
|
||||||
func attackReason(notes Notes) string {
|
|
||||||
return "matched the rule " + notes.RuleID
|
|
||||||
}
|
|
||||||
|
|
||||||
// BanForAdmin bans netblock at now for an admin, with reason, until
|
|
||||||
// expires, or for good when expires is zero, and returns the ban, whose
|
|
||||||
// cause is CauseAdmin. Unlike BanForLimit and BanForAttack, it makes the
|
|
||||||
// ban even while another on netblock is active, since the admin asked
|
|
||||||
// for this one. The ledger fills in the notes' EarlierBans, and counts
|
|
||||||
// the ban among those made.
|
|
||||||
func (l *Ledger) BanForAdmin(
|
|
||||||
netblock netip.Prefix, now, expires time.Time, reason string,
|
|
||||||
) Ban {
|
|
||||||
l.mu.Lock()
|
|
||||||
defer l.mu.Unlock()
|
|
||||||
|
|
||||||
ban := Ban{
|
|
||||||
Netblock: netblock.Masked(), Start: now, Expires: expires, Cause: CauseAdmin,
|
|
||||||
Reason: reason,
|
|
||||||
}
|
|
||||||
|
|
||||||
held, found := l.netblocks.Get(ban.Netblock)
|
|
||||||
if found {
|
|
||||||
ban.Notes.EarlierBans = earlierBans(*held)
|
|
||||||
}
|
|
||||||
|
|
||||||
l.add(ban)
|
|
||||||
l.made[CauseAdmin]++
|
|
||||||
l.markChanged()
|
|
||||||
|
|
||||||
return ban
|
|
||||||
}
|
|
||||||
|
|
||||||
// Lift lifts, at now, every ban active then on a netblock client is in,
|
|
||||||
// as an admin does, and returns those bans. A lifted ban is kept, refuses
|
|
||||||
// nothing, and does not make the netblock's next ban longer.
|
|
||||||
func (l *Ledger) Lift(client netip.Addr, now time.Time) []Ban {
|
|
||||||
l.mu.Lock()
|
|
||||||
defer l.mu.Unlock()
|
|
||||||
|
|
||||||
var lifted []Ban
|
|
||||||
|
|
||||||
for _, bans := range l.covering(client) {
|
|
||||||
for i := range *bans {
|
|
||||||
ban := &(*bans)[i]
|
|
||||||
if ban.ActiveAt(now) {
|
|
||||||
ban.Lifted = now
|
|
||||||
lifted = append(lifted, *ban)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
if len(lifted) > 0 {
|
|
||||||
l.markChanged()
|
|
||||||
}
|
|
||||||
|
|
||||||
return lifted
|
|
||||||
}
|
|
||||||
|
|
||||||
// Covering returns every ban held on a netblock client is in, active or
|
|
||||||
// not, sorted by netblock, and each netblock's bans oldest first. It is
|
|
||||||
// not a request from client, and leaves when the netblocks were last seen
|
|
||||||
// unchanged.
|
|
||||||
func (l *Ledger) Covering(client netip.Addr) []Ban {
|
|
||||||
l.mu.Lock()
|
|
||||||
defer l.mu.Unlock()
|
|
||||||
|
|
||||||
var held []Ban
|
|
||||||
for _, bans := range l.covering(client) {
|
|
||||||
held = append(held, *bans...)
|
|
||||||
}
|
|
||||||
|
|
||||||
slices.SortStableFunc(held, func(a, b Ban) int {
|
|
||||||
return a.Netblock.Compare(b.Netblock)
|
|
||||||
})
|
|
||||||
|
|
||||||
return held
|
|
||||||
}
|
|
||||||
|
|
||||||
// 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)
|
|
||||||
}
|
|
||||||
|
|
||||||
// AddLookup gives the notes of netblock's bans that have no AS number, AS
|
|
||||||
// name or country yet those of a client in it, as the lookup answered
|
|
||||||
// about it. It is not a request from netblock, and leaves when it was last
|
|
||||||
// seen unchanged. It does not have bans.json written at once: the notes
|
|
||||||
// are written with its next write, as the counts in them are.
|
|
||||||
func (l *Ledger) AddLookup(netblock netip.Prefix, asn, asName, country string) {
|
|
||||||
l.mu.Lock()
|
|
||||||
defer l.mu.Unlock()
|
|
||||||
|
|
||||||
bans, found := l.netblocks.Peek(netblock)
|
|
||||||
if !found {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
for i := range *bans {
|
|
||||||
notes := &(*bans)[i].Notes
|
|
||||||
if notes.ASN == "" && notes.ASName == "" && notes.Country == "" {
|
|
||||||
notes.ASN, notes.ASName, notes.Country = asn, asName, country
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// Made returns how many bans for cause have been made since the start:
|
|
||||||
// for CauseLimit and CauseAttack, by the ledger; for CauseAdmin, by an
|
|
||||||
// admin, with BanForAdmin or in an edit of bans.json, as LoadEdit counts
|
|
||||||
// them. The bans read from bans.json at the start are not among them.
|
|
||||||
func (l *Ledger) Made(cause string) int {
|
|
||||||
l.mu.Lock()
|
|
||||||
defer l.mu.Unlock()
|
|
||||||
|
|
||||||
return l.made[cause]
|
|
||||||
}
|
|
||||||
|
|
||||||
// Count returns how many of the bans held are active at now, and how many
|
|
||||||
// of those are permanent. A lifted ban is neither.
|
|
||||||
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) {
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
|
|
||||||
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 at the start 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. A ban without a cause is an admin's, and gets CauseAdmin. 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 whose cause is not CauseAdmin are dropped, as
|
|
||||||
// when they are made.
|
|
||||||
func (l *Ledger) Load(bans []Ban) {
|
|
||||||
l.mu.Lock()
|
|
||||||
defer l.mu.Unlock()
|
|
||||||
|
|
||||||
l.load(bans)
|
|
||||||
}
|
|
||||||
|
|
||||||
// LoadEdit is Load for an admin's edit of bans.json, taken in while
|
|
||||||
// smallwebwaf runs. Each ban in it whose cause is CauseAdmin, and which
|
|
||||||
// the ledger did not hold, with the same netblock and start, is one the
|
|
||||||
// admin made, and is counted among the bans made.
|
|
||||||
func (l *Ledger) LoadEdit(bans []Ban) {
|
|
||||||
l.mu.Lock()
|
|
||||||
defer l.mu.Unlock()
|
|
||||||
|
|
||||||
l.made[CauseAdmin] += l.load(bans)
|
|
||||||
}
|
|
||||||
|
|
||||||
// load does what Load describes, and returns how many of bans are bans
|
|
||||||
// whose cause is CauseAdmin that the ledger did not hold before.
|
|
||||||
func (l *Ledger) load(bans []Ban) int {
|
|
||||||
bans = slices.Clone(bans)
|
|
||||||
added := 0
|
|
||||||
|
|
||||||
for i := range bans {
|
|
||||||
ban := &bans[i]
|
|
||||||
ban.Netblock = ban.Netblock.Masked()
|
|
||||||
ban.Notes.Request = ban.Notes.Request.cut()
|
|
||||||
|
|
||||||
if ban.Cause == "" {
|
|
||||||
ban.Cause = CauseAdmin
|
|
||||||
}
|
|
||||||
|
|
||||||
if ban.Cause == CauseAdmin && !l.holds(ban.Netblock, ban.Start) {
|
|
||||||
added++
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
slices.SortStableFunc(bans, func(a, b Ban) int {
|
|
||||||
return a.Start.Compare(b.Start)
|
|
||||||
})
|
|
||||||
|
|
||||||
l.netblocks.Purge()
|
|
||||||
l.held = 0
|
|
||||||
l.v4Lengths, l.v6Lengths = nil, nil
|
|
||||||
|
|
||||||
for _, ban := range bans {
|
|
||||||
l.add(ban)
|
|
||||||
}
|
|
||||||
|
|
||||||
return added
|
|
||||||
}
|
|
||||||
|
|
||||||
// holds reports whether the ledger holds a ban on netblock that started
|
|
||||||
// at start.
|
|
||||||
func (l *Ledger) holds(netblock netip.Prefix, start time.Time) bool {
|
|
||||||
bans, found := l.netblocks.Peek(netblock)
|
|
||||||
|
|
||||||
return found && slices.ContainsFunc(*bans, func(ban Ban) bool {
|
|
||||||
return ban.Start.Equal(start)
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
// ban bans netblock at now for cause, with reason and notes, as
|
|
||||||
// BanForLimit and BanForAttack describe, and returns the ban, and whether
|
|
||||||
// it made it. Unless keep is true, the ban is not made, only returned: it
|
|
||||||
// is the ban that would have been made.
|
|
||||||
func (l *Ledger) ban(
|
|
||||||
netblock netip.Prefix, now time.Time, cause, reason string, notes Notes, keep bool,
|
|
||||||
) (Ban, bool) {
|
|
||||||
l.mu.Lock()
|
|
||||||
defer l.mu.Unlock()
|
|
||||||
|
|
||||||
// held are the netblock's bans, none of them active.
|
|
||||||
var held []Ban
|
|
||||||
|
|
||||||
bans, found := l.netblocks.Get(netblock)
|
|
||||||
if found {
|
|
||||||
active := activeBan(*bans, now)
|
|
||||||
if active != nil {
|
|
||||||
return *active, false
|
|
||||||
}
|
|
||||||
|
|
||||||
held = *bans
|
|
||||||
notes.EarlierBans = earlierBans(held)
|
|
||||||
}
|
|
||||||
|
|
||||||
notes.Request = notes.Request.cut()
|
|
||||||
ban := Ban{Netblock: netblock, Start: now, Cause: cause, Reason: reason, Notes: notes}
|
|
||||||
|
|
||||||
if cause == CauseAttack {
|
|
||||||
ban.Expires = l.attackExpiry(held, now)
|
|
||||||
} else {
|
|
||||||
ban.Expires = l.limitExpiry(held, now)
|
|
||||||
}
|
|
||||||
|
|
||||||
if !keep {
|
|
||||||
return ban, true
|
|
||||||
}
|
|
||||||
|
|
||||||
l.add(ban)
|
|
||||||
l.made[cause]++
|
|
||||||
l.markChanged()
|
|
||||||
|
|
||||||
return ban, true
|
|
||||||
}
|
|
||||||
|
|
||||||
// earlierBans returns how many bans a netblock with the bans held, oldest
|
|
||||||
// first, has had, by cause: the first ban held counts the bans the
|
|
||||||
// netblock had before that one, since dropped to make room, and each ban
|
|
||||||
// held adds one.
|
|
||||||
func earlierBans(held []Ban) EarlierBans {
|
|
||||||
earlier := held[0].Notes.EarlierBans
|
|
||||||
|
|
||||||
for _, ban := range held {
|
|
||||||
switch ban.Cause {
|
|
||||||
case CauseLimit:
|
|
||||||
earlier.Limit++
|
|
||||||
case CauseAttack:
|
|
||||||
earlier.Attack++
|
|
||||||
case CauseAdmin:
|
|
||||||
earlier.Admin++
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
return earlier
|
|
||||||
}
|
|
||||||
|
|
||||||
// markChanged has Changed receive a value, unless one is waiting already.
|
|
||||||
func (l *Ledger) markChanged() {
|
|
||||||
select {
|
|
||||||
case l.changed <- struct{}{}:
|
|
||||||
default: // a value is waiting already
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// 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
|
|
||||||
}
|
|
||||||
|
|
||||||
// covering returns the bans of each netblock held that client is in,
|
|
||||||
// leaving when the netblocks were last seen unchanged.
|
|
||||||
func (l *Ledger) covering(client netip.Addr) []*[]Ban {
|
|
||||||
lengths := l.v6Lengths
|
|
||||||
if client.Is4() {
|
|
||||||
lengths = l.v4Lengths
|
|
||||||
}
|
|
||||||
|
|
||||||
var found []*[]Ban
|
|
||||||
|
|
||||||
for _, length := range lengths {
|
|
||||||
bans, ok := l.netblocks.Peek(netip.PrefixFrom(client, length).Masked())
|
|
||||||
if ok {
|
|
||||||
found = append(found, bans)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
return found
|
|
||||||
}
|
|
||||||
|
|
||||||
// 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,
|
|
||||||
// unless ban's cause is CauseAdmin, which does not count toward MaxBans.
|
|
||||||
func (l *Ledger) add(ban Ban) {
|
|
||||||
counted := ban.Cause != CauseAdmin
|
|
||||||
if counted && 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)
|
|
||||||
|
|
||||||
if counted {
|
|
||||||
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())
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// limitExpiry returns when a ban for a broken limit made at now ends, or
|
|
||||||
// zero when it is permanent. held are the netblock's bans, none of them
|
|
||||||
// active, of which the one that ended last, other than a ban for a clear
|
|
||||||
// sign of attack or a lifted one, can make the new ban longer. A ban an
|
|
||||||
// admin adds to bans.json can start after another and end before it, so
|
|
||||||
// that one is looked for among them all.
|
|
||||||
func (l *Ledger) limitExpiry(held []Ban, now time.Time) time.Time {
|
|
||||||
length := l.rules.LimitBanDuration
|
|
||||||
|
|
||||||
var last *Ban
|
|
||||||
|
|
||||||
for i, ban := range held {
|
|
||||||
if ban.Cause != CauseAttack && ban.Lifted.IsZero() &&
|
|
||||||
(last == nil || ban.Expires.After(last.Expires)) {
|
|
||||||
last = &held[i]
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
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)
|
|
||||||
}
|
|
||||||
|
|
||||||
// attackExpiry returns when a ban for a clear sign of attack made at now
|
|
||||||
// ends. held are the netblock's bans, none of them active: if one of them
|
|
||||||
// is for a clear sign of attack too, and was not lifted, the new ban is
|
|
||||||
// permanent, and its end zero; otherwise it ends AttackBanDuration later.
|
|
||||||
func (l *Ledger) attackExpiry(held []Ban, now time.Time) time.Time {
|
|
||||||
for _, ban := range held {
|
|
||||||
if ban.Cause == CauseAttack && ban.Lifted.IsZero() {
|
|
||||||
return time.Time{}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
return now.Add(l.rules.AttackBanDuration)
|
|
||||||
}
|
|
||||||
|
|
||||||
// dropOne drops the earliest ban whose cause is not CauseAdmin of the
|
|
||||||
// netblock that has gone longest without a request, of those that hold
|
|
||||||
// such a ban, and the netblock with it if that was its only ban. It is
|
|
||||||
// called with at least one such ban held. It looks at each netblock once
|
|
||||||
// at most, and drops nothing when none holds such a ban.
|
|
||||||
func (l *Ledger) dropOne() {
|
|
||||||
for range l.netblocks.Len() {
|
|
||||||
netblock, bans, _ := l.netblocks.GetOldest()
|
|
||||||
|
|
||||||
i := slices.IndexFunc(*bans, func(ban Ban) bool {
|
|
||||||
return ban.Cause != CauseAdmin
|
|
||||||
})
|
|
||||||
if i < 0 {
|
|
||||||
// Its bans are all an admin's, and never dropped. Get makes
|
|
||||||
// it the most recently seen, so that the next netblock is
|
|
||||||
// looked at; when it was seen matters only for dropping a
|
|
||||||
// ban, and a ban added to it makes it the most recently seen
|
|
||||||
// anyway.
|
|
||||||
l.netblocks.Get(netblock)
|
|
||||||
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
|
|
||||||
if len(*bans) == 1 {
|
|
||||||
l.netblocks.Remove(netblock)
|
|
||||||
} else {
|
|
||||||
*bans = slices.Delete(*bans, i, i+1)
|
|
||||||
}
|
|
||||||
|
|
||||||
l.held--
|
|
||||||
|
|
||||||
return
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// 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)])
|
|
||||||
}
|
|
||||||
@@ -1,522 +0,0 @@
|
|||||||
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 != (bans.EarlierBans{Limit: i}) {
|
|
||||||
t.Fatalf("ban %d lasts %s with earlier bans %+v, want %d hours and %d for a limit",
|
|
||||||
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 != (bans.EarlierBans{Limit: 1}) {
|
|
||||||
t.Errorf("second ban lasts %s with earlier bans %+v, want %s and 1 for a limit",
|
|
||||||
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, made := ledger.BanForLimit(netblock, midnight(), bans.Notes{})
|
|
||||||
if !made {
|
|
||||||
t.Error("the first ban was not made")
|
|
||||||
}
|
|
||||||
|
|
||||||
again, made := ledger.BanForLimit(netblock, midnight().Add(time.Minute), bans.Notes{})
|
|
||||||
|
|
||||||
if made || again != first || len(ledger.Bans(netblock)) != 1 {
|
|
||||||
t.Errorf("a limit broken during a ban gave %+v, made %t, and %d bans, "+
|
|
||||||
"want %+v, not made, and 1", again, made, len(ledger.Bans(netblock)), first)
|
|
||||||
}
|
|
||||||
|
|
||||||
again, made = ledger.BanForAttack(netblock, midnight().Add(time.Minute), bans.Notes{})
|
|
||||||
if made || again != first {
|
|
||||||
t.Errorf("an attack during a ban gave %+v, made %t, want %+v, not made",
|
|
||||||
again, made, 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 != (bans.EarlierBans{Limit: 1}) {
|
|
||||||
t.Errorf("the ledger holds %+v, want only the second ban, "+
|
|
||||||
"with 1 earlier ban for a limit", held)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestRequestDuringAnAttackBanMakesItPermanent(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
ledger := bans.New(defaultRules())
|
|
||||||
netblock := netip.MustParsePrefix("203.0.113.9/32")
|
|
||||||
notes := bans.Notes{RuleID: "env-file", Target: "path"}
|
|
||||||
|
|
||||||
ban, _ := ledger.BanForAttack(netblock, midnight(), notes)
|
|
||||||
if !ban.Expires.Equal(midnight().Add(7*day)) || ban.Cause != bans.CauseAttack ||
|
|
||||||
ban.Notes.RuleID != "env-file" || ledger.Made(bans.CauseAttack) != 1 ||
|
|
||||||
ledger.Made(bans.CauseLimit) != 0 {
|
|
||||||
t.Fatalf("the ban is %+v, with %d made for an attack and %d for a limit, "+
|
|
||||||
"want one for an attack, of seven days", ban,
|
|
||||||
ledger.Made(bans.CauseAttack), ledger.Made(bans.CauseLimit))
|
|
||||||
}
|
|
||||||
|
|
||||||
wantChanged(t, ledger, true)
|
|
||||||
|
|
||||||
// In observe mode the ban refuses nothing, and stays as it is, while
|
|
||||||
// Find tells that the request would have made it permanent.
|
|
||||||
got, _, wouldMakePermanent := ledger.Find(netblock.Addr(), midnight().Add(time.Hour))
|
|
||||||
if got.Permanent() || ledger.Bans(netblock)[0].Permanent() || !wouldMakePermanent {
|
|
||||||
t.Fatalf("a request found under the ban left it %+v, would have made it "+
|
|
||||||
"permanent %t, want it as it was, and true", got, wouldMakePermanent)
|
|
||||||
}
|
|
||||||
|
|
||||||
wantChanged(t, ledger, false)
|
|
||||||
|
|
||||||
// A request it refuses makes it permanent, says so, and makes
|
|
||||||
// bans.json due.
|
|
||||||
got, _, madePermanent := ledger.Check(netblock.Addr(), midnight().Add(time.Hour))
|
|
||||||
if !madePermanent || !got.Permanent() || !ledger.Bans(netblock)[0].Permanent() {
|
|
||||||
t.Fatalf("after a request during the ban, it is %+v, made permanent %t, "+
|
|
||||||
"want it made permanent", got, madePermanent)
|
|
||||||
}
|
|
||||||
|
|
||||||
wantChanged(t, ledger, true)
|
|
||||||
|
|
||||||
// The next request finds it permanent already.
|
|
||||||
_, banned, madePermanent := ledger.Check(netblock.Addr(), midnight().Add(100*365*day))
|
|
||||||
if !banned || madePermanent {
|
|
||||||
t.Errorf("a later request is banned %t, and made the ban permanent %t, "+
|
|
||||||
"want banned by the permanent ban", banned, madePermanent)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestAttackAfterAnAttackBanHasEndedBansPermanently(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
ledger := bans.New(defaultRules())
|
|
||||||
netblock := netip.MustParsePrefix("203.0.113.9/32")
|
|
||||||
|
|
||||||
// A ban for a broken limit before does not count.
|
|
||||||
first, _ := ledger.BanForLimit(netblock, midnight(), bans.Notes{})
|
|
||||||
second, _ := ledger.BanForAttack(netblock, first.Expires, bans.Notes{})
|
|
||||||
|
|
||||||
if second.Expires.Sub(second.Start) != 7*day {
|
|
||||||
t.Fatalf("the first ban for an attack lasts %s, want 7 days",
|
|
||||||
second.Expires.Sub(second.Start))
|
|
||||||
}
|
|
||||||
|
|
||||||
// Once that has run out without a request, the netblock is served, and
|
|
||||||
// its next clear sign of attack bans it for good.
|
|
||||||
_, banned, _ := ledger.Check(netblock.Addr(), second.Expires)
|
|
||||||
if banned {
|
|
||||||
t.Fatal("the ban did not end")
|
|
||||||
}
|
|
||||||
|
|
||||||
// Its notes show the earlier ban for an attack that makes it permanent,
|
|
||||||
// beside the one for a limit.
|
|
||||||
third, _ := ledger.BanForAttack(netblock, second.Expires.Add(30*day), bans.Notes{})
|
|
||||||
if !third.Permanent() ||
|
|
||||||
third.Notes.EarlierBans != (bans.EarlierBans{Limit: 1, Attack: 1}) {
|
|
||||||
t.Errorf("the next ban for an attack is %+v, want a permanent one, "+
|
|
||||||
"with 1 earlier ban for a limit and 1 for an attack", third)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestWouldBanGivesTheBanWithoutMakingIt(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
ledger := bans.New(defaultRules())
|
|
||||||
netblock := netip.MustParsePrefix("203.0.113.9/32")
|
|
||||||
first, _ := ledger.BanForLimit(netblock, midnight(), bans.Notes{})
|
|
||||||
wantChanged(t, ledger, true)
|
|
||||||
|
|
||||||
// While the first ban lasts, none would be made.
|
|
||||||
during, would := ledger.WouldBanForAttack(netblock, midnight(), bans.Notes{})
|
|
||||||
if would || during != first {
|
|
||||||
t.Errorf("during the first ban, would ban %t with %+v, want false with %+v",
|
|
||||||
would, during, first)
|
|
||||||
}
|
|
||||||
|
|
||||||
// As it ends, a clear sign of attack would ban for seven days, and a
|
|
||||||
// limit broken again for three hours, but neither is made.
|
|
||||||
limitNotes := bans.Notes{Limit: 1, Window: "minute"}
|
|
||||||
attack, wouldAttack := ledger.WouldBanForAttack(netblock, first.Expires,
|
|
||||||
bans.Notes{RuleID: "git-dir"})
|
|
||||||
limit, wouldLimit := ledger.WouldBanForLimit(netblock, first.Expires, limitNotes)
|
|
||||||
|
|
||||||
if !wouldAttack || !attack.Expires.Equal(first.Expires.Add(7*day)) ||
|
|
||||||
attack.Reason != "matched the rule git-dir" || !wouldLimit ||
|
|
||||||
!limit.Expires.Equal(first.Expires.Add(3*time.Hour)) ||
|
|
||||||
limit.Reason != "requests per minute over the limit of 1" {
|
|
||||||
t.Errorf("would ban with %+v and %+v, want seven days for the attack and "+
|
|
||||||
"three hours for the limit", attack, limit)
|
|
||||||
}
|
|
||||||
|
|
||||||
if len(ledger.Bans(netblock)) != 1 || ledger.Made(bans.CauseLimit) != 1 ||
|
|
||||||
ledger.Made(bans.CauseAttack) != 0 {
|
|
||||||
t.Errorf("the ledger holds %+v, want the first ban alone", ledger.Bans(netblock))
|
|
||||||
}
|
|
||||||
|
|
||||||
wantChanged(t, ledger, false)
|
|
||||||
|
|
||||||
// The ban made is the one that would have been.
|
|
||||||
made, _ := ledger.BanForLimit(netblock, first.Expires, limitNotes)
|
|
||||||
if made != limit {
|
|
||||||
t.Errorf("the ban made is %+v, want %+v", made, limit)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestWouldBePermanentAnswersAsTheBanWouldBeMade(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
ledger := bans.New(defaultRules())
|
|
||||||
netblock := netip.MustParsePrefix("203.0.113.9/32")
|
|
||||||
now := midnight()
|
|
||||||
|
|
||||||
// Five bans for a limit in a row, of 1, 3, 9, 27 and 81 hours, are not
|
|
||||||
// permanent. The sixth, of 243 hours, would be, while a first ban for
|
|
||||||
// an attack would not.
|
|
||||||
for i := range 5 {
|
|
||||||
if ledger.WouldBePermanent(netblock, now, bans.CauseLimit) {
|
|
||||||
t.Fatalf("ban %d for a limit would be permanent", i+1)
|
|
||||||
}
|
|
||||||
|
|
||||||
ban, _ := ledger.BanForLimit(netblock, now, bans.Notes{})
|
|
||||||
now = ban.Expires
|
|
||||||
}
|
|
||||||
|
|
||||||
if !ledger.WouldBePermanent(netblock, now, bans.CauseLimit) {
|
|
||||||
t.Error("the sixth ban for a limit would not be permanent")
|
|
||||||
}
|
|
||||||
|
|
||||||
if ledger.WouldBePermanent(netblock, now, bans.CauseAttack) {
|
|
||||||
t.Error("a first ban for an attack would be permanent")
|
|
||||||
}
|
|
||||||
|
|
||||||
// Once a first ban for an attack has ended, the next would be permanent.
|
|
||||||
attack, _ := ledger.BanForAttack(netblock, now, bans.Notes{})
|
|
||||||
if !ledger.WouldBePermanent(netblock, attack.Expires, bans.CauseAttack) {
|
|
||||||
t.Error("a second ban for an attack would not be permanent")
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestAttackBanDoesNotLengthenTheNextBanForALimit(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
ledger := bans.New(defaultRules())
|
|
||||||
netblock := netip.MustParsePrefix("203.0.113.9/32")
|
|
||||||
|
|
||||||
// Three times the seven days would be permanent; a limit broken as the
|
|
||||||
// ban for an attack ends bans for an hour, as a first broken limit does.
|
|
||||||
attack, _ := ledger.BanForAttack(netblock, midnight(), bans.Notes{})
|
|
||||||
limit, _ := ledger.BanForLimit(netblock, attack.Expires, bans.Notes{})
|
|
||||||
|
|
||||||
if limit.Expires.Sub(limit.Start) != time.Hour || limit.Cause != bans.CauseLimit {
|
|
||||||
t.Errorf("the ban for a limit is %+v, want one of an hour", limit)
|
|
||||||
}
|
|
||||||
|
|
||||||
// And a request during the ban for a limit leaves it as it is.
|
|
||||||
got, _, madePermanent := ledger.Check(netblock.Addr(), limit.Start)
|
|
||||||
if got.Permanent() || madePermanent {
|
|
||||||
t.Error("a request during a ban for a limit made it permanent")
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
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)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestLookupFillsTheNotesOfTheNetblocksBansWithoutOne(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
ledger := bans.New(defaultRules())
|
|
||||||
netblock := netip.MustParsePrefix("203.0.113.9/32")
|
|
||||||
other := netip.MustParsePrefix("198.51.100.7/32")
|
|
||||||
|
|
||||||
// A ban made with the client's lookup, one made before it came, after
|
|
||||||
// the first ended, and one on another netblock.
|
|
||||||
ledger.BanForLimit(netblock, midnight(), bans.Notes{
|
|
||||||
ASN: "AS64497", ASName: "Other Net", Country: "FR",
|
|
||||||
})
|
|
||||||
ledger.BanForLimit(netblock, midnight().Add(time.Hour), bans.Notes{})
|
|
||||||
ledger.BanForLimit(other, midnight(), bans.Notes{})
|
|
||||||
|
|
||||||
ledger.AddLookup(netblock, "AS64496", "Example Net", "DE")
|
|
||||||
|
|
||||||
held := ledger.Bans(netblock)
|
|
||||||
if len(held) != 2 {
|
|
||||||
t.Fatalf("%s has %d bans, want 2", netblock, len(held))
|
|
||||||
}
|
|
||||||
|
|
||||||
for i, want := range []bans.Notes{
|
|
||||||
{ASN: "AS64497", ASName: "Other Net", Country: "FR"},
|
|
||||||
{ASN: "AS64496", ASName: "Example Net", Country: "DE"},
|
|
||||||
} {
|
|
||||||
got := held[i].Notes
|
|
||||||
if got.ASN != want.ASN || got.ASName != want.ASName || got.Country != want.Country {
|
|
||||||
t.Errorf("ban %d's notes give %q, %q and %q, want %q, %q and %q", i+1,
|
|
||||||
got.ASN, got.ASName, got.Country, want.ASN, want.ASName, want.Country)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
if notes := ledger.Bans(other)[0].Notes; notes.ASN != "" || notes.Country != "" {
|
|
||||||
t.Errorf("the ban on %s has %q and %q, want neither",
|
|
||||||
other, notes.ASN, notes.Country)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// defaultRules are the rules at the settings' defaults.
|
|
||||||
func defaultRules() bans.Rules {
|
|
||||||
return bans.Rules{
|
|
||||||
LimitBanDuration: time.Hour,
|
|
||||||
LimitBanRepeatWindow: day,
|
|
||||||
MaxBanDuration: 7 * day,
|
|
||||||
AttackBanDuration: 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)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@@ -1,314 +0,0 @@
|
|||||||
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 != (bans.EarlierBans{Limit: 1}) {
|
|
||||||
t.Errorf("the next ban lasts %s with earlier bans %+v, want 3h and 1 for a limit",
|
|
||||||
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 cause and no notes.
|
|
||||||
netblock := netip.MustParsePrefix("203.0.113.9/32")
|
|
||||||
nineHours := bans.Ban{
|
|
||||||
Netblock: netblock,
|
|
||||||
Start: midnight(),
|
|
||||||
Expires: midnight().Add(9 * time.Hour),
|
|
||||||
Cause: bans.CauseLimit,
|
|
||||||
Notes: bans.Notes{EarlierBans: bans.EarlierBans{Limit: 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 and it, for a limit, 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 != (bans.EarlierBans{Limit: 3, Admin: 1}) {
|
|
||||||
t.Errorf("the next ban lasts %s with earlier bans %+v, "+
|
|
||||||
"want 27h, 3 for a limit and 1 an admin's",
|
|
||||||
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(),
|
|
||||||
Cause: bans.CauseLimit,
|
|
||||||
}
|
|
||||||
earlier := bans.Ban{
|
|
||||||
Netblock: netip.MustParsePrefix("203.0.113.2/32"),
|
|
||||||
Start: midnight().Add(-time.Hour),
|
|
||||||
Cause: bans.CauseLimit,
|
|
||||||
}
|
|
||||||
|
|
||||||
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(),
|
|
||||||
Cause: bans.CauseLimit,
|
|
||||||
}
|
|
||||||
ledger.Load([]bans.Ban{
|
|
||||||
{
|
|
||||||
Netblock: netip.MustParsePrefix("203.0.113.0/24"),
|
|
||||||
Start: midnight(),
|
|
||||||
Cause: bans.CauseLimit,
|
|
||||||
},
|
|
||||||
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)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
+12
-1005
File diff suppressed because it is too large
Load Diff
+32
-1167
File diff suppressed because it is too large
Load Diff
@@ -1,9 +0,0 @@
|
|||||||
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
|
|
||||||
}
|
|
||||||
@@ -1,233 +0,0 @@
|
|||||||
package lookup
|
|
||||||
|
|
||||||
import (
|
|
||||||
"context"
|
|
||||||
"fmt"
|
|
||||||
"log/slog"
|
|
||||||
"net/netip"
|
|
||||||
"os"
|
|
||||||
"path/filepath"
|
|
||||||
"sync"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
"github.com/fsnotify/fsnotify"
|
|
||||||
"github.com/oschwald/maxminddb-golang/v2"
|
|
||||||
|
|
||||||
"sneak.berlin/go/smallwebwaf/internal/alerts"
|
|
||||||
)
|
|
||||||
|
|
||||||
// quietTime is how long the lookup database must go without a change
|
|
||||||
// before it is read again, so that a file still being copied in is read
|
|
||||||
// only once whole.
|
|
||||||
const quietTime = 2 * time.Second
|
|
||||||
|
|
||||||
// FileParams are what OpenFile needs.
|
|
||||||
type FileParams struct {
|
|
||||||
// Path is the lookup database, the IPinfo Lite file in its .mmdb form
|
|
||||||
// (SWWAF_LOOKUP_DB_PATH).
|
|
||||||
Path string
|
|
||||||
// Now tells the time, normally time.Now.
|
|
||||||
Now func() time.Time
|
|
||||||
// ProcessLog receives each reading of the file, and why a replacement
|
|
||||||
// of it cannot be read.
|
|
||||||
ProcessLog *slog.Logger
|
|
||||||
// Alerts receive a file_error alert for each replacement that cannot
|
|
||||||
// be read.
|
|
||||||
Alerts *alerts.Queue
|
|
||||||
}
|
|
||||||
|
|
||||||
// File looks up clients' AS numbers and countries in the lookup database,
|
|
||||||
// held in memory, and reads it again when it is replaced. It is safe for
|
|
||||||
// concurrent use.
|
|
||||||
type File struct {
|
|
||||||
params FileParams
|
|
||||||
|
|
||||||
mu sync.Mutex
|
|
||||||
// reader is the database in use, and lastRead when it was read.
|
|
||||||
// readFailures are the replacements that could not be read.
|
|
||||||
reader *maxminddb.Reader
|
|
||||||
lastRead time.Time
|
|
||||||
readFailures int
|
|
||||||
}
|
|
||||||
|
|
||||||
// record is what the lookup database holds about a network, of the fields
|
|
||||||
// smallwebwaf reads.
|
|
||||||
type record struct {
|
|
||||||
ASN string `maxminddb:"asn"`
|
|
||||||
ASName string `maxminddb:"as_name"`
|
|
||||||
CountryCode string `maxminddb:"country_code"`
|
|
||||||
}
|
|
||||||
|
|
||||||
// OpenFile reads the lookup database. A file that cannot be read, or that
|
|
||||||
// is not a .mmdb file, is an error.
|
|
||||||
func OpenFile(params FileParams) (*File, error) {
|
|
||||||
reader, err := read(params.Path)
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
f := &File{params: params}
|
|
||||||
f.use(reader)
|
|
||||||
|
|
||||||
return f, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// LookUp returns what the lookup database says about client: its AS
|
|
||||||
// number, such as AS64496, the AS's name, and its country, such as DE,
|
|
||||||
// each "" when the database does not give it, as for an address missing
|
|
||||||
// from it. The database is asked about the client's first address, as
|
|
||||||
// GeoJS is.
|
|
||||||
func (f *File) LookUp(client netip.Prefix) Answer {
|
|
||||||
f.mu.Lock()
|
|
||||||
reader := f.reader
|
|
||||||
f.mu.Unlock()
|
|
||||||
|
|
||||||
var found record
|
|
||||||
|
|
||||||
// A record that cannot be decoded places the client nowhere, as a
|
|
||||||
// missing one does.
|
|
||||||
err := reader.Lookup(client.Addr()).Decode(&found)
|
|
||||||
if err != nil {
|
|
||||||
found = record{}
|
|
||||||
}
|
|
||||||
|
|
||||||
return Answer{
|
|
||||||
Client: client,
|
|
||||||
ASN: found.ASN,
|
|
||||||
ASName: found.ASName,
|
|
||||||
Country: found.CountryCode,
|
|
||||||
Answered: f.params.Now(),
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// LastRead returns when the lookup database in use was read.
|
|
||||||
func (f *File) LastRead() time.Time {
|
|
||||||
f.mu.Lock()
|
|
||||||
defer f.mu.Unlock()
|
|
||||||
|
|
||||||
return f.lastRead
|
|
||||||
}
|
|
||||||
|
|
||||||
// ReadFailures returns how many replacements of the lookup database could
|
|
||||||
// not be read.
|
|
||||||
func (f *File) ReadFailures() int {
|
|
||||||
f.mu.Lock()
|
|
||||||
defer f.mu.Unlock()
|
|
||||||
|
|
||||||
return f.readFailures
|
|
||||||
}
|
|
||||||
|
|
||||||
// Watch watches the directory of the lookup database until ctx is done,
|
|
||||||
// and reads the file again once it has gone without a change for
|
|
||||||
// quietTime, after it is replaced, written or removed, and after Watch
|
|
||||||
// starts watching. If the directory cannot be watched, that is logged, and
|
|
||||||
// the database read at start stays in use.
|
|
||||||
func (f *File) Watch(ctx context.Context) {
|
|
||||||
watcher, err := fsnotify.NewWatcher()
|
|
||||||
if err == nil {
|
|
||||||
defer func() {
|
|
||||||
_ = watcher.Close()
|
|
||||||
}()
|
|
||||||
|
|
||||||
err = watcher.Add(filepath.Dir(f.params.Path))
|
|
||||||
}
|
|
||||||
|
|
||||||
if err != nil {
|
|
||||||
f.params.ProcessLog.Error("cannot watch the lookup database for replacements",
|
|
||||||
"error", err.Error())
|
|
||||||
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
f.params.ProcessLog.Info("watching the lookup database for replacements",
|
|
||||||
"file", f.params.Path)
|
|
||||||
|
|
||||||
f.readAfterChanges(ctx, watcher.Events, watcher.Errors)
|
|
||||||
}
|
|
||||||
|
|
||||||
// readAfterChanges reads the lookup database again once quietTime has
|
|
||||||
// passed without a change to it from events, until ctx is done, and logs
|
|
||||||
// the errors from errs. A change to another file in its directory does not
|
|
||||||
// count. The wait starts at once, as if for a change, so that a file
|
|
||||||
// replaced after OpenFile read it, and before its directory was watched,
|
|
||||||
// is read too.
|
|
||||||
func (f *File) readAfterChanges(
|
|
||||||
ctx context.Context, events <-chan fsnotify.Event, errs <-chan error,
|
|
||||||
) {
|
|
||||||
path := filepath.Clean(f.params.Path)
|
|
||||||
|
|
||||||
quiet := time.NewTimer(quietTime)
|
|
||||||
defer quiet.Stop()
|
|
||||||
|
|
||||||
for {
|
|
||||||
select {
|
|
||||||
case <-ctx.Done():
|
|
||||||
return
|
|
||||||
case event := <-events:
|
|
||||||
if filepath.Clean(event.Name) == path {
|
|
||||||
quiet.Reset(quietTime)
|
|
||||||
}
|
|
||||||
case <-quiet.C:
|
|
||||||
f.readAgain()
|
|
||||||
case err := <-errs:
|
|
||||||
f.params.ProcessLog.Warn("watching the lookup database failed",
|
|
||||||
"error", err.Error())
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// readAgain reads the lookup database again, in place of the one in use,
|
|
||||||
// or, if it cannot be read, counts that, raises a file_error alert for it
|
|
||||||
// and logs it, and the one in use stays in use.
|
|
||||||
func (f *File) readAgain() {
|
|
||||||
reader, err := read(f.params.Path)
|
|
||||||
if err != nil {
|
|
||||||
const kept = "the lookup database cannot be read, " +
|
|
||||||
"and the one read before stays in use"
|
|
||||||
|
|
||||||
f.mu.Lock()
|
|
||||||
f.readFailures++
|
|
||||||
f.mu.Unlock()
|
|
||||||
|
|
||||||
// Raised before it is logged, so that the alert is there once the
|
|
||||||
// log line is.
|
|
||||||
f.params.Alerts.Raise(alerts.Alert{
|
|
||||||
Event: alerts.EventFileError,
|
|
||||||
Reason: kept,
|
|
||||||
Detail: map[string]any{"file": f.params.Path, "error": err.Error()},
|
|
||||||
})
|
|
||||||
f.params.ProcessLog.Error(kept, "error", err.Error())
|
|
||||||
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
f.use(reader)
|
|
||||||
}
|
|
||||||
|
|
||||||
// use puts reader in use, in place of the database read before, and logs
|
|
||||||
// that the file was read.
|
|
||||||
func (f *File) use(reader *maxminddb.Reader) {
|
|
||||||
f.mu.Lock()
|
|
||||||
f.reader = reader
|
|
||||||
f.lastRead = f.params.Now()
|
|
||||||
f.mu.Unlock()
|
|
||||||
|
|
||||||
f.params.ProcessLog.Info("read the lookup database", "file", f.params.Path)
|
|
||||||
}
|
|
||||||
|
|
||||||
// read reads the lookup database at path. The whole file is read into
|
|
||||||
// memory, rather than mapped into it as the reader can, so that a file
|
|
||||||
// overwritten in place cannot change, or end, under a lookup.
|
|
||||||
func read(path string) (*maxminddb.Reader, error) {
|
|
||||||
data, err := os.ReadFile(path) //nolint:gosec // the file the admin names
|
|
||||||
if err != nil {
|
|
||||||
return nil, fmt.Errorf("SWWAF_LOOKUP_DB_PATH cannot be read: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
reader, err := maxminddb.OpenBytes(data)
|
|
||||||
if err != nil {
|
|
||||||
return nil, fmt.Errorf("SWWAF_LOOKUP_DB_PATH %s is not a .mmdb file: %w", path, err)
|
|
||||||
}
|
|
||||||
|
|
||||||
return reader, nil
|
|
||||||
}
|
|
||||||
@@ -1,379 +0,0 @@
|
|||||||
package lookup
|
|
||||||
|
|
||||||
import (
|
|
||||||
"context"
|
|
||||||
"log/slog"
|
|
||||||
"net/netip"
|
|
||||||
"net/url"
|
|
||||||
"os"
|
|
||||||
"path/filepath"
|
|
||||||
"reflect"
|
|
||||||
"testing"
|
|
||||||
"testing/synctest"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
"github.com/fsnotify/fsnotify"
|
|
||||||
"github.com/maxmind/mmdbwriter/mmdbtype"
|
|
||||||
|
|
||||||
"sneak.berlin/go/smallwebwaf/internal/alerts"
|
|
||||||
"sneak.berlin/go/smallwebwaf/internal/lookup/lookuptest"
|
|
||||||
)
|
|
||||||
|
|
||||||
// testNetblock is the netblock the tests' lookup databases place, and
|
|
||||||
// testClient a client in it.
|
|
||||||
const (
|
|
||||||
testNetblock = "203.0.113.0/24"
|
|
||||||
testClient = "203.0.113.9/32"
|
|
||||||
)
|
|
||||||
|
|
||||||
func TestFilePlacesClientsAndCountsAnAddressMissingFromItAsUnknown(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
germany := lookuptest.Network{ASN: "AS64496", ASName: "Example Net", Country: "DE"}
|
|
||||||
northKorea := lookuptest.Network{ASN: "AS64511", ASName: "Other Net", Country: "KP"}
|
|
||||||
path := filepath.Join(t.TempDir(), "ipinfo_lite.mmdb")
|
|
||||||
lookuptest.Write(t, path, map[string]lookuptest.Network{
|
|
||||||
testNetblock: germany,
|
|
||||||
"2001:db8::/32": northKorea,
|
|
||||||
})
|
|
||||||
|
|
||||||
now := time.Date(2026, 10, 7, 0, 0, 0, 0, time.UTC)
|
|
||||||
|
|
||||||
f, err := OpenFile(FileParams{
|
|
||||||
Path: path,
|
|
||||||
Now: func() time.Time { return now },
|
|
||||||
ProcessLog: slog.New(slog.DiscardHandler),
|
|
||||||
Alerts: newQueue(),
|
|
||||||
})
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("open %s: %v", path, err)
|
|
||||||
}
|
|
||||||
|
|
||||||
for client, want := range map[string]lookuptest.Network{
|
|
||||||
testClient: germany,
|
|
||||||
// An IPv6 client is its /64.
|
|
||||||
"2001:db8:1:2::/64": northKorea,
|
|
||||||
"198.51.100.7/32": {},
|
|
||||||
} {
|
|
||||||
prefix := netip.MustParsePrefix(client)
|
|
||||||
|
|
||||||
got := f.LookUp(prefix)
|
|
||||||
if got != (Answer{
|
|
||||||
Client: prefix, ASN: want.ASN, ASName: want.ASName, Country: want.Country,
|
|
||||||
Answered: now,
|
|
||||||
}) {
|
|
||||||
t.Errorf("%s has the answer %+v, want %+v, answered %s", client, got, want, now)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestRecordThatCannotBeReadPlacesTheClientNowhere(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
// The AS number is a number, where a string belongs. The writer writes
|
|
||||||
// a record's fields in the order of their names, so as_name is read
|
|
||||||
// before the AS number fails.
|
|
||||||
path := filepath.Join(t.TempDir(), "ipinfo_lite.mmdb")
|
|
||||||
lookuptest.WriteRecords(t, path, map[string]mmdbtype.Map{
|
|
||||||
testNetblock: {
|
|
||||||
"asn": mmdbtype.Uint32(64496),
|
|
||||||
"as_name": mmdbtype.String("Example Net"),
|
|
||||||
"country_code": mmdbtype.String("DE"),
|
|
||||||
},
|
|
||||||
})
|
|
||||||
|
|
||||||
f := openFile(t, path, newQueue())
|
|
||||||
|
|
||||||
answer := f.LookUp(netip.MustParsePrefix(testClient))
|
|
||||||
if answer.ASN != "" || answer.ASName != "" || answer.Country != "" {
|
|
||||||
t.Errorf("%s is placed %+v, want nowhere", testClient, answer)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestFileThatCannotBeReadIsAnError(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
dir := t.TempDir()
|
|
||||||
missing := filepath.Join(dir, "missing.mmdb")
|
|
||||||
notDatabase := filepath.Join(dir, "not.mmdb")
|
|
||||||
writeFile(t, notDatabase, "not a lookup database\n")
|
|
||||||
|
|
||||||
for path, want := range map[string]string{
|
|
||||||
missing: "SWWAF_LOOKUP_DB_PATH cannot be read: open " + missing +
|
|
||||||
": no such file or directory",
|
|
||||||
notDatabase: "SWWAF_LOOKUP_DB_PATH " + notDatabase +
|
|
||||||
" is not a .mmdb file: error opening database: invalid MaxMind DB file",
|
|
||||||
} {
|
|
||||||
_, err := OpenFile(FileParams{
|
|
||||||
Path: path,
|
|
||||||
Now: time.Now,
|
|
||||||
ProcessLog: slog.New(slog.DiscardHandler),
|
|
||||||
Alerts: newQueue(),
|
|
||||||
})
|
|
||||||
if err == nil || err.Error() != want {
|
|
||||||
t.Errorf("opening %s failed with %v, want %s", path, err, want)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// The tests below run readAfterChanges 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 readAfterChanges waits again, so that every
|
|
||||||
// reading due by then is done. The test sends the changes itself, as the
|
|
||||||
// watch of a directory cannot run in a bubble.
|
|
||||||
|
|
||||||
func TestReplacementCopiedOverTheFileInTwoPartsIsReadOnlyWhole(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
synctest.Test(t, func(t *testing.T) {
|
|
||||||
dir := t.TempDir()
|
|
||||||
path := filepath.Join(dir, "ipinfo_lite.mmdb")
|
|
||||||
writeDatabase(t, path, "DE")
|
|
||||||
|
|
||||||
queue := newQueue()
|
|
||||||
f := openFile(t, path, queue)
|
|
||||||
changes := watch(t, f)
|
|
||||||
|
|
||||||
other := filepath.Join(dir, "replacement.mmdb")
|
|
||||||
writeDatabase(t, other, "KP")
|
|
||||||
|
|
||||||
replacement, err := os.ReadFile(other) //nolint:gosec // a file the test wrote
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("read %s: %v", other, err)
|
|
||||||
}
|
|
||||||
|
|
||||||
// The file in use is overwritten in place, and keeps giving what
|
|
||||||
// it gave. Its first part alone is not a .mmdb file.
|
|
||||||
file, err := os.Create(path) //nolint:gosec // a file the test wrote
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("create %s: %v", path, err)
|
|
||||||
}
|
|
||||||
|
|
||||||
defer func() {
|
|
||||||
_ = file.Close()
|
|
||||||
}()
|
|
||||||
|
|
||||||
half := len(replacement) / 2
|
|
||||||
write(t, file, replacement[:half])
|
|
||||||
|
|
||||||
changes <- fsnotify.Event{Name: path, Op: fsnotify.Write}
|
|
||||||
|
|
||||||
time.Sleep(quietTime - time.Nanosecond)
|
|
||||||
synctest.Wait()
|
|
||||||
wantCountry(t, f, "DE")
|
|
||||||
|
|
||||||
// The second part starts the wait again.
|
|
||||||
write(t, file, replacement[half:])
|
|
||||||
|
|
||||||
changes <- fsnotify.Event{Name: path, Op: fsnotify.Write}
|
|
||||||
|
|
||||||
time.Sleep(quietTime - time.Nanosecond)
|
|
||||||
synctest.Wait()
|
|
||||||
wantCountry(t, f, "DE")
|
|
||||||
|
|
||||||
time.Sleep(time.Nanosecond)
|
|
||||||
synctest.Wait()
|
|
||||||
wantCountry(t, f, "KP")
|
|
||||||
|
|
||||||
if !f.LastRead().Equal(time.Now()) || f.ReadFailures() != 0 {
|
|
||||||
t.Errorf("read at %s, with %d failures; want read now, with none",
|
|
||||||
f.LastRead(), f.ReadFailures())
|
|
||||||
}
|
|
||||||
|
|
||||||
wantAlerts(t, queue)
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestReplacementThatCannotBeReadLeavesTheFileInUseWithOneAlert(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
synctest.Test(t, func(t *testing.T) {
|
|
||||||
path := filepath.Join(t.TempDir(), "ipinfo_lite.mmdb")
|
|
||||||
writeDatabase(t, path, "DE")
|
|
||||||
|
|
||||||
queue := newQueue()
|
|
||||||
f := openFile(t, path, queue)
|
|
||||||
read := f.LastRead()
|
|
||||||
changes := watch(t, f)
|
|
||||||
|
|
||||||
writeFile(t, path, "not a lookup database\n")
|
|
||||||
|
|
||||||
changes <- fsnotify.Event{Name: path, Op: fsnotify.Write}
|
|
||||||
|
|
||||||
// Long after, the replacement has been read once.
|
|
||||||
time.Sleep(time.Hour)
|
|
||||||
synctest.Wait()
|
|
||||||
wantCountry(t, f, "DE")
|
|
||||||
|
|
||||||
if !f.LastRead().Equal(read) || f.ReadFailures() != 1 {
|
|
||||||
t.Errorf("read at %s, with %d failures; want read at %s, with one",
|
|
||||||
f.LastRead(), f.ReadFailures(), read)
|
|
||||||
}
|
|
||||||
|
|
||||||
wantAlerts(t, queue, alerts.Alert{
|
|
||||||
Time: read.Add(quietTime),
|
|
||||||
Event: alerts.EventFileError,
|
|
||||||
Reason: "the lookup database cannot be read, and the one read before stays in use",
|
|
||||||
Detail: map[string]any{
|
|
||||||
"file": path,
|
|
||||||
"error": "SWWAF_LOOKUP_DB_PATH " + path + " is not a .mmdb file: " +
|
|
||||||
"error opening database: invalid MaxMind DB file",
|
|
||||||
},
|
|
||||||
})
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestChangeOfAnotherFileInTheDirectoryIsNoReplacement(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
synctest.Test(t, func(t *testing.T) {
|
|
||||||
dir := t.TempDir()
|
|
||||||
path := filepath.Join(dir, "ipinfo_lite.mmdb")
|
|
||||||
writeDatabase(t, path, "DE")
|
|
||||||
|
|
||||||
f := openFile(t, path, newQueue())
|
|
||||||
changes := watch(t, f)
|
|
||||||
|
|
||||||
// The wait that starts with the watch ends with a reading.
|
|
||||||
time.Sleep(quietTime)
|
|
||||||
synctest.Wait()
|
|
||||||
writeDatabase(t, path, "KP")
|
|
||||||
|
|
||||||
changes <- fsnotify.Event{Name: filepath.Join(dir, "other.mmdb"), Op: fsnotify.Create}
|
|
||||||
|
|
||||||
time.Sleep(quietTime)
|
|
||||||
synctest.Wait()
|
|
||||||
wantCountry(t, f, "DE")
|
|
||||||
|
|
||||||
changes <- fsnotify.Event{Name: path, Op: fsnotify.Write}
|
|
||||||
|
|
||||||
time.Sleep(quietTime)
|
|
||||||
synctest.Wait()
|
|
||||||
wantCountry(t, f, "KP")
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestReplacementSavedBeforeTheWatchStartsIsRead(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
synctest.Test(t, func(t *testing.T) {
|
|
||||||
path := filepath.Join(t.TempDir(), "ipinfo_lite.mmdb")
|
|
||||||
writeDatabase(t, path, "DE")
|
|
||||||
|
|
||||||
f := openFile(t, path, newQueue())
|
|
||||||
|
|
||||||
// Saved after OpenFile read the file, and before its directory was
|
|
||||||
// watched, so that no change is seen for it.
|
|
||||||
writeDatabase(t, path, "KP")
|
|
||||||
watch(t, f)
|
|
||||||
time.Sleep(quietTime)
|
|
||||||
synctest.Wait()
|
|
||||||
wantCountry(t, f, "KP")
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
// newQueue returns a queue of alerts for a webhook that is never sent
|
|
||||||
// them, so that they wait in it for the test to look at.
|
|
||||||
func newQueue() *alerts.Queue {
|
|
||||||
return alerts.New(alerts.Params{
|
|
||||||
WebhookURL: &url.URL{Scheme: "https", Host: "alerts.example"},
|
|
||||||
Events: alerts.Events(),
|
|
||||||
Cooldown: 15 * time.Minute,
|
|
||||||
Now: time.Now,
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
// writeDatabase writes a lookup database at path that places testNetblock
|
|
||||||
// in country, and no other address.
|
|
||||||
func writeDatabase(t *testing.T, path, country string) {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
lookuptest.Write(t, path, map[string]lookuptest.Network{
|
|
||||||
testNetblock: {ASN: "AS64496", ASName: "Example Net", Country: country},
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
// openFile opens the lookup database at path, which raises its alerts to
|
|
||||||
// queue.
|
|
||||||
func openFile(t *testing.T, path string, queue *alerts.Queue) *File {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
f, err := OpenFile(FileParams{
|
|
||||||
Path: path,
|
|
||||||
Now: time.Now,
|
|
||||||
ProcessLog: slog.New(slog.DiscardHandler),
|
|
||||||
Alerts: queue,
|
|
||||||
})
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("open %s: %v", path, err)
|
|
||||||
}
|
|
||||||
|
|
||||||
return f
|
|
||||||
}
|
|
||||||
|
|
||||||
// watch runs f's readAfterChanges until the test ends, and returns the
|
|
||||||
// channel that sends it changes.
|
|
||||||
func watch(t *testing.T, f *File) chan<- fsnotify.Event {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
changes := make(chan fsnotify.Event)
|
|
||||||
ctx, stop := context.WithCancel(t.Context())
|
|
||||||
stopped := make(chan struct{})
|
|
||||||
|
|
||||||
go func() {
|
|
||||||
f.readAfterChanges(ctx, changes, nil)
|
|
||||||
close(stopped)
|
|
||||||
}()
|
|
||||||
|
|
||||||
t.Cleanup(func() {
|
|
||||||
stop()
|
|
||||||
<-stopped
|
|
||||||
})
|
|
||||||
|
|
||||||
return changes
|
|
||||||
}
|
|
||||||
|
|
||||||
// wantCountry checks the country f gives testClient.
|
|
||||||
func wantCountry(t *testing.T, f *File, want string) {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
got := f.LookUp(netip.MustParsePrefix(testClient)).Country
|
|
||||||
if got != want {
|
|
||||||
t.Errorf("%s is in %q, want %q", testClient, got, want)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// wantAlerts checks the alerts waiting in queue, and that it held none
|
|
||||||
// back.
|
|
||||||
func wantAlerts(t *testing.T, queue *alerts.Queue, want ...alerts.Alert) {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
waiting := queue.Snapshot().Waiting[alerts.DestinationWebhook]
|
|
||||||
if len(waiting) != len(want) || (len(want) > 0 && !reflect.DeepEqual(waiting, want)) {
|
|
||||||
t.Errorf("alerts waiting %+v, want %+v", waiting, want)
|
|
||||||
}
|
|
||||||
|
|
||||||
if queue.Suppressed() != 0 {
|
|
||||||
t.Errorf("%d alerts held back, want none", queue.Suppressed())
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// writeFile writes content to the file at path.
|
|
||||||
func writeFile(t *testing.T, path, content string) {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
err := os.WriteFile(path, []byte(content), 0o600)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("write %s: %v", path, err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// write writes data to the end of file.
|
|
||||||
func write(t *testing.T, file *os.File, data []byte) {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
_, err := file.Write(data)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("write: %v", err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@@ -1,531 +0,0 @@
|
|||||||
// Package lookup looks up each client's AS number and country, through
|
|
||||||
// the GeoJS web service or in the lookup database, the IPinfo Lite file
|
|
||||||
// SWWAF_LOOKUP_DB_PATH names. GeoJS's answers are kept in memory, for at
|
|
||||||
// most 100,000 clients and for 7 days each, and 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"
|
|
||||||
"strconv"
|
|
||||||
"strings"
|
|
||||||
"sync"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
"github.com/hashicorp/golang-lru/v2/simplelru"
|
|
||||||
"sneak.berlin/go/smallwebwaf/internal/alerts"
|
|
||||||
"sneak.berlin/go/smallwebwaf/internal/metrics"
|
|
||||||
)
|
|
||||||
|
|
||||||
// URL is GeoJS's endpoint for an address's place and network. 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/geo.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
|
|
||||||
// unknownASN is the AS number GeoJS gives when it knows none.
|
|
||||||
unknownASN = 64512
|
|
||||||
// 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
|
|
||||||
// Timeout is how long a request waits for its client's first answer,
|
|
||||||
// and how long a request to GeoJS may take before it is abandoned
|
|
||||||
// (SWWAF_LOOKUP_TIMEOUT).
|
|
||||||
Timeout time.Duration
|
|
||||||
// Wait is true when a setting needs each request's answer before the
|
|
||||||
// request goes on. Otherwise no request waits for one.
|
|
||||||
Wait bool
|
|
||||||
// Answered, unless nil, is given each answer GeoJS gives, once it is
|
|
||||||
// kept.
|
|
||||||
Answered func(Answer)
|
|
||||||
// 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
|
|
||||||
// Alerts receive a source_failure alert each time GeoJS fails.
|
|
||||||
Alerts *alerts.Queue
|
|
||||||
}
|
|
||||||
|
|
||||||
// GeoJS looks up clients' AS numbers and 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
|
|
||||||
timeout time.Duration
|
|
||||||
wait bool
|
|
||||||
answered func(Answer)
|
|
||||||
now func() time.Time
|
|
||||||
processLog *slog.Logger
|
|
||||||
metrics *metrics.Metrics
|
|
||||||
alerts *alerts.Queue
|
|
||||||
// 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 or the lookup database said about a client: its AS
|
|
||||||
// number, such as AS64496, and the AS's name, both "" when the source knows
|
|
||||||
// no AS number for it; its country, "" when the source cannot place it;
|
|
||||||
// when the source said so; and, for GeoJS's answers, which lookups.json
|
|
||||||
// holds, when the answer was last used. The zero Answer is that of a
|
|
||||||
// client with no answer.
|
|
||||||
//
|
|
||||||
//nolint:tagliatelle // the state files use snake_case, as the request log does
|
|
||||||
type Answer struct {
|
|
||||||
Client netip.Prefix `json:"client"`
|
|
||||||
ASN string `json:"asn"`
|
|
||||||
ASName string `json:"as_name"`
|
|
||||||
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,
|
|
||||||
timeout: params.Timeout,
|
|
||||||
wait: params.Wait,
|
|
||||||
answered: params.Answered,
|
|
||||||
now: params.Now,
|
|
||||||
processLog: params.ProcessLog,
|
|
||||||
metrics: params.Metrics,
|
|
||||||
alerts: params.Alerts,
|
|
||||||
httpClient: &http.Client{
|
|
||||||
CheckRedirect: func(*http.Request, []*http.Request) error {
|
|
||||||
return http.ErrUseLastResponse
|
|
||||||
},
|
|
||||||
},
|
|
||||||
answers: answers,
|
|
||||||
waiting: map[netip.Prefix]*wait{},
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// LookUp returns the answer GeoJS gave about client, with its country as
|
|
||||||
// a two-letter code in capitals, or the zero Answer when there is none
|
|
||||||
// yet. An answer is kept for 7 days. Without one, the client is asked
|
|
||||||
// about in the background, and, while Wait is set, the request waits up
|
|
||||||
// to Timeout for the answer, unless the client has gone without one
|
|
||||||
// before. 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) LookUp(ctx context.Context, client netip.Prefix) Answer {
|
|
||||||
answer, asked := g.answerOrWait(ctx, client)
|
|
||||||
if asked == nil {
|
|
||||||
return answer
|
|
||||||
}
|
|
||||||
|
|
||||||
timer := time.NewTimer(g.timeout)
|
|
||||||
defer timer.Stop()
|
|
||||||
|
|
||||||
select {
|
|
||||||
case <-asked:
|
|
||||||
case <-timer.C:
|
|
||||||
case <-ctx.Done():
|
|
||||||
}
|
|
||||||
|
|
||||||
g.mu.Lock()
|
|
||||||
defer g.mu.Unlock()
|
|
||||||
|
|
||||||
answer, found := g.kept(client)
|
|
||||||
if !found {
|
|
||||||
g.metrics.GeoJSUnanswered.Inc()
|
|
||||||
}
|
|
||||||
|
|
||||||
w, waiting := g.waiting[client]
|
|
||||||
if !found && waiting {
|
|
||||||
w.late = true
|
|
||||||
}
|
|
||||||
|
|
||||||
return answer
|
|
||||||
}
|
|
||||||
|
|
||||||
// Kept returns client's answer, if one is kept, without asking GeoJS.
|
|
||||||
func (g *GeoJS) Kept(client netip.Prefix) (Answer, bool) {
|
|
||||||
g.mu.Lock()
|
|
||||||
defer g.mu.Unlock()
|
|
||||||
|
|
||||||
return g.kept(client)
|
|
||||||
}
|
|
||||||
|
|
||||||
// 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,
|
|
||||||
) (Answer, <-chan struct{}) {
|
|
||||||
g.mu.Lock()
|
|
||||||
defer g.mu.Unlock()
|
|
||||||
|
|
||||||
answer, found := g.kept(client)
|
|
||||||
if found {
|
|
||||||
return answer, 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 !g.wait {
|
|
||||||
return Answer{}, nil // the answer is not needed before the request goes on
|
|
||||||
}
|
|
||||||
|
|
||||||
if w == nil {
|
|
||||||
g.metrics.GeoJSUnanswered.Inc()
|
|
||||||
|
|
||||||
return Answer{}, 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 Answer{}, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
return Answer{}, 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) (Answer, bool) {
|
|
||||||
now := g.now()
|
|
||||||
|
|
||||||
kept, found := g.answers.Get(client)
|
|
||||||
if !found || now.Sub(kept.Answered) >= keepFor {
|
|
||||||
return Answer{}, false
|
|
||||||
}
|
|
||||||
|
|
||||||
kept.Used = now
|
|
||||||
|
|
||||||
return *kept, 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. Each answer kept is given to
|
|
||||||
// Answered, outside the lock, since Answered takes locks of its own.
|
|
||||||
func (g *GeoJS) askAboutWaiting(ctx context.Context) {
|
|
||||||
for {
|
|
||||||
clients := g.nextClients()
|
|
||||||
if len(clients) == 0 {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
given, err := g.request(ctx, clients)
|
|
||||||
kept, answered := g.keep(clients, given, err)
|
|
||||||
|
|
||||||
if g.answered != nil {
|
|
||||||
for _, answer := range kept {
|
|
||||||
g.answered(answer)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
if !answered {
|
|
||||||
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, given being the
|
|
||||||
// answer for each address GeoJS's answer names. It returns the answers it
|
|
||||||
// kept, and reports whether GeoJS answered about all of the clients. Each
|
|
||||||
// client whose address GeoJS's answer names gets its answer. 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, given map[netip.Addr]Answer, err error,
|
|
||||||
) ([]Answer, bool) {
|
|
||||||
g.mu.Lock()
|
|
||||||
defer g.mu.Unlock()
|
|
||||||
|
|
||||||
now := g.now()
|
|
||||||
kept := make([]Answer, 0, len(clients))
|
|
||||||
leftOut := 0
|
|
||||||
|
|
||||||
for _, client := range clients {
|
|
||||||
answer, named := given[client.Addr()]
|
|
||||||
if !named {
|
|
||||||
leftOut++
|
|
||||||
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
|
|
||||||
answer.Client, answer.Answered, answer.Used = client, now, now
|
|
||||||
g.answers.Add(client, &answer)
|
|
||||||
kept = append(kept, answer)
|
|
||||||
|
|
||||||
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())
|
|
||||||
g.alerts.Raise(alerts.Alert{
|
|
||||||
Event: alerts.EventSourceFailure,
|
|
||||||
Reason: "asking GeoJS failed",
|
|
||||||
Detail: map[string]any{
|
|
||||||
"source": "geojs", "error": err.Error(),
|
|
||||||
"asking_again_in": g.retryDelay.String(),
|
|
||||||
},
|
|
||||||
})
|
|
||||||
|
|
||||||
return kept, false
|
|
||||||
}
|
|
||||||
|
|
||||||
g.retryDelay = 0
|
|
||||||
|
|
||||||
return kept, true
|
|
||||||
}
|
|
||||||
|
|
||||||
// request asks GeoJS about clients in one request, and returns the answer
|
|
||||||
// for each address GeoJS's answer names: its AS number and the AS's name,
|
|
||||||
// both "" for the AS number 64512, which GeoJS gives when it knows none,
|
|
||||||
// and its country, in capitals.
|
|
||||||
func (g *GeoJS) request(
|
|
||||||
ctx context.Context, clients []netip.Prefix,
|
|
||||||
) (map[netip.Addr]Answer, error) {
|
|
||||||
addrs := make([]string, 0, len(clients))
|
|
||||||
|
|
||||||
for _, client := range clients {
|
|
||||||
addrs = append(addrs, client.Addr().String())
|
|
||||||
}
|
|
||||||
|
|
||||||
ctx, cancel := context.WithTimeout(ctx, g.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)
|
|
||||||
}
|
|
||||||
|
|
||||||
//nolint:tagliatelle // GeoJS's own names
|
|
||||||
var answers []struct {
|
|
||||||
IP string `json:"ip"`
|
|
||||||
ASN int64 `json:"asn"`
|
|
||||||
ASName string `json:"organization_name"`
|
|
||||||
CountryCode string `json:"country_code"`
|
|
||||||
}
|
|
||||||
|
|
||||||
err = json.NewDecoder(io.LimitReader(res.Body, maxResponseBytes)).Decode(&answers)
|
|
||||||
if err != nil {
|
|
||||||
return nil, fmt.Errorf("read GeoJS's answer: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
given := make(map[netip.Addr]Answer, len(answers))
|
|
||||||
|
|
||||||
for _, item := range answers {
|
|
||||||
addr, err := netip.ParseAddr(item.IP)
|
|
||||||
if err != nil {
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
|
|
||||||
answer := Answer{Country: strings.ToUpper(item.CountryCode)}
|
|
||||||
if item.ASN != 0 && item.ASN != unknownASN {
|
|
||||||
answer.ASN = "AS" + strconv.FormatInt(item.ASN, 10)
|
|
||||||
answer.ASName = item.ASName
|
|
||||||
}
|
|
||||||
|
|
||||||
given[addr] = answer
|
|
||||||
}
|
|
||||||
|
|
||||||
return given, nil
|
|
||||||
}
|
|
||||||
@@ -1,862 +0,0 @@
|
|||||||
package lookup_test
|
|
||||||
|
|
||||||
import (
|
|
||||||
"encoding/json"
|
|
||||||
"log/slog"
|
|
||||||
"net/http"
|
|
||||||
"net/http/httptest"
|
|
||||||
"net/netip"
|
|
||||||
"net/url"
|
|
||||||
"reflect"
|
|
||||||
"slices"
|
|
||||||
"strings"
|
|
||||||
"sync"
|
|
||||||
"testing"
|
|
||||||
"testing/synctest"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
"github.com/prometheus/client_golang/prometheus/testutil"
|
|
||||||
"sneak.berlin/go/smallwebwaf/internal/alerts"
|
|
||||||
"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, and asNumber, kept as asn, and asName the AS it gives them.
|
|
||||||
germany = "DE"
|
|
||||||
asNumber = 64496
|
|
||||||
asn = "AS64496"
|
|
||||||
asName = "Example Net"
|
|
||||||
// unplaced is the address it cannot place, for which it gives the AS
|
|
||||||
// number 64512 and the AS name Unknown, as GeoJS does.
|
|
||||||
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.LookUp(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 TestRequestWaitsAsLongAsTheTimeoutSaysAndGeoJSIsAbandonedAfterIt(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
synctest.Test(t, func(t *testing.T) {
|
|
||||||
// A timeout longer than the default second, and a GeoJS that does
|
|
||||||
// not answer.
|
|
||||||
const longerTimeout = 3 * time.Second
|
|
||||||
|
|
||||||
m := metrics.New(1, "app")
|
|
||||||
g := lookup.New(lookup.Params{
|
|
||||||
URL: lookup.URL,
|
|
||||||
Timeout: longerTimeout,
|
|
||||||
Wait: true,
|
|
||||||
Now: time.Now,
|
|
||||||
ProcessLog: slog.New(slog.DiscardHandler),
|
|
||||||
Metrics: m,
|
|
||||||
Alerts: alerts.New(alerts.Params{}),
|
|
||||||
})
|
|
||||||
g.SetTransport(&standIn{answers: hanging})
|
|
||||||
|
|
||||||
var (
|
|
||||||
request sync.WaitGroup
|
|
||||||
waited time.Duration
|
|
||||||
)
|
|
||||||
|
|
||||||
request.Go(func() {
|
|
||||||
began := time.Now()
|
|
||||||
|
|
||||||
wantCountry(t, g, netip.MustParsePrefix("203.0.113.9/32"), "")
|
|
||||||
|
|
||||||
waited = time.Since(began)
|
|
||||||
})
|
|
||||||
|
|
||||||
// A moment before the timeout runs out, GeoJS is still being asked:
|
|
||||||
// the request to it has not failed.
|
|
||||||
time.Sleep(longerTimeout - time.Millisecond)
|
|
||||||
synctest.Wait()
|
|
||||||
wantFailures(t, m, 0)
|
|
||||||
|
|
||||||
// As it runs out, the client's request goes on, and the request to
|
|
||||||
// GeoJS is abandoned, which counts as a failure.
|
|
||||||
request.Wait()
|
|
||||||
synctest.Wait()
|
|
||||||
|
|
||||||
if waited != longerTimeout {
|
|
||||||
t.Errorf("waited %s for the answer, want %s", waited, longerTimeout)
|
|
||||||
}
|
|
||||||
|
|
||||||
wantFailures(t, m, 1)
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
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 TestAnswerHoldsTheASNumberTheASNameAndTheCountry(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")
|
|
||||||
now := clock.Now()
|
|
||||||
|
|
||||||
// For the client it cannot place, GeoJS gives the AS number 64512
|
|
||||||
// and the AS name Unknown, which count as unknown.
|
|
||||||
for client, want := range map[netip.Prefix]lookup.Answer{
|
|
||||||
placed: {
|
|
||||||
Client: placed, ASN: asn, ASName: asName, Country: germany,
|
|
||||||
Answered: now, Used: now,
|
|
||||||
},
|
|
||||||
notPlaced: {Client: notPlaced, Answered: now, Used: now},
|
|
||||||
} {
|
|
||||||
got := g.LookUp(t.Context(), client)
|
|
||||||
if got != want {
|
|
||||||
t.Errorf("answer for %s\n%+v\nwant\n%+v", client, got, want)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestWithoutWaitTheRequestGoesOnAtOnceAndTheAnswerIsGivenWhenItComes(
|
|
||||||
t *testing.T,
|
|
||||||
) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
synctest.Test(t, func(t *testing.T) {
|
|
||||||
var (
|
|
||||||
mu sync.Mutex
|
|
||||||
given []lookup.Answer
|
|
||||||
)
|
|
||||||
|
|
||||||
geojs := &standIn{answers: answeringSlowly}
|
|
||||||
clock := newClock()
|
|
||||||
m := metrics.New(1, "app")
|
|
||||||
g := lookup.New(lookup.Params{
|
|
||||||
URL: lookup.URL,
|
|
||||||
Timeout: timeout,
|
|
||||||
Answered: func(answer lookup.Answer) {
|
|
||||||
mu.Lock()
|
|
||||||
defer mu.Unlock()
|
|
||||||
|
|
||||||
given = append(given, answer)
|
|
||||||
},
|
|
||||||
Now: clock.Now,
|
|
||||||
ProcessLog: slog.New(slog.DiscardHandler),
|
|
||||||
Metrics: m,
|
|
||||||
Alerts: alerts.New(alerts.Params{}),
|
|
||||||
})
|
|
||||||
g.SetTransport(geojs)
|
|
||||||
|
|
||||||
client := netip.MustParsePrefix("203.0.113.9/32")
|
|
||||||
|
|
||||||
// The request goes on at once, without an answer, and GeoJS is asked
|
|
||||||
// about the client, which it answers most of a second later.
|
|
||||||
began := time.Now()
|
|
||||||
|
|
||||||
got := g.LookUp(t.Context(), client)
|
|
||||||
if took := time.Since(began); took != 0 || got != (lookup.Answer{}) {
|
|
||||||
t.Errorf("waited %s for %+v, want no wait and no answer", took, got)
|
|
||||||
}
|
|
||||||
|
|
||||||
waitForRequests(t, geojs, 1)
|
|
||||||
wantAsked(t, geojs, 0, "203.0.113.9")
|
|
||||||
|
|
||||||
time.Sleep(timeout)
|
|
||||||
synctest.Wait()
|
|
||||||
|
|
||||||
now := clock.Now()
|
|
||||||
want := lookup.Answer{
|
|
||||||
Client: client, ASN: asn, ASName: asName, Country: germany,
|
|
||||||
Answered: now, Used: now,
|
|
||||||
}
|
|
||||||
|
|
||||||
mu.Lock()
|
|
||||||
if !slices.Equal(given, []lookup.Answer{want}) {
|
|
||||||
t.Errorf("answers given %+v, want only %+v", given, want)
|
|
||||||
}
|
|
||||||
mu.Unlock()
|
|
||||||
|
|
||||||
wantCountry(t, g, client, germany)
|
|
||||||
wantUnanswered(t, m, 0)
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
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,
|
|
||||||
Timeout: timeout,
|
|
||||||
Wait: true,
|
|
||||||
Now: time.Now,
|
|
||||||
ProcessLog: slog.New(slog.NewTextHandler(&log, nil)),
|
|
||||||
Metrics: metrics.New(1, "app"),
|
|
||||||
Alerts: alerts.New(alerts.Params{}),
|
|
||||||
})
|
|
||||||
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 TestFailureRaisesASourceFailureAlertOncePerCooldown(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
synctest.Test(t, func(t *testing.T) {
|
|
||||||
geojs, clock, g, queue := startWithAlerts()
|
|
||||||
clients := newClients()
|
|
||||||
|
|
||||||
geojs.set(failing)
|
|
||||||
|
|
||||||
wantCountry(t, g, clients(), "")
|
|
||||||
|
|
||||||
want := alerts.Alert{
|
|
||||||
Time: clock.Now(),
|
|
||||||
Event: alerts.EventSourceFailure,
|
|
||||||
Reason: "asking GeoJS failed",
|
|
||||||
Detail: map[string]any{
|
|
||||||
"source": "geojs",
|
|
||||||
"error": "GeoJS answered 503 Service Unavailable",
|
|
||||||
"asking_again_in": "1s",
|
|
||||||
},
|
|
||||||
}
|
|
||||||
|
|
||||||
// The next failure, a second later, is a repeat within the
|
|
||||||
// cooldown.
|
|
||||||
clock.advance(time.Second)
|
|
||||||
wantCountry(t, g, clients(), "")
|
|
||||||
wantRequests(t, geojs, 2)
|
|
||||||
|
|
||||||
waiting := queue.Snapshot().Waiting[alerts.DestinationWebhook]
|
|
||||||
if len(waiting) != 1 || !reflect.DeepEqual(waiting[0], want) {
|
|
||||||
t.Errorf("alerts waiting %+v, want only %+v", waiting, want)
|
|
||||||
}
|
|
||||||
|
|
||||||
if queue.Suppressed() != 1 {
|
|
||||||
t.Errorf("%d alerts held back, want the repeat", queue.Suppressed())
|
|
||||||
}
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
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, "app")
|
|
||||||
g := lookup.New(lookup.Params{
|
|
||||||
URL: lookup.URL,
|
|
||||||
Timeout: timeout,
|
|
||||||
Wait: true,
|
|
||||||
Now: time.Now,
|
|
||||||
ProcessLog: slog.New(slog.DiscardHandler),
|
|
||||||
Metrics: m,
|
|
||||||
Alerts: alerts.New(alerts.Params{}),
|
|
||||||
})
|
|
||||||
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]any, 0, len(addrs))
|
|
||||||
|
|
||||||
for _, addr := range addrs {
|
|
||||||
item := map[string]any{
|
|
||||||
"ip": addr, "asn": asNumber, "organization_name": asName,
|
|
||||||
"country_code": germany,
|
|
||||||
}
|
|
||||||
|
|
||||||
switch {
|
|
||||||
case addr == unplaced:
|
|
||||||
item = map[string]any{"ip": addr, "asn": 64512, "organization_name": "Unknown"}
|
|
||||||
case addr == leftOut && answers == answeringWithoutLeftOut:
|
|
||||||
continue
|
|
||||||
case answers == answeringInLowerCase:
|
|
||||||
item["country_code"] = strings.ToLower(germany)
|
|
||||||
}
|
|
||||||
|
|
||||||
list = append(list, item)
|
|
||||||
}
|
|
||||||
|
|
||||||
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, for which a request waits for its
|
|
||||||
// client's first answer.
|
|
||||||
func start() (*standIn, *testClock, *lookup.GeoJS) {
|
|
||||||
geojs, clock, g, _ := startWithAlerts()
|
|
||||||
|
|
||||||
return geojs, clock, g
|
|
||||||
}
|
|
||||||
|
|
||||||
// startWithAlerts is start, and returns the queue of the alerts GeoJS
|
|
||||||
// raises as well, for a webhook that is never sent them, with the default
|
|
||||||
// cooldown, by the same clock.
|
|
||||||
func startWithAlerts() (*standIn, *testClock, *lookup.GeoJS, *alerts.Queue) {
|
|
||||||
geojs := &standIn{}
|
|
||||||
clock := newClock()
|
|
||||||
queue := alerts.New(alerts.Params{
|
|
||||||
WebhookURL: &url.URL{Scheme: "https", Host: "alerts.example"},
|
|
||||||
Events: alerts.Events(),
|
|
||||||
Cooldown: 15 * time.Minute,
|
|
||||||
Now: clock.Now,
|
|
||||||
})
|
|
||||||
g := lookup.New(lookup.Params{
|
|
||||||
URL: lookup.URL,
|
|
||||||
Timeout: timeout,
|
|
||||||
Wait: true,
|
|
||||||
Now: clock.Now,
|
|
||||||
ProcessLog: slog.New(slog.DiscardHandler),
|
|
||||||
Metrics: metrics.New(1, "app"),
|
|
||||||
Alerts: queue,
|
|
||||||
})
|
|
||||||
g.SetTransport(geojs)
|
|
||||||
|
|
||||||
return geojs, clock, g, queue
|
|
||||||
}
|
|
||||||
|
|
||||||
// newClock returns a clock set to the start of a day.
|
|
||||||
func newClock() *testClock {
|
|
||||||
return &testClock{now: time.Date(2026, 10, 4, 0, 0, 0, 0, time.UTC)}
|
|
||||||
}
|
|
||||||
|
|
||||||
// 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.LookUp(t.Context(), client).Country
|
|
||||||
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)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// wantFailures checks how many requests to GeoJS m counts as failed.
|
|
||||||
func wantFailures(t *testing.T, m *metrics.Metrics, want float64) {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
got := testutil.ToFloat64(m.GeoJSFailures)
|
|
||||||
if got != want {
|
|
||||||
t.Errorf("%v requests to GeoJS failed, 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)
|
|
||||||
}
|
|
||||||
@@ -1,83 +0,0 @@
|
|||||||
// Package lookuptest writes lookup databases, IPinfo Lite files in their
|
|
||||||
// .mmdb form, for the tests of the packages that read them.
|
|
||||||
package lookuptest
|
|
||||||
|
|
||||||
import (
|
|
||||||
"bytes"
|
|
||||||
"net"
|
|
||||||
"os"
|
|
||||||
"testing"
|
|
||||||
|
|
||||||
"github.com/maxmind/mmdbwriter"
|
|
||||||
"github.com/maxmind/mmdbwriter/mmdbtype"
|
|
||||||
)
|
|
||||||
|
|
||||||
// fileMode is the mode of the files written: read and written by their
|
|
||||||
// owner alone.
|
|
||||||
const fileMode = 0o600
|
|
||||||
|
|
||||||
// Network is what a lookup database holds about a netblock, of the fields
|
|
||||||
// smallwebwaf reads: its AS number, such as AS64496, the AS's name, and
|
|
||||||
// its country, such as DE.
|
|
||||||
type Network struct {
|
|
||||||
ASN string
|
|
||||||
ASName string
|
|
||||||
Country string
|
|
||||||
}
|
|
||||||
|
|
||||||
// Write writes a lookup database at path that places each netblock in
|
|
||||||
// networks, such as 203.0.113.0/24, as its Network says, and no other
|
|
||||||
// address.
|
|
||||||
func Write(tb testing.TB, path string, networks map[string]Network) {
|
|
||||||
tb.Helper()
|
|
||||||
|
|
||||||
records := make(map[string]mmdbtype.Map, len(networks))
|
|
||||||
for netblock, network := range networks {
|
|
||||||
records[netblock] = mmdbtype.Map{
|
|
||||||
"asn": mmdbtype.String(network.ASN),
|
|
||||||
"as_name": mmdbtype.String(network.ASName),
|
|
||||||
"country_code": mmdbtype.String(network.Country),
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
WriteRecords(tb, path, records)
|
|
||||||
}
|
|
||||||
|
|
||||||
// WriteRecords writes a lookup database at path that holds each record in
|
|
||||||
// records for its netblock, and nothing for any other address.
|
|
||||||
func WriteRecords(tb testing.TB, path string, records map[string]mmdbtype.Map) {
|
|
||||||
tb.Helper()
|
|
||||||
|
|
||||||
tree, err := mmdbwriter.New(mmdbwriter.Options{
|
|
||||||
DatabaseType: "ipinfo_lite",
|
|
||||||
// The tests' clients are in the netblocks kept for documentation.
|
|
||||||
IncludeReservedNetworks: true,
|
|
||||||
})
|
|
||||||
if err != nil {
|
|
||||||
tb.Fatalf("new lookup database: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
for netblock, record := range records {
|
|
||||||
_, network, err := net.ParseCIDR(netblock)
|
|
||||||
if err != nil {
|
|
||||||
tb.Fatalf("netblock %q: %v", netblock, err)
|
|
||||||
}
|
|
||||||
|
|
||||||
err = tree.Insert(network, record)
|
|
||||||
if err != nil {
|
|
||||||
tb.Fatalf("insert %s: %v", netblock, err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
var database bytes.Buffer
|
|
||||||
|
|
||||||
_, err = tree.WriteTo(&database)
|
|
||||||
if err != nil {
|
|
||||||
tb.Fatalf("write the lookup database: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
err = os.WriteFile(path, database.Bytes(), fileMode)
|
|
||||||
if err != nil {
|
|
||||||
tb.Fatalf("write %s: %v", path, err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@@ -1,102 +0,0 @@
|
|||||||
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, Answered: asked, Used: asked},
|
|
||||||
{
|
|
||||||
Client: placed, ASN: asn, ASName: asName, 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)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@@ -1,155 +0,0 @@
|
|||||||
package metrics
|
|
||||||
|
|
||||||
import (
|
|
||||||
"sync"
|
|
||||||
|
|
||||||
"github.com/prometheus/client_golang/prometheus"
|
|
||||||
"sneak.berlin/go/smallwebwaf/internal/requestlog"
|
|
||||||
)
|
|
||||||
|
|
||||||
// other is the label value under which the countries or AS numbers
|
|
||||||
// outside the busiest are counted.
|
|
||||||
const other = "other"
|
|
||||||
|
|
||||||
// busiest are the metrics by one thing the lookup finds of the client,
|
|
||||||
// its country or its AS number, for requests whose client's is known.
|
|
||||||
// The topN busiest countries or AS numbers, 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. One 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 busiest struct {
|
|
||||||
topN int
|
|
||||||
|
|
||||||
requests *prometheus.CounterVec
|
|
||||||
requestBytes *prometheus.CounterVec
|
|
||||||
responseBytes *prometheus.CounterVec
|
|
||||||
// refused are, by country, the requests the country lists refused; nil
|
|
||||||
// by AS number.
|
|
||||||
refused *prometheus.CounterVec
|
|
||||||
|
|
||||||
mu sync.Mutex
|
|
||||||
// seen is each country's or AS number's requests since the start, by
|
|
||||||
// which they are ranked.
|
|
||||||
seen map[string]int64
|
|
||||||
// top are those 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) *busiest {
|
|
||||||
countries := newBusiest(topN, "country", "the client's country")
|
|
||||||
countries.refused = counterVec("smallwebwaf_country_list_refusals_total",
|
|
||||||
"Requests the country lists refused, by the client's country.",
|
|
||||||
[]string{"country"})
|
|
||||||
|
|
||||||
return countries
|
|
||||||
}
|
|
||||||
|
|
||||||
// newASNs returns the metrics by AS number, with series of their own for
|
|
||||||
// the topN busiest AS numbers.
|
|
||||||
func newASNs(topN int) *busiest {
|
|
||||||
return newBusiest(topN, "asn", "the client's AS number")
|
|
||||||
}
|
|
||||||
|
|
||||||
// newBusiest returns the metrics by label, which is described as
|
|
||||||
// description, with series of their own for the topN busiest values.
|
|
||||||
func newBusiest(topN int, label, description string) *busiest {
|
|
||||||
by := []string{label}
|
|
||||||
|
|
||||||
return &busiest{
|
|
||||||
topN: topN,
|
|
||||||
requests: counterVec("smallwebwaf_"+label+"_requests_total",
|
|
||||||
"Requests, by "+description+".", by),
|
|
||||||
requestBytes: counterVec("smallwebwaf_"+label+"_request_bytes_total",
|
|
||||||
"Request body bytes, by "+description+".", by),
|
|
||||||
responseBytes: counterVec("smallwebwaf_"+label+"_response_bytes_total",
|
|
||||||
"Response body bytes, by "+description+".", by),
|
|
||||||
seen: map[string]int64{},
|
|
||||||
top: map[string]bool{},
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// Describe and Collect make the metrics a prometheus.Collector, so that
|
|
||||||
// they are registered together.
|
|
||||||
func (b *busiest) Describe(ch chan<- *prometheus.Desc) {
|
|
||||||
for _, vec := range b.vecs() {
|
|
||||||
vec.Describe(ch)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// Collect is the other half of prometheus.Collector, with Describe.
|
|
||||||
func (b *busiest) Collect(ch chan<- prometheus.Metric) {
|
|
||||||
for _, vec := range b.vecs() {
|
|
||||||
vec.Collect(ch)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// vecs returns the metrics: by AS number, those of requests and bytes; by
|
|
||||||
// country, the refusals by the country lists as well.
|
|
||||||
func (b *busiest) vecs() []*prometheus.CounterVec {
|
|
||||||
vecs := []*prometheus.CounterVec{b.requests, b.requestBytes, b.responseBytes}
|
|
||||||
if b.refused != nil {
|
|
||||||
vecs = append(vecs, b.refused)
|
|
||||||
}
|
|
||||||
|
|
||||||
return vecs
|
|
||||||
}
|
|
||||||
|
|
||||||
// add counts a request from its log line, whose client's country or AS
|
|
||||||
// number, value, is known.
|
|
||||||
func (b *busiest) add(value string, line *requestlog.Line) {
|
|
||||||
b.mu.Lock()
|
|
||||||
defer b.mu.Unlock()
|
|
||||||
|
|
||||||
b.seen[value]++
|
|
||||||
|
|
||||||
label := b.label(value)
|
|
||||||
b.requests.WithLabelValues(label).Inc()
|
|
||||||
b.requestBytes.WithLabelValues(label).Add(float64(line.RequestBytes))
|
|
||||||
b.responseBytes.WithLabelValues(label).Add(float64(line.ResponseBytes))
|
|
||||||
|
|
||||||
if b.refused != nil && line.Action == requestlog.ActionCountryDenied {
|
|
||||||
b.refused.WithLabelValues(label).Inc()
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// label returns the label a request from value is counted under: value
|
|
||||||
// while it is one of the busiest, other while it is not. A value busier
|
|
||||||
// than the least busy of them takes its place, and that one's series are
|
|
||||||
// dropped.
|
|
||||||
func (b *busiest) label(value string) string {
|
|
||||||
if b.top[value] {
|
|
||||||
return value
|
|
||||||
}
|
|
||||||
|
|
||||||
if len(b.top) < b.topN {
|
|
||||||
b.top[value] = true
|
|
||||||
|
|
||||||
return value
|
|
||||||
}
|
|
||||||
|
|
||||||
least := ""
|
|
||||||
|
|
||||||
for top := range b.top {
|
|
||||||
if least == "" || b.seen[top] < b.seen[least] {
|
|
||||||
least = top
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
if b.seen[value] <= b.seen[least] {
|
|
||||||
return other
|
|
||||||
}
|
|
||||||
|
|
||||||
delete(b.top, least)
|
|
||||||
|
|
||||||
for _, vec := range b.vecs() {
|
|
||||||
vec.DeleteLabelValues(least)
|
|
||||||
}
|
|
||||||
|
|
||||||
b.top[value] = true
|
|
||||||
|
|
||||||
return value
|
|
||||||
}
|
|
||||||
@@ -1,409 +0,0 @@
|
|||||||
// 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/alerts"
|
|
||||||
"sneak.berlin/go/smallwebwaf/internal/bans"
|
|
||||||
"sneak.berlin/go/smallwebwaf/internal/ratelimit"
|
|
||||||
"sneak.berlin/go/smallwebwaf/internal/remotelog"
|
|
||||||
"sneak.berlin/go/smallwebwaf/internal/requestlog"
|
|
||||||
"sneak.berlin/go/smallwebwaf/internal/rules"
|
|
||||||
)
|
|
||||||
|
|
||||||
// Metrics are smallwebwaf's metrics. They are safe for concurrent use.
|
|
||||||
type Metrics struct {
|
|
||||||
// registry gives every metric registered with it the label instance.
|
|
||||||
registry prometheus.Registerer
|
|
||||||
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
|
|
||||||
// ruleMatches are made by AddRules.
|
|
||||||
ruleMatches *prometheus.CounterVec
|
|
||||||
countries *busiest
|
|
||||||
asns *busiest
|
|
||||||
|
|
||||||
// GeoJSRequests are the requests to GeoJS, and GeoJSFailures those
|
|
||||||
// that failed. GeoJSUnanswered are the requests that needed their
|
|
||||||
// client's answer, for a setting that acts on it, and went on without
|
|
||||||
// it because GeoJS had not given 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 and how many AS numbers get series of their
|
|
||||||
// own (SWWAF_METRICS_TOP_N). Every metric carries instanceName
|
|
||||||
// (SWWAF_INSTANCE_NAME) as its label instance.
|
|
||||||
func New(topN int, instanceName string) *Metrics {
|
|
||||||
byStatus := []string{"status_class", "action"}
|
|
||||||
byFile := []string{"file"}
|
|
||||||
registry := prometheus.NewRegistry()
|
|
||||||
|
|
||||||
m := &Metrics{
|
|
||||||
registry: prometheus.WrapRegistererWith(
|
|
||||||
prometheus.Labels{"instance": instanceName}, registry),
|
|
||||||
handler: promhttp.HandlerFor(registry, promhttp.HandlerOpts{}),
|
|
||||||
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),
|
|
||||||
asns: newASNs(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 that needed their client's answer from GeoJS and " +
|
|
||||||
"went on without it, because GeoJS had not given 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 <name>.bad because they did not parse.",
|
|
||||||
byFile),
|
|
||||||
}
|
|
||||||
|
|
||||||
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, m.asns,
|
|
||||||
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,
|
|
||||||
// by cause, 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,
|
|
||||||
) {
|
|
||||||
for _, cause := range []string{bans.CauseLimit, bans.CauseAttack, bans.CauseAdmin} {
|
|
||||||
m.registry.MustRegister(prometheus.NewCounterFunc(prometheus.CounterOpts{
|
|
||||||
Name: "smallwebwaf_bans_made_total",
|
|
||||||
Help: "Bans made, by cause.",
|
|
||||||
ConstLabels: prometheus.Labels{"cause": cause},
|
|
||||||
}, func() float64 {
|
|
||||||
return float64(ledger.Made(cause))
|
|
||||||
}))
|
|
||||||
}
|
|
||||||
|
|
||||||
m.registry.MustRegister(
|
|
||||||
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 not lifted.",
|
|
||||||
}, 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())
|
|
||||||
}),
|
|
||||||
)
|
|
||||||
}
|
|
||||||
|
|
||||||
// AddRules adds the metrics of the rule files: the requests that matched
|
|
||||||
// each rule, which RuleMatched counts, and the rules loaded from
|
|
||||||
// ruleFiles, read as the metrics are asked for. It is called once, before
|
|
||||||
// RuleMatched.
|
|
||||||
func (m *Metrics) AddRules(ruleFiles *rules.Files) {
|
|
||||||
m.ruleMatches = counterVec("smallwebwaf_rule_matches_total",
|
|
||||||
"Requests that matched a rule of the rule files, by its id and action.",
|
|
||||||
[]string{"rule_id", "action"})
|
|
||||||
|
|
||||||
m.registry.MustRegister(m.ruleMatches,
|
|
||||||
prometheus.NewGaugeFunc(prometheus.GaugeOpts{
|
|
||||||
Name: "smallwebwaf_rules_loaded",
|
|
||||||
Help: "Rules loaded from the rule files.",
|
|
||||||
}, func() float64 {
|
|
||||||
return float64(ruleFiles.Len())
|
|
||||||
}))
|
|
||||||
}
|
|
||||||
|
|
||||||
// AddRemoteLog adds the metrics of sending the log lines to
|
|
||||||
// SWWAF_LOG_REMOTE_URL, read from remote as the metrics are asked for: the
|
|
||||||
// lines sent, those dropped, and those waiting in the buffer.
|
|
||||||
func (m *Metrics) AddRemoteLog(remote *remotelog.Sender) {
|
|
||||||
m.registry.MustRegister(
|
|
||||||
prometheus.NewCounterFunc(prometheus.CounterOpts{
|
|
||||||
Name: "smallwebwaf_remote_log_lines_sent_total",
|
|
||||||
Help: "Log lines sent to SWWAF_LOG_REMOTE_URL.",
|
|
||||||
}, func() float64 {
|
|
||||||
return float64(remote.Sent())
|
|
||||||
}),
|
|
||||||
prometheus.NewCounterFunc(prometheus.CounterOpts{
|
|
||||||
Name: "smallwebwaf_remote_log_lines_dropped_total",
|
|
||||||
Help: "Log lines dropped: the oldest in a full buffer, and those " +
|
|
||||||
"whose sending failed.",
|
|
||||||
}, func() float64 {
|
|
||||||
return float64(remote.Dropped())
|
|
||||||
}),
|
|
||||||
prometheus.NewGaugeFunc(prometheus.GaugeOpts{
|
|
||||||
Name: "smallwebwaf_remote_log_buffer_depth",
|
|
||||||
Help: "Log lines in the buffer, waiting to be sent.",
|
|
||||||
}, func() float64 {
|
|
||||||
return float64(remote.Depth())
|
|
||||||
}),
|
|
||||||
)
|
|
||||||
}
|
|
||||||
|
|
||||||
// AddLookupFile adds the metrics of the lookup database, read as the
|
|
||||||
// metrics are asked for: when the file in use was read, which lastRead
|
|
||||||
// returns, and the replacements of it that could not be read, which
|
|
||||||
// readFailures returns. The lookup package's File, which has both, cannot
|
|
||||||
// be named here: that package counts GeoJS's requests in these metrics.
|
|
||||||
func (m *Metrics) AddLookupFile(lastRead func() time.Time, readFailures func() int) {
|
|
||||||
m.registry.MustRegister(
|
|
||||||
prometheus.NewGaugeFunc(prometheus.GaugeOpts{
|
|
||||||
Name: "smallwebwaf_lookup_database_last_read_timestamp_seconds",
|
|
||||||
Help: "When the lookup database in use was read, in seconds since 1970.",
|
|
||||||
}, func() float64 {
|
|
||||||
return float64(lastRead().Unix())
|
|
||||||
}),
|
|
||||||
prometheus.NewCounterFunc(prometheus.CounterOpts{
|
|
||||||
Name: "smallwebwaf_lookup_database_read_failures_total",
|
|
||||||
Help: "Replacements of the lookup database that could not be read.",
|
|
||||||
}, func() float64 {
|
|
||||||
return float64(readFailures())
|
|
||||||
}),
|
|
||||||
)
|
|
||||||
}
|
|
||||||
|
|
||||||
// AddAlerts adds the metrics of the alerts sent to each destination set,
|
|
||||||
// read from queue as the metrics are asked for, by destination: the
|
|
||||||
// alerts sent, the requests to the destination that failed, the alerts
|
|
||||||
// held back, which are the same for every destination, and those
|
|
||||||
// dropped. With no destination set, it adds none.
|
|
||||||
func (m *Metrics) AddAlerts(queue *alerts.Queue) {
|
|
||||||
for _, name := range queue.DestinationsSet() {
|
|
||||||
destination := prometheus.Labels{"destination": name}
|
|
||||||
|
|
||||||
m.registry.MustRegister(
|
|
||||||
prometheus.NewCounterFunc(prometheus.CounterOpts{
|
|
||||||
Name: "smallwebwaf_alerts_sent_total",
|
|
||||||
Help: "Alerts the destination took.",
|
|
||||||
ConstLabels: destination,
|
|
||||||
}, func() float64 {
|
|
||||||
return float64(queue.Counts(name).Sent)
|
|
||||||
}),
|
|
||||||
prometheus.NewCounterFunc(prometheus.CounterOpts{
|
|
||||||
Name: "smallwebwaf_alerts_failed_total",
|
|
||||||
Help: "Requests to the destination that failed.",
|
|
||||||
ConstLabels: destination,
|
|
||||||
}, func() float64 {
|
|
||||||
return float64(queue.Counts(name).Failed)
|
|
||||||
}),
|
|
||||||
prometheus.NewCounterFunc(prometheus.CounterOpts{
|
|
||||||
Name: "smallwebwaf_alerts_suppressed_total",
|
|
||||||
Help: "Alerts held back: repeats within SWWAF_ALERT_COOLDOWN, and " +
|
|
||||||
"alerts past SWWAF_ALERT_MAX_PER_HOUR, for the hour's summary.",
|
|
||||||
ConstLabels: destination,
|
|
||||||
}, func() float64 {
|
|
||||||
return float64(queue.Suppressed())
|
|
||||||
}),
|
|
||||||
prometheus.NewCounterFunc(prometheus.CounterOpts{
|
|
||||||
Name: "smallwebwaf_alerts_dropped_total",
|
|
||||||
Help: "Alerts dropped, the oldest first, from a full queue, and " +
|
|
||||||
"alerts given up as the destination refused them.",
|
|
||||||
ConstLabels: destination,
|
|
||||||
}, func() float64 {
|
|
||||||
return float64(queue.Counts(name).Dropped)
|
|
||||||
}),
|
|
||||||
)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// 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.Country, line)
|
|
||||||
}
|
|
||||||
|
|
||||||
if line.ASN != "" {
|
|
||||||
m.asns.add(line.ASN, line)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// RuleMatched counts a request that matched the rule id, whose action is
|
|
||||||
// action.
|
|
||||||
func (m *Metrics) RuleMatched(id, action string) {
|
|
||||||
m.ruleMatches.WithLabelValues(id, action).Inc()
|
|
||||||
}
|
|
||||||
|
|
||||||
// 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)
|
|
||||||
}
|
|
||||||
@@ -1,325 +0,0 @@
|
|||||||
package proxy
|
|
||||||
|
|
||||||
import (
|
|
||||||
"bytes"
|
|
||||||
"crypto/subtle"
|
|
||||||
"encoding/json"
|
|
||||||
"errors"
|
|
||||||
"fmt"
|
|
||||||
"io"
|
|
||||||
"net/http"
|
|
||||||
"net/netip"
|
|
||||||
"os"
|
|
||||||
"strings"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
"sneak.berlin/go/smallwebwaf/internal/bans"
|
|
||||||
"sneak.berlin/go/smallwebwaf/internal/config"
|
|
||||||
"sneak.berlin/go/smallwebwaf/internal/ratelimit"
|
|
||||||
"sneak.berlin/go/smallwebwaf/internal/requestlog"
|
|
||||||
"sneak.berlin/go/smallwebwaf/internal/state"
|
|
||||||
)
|
|
||||||
|
|
||||||
// banBodyMaxBytes is the most of the body of a request to add a ban that
|
|
||||||
// is read; its three fields need far less.
|
|
||||||
const banBodyMaxBytes = 4 << 10
|
|
||||||
|
|
||||||
// permanent is how the log line and the ban endpoint name a ban that
|
|
||||||
// never ends.
|
|
||||||
const permanent = "permanent"
|
|
||||||
|
|
||||||
var (
|
|
||||||
errNotBanToAdd = errors.New(
|
|
||||||
"the body is not a JSON object of netblock, duration and reason")
|
|
||||||
errNotNetblock = errors.New(
|
|
||||||
"is not an address or a netblock, such as 203.0.113.9 or 203.0.113.0/24")
|
|
||||||
errMappedNetblock = errors.New(
|
|
||||||
"is IPv4-mapped: give the IPv4 netblock, such as 203.0.113.0/24")
|
|
||||||
errZone = errors.New("has a zone, which a netblock cannot have")
|
|
||||||
errNotDuration = errors.New(
|
|
||||||
"is not a duration above zero, such as 1h or 7d, or permanent")
|
|
||||||
errNotAddress = errors.New("is not an address, such as 203.0.113.9")
|
|
||||||
)
|
|
||||||
|
|
||||||
// answerAdmin answers a request for smallwebwaf itself, under
|
|
||||||
// /_smallwebwaf/, once it has passed the checks. Each endpoint needs a
|
|
||||||
// token, sent as Authorization: Bearer <token>: the metrics
|
|
||||||
// SWWAF_METRICS_TOKEN, the others SWWAF_ADMIN_TOKEN. A request without
|
|
||||||
// it is refused with 401. An endpoint whose token is unset answers 404,
|
|
||||||
// as any other request under /_smallwebwaf/ does.
|
|
||||||
func (rq *request) answerAdmin() {
|
|
||||||
rq.line.Action = requestlog.ActionAdmin
|
|
||||||
rq.startClientResponseTimeout()
|
|
||||||
|
|
||||||
token, answer := rq.endpoint()
|
|
||||||
|
|
||||||
switch {
|
|
||||||
case token == "":
|
|
||||||
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:
|
|
||||||
answer()
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// endpoint returns the token the request's endpoint needs, and what
|
|
||||||
// answers the request there; "" when there is no such endpoint.
|
|
||||||
func (rq *request) endpoint() (string, func()) {
|
|
||||||
cfg := rq.h.config
|
|
||||||
method, path := rq.in.Method, rq.in.URL.Path
|
|
||||||
|
|
||||||
switch {
|
|
||||||
case method == http.MethodGet && path == MetricsPath:
|
|
||||||
return cfg.MetricsToken, func() { rq.h.metrics.ServeHTTP(rq.out, rq.in) }
|
|
||||||
case method == http.MethodGet && path == BansPath:
|
|
||||||
return cfg.AdminToken, rq.listBans
|
|
||||||
case method == http.MethodPost && path == BansPath:
|
|
||||||
return cfg.AdminToken, rq.addBan
|
|
||||||
case method == http.MethodDelete && strings.HasPrefix(path, BansPath+"/"):
|
|
||||||
return cfg.AdminToken, rq.liftBans
|
|
||||||
case method == http.MethodGet && strings.HasPrefix(path, ClientsPath):
|
|
||||||
return cfg.AdminToken, rq.showClient
|
|
||||||
default:
|
|
||||||
return "", nil
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// hasToken reports whether r carries token, as Authorization: Bearer
|
|
||||||
// <token>.
|
|
||||||
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
|
|
||||||
}
|
|
||||||
|
|
||||||
// listBans answers GET BansPath with every ban held.
|
|
||||||
func (rq *request) listBans() {
|
|
||||||
rq.answerBans(rq.h.ledger.Snapshot())
|
|
||||||
}
|
|
||||||
|
|
||||||
// banToAdd is the body of POST BansPath.
|
|
||||||
type banToAdd struct {
|
|
||||||
// Netblock is a netblock, or a client's address, which stands for the
|
|
||||||
// netblock a ban on that client covers.
|
|
||||||
Netblock string `json:"netblock"`
|
|
||||||
// Duration is how long the ban lasts, as a setting gives a duration,
|
|
||||||
// or permanent.
|
|
||||||
Duration string `json:"duration"`
|
|
||||||
Reason string `json:"reason"`
|
|
||||||
}
|
|
||||||
|
|
||||||
// addBan answers POST BansPath: it bans the netblock the body names, as
|
|
||||||
// an admin, from now for the duration the body gives, with its reason,
|
|
||||||
// and answers with that ban.
|
|
||||||
func (rq *request) addBan() {
|
|
||||||
// The body must arrive within SWWAF_CLIENT_REQUEST_TIMEOUT, as any
|
|
||||||
// other request's must.
|
|
||||||
rq.stopReadingBody(rq.clientRequestDeadline())
|
|
||||||
|
|
||||||
toAdd, err := rq.readBanToAdd()
|
|
||||||
if refused := rq.refused.Load(); refused != nil {
|
|
||||||
rq.answer(*refused) // the body is over SWWAF_REQUEST_MAX_BYTES
|
|
||||||
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
if errors.Is(err, os.ErrDeadlineExceeded) {
|
|
||||||
rq.answer(refusal{
|
|
||||||
status: http.StatusRequestTimeout,
|
|
||||||
action: requestlog.ActionTimedOut,
|
|
||||||
limit: "SWWAF_CLIENT_REQUEST_TIMEOUT",
|
|
||||||
})
|
|
||||||
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
var (
|
|
||||||
netblock netip.Prefix
|
|
||||||
expires time.Time
|
|
||||||
now = rq.h.now()
|
|
||||||
)
|
|
||||||
|
|
||||||
if err == nil {
|
|
||||||
netblock, err = rq.h.banNetblock(toAdd.Netblock)
|
|
||||||
}
|
|
||||||
|
|
||||||
if err == nil {
|
|
||||||
expires, err = expiry(toAdd.Duration, now)
|
|
||||||
}
|
|
||||||
|
|
||||||
if err != nil {
|
|
||||||
http.Error(rq.out, err.Error(), http.StatusBadRequest)
|
|
||||||
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
ban := rq.h.ledger.BanForAdmin(netblock, now, expires, toAdd.Reason)
|
|
||||||
rq.answerBans([]bans.Ban{ban})
|
|
||||||
}
|
|
||||||
|
|
||||||
// readBanToAdd reads the body of POST BansPath: a JSON object with
|
|
||||||
// nothing but whitespace after it, in at most banBodyMaxBytes.
|
|
||||||
func (rq *request) readBanToAdd() (banToAdd, error) {
|
|
||||||
var body io.ReadCloser = http.NoBody
|
|
||||||
if rq.body != nil {
|
|
||||||
body = rq.body
|
|
||||||
}
|
|
||||||
|
|
||||||
data, err := io.ReadAll(http.MaxBytesReader(nil, body, banBodyMaxBytes))
|
|
||||||
if err != nil {
|
|
||||||
return banToAdd{}, fmt.Errorf("%w: %w", errNotBanToAdd, err)
|
|
||||||
}
|
|
||||||
|
|
||||||
var toAdd banToAdd
|
|
||||||
|
|
||||||
decoder := json.NewDecoder(bytes.NewReader(data))
|
|
||||||
decoder.DisallowUnknownFields()
|
|
||||||
|
|
||||||
err = decoder.Decode(&toAdd)
|
|
||||||
if err != nil {
|
|
||||||
return banToAdd{}, fmt.Errorf("%w: %w", errNotBanToAdd, err)
|
|
||||||
}
|
|
||||||
|
|
||||||
// Token returns io.EOF only when nothing but whitespace is left.
|
|
||||||
_, err = decoder.Token()
|
|
||||||
if !errors.Is(err, io.EOF) {
|
|
||||||
return banToAdd{}, fmt.Errorf("%w: more follows the object", errNotBanToAdd)
|
|
||||||
}
|
|
||||||
|
|
||||||
return toAdd, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// banNetblock reads value, a netblock such as 203.0.113.0/24, or a
|
|
||||||
// client's address, which stands for the netblock a ban on that client
|
|
||||||
// covers. An IPv4-mapped netblock, such as ::ffff:203.0.113.0/120, is
|
|
||||||
// refused, since a client's address is looked up as IPv4 and a ban on it
|
|
||||||
// would refuse nothing, and so is a value with a zone.
|
|
||||||
func (h *handler) banNetblock(value string) (netip.Prefix, error) {
|
|
||||||
netblock, err := netip.ParsePrefix(value)
|
|
||||||
if err == nil {
|
|
||||||
if netblock.Addr().Is4In6() {
|
|
||||||
return netip.Prefix{}, fmt.Errorf("netblock %q %w", value, errMappedNetblock)
|
|
||||||
}
|
|
||||||
|
|
||||||
return netblock, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// ParsePrefix refuses a zone, but ParseAddr reads the /48 of
|
|
||||||
// 2001:db8::1%x/48 as part of the zone.
|
|
||||||
addr, err := netip.ParseAddr(value)
|
|
||||||
if err != nil {
|
|
||||||
return netip.Prefix{}, fmt.Errorf("netblock %q %w", value, errNotNetblock)
|
|
||||||
}
|
|
||||||
|
|
||||||
if addr.Zone() != "" {
|
|
||||||
return netip.Prefix{}, fmt.Errorf("netblock %q %w", value, errZone)
|
|
||||||
}
|
|
||||||
|
|
||||||
return h.netblock(addr), nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// expiry returns when a ban made at now for duration ends: duration
|
|
||||||
// later, for a duration as a setting gives one, or zero for permanent.
|
|
||||||
func expiry(duration string, now time.Time) (time.Time, error) {
|
|
||||||
if duration == permanent {
|
|
||||||
return time.Time{}, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
length, err := config.ParseDurationNotOff(duration)
|
|
||||||
if err != nil {
|
|
||||||
return time.Time{}, fmt.Errorf("duration %q %w", duration, errNotDuration)
|
|
||||||
}
|
|
||||||
|
|
||||||
return now.Add(length), nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// liftBans answers DELETE BansPath/<client>: it lifts every ban active on
|
|
||||||
// a netblock the client's address is in, and answers with those bans, or
|
|
||||||
// with 404 when none is active.
|
|
||||||
func (rq *request) liftBans() {
|
|
||||||
client, err := pathAddress(rq.in.URL.Path, BansPath+"/")
|
|
||||||
if err != nil {
|
|
||||||
http.Error(rq.out, err.Error(), http.StatusBadRequest)
|
|
||||||
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
lifted := rq.h.ledger.Lift(client, rq.h.now())
|
|
||||||
if len(lifted) == 0 {
|
|
||||||
http.Error(rq.out, "no ban is active on "+client.String(), http.StatusNotFound)
|
|
||||||
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
rq.answerBans(lifted)
|
|
||||||
}
|
|
||||||
|
|
||||||
// clientAnswer is the answer to GET ClientsPath<ip>: the client the
|
|
||||||
// address is, as clients.json holds it, or null when the table of
|
|
||||||
// clients does not hold it, and the bans on each netblock the address is
|
|
||||||
// in, as bans.json lists them.
|
|
||||||
type clientAnswer struct {
|
|
||||||
Client *ratelimit.Client `json:"client"`
|
|
||||||
Bans []state.BanEntry `json:"bans"`
|
|
||||||
}
|
|
||||||
|
|
||||||
// showClient answers GET ClientsPath<ip> with what smallwebwaf knows of
|
|
||||||
// the client: its counters, its history, which holds its country as last
|
|
||||||
// looked up and its offences, and its bans with their notes.
|
|
||||||
func (rq *request) showClient() {
|
|
||||||
addr, err := pathAddress(rq.in.URL.Path, ClientsPath)
|
|
||||||
if err != nil {
|
|
||||||
http.Error(rq.out, err.Error(), http.StatusBadRequest)
|
|
||||||
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
answer := clientAnswer{Bans: state.BanEntries(rq.h.ledger.Covering(addr))}
|
|
||||||
|
|
||||||
client, seen := rq.h.limiter.Client(clientGroup(addr))
|
|
||||||
if seen {
|
|
||||||
answer.Client = &client
|
|
||||||
}
|
|
||||||
|
|
||||||
rq.answerJSON(answer)
|
|
||||||
}
|
|
||||||
|
|
||||||
// pathAddress reads the client's address that follows prefix in path.
|
|
||||||
func pathAddress(path, prefix string) (netip.Addr, error) {
|
|
||||||
value := strings.TrimPrefix(path, prefix)
|
|
||||||
|
|
||||||
addr, err := netip.ParseAddr(value)
|
|
||||||
if err != nil {
|
|
||||||
return netip.Addr{}, fmt.Errorf("%q %w", value, errNotAddress)
|
|
||||||
}
|
|
||||||
|
|
||||||
return addr.Unmap(), nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// answerBans answers with held under bans, as bans.json lists them.
|
|
||||||
func (rq *request) answerBans(held []bans.Ban) {
|
|
||||||
rq.answerJSON(struct {
|
|
||||||
Bans []state.BanEntry `json:"bans"`
|
|
||||||
}{state.BanEntries(held)})
|
|
||||||
}
|
|
||||||
|
|
||||||
// answerJSON answers with value as indented JSON.
|
|
||||||
func (rq *request) answerJSON(value any) {
|
|
||||||
body, err := json.MarshalIndent(value, "", " ")
|
|
||||||
if err != nil {
|
|
||||||
rq.h.processLog.Error("encoding an answer failed", "error", err.Error())
|
|
||||||
http.Error(rq.out, http.StatusText(http.StatusInternalServerError),
|
|
||||||
http.StatusInternalServerError)
|
|
||||||
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
rq.out.Header().Set("Content-Type", "application/json")
|
|
||||||
_, _ = rq.out.Write(append(body, '\n'))
|
|
||||||
}
|
|
||||||
@@ -1,535 +0,0 @@
|
|||||||
package proxy_test
|
|
||||||
|
|
||||||
import (
|
|
||||||
"encoding/json"
|
|
||||||
"net/http"
|
|
||||||
"net/netip"
|
|
||||||
"slices"
|
|
||||||
"strconv"
|
|
||||||
"strings"
|
|
||||||
"testing"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
"sneak.berlin/go/smallwebwaf/internal/bans"
|
|
||||||
"sneak.berlin/go/smallwebwaf/internal/proxy"
|
|
||||||
"sneak.berlin/go/smallwebwaf/internal/ratelimit"
|
|
||||||
"sneak.berlin/go/smallwebwaf/internal/requestlog"
|
|
||||||
"sneak.berlin/go/smallwebwaf/internal/state"
|
|
||||||
)
|
|
||||||
|
|
||||||
const (
|
|
||||||
adminToken = "SWWAF_ADMIN_TOKEN" //nolint:gosec // the setting's name
|
|
||||||
// adminSecret is the SWWAF_ADMIN_TOKEN the tests set, and adminBearer
|
|
||||||
// how a request carries it.
|
|
||||||
adminSecret = "fedcba9876543210fedcba9876543210"
|
|
||||||
adminBearer = "Bearer " + adminSecret
|
|
||||||
// adminClient is the client the tests' admin sends its requests from.
|
|
||||||
adminClient = "192.0.2.10"
|
|
||||||
// banOtherClient is the body of a request to ban otherClient for an
|
|
||||||
// hour.
|
|
||||||
banOtherClient = `{"netblock": "` + otherClient + `", "duration": "1h", ` +
|
|
||||||
`"reason": "probes for logins"}`
|
|
||||||
)
|
|
||||||
|
|
||||||
func TestAdminEndpointsAreOffWhileTheTokenIsUnset(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
// The metrics token is set, and opens none of them.
|
|
||||||
s, clk, server := startWithClock(t, "", map[string]string{metricsToken: token})
|
|
||||||
server.Ledger.BanForLimit(netip.MustParsePrefix(otherClient+"/32"), clk.Now(),
|
|
||||||
bans.Notes{})
|
|
||||||
before := server.Ledger.Snapshot()
|
|
||||||
|
|
||||||
// An empty token does not match the unset one either.
|
|
||||||
for _, authorization := range []string{adminBearer, bearer, "Bearer ", ""} {
|
|
||||||
for _, e := range adminEndpoints() {
|
|
||||||
s.adminRequest(adminClient, authorization, e.method, e.path, e.body,
|
|
||||||
http.StatusNotFound, requestlog.ActionAdmin)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
if after := server.Ledger.Snapshot(); !slices.Equal(after, before) {
|
|
||||||
t.Errorf("the bans are now\n%+v\nwant them unchanged\n%+v", after, before)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestAdminEndpointsNeedTheAdminToken(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
s, _, server := startWithClock(t, "", map[string]string{
|
|
||||||
adminToken: adminSecret,
|
|
||||||
metricsToken: token,
|
|
||||||
})
|
|
||||||
|
|
||||||
// Listing the bans, banning otherClient, lifting that ban, and asking
|
|
||||||
// about otherClient, in that order. Without the admin token, with the
|
|
||||||
// metrics token, or with one that differs, each is refused, and
|
|
||||||
// changes nothing; with the admin token, it is answered.
|
|
||||||
for _, e := range adminEndpoints() {
|
|
||||||
before := server.Ledger.Snapshot()
|
|
||||||
|
|
||||||
for _, authorization := range []string{
|
|
||||||
"", bearer, "Bearer " + strings.ToUpper(adminSecret), "Basic " + adminSecret,
|
|
||||||
} {
|
|
||||||
got := s.adminRequest(adminClient, authorization, e.method, e.path, e.body,
|
|
||||||
http.StatusUnauthorized, requestlog.ActionAdmin)
|
|
||||||
if got.header.Get("WWW-Authenticate") != "Bearer" {
|
|
||||||
t.Errorf("%s %s with %q was answered without WWW-Authenticate: Bearer",
|
|
||||||
e.method, e.path, authorization)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
if after := server.Ledger.Snapshot(); !slices.Equal(after, before) {
|
|
||||||
t.Errorf("%s %s without the token changed the bans to\n%+v\nfrom\n%+v",
|
|
||||||
e.method, e.path, after, before)
|
|
||||||
}
|
|
||||||
|
|
||||||
got := s.admin(e.method, e.path, e.body, http.StatusOK)
|
|
||||||
if got.header.Get("Content-Type") != "application/json" {
|
|
||||||
t.Errorf("%s %s answered %q", e.method, e.path, got.header.Get("Content-Type"))
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// Any other request under /_smallwebwaf/ is not found.
|
|
||||||
for _, e := range []adminEndpoint{
|
|
||||||
{http.MethodPut, proxy.BansPath, banOtherClient},
|
|
||||||
{http.MethodDelete, proxy.BansPath, ""},
|
|
||||||
{http.MethodGet, proxy.BansPath + "/" + otherClient, ""},
|
|
||||||
{http.MethodPost, proxy.ClientsPath + otherClient, ""},
|
|
||||||
{http.MethodGet, strings.TrimSuffix(proxy.ClientsPath, "/"), ""},
|
|
||||||
} {
|
|
||||||
s.admin(e.method, e.path, e.body, http.StatusNotFound)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestBanAddedListedAndLiftedThroughTheEndpoints(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
s, clk, _ := startWithClock(t, "", map[string]string{
|
|
||||||
adminToken: adminSecret,
|
|
||||||
banScopeV4Prefix: "24",
|
|
||||||
})
|
|
||||||
|
|
||||||
// A ban on otherClient bans the /24 a ban on that client covers, so it
|
|
||||||
// refuses client too, for an hour.
|
|
||||||
start := clk.Now()
|
|
||||||
expires := start.Add(time.Hour)
|
|
||||||
want := state.BanEntry{
|
|
||||||
Netblock: netip.MustParsePrefix("203.0.113.0/24"),
|
|
||||||
Start: start,
|
|
||||||
Expires: &expires,
|
|
||||||
Cause: bans.CauseAdmin,
|
|
||||||
Reason: "probes for logins",
|
|
||||||
}
|
|
||||||
|
|
||||||
wantBans(t, s.admin(http.MethodPost, proxy.BansPath, banOtherClient, http.StatusOK),
|
|
||||||
want)
|
|
||||||
|
|
||||||
line := s.get(client, http.StatusForbidden, requestlog.ActionBanned)
|
|
||||||
if line.BanExpires != requestlog.FormatTime(expires) {
|
|
||||||
t.Errorf("the ban ends at %s, want %s", line.BanExpires, expires)
|
|
||||||
}
|
|
||||||
|
|
||||||
// Its notes count the request it refused.
|
|
||||||
want.Notes.Requests, want.Notes.Refused = 1, 1
|
|
||||||
|
|
||||||
wantBans(t, s.admin(http.MethodGet, proxy.BansPath, "", http.StatusOK), want)
|
|
||||||
|
|
||||||
// Ten minutes on, lifting the bans on client lifts that one, which is
|
|
||||||
// kept, marked lifted.
|
|
||||||
clk.advance(10 * time.Minute)
|
|
||||||
|
|
||||||
lifted := clk.Now()
|
|
||||||
want.Lifted = &lifted
|
|
||||||
|
|
||||||
wantBans(t, s.admin(http.MethodDelete, proxy.BansPath+"/"+client, "",
|
|
||||||
http.StatusOK), want)
|
|
||||||
s.get(client, http.StatusOK, requestlog.ActionForward)
|
|
||||||
wantBans(t, s.admin(http.MethodGet, proxy.BansPath, "", http.StatusOK), want)
|
|
||||||
|
|
||||||
// No ban on it is active any more.
|
|
||||||
s.admin(http.MethodDelete, proxy.BansPath+"/"+client, "", http.StatusNotFound)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestBanToAddGivesItsNetblockAndDuration(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
s, clk, _ := startWithClock(t, "", map[string]string{
|
|
||||||
adminToken: adminSecret,
|
|
||||||
banScopeV4Prefix: "24",
|
|
||||||
})
|
|
||||||
start := clk.Now()
|
|
||||||
|
|
||||||
for _, tc := range []struct {
|
|
||||||
netblock, duration string
|
|
||||||
want string
|
|
||||||
length time.Duration // 0 for a permanent ban
|
|
||||||
}{
|
|
||||||
// An address stands for the netblock a ban on that client covers.
|
|
||||||
{client, "7d", "203.0.113.0/24", 7 * 24 * time.Hour},
|
|
||||||
{"::ffff:198.51.100.7", "90m", "198.51.100.0/24", 90 * time.Minute},
|
|
||||||
{"2001:db8:5::1", "permanent", "2001:db8:5::/64", 0},
|
|
||||||
// A netblock stands for itself, its bits past its length cleared.
|
|
||||||
{"198.51.100.7/16", "1h", "198.51.0.0/16", time.Hour},
|
|
||||||
{"2001:db8:6::/48", "1h", "2001:db8:6::/48", time.Hour},
|
|
||||||
} {
|
|
||||||
// Whitespace may follow the object.
|
|
||||||
body := `{"netblock": "` + tc.netblock + `", "duration": "` + tc.duration + `"}` +
|
|
||||||
"\r\n"
|
|
||||||
want := state.BanEntry{
|
|
||||||
Netblock: netip.MustParsePrefix(tc.want), Start: start, Cause: bans.CauseAdmin,
|
|
||||||
}
|
|
||||||
|
|
||||||
if tc.length != 0 {
|
|
||||||
expires := start.Add(tc.length)
|
|
||||||
want.Expires = &expires
|
|
||||||
}
|
|
||||||
|
|
||||||
wantBans(t, s.admin(http.MethodPost, proxy.BansPath, body, http.StatusOK), want)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestBanToAddThatCannotBeReadIsRefused(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
s, _, server := startWithClock(t, "", map[string]string{adminToken: adminSecret})
|
|
||||||
|
|
||||||
for _, tc := range []struct{ body, want string }{
|
|
||||||
{"", "the body is not a JSON object of netblock, duration and reason: EOF"},
|
|
||||||
{"netblock=203.0.113.9", "the body is not a JSON object"},
|
|
||||||
{
|
|
||||||
`{"netblock": "203.0.113.9", "duration": "1h", "until": "2027"}`,
|
|
||||||
`unknown field "until"`,
|
|
||||||
},
|
|
||||||
{
|
|
||||||
`{"netblock": "203.0.113", "duration": "1h"}`,
|
|
||||||
`netblock "203.0.113" is not an address or a netblock`,
|
|
||||||
},
|
|
||||||
// A client's address is looked up as IPv4, so a ban on an
|
|
||||||
// IPv4-mapped netblock would refuse nothing.
|
|
||||||
{
|
|
||||||
`{"netblock": "::ffff:203.0.113.0/120", "duration": "1h"}`,
|
|
||||||
`netblock "::ffff:203.0.113.0/120" is IPv4-mapped`,
|
|
||||||
},
|
|
||||||
// Read as an address, its zone would be "x/48", and its ban on the
|
|
||||||
// /64 around it.
|
|
||||||
{
|
|
||||||
`{"netblock": "2001:db8::1%x/48", "duration": "1h"}`,
|
|
||||||
`netblock "2001:db8::1%x/48" has a zone`,
|
|
||||||
},
|
|
||||||
{
|
|
||||||
`{"netblock": "fe80::1%eth0", "duration": "1h"}`,
|
|
||||||
`netblock "fe80::1%eth0" has a zone`,
|
|
||||||
},
|
|
||||||
// Anything but whitespace after the object.
|
|
||||||
{
|
|
||||||
`{"netblock": "203.0.113.9", "duration": "1h"}` +
|
|
||||||
`{"netblock": "198.51.100.0/24", "duration": "1h"}`,
|
|
||||||
"more follows the object",
|
|
||||||
},
|
|
||||||
{`{"netblock": "203.0.113.9", "duration": "1h"} x`, "more follows the object"},
|
|
||||||
{`{"duration": "1h"}`, `netblock "" is not an address or a netblock`},
|
|
||||||
{`{"netblock": "203.0.113.9"}`, `duration "" is not a duration above zero`},
|
|
||||||
{
|
|
||||||
`{"netblock": "203.0.113.9", "duration": "off"}`,
|
|
||||||
`duration "off" is not a duration above zero`,
|
|
||||||
},
|
|
||||||
{
|
|
||||||
`{"netblock": "203.0.113.9", "duration": "0s"}`,
|
|
||||||
`duration "0s" is not a duration above zero`,
|
|
||||||
},
|
|
||||||
{
|
|
||||||
`{"netblock": "203.0.113.9", "duration": "forever"}`,
|
|
||||||
`duration "forever" is not a duration above zero, such as 1h or 7d, ` +
|
|
||||||
`or permanent`,
|
|
||||||
},
|
|
||||||
// Over the 4 KiB read of a body, even when the object comes first.
|
|
||||||
{
|
|
||||||
`{"netblock": "203.0.113.9", "duration": "1h", "reason": "` +
|
|
||||||
strings.Repeat("x", 4<<10) + `"}`,
|
|
||||||
"request body too large",
|
|
||||||
},
|
|
||||||
{
|
|
||||||
`{"netblock": "203.0.113.9", "duration": "1h"}` + strings.Repeat(" ", 4<<10),
|
|
||||||
"request body too large",
|
|
||||||
},
|
|
||||||
} {
|
|
||||||
got := s.admin(http.MethodPost, proxy.BansPath, tc.body, http.StatusBadRequest)
|
|
||||||
if !strings.Contains(string(got.body), tc.want) {
|
|
||||||
t.Errorf("%.80s was answered %q, want it to say %q", tc.body, got.body, tc.want)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
if held := server.Ledger.Snapshot(); len(held) != 0 {
|
|
||||||
t.Errorf("the ledger holds %+v, want no ban", held)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestBanToAddOverTheRequestSizeLimitIsRefused(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
s, _, server := startWithClock(t, "", map[string]string{
|
|
||||||
adminToken: adminSecret,
|
|
||||||
requestMaxBytes: "16",
|
|
||||||
})
|
|
||||||
|
|
||||||
// Sent in a chunk, its length is not announced, so that it is found
|
|
||||||
// over SWWAF_REQUEST_MAX_BYTES only as it is read.
|
|
||||||
chunk := `{"netblock": "203.0.113.9", "duration": "1h"}`
|
|
||||||
s.adminRequest(adminClient, adminBearer+"\r\nTransfer-Encoding: chunked",
|
|
||||||
http.MethodPost, proxy.BansPath,
|
|
||||||
strconv.FormatInt(int64(len(chunk)), 16)+"\r\n"+chunk+"\r\n0\r\n\r\n",
|
|
||||||
http.StatusRequestEntityTooLarge, requestlog.ActionTooLarge)
|
|
||||||
|
|
||||||
if held := server.Ledger.Snapshot(); len(held) != 0 {
|
|
||||||
t.Errorf("the ledger holds %+v, want no ban", held)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestBanToAddSlowerThanTheClientRequestTimeoutIsRefused(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
s, _, server := startWithClock(t, "", map[string]string{
|
|
||||||
adminToken: adminSecret,
|
|
||||||
metricsToken: token,
|
|
||||||
clientRequestTimeout: shortTimeoutSetting,
|
|
||||||
})
|
|
||||||
|
|
||||||
// The chunk announces 256 bytes and the rest of it never comes, so only
|
|
||||||
// the timeout ends the wait. A hold-up of the test process can only
|
|
||||||
// make the answer later, so the time is checked only for not being
|
|
||||||
// shorter than the timeout.
|
|
||||||
start := time.Now()
|
|
||||||
|
|
||||||
s.adminRequest(adminClient, adminBearer+"\r\nTransfer-Encoding: chunked",
|
|
||||||
http.MethodPost, proxy.BansPath, "100\r\n"+`{"netblock": "203.0.113.9", `,
|
|
||||||
http.StatusRequestTimeout, requestlog.ActionTimedOut)
|
|
||||||
|
|
||||||
if took := time.Since(start); took < shortTimeout {
|
|
||||||
t.Errorf("answered after %s, before the timeout of %s ran out", took, shortTimeout)
|
|
||||||
}
|
|
||||||
|
|
||||||
wantLimitHits(t, s.addr, clientRequestTimeout, 1)
|
|
||||||
|
|
||||||
if held := server.Ledger.Snapshot(); len(held) != 0 {
|
|
||||||
t.Errorf("the ledger holds %+v, want no ban", held)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestClientEndpointShowsTheClientAndItsBans(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
s, clk, _ := startWithClock(t, "", map[string]string{
|
|
||||||
adminToken: adminSecret,
|
|
||||||
rateLimitPerMinute: "2",
|
|
||||||
rateLimitExemptNets: adminClient,
|
|
||||||
})
|
|
||||||
start := clk.Now()
|
|
||||||
|
|
||||||
// Two of otherClient's requests are let through; the third breaks the
|
|
||||||
// limit of two a minute, and bans it.
|
|
||||||
s.get(otherClient, http.StatusOK, requestlog.ActionForward)
|
|
||||||
s.get(otherClient, http.StatusOK, requestlog.ActionForward)
|
|
||||||
s.get(otherClient, http.StatusForbidden, requestlog.ActionRateLimited)
|
|
||||||
|
|
||||||
// Asked about by its address in IPv6 form too.
|
|
||||||
for _, addr := range []string{otherClient, "::ffff:" + otherClient} {
|
|
||||||
var got struct {
|
|
||||||
Client *ratelimit.Client `json:"client"`
|
|
||||||
Bans []state.BanEntry `json:"bans"`
|
|
||||||
}
|
|
||||||
|
|
||||||
decode(t, s.admin(http.MethodGet, proxy.ClientsPath+addr, "", http.StatusOK), &got)
|
|
||||||
|
|
||||||
if got.Client == nil {
|
|
||||||
t.Fatalf("%s: no client", addr)
|
|
||||||
}
|
|
||||||
|
|
||||||
history := got.Client.History
|
|
||||||
if got.Client.Client != netip.MustParsePrefix(otherClient+"/32") ||
|
|
||||||
history.Requests != 3 || history.Forwarded != 2 || history.Refused != 1 ||
|
|
||||||
history.Offences.Limit != 1 || !history.FirstSeen.Equal(start) {
|
|
||||||
t.Errorf("%s: client %+v", addr, got.Client)
|
|
||||||
}
|
|
||||||
|
|
||||||
if len(got.Bans) != 1 || got.Bans[0].Cause != bans.CauseLimit ||
|
|
||||||
got.Bans[0].Reason != "requests per minute over the limit of 2" ||
|
|
||||||
got.Bans[0].Notes.Count != 3 {
|
|
||||||
t.Errorf("%s: bans %+v, want the one for the broken limit", addr, got.Bans)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// Of an address no request came from and no ban covers, nothing is
|
|
||||||
// known.
|
|
||||||
got := s.admin(http.MethodGet, proxy.ClientsPath+"198.51.100.99", "", http.StatusOK)
|
|
||||||
if string(got.body) != "{\n \"client\": null,\n \"bans\": []\n}\n" {
|
|
||||||
t.Errorf("an unknown client is answered\n%s", got.body)
|
|
||||||
}
|
|
||||||
|
|
||||||
s.admin(http.MethodGet, proxy.ClientsPath+"203.0.113", "", http.StatusBadRequest)
|
|
||||||
s.admin(http.MethodDelete, proxy.BansPath+"/203.0.113.0/24", "",
|
|
||||||
http.StatusBadRequest)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestBannedClientIsRefusedAtTheEndpointsEvenWithTheToken(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
s, _, _ := startWithClock(t, "", map[string]string{adminToken: adminSecret})
|
|
||||||
|
|
||||||
s.admin(http.MethodPost, proxy.BansPath, banOtherClient, http.StatusOK)
|
|
||||||
|
|
||||||
// otherClient cannot lift its own ban either.
|
|
||||||
for _, e := range adminEndpoints() {
|
|
||||||
s.adminRequest(otherClient, adminBearer, e.method, e.path, e.body,
|
|
||||||
http.StatusForbidden, requestlog.ActionBanned)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestAdminRequestsCountTowardTheLimits(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
s, _, _ := startWithClock(t, "", map[string]string{
|
|
||||||
adminToken: adminSecret,
|
|
||||||
rateLimitPerMinute: "2",
|
|
||||||
})
|
|
||||||
|
|
||||||
// A request refused for a missing token and one answered count toward
|
|
||||||
// the limit of two a minute, so the next breaks it.
|
|
||||||
s.adminRequest(client, "", http.MethodGet, proxy.BansPath, "",
|
|
||||||
http.StatusUnauthorized, requestlog.ActionAdmin)
|
|
||||||
s.adminRequest(client, adminBearer, http.MethodGet, proxy.BansPath, "",
|
|
||||||
http.StatusOK, requestlog.ActionAdmin)
|
|
||||||
s.adminRequest(client, adminBearer, http.MethodGet, proxy.BansPath, "",
|
|
||||||
http.StatusForbidden, requestlog.ActionRateLimited)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestClientInAllowNetsSkipsTheChecksButNeedsTheToken(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
const allowed = "192.0.2.60" // in SWWAF_ALLOW_NETS
|
|
||||||
|
|
||||||
s, clk, server := startWithClock(t, "", map[string]string{
|
|
||||||
adminToken: adminSecret,
|
|
||||||
allowNets: allowed,
|
|
||||||
rateLimitPerMinute: "1",
|
|
||||||
})
|
|
||||||
|
|
||||||
// A ban on it refuses nothing, and its requests are not counted.
|
|
||||||
server.Ledger.BanForAdmin(netip.MustParsePrefix(allowed+"/32"), clk.Now(),
|
|
||||||
time.Time{}, "")
|
|
||||||
|
|
||||||
for range 2 {
|
|
||||||
s.adminRequest(allowed, "", http.MethodGet, proxy.BansPath, "",
|
|
||||||
http.StatusUnauthorized, requestlog.ActionAdmin)
|
|
||||||
s.adminRequest(allowed, adminBearer, http.MethodGet, proxy.BansPath, "",
|
|
||||||
http.StatusOK, requestlog.ActionAdmin)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestAdminEndpointsNeedTheTokenInObserveMode(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
s, _, server := startWithClock(t, "", map[string]string{
|
|
||||||
adminToken: adminSecret,
|
|
||||||
mode: observe,
|
|
||||||
})
|
|
||||||
|
|
||||||
for _, e := range adminEndpoints() {
|
|
||||||
s.adminRequest(adminClient, "", e.method, e.path, e.body,
|
|
||||||
http.StatusUnauthorized, requestlog.ActionAdmin)
|
|
||||||
}
|
|
||||||
|
|
||||||
if held := server.Ledger.Snapshot(); len(held) != 0 {
|
|
||||||
t.Errorf("the ledger holds %+v, want no ban", held)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// adminEndpoint is a request to an endpoint SWWAF_ADMIN_TOKEN opens.
|
|
||||||
type adminEndpoint struct {
|
|
||||||
method, path, body string
|
|
||||||
}
|
|
||||||
|
|
||||||
// adminEndpoints returns a request to each endpoint SWWAF_ADMIN_TOKEN
|
|
||||||
// opens: listing the bans, banning otherClient for an hour, lifting the
|
|
||||||
// bans on otherClient, and asking about otherClient.
|
|
||||||
func adminEndpoints() []adminEndpoint {
|
|
||||||
return []adminEndpoint{
|
|
||||||
{http.MethodGet, proxy.BansPath, ""},
|
|
||||||
{http.MethodPost, proxy.BansPath, banOtherClient},
|
|
||||||
{http.MethodDelete, proxy.BansPath + "/" + otherClient, ""},
|
|
||||||
{http.MethodGet, proxy.ClientsPath + otherClient, ""},
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// admin sends a request with method for path, with body, from
|
|
||||||
// adminClient, with the admin token, and checks that it is answered with
|
|
||||||
// status, its log line's action admin. It returns the answer.
|
|
||||||
func (s *sender) admin(method, path, body string, status int) answer {
|
|
||||||
s.t.Helper()
|
|
||||||
|
|
||||||
return s.adminRequest(adminClient, adminBearer, method, path, body, status,
|
|
||||||
requestlog.ActionAdmin)
|
|
||||||
}
|
|
||||||
|
|
||||||
// adminRequest sends a request with method for path, with body, from the
|
|
||||||
// client at from, with authorization as its Authorization header unless
|
|
||||||
// it is "", and checks its answer's status and its log line's action, as
|
|
||||||
// request does. authorization may end in more header lines. A body that
|
|
||||||
// is not "" has its length announced, unless authorization names
|
|
||||||
// Transfer-Encoding. It returns the answer.
|
|
||||||
func (s *sender) adminRequest(
|
|
||||||
from, authorization, method, path, body string, status int, action string,
|
|
||||||
) answer {
|
|
||||||
s.t.Helper()
|
|
||||||
|
|
||||||
var header []string
|
|
||||||
|
|
||||||
if authorization != "" {
|
|
||||||
header = append(header, "Authorization: "+authorization)
|
|
||||||
}
|
|
||||||
|
|
||||||
if body != "" && !strings.Contains(authorization, "Transfer-Encoding") {
|
|
||||||
header = append(header, "Content-Length: "+strconv.Itoa(len(body)))
|
|
||||||
}
|
|
||||||
|
|
||||||
_, got := s.requestWithBody(method, from, path, strings.Join(header, "\r\n"),
|
|
||||||
body, status, action)
|
|
||||||
|
|
||||||
return got
|
|
||||||
}
|
|
||||||
|
|
||||||
// wantBans checks that a ban endpoint answered with want, and no other
|
|
||||||
// ban.
|
|
||||||
func wantBans(t *testing.T, got answer, want ...state.BanEntry) {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
var decoded struct {
|
|
||||||
Bans []state.BanEntry `json:"bans"`
|
|
||||||
}
|
|
||||||
|
|
||||||
decode(t, got, &decoded)
|
|
||||||
|
|
||||||
gotJSON, err := json.Marshal(decoded.Bans)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("encode %+v: %v", decoded.Bans, err)
|
|
||||||
}
|
|
||||||
|
|
||||||
wantJSON, err := json.Marshal(want)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("encode %+v: %v", want, err)
|
|
||||||
}
|
|
||||||
|
|
||||||
if string(gotJSON) != string(wantJSON) {
|
|
||||||
t.Errorf("bans\n%s\nwant\n%s", gotJSON, wantJSON)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// decode reads the JSON answer of an endpoint into value.
|
|
||||||
func decode(t *testing.T, got answer, value any) {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
err := json.Unmarshal(got.body, value)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("decode %s: %v", got.body, err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@@ -1,264 +0,0 @@
|
|||||||
package proxy_test
|
|
||||||
|
|
||||||
import (
|
|
||||||
"maps"
|
|
||||||
"net/http"
|
|
||||||
"net/netip"
|
|
||||||
"reflect"
|
|
||||||
"testing"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
"sneak.berlin/go/smallwebwaf/internal/alerts"
|
|
||||||
"sneak.berlin/go/smallwebwaf/internal/bans"
|
|
||||||
"sneak.berlin/go/smallwebwaf/internal/proxy"
|
|
||||||
"sneak.berlin/go/smallwebwaf/internal/requestlog"
|
|
||||||
)
|
|
||||||
|
|
||||||
const (
|
|
||||||
alertWebhookURL = "SWWAF_ALERT_WEBHOOK_URL"
|
|
||||||
alertMaxPerHour = "SWWAF_ALERT_MAX_PER_HOUR"
|
|
||||||
// alertInstance is the instance every alert of these tests gives.
|
|
||||||
alertInstance = "fsn1app1/gitea"
|
|
||||||
)
|
|
||||||
|
|
||||||
func TestBanForABrokenLimitRaisesABanAlertWithItsNotes(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
s, clk, server, queue := startWithAlerts(t, map[string]string{
|
|
||||||
rateLimitPerMinute: "1",
|
|
||||||
banScopeV4Prefix: "24",
|
|
||||||
})
|
|
||||||
start := clk.Now()
|
|
||||||
|
|
||||||
s.get(client, http.StatusOK, requestlog.ActionForward)
|
|
||||||
s.get(client, http.StatusForbidden, requestlog.ActionRateLimited)
|
|
||||||
|
|
||||||
netblock := netip.MustParsePrefix("203.0.113.0/24")
|
|
||||||
ban := server.Ledger.Bans(netblock)[0]
|
|
||||||
|
|
||||||
// A request refused under the ban raises no other alert.
|
|
||||||
clk.advance(time.Minute)
|
|
||||||
s.get(client, http.StatusForbidden, requestlog.ActionBanned)
|
|
||||||
|
|
||||||
wantAlerts(t, queue, banAlert(alerts.EventBan, start, client, bans.Ban{
|
|
||||||
Netblock: netblock, Cause: bans.CauseLimit,
|
|
||||||
Reason: "requests per minute over the limit of 1", Notes: ban.Notes,
|
|
||||||
}, requestlog.FormatTime(start.Add(time.Hour))))
|
|
||||||
|
|
||||||
if ban.Notes.Limit != 1 || ban.Notes.Request.Path != "/" {
|
|
||||||
t.Errorf("the alert's notes are %+v, want those of the broken limit", ban.Notes)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestAttackBanRaisesABanAlertThenAPermanentBanAlert(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
s, clk, server, queue := startWithAlerts(t, map[string]string{
|
|
||||||
rulesDir: writeRules(t, testRules),
|
|
||||||
})
|
|
||||||
start := clk.Now()
|
|
||||||
netblock := netip.MustParsePrefix(client + "/32")
|
|
||||||
other := netip.MustParsePrefix(otherClient + "/32")
|
|
||||||
|
|
||||||
// The probe bans the client for seven days, and its next request makes
|
|
||||||
// the ban permanent. The request after that changes nothing.
|
|
||||||
s.request(client, "/.env", http.StatusForbidden, requestlog.ActionBanned)
|
|
||||||
|
|
||||||
attackBan := server.Ledger.Bans(netblock)[0]
|
|
||||||
|
|
||||||
clk.advance(time.Minute)
|
|
||||||
s.get(client, http.StatusForbidden, requestlog.ActionBanned)
|
|
||||||
|
|
||||||
permanentBan := server.Ledger.Bans(netblock)[0]
|
|
||||||
|
|
||||||
s.get(client, http.StatusForbidden, requestlog.ActionBanned)
|
|
||||||
|
|
||||||
// Another client's probe after its first ban has run out without a
|
|
||||||
// request makes a permanent ban at once.
|
|
||||||
s.request(otherClient, "/.env", http.StatusForbidden, requestlog.ActionBanned)
|
|
||||||
clk.advance(7 * 24 * time.Hour)
|
|
||||||
s.request(otherClient, "/.env", http.StatusForbidden, requestlog.ActionBanned)
|
|
||||||
|
|
||||||
otherBans := server.Ledger.Bans(other)
|
|
||||||
|
|
||||||
wantAlerts(t, queue,
|
|
||||||
attackAlert(alerts.EventBan, start, client, attackBan,
|
|
||||||
requestlog.FormatTime(start.Add(7*24*time.Hour))),
|
|
||||||
attackAlert(alerts.EventPermanentBan, start.Add(time.Minute), client,
|
|
||||||
permanentBan, "permanent"),
|
|
||||||
attackAlert(alerts.EventBan, start.Add(time.Minute), otherClient, otherBans[0],
|
|
||||||
requestlog.FormatTime(start.Add(time.Minute+7*24*time.Hour))),
|
|
||||||
attackAlert(alerts.EventPermanentBan, start.Add(time.Minute+7*24*time.Hour),
|
|
||||||
otherClient, otherBans[1], "permanent"),
|
|
||||||
)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestObserveModeRaisesTheBanAlertsItWouldHave(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
s, clk, server, queue := startWithAlerts(t, map[string]string{
|
|
||||||
mode: observe,
|
|
||||||
rateLimitPerMinute: "2",
|
|
||||||
rulesDir: writeRules(t, testRules),
|
|
||||||
})
|
|
||||||
start := clk.Now()
|
|
||||||
|
|
||||||
// A ban for a clear sign of attack, which a request under it would make
|
|
||||||
// permanent.
|
|
||||||
group := netip.MustParsePrefix(ipv6Group)
|
|
||||||
attackBan, _ := server.Ledger.BanForAttack(group, start, bans.Notes{RuleID: "probe"})
|
|
||||||
|
|
||||||
// The third request breaks the limit, and so does the fourth, within the
|
|
||||||
// cooldown, which raises nothing. The probe is a clear sign of attack.
|
|
||||||
for range 4 {
|
|
||||||
s.get(client, http.StatusOK, requestlog.ActionForward)
|
|
||||||
}
|
|
||||||
|
|
||||||
s.request(otherClient, "/.env", http.StatusOK, requestlog.ActionForward)
|
|
||||||
line := s.get(ipv6Client, http.StatusOK, requestlog.ActionForward)
|
|
||||||
|
|
||||||
// No ban is made, and none made permanent.
|
|
||||||
if held := server.Ledger.Snapshot(); len(held) != 1 || held[0] != attackBan ||
|
|
||||||
line.BanExpires != requestlog.FormatTime(attackBan.Expires) {
|
|
||||||
t.Errorf("the ledger holds %+v, and the log line gives %s, want the ban "+
|
|
||||||
"for the attack alone, as it was", held, line.BanExpires)
|
|
||||||
}
|
|
||||||
|
|
||||||
waiting := queue.Snapshot().Waiting[alerts.DestinationWebhook]
|
|
||||||
if len(waiting) != 3 || queue.Suppressed() != 0 {
|
|
||||||
t.Fatalf("%d alerts wait and %d are held back, want 3 and 0: %+v",
|
|
||||||
len(waiting), queue.Suppressed(), waiting)
|
|
||||||
}
|
|
||||||
|
|
||||||
limitNotes, _ := waiting[0].Detail["notes"].(bans.Notes)
|
|
||||||
attackNotes, _ := waiting[1].Detail["notes"].(bans.Notes)
|
|
||||||
|
|
||||||
if limitNotes.Limit != 2 || limitNotes.Request.Path != "/" ||
|
|
||||||
attackNotes.Request.Path != "/.env" {
|
|
||||||
t.Errorf("the notes are %+v and %+v, want those of the broken limit and "+
|
|
||||||
"of the probe", limitNotes, attackNotes)
|
|
||||||
}
|
|
||||||
|
|
||||||
// Each alert is the one enforce mode would have raised, with mode
|
|
||||||
// observe in its detail.
|
|
||||||
want := []alerts.Alert{
|
|
||||||
banAlert(alerts.EventBan, start, client, bans.Ban{
|
|
||||||
Netblock: netip.MustParsePrefix(client + "/32"), Cause: bans.CauseLimit,
|
|
||||||
Reason: "requests per minute over the limit of 2", Notes: limitNotes,
|
|
||||||
}, requestlog.FormatTime(start.Add(time.Hour))),
|
|
||||||
attackAlert(alerts.EventBan, start, otherClient, bans.Ban{
|
|
||||||
Netblock: netip.MustParsePrefix(otherClient + "/32"), Notes: attackNotes,
|
|
||||||
}, requestlog.FormatTime(start.Add(7*24*time.Hour))),
|
|
||||||
attackAlert(alerts.EventPermanentBan, start, ipv6Client, attackBan, permanent),
|
|
||||||
}
|
|
||||||
for _, alert := range want {
|
|
||||||
alert.Detail["mode"] = observe
|
|
||||||
}
|
|
||||||
|
|
||||||
wantAlerts(t, queue, want...)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestObserveModeWorksOutABanOnlyWhenItsAlertWouldBeSent(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
s, _, _, queue := startWithAlerts(t, map[string]string{
|
|
||||||
mode: observe,
|
|
||||||
rateLimitPerMinute: "2",
|
|
||||||
rulesDir: writeRules(t, testRules),
|
|
||||||
alertMaxPerHour: "2",
|
|
||||||
})
|
|
||||||
|
|
||||||
// The client's third request breaks the limit, and raises the first
|
|
||||||
// alert of the hour. Its fourth is within the cooldown.
|
|
||||||
for range 4 {
|
|
||||||
s.get(client, http.StatusOK, requestlog.ActionForward)
|
|
||||||
}
|
|
||||||
|
|
||||||
// The other client's first probe raises the second. Its second probe is
|
|
||||||
// within the cooldown.
|
|
||||||
for range 2 {
|
|
||||||
s.request(otherClient, "/.env", http.StatusOK, requestlog.ActionForward)
|
|
||||||
}
|
|
||||||
|
|
||||||
// The IPv6 client's third request breaks the limit past the two alerts
|
|
||||||
// an hour.
|
|
||||||
for range 3 {
|
|
||||||
s.get(ipv6Client, http.StatusOK, requestlog.ActionForward)
|
|
||||||
}
|
|
||||||
|
|
||||||
// Had the ban been worked out for any of the requests within the
|
|
||||||
// cooldown or past the two an hour, its alert would have been raised,
|
|
||||||
// held back and counted.
|
|
||||||
waiting := queue.Snapshot().Waiting[alerts.DestinationWebhook]
|
|
||||||
if len(waiting) != 2 || queue.Suppressed() != 0 {
|
|
||||||
t.Errorf("%d alerts wait and %d are held back, want 2 and 0: %+v",
|
|
||||||
len(waiting), queue.Suppressed(), waiting)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// startWithAlerts is startWithClock with alerts to a webhook, which is
|
|
||||||
// never sent them, and returns the queue they wait in as well.
|
|
||||||
func startWithAlerts(
|
|
||||||
t *testing.T, env map[string]string,
|
|
||||||
) (*sender, *clock, *proxy.Server, *alerts.Queue) {
|
|
||||||
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,
|
|
||||||
alertWebhookURL: "https://alerts.example/smallwebwaf",
|
|
||||||
instanceName: alertInstance,
|
|
||||||
}
|
|
||||||
maps.Copy(settings, env)
|
|
||||||
|
|
||||||
addr, out, server, queue := startProxyWithAlerts(t, app.URL, "", clk.Now, settings)
|
|
||||||
|
|
||||||
return &sender{t: t, addr: addr, out: out}, clk, server, queue
|
|
||||||
}
|
|
||||||
|
|
||||||
// banAlert returns the alert for event, raised by a request from client at
|
|
||||||
// the time raised, for ban, with its netblock, cause, reason and notes,
|
|
||||||
// which ends at expires, as the log line gives it.
|
|
||||||
func banAlert(
|
|
||||||
event string, raised time.Time, client string, ban bans.Ban, expires string,
|
|
||||||
) alerts.Alert {
|
|
||||||
return alerts.Alert{
|
|
||||||
Instance: alertInstance,
|
|
||||||
Time: raised,
|
|
||||||
Event: event,
|
|
||||||
Client: netip.MustParseAddr(client),
|
|
||||||
Netblock: ban.Netblock,
|
|
||||||
Reason: ban.Reason,
|
|
||||||
Detail: map[string]any{
|
|
||||||
"cause": ban.Cause, "ban_expires": expires, "notes": ban.Notes,
|
|
||||||
},
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// attackAlert is banAlert for a ban for the probe rule of testRules, with
|
|
||||||
// the netblock and the notes of ban.
|
|
||||||
func attackAlert(
|
|
||||||
event string, raised time.Time, client string, ban bans.Ban, expires string,
|
|
||||||
) alerts.Alert {
|
|
||||||
return banAlert(event, raised, client, bans.Ban{
|
|
||||||
Netblock: ban.Netblock, Cause: bans.CauseAttack, Reason: "matched the rule probe",
|
|
||||||
Notes: ban.Notes,
|
|
||||||
}, expires)
|
|
||||||
}
|
|
||||||
|
|
||||||
// wantAlerts checks the alerts waiting in queue, in order.
|
|
||||||
func wantAlerts(t *testing.T, queue *alerts.Queue, want ...alerts.Alert) {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
got := queue.Snapshot().Waiting[alerts.DestinationWebhook]
|
|
||||||
if len(got) != len(want) {
|
|
||||||
t.Fatalf("%d alerts wait, want %d: %+v", len(got), len(want), got)
|
|
||||||
}
|
|
||||||
|
|
||||||
for i := range want {
|
|
||||||
if !reflect.DeepEqual(got[i], want[i]) {
|
|
||||||
t.Errorf("alert %d is\n%+v\nwant\n%+v", i, got[i], want[i])
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@@ -1,221 +0,0 @@
|
|||||||
package proxy
|
|
||||||
|
|
||||||
import (
|
|
||||||
"net/netip"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
"sneak.berlin/go/smallwebwaf/internal/alerts"
|
|
||||||
"sneak.berlin/go/smallwebwaf/internal/bans"
|
|
||||||
"sneak.berlin/go/smallwebwaf/internal/requestlog"
|
|
||||||
"sneak.berlin/go/smallwebwaf/internal/rules"
|
|
||||||
)
|
|
||||||
|
|
||||||
// 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. A
|
|
||||||
// request that makes the ban permanent, or in observe mode would have,
|
|
||||||
// raises the alert for it.
|
|
||||||
func (rq *request) banned(now time.Time) bool {
|
|
||||||
check := rq.h.ledger.Check
|
|
||||||
if rq.h.config.Observe {
|
|
||||||
check = rq.h.ledger.Find // the ban refuses nothing, and stays as it is
|
|
||||||
}
|
|
||||||
|
|
||||||
ban, banned, madePermanent := check(rq.client, now)
|
|
||||||
if banned {
|
|
||||||
rq.line.BanExpires = banExpires(ban)
|
|
||||||
}
|
|
||||||
|
|
||||||
if madePermanent {
|
|
||||||
ban.Expires = time.Time{} // the ban made permanent, which Find leaves as it is
|
|
||||||
rq.alertBan(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, and raises the alert for the ban it would
|
|
||||||
// have made, if that alert would be sent.
|
|
||||||
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
|
|
||||||
|
|
||||||
netblock := rq.h.netblock(rq.client)
|
|
||||||
if rq.h.config.Observe && !rq.wouldAlertBan(netblock, now, bans.CauseLimit) {
|
|
||||||
return true
|
|
||||||
}
|
|
||||||
|
|
||||||
notes := bans.Notes{
|
|
||||||
ASN: rq.line.ASN,
|
|
||||||
ASName: rq.line.ASName,
|
|
||||||
Country: rq.line.Country,
|
|
||||||
Limit: hit.Limit,
|
|
||||||
Window: hit.Window,
|
|
||||||
Count: hit.Requests,
|
|
||||||
Request: rq.noted(now),
|
|
||||||
Requests: rq.netblockRequests(netblock),
|
|
||||||
}
|
|
||||||
|
|
||||||
if rq.h.config.Observe {
|
|
||||||
ban, wouldBan := rq.h.ledger.WouldBanForLimit(netblock, now, notes)
|
|
||||||
if wouldBan {
|
|
||||||
rq.alertBan(ban)
|
|
||||||
}
|
|
||||||
|
|
||||||
return true
|
|
||||||
}
|
|
||||||
|
|
||||||
ban, made := rq.h.ledger.BanForLimit(netblock, now, notes)
|
|
||||||
rq.h.limiter.Reset(group)
|
|
||||||
rq.line.BanExpires = banExpires(ban)
|
|
||||||
|
|
||||||
if made {
|
|
||||||
rq.alertBan(ban)
|
|
||||||
}
|
|
||||||
|
|
||||||
return true
|
|
||||||
}
|
|
||||||
|
|
||||||
// banForAttack bans the client's netblock at now for a clear sign of
|
|
||||||
// attack, the match of rule, a ban rule. In observe mode it makes no ban,
|
|
||||||
// and raises the alert for the ban it would have made, if that alert
|
|
||||||
// would be sent.
|
|
||||||
func (rq *request) banForAttack(now time.Time, rule rules.Rule) {
|
|
||||||
netblock := rq.h.netblock(rq.client)
|
|
||||||
if rq.h.config.Observe && !rq.wouldAlertBan(netblock, now, bans.CauseAttack) {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
notes := bans.Notes{
|
|
||||||
ASN: rq.line.ASN,
|
|
||||||
ASName: rq.line.ASName,
|
|
||||||
Country: rq.line.Country,
|
|
||||||
RuleID: rule.ID,
|
|
||||||
Target: rule.Target,
|
|
||||||
Request: rq.noted(now),
|
|
||||||
Requests: rq.netblockRequests(netblock),
|
|
||||||
}
|
|
||||||
|
|
||||||
if rq.h.config.Observe {
|
|
||||||
ban, wouldBan := rq.h.ledger.WouldBanForAttack(netblock, now, notes)
|
|
||||||
if wouldBan {
|
|
||||||
rq.alertBan(ban)
|
|
||||||
}
|
|
||||||
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
ban, made := rq.h.ledger.BanForAttack(netblock, now, notes)
|
|
||||||
rq.line.BanExpires = banExpires(ban)
|
|
||||||
|
|
||||||
if made {
|
|
||||||
rq.alertBan(ban)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// wouldAlertBan reports whether the alert for a ban on netblock for cause
|
|
||||||
// made at now would be sent. In observe mode the ban the request would
|
|
||||||
// have made is worked out only then, at most once per
|
|
||||||
// SWWAF_ALERT_COOLDOWN and never with no webhook set: its notes count the
|
|
||||||
// netblock's requests, which can mean going through every client.
|
|
||||||
func (rq *request) wouldAlertBan(
|
|
||||||
netblock netip.Prefix, now time.Time, cause string,
|
|
||||||
) bool {
|
|
||||||
event := alerts.EventBan
|
|
||||||
if rq.h.ledger.WouldBePermanent(netblock, now, cause) {
|
|
||||||
event = alerts.EventPermanentBan
|
|
||||||
}
|
|
||||||
|
|
||||||
return rq.h.alerts.WouldSend(event, netblock)
|
|
||||||
}
|
|
||||||
|
|
||||||
// alertBan raises the alert for ban, which the request made, or made
|
|
||||||
// permanent: permanent_ban for a permanent ban, ban for another. Its
|
|
||||||
// detail gives the ban's cause, when it ends, and its notes, and in
|
|
||||||
// observe mode, where ban is the ban that would have been made, or made
|
|
||||||
// permanent, mode, observe.
|
|
||||||
func (rq *request) alertBan(ban bans.Ban) {
|
|
||||||
event := alerts.EventBan
|
|
||||||
if ban.Permanent() {
|
|
||||||
event = alerts.EventPermanentBan
|
|
||||||
}
|
|
||||||
|
|
||||||
detail := map[string]any{
|
|
||||||
"cause": ban.Cause, "ban_expires": banExpires(ban), "notes": ban.Notes,
|
|
||||||
}
|
|
||||||
if rq.h.config.Observe {
|
|
||||||
detail["mode"] = "observe"
|
|
||||||
}
|
|
||||||
|
|
||||||
rq.h.alerts.Raise(alerts.Alert{
|
|
||||||
Event: event,
|
|
||||||
Client: rq.client,
|
|
||||||
Netblock: ban.Netblock,
|
|
||||||
ASN: ban.Notes.ASN,
|
|
||||||
ASName: ban.Notes.ASName,
|
|
||||||
Country: ban.Notes.Country,
|
|
||||||
Reason: ban.Reason,
|
|
||||||
Detail: detail,
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
// noted is the request, refused at now with SWWAF_BAN_RESPONSE, or in
|
|
||||||
// observe mode as it would have been, as the notes of the ban it makes
|
|
||||||
// keep it.
|
|
||||||
func (rq *request) noted(now time.Time) bans.Request {
|
|
||||||
return 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(),
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// netblockRequests is how many requests netblock has sent since it was
|
|
||||||
// first seen, this one included: the histories count it only once it has
|
|
||||||
// ended.
|
|
||||||
func (rq *request) netblockRequests(netblock netip.Prefix) int64 {
|
|
||||||
return rq.h.limiter.Requests(netblock) + 1
|
|
||||||
}
|
|
||||||
|
|
||||||
// netblock is the netblock a ban on client covers: its IPv4 address,
|
|
||||||
// widened to SWWAF_BAN_SCOPE_V4_PREFIX, or the IPv6 group clientGroup
|
|
||||||
// counts it in.
|
|
||||||
func (h *handler) netblock(client netip.Addr) netip.Prefix {
|
|
||||||
addr := client.Unmap()
|
|
||||||
if addr.Is4() {
|
|
||||||
return netip.PrefixFrom(addr, 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)
|
|
||||||
}
|
|
||||||
@@ -1,474 +0,0 @@
|
|||||||
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),
|
|
||||||
Cause: bans.CauseLimit,
|
|
||||||
Reason: "requests per minute over the limit of 1",
|
|
||||||
Notes: bans.Notes{
|
|
||||||
ASN: asnDE,
|
|
||||||
ASName: asNameDE,
|
|
||||||
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: bans.EarlierBans{},
|
|
||||||
},
|
|
||||||
}
|
|
||||||
|
|
||||||
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 != (bans.EarlierBans{Limit: 1}) {
|
|
||||||
t.Errorf("bans %+v, want two, the second with one earlier ban for a limit", 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' AS numbers and 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()
|
|
||||||
|
|
||||||
line, got := s.requestWithBody(http.MethodGet, from, path, header, "", status, action)
|
|
||||||
|
|
||||||
return line, string(got.body)
|
|
||||||
}
|
|
||||||
|
|
||||||
// requestWithBody is requestWithHeader for a request with method, whose
|
|
||||||
// body is sent as it is after the headers, header holding its
|
|
||||||
// Content-Length or Transfer-Encoding. header may hold several lines,
|
|
||||||
// separated by "\r\n". It returns the whole answer.
|
|
||||||
func (s *sender) requestWithBody(
|
|
||||||
method, from, path, header, body string, status int, action string,
|
|
||||||
) (logLine, answer) {
|
|
||||||
s.t.Helper()
|
|
||||||
|
|
||||||
if header != "" {
|
|
||||||
header += "\r\n"
|
|
||||||
}
|
|
||||||
|
|
||||||
conn := dial(s.t, s.addr)
|
|
||||||
send(s.t, conn, method+" "+path+" HTTP/1.1\r\nHost: "+appHost+
|
|
||||||
"\r\nUser-Agent: "+userAgent+"\r\n"+forwardedFor+": "+from+"\r\n"+
|
|
||||||
header+"\r\n"+body)
|
|
||||||
|
|
||||||
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, got
|
|
||||||
}
|
|
||||||
@@ -46,7 +46,6 @@ func (b *requestBody) Read(p []byte) (int, error) {
|
|||||||
b.rq.refuse(refusal{
|
b.rq.refuse(refusal{
|
||||||
status: http.StatusRequestEntityTooLarge,
|
status: http.StatusRequestEntityTooLarge,
|
||||||
action: requestlog.ActionTooLarge,
|
action: requestlog.ActionTooLarge,
|
||||||
limit: "SWWAF_REQUEST_MAX_BYTES",
|
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -82,7 +81,6 @@ func (b *responseBody) Read(p []byte) (int, error) {
|
|||||||
b.rq.refuse(refusal{
|
b.rq.refuse(refusal{
|
||||||
status: http.StatusBadGateway,
|
status: http.StatusBadGateway,
|
||||||
action: requestlog.ActionTooLarge,
|
action: requestlog.ActionTooLarge,
|
||||||
limit: "SWWAF_RESPONSE_MAX_BYTES",
|
|
||||||
})
|
})
|
||||||
|
|
||||||
return n, errResponseTooLarge
|
return n, errResponseTooLarge
|
||||||
|
|||||||
@@ -1,7 +1,6 @@
|
|||||||
package proxy
|
package proxy
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"crypto/rand"
|
|
||||||
"net/http"
|
"net/http"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
"slices"
|
"slices"
|
||||||
@@ -49,33 +48,6 @@ func clientAddress(
|
|||||||
return client
|
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.
|
// ipv6GroupPrefix is the length of the IPv6 netblock that is one client.
|
||||||
const ipv6GroupPrefix = 64
|
const ipv6GroupPrefix = 64
|
||||||
|
|
||||||
|
|||||||
@@ -14,14 +14,10 @@ const (
|
|||||||
appHost = "app.example"
|
appHost = "app.example"
|
||||||
// client is the client's address, as a proxy names it.
|
// client is the client's address, as a proxy names it.
|
||||||
client = "203.0.113.9"
|
client = "203.0.113.9"
|
||||||
// forwardedFor is the header that lists the client and its proxies,
|
// 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"
|
||||||
forwardedFor = "X-Forwarded-For"
|
// secure is the scheme a client reached traefik with.
|
||||||
forwardedProto = "X-Forwarded-Proto"
|
|
||||||
// secure is the scheme a client reached traefik with, and plain the
|
|
||||||
// one smallwebwaf serves.
|
|
||||||
secure = "https"
|
secure = "https"
|
||||||
plain = "http"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
// appHeaders is what the app tells about the headers it received.
|
// appHeaders is what the app tells about the headers it received.
|
||||||
@@ -69,13 +65,13 @@ func TestClientAddressAndForwardedHeaders(t *testing.T) {
|
|||||||
func clientAddressCases() []clientAddressCase {
|
func clientAddressCases() []clientAddressCase {
|
||||||
trusted := map[string]string{trustedProxies: trustLocalhost}
|
trusted := map[string]string{trustedProxies: trustLocalhost}
|
||||||
forged := http.Header{
|
forged := http.Header{
|
||||||
forwardedFor: {client},
|
forwardedFor: {client},
|
||||||
"X-Forwarded-Host": {"forged.example"},
|
"X-Forwarded-Host": {"forged.example"},
|
||||||
forwardedProto: {secure},
|
"X-Forwarded-Proto": {secure},
|
||||||
"X-Real-Ip": {client},
|
"X-Real-Ip": {client},
|
||||||
}
|
}
|
||||||
replaced := appHeaders{
|
replaced := appHeaders{
|
||||||
ForwardedFor: localhost, ForwardedHost: appHost, ForwardedProto: plain,
|
ForwardedFor: localhost, ForwardedHost: appHost, ForwardedProto: "http",
|
||||||
}
|
}
|
||||||
|
|
||||||
return []clientAddressCase{{
|
return []clientAddressCase{{
|
||||||
@@ -91,10 +87,10 @@ func clientAddressCases() []clientAddressCase {
|
|||||||
"outside the trusted proxies from the right",
|
"outside the trusted proxies from the right",
|
||||||
env: trusted,
|
env: trusted,
|
||||||
header: http.Header{
|
header: http.Header{
|
||||||
forwardedFor: {"198.51.100.7, " + client + ", 10.0.0.2"},
|
forwardedFor: {"198.51.100.7, " + client + ", 10.0.0.2"},
|
||||||
"X-Forwarded-Host": {appHost},
|
"X-Forwarded-Host": {appHost},
|
||||||
forwardedProto: {secure},
|
"X-Forwarded-Proto": {secure},
|
||||||
"X-Real-Ip": {client},
|
"X-Real-Ip": {client},
|
||||||
},
|
},
|
||||||
wantClient: client,
|
wantClient: client,
|
||||||
wantApp: appHeaders{
|
wantApp: appHeaders{
|
||||||
@@ -142,8 +138,8 @@ func requestWithHeaders(
|
|||||||
Host: r.Host,
|
Host: r.Host,
|
||||||
ForwardedFor: r.Header.Get(forwardedFor),
|
ForwardedFor: r.Header.Get(forwardedFor),
|
||||||
ForwardedHost: r.Header.Get("X-Forwarded-Host"),
|
ForwardedHost: r.Header.Get("X-Forwarded-Host"),
|
||||||
ForwardedProto: r.Header.Get(forwardedProto),
|
ForwardedProto: r.Header.Get("X-Forwarded-Proto"),
|
||||||
RealIP: r.Header.Get("X-Real-IP"),
|
RealIP: r.Header.Get("X-Real-Ip"),
|
||||||
})
|
})
|
||||||
})
|
})
|
||||||
addr, out := startProxy(t, app.URL, env)
|
addr, out := startProxy(t, app.URL, env)
|
||||||
|
|||||||
@@ -1,21 +0,0 @@
|
|||||||
package proxy
|
|
||||||
|
|
||||||
import (
|
|
||||||
"slices"
|
|
||||||
)
|
|
||||||
|
|
||||||
// countryDenied reports whether the country lists refuse the request, by
|
|
||||||
// the client's country as it was looked up. A client without a country,
|
|
||||||
// or whose country cannot be found, is refused only by
|
|
||||||
// SWWAF_EXCLUSIVELY_ALLOWED_COUNTRIES.
|
|
||||||
func (rq *request) countryDenied() bool {
|
|
||||||
denied := rq.h.config.DeniedCountries
|
|
||||||
allowed := rq.h.config.ExclusivelyAllowedCountries
|
|
||||||
country := rq.line.Country
|
|
||||||
|
|
||||||
if slices.Contains(denied, country) {
|
|
||||||
return true
|
|
||||||
}
|
|
||||||
|
|
||||||
return len(allowed) > 0 && !slices.Contains(allowed, country)
|
|
||||||
}
|
|
||||||
@@ -1,373 +0,0 @@
|
|||||||
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)
|
|
||||||
|
|
||||||
// The AS number GeoJS gives unplaced, 64512, counts as unknown.
|
|
||||||
for i, sent := range []struct{ client, asn, asName, country string }{
|
|
||||||
{fromDE, asnDE, asNameDE, "DE"}, {fromKP, asnKP, asNameKP, "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.ASN != sent.asn || line.ASName != sent.asName ||
|
|
||||||
line.Country != sent.country {
|
|
||||||
t.Errorf("log line has %q, %q and %q, want %q, %q and %q",
|
|
||||||
line.ASN, line.ASName, line.Country,
|
|
||||||
sent.asn, sent.asName, 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 := []geojsAnswer{{IP: r.URL.Query().Get("ip"), CountryCode: "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 TestPrivateAddressIsNeverLookedUp(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
for _, tc := range []struct {
|
|
||||||
name string
|
|
||||||
env map[string]string
|
|
||||||
}{
|
|
||||||
{"no setting needs the lookup", nil},
|
|
||||||
{"a country list is set", map[string]string{deniedCountries: "kp"}},
|
|
||||||
} {
|
|
||||||
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)
|
|
||||||
|
|
||||||
// "" sends no X-Forwarded-For: the client is 127.0.0.1.
|
|
||||||
for i, sent := range []string{
|
|
||||||
"10.0.0.5", "192.168.1.9", "fd00::5", "", "169.254.0.9", "fe80::9",
|
|
||||||
} {
|
|
||||||
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)
|
|
||||||
|
|
||||||
for _, field := range []string{"asn", "as_name", "country"} {
|
|
||||||
value, present := line.fields[field]
|
|
||||||
if !present || value != "" {
|
|
||||||
t.Errorf("log line for %q has %s %v, want an empty one",
|
|
||||||
line.ClientIP, field, value)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// GeoJS is asked about up to 200 waiting clients at once, so once it
|
|
||||||
// has been asked about fromDE, which comes last, it has been asked
|
|
||||||
// about every client before it that waited for an answer.
|
|
||||||
req := newRequest(t, http.MethodGet, addr, "/", http.NoBody)
|
|
||||||
req.Header.Set(forwardedFor, fromDE)
|
|
||||||
wantStatus(t, do(t, req), http.StatusOK)
|
|
||||||
|
|
||||||
waitUntil(func() bool { return slices.Contains(asked(), fromDE) })
|
|
||||||
|
|
||||||
if got := asked(); !slices.Equal(got, []string{fromDE}) {
|
|
||||||
t.Errorf("GeoJS was asked about %v, want %s alone", got, fromDE)
|
|
||||||
}
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
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,
|
|
||||||
// each in an AS of its own, 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()
|
|
||||||
|
|
||||||
geojsURL, asked, release := startHeldGeoJS(t)
|
|
||||||
release()
|
|
||||||
|
|
||||||
return geojsURL, asked
|
|
||||||
}
|
|
||||||
|
|
||||||
// startHeldGeoJS is startGeoJS for a stand-in that answers nothing until
|
|
||||||
// release is called. Each request to it waits until then.
|
|
||||||
func startHeldGeoJS(t *testing.T) (string, func() []string, func()) {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
var (
|
|
||||||
asked struct {
|
|
||||||
mu sync.Mutex
|
|
||||||
addrs []string
|
|
||||||
}
|
|
||||||
released = make(chan struct{})
|
|
||||||
once sync.Once
|
|
||||||
)
|
|
||||||
|
|
||||||
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()
|
|
||||||
|
|
||||||
<-released
|
|
||||||
|
|
||||||
answers := make([]geojsAnswer, 0, len(addrs))
|
|
||||||
for _, addr := range addrs {
|
|
||||||
answers = append(answers, answerAbout(addr))
|
|
||||||
}
|
|
||||||
|
|
||||||
err := json.NewEncoder(w).Encode(answers)
|
|
||||||
if err != nil {
|
|
||||||
http.Error(w, err.Error(), http.StatusInternalServerError)
|
|
||||||
}
|
|
||||||
}))
|
|
||||||
t.Cleanup(geojs.Close)
|
|
||||||
|
|
||||||
release := func() { once.Do(func() { close(released) }) }
|
|
||||||
// Run before geojs.Close, which waits for every request to be answered.
|
|
||||||
t.Cleanup(release)
|
|
||||||
|
|
||||||
return geojs.URL, func() []string {
|
|
||||||
asked.mu.Lock()
|
|
||||||
defer asked.mu.Unlock()
|
|
||||||
|
|
||||||
return slices.Clone(asked.addrs)
|
|
||||||
}, release
|
|
||||||
}
|
|
||||||
|
|
||||||
// The AS numbers and names the stand-in for GeoJS gives fromDE and
|
|
||||||
// fromKP, as they are logged.
|
|
||||||
const (
|
|
||||||
asnDE = "AS64496"
|
|
||||||
asNameDE = "Example Net"
|
|
||||||
asnKP = "AS64511"
|
|
||||||
asNameKP = "Other Net"
|
|
||||||
)
|
|
||||||
|
|
||||||
// geojsAnswer is an answer of GeoJS about one address, with the fields
|
|
||||||
// smallwebwaf reads.
|
|
||||||
//
|
|
||||||
//nolint:tagliatelle // GeoJS's own names
|
|
||||||
type geojsAnswer struct {
|
|
||||||
IP string `json:"ip"`
|
|
||||||
ASN int `json:"asn"`
|
|
||||||
ASName string `json:"organization_name"`
|
|
||||||
CountryCode string `json:"country_code,omitempty"`
|
|
||||||
}
|
|
||||||
|
|
||||||
// answerAbout is what the stand-in for GeoJS answers about addr: for an
|
|
||||||
// address it cannot place, the AS number 64512 and the AS name Unknown
|
|
||||||
// with no country, as GeoJS does.
|
|
||||||
func answerAbout(addr string) geojsAnswer {
|
|
||||||
switch addr {
|
|
||||||
case fromDE:
|
|
||||||
return geojsAnswer{IP: addr, ASN: 64496, ASName: asNameDE, CountryCode: "DE"}
|
|
||||||
case fromKP:
|
|
||||||
return geojsAnswer{IP: addr, ASN: 64511, ASName: asNameKP, CountryCode: "KP"}
|
|
||||||
}
|
|
||||||
|
|
||||||
return geojsAnswer{IP: addr, ASN: 64512, ASName: "Unknown"}
|
|
||||||
}
|
|
||||||
@@ -1,56 +0,0 @@
|
|||||||
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())
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@@ -1,128 +0,0 @@
|
|||||||
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 client is not looked up. GeoJS
|
|
||||||
// answers about the client at its first request, and its later ones
|
|
||||||
// use that answer.
|
|
||||||
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),
|
|
||||||
ASN: asnDE,
|
|
||||||
ASName: asNameDE,
|
|
||||||
Country: "DE",
|
|
||||||
LookedUp: start,
|
|
||||||
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{}
|
|
||||||
}
|
|
||||||
@@ -55,7 +55,6 @@ func TestRequestBodyLimit(t *testing.T) {
|
|||||||
})
|
})
|
||||||
addr, out := startProxy(t, app.URL, map[string]string{
|
addr, out := startProxy(t, app.URL, map[string]string{
|
||||||
requestMaxBytes: sizeLimitSetting,
|
requestMaxBytes: sizeLimitSetting,
|
||||||
metricsToken: token,
|
|
||||||
})
|
})
|
||||||
|
|
||||||
var body io.Reader = bytes.NewReader(make([]byte, tc.size))
|
var body io.Reader = bytes.NewReader(make([]byte, tc.size))
|
||||||
@@ -67,13 +66,6 @@ func TestRequestBodyLimit(t *testing.T) {
|
|||||||
tc.want)
|
tc.want)
|
||||||
wantLine(t, out.requestLine(t), tc.want, tc.action)
|
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 {
|
if tc.refusedBeforeApp && calls.Load() != 0 {
|
||||||
t.Errorf("the app was called %d times, want never", calls.Load())
|
t.Errorf("the app was called %d times, want never", calls.Load())
|
||||||
}
|
}
|
||||||
@@ -114,7 +106,6 @@ func TestResponseBodyLimit(t *testing.T) {
|
|||||||
})
|
})
|
||||||
addr, out := startProxy(t, app.URL, map[string]string{
|
addr, out := startProxy(t, app.URL, map[string]string{
|
||||||
responseMaxBytes: sizeLimitSetting,
|
responseMaxBytes: sizeLimitSetting,
|
||||||
metricsToken: token,
|
|
||||||
})
|
})
|
||||||
|
|
||||||
got := get(t, addr, "/download")
|
got := get(t, addr, "/download")
|
||||||
@@ -132,13 +123,6 @@ func TestResponseBodyLimit(t *testing.T) {
|
|||||||
if line.UpstreamStatus != http.StatusOK {
|
if line.UpstreamStatus != http.StatusOK {
|
||||||
t.Errorf("log line has upstream_status %d", line.UpstreamStatus)
|
t.Errorf("log line has upstream_status %d", line.UpstreamStatus)
|
||||||
}
|
}
|
||||||
|
|
||||||
hits := 0
|
|
||||||
if tc.action == requestlog.ActionTooLarge {
|
|
||||||
hits = 1
|
|
||||||
}
|
|
||||||
|
|
||||||
wantLimitHits(t, addr, responseMaxBytes, hits)
|
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -1,70 +0,0 @@
|
|||||||
package proxy
|
|
||||||
|
|
||||||
import (
|
|
||||||
"context"
|
|
||||||
"net/http"
|
|
||||||
"net/netip"
|
|
||||||
|
|
||||||
"sneak.berlin/go/smallwebwaf/internal/lookup"
|
|
||||||
)
|
|
||||||
|
|
||||||
// The headers in which the app is passed the client's AS number and
|
|
||||||
// country while SWWAF_ADD_LOOKUP_HEADERS is set. Go writes every header
|
|
||||||
// name in this form, as it sends it and as it receives it, so X-Client-ASN
|
|
||||||
// arrives as X-Client-Asn, and Del removes a client's own whatever their
|
|
||||||
// case; header names are not case-sensitive.
|
|
||||||
const (
|
|
||||||
asnHeader = "X-Client-Asn"
|
|
||||||
countryHeader = "X-Client-Country"
|
|
||||||
)
|
|
||||||
|
|
||||||
// lookUp looks up the client's AS number and country, in the lookup
|
|
||||||
// database or through GeoJS, and notes them for the log line, unless
|
|
||||||
// SWWAF_LOOKUP_SOURCE is off or the client is on a private, loopback or
|
|
||||||
// link-local address, which no lookup can place. The lookup database
|
|
||||||
// answers at once. With GeoJS, while a setting needs the answer, a new
|
|
||||||
// client's request waits for it. ctx is the request's own context.
|
|
||||||
func (rq *request) lookUp(ctx context.Context) {
|
|
||||||
if rq.h.config.LookupSource == "off" || !canBePlaced(rq.client) {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
if rq.h.config.LookupSource == "file" {
|
|
||||||
rq.lookupAnswer = rq.h.lookupFile.LookUp(clientGroup(rq.client))
|
|
||||||
} else {
|
|
||||||
rq.lookupAnswer = rq.h.geojs.LookUp(ctx, clientGroup(rq.client))
|
|
||||||
}
|
|
||||||
|
|
||||||
rq.lookedUp = true
|
|
||||||
rq.line.ASN = rq.lookupAnswer.ASN
|
|
||||||
rq.line.ASName = rq.lookupAnswer.ASName
|
|
||||||
rq.line.Country = rq.lookupAnswer.Country
|
|
||||||
}
|
|
||||||
|
|
||||||
// addLookup adds answer, an answer about a client from the lookup
|
|
||||||
// database or GeoJS, to the client's history, and to the notes of the bans
|
|
||||||
// on its netblock that have no AS number, AS name or country yet.
|
|
||||||
func (h *handler) addLookup(answer lookup.Answer) {
|
|
||||||
h.limiter.AddLookup(answer.Client, answer.Answered,
|
|
||||||
answer.ASN, answer.ASName, answer.Country)
|
|
||||||
h.ledger.AddLookup(h.netblock(answer.Client.Addr()),
|
|
||||||
answer.ASN, answer.ASName, answer.Country)
|
|
||||||
}
|
|
||||||
|
|
||||||
// setLookupHeaders sets the headers in which the app is passed the
|
|
||||||
// client's AS number and country, leaving out one that is unknown.
|
|
||||||
func setLookupHeaders(header http.Header, asn, country string) {
|
|
||||||
if asn != "" {
|
|
||||||
header.Set(asnHeader, asn)
|
|
||||||
}
|
|
||||||
|
|
||||||
if country != "" {
|
|
||||||
header.Set(countryHeader, country)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// canBePlaced reports whether a lookup can place addr: private, loopback
|
|
||||||
// and link-local addresses have no AS number or country.
|
|
||||||
func canBePlaced(addr netip.Addr) bool {
|
|
||||||
return !addr.IsPrivate() && !addr.IsLoopback() && !addr.IsLinkLocalUnicast()
|
|
||||||
}
|
|
||||||
@@ -1,373 +0,0 @@
|
|||||||
package proxy_test
|
|
||||||
|
|
||||||
import (
|
|
||||||
"net"
|
|
||||||
"net/http"
|
|
||||||
"net/http/httptest"
|
|
||||||
"net/netip"
|
|
||||||
"path/filepath"
|
|
||||||
"slices"
|
|
||||||
"sync"
|
|
||||||
"testing"
|
|
||||||
"testing/synctest"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
"sneak.berlin/go/smallwebwaf/internal/alerts"
|
|
||||||
"sneak.berlin/go/smallwebwaf/internal/lookup"
|
|
||||||
"sneak.berlin/go/smallwebwaf/internal/lookup/lookuptest"
|
|
||||||
"sneak.berlin/go/smallwebwaf/internal/requestlog"
|
|
||||||
)
|
|
||||||
|
|
||||||
// asnAndCountry is what a lookup gives a client: its AS number, AS name
|
|
||||||
// and country.
|
|
||||||
type asnAndCountry struct{ asn, asName, country string }
|
|
||||||
|
|
||||||
func TestEveryClientIsLookedUpWithoutWaitingWhileNoSettingNeedsIt(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
// The stand-in for GeoJS answers only once released. A request that
|
|
||||||
// waited for it would wait an hour, and get no answer within
|
|
||||||
// waitLimit.
|
|
||||||
geojsURL, asked, release := startHeldGeoJS(t)
|
|
||||||
s, _, server := startWithClock(t, geojsURL, map[string]string{
|
|
||||||
lookupTimeout: "1h",
|
|
||||||
rateLimitPerMinute: "1",
|
|
||||||
})
|
|
||||||
|
|
||||||
// fromDE's second request breaks the limit and bans it, and fromKP
|
|
||||||
// comes too. None waits for GeoJS.
|
|
||||||
for _, line := range []logLine{
|
|
||||||
s.get(fromDE, http.StatusOK, requestlog.ActionForward),
|
|
||||||
s.get(fromDE, http.StatusForbidden, requestlog.ActionRateLimited),
|
|
||||||
s.get(fromKP, http.StatusOK, requestlog.ActionForward),
|
|
||||||
} {
|
|
||||||
got := asnAndCountry{line.ASN, line.ASName, line.Country}
|
|
||||||
if got != (asnAndCountry{}) {
|
|
||||||
t.Errorf("log line has %+v before GeoJS answered, want nothing", got)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// Once GeoJS answers, each answer reaches the client's history, and
|
|
||||||
// fromDE's reaches the notes of its ban.
|
|
||||||
release()
|
|
||||||
|
|
||||||
netblock := netip.MustParsePrefix(fromDE + "/32")
|
|
||||||
|
|
||||||
waitUntil(func() bool {
|
|
||||||
return historyOf(t, server, fromDE).ASN != "" &&
|
|
||||||
historyOf(t, server, fromKP).ASN != "" &&
|
|
||||||
server.Ledger.Bans(netblock)[0].Notes.ASN != ""
|
|
||||||
})
|
|
||||||
|
|
||||||
de := asnAndCountry{asnDE, asNameDE, "DE"}
|
|
||||||
|
|
||||||
for addr, want := range map[string]asnAndCountry{
|
|
||||||
fromDE: de, fromKP: {asnKP, asNameKP, "KP"},
|
|
||||||
} {
|
|
||||||
h := historyOf(t, server, addr)
|
|
||||||
if got := (asnAndCountry{h.ASN, h.ASName, h.Country}); got != want {
|
|
||||||
t.Errorf("%s's history has %+v, want %+v", addr, got, want)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
notes := server.Ledger.Bans(netblock)[0].Notes
|
|
||||||
if got := (asnAndCountry{notes.ASN, notes.ASName, notes.Country}); got != de {
|
|
||||||
t.Errorf("the ban's notes have %+v, want %+v", got, de)
|
|
||||||
}
|
|
||||||
|
|
||||||
// GeoJS was asked about each client once, fromKP after fromDE, whose
|
|
||||||
// request was under way when fromKP came.
|
|
||||||
if got := asked(); !slices.Equal(got, []string{fromDE, fromKP}) {
|
|
||||||
t.Errorf("GeoJS was asked about %v, want %s and %s", got, fromDE, fromKP)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestASNumberAndNameInTheLogLineTheHistoryTheBanNotesAndTheAlert(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
geojsURL, _ := startGeoJS(t)
|
|
||||||
app := startApp(t, func(http.ResponseWriter, *http.Request) {})
|
|
||||||
clk := &clock{now: time.Date(2026, 10, 6, 0, 0, 0, 0, time.UTC)}
|
|
||||||
addr, out, server, queue := startProxyWithAlerts(t, app.URL, geojsURL, clk.Now,
|
|
||||||
map[string]string{
|
|
||||||
trustedProxies: trustLocalhost,
|
|
||||||
alertWebhookURL: "https://alerts.example/smallwebwaf",
|
|
||||||
rateLimitPerMinute: "1",
|
|
||||||
})
|
|
||||||
s := &sender{t: t, addr: addr, out: out}
|
|
||||||
|
|
||||||
// The answer is kept before the requests, so GeoJS is not asked, and
|
|
||||||
// gives no answer of its own.
|
|
||||||
netblock := netip.MustParsePrefix(fromDE + "/32")
|
|
||||||
server.GeoJS.Load([]lookup.Answer{{
|
|
||||||
Client: netblock, ASN: asnDE, ASName: asNameDE, Country: "DE",
|
|
||||||
Answered: clk.Now(), Used: clk.Now(),
|
|
||||||
}})
|
|
||||||
|
|
||||||
want := asnAndCountry{asnDE, asNameDE, "DE"}
|
|
||||||
|
|
||||||
for _, line := range []logLine{
|
|
||||||
s.get(fromDE, http.StatusOK, requestlog.ActionForward),
|
|
||||||
s.get(fromDE, http.StatusForbidden, requestlog.ActionRateLimited),
|
|
||||||
} {
|
|
||||||
if got := (asnAndCountry{line.ASN, line.ASName, line.Country}); got != want {
|
|
||||||
t.Errorf("log line has %+v, want %+v", got, want)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
h := historyOf(t, server, fromDE)
|
|
||||||
if got := (asnAndCountry{h.ASN, h.ASName, h.Country}); got != want ||
|
|
||||||
!h.LookedUp.Equal(clk.Now()) {
|
|
||||||
t.Errorf("history has %+v, looked up at %s; want %+v, at %s",
|
|
||||||
got, h.LookedUp, want, clk.Now())
|
|
||||||
}
|
|
||||||
|
|
||||||
notes := server.Ledger.Bans(netblock)[0].Notes
|
|
||||||
if got := (asnAndCountry{notes.ASN, notes.ASName, notes.Country}); got != want {
|
|
||||||
t.Errorf("the ban's notes have %+v, want %+v", got, want)
|
|
||||||
}
|
|
||||||
|
|
||||||
waiting := queue.Snapshot().Waiting[alerts.DestinationWebhook]
|
|
||||||
if len(waiting) != 1 {
|
|
||||||
t.Fatalf("alerts waiting %+v, want the ban's alone", waiting)
|
|
||||||
}
|
|
||||||
|
|
||||||
alert := waiting[0]
|
|
||||||
if got := (asnAndCountry{alert.ASN, alert.ASName, alert.Country}); got != want {
|
|
||||||
t.Errorf("the ban's alert has %+v, want %+v", got, want)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestLookupSourceOffLooksNoClientUp(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
geojsURL, asked := startGeoJS(t)
|
|
||||||
s, clk, server := startWithClock(t, geojsURL, map[string]string{lookupSource: "off"})
|
|
||||||
|
|
||||||
// Even an answer kept from before is not used.
|
|
||||||
server.GeoJS.Load([]lookup.Answer{{
|
|
||||||
Client: netip.MustParsePrefix(fromDE + "/32"), ASN: asnDE, ASName: asNameDE,
|
|
||||||
Country: "DE", Answered: clk.Now(), Used: clk.Now(),
|
|
||||||
}})
|
|
||||||
|
|
||||||
for _, from := range []string{fromDE, fromKP} {
|
|
||||||
line := s.get(from, http.StatusOK, requestlog.ActionForward)
|
|
||||||
|
|
||||||
got := asnAndCountry{line.ASN, line.ASName, line.Country}
|
|
||||||
if got != (asnAndCountry{}) {
|
|
||||||
t.Errorf("log line has %+v, want nothing", got)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
if h := historyOf(t, server, fromDE); h.ASN != "" || !h.LookedUp.IsZero() {
|
|
||||||
t.Errorf("history has %q, looked up at %s, want no lookup", h.ASN, h.LookedUp)
|
|
||||||
}
|
|
||||||
|
|
||||||
if len(asked()) != 0 {
|
|
||||||
t.Errorf("GeoJS was asked about %v, want nothing", asked())
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestClientsAreLookedUpInTheLookupDatabaseAndGeoJSIsNotAsked(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
geojsURL, asked := startGeoJS(t)
|
|
||||||
path := filepath.Join(t.TempDir(), "ipinfo_lite.mmdb")
|
|
||||||
lookuptest.Write(t, path, map[string]lookuptest.Network{
|
|
||||||
fromDE + "/32": {ASN: asnDE, ASName: asNameDE, Country: "DE"},
|
|
||||||
fromKP + "/32": {ASN: asnKP, ASName: asNameKP, Country: "KP"},
|
|
||||||
})
|
|
||||||
s, clk, server := startWithClock(t, geojsURL, map[string]string{
|
|
||||||
lookupSource: "file",
|
|
||||||
lookupDBPath: path,
|
|
||||||
allowedCountries: "DE",
|
|
||||||
rateLimitPerMinute: "1",
|
|
||||||
})
|
|
||||||
|
|
||||||
// fromDE's second request breaks the limit and bans it. The list
|
|
||||||
// refuses fromKP, and unplaced, which the file does not hold.
|
|
||||||
de := asnAndCountry{asnDE, asNameDE, "DE"}
|
|
||||||
|
|
||||||
for _, tc := range []struct {
|
|
||||||
line logLine
|
|
||||||
want asnAndCountry
|
|
||||||
}{
|
|
||||||
{s.get(fromDE, http.StatusOK, requestlog.ActionForward), de},
|
|
||||||
{s.get(fromDE, http.StatusForbidden, requestlog.ActionRateLimited), de},
|
|
||||||
{
|
|
||||||
s.get(fromKP, http.StatusForbidden, requestlog.ActionCountryDenied),
|
|
||||||
asnAndCountry{asnKP, asNameKP, "KP"},
|
|
||||||
},
|
|
||||||
{
|
|
||||||
s.get(unplaced, http.StatusForbidden, requestlog.ActionCountryDenied),
|
|
||||||
asnAndCountry{},
|
|
||||||
},
|
|
||||||
} {
|
|
||||||
got := asnAndCountry{tc.line.ASN, tc.line.ASName, tc.line.Country}
|
|
||||||
if got != tc.want {
|
|
||||||
t.Errorf("log line has %+v, want %+v", got, tc.want)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
h := historyOf(t, server, fromDE)
|
|
||||||
if got := (asnAndCountry{h.ASN, h.ASName, h.Country}); got != de ||
|
|
||||||
!h.LookedUp.Equal(clk.Now()) {
|
|
||||||
t.Errorf("history has %+v, looked up at %s; want %+v, at %s",
|
|
||||||
got, h.LookedUp, de, clk.Now())
|
|
||||||
}
|
|
||||||
|
|
||||||
notes := server.Ledger.Bans(netip.MustParsePrefix(fromDE + "/32"))[0].Notes
|
|
||||||
if got := (asnAndCountry{notes.ASN, notes.ASName, notes.Country}); got != de {
|
|
||||||
t.Errorf("the ban's notes have %+v, want %+v", got, de)
|
|
||||||
}
|
|
||||||
|
|
||||||
if len(asked()) != 0 {
|
|
||||||
t.Errorf("GeoJS was asked about %v, want nothing", asked())
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestLookupHeadersArePassedToTheAppAndTheClientsOwnRemoved(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
var (
|
|
||||||
mu sync.Mutex
|
|
||||||
got [][2][]string // each request's X-Client-ASN and X-Client-Country
|
|
||||||
)
|
|
||||||
|
|
||||||
app := startApp(t, func(_ http.ResponseWriter, r *http.Request) {
|
|
||||||
mu.Lock()
|
|
||||||
defer mu.Unlock()
|
|
||||||
|
|
||||||
got = append(got, [2][]string{
|
|
||||||
r.Header.Values("X-Client-Asn"), r.Header.Values("X-Client-Country"),
|
|
||||||
})
|
|
||||||
})
|
|
||||||
geojsURL, _ := startGeoJS(t)
|
|
||||||
addr, out, _ := startProxyWithClock(t, app.URL, geojsURL, time.Now, map[string]string{
|
|
||||||
trustedProxies: trustLocalhost,
|
|
||||||
addLookupHeaders: "true",
|
|
||||||
})
|
|
||||||
s := &sender{t: t, addr: addr, out: out}
|
|
||||||
|
|
||||||
// Each client sends headers of its own. fromDE's first request waits
|
|
||||||
// for its answer, which the app is passed; unplaced has none to pass,
|
|
||||||
// and a client on a private address is not looked up.
|
|
||||||
for _, from := range []string{fromDE, unplaced, "10.0.0.8"} {
|
|
||||||
s.requestWithHeader(from, "/", clientsOwnLookupHeaders,
|
|
||||||
http.StatusOK, requestlog.ActionForward)
|
|
||||||
}
|
|
||||||
|
|
||||||
mu.Lock()
|
|
||||||
defer mu.Unlock()
|
|
||||||
|
|
||||||
want := [][2][]string{{{asnDE}, {"DE"}}, {nil, nil}, {nil, nil}}
|
|
||||||
if !slices.EqualFunc(got, want, func(a, b [2][]string) bool {
|
|
||||||
return slices.Equal(a[0], b[0]) && slices.Equal(a[1], b[1])
|
|
||||||
}) {
|
|
||||||
t.Errorf("the app was passed %v, want %v", got, want)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestClientsOwnLookupHeadersAreRemovedWhileTheSettingIsOff(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
var (
|
|
||||||
mu sync.Mutex
|
|
||||||
asn, country []string
|
|
||||||
)
|
|
||||||
|
|
||||||
app := startApp(t, func(_ http.ResponseWriter, r *http.Request) {
|
|
||||||
mu.Lock()
|
|
||||||
defer mu.Unlock()
|
|
||||||
|
|
||||||
asn, country = r.Header.Values("X-Client-Asn"), r.Header.Values("X-Client-Country")
|
|
||||||
})
|
|
||||||
geojsURL, _ := startGeoJS(t)
|
|
||||||
addr, out, _ := startProxyWithClock(t, app.URL, geojsURL, time.Now, map[string]string{
|
|
||||||
trustedProxies: trustLocalhost,
|
|
||||||
})
|
|
||||||
s := &sender{t: t, addr: addr, out: out}
|
|
||||||
|
|
||||||
s.requestWithHeader(fromDE, "/", clientsOwnLookupHeaders,
|
|
||||||
http.StatusOK, requestlog.ActionForward)
|
|
||||||
|
|
||||||
mu.Lock()
|
|
||||||
defer mu.Unlock()
|
|
||||||
|
|
||||||
if asn != nil || country != nil {
|
|
||||||
t.Errorf("the app was passed X-Client-ASN %v and X-Client-Country %v, want neither",
|
|
||||||
asn, country)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestRequestWaitsAsLongAsTheLookupTimeoutSays(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
// The test runs in a synctest bubble, where the time package runs on a
|
|
||||||
// clock of the test's own: the wait lasts exactly as long as it should,
|
|
||||||
// however slowly the test process runs. Nothing in it may wait on the
|
|
||||||
// network, which would keep that clock from moving on: the request is
|
|
||||||
// handed to the proxy's handler, and GeoJS is one that never answers.
|
|
||||||
synctest.Test(t, func(t *testing.T) {
|
|
||||||
// Not the default second. The exclusive list needs the answer, and
|
|
||||||
// the app is never reached.
|
|
||||||
const timeout = 3 * time.Second
|
|
||||||
|
|
||||||
server, out, _ := newProxy(t, "http://app.invalid", unansweredGeoJSURL,
|
|
||||||
time.Now, map[string]string{
|
|
||||||
lookupTimeout: timeout.String(),
|
|
||||||
allowedCountries: "DE",
|
|
||||||
})
|
|
||||||
|
|
||||||
req := httptest.NewRequestWithContext(t.Context(), http.MethodGet, "/",
|
|
||||||
http.NoBody)
|
|
||||||
req.RemoteAddr = net.JoinHostPort(fromDE, "1234")
|
|
||||||
began := time.Now()
|
|
||||||
|
|
||||||
server.Handler.ServeHTTP(httptest.NewRecorder(), req)
|
|
||||||
|
|
||||||
if waited := time.Since(began); waited != timeout {
|
|
||||||
t.Errorf("the request waited %s for its answer, want %s", waited, timeout)
|
|
||||||
}
|
|
||||||
|
|
||||||
// Without an answer, the client is in no country the list allows.
|
|
||||||
wantLine(t, out.requestLine(t), http.StatusForbidden,
|
|
||||||
requestlog.ActionCountryDenied)
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
// unansweredGeoJSURL is where a GeoJS that never answers is asked: a
|
|
||||||
// request to it waits, without the network, until it is abandoned.
|
|
||||||
// TestMain registers it with Go's default transport, through which GeoJS
|
|
||||||
// is asked.
|
|
||||||
const unansweredGeoJSURL = "unanswered://geojs/v1/ip/geo.json"
|
|
||||||
|
|
||||||
func TestMain(m *testing.M) {
|
|
||||||
transport, _ := http.DefaultTransport.(*http.Transport)
|
|
||||||
transport.RegisterProtocol("unanswered", unansweredGeoJS{})
|
|
||||||
|
|
||||||
m.Run()
|
|
||||||
}
|
|
||||||
|
|
||||||
// unansweredGeoJS is the GeoJS at unansweredGeoJSURL.
|
|
||||||
type unansweredGeoJS struct{}
|
|
||||||
|
|
||||||
// RoundTrip waits until req is abandoned.
|
|
||||||
func (unansweredGeoJS) RoundTrip(req *http.Request) (*http.Response, error) {
|
|
||||||
<-req.Context().Done()
|
|
||||||
|
|
||||||
return nil, req.Context().Err()
|
|
||||||
}
|
|
||||||
|
|
||||||
// clientsOwnLookupHeaders are the X-Client-ASN and X-Client-Country a
|
|
||||||
// client sends of its own, each twice, in two cases.
|
|
||||||
const clientsOwnLookupHeaders = "X-Client-ASN: AS1\r\nx-client-asn: AS2\r\n" +
|
|
||||||
"X-CLIENT-COUNTRY: KP\r\nx-client-country: CN"
|
|
||||||
|
|
||||||
// waitUntil waits until done reports true, for at most waitLimit.
|
|
||||||
func waitUntil(done func() bool) {
|
|
||||||
deadline := time.Now().Add(waitLimit)
|
|
||||||
for !done() && time.Now().Before(deadline) {
|
|
||||||
time.Sleep(pollInterval)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@@ -1,545 +0,0 @@
|
|||||||
package proxy_test
|
|
||||||
|
|
||||||
import (
|
|
||||||
"io"
|
|
||||||
"net/http"
|
|
||||||
"net/http/httptest"
|
|
||||||
"net/netip"
|
|
||||||
"strconv"
|
|
||||||
"strings"
|
|
||||||
"sync"
|
|
||||||
"sync/atomic"
|
|
||||||
"testing"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
"sneak.berlin/go/smallwebwaf/internal/bans"
|
|
||||||
"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",instance="app",status_class="2xx"}`
|
|
||||||
notFound := `{action="admin",instance="app",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{instance="app"}`, 2)
|
|
||||||
wantMetric(t, metrics,
|
|
||||||
`smallwebwaf_upstream_duration_seconds_count{instance="app"}`, 1)
|
|
||||||
wantMetric(t, metrics, `smallwebwaf_requests_in_flight{instance="app"}`, 1)
|
|
||||||
metric(t, metrics, `go_goroutines{instance="app"}`)
|
|
||||||
metric(t, metrics, `process_start_time_seconds{instance="app"}`)
|
|
||||||
|
|
||||||
// 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{instance="app"}`, 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{instance="app"}`, 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",instance="app",status_class="none"}`, 1)
|
|
||||||
wantMetric(t, metrics,
|
|
||||||
`smallwebwaf_rate_limit_hits_total{instance="app",window="minute"}`, 1)
|
|
||||||
wantMetric(t, metrics, `smallwebwaf_offences_total{instance="app",kind="limit"}`, 1)
|
|
||||||
wantMetric(t, metrics, `smallwebwaf_bans_made_total{cause="limit",instance="app"}`, 1)
|
|
||||||
wantMetric(t, metrics, `smallwebwaf_active_bans{instance="app"}`, 1)
|
|
||||||
wantMetric(t, metrics, `smallwebwaf_permanent_bans{instance="app"}`, 0)
|
|
||||||
|
|
||||||
clk.advance(time.Hour)
|
|
||||||
wantMetric(t, s.scrape(scraper), `smallwebwaf_active_bans{instance="app"}`, 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{instance="app",window="minute"}`, 2)
|
|
||||||
wantMetric(t, metrics, `smallwebwaf_offences_total{instance="app",kind="limit"}`, 2)
|
|
||||||
wantMetric(t, metrics, `smallwebwaf_bans_made_total{cause="limit",instance="app"}`, 2)
|
|
||||||
wantMetric(t, metrics, `smallwebwaf_active_bans{instance="app"}`, 1)
|
|
||||||
wantMetric(t, metrics, `smallwebwaf_permanent_bans{instance="app"}`, 1)
|
|
||||||
// denied, client, and the scraper as of its earlier requests.
|
|
||||||
wantMetric(t, metrics, `smallwebwaf_tracked_clients{instance="app"}`, 3)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestMetricsCountTheBansAnAdminMakes(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
const scraper = "192.0.2.200" // in SWWAF_RATE_LIMIT_EXEMPT_NETS
|
|
||||||
|
|
||||||
s, clk, server := startWithClock(t, "", map[string]string{
|
|
||||||
metricsToken: token,
|
|
||||||
adminToken: adminSecret,
|
|
||||||
rateLimitExemptNets: scraper,
|
|
||||||
})
|
|
||||||
|
|
||||||
const admins = `smallwebwaf_bans_made_total{cause="admin",instance="app"}`
|
|
||||||
|
|
||||||
wantMetric(t, s.scrape(scraper), admins, 0)
|
|
||||||
|
|
||||||
// As an admin's edit of bans.json that adds a ban is taken in.
|
|
||||||
server.Ledger.LoadEdit([]bans.Ban{{
|
|
||||||
Netblock: netip.MustParsePrefix(client + "/32"),
|
|
||||||
Start: clk.Now(),
|
|
||||||
}})
|
|
||||||
|
|
||||||
wantMetric(t, s.scrape(scraper), admins, 1)
|
|
||||||
|
|
||||||
// And a ban made through the endpoint.
|
|
||||||
s.admin(http.MethodPost, proxy.BansPath, banOtherClient, http.StatusOK)
|
|
||||||
wantMetric(t, s.scrape(scraper), admins, 2)
|
|
||||||
}
|
|
||||||
|
|
||||||
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",
|
|
||||||
}
|
|
||||||
geojsURL, _ := startGeoJS(t)
|
|
||||||
addr, out, server := startProxyWithClock(t, app.URL, geojsURL, 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",instance="app"}`, 3)
|
|
||||||
wantMetric(t, metrics,
|
|
||||||
`smallwebwaf_country_requests_total{country="DE",instance="app"}`, 2)
|
|
||||||
wantMetric(t, metrics,
|
|
||||||
`smallwebwaf_country_requests_total{country="other",instance="app"}`, 1)
|
|
||||||
wantMetric(t, metrics,
|
|
||||||
`smallwebwaf_country_list_refusals_total{country="KP",instance="app"}`, 3)
|
|
||||||
wantMetric(t, metrics,
|
|
||||||
`smallwebwaf_country_request_bytes_total{country="KP",instance="app"}`, 0)
|
|
||||||
wantMetric(t, metrics,
|
|
||||||
`smallwebwaf_country_request_bytes_total{country="DE",instance="app"}`, 6)
|
|
||||||
wantMetric(t, metrics,
|
|
||||||
`smallwebwaf_country_response_bytes_total{country="KP",instance="app"}`,
|
|
||||||
float64(3*len("Forbidden\n")))
|
|
||||||
wantMetric(t, metrics,
|
|
||||||
`smallwebwaf_country_response_bytes_total{country="other",instance="app"}`,
|
|
||||||
float64(len("hello")))
|
|
||||||
wantNoSeries(t, metrics,
|
|
||||||
`smallwebwaf_country_requests_total{country="FR",instance="app"}`)
|
|
||||||
|
|
||||||
// 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",instance="app"}`, 3)
|
|
||||||
wantMetric(t, metrics,
|
|
||||||
`smallwebwaf_country_requests_total{country="FR",instance="app"}`, 2)
|
|
||||||
wantMetric(t, metrics,
|
|
||||||
`smallwebwaf_country_requests_total{country="other",instance="app"}`, 2)
|
|
||||||
wantNoSeries(t, metrics,
|
|
||||||
`smallwebwaf_country_requests_total{country="DE",instance="app"}`)
|
|
||||||
wantNoSeries(t, metrics,
|
|
||||||
`smallwebwaf_country_request_bytes_total{country="DE",instance="app"}`)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestMetricsByASNumberKeepTheBusiestAndCountTheRestAsOther(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
geojsURL, _ := startGeoJS(t)
|
|
||||||
s, clk, server := startWithClock(t, geojsURL, map[string]string{
|
|
||||||
metricsToken: token,
|
|
||||||
metricsTopN: "1",
|
|
||||||
})
|
|
||||||
|
|
||||||
// The answers are kept before the requests, so that GeoJS gives none
|
|
||||||
// of its own. Each client is in an AS of its own.
|
|
||||||
answer := func(addr, asn string) lookup.Answer {
|
|
||||||
return lookup.Answer{
|
|
||||||
Client: netip.MustParsePrefix(addr + "/32"), ASN: asn,
|
|
||||||
Answered: clk.Now(), Used: clk.Now(),
|
|
||||||
}
|
|
||||||
}
|
|
||||||
server.GeoJS.Load([]lookup.Answer{
|
|
||||||
answer(fromDE, "AS64501"), answer(fromKP, "AS64502"),
|
|
||||||
})
|
|
||||||
|
|
||||||
// With one AS number of its own, the other is counted as other. The
|
|
||||||
// metrics are asked for from a private address, which has no AS number.
|
|
||||||
s.get(fromDE, http.StatusOK, requestlog.ActionForward)
|
|
||||||
s.get(fromDE, http.StatusOK, requestlog.ActionForward)
|
|
||||||
s.get(fromKP, http.StatusOK, requestlog.ActionForward)
|
|
||||||
|
|
||||||
metrics := s.scrape("10.0.0.9")
|
|
||||||
wantMetric(t, metrics,
|
|
||||||
`smallwebwaf_asn_requests_total{asn="AS64501",instance="app"}`, 2)
|
|
||||||
wantMetric(t, metrics,
|
|
||||||
`smallwebwaf_asn_requests_total{asn="other",instance="app"}`, 1)
|
|
||||||
wantMetric(t, metrics,
|
|
||||||
`smallwebwaf_asn_request_bytes_total{asn="AS64501",instance="app"}`, 0)
|
|
||||||
wantMetric(t, metrics,
|
|
||||||
`smallwebwaf_asn_response_bytes_total{asn="other",instance="app"}`, 0)
|
|
||||||
wantNoSeries(t, metrics,
|
|
||||||
`smallwebwaf_asn_requests_total{asn="AS64502",instance="app"}`)
|
|
||||||
}
|
|
||||||
|
|
||||||
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{instance="app"}`) == 0 &&
|
|
||||||
time.Now().Before(deadline) {
|
|
||||||
time.Sleep(pollInterval)
|
|
||||||
|
|
||||||
metrics = scrape(t, addr)
|
|
||||||
}
|
|
||||||
|
|
||||||
wantMetric(t, metrics, `smallwebwaf_geojs_requests_total{instance="app"}`, 1)
|
|
||||||
wantMetric(t, metrics, `smallwebwaf_geojs_failures_total{instance="app"}`, 1)
|
|
||||||
wantMetric(t, metrics, `smallwebwaf_geojs_unanswered_total{instance="app"}`, 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{instance="app",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{instance="app",limit="` +
|
|
||||||
limit + `"}`
|
|
||||||
metrics := scrape(t, addr)
|
|
||||||
|
|
||||||
if hits == 0 {
|
|
||||||
wantNoSeries(t, metrics, series)
|
|
||||||
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
wantMetric(t, metrics, series, float64(hits))
|
|
||||||
}
|
|
||||||
@@ -1,203 +0,0 @@
|
|||||||
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),
|
|
||||||
Cause: bans.CauseAdmin,
|
|
||||||
}
|
|
||||||
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)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@@ -5,8 +5,8 @@ import (
|
|||||||
"bytes"
|
"bytes"
|
||||||
"errors"
|
"errors"
|
||||||
"io"
|
"io"
|
||||||
|
"net"
|
||||||
"net/http"
|
"net/http"
|
||||||
"reflect"
|
|
||||||
"slices"
|
"slices"
|
||||||
"strings"
|
"strings"
|
||||||
"sync/atomic"
|
"sync/atomic"
|
||||||
@@ -15,7 +15,6 @@ import (
|
|||||||
|
|
||||||
"sneak.berlin/go/smallwebwaf/internal/config"
|
"sneak.berlin/go/smallwebwaf/internal/config"
|
||||||
"sneak.berlin/go/smallwebwaf/internal/proxy"
|
"sneak.berlin/go/smallwebwaf/internal/proxy"
|
||||||
"sneak.berlin/go/smallwebwaf/internal/ratelimit"
|
|
||||||
"sneak.berlin/go/smallwebwaf/internal/requestlog"
|
"sneak.berlin/go/smallwebwaf/internal/requestlog"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -117,32 +116,27 @@ func wantAnswer(t *testing.T, got answer, body []byte) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// wantRequestFields checks the log line's fields about the request. Its
|
// wantRequestFields checks the log line's fields about the request.
|
||||||
// time, its id and its timings are checked only for being there.
|
|
||||||
func wantRequestFields(t *testing.T, line logLine, host string, sent, received int) {
|
func wantRequestFields(t *testing.T, line logLine, host string, sent, received int) {
|
||||||
t.Helper()
|
t.Helper()
|
||||||
|
|
||||||
want := withTimings(line, requestlog.Line{
|
want := requestlog.Line{
|
||||||
Type: requestType, Time: line.Time, Instance: "app",
|
Type: "request", Time: line.Time, ClientIP: localhost, PeerIP: localhost,
|
||||||
ClientIP: localhost, Method: http.MethodPatch, Scheme: plain, Host: host,
|
Method: http.MethodPatch, Host: host, Path: rawPath, Query: rawQuery,
|
||||||
Path: rawPath, Query: rawQuery, Protocol: protocol,
|
Protocol: "HTTP/1.1", Status: http.StatusTeapot,
|
||||||
Status: http.StatusTeapot, RequestBytes: int64(sent),
|
UpstreamStatus: http.StatusTeapot, RequestBytes: int64(sent),
|
||||||
ResponseBytes: int64(received), UserAgent: "test-agent",
|
ResponseBytes: int64(received), UserAgent: "test-agent",
|
||||||
RequestID: line.RequestID, PeerIP: localhost, ClientGroup: localhost + "/32",
|
Action: requestlog.ActionForward, DurationTotal: line.DurationTotal,
|
||||||
ContentLength: int64(sent), ResponseContentType: "text/plain; charset=utf-8",
|
DurationUpstreamTotal: line.DurationUpstreamTotal,
|
||||||
UpstreamStatus: http.StatusTeapot, Action: requestlog.ActionForward,
|
}
|
||||||
Counts: ratelimit.Counts{Minute: 1, Hour: 1, Day: 1},
|
if line.Line != want {
|
||||||
})
|
|
||||||
if !reflect.DeepEqual(line.Line, want) {
|
|
||||||
t.Errorf("log line\n%+v\nwant\n%+v", line.Line, want)
|
t.Errorf("log line\n%+v\nwant\n%+v", line.Line, want)
|
||||||
}
|
}
|
||||||
|
|
||||||
_, err := time.Parse(time.RFC3339, line.Time)
|
_, err := time.Parse(time.RFC3339, line.Time)
|
||||||
if err != nil || line.RequestID == "" || line.DurationTotal <= 0 ||
|
if err != nil || line.DurationTotal <= 0 || line.DurationUpstreamTotal <= 0 {
|
||||||
line.DurationUpstreamTotal == nil || *line.DurationUpstreamTotal <= 0 {
|
t.Errorf("log line has time %q and durations %v and %v",
|
||||||
t.Errorf("log line has time %q, request_id %q and durations %v and %v",
|
line.Time, line.DurationTotal, line.DurationUpstreamTotal)
|
||||||
line.Time, line.RequestID, line.DurationTotal,
|
|
||||||
line.fields["duration_upstream_total"])
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -247,17 +241,14 @@ func TestUpgradedConnectionOutlastsTheTimeouts(t *testing.T) {
|
|||||||
t.Fatalf("read the answer to the upgrade: %v", err)
|
t.Fatalf("read the answer to the upgrade: %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
answered := time.Now()
|
|
||||||
_ = res.Body.Close()
|
_ = res.Body.Close()
|
||||||
|
|
||||||
if res.StatusCode != http.StatusSwitchingProtocols {
|
if res.StatusCode != http.StatusSwitchingProtocols {
|
||||||
t.Fatalf("status %d, want %d", res.StatusCode, http.StatusSwitchingProtocols)
|
t.Fatalf("status %d, want %d", res.StatusCode, http.StatusSwitchingProtocols)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Every timeout started before the upgrade was answered, the response
|
// Wait past every timeout, then use the connection.
|
||||||
// timeouts last, at the end of the request: wait until just past
|
time.Sleep(3 * shortTimeout)
|
||||||
// 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")
|
send(t, conn, "still here\n")
|
||||||
|
|
||||||
echoed, err := reader.ReadString('\n')
|
echoed, err := reader.ReadString('\n')
|
||||||
@@ -304,7 +295,7 @@ func echoAfterUpgrade(w http.ResponseWriter, r *http.Request) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestServerHasTheDefaultLimits(t *testing.T) {
|
func TestServerHasTheFixedLimits(t *testing.T) {
|
||||||
t.Parallel()
|
t.Parallel()
|
||||||
|
|
||||||
cfg, err := config.FromEnvironment(func(string) (string, bool) { return "", false })
|
cfg, err := config.FromEnvironment(func(string) (string, bool) { return "", false })
|
||||||
@@ -315,7 +306,7 @@ func TestServerHasTheDefaultLimits(t *testing.T) {
|
|||||||
server := proxy.New(proxy.Params{
|
server := proxy.New(proxy.Params{
|
||||||
Config: cfg,
|
Config: cfg,
|
||||||
RequestLog: io.Discard,
|
RequestLog: io.Discard,
|
||||||
ProcessLog: requestlog.NewProcessLogger(io.Discard, cfg.InstanceName),
|
ProcessLog: requestlog.NewProcessLogger(io.Discard),
|
||||||
})
|
})
|
||||||
|
|
||||||
if server.Addr != ":8080" || server.MaxHeaderBytes != 28<<10 ||
|
if server.Addr != ":8080" || server.MaxHeaderBytes != 28<<10 ||
|
||||||
@@ -326,65 +317,56 @@ func TestServerHasTheDefaultLimits(t *testing.T) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestRefusesHeadersOverTheLimit(t *testing.T) {
|
func TestRefusesHeadersOver32KiB(t *testing.T) {
|
||||||
t.Parallel()
|
t.Parallel()
|
||||||
|
|
||||||
|
var calls atomic.Int32
|
||||||
|
|
||||||
|
app := startApp(t, func(http.ResponseWriter, *http.Request) {
|
||||||
|
calls.Add(1)
|
||||||
|
})
|
||||||
|
addr, _ := startProxy(t, app.URL, nil)
|
||||||
|
|
||||||
|
// 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 _, tc := range []struct {
|
for _, tc := range []struct {
|
||||||
name string
|
size int
|
||||||
env map[string]string
|
want int
|
||||||
limit int
|
|
||||||
}{
|
}{
|
||||||
{"by default", nil, 32 << 10},
|
{size: 32 << 10, want: http.StatusOK},
|
||||||
{"as set", map[string]string{clientHeaderMaxBytes: "8K"}, 8 << 10},
|
{size: 32<<10 + 1, want: http.StatusRequestHeaderFieldsTooLarge},
|
||||||
} {
|
} {
|
||||||
t.Run(tc.name, func(t *testing.T) {
|
conn := dial(t, addr)
|
||||||
t.Parallel()
|
send(t, conn, start+strings.Repeat("a", tc.size-len(start)-len(end))+end)
|
||||||
|
wantStatus(t, readResponse(t, conn), tc.want)
|
||||||
|
}
|
||||||
|
|
||||||
var calls atomic.Int32
|
if calls.Load() != 1 {
|
||||||
|
t.Errorf("the app was called %d times, want once", calls.Load())
|
||||||
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) {
|
func TestAnswers502WhenTheAppCannotBeReached(t *testing.T) {
|
||||||
t.Parallel()
|
t.Parallel()
|
||||||
|
|
||||||
// No test can listen on port 1: listening on port 0 gets one from 32768 up.
|
listener, err := (&net.ListenConfig{}).Listen(t.Context(), "tcp", localhost+":0")
|
||||||
addr, out := startProxy(t, "http://"+localhost+":1", nil)
|
if err != nil {
|
||||||
|
t.Fatalf("listen: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
closedAddr := listener.Addr().String()
|
||||||
|
_ = listener.Close()
|
||||||
|
|
||||||
|
addr, out := startProxy(t, "http://"+closedAddr, nil)
|
||||||
|
|
||||||
wantStatus(t, get(t, addr, "/"), http.StatusBadGateway)
|
wantStatus(t, get(t, addr, "/"), http.StatusBadGateway)
|
||||||
|
wantLine(t, out.requestLine(t), http.StatusBadGateway,
|
||||||
line := out.requestLine(t)
|
requestlog.ActionUpstreamError)
|
||||||
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 {
|
logged := slices.ContainsFunc(out.lines(t), func(line map[string]any) bool {
|
||||||
return line["type"] == "process" && line["msg"] == "request to the app failed"
|
return line["type"] == "process" && line["msg"] == "request to the app failed"
|
||||||
|
|||||||
+35
-153
@@ -8,17 +8,23 @@ import (
|
|||||||
"log"
|
"log"
|
||||||
"log/slog"
|
"log/slog"
|
||||||
"net/http"
|
"net/http"
|
||||||
"strings"
|
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"sneak.berlin/go/smallwebwaf/internal/alerts"
|
|
||||||
"sneak.berlin/go/smallwebwaf/internal/bans"
|
|
||||||
"sneak.berlin/go/smallwebwaf/internal/config"
|
"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/ratelimit"
|
||||||
"sneak.berlin/go/smallwebwaf/internal/requestlog"
|
)
|
||||||
"sneak.berlin/go/smallwebwaf/internal/rules"
|
|
||||||
|
// The request line and headers a client may send, and how long a
|
||||||
|
// kept-open client connection may wait for its next request, are fixed
|
||||||
|
// rather than settings. The limit on the request line and headers is
|
||||||
|
// 32 KiB, but Go's server reads 4 KiB past its MaxHeaderBytes before it
|
||||||
|
// refuses, so MaxHeaderBytes is set 4 KiB lower. The idle time 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.
|
||||||
|
const (
|
||||||
|
requestHeaderMaxBytes = 32<<10 - 4<<10
|
||||||
|
clientIdleTimeout = 120 * time.Second
|
||||||
)
|
)
|
||||||
|
|
||||||
// How smallwebwaf keeps connections to the app open between requests.
|
// How smallwebwaf keeps connections to the app open between requests.
|
||||||
@@ -27,26 +33,6 @@ const (
|
|||||||
appIdleConnTimeout = 90 * time.Second
|
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"
|
|
||||||
|
|
||||||
// BansPath is where an admin lists and adds bans, and, followed by / and
|
|
||||||
// a client's address, lifts them, with SWWAF_ADMIN_TOKEN.
|
|
||||||
const BansPath = "/_smallwebwaf/bans"
|
|
||||||
|
|
||||||
// ClientsPath is where an admin asks what smallwebwaf knows of a client,
|
|
||||||
// by the client's address after it, with SWWAF_ADMIN_TOKEN.
|
|
||||||
const ClientsPath = "/_smallwebwaf/clients/"
|
|
||||||
|
|
||||||
// Params are what New needs.
|
// Params are what New needs.
|
||||||
type Params struct {
|
type Params struct {
|
||||||
Config *config.Config
|
Config *config.Config
|
||||||
@@ -54,105 +40,34 @@ type Params struct {
|
|||||||
RequestLog io.Writer
|
RequestLog io.Writer
|
||||||
// ProcessLog receives the process's own messages.
|
// ProcessLog receives the process's own messages.
|
||||||
ProcessLog *slog.Logger
|
ProcessLog *slog.Logger
|
||||||
// GeoJSURL is where clients' AS numbers and countries are looked up
|
|
||||||
// while SWWAF_LOOKUP_SOURCE is geojs, normally lookup.URL.
|
|
||||||
GeoJSURL string
|
|
||||||
// LookupFile is the lookup database they are looked up in while
|
|
||||||
// SWWAF_LOOKUP_SOURCE is file, and nil otherwise.
|
|
||||||
LookupFile *lookup.File
|
|
||||||
// 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
|
|
||||||
// Rules are the rule files' rules, which each request is checked
|
|
||||||
// against.
|
|
||||||
Rules *rules.Files
|
|
||||||
// Alerts receive the alert for each ban the proxy makes or makes
|
|
||||||
// permanent, and for GeoJS failing.
|
|
||||||
Alerts *alerts.Queue
|
|
||||||
}
|
|
||||||
|
|
||||||
// Server is the server smallwebwaf runs, with the parts of the proxy
|
|
||||||
// whose state the state files keep, the lookup database, nil unless
|
|
||||||
// SWWAF_LOOKUP_SOURCE is file, and the metrics.
|
|
||||||
type Server struct {
|
|
||||||
*http.Server
|
|
||||||
|
|
||||||
Ledger *bans.Ledger
|
|
||||||
Limiter *ratelimit.Limiter
|
|
||||||
GeoJS *lookup.GeoJS
|
|
||||||
LookupFile *lookup.File
|
|
||||||
Metrics *metrics.Metrics
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// New returns the server smallwebwaf runs: each request it reads passes
|
// New returns the server smallwebwaf runs: each request it reads passes
|
||||||
// through the proxy. Go's server itself refuses a request line and
|
// through the proxy. Go's server itself refuses headers over 32 KiB, with
|
||||||
// headers over SWWAF_CLIENT_REQUEST_HEADER_MAX_BYTES, with 431, closes a
|
// 431, closes a connection idle for 120 seconds, and applies
|
||||||
// connection idle for SWWAF_CLIENT_IDLE_TIMEOUT, and applies
|
|
||||||
// SWWAF_CLIENT_REQUEST_TIMEOUT while the headers arrive; the proxy
|
// SWWAF_CLIENT_REQUEST_TIMEOUT while the headers arrive; the proxy
|
||||||
// applies the timeouts and size limits from then on.
|
// applies the timeouts and size limits from then on.
|
||||||
func New(params Params) *Server {
|
func New(params Params) *http.Server {
|
||||||
errorLog := slog.NewLogLogger(params.ProcessLog.Handler(), slog.LevelWarn)
|
errorLog := slog.NewLogLogger(params.ProcessLog.Handler(), slog.LevelWarn)
|
||||||
m := metrics.New(params.Config.MetricsTopN, params.Config.InstanceName)
|
|
||||||
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,
|
|
||||||
AttackBanDuration: params.Config.AttackBanDuration,
|
|
||||||
MaxBans: params.Config.MaxBans,
|
|
||||||
}),
|
|
||||||
lookupFile: params.LookupFile,
|
|
||||||
rules: params.Rules,
|
|
||||||
alerts: params.Alerts,
|
|
||||||
}
|
|
||||||
h.geojs = lookup.New(lookup.Params{
|
|
||||||
URL: params.GeoJSURL,
|
|
||||||
Timeout: params.Config.LookupTimeout,
|
|
||||||
// The country lists and the headers act on the answer before the
|
|
||||||
// request goes on.
|
|
||||||
Wait: len(params.Config.DeniedCountries) > 0 ||
|
|
||||||
len(params.Config.ExclusivelyAllowedCountries) > 0 ||
|
|
||||||
params.Config.AddLookupHeaders,
|
|
||||||
Answered: h.addLookup,
|
|
||||||
Now: params.Now,
|
|
||||||
ProcessLog: params.ProcessLog,
|
|
||||||
Metrics: m,
|
|
||||||
Alerts: params.Alerts,
|
|
||||||
})
|
|
||||||
m.AddBansAndClients(h.ledger, h.limiter, params.Now)
|
|
||||||
m.AddRules(params.Rules)
|
|
||||||
|
|
||||||
return &Server{
|
return &http.Server{
|
||||||
Server: &http.Server{
|
Addr: params.Config.ListenAddr,
|
||||||
Addr: params.Config.ListenAddr,
|
Handler: &handler{
|
||||||
Handler: h,
|
config: params.Config,
|
||||||
ReadHeaderTimeout: params.Config.ClientRequestTimeout,
|
requestLog: params.RequestLog,
|
||||||
// Off is an IdleTimeout of 0, which Go's server replaces with
|
processLog: params.ProcessLog,
|
||||||
// ReadTimeout: no limit, as long as ReadTimeout stays unset.
|
errorLog: errorLog,
|
||||||
IdleTimeout: params.Config.ClientIdleTimeout,
|
transport: newTransport(),
|
||||||
// Go's server reads 4 KiB past MaxHeaderBytes before it
|
limiter: ratelimit.New(ratelimit.Limits{
|
||||||
// refuses, so the limit a client meets is the setting.
|
PerMinute: params.Config.RateLimitPerMinute,
|
||||||
MaxHeaderBytes: int(params.Config.ClientRequestHeaderMaxBytes - 4<<10),
|
PerHour: params.Config.RateLimitPerHour,
|
||||||
ErrorLog: errorLog,
|
PerDay: params.Config.RateLimitPerDay,
|
||||||
|
}),
|
||||||
},
|
},
|
||||||
Ledger: h.ledger,
|
ReadHeaderTimeout: params.Config.ClientRequestTimeout,
|
||||||
Limiter: h.limiter,
|
IdleTimeout: clientIdleTimeout,
|
||||||
GeoJS: h.geojs,
|
MaxHeaderBytes: requestHeaderMaxBytes,
|
||||||
LookupFile: h.lookupFile,
|
ErrorLog: errorLog,
|
||||||
Metrics: m,
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -164,14 +79,7 @@ type handler struct {
|
|||||||
processLog *slog.Logger
|
processLog *slog.Logger
|
||||||
errorLog *log.Logger
|
errorLog *log.Logger
|
||||||
transport http.RoundTripper
|
transport http.RoundTripper
|
||||||
now func() time.Time
|
|
||||||
metrics *metrics.Metrics
|
|
||||||
limiter *ratelimit.Limiter
|
limiter *ratelimit.Limiter
|
||||||
ledger *bans.Ledger
|
|
||||||
geojs *lookup.GeoJS
|
|
||||||
lookupFile *lookup.File
|
|
||||||
rules *rules.Files
|
|
||||||
alerts *alerts.Queue
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// newTransport returns what carries requests to the app. It never goes
|
// newTransport returns what carries requests to the app. It never goes
|
||||||
@@ -188,43 +96,17 @@ func newTransport() *http.Transport {
|
|||||||
|
|
||||||
// ServeHTTP handles one request: it works out the client, runs the
|
// ServeHTTP handles one request: it works out the client, runs the
|
||||||
// checks, passes the request to the app and the answer back within 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
|
// limits, and writes the request's log line.
|
||||||
// request's log line.
|
|
||||||
func (h *handler) ServeHTTP(w http.ResponseWriter, r *http.Request) {
|
func (h *handler) ServeHTTP(w http.ResponseWriter, r *http.Request) {
|
||||||
rq := h.newRequest(w, r)
|
rq := h.newRequest(w, r)
|
||||||
defer rq.finish()
|
defer rq.finish()
|
||||||
|
|
||||||
// The health endpoint is answered at once, before any check, so that
|
refused := rq.check()
|
||||||
// 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 {
|
if refused != nil {
|
||||||
rq.answer(*refused)
|
rq.answer(*refused)
|
||||||
|
|
||||||
return
|
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())
|
rq.forward(r.Context())
|
||||||
}
|
}
|
||||||
|
|||||||
+21
-165
@@ -14,23 +14,18 @@ import (
|
|||||||
"testing"
|
"testing"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"sneak.berlin/go/smallwebwaf/internal/alerts"
|
|
||||||
"sneak.berlin/go/smallwebwaf/internal/config"
|
"sneak.berlin/go/smallwebwaf/internal/config"
|
||||||
"sneak.berlin/go/smallwebwaf/internal/lookup"
|
|
||||||
"sneak.berlin/go/smallwebwaf/internal/proxy"
|
"sneak.berlin/go/smallwebwaf/internal/proxy"
|
||||||
"sneak.berlin/go/smallwebwaf/internal/requestlog"
|
"sneak.berlin/go/smallwebwaf/internal/requestlog"
|
||||||
"sneak.berlin/go/smallwebwaf/internal/rules"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
const (
|
const (
|
||||||
// shortTimeout is what a test sets a timeout to, to see it run out.
|
// 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
|
shortTimeout = 300 * time.Millisecond
|
||||||
// or the app's buffers filling, so it is as long as the hold-up of the
|
// shortTimeoutSetting is shortTimeout as a setting's value.
|
||||||
// test process that wantTimedOut allows: a shorter one can run out
|
shortTimeoutSetting = "300ms"
|
||||||
// first on a busy host.
|
|
||||||
shortTimeout = waitLimit / 2
|
|
||||||
// longTimeoutSetting is a timeout that does not run out in a test.
|
// longTimeoutSetting is a timeout that does not run out in a test.
|
||||||
longTimeoutSetting = "1m"
|
longTimeoutSetting = "10s"
|
||||||
// waitLimit bounds how long a test waits for what should happen.
|
// waitLimit bounds how long a test waits for what should happen.
|
||||||
waitLimit = 10 * time.Second
|
waitLimit = 10 * time.Second
|
||||||
// pollInterval is how often a test looks for a log line.
|
// pollInterval is how often a test looks for a log line.
|
||||||
@@ -38,51 +33,18 @@ const (
|
|||||||
// localhost is where every test server listens, and so the address
|
// localhost is where every test server listens, and so the address
|
||||||
// smallwebwaf sees each test's requests come from.
|
// smallwebwaf sees each test's requests come from.
|
||||||
localhost = "127.0.0.1"
|
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.
|
// The settings the tests set.
|
||||||
const (
|
const (
|
||||||
clientRequestTimeout = "SWWAF_CLIENT_REQUEST_TIMEOUT"
|
clientRequestTimeout = "SWWAF_CLIENT_REQUEST_TIMEOUT"
|
||||||
clientHeaderMaxBytes = "SWWAF_CLIENT_REQUEST_HEADER_MAX_BYTES"
|
|
||||||
clientIdleTimeout = "SWWAF_CLIENT_IDLE_TIMEOUT"
|
|
||||||
clientResponseTimeout = "SWWAF_CLIENT_RESPONSE_TIMEOUT"
|
clientResponseTimeout = "SWWAF_CLIENT_RESPONSE_TIMEOUT"
|
||||||
upstreamRequestTimeout = "SWWAF_UPSTREAM_REQUEST_TIMEOUT"
|
upstreamRequestTimeout = "SWWAF_UPSTREAM_REQUEST_TIMEOUT"
|
||||||
upstreamResponseTimeout = "SWWAF_UPSTREAM_RESPONSE_TIMEOUT"
|
upstreamResponseTimeout = "SWWAF_UPSTREAM_RESPONSE_TIMEOUT"
|
||||||
mode = "SWWAF_MODE"
|
|
||||||
requestMaxBytes = "SWWAF_REQUEST_MAX_BYTES"
|
requestMaxBytes = "SWWAF_REQUEST_MAX_BYTES"
|
||||||
responseMaxBytes = "SWWAF_RESPONSE_MAX_BYTES"
|
responseMaxBytes = "SWWAF_RESPONSE_MAX_BYTES"
|
||||||
trustedProxies = "SWWAF_TRUSTED_PROXIES"
|
trustedProxies = "SWWAF_TRUSTED_PROXIES"
|
||||||
allowNets = "SWWAF_ALLOW_NETS"
|
|
||||||
rateLimitExemptNets = "SWWAF_RATE_LIMIT_EXEMPT_NETS"
|
|
||||||
denyNets = "SWWAF_DENY_NETS"
|
|
||||||
rateLimitPerMinute = "SWWAF_RATE_LIMIT_PER_MINUTE"
|
rateLimitPerMinute = "SWWAF_RATE_LIMIT_PER_MINUTE"
|
||||||
rateLimitPerDay = "SWWAF_RATE_LIMIT_PER_DAY"
|
|
||||||
rateLimitExemptPaths = "SWWAF_RATE_LIMIT_EXEMPT_PATHS"
|
|
||||||
lookupSource = "SWWAF_LOOKUP_SOURCE"
|
|
||||||
lookupDBPath = "SWWAF_LOOKUP_DB_PATH"
|
|
||||||
lookupTimeout = "SWWAF_LOOKUP_TIMEOUT"
|
|
||||||
addLookupHeaders = "SWWAF_ADD_LOOKUP_HEADERS"
|
|
||||||
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"
|
|
||||||
attackBanDuration = "SWWAF_ATTACK_BAN_DURATION"
|
|
||||||
rulesDir = "SWWAF_RULES_DIR"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
// output collects what smallwebwaf writes on stdout.
|
// output collects what smallwebwaf writes on stdout.
|
||||||
@@ -99,14 +61,6 @@ func (o *output) Write(p []byte) (int, error) {
|
|||||||
return o.buf.Write(p)
|
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.
|
// lines returns every line written so far, decoded.
|
||||||
func (o *output) lines(t *testing.T) []map[string]any {
|
func (o *output) lines(t *testing.T) []map[string]any {
|
||||||
t.Helper()
|
t.Helper()
|
||||||
@@ -146,7 +100,7 @@ func (o *output) requestLines(t *testing.T, count int) []logLine {
|
|||||||
var found []logLine
|
var found []logLine
|
||||||
|
|
||||||
for _, fields := range o.lines(t) {
|
for _, fields := range o.lines(t) {
|
||||||
if fields["type"] == requestType {
|
if fields["type"] == "request" {
|
||||||
found = append(found, decodeLine(t, fields))
|
found = append(found, decodeLine(t, fields))
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -205,46 +159,24 @@ func startApp(t *testing.T, app http.HandlerFunc) *httptest.Server {
|
|||||||
func startProxy(t *testing.T, appURL string, env map[string]string) (string, *output) {
|
func startProxy(t *testing.T, appURL string, env map[string]string) (string, *output) {
|
||||||
t.Helper()
|
t.Helper()
|
||||||
|
|
||||||
return startProxyWithGeoJS(t, appURL, "", env)
|
settings := map[string]string{"SWWAF_UPSTREAM_URL": appURL}
|
||||||
}
|
maps.Copy(settings, env)
|
||||||
|
|
||||||
// startProxyWithGeoJS is startProxy with clients' AS numbers and
|
cfg, err := config.FromEnvironment(func(name string) (string, bool) {
|
||||||
// countries looked up at geojsURL.
|
value, ok := settings[name]
|
||||||
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 value, ok
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("settings %v: %v", settings, err)
|
||||||
|
}
|
||||||
|
|
||||||
return addr, out
|
out := &output{}
|
||||||
}
|
server := proxy.New(proxy.Params{
|
||||||
|
Config: cfg,
|
||||||
// startProxyWithClock is startProxyWithGeoJS with requests counted and
|
RequestLog: out,
|
||||||
// bans made by the time now tells, and returns the server as well. Unless
|
ProcessLog: requestlog.NewProcessLogger(out),
|
||||||
// env sets SWWAF_RULES_DIR, it is an empty directory, of no rules, and
|
})
|
||||||
// unless it sets SWWAF_INSTANCE_NAME, that is app, the label instance of
|
|
||||||
// every metric.
|
|
||||||
func startProxyWithClock(
|
|
||||||
t *testing.T, appURL, geojsURL string, now func() time.Time,
|
|
||||||
env map[string]string,
|
|
||||||
) (string, *output, *proxy.Server) {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
addr, out, server, _ := startProxyWithAlerts(t, appURL, geojsURL, now, env)
|
|
||||||
|
|
||||||
return addr, out, server
|
|
||||||
}
|
|
||||||
|
|
||||||
// startProxyWithAlerts is startProxyWithClock, and returns the queue of
|
|
||||||
// the alerts the proxy raises as well, as newProxy makes them.
|
|
||||||
func startProxyWithAlerts(
|
|
||||||
t *testing.T, appURL, geojsURL string, now func() time.Time,
|
|
||||||
env map[string]string,
|
|
||||||
) (string, *output, *proxy.Server, *alerts.Queue) {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
server, out, alertQueue := newProxy(t, appURL, geojsURL, now, env)
|
|
||||||
|
|
||||||
listener, err := (&net.ListenConfig{}).Listen(t.Context(), "tcp", localhost+":0")
|
listener, err := (&net.ListenConfig{}).Listen(t.Context(), "tcp", localhost+":0")
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -259,83 +191,7 @@ func startProxyWithAlerts(
|
|||||||
_ = server.Close()
|
_ = server.Close()
|
||||||
})
|
})
|
||||||
|
|
||||||
return listener.Addr().String(), out, server, alertQueue
|
return listener.Addr().String(), out
|
||||||
}
|
|
||||||
|
|
||||||
// newProxy makes the server startProxyWithClock starts, without starting
|
|
||||||
// it, and returns it, what it writes, and the queue of the alerts the
|
|
||||||
// proxy raises, as the settings in env make it. No alert is sent from the
|
|
||||||
// queue: they wait in it, for the test to look at. With no geojsURL, there
|
|
||||||
// is no stand-in for GeoJS to look clients up at, and SWWAF_LOOKUP_SOURCE
|
|
||||||
// is off unless env sets it. While it is file, the lookup database
|
|
||||||
// SWWAF_LOOKUP_DB_PATH names is read.
|
|
||||||
func newProxy(
|
|
||||||
t *testing.T, appURL, geojsURL string, now func() time.Time,
|
|
||||||
env map[string]string,
|
|
||||||
) (*proxy.Server, *output, *alerts.Queue) {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
settings := map[string]string{
|
|
||||||
"SWWAF_UPSTREAM_URL": appURL, rulesDir: t.TempDir(), instanceName: "app",
|
|
||||||
}
|
|
||||||
if geojsURL == "" {
|
|
||||||
settings[lookupSource] = "off"
|
|
||||||
}
|
|
||||||
|
|
||||||
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{}
|
|
||||||
processLog := requestlog.NewProcessLogger(out, cfg.InstanceName)
|
|
||||||
|
|
||||||
ruleFiles, err := rules.Load(rules.Params{
|
|
||||||
Dir: cfg.RulesDir, Enabled: cfg.RulesEnabled, ProcessLog: processLog,
|
|
||||||
})
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("rule files: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
alertQueue := alerts.New(alerts.Params{
|
|
||||||
WebhookURL: cfg.AlertWebhookURL,
|
|
||||||
Events: cfg.AlertEvents,
|
|
||||||
Cooldown: cfg.AlertCooldown,
|
|
||||||
MaxPerHour: cfg.AlertMaxPerHour,
|
|
||||||
Instance: cfg.InstanceName,
|
|
||||||
Now: now,
|
|
||||||
ProcessLog: processLog,
|
|
||||||
})
|
|
||||||
|
|
||||||
var lookupFile *lookup.File
|
|
||||||
|
|
||||||
if cfg.LookupSource == "file" {
|
|
||||||
lookupFile, err = lookup.OpenFile(lookup.FileParams{
|
|
||||||
Path: cfg.LookupDBPath, Now: now, ProcessLog: processLog, Alerts: alertQueue,
|
|
||||||
})
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("lookup database: %v", err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
server := proxy.New(proxy.Params{
|
|
||||||
Config: cfg,
|
|
||||||
RequestLog: out,
|
|
||||||
ProcessLog: processLog,
|
|
||||||
GeoJSURL: geojsURL,
|
|
||||||
LookupFile: lookupFile,
|
|
||||||
Now: now,
|
|
||||||
Rules: ruleFiles,
|
|
||||||
Alerts: alertQueue,
|
|
||||||
})
|
|
||||||
|
|
||||||
return server, out, alertQueue
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// newClient returns an HTTP client that sends requests as they are made,
|
// newClient returns an HTTP client that sends requests as they are made,
|
||||||
|
|||||||
@@ -5,16 +5,10 @@ import (
|
|||||||
"sync/atomic"
|
"sync/atomic"
|
||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
"sneak.berlin/go/smallwebwaf/internal/lookup"
|
|
||||||
"sneak.berlin/go/smallwebwaf/internal/ratelimit"
|
|
||||||
"sneak.berlin/go/smallwebwaf/internal/requestlog"
|
"sneak.berlin/go/smallwebwaf/internal/requestlog"
|
||||||
)
|
)
|
||||||
|
|
||||||
// minute is the window of SWWAF_RATE_LIMIT_PER_MINUTE, as a log line's
|
func TestRateLimitRefusesWith429BeforeTheApp(t *testing.T) {
|
||||||
// limit_hit names it.
|
|
||||||
const minute = "minute"
|
|
||||||
|
|
||||||
func TestRateLimitRefusesBeforeTheApp(t *testing.T) {
|
|
||||||
t.Parallel()
|
t.Parallel()
|
||||||
|
|
||||||
var calls atomic.Int32
|
var calls atomic.Int32
|
||||||
@@ -30,19 +24,19 @@ func TestRateLimitRefusesBeforeTheApp(t *testing.T) {
|
|||||||
const otherClient = "203.0.113.10"
|
const otherClient = "203.0.113.10"
|
||||||
|
|
||||||
// With a limit of one request a minute, a client's second request is
|
// 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
|
// refused. A client is one IPv4 address, or one IPv6 /64; an IPv4
|
||||||
// IPv6 /64; an IPv4 address in IPv6 form is that IPv4 address.
|
// address in IPv6 form is that IPv4 address.
|
||||||
requests := []struct {
|
requests := []struct {
|
||||||
client string // as X-Forwarded-For names it
|
client string // as X-Forwarded-For names it
|
||||||
logged string // as the log line's client_ip names it
|
logged string // as the log line's client_ip names it
|
||||||
want int
|
want int
|
||||||
}{
|
}{
|
||||||
{client, client, http.StatusOK},
|
{client, client, http.StatusOK},
|
||||||
{client, client, http.StatusForbidden},
|
{client, client, http.StatusTooManyRequests},
|
||||||
{otherClient, otherClient, http.StatusOK},
|
{otherClient, otherClient, http.StatusOK},
|
||||||
{"::ffff:" + otherClient, otherClient, http.StatusForbidden},
|
{"::ffff:" + otherClient, otherClient, http.StatusTooManyRequests},
|
||||||
{"2001:db8::1", "2001:db8::1", http.StatusOK},
|
{"2001:db8::1", "2001:db8::1", http.StatusOK},
|
||||||
{"2001:db8::8000:0:0:1", "2001:db8::8000:0:0:1", http.StatusForbidden},
|
{"2001:db8::8000:0:0:1", "2001:db8::8000:0:0:1", http.StatusTooManyRequests},
|
||||||
{"2001:db8:0:1::1", "2001:db8:0:1::1", http.StatusOK},
|
{"2001:db8:0:1::1", "2001:db8:0:1::1", http.StatusOK},
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -59,9 +53,9 @@ func TestRateLimitRefusesBeforeTheApp(t *testing.T) {
|
|||||||
if sent.want == http.StatusOK {
|
if sent.want == http.StatusOK {
|
||||||
wantLine(t, line, http.StatusOK, requestlog.ActionForward)
|
wantLine(t, line, http.StatusOK, requestlog.ActionForward)
|
||||||
} else {
|
} else {
|
||||||
wantLine(t, line, http.StatusForbidden, requestlog.ActionRateLimited)
|
wantLine(t, line, http.StatusTooManyRequests, requestlog.ActionRateLimited)
|
||||||
|
|
||||||
if line.LimitHit != minute {
|
if line.LimitHit != "minute" {
|
||||||
t.Errorf("log line has limit_hit %q, want minute", line.LimitHit)
|
t.Errorf("log line has limit_hit %q, want minute", line.LimitHit)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -71,90 +65,3 @@ func TestRateLimitRefusesBeforeTheApp(t *testing.T) {
|
|||||||
t.Errorf("the app was called %d times, want 4", calls.Load())
|
t.Errorf("the app was called %d times, want 4", calls.Load())
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestRateLimitExemptPathsAreNeitherCountedNorRefused(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
const denied = "192.0.2.50" // in SWWAF_DENY_NETS
|
|
||||||
|
|
||||||
geojsURL, _ := startGeoJS(t)
|
|
||||||
s, _, server := startWithClock(t, geojsURL, map[string]string{
|
|
||||||
rateLimitPerMinute: "1",
|
|
||||||
rateLimitExemptPaths: "/assets/,/favicon.ico",
|
|
||||||
denyNets: denied,
|
|
||||||
deniedCountries: "kp",
|
|
||||||
})
|
|
||||||
|
|
||||||
// The answers are kept before the requests, so that none waits for
|
|
||||||
// GeoJS.
|
|
||||||
server.GeoJS.Load([]lookup.Answer{
|
|
||||||
keptAnswer(client, "DE"), keptAnswer(fromKP, "KP"),
|
|
||||||
})
|
|
||||||
|
|
||||||
// With a limit of one request a minute, the requests for paths under a
|
|
||||||
// prefix are not counted, so client's first request for / is within
|
|
||||||
// the limit; and once client has reached it, they are not refused.
|
|
||||||
s.request(client, "/assets/app.js", http.StatusOK, requestlog.ActionForward)
|
|
||||||
s.request(client, "/favicon.ico?v=2", http.StatusOK, requestlog.ActionForward)
|
|
||||||
s.get(client, http.StatusOK, requestlog.ActionForward)
|
|
||||||
|
|
||||||
line := s.request(client, "/assets/app.js", http.StatusOK, requestlog.ActionForward)
|
|
||||||
if line.LimitHit != "" || line.Counts != (ratelimit.Counts{}) {
|
|
||||||
t.Errorf("log line has limit_hit %q and counts %+v, want neither",
|
|
||||||
line.LimitHit, line.Counts)
|
|
||||||
}
|
|
||||||
|
|
||||||
// A path outside every prefix is counted: /assets is not under
|
|
||||||
// /assets/, and breaks the limit.
|
|
||||||
s.request(client, "/assets", http.StatusForbidden, requestlog.ActionRateLimited)
|
|
||||||
|
|
||||||
// A ban, SWWAF_DENY_NETS and the country lists still refuse a path
|
|
||||||
// under a prefix.
|
|
||||||
s.request(client, "/assets/app.js", http.StatusForbidden, requestlog.ActionBanned)
|
|
||||||
s.request(denied, "/assets/app.js", http.StatusForbidden, requestlog.ActionDenied)
|
|
||||||
s.request(fromKP, "/assets/app.js",
|
|
||||||
http.StatusForbidden, requestlog.ActionCountryDenied)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestRateLimitCountsPathsThatAreNotExempt(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
for _, sent := range []string{
|
|
||||||
// A prefix matches only at the start of the path.
|
|
||||||
"/static/assets/app.js",
|
|
||||||
// A prefix matches the path as sent: a router that matches the
|
|
||||||
// path as received does not take /%61ssets/x for a path under
|
|
||||||
// /assets/.
|
|
||||||
"/%61ssets/x",
|
|
||||||
// .. once percent-decoded: an app may act on these as /login, the
|
|
||||||
// last as a path under /sneak/app/ or as /assets/x.
|
|
||||||
"/assets/../login",
|
|
||||||
"/assets/%2e%2e/login",
|
|
||||||
"/assets/..%2Flogin",
|
|
||||||
"/assets/..;/login",
|
|
||||||
"/sneak/app/src/branch/main/..%2F..%2F..%2F..%2F..%2F..%2Fassets/x",
|
|
||||||
// Not under /assets/ as sent: Go's router takes /assets%2Fx for one
|
|
||||||
// path segment, not a path under /assets/.
|
|
||||||
"/assets%2Fx",
|
|
||||||
"/assets%2fx",
|
|
||||||
// Under /assets/ as sent, but holding an encoded slash, in either
|
|
||||||
// case, or a backslash: never exempt, whatever the prefix.
|
|
||||||
"/assets/x%2Fy",
|
|
||||||
"/assets/x%2fy",
|
|
||||||
`/assets/x\y`,
|
|
||||||
} {
|
|
||||||
t.Run(sent, func(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
s, _, _ := startWithClock(t, "", map[string]string{
|
|
||||||
rateLimitPerMinute: "1",
|
|
||||||
rateLimitExemptPaths: "/assets/",
|
|
||||||
})
|
|
||||||
|
|
||||||
// Counted, the second request breaks the limit of one request
|
|
||||||
// a minute.
|
|
||||||
s.request(client, sent, http.StatusOK, requestlog.ActionForward)
|
|
||||||
s.request(client, sent, http.StatusForbidden, requestlog.ActionRateLimited)
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|||||||
+48
-313
@@ -7,16 +7,11 @@ import (
|
|||||||
"net/http/httptrace"
|
"net/http/httptrace"
|
||||||
"net/http/httputil"
|
"net/http/httputil"
|
||||||
"net/netip"
|
"net/netip"
|
||||||
"net/url"
|
|
||||||
"os"
|
"os"
|
||||||
"slices"
|
|
||||||
"strings"
|
|
||||||
"sync"
|
"sync"
|
||||||
"sync/atomic"
|
"sync/atomic"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"sneak.berlin/go/smallwebwaf/internal/lookup"
|
|
||||||
"sneak.berlin/go/smallwebwaf/internal/ratelimit"
|
|
||||||
"sneak.berlin/go/smallwebwaf/internal/requestlog"
|
"sneak.berlin/go/smallwebwaf/internal/requestlog"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -26,13 +21,10 @@ const flushAfterEachWrite time.Duration = -1
|
|||||||
|
|
||||||
// refusal is smallwebwaf refusing a request, or refusing to go on with it:
|
// 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,
|
// 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
|
// and the action the log line names.
|
||||||
// names, and the setting whose size or time limit the request passed, if
|
|
||||||
// that is why.
|
|
||||||
type refusal struct {
|
type refusal struct {
|
||||||
status int
|
status int
|
||||||
action string
|
action string
|
||||||
limit string
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// request is one request on its way through smallwebwaf, from the moment
|
// request is one request on its way through smallwebwaf, from the moment
|
||||||
@@ -49,15 +41,8 @@ type request struct {
|
|||||||
client netip.Addr
|
client netip.Addr
|
||||||
peer netip.Addr
|
peer netip.Addr
|
||||||
peerTrusted bool
|
peerTrusted bool
|
||||||
// lookedUp is true once the client's AS number and country have been
|
start time.Time
|
||||||
// looked up, whether or not an answer was there, and lookupAnswer is
|
// upstreamStart is when the request was handed to the app.
|
||||||
// what the lookup gave then, the zero Answer while GeoJS had given none.
|
|
||||||
lookedUp bool
|
|
||||||
lookupAnswer lookup.Answer
|
|
||||||
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
|
upstreamStart time.Time
|
||||||
// cancel ends the request to the app.
|
// cancel ends the request to the app.
|
||||||
cancel context.CancelFunc
|
cancel context.CancelFunc
|
||||||
@@ -67,34 +52,24 @@ type request struct {
|
|||||||
complete bool
|
complete bool
|
||||||
|
|
||||||
// mu guards what follows. The timeouts run on goroutines of their
|
// mu guards what follows. The timeouts run on goroutines of their
|
||||||
// own, and the transport starts and stops them, and notes the times
|
// own, and the transport starts and stops them from its own; once
|
||||||
// below, from its own; once timersStopped is set, none of the timeouts
|
// timersStopped is set, none of them acts any more.
|
||||||
// acts any more.
|
|
||||||
mu sync.Mutex
|
mu sync.Mutex
|
||||||
timersStopped bool
|
timersStopped bool
|
||||||
clientRequestTimer *time.Timer
|
clientRequestTimer *time.Timer
|
||||||
upstreamRequestTimer *time.Timer
|
upstreamRequestTimer *time.Timer
|
||||||
upstreamResponseTimer *time.Timer
|
upstreamResponseTimer *time.Timer
|
||||||
// connected is when there was a connection to the app, requestSent
|
// requestSent is when the app had been sent the whole request.
|
||||||
// when the app had been sent the whole request, and answerStarted
|
requestSent time.Time
|
||||||
// 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
|
// newRequest starts handling r: it notes the time and works out the
|
||||||
// under way, works out the client, and starts the log line with what is
|
// client.
|
||||||
// known of the request.
|
|
||||||
func (h *handler) newRequest(w http.ResponseWriter, r *http.Request) *request {
|
func (h *handler) newRequest(w http.ResponseWriter, r *http.Request) *request {
|
||||||
h.metrics.RequestStarted()
|
|
||||||
|
|
||||||
start := time.Now()
|
start := time.Now()
|
||||||
peer := peerAddress(r)
|
peer := peerAddress(r)
|
||||||
trusted := h.config.TrustedProxies
|
trusted := h.config.TrustedProxies
|
||||||
peerTrusted := isInside(peer, trusted)
|
client := clientAddress(peer, r.Header.Values("X-Forwarded-For"), trusted)
|
||||||
forwardedFor := r.Header.Values("X-Forwarded-For")
|
|
||||||
client := clientAddress(peer, forwardedFor, trusted)
|
|
||||||
|
|
||||||
rq := &request{
|
rq := &request{
|
||||||
h: h,
|
h: h,
|
||||||
@@ -103,37 +78,22 @@ func (h *handler) newRequest(w http.ResponseWriter, r *http.Request) *request {
|
|||||||
out: &responseWriter{ResponseWriter: w},
|
out: &responseWriter{ResponseWriter: w},
|
||||||
client: client,
|
client: client,
|
||||||
peer: peer,
|
peer: peer,
|
||||||
peerTrusted: peerTrusted,
|
peerTrusted: isInside(peer, trusted),
|
||||||
start: start,
|
start: start,
|
||||||
line: requestlog.Line{
|
line: requestlog.Line{
|
||||||
Time: requestlog.FormatTime(start),
|
Time: requestlog.FormatTime(start),
|
||||||
Instance: h.config.InstanceName,
|
ClientIP: client.String(),
|
||||||
ClientIP: client.String(),
|
PeerIP: peer.String(),
|
||||||
Method: r.Method,
|
Method: r.Method,
|
||||||
Scheme: scheme(r, peerTrusted),
|
Host: r.Host,
|
||||||
Host: r.Host,
|
Path: r.URL.EscapedPath(),
|
||||||
Path: r.URL.EscapedPath(),
|
Query: r.URL.RawQuery,
|
||||||
Query: r.URL.RawQuery,
|
Protocol: r.Proto,
|
||||||
Protocol: r.Proto,
|
Referer: r.Referer(),
|
||||||
Referer: r.Referer(),
|
UserAgent: r.UserAgent(),
|
||||||
UserAgent: r.UserAgent(),
|
Action: requestlog.ActionForward,
|
||||||
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 {
|
if r.Body != http.NoBody {
|
||||||
rq.body = &requestBody{body: limitBody(r.Body, h.config.RequestMaxBytes), rq: rq}
|
rq.body = &requestBody{body: limitBody(r.Body, h.config.RequestMaxBytes), rq: rq}
|
||||||
}
|
}
|
||||||
@@ -141,47 +101,19 @@ func (h *handler) newRequest(w http.ResponseWriter, r *http.Request) *request {
|
|||||||
return 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
|
// 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
|
// 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
|
// returns nil to let the request through. The rate limits come first, so
|
||||||
// first, answered with SWWAF_BAN_RESPONSE, or 403 for a block rule, and
|
// that every request is counted, one refused for its size too.
|
||||||
// then the size limit, so that a request the rate limits count is counted
|
func (rq *request) check() *refusal {
|
||||||
// even when it is refused for its size. In observe mode a request
|
limitHit := rq.h.limiter.Count(clientGroup(rq.client), rq.start)
|
||||||
// checkClient refuses goes on to the size limit like any other. ctx is
|
if limitHit != "" {
|
||||||
// the request's own context.
|
rq.line.LimitHit = limitHit
|
||||||
func (rq *request) check(ctx context.Context) *refusal {
|
|
||||||
action := rq.checkClient(ctx)
|
|
||||||
|
|
||||||
switch {
|
return &refusal{
|
||||||
case action == "":
|
status: http.StatusTooManyRequests,
|
||||||
case rq.h.config.Observe:
|
action: requestlog.ActionRateLimited,
|
||||||
// The log line names what enforce mode would have done.
|
}
|
||||||
rq.line.WouldAction = action
|
|
||||||
case action == requestlog.ActionRuleBlocked:
|
|
||||||
return &refusal{status: http.StatusForbidden, action: action}
|
|
||||||
default:
|
|
||||||
return rq.banResponse(action)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
maxBytes := rq.h.config.RequestMaxBytes
|
maxBytes := rq.h.config.RequestMaxBytes
|
||||||
@@ -189,79 +121,12 @@ func (rq *request) check(ctx context.Context) *refusal {
|
|||||||
return &refusal{
|
return &refusal{
|
||||||
status: http.StatusRequestEntityTooLarge,
|
status: http.StatusRequestEntityTooLarge,
|
||||||
action: requestlog.ActionTooLarge,
|
action: requestlog.ActionTooLarge,
|
||||||
limit: "SWWAF_REQUEST_MAX_BYTES",
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
return nil
|
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, and is not looked up. 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, then the lookup of
|
|
||||||
// its AS number and country, 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 or the
|
|
||||||
// request's path is exempt under SWWAF_RATE_LIMIT_EXEMPT_PATHS, so that
|
|
||||||
// every other request is counted, and last the rule files. 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
|
|
||||||
}
|
|
||||||
|
|
||||||
rq.lookUp(ctx)
|
|
||||||
|
|
||||||
if rq.countryDenied() {
|
|
||||||
return requestlog.ActionCountryDenied
|
|
||||||
}
|
|
||||||
|
|
||||||
exempt := isInside(rq.client, cfg.RateLimitExemptNets) ||
|
|
||||||
pathExempt(rq.in.URL, cfg.RateLimitExemptPaths)
|
|
||||||
if !exempt && rq.limitBroken(now) {
|
|
||||||
return requestlog.ActionRateLimited
|
|
||||||
}
|
|
||||||
|
|
||||||
return rq.checkRules(now)
|
|
||||||
}
|
|
||||||
|
|
||||||
// pathExempt reports whether the rate limits leave out a request for u
|
|
||||||
// because of SWWAF_RATE_LIMIT_EXEMPT_PATHS: whether its path as sent, the
|
|
||||||
// path the app receives, not percent-decoded, starts with one of
|
|
||||||
// prefixes, so that /%61ssets/x is not under /assets/ for an app whose
|
|
||||||
// router matches the path as received. A request whose decoded path
|
|
||||||
// contains .. anywhere or a backslash, or whose path as sent holds an
|
|
||||||
// encoded slash (%2F or %2f), never is, since an app may act on it as a
|
|
||||||
// path outside every prefix: /assets/..%2Flogin as /login, or /assets%2Fx
|
|
||||||
// as one path segment, as Go's router does.
|
|
||||||
func pathExempt(u *url.URL, prefixes []string) bool {
|
|
||||||
decoded := u.Path
|
|
||||||
// EscapedPath is the path as the app receives it, not decoded.
|
|
||||||
sent := u.EscapedPath()
|
|
||||||
|
|
||||||
if strings.Contains(decoded, "..") || strings.Contains(decoded, `\`) ||
|
|
||||||
strings.Contains(strings.ToLower(sent), "%2f") {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
|
|
||||||
return slices.ContainsFunc(prefixes, func(prefix string) bool {
|
|
||||||
return strings.HasPrefix(sent, prefix)
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
// forward passes the request to the app and the app's answer back. ctx
|
// forward passes the request to the app and the app's answer back. ctx
|
||||||
// is the request's own context.
|
// is the request's own context.
|
||||||
func (rq *request) forward(ctx context.Context) {
|
func (rq *request) forward(ctx context.Context) {
|
||||||
@@ -270,9 +135,7 @@ func (rq *request) forward(ctx context.Context) {
|
|||||||
|
|
||||||
rq.cancel = cancel
|
rq.cancel = cancel
|
||||||
ctx = httptrace.WithClientTrace(ctx, &httptrace.ClientTrace{
|
ctx = httptrace.WithClientTrace(ctx, &httptrace.ClientTrace{
|
||||||
GotConn: rq.gotConn,
|
WroteRequest: rq.wroteRequest,
|
||||||
WroteRequest: rq.wroteRequest,
|
|
||||||
GotFirstResponseByte: rq.gotFirstResponseByte,
|
|
||||||
})
|
})
|
||||||
|
|
||||||
out := rq.in.WithContext(ctx)
|
out := rq.in.WithContext(ctx)
|
||||||
@@ -295,10 +158,7 @@ func (rq *request) forward(ctx context.Context) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// rewrite makes the request the app receives: the client's request,
|
// rewrite makes the request the app receives: the client's request,
|
||||||
// unchanged, sent to SWWAF_UPSTREAM_URL, with the forwarded headers and
|
// unchanged, sent to SWWAF_UPSTREAM_URL, with the forwarded headers set.
|
||||||
// the request's id set, without any X-Client-ASN or X-Client-Country the
|
|
||||||
// client sent, whatever SWWAF_ADD_LOOKUP_HEADERS says, and, while it is
|
|
||||||
// set, with the client's AS number and country in them.
|
|
||||||
func (rq *request) rewrite(pr *httputil.ProxyRequest) {
|
func (rq *request) rewrite(pr *httputil.ProxyRequest) {
|
||||||
upstream := rq.h.config.UpstreamURL
|
upstream := rq.h.config.UpstreamURL
|
||||||
pr.Out.URL.Scheme = upstream.Scheme
|
pr.Out.URL.Scheme = upstream.Scheme
|
||||||
@@ -307,13 +167,6 @@ func (rq *request) rewrite(pr *httputil.ProxyRequest) {
|
|||||||
// the query as the client sent it.
|
// the query as the client sent it.
|
||||||
pr.Out.URL.RawQuery = pr.In.URL.RawQuery
|
pr.Out.URL.RawQuery = pr.In.URL.RawQuery
|
||||||
setForwardedHeaders(pr.In, pr.Out, rq.peer, rq.peerTrusted)
|
setForwardedHeaders(pr.In, pr.Out, rq.peer, rq.peerTrusted)
|
||||||
pr.Out.Header.Set(requestIDHeader, rq.line.RequestID)
|
|
||||||
pr.Out.Header.Del(asnHeader)
|
|
||||||
pr.Out.Header.Del(countryHeader)
|
|
||||||
|
|
||||||
if rq.h.config.AddLookupHeaders {
|
|
||||||
setLookupHeaders(pr.Out.Header, rq.line.ASN, rq.line.Country)
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// modifyResponse looks at the app's answer before ReverseProxy passes it
|
// modifyResponse looks at the app's answer before ReverseProxy passes it
|
||||||
@@ -327,18 +180,13 @@ func (rq *request) modifyResponse(res *http.Response) error {
|
|||||||
// connection it takes over, not through rq.out.
|
// connection it takes over, not through rq.out.
|
||||||
rq.stopTimers()
|
rq.stopTimers()
|
||||||
rq.out.status = res.StatusCode
|
rq.out.status = res.StatusCode
|
||||||
rq.line.Websocket = true
|
|
||||||
|
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
maxBytes := rq.h.config.ResponseMaxBytes
|
maxBytes := rq.h.config.ResponseMaxBytes
|
||||||
if maxBytes > 0 && res.Body != http.NoBody && res.ContentLength > maxBytes {
|
if maxBytes > 0 && res.Body != http.NoBody && res.ContentLength > maxBytes {
|
||||||
rq.refuse(refusal{
|
rq.refuse(refusal{status: http.StatusBadGateway, action: requestlog.ActionTooLarge})
|
||||||
status: http.StatusBadGateway,
|
|
||||||
action: requestlog.ActionTooLarge,
|
|
||||||
limit: "SWWAF_RESPONSE_MAX_BYTES",
|
|
||||||
})
|
|
||||||
|
|
||||||
return errResponseTooLarge
|
return errResponseTooLarge
|
||||||
}
|
}
|
||||||
@@ -378,13 +226,6 @@ func (rq *request) answer(r refusal) {
|
|||||||
return // too late to answer: the connection can only be cut
|
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
|
// 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
|
// sending until its time is up, so that Go's server can read the
|
||||||
// rest of the body and end the request cleanly.
|
// rest of the body and end the request cleanly.
|
||||||
@@ -404,18 +245,13 @@ func (rq *request) answer(r refusal) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// refuse records r, unless an earlier refusal was, and ends the request
|
// refuse records r, unless an earlier refusal was, and ends the request
|
||||||
// to the app, if one was made: smallwebwaf reads the body of a request
|
// to the app.
|
||||||
// it answers itself too.
|
|
||||||
func (rq *request) refuse(r refusal) {
|
func (rq *request) refuse(r refusal) {
|
||||||
rq.refused.CompareAndSwap(nil, &r)
|
rq.refused.CompareAndSwap(nil, &r)
|
||||||
|
rq.cancel()
|
||||||
if rq.cancel != nil {
|
|
||||||
rq.cancel()
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// finish ends the request's timeouts, counts it in the metrics and writes
|
// finish ends the request's timeouts and writes its log line.
|
||||||
// its log line.
|
|
||||||
func (rq *request) finish() {
|
func (rq *request) finish() {
|
||||||
rq.stopTimers()
|
rq.stopTimers()
|
||||||
|
|
||||||
@@ -427,110 +263,35 @@ func (rq *request) finish() {
|
|||||||
line := &rq.line
|
line := &rq.line
|
||||||
line.Status = rq.out.status
|
line.Status = rq.out.status
|
||||||
line.ResponseBytes = rq.out.bytes
|
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 {
|
if rq.body != nil {
|
||||||
line.RequestBytes = rq.body.bytes.Load()
|
line.RequestBytes = rq.body.bytes.Load()
|
||||||
}
|
}
|
||||||
|
|
||||||
// limit is the setting whose size or time limit the request passed.
|
|
||||||
var limit string
|
|
||||||
|
|
||||||
switch {
|
switch {
|
||||||
case refused != nil:
|
case refused != nil:
|
||||||
line.Action = refused.action
|
line.Action = refused.action
|
||||||
limit = refused.limit
|
|
||||||
case errors.Is(rq.out.err, os.ErrDeadlineExceeded):
|
case errors.Is(rq.out.err, os.ErrDeadlineExceeded):
|
||||||
// The client took longer than SWWAF_CLIENT_RESPONSE_TIMEOUT to
|
// The client took longer than SWWAF_CLIENT_RESPONSE_TIMEOUT to
|
||||||
// take the response.
|
// take the response.
|
||||||
line.Action = requestlog.ActionTimedOut
|
line.Action = requestlog.ActionTimedOut
|
||||||
limit = "SWWAF_CLIENT_RESPONSE_TIMEOUT"
|
|
||||||
case !rq.complete && (rq.out.err != nil || rq.in.Context().Err() != nil):
|
case !rq.complete && (rq.out.err != nil || rq.in.Context().Err() != nil):
|
||||||
line.Aborted = true
|
line.Aborted = true
|
||||||
}
|
}
|
||||||
|
|
||||||
now := time.Now()
|
now := time.Now()
|
||||||
duration := now.Sub(rq.start)
|
line.DurationTotal = requestlog.Milliseconds(now.Sub(rq.start))
|
||||||
line.DurationTotal = requestlog.Milliseconds(duration)
|
|
||||||
line.DurationChecks = timing(rq.start, rq.checked)
|
|
||||||
|
|
||||||
var upstreamDuration time.Duration
|
|
||||||
|
|
||||||
if !rq.upstreamStart.IsZero() {
|
if !rq.upstreamStart.IsZero() {
|
||||||
upstreamDuration = now.Sub(rq.upstreamStart)
|
line.DurationUpstreamTotal = requestlog.Milliseconds(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)
|
err := requestlog.Write(rq.h.requestLog, line)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
rq.h.processLog.Error("writing the request log failed", "error", err.Error())
|
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, and then, for a client that was looked up, the lookup
|
|
||||||
// database's answer about it, or the answer from GeoJS kept about it, to
|
|
||||||
// that history and to the notes of the bans on its netblock: an answer
|
|
||||||
// may have come before either was there, and one from GeoJS that comes
|
|
||||||
// later is added when it comes.
|
|
||||||
func (rq *request) addToHistory() {
|
|
||||||
var requestBytes int64
|
|
||||||
if rq.body != nil {
|
|
||||||
requestBytes = rq.body.bytes.Load()
|
|
||||||
}
|
|
||||||
|
|
||||||
forwarded := !rq.upstreamStart.IsZero()
|
|
||||||
group := clientGroup(rq.client)
|
|
||||||
|
|
||||||
rq.h.limiter.AddToHistory(group, rq.h.now(), ratelimit.Request{
|
|
||||||
Forwarded: forwarded,
|
|
||||||
Refused: !forwarded && rq.refused.Load() != nil,
|
|
||||||
Status: rq.out.status,
|
|
||||||
RequestBytes: requestBytes,
|
|
||||||
ResponseBytes: rq.out.bytes,
|
|
||||||
BrokeLimit: rq.line.Offence == requestlog.OffenceLimit,
|
|
||||||
})
|
|
||||||
|
|
||||||
if !rq.lookedUp {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
// The lookup database's answer was there at once.
|
|
||||||
if rq.h.config.LookupSource == "file" {
|
|
||||||
rq.h.addLookup(rq.lookupAnswer)
|
|
||||||
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
answer, kept := rq.h.geojs.Kept(group)
|
|
||||||
if kept {
|
|
||||||
rq.h.addLookup(answer)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// clientRequestDeadline is when the client must have sent its whole
|
// clientRequestDeadline is when the client must have sent its whole
|
||||||
// request, or zero when SWWAF_CLIENT_REQUEST_TIMEOUT is off.
|
// request, or zero when SWWAF_CLIENT_REQUEST_TIMEOUT is off.
|
||||||
func (rq *request) clientRequestDeadline() time.Time {
|
func (rq *request) clientRequestDeadline() time.Time {
|
||||||
@@ -563,26 +324,21 @@ func (rq *request) startRequestTimers() {
|
|||||||
|
|
||||||
if rq.body != nil && rq.h.config.ClientRequestTimeout > 0 {
|
if rq.body != nil && rq.h.config.ClientRequestTimeout > 0 {
|
||||||
rq.clientRequestTimer = time.AfterFunc(
|
rq.clientRequestTimer = time.AfterFunc(
|
||||||
time.Until(rq.clientRequestDeadline()), func() {
|
time.Until(rq.clientRequestDeadline()), rq.requestTimedOut)
|
||||||
rq.requestTimedOut("SWWAF_CLIENT_REQUEST_TIMEOUT")
|
|
||||||
})
|
|
||||||
}
|
}
|
||||||
|
|
||||||
timeout := rq.h.config.UpstreamRequestTimeout
|
timeout := rq.h.config.UpstreamRequestTimeout
|
||||||
if timeout > 0 {
|
if timeout > 0 {
|
||||||
rq.upstreamRequestTimer = time.AfterFunc(timeout, func() {
|
rq.upstreamRequestTimer = time.AfterFunc(timeout, rq.requestTimedOut)
|
||||||
rq.requestTimedOut("SWWAF_UPSTREAM_REQUEST_TIMEOUT")
|
|
||||||
})
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// requestTimedOut is called when limit, SWWAF_CLIENT_REQUEST_TIMEOUT or
|
// requestTimedOut is called when a request timeout runs out while the
|
||||||
// SWWAF_UPSTREAM_REQUEST_TIMEOUT, runs out while the request is still on
|
// request is still on its way to the app. The answer names the side
|
||||||
// its way to the app. The answer names the side smallwebwaf was waiting
|
// smallwebwaf was waiting on at that moment: 408 when it was waiting for
|
||||||
// on at that moment: 408 when it was waiting for the client to send more
|
// the client to send more of its body, 504 when it was waiting for the
|
||||||
// of its body, 504 when it was waiting for the app to be reached or to
|
// app to be reached or to take what it had.
|
||||||
// take what it had.
|
func (rq *request) requestTimedOut() {
|
||||||
func (rq *request) requestTimedOut(limit string) {
|
|
||||||
rq.mu.Lock()
|
rq.mu.Lock()
|
||||||
defer rq.mu.Unlock()
|
defer rq.mu.Unlock()
|
||||||
|
|
||||||
@@ -594,7 +350,6 @@ func (rq *request) requestTimedOut(limit string) {
|
|||||||
rq.refuse(refusal{
|
rq.refuse(refusal{
|
||||||
status: http.StatusGatewayTimeout,
|
status: http.StatusGatewayTimeout,
|
||||||
action: requestlog.ActionTimedOut,
|
action: requestlog.ActionTimedOut,
|
||||||
limit: limit,
|
|
||||||
})
|
})
|
||||||
|
|
||||||
return
|
return
|
||||||
@@ -603,7 +358,6 @@ func (rq *request) requestTimedOut(limit string) {
|
|||||||
rq.refuse(refusal{
|
rq.refuse(refusal{
|
||||||
status: http.StatusRequestTimeout,
|
status: http.StatusRequestTimeout,
|
||||||
action: requestlog.ActionTimedOut,
|
action: requestlog.ActionTimedOut,
|
||||||
limit: limit,
|
|
||||||
})
|
})
|
||||||
// The transport gives up on the app only once its Read of the
|
// 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
|
// client's body returns, so that Read is ended now. The lock keeps
|
||||||
@@ -619,24 +373,6 @@ func (rq *request) bodyReceived() {
|
|||||||
stopTimer(rq.clientRequestTimer)
|
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:
|
// wroteRequest is called once the app has been sent the whole request:
|
||||||
// the request timeouts end and SWWAF_UPSTREAM_RESPONSE_TIMEOUT starts.
|
// the request timeouts end and SWWAF_UPSTREAM_RESPONSE_TIMEOUT starts.
|
||||||
func (rq *request) wroteRequest(info httptrace.WroteRequestInfo) {
|
func (rq *request) wroteRequest(info httptrace.WroteRequestInfo) {
|
||||||
@@ -671,7 +407,6 @@ func (rq *request) responseTimedOut() {
|
|||||||
rq.refuse(refusal{
|
rq.refuse(refusal{
|
||||||
status: http.StatusGatewayTimeout,
|
status: http.StatusGatewayTimeout,
|
||||||
action: requestlog.ActionTimedOut,
|
action: requestlog.ActionTimedOut,
|
||||||
limit: "SWWAF_UPSTREAM_RESPONSE_TIMEOUT",
|
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -1,368 +0,0 @@
|
|||||||
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"])
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@@ -1,233 +0,0 @@
|
|||||||
package proxy_test
|
|
||||||
|
|
||||||
import (
|
|
||||||
"net/http"
|
|
||||||
"net/netip"
|
|
||||||
"os"
|
|
||||||
"path/filepath"
|
|
||||||
"slices"
|
|
||||||
"testing"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
"sneak.berlin/go/smallwebwaf/internal/bans"
|
|
||||||
"sneak.berlin/go/smallwebwaf/internal/requestlog"
|
|
||||||
)
|
|
||||||
|
|
||||||
// testRules are the rules most tests here load: a block rule for
|
|
||||||
// /blocked and a ban rule for /.env.
|
|
||||||
const testRules = `
|
|
||||||
blocked path block ^/blocked$
|
|
||||||
probe path ban ^/\.env$
|
|
||||||
`
|
|
||||||
|
|
||||||
func TestEachRuleAction(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
s, clk, server := startWithClock(t, "", map[string]string{
|
|
||||||
rulesDir: writeRules(t, "noted path log ^/\n"+testRules),
|
|
||||||
banResponse: "429",
|
|
||||||
})
|
|
||||||
start := clk.Now()
|
|
||||||
|
|
||||||
// A log rule notes its match, and lets the request through.
|
|
||||||
line := s.get(client, http.StatusOK, requestlog.ActionForward)
|
|
||||||
wantRuleIDs(t, line, "noted")
|
|
||||||
|
|
||||||
// A block rule refuses with 403, whatever SWWAF_BAN_RESPONSE is, and
|
|
||||||
// bans no one.
|
|
||||||
line = s.request(client, "/blocked", http.StatusForbidden,
|
|
||||||
requestlog.ActionRuleBlocked)
|
|
||||||
wantRuleIDs(t, line, "noted", "blocked")
|
|
||||||
s.get(client, http.StatusOK, requestlog.ActionForward)
|
|
||||||
|
|
||||||
// A ban rule refuses with SWWAF_BAN_RESPONSE, and bans the client for
|
|
||||||
// seven days, the default.
|
|
||||||
line = s.request(client, "/.env", http.StatusTooManyRequests, requestlog.ActionBanned)
|
|
||||||
wantRuleIDs(t, line, "noted", "probe")
|
|
||||||
|
|
||||||
if line.BanExpires != requestlog.FormatTime(start.Add(7*24*time.Hour)) {
|
|
||||||
t.Errorf("log line has ban_expires %q, want seven days on", line.BanExpires)
|
|
||||||
}
|
|
||||||
|
|
||||||
netblock := netip.MustParsePrefix(client + "/32")
|
|
||||||
want := bans.Ban{
|
|
||||||
Netblock: netblock,
|
|
||||||
Start: start,
|
|
||||||
Expires: start.Add(7 * 24 * time.Hour),
|
|
||||||
Cause: bans.CauseAttack,
|
|
||||||
Reason: "matched the rule probe",
|
|
||||||
Notes: bans.Notes{
|
|
||||||
RuleID: "probe",
|
|
||||||
Target: "path",
|
|
||||||
Request: bans.Request{
|
|
||||||
Time: start,
|
|
||||||
Method: http.MethodGet,
|
|
||||||
Host: appHost,
|
|
||||||
Path: "/.env",
|
|
||||||
Status: http.StatusTooManyRequests,
|
|
||||||
UserAgent: userAgent,
|
|
||||||
},
|
|
||||||
// The four requests up to and including the probe.
|
|
||||||
Requests: 4,
|
|
||||||
},
|
|
||||||
}
|
|
||||||
|
|
||||||
got := server.Ledger.Bans(netblock)
|
|
||||||
if len(got) != 1 || got[0] != want {
|
|
||||||
t.Fatalf("bans\n%+v\nwant\n%+v", got, want)
|
|
||||||
}
|
|
||||||
|
|
||||||
// The next request is refused under the ban, without being checked
|
|
||||||
// against the rules, and makes the ban permanent.
|
|
||||||
clk.advance(time.Hour)
|
|
||||||
|
|
||||||
line = s.get(client, http.StatusTooManyRequests, requestlog.ActionBanned)
|
|
||||||
wantRuleIDs(t, line)
|
|
||||||
|
|
||||||
if line.BanExpires != permanent {
|
|
||||||
t.Errorf("log line has ban_expires %q, want permanent", line.BanExpires)
|
|
||||||
}
|
|
||||||
|
|
||||||
clk.advance(365 * 24 * time.Hour)
|
|
||||||
s.get(client, http.StatusTooManyRequests, requestlog.ActionBanned)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestNextClearSignOfAttackAfterABanBansPermanently(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
s, clk, _ := startWithClock(t, "", map[string]string{
|
|
||||||
rulesDir: writeRules(t, testRules),
|
|
||||||
attackBanDuration: "1h",
|
|
||||||
})
|
|
||||||
|
|
||||||
// The first probe bans for SWWAF_ATTACK_BAN_DURATION.
|
|
||||||
line := s.request(client, "/.env", http.StatusForbidden, requestlog.ActionBanned)
|
|
||||||
if line.BanExpires != requestlog.FormatTime(clk.Now().Add(time.Hour)) {
|
|
||||||
t.Errorf("log line has ban_expires %q, want an hour on", line.BanExpires)
|
|
||||||
}
|
|
||||||
|
|
||||||
// Once that ban has run out without a request, the client is served,
|
|
||||||
// and its next probe bans it for good.
|
|
||||||
clk.advance(time.Hour)
|
|
||||||
s.get(client, http.StatusOK, requestlog.ActionForward)
|
|
||||||
|
|
||||||
line = s.request(client, "/.env", http.StatusForbidden, requestlog.ActionBanned)
|
|
||||||
if line.BanExpires != permanent {
|
|
||||||
t.Errorf("log line has ban_expires %q, want permanent", line.BanExpires)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestRulesComeAfterTheOtherChecks(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
const (
|
|
||||||
allowed = "192.0.2.60" // in SWWAF_ALLOW_NETS
|
|
||||||
exempt = "192.0.2.50" // in SWWAF_RATE_LIMIT_EXEMPT_NETS
|
|
||||||
)
|
|
||||||
|
|
||||||
s, _, server := startWithClock(t, "", map[string]string{
|
|
||||||
rulesDir: writeRules(t, testRules),
|
|
||||||
allowNets: allowed,
|
|
||||||
rateLimitExemptNets: exempt,
|
|
||||||
rateLimitPerMinute: "1",
|
|
||||||
})
|
|
||||||
|
|
||||||
// A client in SWWAF_ALLOW_NETS is not checked.
|
|
||||||
line := s.request(allowed, "/.env", http.StatusOK, requestlog.ActionForward)
|
|
||||||
wantRuleIDs(t, line)
|
|
||||||
|
|
||||||
// A probe over the rate limit breaks the limit before any rule sees
|
|
||||||
// it.
|
|
||||||
s.get(client, http.StatusOK, requestlog.ActionForward)
|
|
||||||
line = s.request(client, "/.env", http.StatusForbidden, requestlog.ActionRateLimited)
|
|
||||||
wantRuleIDs(t, line)
|
|
||||||
|
|
||||||
limitBan := server.Ledger.Bans(netip.MustParsePrefix(client + "/32"))
|
|
||||||
if len(limitBan) != 1 || limitBan[0].Cause != bans.CauseLimit {
|
|
||||||
t.Errorf("bans %+v, want one for a broken limit", limitBan)
|
|
||||||
}
|
|
||||||
|
|
||||||
// A client the rate limits do not apply to is still checked.
|
|
||||||
s.get(exempt, http.StatusOK, requestlog.ActionForward)
|
|
||||||
s.request(exempt, "/.env", http.StatusForbidden, requestlog.ActionBanned)
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestObserveModeLogsWhatTheRulesWouldDo(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
s, _, server := startWithClock(t, "", map[string]string{
|
|
||||||
rulesDir: writeRules(t, testRules),
|
|
||||||
mode: observe,
|
|
||||||
})
|
|
||||||
|
|
||||||
line := s.request(client, "/blocked", http.StatusOK, requestlog.ActionForward)
|
|
||||||
wantWouldAction(t, line, requestlog.ActionRuleBlocked)
|
|
||||||
wantRuleIDs(t, line, "blocked")
|
|
||||||
|
|
||||||
line = s.request(client, "/.env", http.StatusOK, requestlog.ActionForward)
|
|
||||||
wantWouldAction(t, line, requestlog.ActionBanned)
|
|
||||||
wantRuleIDs(t, line, "probe")
|
|
||||||
|
|
||||||
if line.BanExpires != "" {
|
|
||||||
t.Errorf("log line has ban_expires %q, want none", line.BanExpires)
|
|
||||||
}
|
|
||||||
|
|
||||||
// No ban was made.
|
|
||||||
line = s.get(client, http.StatusOK, requestlog.ActionForward)
|
|
||||||
wantWouldAction(t, line, "")
|
|
||||||
|
|
||||||
if got := server.Ledger.Snapshot(); len(got) != 0 {
|
|
||||||
t.Errorf("bans %+v, want none", got)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestMetricsCountRuleMatchesAndBansForAnAttack(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
const scraper = "192.0.2.200"
|
|
||||||
|
|
||||||
s, _, _ := startWithClock(t, "", map[string]string{
|
|
||||||
rulesDir: writeRules(t, testRules),
|
|
||||||
metricsToken: token,
|
|
||||||
})
|
|
||||||
|
|
||||||
s.request(client, "/blocked", http.StatusForbidden, requestlog.ActionRuleBlocked)
|
|
||||||
s.request(client, "/.env", http.StatusForbidden, requestlog.ActionBanned)
|
|
||||||
s.get(client, http.StatusForbidden, requestlog.ActionBanned)
|
|
||||||
|
|
||||||
metrics := s.scrape(scraper)
|
|
||||||
wantMetric(t, metrics,
|
|
||||||
`smallwebwaf_rule_matches_total{action="block",instance="app",rule_id="blocked"}`, 1)
|
|
||||||
wantMetric(t, metrics,
|
|
||||||
`smallwebwaf_rule_matches_total{action="ban",instance="app",rule_id="probe"}`, 1)
|
|
||||||
wantMetric(t, metrics, `smallwebwaf_rules_loaded{instance="app"}`, 2)
|
|
||||||
wantMetric(t, metrics, `smallwebwaf_requests_total{action="rule_blocked",`+
|
|
||||||
`instance="app",status_class="4xx"}`, 1)
|
|
||||||
wantMetric(t, metrics, `smallwebwaf_bans_made_total{cause="attack",instance="app"}`, 1)
|
|
||||||
wantMetric(t, metrics, `smallwebwaf_bans_made_total{cause="limit",instance="app"}`, 0)
|
|
||||||
wantMetric(t, metrics, `smallwebwaf_permanent_bans{instance="app"}`, 1)
|
|
||||||
}
|
|
||||||
|
|
||||||
// writeRules writes content as a rule file into a new directory, and
|
|
||||||
// returns the directory.
|
|
||||||
func writeRules(t *testing.T, content string) string {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
dir := t.TempDir()
|
|
||||||
|
|
||||||
err := os.WriteFile(filepath.Join(dir, "test.rules"), []byte(content), 0o600)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("write the rule file: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
return dir
|
|
||||||
}
|
|
||||||
|
|
||||||
// wantRuleIDs checks the request log line's rule_ids.
|
|
||||||
func wantRuleIDs(t *testing.T, line logLine, want ...string) {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
if !slices.Equal(line.RuleIDs, want) {
|
|
||||||
t.Errorf("log line has rule_ids %v, want %v", line.RuleIDs, want)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@@ -1,39 +0,0 @@
|
|||||||
package proxy
|
|
||||||
|
|
||||||
import (
|
|
||||||
"time"
|
|
||||||
|
|
||||||
"sneak.berlin/go/smallwebwaf/internal/requestlog"
|
|
||||||
"sneak.berlin/go/smallwebwaf/internal/rules"
|
|
||||||
)
|
|
||||||
|
|
||||||
// checkRules checks the request against the rules of the rule files at
|
|
||||||
// now, notes the ids of those it matches in the log line, and returns the
|
|
||||||
// action of the rule that refuses it, ActionRuleBlocked for a block rule
|
|
||||||
// and ActionBanned for a ban rule, or "" when none does. A ban rule bans
|
|
||||||
// the client's netblock for a clear sign of attack, or in observe mode
|
|
||||||
// raises the alert for the ban it would have made.
|
|
||||||
func (rq *request) checkRules(now time.Time) string {
|
|
||||||
matched := rq.h.rules.Match(rq.in)
|
|
||||||
|
|
||||||
for _, rule := range matched {
|
|
||||||
rq.line.RuleIDs = append(rq.line.RuleIDs, rule.ID)
|
|
||||||
rq.h.metrics.RuleMatched(rule.ID, rule.Action)
|
|
||||||
}
|
|
||||||
|
|
||||||
if len(matched) == 0 {
|
|
||||||
return ""
|
|
||||||
}
|
|
||||||
|
|
||||||
// Only the last rule matched can refuse the request.
|
|
||||||
switch last := matched[len(matched)-1]; last.Action {
|
|
||||||
case rules.ActionBlock:
|
|
||||||
return requestlog.ActionRuleBlocked
|
|
||||||
case rules.ActionBan:
|
|
||||||
rq.banForAttack(now, last)
|
|
||||||
|
|
||||||
return requestlog.ActionBanned
|
|
||||||
default:
|
|
||||||
return ""
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@@ -1,184 +0,0 @@
|
|||||||
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)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@@ -5,10 +5,8 @@ import (
|
|||||||
"io"
|
"io"
|
||||||
"net"
|
"net"
|
||||||
"net/http"
|
"net/http"
|
||||||
"net/http/httptest"
|
|
||||||
"strconv"
|
"strconv"
|
||||||
"sync"
|
"sync"
|
||||||
"sync/atomic"
|
|
||||||
"testing"
|
"testing"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
@@ -28,36 +26,41 @@ func TestRequestTimeouts(t *testing.T) {
|
|||||||
|
|
||||||
for _, tc := range []struct {
|
for _, tc := range []struct {
|
||||||
name string
|
name string
|
||||||
// limit is the setting set to shortTimeout, which runs out; long
|
env map[string]string
|
||||||
// is one set to longTimeoutSetting, which does not, or "".
|
|
||||||
limit, long string
|
|
||||||
// appTakesNothing has the app never read, while the client sends
|
// appTakesNothing has the app never read, while the client sends
|
||||||
// as fast as it can; otherwise the app reads, and the client
|
// as fast as it can; otherwise the app reads, and the client
|
||||||
// stops sending halfway.
|
// stops sending halfway. smallwebwaf then waits on the client
|
||||||
|
// only once it has connected to the app and passed on the first
|
||||||
|
// bytes; a test process held up for shortTimeout before that
|
||||||
|
// gets 504, which is the right answer, and the case fails.
|
||||||
appTakesNothing bool
|
appTakesNothing bool
|
||||||
want int
|
want int
|
||||||
}{
|
}{
|
||||||
{
|
{
|
||||||
name: "client request timeout, waiting on the client",
|
name: "client request timeout, waiting on the client",
|
||||||
limit: clientRequestTimeout,
|
env: map[string]string{clientRequestTimeout: shortTimeoutSetting},
|
||||||
want: http.StatusRequestTimeout,
|
want: http.StatusRequestTimeout,
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
name: "upstream request timeout, waiting on the client",
|
name: "upstream request timeout, waiting on the client",
|
||||||
limit: upstreamRequestTimeout,
|
env: map[string]string{
|
||||||
long: clientRequestTimeout,
|
upstreamRequestTimeout: shortTimeoutSetting,
|
||||||
want: http.StatusRequestTimeout,
|
clientRequestTimeout: longTimeoutSetting,
|
||||||
|
},
|
||||||
|
want: http.StatusRequestTimeout,
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
name: "upstream request timeout, waiting on the app",
|
name: "upstream request timeout, waiting on the app",
|
||||||
limit: upstreamRequestTimeout,
|
env: map[string]string{upstreamRequestTimeout: shortTimeoutSetting},
|
||||||
appTakesNothing: true,
|
appTakesNothing: true,
|
||||||
want: http.StatusGatewayTimeout,
|
want: http.StatusGatewayTimeout,
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
name: "client request timeout, waiting on the app",
|
name: "client request timeout, waiting on the app",
|
||||||
limit: clientRequestTimeout,
|
env: map[string]string{
|
||||||
long: upstreamRequestTimeout,
|
clientRequestTimeout: shortTimeoutSetting,
|
||||||
|
upstreamRequestTimeout: longTimeoutSetting,
|
||||||
|
},
|
||||||
appTakesNothing: true,
|
appTakesNothing: true,
|
||||||
want: http.StatusGatewayTimeout,
|
want: http.StatusGatewayTimeout,
|
||||||
},
|
},
|
||||||
@@ -66,53 +69,32 @@ func TestRequestTimeouts(t *testing.T) {
|
|||||||
t.Parallel()
|
t.Parallel()
|
||||||
|
|
||||||
var (
|
var (
|
||||||
app *httptest.Server
|
|
||||||
appURL string
|
appURL string
|
||||||
sendRequest func(*testing.T, string) net.Conn
|
sendRequest func(*testing.T, string) net.Conn
|
||||||
appGotBody atomic.Bool
|
|
||||||
)
|
)
|
||||||
|
|
||||||
if tc.appTakesNothing {
|
if tc.appTakesNothing {
|
||||||
appURL, sendRequest = startAppThatTakesNothing(t), sendLargeBody
|
appURL, sendRequest = startAppThatTakesNothing(t), sendLargeBody
|
||||||
} else {
|
} else {
|
||||||
app = startApp(t, func(_ http.ResponseWriter, r *http.Request) {
|
appURL, sendRequest = startApp(t, readBody).URL, sendPartOfBody
|
||||||
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}
|
addr, out := startProxy(t, appURL, tc.env)
|
||||||
if tc.long != "" {
|
|
||||||
env[tc.long] = longTimeoutSetting
|
|
||||||
}
|
|
||||||
|
|
||||||
addr, out := startProxy(t, appURL, env)
|
|
||||||
start := time.Now()
|
start := time.Now()
|
||||||
got := readResponse(t, sendRequest(t, addr))
|
conn := sendRequest(t, addr)
|
||||||
|
|
||||||
|
wantStatus(t, readResponse(t, conn), tc.want)
|
||||||
wantTimedOut(t, start)
|
wantTimedOut(t, start)
|
||||||
|
wantLine(t, out.requestLine(t), tc.want, requestlog.ActionTimedOut)
|
||||||
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)
|
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// readBody is an app that reads the request body, then answers.
|
||||||
|
func readBody(_ http.ResponseWriter, r *http.Request) {
|
||||||
|
_, _ = io.Copy(io.Discard, r.Body)
|
||||||
|
}
|
||||||
|
|
||||||
// startAppThatTakesNothing starts an app that accepts connections and
|
// startAppThatTakesNothing starts an app that accepts connections and
|
||||||
// never reads from them, and returns its URL.
|
// never reads from them, and returns its URL.
|
||||||
func startAppThatTakesNothing(t *testing.T) string {
|
func startAppThatTakesNothing(t *testing.T) string {
|
||||||
@@ -202,7 +184,6 @@ func TestAppTooSlowToAnswer(t *testing.T) {
|
|||||||
})
|
})
|
||||||
addr, out := startProxy(t, app.URL, map[string]string{
|
addr, out := startProxy(t, app.URL, map[string]string{
|
||||||
upstreamResponseTimeout: shortTimeoutSetting,
|
upstreamResponseTimeout: shortTimeoutSetting,
|
||||||
metricsToken: token,
|
|
||||||
})
|
})
|
||||||
|
|
||||||
start := time.Now()
|
start := time.Now()
|
||||||
@@ -218,8 +199,6 @@ func TestAppTooSlowToAnswer(t *testing.T) {
|
|||||||
t.Errorf("log line has upstream_status %v for an app that never answered",
|
t.Errorf("log line has upstream_status %v for an app that never answered",
|
||||||
line.fields["upstream_status"])
|
line.fields["upstream_status"])
|
||||||
}
|
}
|
||||||
|
|
||||||
wantLimitHits(t, addr, upstreamResponseTimeout, 1)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestAppTooSlowToFinishItsAnswer(t *testing.T) {
|
func TestAppTooSlowToFinishItsAnswer(t *testing.T) {
|
||||||
@@ -268,7 +247,6 @@ func TestClientTooSlowToTakeTheAnswer(t *testing.T) {
|
|||||||
})
|
})
|
||||||
addr, out := startProxy(t, app.URL, map[string]string{
|
addr, out := startProxy(t, app.URL, map[string]string{
|
||||||
clientResponseTimeout: shortTimeoutSetting,
|
clientResponseTimeout: shortTimeoutSetting,
|
||||||
metricsToken: token,
|
|
||||||
})
|
})
|
||||||
|
|
||||||
start := time.Now()
|
start := time.Now()
|
||||||
@@ -280,28 +258,4 @@ func TestClientTooSlowToTakeTheAnswer(t *testing.T) {
|
|||||||
line := out.requestLine(t)
|
line := out.requestLine(t)
|
||||||
wantTimedOut(t, start)
|
wantTimedOut(t, start)
|
||||||
wantLine(t, line, http.StatusOK, requestlog.ActionTimedOut)
|
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)
|
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -1,154 +0,0 @@
|
|||||||
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{
|
|
||||||
{Forwarded: true, Status: 200, RequestBytes: 10, ResponseBytes: 100},
|
|
||||||
{Forwarded: true, Status: 101},
|
|
||||||
{Forwarded: true, Status: 304, RequestBytes: 5},
|
|
||||||
{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),
|
|
||||||
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 TestLookupReachesTheHistoryOfAClientInTheTable(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
limiter := ratelimit.New(ratelimit.Limits{})
|
|
||||||
client := netip.MustParsePrefix("203.0.113.9/32")
|
|
||||||
other := netip.MustParsePrefix("198.51.100.7/32")
|
|
||||||
start := midnight()
|
|
||||||
|
|
||||||
limiter.AddToHistory(client, start, ratelimit.Request{Forwarded: true})
|
|
||||||
limiter.AddLookup(client, start, "AS64496", "Example Net", "DE")
|
|
||||||
|
|
||||||
// A later answer replaces it, and one for a client the table does not
|
|
||||||
// hold adds no client.
|
|
||||||
limiter.AddLookup(client, start.Add(time.Hour), "AS64497", "Other Net", "FR")
|
|
||||||
limiter.AddLookup(other, start, "AS64496", "Example Net", "DE")
|
|
||||||
|
|
||||||
want := ratelimit.History{
|
|
||||||
FirstSeen: start,
|
|
||||||
LastSeen: start,
|
|
||||||
ASN: "AS64497",
|
|
||||||
ASName: "Other Net",
|
|
||||||
Country: "FR",
|
|
||||||
LookedUp: start.Add(time.Hour),
|
|
||||||
Requests: 1,
|
|
||||||
Forwarded: 1,
|
|
||||||
}
|
|
||||||
|
|
||||||
got := historyOf(t, limiter, client)
|
|
||||||
if got != want {
|
|
||||||
t.Errorf("history\n%+v\nwant\n%+v", got, want)
|
|
||||||
}
|
|
||||||
|
|
||||||
if clients := limiter.Snapshot(); len(clients) != 1 {
|
|
||||||
t.Errorf("the table holds %+v, want %s alone", clients, client)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
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{}
|
|
||||||
}
|
|
||||||
+45
-332
@@ -1,15 +1,11 @@
|
|||||||
// Package ratelimit keeps the table of clients: each client's requests
|
// Package ratelimit counts each client's requests over a minute, an hour
|
||||||
// counted over a minute, an hour and a day, as the "Counting method"
|
// and a day, as the "Counting method" section of SPEC.md describes, and
|
||||||
// section of SPEC.md describes, which tell when a request takes the client
|
// tells when a request takes a client over a rate limit. The counts are
|
||||||
// over a rate limit, and each client's history since it was first seen.
|
// kept in memory only, for at most 20,000 clients.
|
||||||
// At most 20,000 clients are kept, in memory, and written to clients.json
|
|
||||||
// and read from it by the state package.
|
|
||||||
package ratelimit
|
package ratelimit
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"net/http"
|
|
||||||
"net/netip"
|
"net/netip"
|
||||||
"slices"
|
|
||||||
"sync"
|
"sync"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
@@ -17,8 +13,7 @@ import (
|
|||||||
)
|
)
|
||||||
|
|
||||||
// maxClients is how many clients are kept. Past it, the least recently
|
// 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
|
// seen client is dropped, and starts afresh if it comes back.
|
||||||
// back.
|
|
||||||
const maxClients = 20000
|
const maxClients = 20000
|
||||||
|
|
||||||
const day = 24 * time.Hour
|
const day = 24 * time.Hour
|
||||||
@@ -31,101 +26,20 @@ type Limits struct {
|
|||||||
PerDay int64
|
PerDay int64
|
||||||
}
|
}
|
||||||
|
|
||||||
// Limiter counts each client's requests against the limits, and keeps
|
// Limiter counts each client's requests against the limits. It is safe
|
||||||
// its history. It is safe for concurrent use.
|
// for concurrent use.
|
||||||
type Limiter struct {
|
type Limiter struct {
|
||||||
// windows are the minute, the hour and the day, in the order of
|
|
||||||
// Client.buckets.
|
|
||||||
windows [3]window
|
windows [3]window
|
||||||
|
|
||||||
mu sync.Mutex
|
mu sync.Mutex
|
||||||
clients *simplelru.LRU[netip.Prefix, *Client]
|
// clients holds each client's buckets, one pair for each of windows,
|
||||||
}
|
// in the same order.
|
||||||
|
clients *simplelru.LRU[netip.Prefix, *[3]buckets]
|
||||||
// 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"`
|
|
||||||
// ASN, ASName and Country are the client's AS number, AS name and
|
|
||||||
// country as last looked up, each empty when the lookup could not
|
|
||||||
// find it, and LookedUp is when the lookup gave that answer; all are
|
|
||||||
// empty while the client never was looked up.
|
|
||||||
ASN string `json:"asn,omitempty"`
|
|
||||||
ASName string `json:"as_name,omitempty"`
|
|
||||||
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 {
|
|
||||||
// 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.
|
// New returns a Limiter for limits, with no client counted yet.
|
||||||
func New(limits Limits) *Limiter {
|
func New(limits Limits) *Limiter {
|
||||||
clients, err := simplelru.NewLRU[netip.Prefix, *Client](maxClients, nil)
|
clients, err := simplelru.NewLRU[netip.Prefix, *[3]buckets](maxClients, nil)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
panic(err) // NewLRU fails only for a size below one
|
panic(err) // NewLRU fails only for a size below one
|
||||||
}
|
}
|
||||||
@@ -140,222 +54,30 @@ func New(limits Limits) *Limiter {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// 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
|
// 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
|
// not it is refused. It returns the window whose limit the request takes
|
||||||
// reports whether the request takes the client over a limit, and the
|
// the client over, "minute", "hour" or "day", the shortest if it is over
|
||||||
// window whose limit it goes over, the shortest if it is over several.
|
// several, or "" if it is within every limit.
|
||||||
func (l *Limiter) Count(client netip.Prefix, now time.Time) (Counts, Hit, bool) {
|
func (l *Limiter) Count(client netip.Prefix, now time.Time) string {
|
||||||
l.mu.Lock()
|
l.mu.Lock()
|
||||||
defer l.mu.Unlock()
|
defer l.mu.Unlock()
|
||||||
|
|
||||||
var (
|
counts, seen := l.clients.Get(client)
|
||||||
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
|
|
||||||
|
|
||||||
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++
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// AddLookup gives client's history its AS number, AS name and country, as
|
|
||||||
// the lookup gave them at lookedUp, if the table of clients holds the
|
|
||||||
// client.
|
|
||||||
// It does not make the client the most recently seen.
|
|
||||||
func (l *Limiter) AddLookup(
|
|
||||||
client netip.Prefix, lookedUp time.Time, asn, asName, country string,
|
|
||||||
) {
|
|
||||||
l.mu.Lock()
|
|
||||||
defer l.mu.Unlock()
|
|
||||||
|
|
||||||
c, held := l.clients.Peek(client)
|
|
||||||
if !held {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
h := &c.History
|
|
||||||
h.ASN, h.ASName, h.Country, h.LookedUp = asn, asName, country, lookedUp
|
|
||||||
}
|
|
||||||
|
|
||||||
// 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
|
|
||||||
}
|
|
||||||
|
|
||||||
// Client returns client as the table holds it, and whether it does. It is
|
|
||||||
// not a request from client, and leaves when it was last seen unchanged.
|
|
||||||
func (l *Limiter) Client(client netip.Prefix) (Client, bool) {
|
|
||||||
l.mu.Lock()
|
|
||||||
defer l.mu.Unlock()
|
|
||||||
|
|
||||||
c, seen := l.clients.Peek(client)
|
|
||||||
if !seen {
|
if !seen {
|
||||||
return Client{}, false
|
counts = &[3]buckets{}
|
||||||
|
l.clients.Add(client, counts)
|
||||||
}
|
}
|
||||||
|
|
||||||
return *c, true
|
limitHit := ""
|
||||||
}
|
|
||||||
|
|
||||||
// Len returns how many clients are in the table.
|
for i, w := range l.windows {
|
||||||
func (l *Limiter) Len() int {
|
requests := counts[i].add(now, w.length)
|
||||||
l.mu.Lock()
|
if limitHit == "" && w.limit > 0 && requests > float64(w.limit) {
|
||||||
defer l.mu.Unlock()
|
limitHit = w.name
|
||||||
|
|
||||||
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
|
return limitHit
|
||||||
}
|
|
||||||
|
|
||||||
// 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
|
// window is a length of time over which requests are counted, and the
|
||||||
@@ -366,6 +88,14 @@ type window struct {
|
|||||||
limit int64
|
limit int64
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// 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
|
||||||
|
current int64
|
||||||
|
previous int64
|
||||||
|
}
|
||||||
|
|
||||||
// add counts a request at now in a window of length, and returns the
|
// 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
|
// 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
|
// under way, and those in the bucket before it weighted by how much of
|
||||||
@@ -376,44 +106,27 @@ type window struct {
|
|||||||
// bucket. A request dated more than a second before it means the clock
|
// 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
|
// was set back, and the buckets start afresh: otherwise the bucket before
|
||||||
// would keep its full weight until the clock caught up.
|
// would keep its full weight until the clock caught up.
|
||||||
func (b *Buckets) add(now time.Time, length time.Duration) float64 {
|
func (b *buckets) add(now time.Time, length time.Duration) float64 {
|
||||||
if now.Before(b.Start.Add(-time.Second)) {
|
if now.Before(b.start.Add(-time.Second)) {
|
||||||
*b = Buckets{}
|
*b = buckets{}
|
||||||
}
|
}
|
||||||
|
|
||||||
start := now.Truncate(length)
|
start := now.Truncate(length)
|
||||||
if start.After(b.Start) {
|
if start.After(b.start) {
|
||||||
if start.Equal(b.Start.Add(length)) {
|
if start.Equal(b.start.Add(length)) {
|
||||||
b.Previous = b.Current
|
b.previous = b.current
|
||||||
} else {
|
} else {
|
||||||
b.Previous = 0
|
b.previous = 0
|
||||||
}
|
}
|
||||||
|
|
||||||
b.Start = start
|
b.start = start
|
||||||
b.Current = 0
|
b.current = 0
|
||||||
}
|
}
|
||||||
|
|
||||||
b.Current++
|
b.current++
|
||||||
|
|
||||||
elapsed := max(now.Sub(b.Start), 0)
|
elapsed := max(now.Sub(b.start), 0)
|
||||||
covered := 1 - float64(elapsed)/float64(length)
|
covered := 1 - float64(elapsed)/float64(length)
|
||||||
|
|
||||||
return float64(b.Previous)*covered + float64(b.Current)
|
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++
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -54,75 +54,6 @@ func TestEachWindowRefusesAtItsLimitAndLetsTheClientBack(t *testing.T) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
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) {
|
func TestClientBackAfterAWholeBucketIsWithinTheLimitAtOnce(t *testing.T) {
|
||||||
t.Parallel()
|
t.Parallel()
|
||||||
|
|
||||||
@@ -261,9 +192,9 @@ func wantCount(
|
|||||||
) {
|
) {
|
||||||
t.Helper()
|
t.Helper()
|
||||||
|
|
||||||
_, hit, _ := limiter.Count(client, now)
|
got := limiter.Count(client, now)
|
||||||
if hit.Window != want {
|
if got != want {
|
||||||
t.Errorf("request from %s at %s is over %q, want %q",
|
t.Errorf("request from %s at %s is over %q, want %q",
|
||||||
client, now.Format(time.RFC3339), hit.Window, want)
|
client, now.Format(time.RFC3339), got, want)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -1,122 +0,0 @@
|
|||||||
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)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@@ -1,310 +0,0 @@
|
|||||||
// Package remotelog sends the lines smallwebwaf writes on stdout to the
|
|
||||||
// remote log endpoint, SWWAF_LOG_REMOTE_URL, as the "Request log" section
|
|
||||||
// of SPEC.md describes: each line as the message of an RFC 5424 syslog
|
|
||||||
// record, over UDP, TCP or TLS. Lines wait in a bounded buffer, so a slow
|
|
||||||
// or unreachable endpoint never holds up a request or stdout.
|
|
||||||
package remotelog
|
|
||||||
|
|
||||||
import (
|
|
||||||
"bytes"
|
|
||||||
"context"
|
|
||||||
"crypto/tls"
|
|
||||||
"crypto/x509"
|
|
||||||
"errors"
|
|
||||||
"fmt"
|
|
||||||
"log/slog"
|
|
||||||
"net"
|
|
||||||
"net/url"
|
|
||||||
"os"
|
|
||||||
"strconv"
|
|
||||||
"sync/atomic"
|
|
||||||
"syscall"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
"sneak.berlin/go/smallwebwaf/internal/requestlog"
|
|
||||||
)
|
|
||||||
|
|
||||||
// The forms of SWWAF_LOG_REMOTE_URL, by its scheme.
|
|
||||||
const (
|
|
||||||
SchemeUDP = "syslog+udp"
|
|
||||||
SchemeTCP = "syslog+tcp"
|
|
||||||
SchemeTLS = "syslog+tls"
|
|
||||||
)
|
|
||||||
|
|
||||||
// A record's priority is the number of its facility times the number of
|
|
||||||
// severities there are, plus the number of its severity. Every record's
|
|
||||||
// severity is informational.
|
|
||||||
const (
|
|
||||||
severities = 8
|
|
||||||
informational = 6
|
|
||||||
)
|
|
||||||
|
|
||||||
const (
|
|
||||||
// dialTimeout bounds connecting to the endpoint, the TLS handshake
|
|
||||||
// included.
|
|
||||||
dialTimeout = 10 * time.Second
|
|
||||||
// After a failed attempt to connect, or a connection on which a record
|
|
||||||
// fails, the next attempt to connect is made a second later, and
|
|
||||||
// retryDelayFactor times as long after each further failure in a row,
|
|
||||||
// up to a minute. A connection that fails after it has stayed up for
|
|
||||||
// resetRetryDelayAfter ends the row.
|
|
||||||
firstRetryDelay = time.Second
|
|
||||||
retryDelayFactor = 2
|
|
||||||
maxRetryDelay = time.Minute
|
|
||||||
resetRetryDelayAfter = time.Minute
|
|
||||||
)
|
|
||||||
|
|
||||||
// Params are what New needs.
|
|
||||||
type Params struct {
|
|
||||||
// URL is the endpoint (SWWAF_LOG_REMOTE_URL): SchemeUDP, SchemeTCP or
|
|
||||||
// SchemeTLS, a host and a port.
|
|
||||||
URL *url.URL
|
|
||||||
// RootCAs are the certificates a SchemeTLS endpoint's certificate
|
|
||||||
// must chain to (SWWAF_LOG_REMOTE_TLS_CA_FILE), nil for the host's.
|
|
||||||
RootCAs *x509.CertPool
|
|
||||||
// Buffer is the most lines held while they wait to be sent
|
|
||||||
// (SWWAF_LOG_REMOTE_BUFFER).
|
|
||||||
Buffer int
|
|
||||||
// Facility is the number of the records' syslog facility
|
|
||||||
// (SWWAF_LOG_REMOTE_FACILITY), and AppName their APP-NAME
|
|
||||||
// (SWWAF_LOG_REMOTE_APP_NAME).
|
|
||||||
Facility int
|
|
||||||
AppName string
|
|
||||||
}
|
|
||||||
|
|
||||||
// Sender sends lines to the endpoint. Write puts them in its buffer, and
|
|
||||||
// Run sends them from there.
|
|
||||||
type Sender struct {
|
|
||||||
url *url.URL
|
|
||||||
tlsConfig *tls.Config
|
|
||||||
// beforeTime and afterTime are the parts of every record's header
|
|
||||||
// before and after its time, as RFC 5424 lays the header out.
|
|
||||||
beforeTime string
|
|
||||||
afterTime string
|
|
||||||
// records is the buffer: each line's record, framed to be sent.
|
|
||||||
records chan []byte
|
|
||||||
sent atomic.Int64
|
|
||||||
dropped atomic.Int64
|
|
||||||
}
|
|
||||||
|
|
||||||
// New returns a Sender for the endpoint params.URL.
|
|
||||||
func New(params Params) *Sender {
|
|
||||||
hostname, err := os.Hostname()
|
|
||||||
if err != nil || hostname == "" {
|
|
||||||
hostname = "-" // RFC 5424's value for a field that has none
|
|
||||||
}
|
|
||||||
|
|
||||||
priority := params.Facility*severities + informational
|
|
||||||
|
|
||||||
return &Sender{
|
|
||||||
url: params.URL,
|
|
||||||
tlsConfig: &tls.Config{
|
|
||||||
RootCAs: params.RootCAs,
|
|
||||||
MinVersion: tls.VersionTLS12,
|
|
||||||
},
|
|
||||||
// The 1 is the version of the format. The process id, the message
|
|
||||||
// id and the structured data have no value.
|
|
||||||
beforeTime: "<" + strconv.Itoa(priority) + ">1 ",
|
|
||||||
afterTime: " " + hostname + " " + params.AppName + " - - - ",
|
|
||||||
records: make(chan []byte, params.Buffer),
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// Write puts each line in p in the buffer, as the message of a record of
|
|
||||||
// its own, and never waits: when the buffer is full, the oldest record in
|
|
||||||
// it is dropped to make room. It is safe for concurrent use.
|
|
||||||
func (s *Sender) Write(p []byte) (int, error) {
|
|
||||||
at := requestlog.FormatTime(time.Now())
|
|
||||||
|
|
||||||
for line := range bytes.Lines(p) {
|
|
||||||
line = bytes.TrimSuffix(line, []byte("\n"))
|
|
||||||
if len(line) > 0 {
|
|
||||||
s.put(s.record(at, line))
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
return len(p), nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// Sent is how many records have been sent.
|
|
||||||
func (s *Sender) Sent() int64 {
|
|
||||||
return s.sent.Load()
|
|
||||||
}
|
|
||||||
|
|
||||||
// Dropped is how many records were dropped: the oldest in a full buffer,
|
|
||||||
// and those whose sending failed.
|
|
||||||
func (s *Sender) Dropped() int64 {
|
|
||||||
return s.dropped.Load()
|
|
||||||
}
|
|
||||||
|
|
||||||
// Depth is how many records are in the buffer.
|
|
||||||
func (s *Sender) Depth() int {
|
|
||||||
return len(s.records)
|
|
||||||
}
|
|
||||||
|
|
||||||
// Run connects to the endpoint and sends each record as it comes into the
|
|
||||||
// buffer, until ctx is done. Then it sends the records still in the buffer,
|
|
||||||
// on the connection open at that time or, if there is none, on a new one,
|
|
||||||
// until none is left or one fails, and returns. How long it may take over
|
|
||||||
// that is for the caller to bound.
|
|
||||||
//
|
|
||||||
// A connection on which a record fails is closed and the record dropped.
|
|
||||||
// That failure, like a failed attempt to connect, is logged to processLog
|
|
||||||
// and followed by the next attempt after firstRetryDelay, retryDelayFactor
|
|
||||||
// times as long after each further failure in a row up to maxRetryDelay,
|
|
||||||
// and firstRetryDelay again after a connection that stayed up for
|
|
||||||
// resetRetryDelayAfter. Meanwhile the records wait in the buffer.
|
|
||||||
func (s *Sender) Run(ctx context.Context, processLog *slog.Logger) {
|
|
||||||
conn := s.send(ctx, processLog)
|
|
||||||
if conn == nil && len(s.records) > 0 {
|
|
||||||
conn, _ = s.dial(context.WithoutCancel(ctx))
|
|
||||||
}
|
|
||||||
|
|
||||||
if conn == nil {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
defer func() {
|
|
||||||
_ = conn.Close()
|
|
||||||
}()
|
|
||||||
|
|
||||||
for {
|
|
||||||
select {
|
|
||||||
case record := <-s.records:
|
|
||||||
if s.write(conn, record) != nil {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
default:
|
|
||||||
return
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// record returns line as an RFC 5424 record made at the time at, framed
|
|
||||||
// for the endpoint: on its own over UDP, since each datagram holds one,
|
|
||||||
// and over TCP and TLS after its length in bytes and a space, the
|
|
||||||
// octet-counted framing of RFC 6587 and RFC 5425.
|
|
||||||
func (s *Sender) record(at string, line []byte) []byte {
|
|
||||||
record := make([]byte, 0, len(s.beforeTime)+len(at)+len(s.afterTime)+len(line))
|
|
||||||
record = append(record, s.beforeTime...)
|
|
||||||
record = append(record, at...)
|
|
||||||
record = append(record, s.afterTime...)
|
|
||||||
record = append(record, line...)
|
|
||||||
|
|
||||||
if s.url.Scheme == SchemeUDP {
|
|
||||||
return record
|
|
||||||
}
|
|
||||||
|
|
||||||
return append([]byte(strconv.Itoa(len(record))+" "), record...)
|
|
||||||
}
|
|
||||||
|
|
||||||
// put adds record to the buffer, first dropping the oldest record in it
|
|
||||||
// while it is full.
|
|
||||||
func (s *Sender) put(record []byte) {
|
|
||||||
for {
|
|
||||||
select {
|
|
||||||
case s.records <- record:
|
|
||||||
return
|
|
||||||
default:
|
|
||||||
}
|
|
||||||
|
|
||||||
select {
|
|
||||||
case <-s.records:
|
|
||||||
s.dropped.Add(1)
|
|
||||||
default:
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// send connects to the endpoint and sends each record as it comes into
|
|
||||||
// the buffer, until ctx is done, and returns the connection then open, or
|
|
||||||
// nil.
|
|
||||||
func (s *Sender) send(ctx context.Context, processLog *slog.Logger) net.Conn {
|
|
||||||
delay := firstRetryDelay
|
|
||||||
|
|
||||||
for {
|
|
||||||
conn, err := s.dial(ctx)
|
|
||||||
if ctx.Err() != nil {
|
|
||||||
return conn
|
|
||||||
}
|
|
||||||
|
|
||||||
if err == nil {
|
|
||||||
connected := time.Now()
|
|
||||||
|
|
||||||
err = s.sendOn(ctx, conn)
|
|
||||||
if err == nil {
|
|
||||||
return conn
|
|
||||||
}
|
|
||||||
|
|
||||||
_ = conn.Close()
|
|
||||||
|
|
||||||
if time.Since(connected) >= resetRetryDelayAfter {
|
|
||||||
delay = firstRetryDelay
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
processLog.Warn("sending to SWWAF_LOG_REMOTE_URL failed",
|
|
||||||
"error", err.Error(), "connecting_again_in", delay.String())
|
|
||||||
|
|
||||||
select {
|
|
||||||
case <-time.After(delay):
|
|
||||||
case <-ctx.Done():
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
delay = min(retryDelayFactor*delay, maxRetryDelay)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// sendOn sends each record on conn as it comes into the buffer, until one
|
|
||||||
// fails, whose error it returns, or ctx is done.
|
|
||||||
func (s *Sender) sendOn(ctx context.Context, conn net.Conn) error {
|
|
||||||
for {
|
|
||||||
select {
|
|
||||||
case record := <-s.records:
|
|
||||||
err := s.write(conn, record)
|
|
||||||
if err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
case <-ctx.Done():
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// write sends record on conn, and counts it as sent or, if that fails,
|
|
||||||
// as dropped. A record too long for one UDP datagram is dropped without
|
|
||||||
// an error, since the connection has not failed: a long request must not
|
|
||||||
// hold up the lines after it.
|
|
||||||
func (s *Sender) write(conn net.Conn, record []byte) error {
|
|
||||||
_, err := conn.Write(record)
|
|
||||||
if err != nil {
|
|
||||||
s.dropped.Add(1)
|
|
||||||
|
|
||||||
if errors.Is(err, syscall.EMSGSIZE) {
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
return fmt.Errorf("send a record: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
s.sent.Add(1)
|
|
||||||
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// dial connects to the endpoint.
|
|
||||||
func (s *Sender) dial(ctx context.Context) (net.Conn, error) {
|
|
||||||
dialer := &net.Dialer{Timeout: dialTimeout}
|
|
||||||
|
|
||||||
switch s.url.Scheme {
|
|
||||||
case SchemeUDP:
|
|
||||||
return dialer.DialContext(ctx, "udp", s.url.Host)
|
|
||||||
case SchemeTLS:
|
|
||||||
tlsDialer := &tls.Dialer{NetDialer: dialer, Config: s.tlsConfig}
|
|
||||||
|
|
||||||
return tlsDialer.DialContext(ctx, "tcp", s.url.Host)
|
|
||||||
default:
|
|
||||||
return dialer.DialContext(ctx, "tcp", s.url.Host)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@@ -1,640 +0,0 @@
|
|||||||
package remotelog_test
|
|
||||||
|
|
||||||
import (
|
|
||||||
"bufio"
|
|
||||||
"bytes"
|
|
||||||
"context"
|
|
||||||
"crypto/ecdsa"
|
|
||||||
"crypto/elliptic"
|
|
||||||
"crypto/rand"
|
|
||||||
"crypto/tls"
|
|
||||||
"crypto/x509"
|
|
||||||
"crypto/x509/pkix"
|
|
||||||
"encoding/json"
|
|
||||||
"fmt"
|
|
||||||
"io"
|
|
||||||
"log/slog"
|
|
||||||
"math/big"
|
|
||||||
"net"
|
|
||||||
"net/url"
|
|
||||||
"os"
|
|
||||||
"slices"
|
|
||||||
"strconv"
|
|
||||||
"strings"
|
|
||||||
"sync"
|
|
||||||
"sync/atomic"
|
|
||||||
"testing"
|
|
||||||
"testing/synctest"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
"sneak.berlin/go/smallwebwaf/internal/remotelog"
|
|
||||||
)
|
|
||||||
|
|
||||||
// The tests run in a synctest bubble, where the time package runs on a
|
|
||||||
// clock of the test's own, which starts at 2000-01-01T00:00:00Z: a wait
|
|
||||||
// lasts exactly as long as it should, however slowly the test process
|
|
||||||
// runs, and synctest.Wait returns once the sender has done all it can
|
|
||||||
// before time passes. The endpoint is a listener on the loopback address.
|
|
||||||
// A test reads from it only once the records are on their way, and checks
|
|
||||||
// the sender's counts first, since a goroutine of the bubble that waits on
|
|
||||||
// the network keeps that clock from moving on. For the same reason the
|
|
||||||
// endpoint that refuses connections, a tlsEndpoint, runs outside the
|
|
||||||
// bubble: a sender connecting over TLS waits on the endpoint's answer.
|
|
||||||
|
|
||||||
const (
|
|
||||||
// started is the time a record made as a test starts gives.
|
|
||||||
started = "2000-01-01T00:00:00.000Z"
|
|
||||||
appName = "fsn1app1/gitea"
|
|
||||||
// local0 is the number of the default facility, and local0Info the
|
|
||||||
// priority of its records.
|
|
||||||
local0 = 16
|
|
||||||
local0Info = "<134>"
|
|
||||||
// loopback is where the endpoints listen.
|
|
||||||
loopback = "127.0.0.1:0"
|
|
||||||
)
|
|
||||||
|
|
||||||
func TestRecordsOverUDPGoOnePerDatagram(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
synctest.Test(t, func(t *testing.T) {
|
|
||||||
endpoint, err := (&net.ListenConfig{}).ListenPacket(t.Context(), "udp", loopback)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("listen: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
t.Cleanup(func() { _ = endpoint.Close() })
|
|
||||||
|
|
||||||
sender, _, _ := run(t, params(remotelog.SchemeUDP, endpoint.LocalAddr()))
|
|
||||||
|
|
||||||
_, _ = sender.Write([]byte(`{"type":"request"}` + "\n" + `{"type":"process"}` + "\n"))
|
|
||||||
|
|
||||||
synctest.Wait()
|
|
||||||
wantCounts(t, sender, 2, 0, 0)
|
|
||||||
|
|
||||||
for _, line := range []string{`{"type":"request"}`, `{"type":"process"}`} {
|
|
||||||
datagram := make([]byte, 1024)
|
|
||||||
|
|
||||||
n, _, err := endpoint.ReadFrom(datagram)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("read: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
want := record(t, local0Info, appName, line)
|
|
||||||
if string(datagram[:n]) != want {
|
|
||||||
t.Errorf("datagram %q, want %q", datagram[:n], want)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestRecordsOverTCPAreOctetCountedWithTheirFacility(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
synctest.Test(t, func(t *testing.T) {
|
|
||||||
endpoint := listen(t)
|
|
||||||
endpointParams := params(remotelog.SchemeTCP, endpoint.Addr())
|
|
||||||
endpointParams.Facility = 19 // local3
|
|
||||||
endpointParams.AppName = "gitea"
|
|
||||||
sender, _, _ := run(t, endpointParams)
|
|
||||||
|
|
||||||
_, _ = sender.Write([]byte("first\nsecond\n"))
|
|
||||||
|
|
||||||
synctest.Wait()
|
|
||||||
wantCounts(t, sender, 2, 0, 0)
|
|
||||||
|
|
||||||
frames := bufio.NewReader(accept(t, endpoint))
|
|
||||||
wantFrame(t, frames, record(t, "<158>", "gitea", "first"))
|
|
||||||
wantFrame(t, frames, record(t, "<158>", "gitea", "second"))
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestStalledEndpointHoldsUpNoWriteAndOldestRecordsAreDropped(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
synctest.Test(t, func(t *testing.T) {
|
|
||||||
certificate, roots := testCertificate(t)
|
|
||||||
endpoint := listen(t)
|
|
||||||
endpointParams := params(remotelog.SchemeTLS, endpoint.Addr())
|
|
||||||
endpointParams.RootCAs = roots
|
|
||||||
endpointParams.Buffer = 3
|
|
||||||
sender, _, _ := run(t, endpointParams)
|
|
||||||
|
|
||||||
// The sender connects, and its TLS handshake waits for an answer
|
|
||||||
// the endpoint does not give yet.
|
|
||||||
conn := accept(t, endpoint)
|
|
||||||
|
|
||||||
var stdout bytes.Buffer
|
|
||||||
|
|
||||||
out := io.MultiWriter(&stdout, sender)
|
|
||||||
for i := range 5 {
|
|
||||||
_, _ = fmt.Fprintf(out, "line %d\n", i+1)
|
|
||||||
}
|
|
||||||
|
|
||||||
if stdout.String() != "line 1\nline 2\nline 3\nline 4\nline 5\n" {
|
|
||||||
t.Errorf("stdout has %q", stdout.String())
|
|
||||||
}
|
|
||||||
|
|
||||||
wantCounts(t, sender, 0, 2, 3)
|
|
||||||
|
|
||||||
// Once the endpoint answers, the three newest records are sent.
|
|
||||||
server := tls.Server(conn, &tls.Config{
|
|
||||||
Certificates: []tls.Certificate{certificate},
|
|
||||||
MinVersion: tls.VersionTLS12,
|
|
||||||
})
|
|
||||||
|
|
||||||
err := server.HandshakeContext(t.Context())
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("handshake: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
synctest.Wait()
|
|
||||||
wantCounts(t, sender, 3, 2, 0)
|
|
||||||
|
|
||||||
frames := bufio.NewReader(server)
|
|
||||||
for _, line := range []string{"line 3", "line 4", "line 5"} {
|
|
||||||
wantFrame(t, frames, record(t, local0Info, appName, line))
|
|
||||||
}
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestReconnectsWithBackoffAfterTheEndpointGoesAway(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
certificate, roots := testCertificate(t)
|
|
||||||
endpoint := startTLSEndpoint(t, certificate)
|
|
||||||
|
|
||||||
synctest.Test(t, func(t *testing.T) {
|
|
||||||
endpointParams := params(remotelog.SchemeTLS, endpoint.addr)
|
|
||||||
endpointParams.RootCAs = roots
|
|
||||||
sender, logged, _ := run(t, endpointParams)
|
|
||||||
|
|
||||||
_, _ = sender.Write([]byte("one\n"))
|
|
||||||
|
|
||||||
synctest.Wait()
|
|
||||||
wantCounts(t, sender, 1, 0, 0)
|
|
||||||
|
|
||||||
conn := endpoint.next(t)
|
|
||||||
wantFrame(t, bufio.NewReader(conn), record(t, local0Info, appName, "one"))
|
|
||||||
|
|
||||||
// The endpoint goes away: it closes the connection, and refuses the
|
|
||||||
// next ones. The sender notices when a record fails, and tries to
|
|
||||||
// connect again a second later, then two seconds after that.
|
|
||||||
endpoint.refusing.Store(true)
|
|
||||||
|
|
||||||
_ = conn.Close()
|
|
||||||
|
|
||||||
writeUntilDropped(t, sender, 1)
|
|
||||||
sent := sender.Sent()
|
|
||||||
|
|
||||||
_, _ = sender.Write([]byte("two\n"))
|
|
||||||
|
|
||||||
time.Sleep(time.Second)
|
|
||||||
synctest.Wait()
|
|
||||||
|
|
||||||
endpoint.refusing.Store(false)
|
|
||||||
|
|
||||||
time.Sleep(2*time.Second - time.Nanosecond)
|
|
||||||
synctest.Wait()
|
|
||||||
wantCounts(t, sender, sent, 1, 1)
|
|
||||||
|
|
||||||
// The endpoint is back, and the record waiting is sent.
|
|
||||||
time.Sleep(time.Nanosecond)
|
|
||||||
synctest.Wait()
|
|
||||||
wantCounts(t, sender, sent+1, 1, 0)
|
|
||||||
|
|
||||||
conn = endpoint.next(t)
|
|
||||||
wantFrame(t, bufio.NewReader(conn), record(t, local0Info, appName, "two"))
|
|
||||||
wantRetries(t, logged, "1s", "2s")
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestAConnectionClosedAtOnceIsMadeAgainAfterAGrowingDelay(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
synctest.Test(t, func(t *testing.T) {
|
|
||||||
endpoint := listen(t)
|
|
||||||
sender, logged, _ := run(t, params(remotelog.SchemeTCP, endpoint.Addr()))
|
|
||||||
|
|
||||||
// The endpoint closes each connection as soon as it takes it. The
|
|
||||||
// sender notices when a record fails, and connects again a second
|
|
||||||
// later, then two seconds after that, then four.
|
|
||||||
delays := []time.Duration{time.Second, 2 * time.Second, 4 * time.Second}
|
|
||||||
for i, delay := range delays {
|
|
||||||
_ = accept(t, endpoint).Close()
|
|
||||||
|
|
||||||
writeUntilDropped(t, sender, int64(i+1))
|
|
||||||
wantConnectedAgainAfter(t, sender, delay)
|
|
||||||
}
|
|
||||||
|
|
||||||
wantRetries(t, logged, "1s", "2s", "4s")
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestTheDelayStartsAgainAfterAConnectionThatStayedUpAMinute(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
synctest.Test(t, func(t *testing.T) {
|
|
||||||
endpoint := listen(t)
|
|
||||||
sender, logged, _ := run(t, params(remotelog.SchemeTCP, endpoint.Addr()))
|
|
||||||
|
|
||||||
_ = accept(t, endpoint).Close()
|
|
||||||
|
|
||||||
writeUntilDropped(t, sender, 1)
|
|
||||||
wantConnectedAgainAfter(t, sender, time.Second)
|
|
||||||
|
|
||||||
// A connection that fails just short of a minute after it was made
|
|
||||||
// leaves the delay growing.
|
|
||||||
conn := accept(t, endpoint)
|
|
||||||
|
|
||||||
time.Sleep(time.Minute - time.Nanosecond)
|
|
||||||
|
|
||||||
_ = conn.Close()
|
|
||||||
|
|
||||||
writeUntilDropped(t, sender, 2)
|
|
||||||
wantConnectedAgainAfter(t, sender, 2*time.Second)
|
|
||||||
|
|
||||||
// One that fails a minute after it was made starts it again from a
|
|
||||||
// second.
|
|
||||||
conn = accept(t, endpoint)
|
|
||||||
|
|
||||||
time.Sleep(time.Minute)
|
|
||||||
|
|
||||||
_ = conn.Close()
|
|
||||||
|
|
||||||
writeUntilDropped(t, sender, 3)
|
|
||||||
wantConnectedAgainAfter(t, sender, time.Second)
|
|
||||||
wantRetries(t, logged, "1s", "2s", "1s")
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestALineTooLongForADatagramIsDroppedAlone(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
synctest.Test(t, func(t *testing.T) {
|
|
||||||
endpoint, err := (&net.ListenConfig{}).ListenPacket(t.Context(), "udp", loopback)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("listen: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
t.Cleanup(func() { _ = endpoint.Close() })
|
|
||||||
|
|
||||||
sender, logged, _ := run(t, params(remotelog.SchemeUDP, endpoint.LocalAddr()))
|
|
||||||
|
|
||||||
// With its header, the first line's record is longer than the 65507
|
|
||||||
// bytes a UDP datagram over IPv4 holds. It is dropped, nothing is
|
|
||||||
// logged, and the next line is sent at once.
|
|
||||||
_, _ = sender.Write([]byte(strings.Repeat("x", 65507) + "\nnext\n"))
|
|
||||||
|
|
||||||
synctest.Wait()
|
|
||||||
wantCounts(t, sender, 1, 1, 0)
|
|
||||||
wantRetries(t, logged)
|
|
||||||
|
|
||||||
datagram := make([]byte, 1024)
|
|
||||||
|
|
||||||
n, _, err := endpoint.ReadFrom(datagram)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("read: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
want := record(t, local0Info, appName, "next")
|
|
||||||
if string(datagram[:n]) != want {
|
|
||||||
t.Errorf("datagram %q, want %q", datagram[:n], want)
|
|
||||||
}
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestRecordsWaitingAtTheStopAreSent(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
certificate, roots := testCertificate(t)
|
|
||||||
endpoint := startTLSEndpoint(t, certificate)
|
|
||||||
|
|
||||||
synctest.Test(t, func(t *testing.T) {
|
|
||||||
// The endpoint refuses the sender's first connection: it fails to
|
|
||||||
// connect, and waits a second to try again.
|
|
||||||
endpoint.refusing.Store(true)
|
|
||||||
|
|
||||||
endpointParams := params(remotelog.SchemeTLS, endpoint.addr)
|
|
||||||
endpointParams.RootCAs = roots
|
|
||||||
sender, logged, stop := run(t, endpointParams)
|
|
||||||
|
|
||||||
synctest.Wait()
|
|
||||||
wantRetries(t, logged, "1s")
|
|
||||||
|
|
||||||
_, _ = sender.Write([]byte("one\ntwo\n"))
|
|
||||||
|
|
||||||
endpoint.refusing.Store(false)
|
|
||||||
|
|
||||||
// Stopped before that second is over, it connects to send them.
|
|
||||||
stop()
|
|
||||||
wantCounts(t, sender, 2, 0, 0)
|
|
||||||
|
|
||||||
frames := bufio.NewReader(endpoint.next(t))
|
|
||||||
wantFrame(t, frames, record(t, local0Info, appName, "one"))
|
|
||||||
wantFrame(t, frames, record(t, local0Info, appName, "two"))
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
// output collects what the sender logs.
|
|
||||||
type output struct {
|
|
||||||
mu sync.Mutex
|
|
||||||
buf bytes.Buffer
|
|
||||||
}
|
|
||||||
|
|
||||||
// Write adds lines the sender logs.
|
|
||||||
func (o *output) Write(p []byte) (int, error) {
|
|
||||||
o.mu.Lock()
|
|
||||||
defer o.mu.Unlock()
|
|
||||||
|
|
||||||
return o.buf.Write(p)
|
|
||||||
}
|
|
||||||
|
|
||||||
// text returns everything logged so far.
|
|
||||||
func (o *output) text() string {
|
|
||||||
o.mu.Lock()
|
|
||||||
defer o.mu.Unlock()
|
|
||||||
|
|
||||||
return o.buf.String()
|
|
||||||
}
|
|
||||||
|
|
||||||
// params returns the settings of a Sender for the endpoint at addr, in
|
|
||||||
// the form scheme names: room for ten lines, the default facility, and
|
|
||||||
// appName.
|
|
||||||
func params(scheme string, addr net.Addr) remotelog.Params {
|
|
||||||
return remotelog.Params{
|
|
||||||
URL: &url.URL{Scheme: scheme, Host: addr.String()},
|
|
||||||
Buffer: 10,
|
|
||||||
Facility: local0,
|
|
||||||
AppName: appName,
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// run runs a Sender with settings until the test ends or the function
|
|
||||||
// it returns is called, which waits for Run to return. It returns the
|
|
||||||
// Sender, and what it logs.
|
|
||||||
func run(t *testing.T, settings remotelog.Params) (*remotelog.Sender, *output, func()) {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
sender := remotelog.New(settings)
|
|
||||||
logged := &output{}
|
|
||||||
ctx, cancel := context.WithCancel(t.Context())
|
|
||||||
ran := make(chan struct{})
|
|
||||||
|
|
||||||
go func() {
|
|
||||||
sender.Run(ctx, slog.New(slog.NewJSONHandler(logged, nil)))
|
|
||||||
close(ran)
|
|
||||||
}()
|
|
||||||
|
|
||||||
stop := func() {
|
|
||||||
cancel()
|
|
||||||
<-ran
|
|
||||||
}
|
|
||||||
t.Cleanup(stop)
|
|
||||||
|
|
||||||
return sender, logged, stop
|
|
||||||
}
|
|
||||||
|
|
||||||
// listen returns a TCP listener on the loopback address, closed when the
|
|
||||||
// test ends.
|
|
||||||
func listen(t *testing.T) net.Listener {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
listener, err := (&net.ListenConfig{}).Listen(t.Context(), "tcp", loopback)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("listen: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
t.Cleanup(func() { _ = listener.Close() })
|
|
||||||
|
|
||||||
return listener
|
|
||||||
}
|
|
||||||
|
|
||||||
// accept returns the next connection to listener, closed when the test
|
|
||||||
// ends.
|
|
||||||
func accept(t *testing.T, listener net.Listener) net.Conn {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
conn, err := listener.Accept()
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("accept: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
t.Cleanup(func() { _ = conn.Close() })
|
|
||||||
|
|
||||||
return conn
|
|
||||||
}
|
|
||||||
|
|
||||||
// tlsEndpoint is a syslog+tls endpoint on the loopback address, which a
|
|
||||||
// test starts outside its bubble. It keeps its listener until the test
|
|
||||||
// ends, and either takes each connection or refuses it.
|
|
||||||
type tlsEndpoint struct {
|
|
||||||
addr net.Addr
|
|
||||||
// refusing is set while the endpoint closes each connection before the
|
|
||||||
// TLS handshake, which fails the sender's attempt to connect.
|
|
||||||
refusing atomic.Bool
|
|
||||||
// conns are the connections it has taken, after the handshake.
|
|
||||||
conns chan net.Conn
|
|
||||||
}
|
|
||||||
|
|
||||||
// startTLSEndpoint starts a tlsEndpoint with certificate, which takes
|
|
||||||
// connections until it is told to refuse them.
|
|
||||||
func startTLSEndpoint(t *testing.T, certificate tls.Certificate) *tlsEndpoint {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
listener := listen(t)
|
|
||||||
endpoint := &tlsEndpoint{addr: listener.Addr(), conns: make(chan net.Conn, 10)}
|
|
||||||
config := &tls.Config{
|
|
||||||
Certificates: []tls.Certificate{certificate},
|
|
||||||
MinVersion: tls.VersionTLS12,
|
|
||||||
}
|
|
||||||
|
|
||||||
go func() {
|
|
||||||
for {
|
|
||||||
conn, err := listener.Accept()
|
|
||||||
if err != nil {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
server := tls.Server(conn, config)
|
|
||||||
if endpoint.refusing.Load() || server.HandshakeContext(t.Context()) != nil {
|
|
||||||
_ = conn.Close()
|
|
||||||
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
|
|
||||||
endpoint.conns <- server
|
|
||||||
}
|
|
||||||
}()
|
|
||||||
|
|
||||||
return endpoint
|
|
||||||
}
|
|
||||||
|
|
||||||
// next returns the next connection the endpoint has taken, closed when
|
|
||||||
// the test ends.
|
|
||||||
func (e *tlsEndpoint) next(t *testing.T) net.Conn {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
conn := <-e.conns
|
|
||||||
|
|
||||||
t.Cleanup(func() { _ = conn.Close() })
|
|
||||||
|
|
||||||
return conn
|
|
||||||
}
|
|
||||||
|
|
||||||
// record returns the record of line made as the test started, with the
|
|
||||||
// priority and the app name given.
|
|
||||||
func record(t *testing.T, priority, app, line string) string {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
hostname, err := os.Hostname()
|
|
||||||
if err != nil || hostname == "" {
|
|
||||||
hostname = "-"
|
|
||||||
}
|
|
||||||
|
|
||||||
return priority + "1 " + started + " " + hostname + " " + app + " - - - " + line
|
|
||||||
}
|
|
||||||
|
|
||||||
// wantFrame reads the next octet-counted frame from frames, and checks
|
|
||||||
// that it holds want.
|
|
||||||
func wantFrame(t *testing.T, frames *bufio.Reader, want string) {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
count, err := frames.ReadString(' ')
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("read a frame's length: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
length, err := strconv.Atoi(strings.TrimSuffix(count, " "))
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("frame starts %q, not with its length", count)
|
|
||||||
}
|
|
||||||
|
|
||||||
got := make([]byte, length)
|
|
||||||
|
|
||||||
_, err = io.ReadFull(frames, got)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("read a frame: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
if string(got) != want {
|
|
||||||
t.Errorf("frame %q, want %q", got, want)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// wantCounts checks the records sender has sent, dropped and holds in
|
|
||||||
// its buffer.
|
|
||||||
func wantCounts(
|
|
||||||
t *testing.T, sender *remotelog.Sender, sent, dropped int64, depth int,
|
|
||||||
) {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
if sender.Sent() != sent || sender.Dropped() != dropped || sender.Depth() != depth {
|
|
||||||
t.Fatalf("sent %d, dropped %d, %d in the buffer; want %d, %d and %d",
|
|
||||||
sender.Sent(), sender.Dropped(), sender.Depth(), sent, dropped, depth)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// writeUntilDropped writes a line at a time until the count of records
|
|
||||||
// sender has dropped reaches dropped. The records it sends on a
|
|
||||||
// connection the endpoint has closed are lost before one fails; how many
|
|
||||||
// depends on when the endpoint's host answers that the connection is
|
|
||||||
// gone.
|
|
||||||
func writeUntilDropped(t *testing.T, sender *remotelog.Sender, dropped int64) {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
for sender.Dropped() < dropped {
|
|
||||||
_, _ = sender.Write([]byte("lost\n"))
|
|
||||||
|
|
||||||
synctest.Wait()
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// wantConnectedAgainAfter writes a line while the sender waits to connect
|
|
||||||
// again, and checks that it connects, and takes the line from the buffer,
|
|
||||||
// only once delay is over.
|
|
||||||
func wantConnectedAgainAfter(
|
|
||||||
t *testing.T, sender *remotelog.Sender, delay time.Duration,
|
|
||||||
) {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
_, _ = sender.Write([]byte("waiting\n"))
|
|
||||||
|
|
||||||
time.Sleep(delay - time.Nanosecond)
|
|
||||||
synctest.Wait()
|
|
||||||
|
|
||||||
if sender.Depth() != 1 {
|
|
||||||
t.Fatalf("connected again before %v", delay)
|
|
||||||
}
|
|
||||||
|
|
||||||
time.Sleep(time.Nanosecond)
|
|
||||||
synctest.Wait()
|
|
||||||
|
|
||||||
if sender.Depth() != 0 {
|
|
||||||
t.Fatalf("not connected again after %v", delay)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// wantRetries checks that the sender logged a failure, of an attempt to
|
|
||||||
// connect or of a connection, for each of delays, the time until the next
|
|
||||||
// attempt, in order, and logged nothing else.
|
|
||||||
func wantRetries(t *testing.T, logged *output, delays ...string) {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
var got []string
|
|
||||||
|
|
||||||
for line := range strings.Lines(logged.text()) {
|
|
||||||
var fields map[string]any
|
|
||||||
|
|
||||||
err := json.Unmarshal([]byte(line), &fields)
|
|
||||||
if err != nil || fields["msg"] != "sending to SWWAF_LOG_REMOTE_URL failed" {
|
|
||||||
t.Fatalf("logged %q", line)
|
|
||||||
}
|
|
||||||
|
|
||||||
delay, _ := fields["connecting_again_in"].(string)
|
|
||||||
got = append(got, delay)
|
|
||||||
}
|
|
||||||
|
|
||||||
if !slices.Equal(got, delays) {
|
|
||||||
t.Errorf("logged failures to connect again in %v, want %v", got, delays)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// testCertificate returns a certificate for 127.0.0.1 that is its own
|
|
||||||
// CA, and a pool that holds it. It is valid on the bubble's clock, which
|
|
||||||
// starts at 2000-01-01T00:00:00Z.
|
|
||||||
func testCertificate(t *testing.T) (tls.Certificate, *x509.CertPool) {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
key, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("generate a key: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
template := &x509.Certificate{
|
|
||||||
SerialNumber: big.NewInt(1),
|
|
||||||
Subject: pkix.Name{CommonName: "smallwebwaf test CA"},
|
|
||||||
NotBefore: time.Date(1999, 12, 31, 0, 0, 0, 0, time.UTC),
|
|
||||||
NotAfter: time.Date(2000, 1, 2, 0, 0, 0, 0, time.UTC),
|
|
||||||
IsCA: true,
|
|
||||||
BasicConstraintsValid: true,
|
|
||||||
KeyUsage: x509.KeyUsageCertSign | x509.KeyUsageDigitalSignature,
|
|
||||||
ExtKeyUsage: []x509.ExtKeyUsage{x509.ExtKeyUsageServerAuth},
|
|
||||||
IPAddresses: []net.IP{net.IPv4(127, 0, 0, 1)},
|
|
||||||
}
|
|
||||||
|
|
||||||
der, err := x509.CreateCertificate(rand.Reader, template, template,
|
|
||||||
&key.PublicKey, key)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("create a certificate: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
certificate, err := x509.ParseCertificate(der)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("parse the certificate: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
roots := x509.NewCertPool()
|
|
||||||
roots.AddCert(certificate)
|
|
||||||
|
|
||||||
return tls.Certificate{Certificate: [][]byte{der}, PrivateKey: key}, roots
|
|
||||||
}
|
|
||||||
@@ -9,8 +9,6 @@ import (
|
|||||||
"io"
|
"io"
|
||||||
"log/slog"
|
"log/slog"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"sneak.berlin/go/smallwebwaf/internal/ratelimit"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
// The action a request line names: what smallwebwaf did with the
|
// The action a request line names: what smallwebwaf did with the
|
||||||
@@ -26,122 +24,42 @@ const (
|
|||||||
// for, or whose answer could not be passed on.
|
// for, or whose answer could not be passed on.
|
||||||
ActionUpstreamError = "upstream_error"
|
ActionUpstreamError = "upstream_error"
|
||||||
// ActionRateLimited is a request refused because it took its client
|
// ActionRateLimited is a request refused because it took its client
|
||||||
// over a rate limit, which bans the client.
|
// over a rate limit, or came while the client was over one.
|
||||||
ActionRateLimited = "rate_limited"
|
ActionRateLimited = "rate_limited"
|
||||||
// ActionBanned is a request refused because a ban covers its client,
|
|
||||||
// or because it matched a ban rule, which bans the client.
|
|
||||||
ActionBanned = "banned"
|
|
||||||
// ActionRuleBlocked is a request refused because it matched a block
|
|
||||||
// rule.
|
|
||||||
ActionRuleBlocked = "rule_blocked"
|
|
||||||
// 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.
|
// timeLayout is RFC 3339 with milliseconds.
|
||||||
const timeLayout = "2006-01-02T15:04:05.000Z07:00"
|
const timeLayout = "2006-01-02T15:04:05.000Z07:00"
|
||||||
|
|
||||||
// Line is one request's line in the request log. The field names, and
|
// Line is one request's line in the request log. The field names are
|
||||||
// their order, are those of the "Request log" section of SPEC.md. A field
|
// those of the "Request log" section of SPEC.md.
|
||||||
// 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
|
//nolint:tagliatelle // SPEC.md's request log names its fields in snake_case
|
||||||
type Line struct {
|
type Line struct {
|
||||||
Type string `json:"type"`
|
Type string `json:"type"`
|
||||||
|
Time string `json:"time"`
|
||||||
// The standard web log fields. Scheme is how the client reached
|
ClientIP string `json:"client_ip"`
|
||||||
// smallwebwaf, or the trusted proxy in front of it.
|
PeerIP string `json:"peer_ip"`
|
||||||
Time string `json:"time"`
|
Method string `json:"method"`
|
||||||
Instance string `json:"instance"`
|
Host string `json:"host"`
|
||||||
ClientIP string `json:"client_ip"`
|
Path string `json:"path"`
|
||||||
Method string `json:"method"`
|
Query string `json:"query"`
|
||||||
Scheme string `json:"scheme"`
|
Protocol string `json:"protocol"`
|
||||||
Host string `json:"host"`
|
Status int `json:"status"`
|
||||||
Path string `json:"path"`
|
UpstreamStatus int `json:"upstream_status,omitempty"`
|
||||||
Query string `json:"query"`
|
RequestBytes int64 `json:"request_bytes"`
|
||||||
Protocol string `json:"protocol"`
|
ResponseBytes int64 `json:"response_bytes"`
|
||||||
Status int `json:"status"`
|
Referer string `json:"referer"`
|
||||||
RequestBytes int64 `json:"request_bytes"`
|
UserAgent string `json:"user_agent"`
|
||||||
ResponseBytes int64 `json:"response_bytes"`
|
Action string `json:"action"`
|
||||||
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. ASN, ASName and Country are the client's AS
|
|
||||||
// number, AS name and country, as looked up.
|
|
||||||
RequestID string `json:"request_id"`
|
|
||||||
PeerIP string `json:"peer_ip"`
|
|
||||||
ForwardedFor string `json:"forwarded_for,omitempty"`
|
|
||||||
ClientGroup string `json:"client_group"`
|
|
||||||
ASN string `json:"asn"`
|
|
||||||
ASName string `json:"as_name"`
|
|
||||||
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, ActionRateLimited or
|
|
||||||
// ActionRuleBlocked.
|
|
||||||
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"`
|
|
||||||
// RuleIDs are the ids of the rule file rules the request matched.
|
|
||||||
RuleIDs []string `json:"rule_ids,omitempty"`
|
|
||||||
// LimitHit is the window whose rate limit the request went over:
|
// LimitHit is the window whose rate limit the request went over:
|
||||||
// minute, hour or day.
|
// minute, hour or day.
|
||||||
LimitHit string `json:"limit_hit,omitempty"`
|
LimitHit string `json:"limit_hit,omitempty"`
|
||||||
// Offence is the offence the request was held as, OffenceLimit.
|
// Aborted is true when the client went away early.
|
||||||
Offence string `json:"offence,omitempty"`
|
Aborted bool `json:"aborted,omitempty"`
|
||||||
// BanExpires is when the ban the request made, or was refused under,
|
// DurationTotal and DurationUpstreamTotal are in milliseconds.
|
||||||
// ends: a time, or "permanent".
|
DurationTotal float64 `json:"duration_total"`
|
||||||
BanExpires string `json:"ban_expires,omitempty"`
|
DurationUpstreamTotal float64 `json:"duration_upstream_total,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".
|
// Write writes line to w as one JSON line marked "type":"request".
|
||||||
@@ -174,8 +92,8 @@ func Milliseconds(d time.Duration) float64 {
|
|||||||
|
|
||||||
// NewProcessLogger returns the logger for the process's own messages:
|
// NewProcessLogger returns the logger for the process's own messages:
|
||||||
// JSON lines on w, marked "type":"process", with the time in the same form
|
// JSON lines on w, marked "type":"process", with the time in the same form
|
||||||
// as a request line's, and instanceName, SWWAF_INSTANCE_NAME, as instance.
|
// as a request line's.
|
||||||
func NewProcessLogger(w io.Writer, instanceName string) *slog.Logger {
|
func NewProcessLogger(w io.Writer) *slog.Logger {
|
||||||
handler := slog.NewJSONHandler(w, &slog.HandlerOptions{
|
handler := slog.NewJSONHandler(w, &slog.HandlerOptions{
|
||||||
ReplaceAttr: func(groups []string, attr slog.Attr) slog.Attr {
|
ReplaceAttr: func(groups []string, attr slog.Attr) slog.Attr {
|
||||||
if attr.Key == slog.TimeKey && len(groups) == 0 {
|
if attr.Key == slog.TimeKey && len(groups) == 0 {
|
||||||
@@ -186,5 +104,5 @@ func NewProcessLogger(w io.Writer, instanceName string) *slog.Logger {
|
|||||||
},
|
},
|
||||||
})
|
})
|
||||||
|
|
||||||
return slog.New(handler).With("type", "process", "instance", instanceName)
|
return slog.New(handler).With("type", "process")
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -50,12 +50,7 @@ func TestWriteWritesOneJSONLineMarkedRequest(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
unset := []string{
|
unset := []string{
|
||||||
"forwarded_for", "content_type", "content_length", "request_headers",
|
"upstream_status", "limit_hit", "aborted", "duration_upstream_total",
|
||||||
"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 {
|
for _, name := range unset {
|
||||||
_, present := fields[name]
|
_, present := fields[name]
|
||||||
@@ -65,12 +60,12 @@ func TestWriteWritesOneJSONLineMarkedRequest(t *testing.T) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
func TestProcessLinesAreMarkedProcessAndGiveTheInstance(t *testing.T) {
|
func TestProcessLinesAreMarkedProcess(t *testing.T) {
|
||||||
t.Parallel()
|
t.Parallel()
|
||||||
|
|
||||||
var out bytes.Buffer
|
var out bytes.Buffer
|
||||||
|
|
||||||
requestlog.NewProcessLogger(&out, "fsn1app1/gitea").Info("starting", "version", "v1")
|
requestlog.NewProcessLogger(&out).Info("starting", "version", "v1")
|
||||||
|
|
||||||
var fields map[string]any
|
var fields map[string]any
|
||||||
|
|
||||||
@@ -79,9 +74,8 @@ func TestProcessLinesAreMarkedProcessAndGiveTheInstance(t *testing.T) {
|
|||||||
t.Fatalf("decode %q: %v", out.String(), err)
|
t.Fatalf("decode %q: %v", out.String(), err)
|
||||||
}
|
}
|
||||||
|
|
||||||
if fields["type"] != "process" || fields["instance"] != "fsn1app1/gitea" ||
|
if fields["type"] != "process" || fields["msg"] != "starting" ||
|
||||||
fields["msg"] != "starting" || fields["level"] != "INFO" ||
|
fields["level"] != "INFO" || fields["version"] != "v1" {
|
||||||
fields["version"] != "v1" {
|
|
||||||
t.Errorf("process line %v", fields)
|
t.Errorf("process line %v", fields)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -1,483 +0,0 @@
|
|||||||
// Package rules reads the rule files: the plain text files in
|
|
||||||
// SWWAF_RULES_DIR, one rule to a line, that each request is checked
|
|
||||||
// against, as the "Rule files" section of SPEC.md describes. They are read
|
|
||||||
// at start, and again once the directory has had no change for a short
|
|
||||||
// time after one is edited, added or removed.
|
|
||||||
package rules
|
|
||||||
|
|
||||||
import (
|
|
||||||
"context"
|
|
||||||
"encoding/hex"
|
|
||||||
"errors"
|
|
||||||
"fmt"
|
|
||||||
"log/slog"
|
|
||||||
"net/http"
|
|
||||||
"os"
|
|
||||||
"path/filepath"
|
|
||||||
"regexp"
|
|
||||||
"slices"
|
|
||||||
"strings"
|
|
||||||
"sync/atomic"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
"github.com/fsnotify/fsnotify"
|
|
||||||
|
|
||||||
"sneak.berlin/go/smallwebwaf/internal/alerts"
|
|
||||||
"sneak.berlin/go/smallwebwaf/internal/config"
|
|
||||||
)
|
|
||||||
|
|
||||||
// The actions a rule takes when it matches.
|
|
||||||
const (
|
|
||||||
// ActionLog notes the match in the request log, and does nothing else.
|
|
||||||
ActionLog = "log"
|
|
||||||
// ActionBlock refuses the request with 403.
|
|
||||||
ActionBlock = "block"
|
|
||||||
// ActionBan refuses the request and bans the client's netblock: the
|
|
||||||
// request is a clear sign of attack.
|
|
||||||
ActionBan = "ban"
|
|
||||||
)
|
|
||||||
|
|
||||||
// extension ends the name of every rule file.
|
|
||||||
const extension = ".rules"
|
|
||||||
|
|
||||||
// quietTime is how long SWWAF_RULES_DIR must go without a change before
|
|
||||||
// the rule files are read again, so that a file still being written, such
|
|
||||||
// as one saved in place, appended to or copied in with scp, is read only
|
|
||||||
// once whole.
|
|
||||||
const quietTime = 2 * time.Second
|
|
||||||
|
|
||||||
// headerTarget starts the target that is one request header,
|
|
||||||
// header:<Name>.
|
|
||||||
const headerTarget = "header:"
|
|
||||||
|
|
||||||
// escapeLength is the length of a percent escape, such as %2e.
|
|
||||||
const escapeLength = 3
|
|
||||||
|
|
||||||
var (
|
|
||||||
// ruleLine is a rule: four fields separated by spaces or tabs, of
|
|
||||||
// which the fourth, the regex, runs to the end of the line.
|
|
||||||
ruleLine = regexp.MustCompile(`^([^ \t]+)[ \t]+([^ \t]+)[ \t]+([^ \t]+)[ \t]+(.+)$`)
|
|
||||||
// idChars are the characters of a rule's id.
|
|
||||||
idChars = regexp.MustCompile(`^[A-Za-z0-9_-]+$`)
|
|
||||||
)
|
|
||||||
|
|
||||||
var (
|
|
||||||
errNotRule = errors.New(
|
|
||||||
"is not a rule: an id, a target, an action and a regex, " +
|
|
||||||
"separated by spaces or tabs")
|
|
||||||
errNotID = errors.New("is not an id of letters, digits, - and _")
|
|
||||||
errNotTarget = errors.New(
|
|
||||||
"is not path, query, uri, method, host, user_agent, referer or header:<Name>")
|
|
||||||
errNotHeaderName = errors.New(
|
|
||||||
"has a character after header: that no header name can have")
|
|
||||||
errHeaderTakenOut = errors.New(
|
|
||||||
"names a header that Go's HTTP server takes out of every request, " +
|
|
||||||
"so a rule never sees it")
|
|
||||||
errNotAction = errors.New("is not log, block or ban")
|
|
||||||
errNotRegex = errors.New("does not compile")
|
|
||||||
errUsedTwice = errors.New("is already the id of the rule at")
|
|
||||||
)
|
|
||||||
|
|
||||||
// Rule is one rule of a rule file.
|
|
||||||
type Rule struct {
|
|
||||||
// ID names the rule in the request log, the metrics and ban notes.
|
|
||||||
ID string
|
|
||||||
// Target is what the regex is matched against, such as path or
|
|
||||||
// header:Accept.
|
|
||||||
Target string
|
|
||||||
// Action is ActionLog, ActionBlock or ActionBan.
|
|
||||||
Action string
|
|
||||||
|
|
||||||
regex *regexp.Regexp
|
|
||||||
}
|
|
||||||
|
|
||||||
// Params are what Load needs.
|
|
||||||
type Params struct {
|
|
||||||
// Dir is the directory of the rule files (SWWAF_RULES_DIR).
|
|
||||||
Dir string
|
|
||||||
// Enabled is SWWAF_RULES_ENABLED: while it is false, no file is read
|
|
||||||
// and no rule loaded.
|
|
||||||
Enabled bool
|
|
||||||
// ProcessLog receives how many rules were read, and the error in a
|
|
||||||
// rule file edited while smallwebwaf runs.
|
|
||||||
ProcessLog *slog.Logger
|
|
||||||
// Alerts receive a file_error alert for that error.
|
|
||||||
Alerts *alerts.Queue
|
|
||||||
}
|
|
||||||
|
|
||||||
// Files are the rule files of a running smallwebwaf, and the rules read
|
|
||||||
// from them. They are safe for concurrent use.
|
|
||||||
type Files struct {
|
|
||||||
params Params
|
|
||||||
// rules are the rules loaded, in the order of their files' names, and
|
|
||||||
// then of their lines.
|
|
||||||
rules atomic.Pointer[[]Rule]
|
|
||||||
}
|
|
||||||
|
|
||||||
// Load reads the rules of every *.rules file in Dir, in the order of the
|
|
||||||
// files' names, unless Enabled is false. A Dir that cannot be read is an
|
|
||||||
// error, and so is a line that is not a rule, a header name with a
|
|
||||||
// character no header name can have, a rule for the Host or the
|
|
||||||
// Transfer-Encoding header, which Go's HTTP server takes out of every
|
|
||||||
// request, a regex that does not compile and an id used twice, each named
|
|
||||||
// with its file and line.
|
|
||||||
func Load(params Params) (*Files, error) {
|
|
||||||
f := &Files{params: params}
|
|
||||||
f.rules.Store(&[]Rule{})
|
|
||||||
|
|
||||||
if !params.Enabled {
|
|
||||||
return f, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
rules, _, err := read(params.Dir)
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
f.rules.Store(&rules)
|
|
||||||
f.logRead(len(rules))
|
|
||||||
|
|
||||||
return f, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// Match checks r against the rules, in order, and returns those it
|
|
||||||
// matches, up to the first whose action refuses it, block or ban, which
|
|
||||||
// is then the last one returned.
|
|
||||||
func (f *Files) Match(r *http.Request) []Rule {
|
|
||||||
var matched []Rule
|
|
||||||
|
|
||||||
for _, rule := range *f.rules.Load() {
|
|
||||||
if !rule.matches(r) {
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
|
|
||||||
matched = append(matched, rule)
|
|
||||||
if rule.Action != ActionLog {
|
|
||||||
break
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
return matched
|
|
||||||
}
|
|
||||||
|
|
||||||
// Len returns how many rules are loaded.
|
|
||||||
func (f *Files) Len() int {
|
|
||||||
return len(*f.rules.Load())
|
|
||||||
}
|
|
||||||
|
|
||||||
// Watch watches Dir until ctx is done, and reads the rule files again
|
|
||||||
// once Dir has had no change for quietTime, after one is edited, added or
|
|
||||||
// removed, and after Watch starts watching. If they then hold an error,
|
|
||||||
// the rules stay as they were, the error is logged with its file and
|
|
||||||
// line, and the files are read again after the next change. If Dir cannot
|
|
||||||
// be watched, that is logged, and the rules stay as they were loaded.
|
|
||||||
// While Enabled is false, Watch returns at once.
|
|
||||||
func (f *Files) Watch(ctx context.Context) {
|
|
||||||
if !f.params.Enabled {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
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 rule files for edits",
|
|
||||||
"error", err.Error())
|
|
||||||
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
f.params.ProcessLog.Info("watching the rule files for edits",
|
|
||||||
"directory", f.params.Dir)
|
|
||||||
|
|
||||||
f.readAfterChanges(ctx, watcher.Events, watcher.Errors)
|
|
||||||
}
|
|
||||||
|
|
||||||
// readAfterChanges reads the rule files again once quietTime has passed
|
|
||||||
// without a change from events, until ctx is done, and logs the errors
|
|
||||||
// from errs. The wait starts at once, as if for a change, so that an edit
|
|
||||||
// saved after Load read the files, and before Dir was watched, is taken
|
|
||||||
// in too.
|
|
||||||
func (f *Files) readAfterChanges(
|
|
||||||
ctx context.Context, events <-chan fsnotify.Event, errs <-chan error,
|
|
||||||
) {
|
|
||||||
quiet := time.NewTimer(quietTime)
|
|
||||||
defer quiet.Stop()
|
|
||||||
|
|
||||||
for {
|
|
||||||
select {
|
|
||||||
case <-ctx.Done():
|
|
||||||
return
|
|
||||||
case <-events:
|
|
||||||
quiet.Reset(quietTime)
|
|
||||||
case <-quiet.C:
|
|
||||||
f.readAgain()
|
|
||||||
case err := <-errs:
|
|
||||||
f.params.ProcessLog.Warn("watching the rule files failed",
|
|
||||||
"error", err.Error())
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// readAgain reads the rule files again, in place of the rules loaded, or
|
|
||||||
// logs the error that keeps the rules as they were, and raises a
|
|
||||||
// file_error alert for it, for the file it is in.
|
|
||||||
func (f *Files) readAgain() {
|
|
||||||
rules, path, err := read(f.params.Dir)
|
|
||||||
if err != nil {
|
|
||||||
const kept = "a rule file has an error, and the rules stay as they were"
|
|
||||||
|
|
||||||
// Raised before it is logged, so that the alert is there once the
|
|
||||||
// log line is.
|
|
||||||
f.params.Alerts.Raise(alerts.Alert{
|
|
||||||
Event: alerts.EventFileError,
|
|
||||||
Reason: kept,
|
|
||||||
Detail: map[string]any{"file": path, "error": err.Error()},
|
|
||||||
})
|
|
||||||
f.params.ProcessLog.Error(kept, "error", err.Error())
|
|
||||||
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
f.rules.Store(&rules)
|
|
||||||
f.logRead(len(rules))
|
|
||||||
}
|
|
||||||
|
|
||||||
// logRead logs that the rule files were read, and how many rules they
|
|
||||||
// hold, which can be none.
|
|
||||||
func (f *Files) logRead(count int) {
|
|
||||||
f.params.ProcessLog.Info("read the rule files",
|
|
||||||
"directory", f.params.Dir, "rules", count)
|
|
||||||
}
|
|
||||||
|
|
||||||
// read returns the rules of every rule file in dir, in the order of the
|
|
||||||
// files' names, and then of their lines, or an error, with the path of the
|
|
||||||
// rule file it is in, or dir. A file whose name starts with a dot, such as
|
|
||||||
// an editor's lock file .#50-app.rules, is not a rule file, as a shell's
|
|
||||||
// *.rules would not match it.
|
|
||||||
func read(dir string) ([]Rule, string, error) {
|
|
||||||
entries, err := os.ReadDir(dir)
|
|
||||||
if err != nil {
|
|
||||||
return nil, dir, fmt.Errorf("SWWAF_RULES_DIR cannot be read: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
var rules []Rule
|
|
||||||
|
|
||||||
// places are where each id is, as "<file>, line <n>".
|
|
||||||
places := map[string]string{}
|
|
||||||
|
|
||||||
for _, entry := range entries {
|
|
||||||
name := entry.Name()
|
|
||||||
if entry.IsDir() || strings.HasPrefix(name, ".") || filepath.Ext(name) != extension {
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
|
|
||||||
path := filepath.Join(dir, name)
|
|
||||||
|
|
||||||
rules, err = readFile(path, rules, places)
|
|
||||||
if err != nil {
|
|
||||||
return nil, path, err
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
return rules, "", nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// readFile appends the rules of the rule file at path to rules. places
|
|
||||||
// are where each id read so far is, and gain those of the file.
|
|
||||||
func readFile(path string, rules []Rule, places map[string]string) ([]Rule, error) {
|
|
||||||
data, err := os.ReadFile(path) //nolint:gosec // a rule file, in SWWAF_RULES_DIR
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
number := 0
|
|
||||||
|
|
||||||
for line := range strings.Lines(string(data)) {
|
|
||||||
number++
|
|
||||||
place := fmt.Sprintf("%s, line %d", path, number)
|
|
||||||
|
|
||||||
text := strings.TrimSuffix(strings.TrimSuffix(line, "\n"), "\r")
|
|
||||||
|
|
||||||
rule, isRule, err := parse(text)
|
|
||||||
if err != nil {
|
|
||||||
return nil, fmt.Errorf("%s: %w", place, err)
|
|
||||||
}
|
|
||||||
|
|
||||||
if !isRule {
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
|
|
||||||
first, used := places[rule.ID]
|
|
||||||
if used {
|
|
||||||
return nil, fmt.Errorf("%s: the id %q %w %s", place, rule.ID, errUsedTwice, first)
|
|
||||||
}
|
|
||||||
|
|
||||||
places[rule.ID] = place
|
|
||||||
rules = append(rules, rule)
|
|
||||||
}
|
|
||||||
|
|
||||||
return rules, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// parse reads a line of a rule file. It returns false for a blank line
|
|
||||||
// and for a comment, a line that starts with #. Spaces and tabs at the
|
|
||||||
// end of the line are not part of its regex, so a line with only those
|
|
||||||
// after its action has no regex, and is not a rule.
|
|
||||||
func parse(line string) (Rule, bool, error) {
|
|
||||||
line = strings.Trim(line, " \t")
|
|
||||||
if line == "" || strings.HasPrefix(line, "#") {
|
|
||||||
return Rule{}, false, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
fields := ruleLine.FindStringSubmatch(line)
|
|
||||||
if fields == nil {
|
|
||||||
return Rule{}, false, errNotRule
|
|
||||||
}
|
|
||||||
|
|
||||||
rule := Rule{ID: fields[1], Target: fields[2], Action: fields[3]}
|
|
||||||
headerName, isHeader := strings.CutPrefix(rule.Target, headerTarget)
|
|
||||||
|
|
||||||
switch {
|
|
||||||
case !idChars.MatchString(rule.ID):
|
|
||||||
return Rule{}, false, fmt.Errorf("the id %q %w", rule.ID, errNotID)
|
|
||||||
case !isTarget(rule.Target):
|
|
||||||
return Rule{}, false, fmt.Errorf("the target %q %w", rule.Target, errNotTarget)
|
|
||||||
case isHeader && !config.IsHeaderName(headerName):
|
|
||||||
return Rule{}, false, fmt.Errorf("the target %q %w", rule.Target, errNotHeaderName)
|
|
||||||
case strings.EqualFold(rule.Target, headerTarget+"Host"):
|
|
||||||
return Rule{}, false, fmt.Errorf(
|
|
||||||
"the target %q %w; the request's host is the target host",
|
|
||||||
rule.Target, errHeaderTakenOut)
|
|
||||||
case strings.EqualFold(rule.Target, headerTarget+"Transfer-Encoding"):
|
|
||||||
return Rule{}, false, fmt.Errorf("the target %q %w", rule.Target, errHeaderTakenOut)
|
|
||||||
case !slices.Contains([]string{ActionLog, ActionBlock, ActionBan}, rule.Action):
|
|
||||||
return Rule{}, false, fmt.Errorf("the action %q %w", rule.Action, errNotAction)
|
|
||||||
}
|
|
||||||
|
|
||||||
regex, err := regexp.Compile(fields[4])
|
|
||||||
if err != nil {
|
|
||||||
return Rule{}, false, fmt.Errorf("the regex %w: %w", errNotRegex, err)
|
|
||||||
}
|
|
||||||
|
|
||||||
rule.regex = regex
|
|
||||||
|
|
||||||
return rule, true, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// isTarget reports whether target is one a rule may have.
|
|
||||||
func isTarget(target string) bool {
|
|
||||||
switch target {
|
|
||||||
case "path", "query", "uri", "method", "host", "user_agent", "referer":
|
|
||||||
return true
|
|
||||||
}
|
|
||||||
|
|
||||||
name, isHeader := strings.CutPrefix(target, headerTarget)
|
|
||||||
|
|
||||||
return isHeader && name != ""
|
|
||||||
}
|
|
||||||
|
|
||||||
// matches reports whether the rule's regex matches its target in r. For
|
|
||||||
// uri it is matched against the path and query as received, and against
|
|
||||||
// them once percent-decoded, so that an encoded probe cannot slip past.
|
|
||||||
func (rule Rule) matches(r *http.Request) bool {
|
|
||||||
if rule.Target == "uri" {
|
|
||||||
uri := pathAndQuery(r)
|
|
||||||
|
|
||||||
return rule.regex.MatchString(uri) || rule.regex.MatchString(decodeOnce(uri))
|
|
||||||
}
|
|
||||||
|
|
||||||
return rule.regex.MatchString(value(rule.Target, r))
|
|
||||||
}
|
|
||||||
|
|
||||||
// value returns what a rule with target, other than uri, is matched
|
|
||||||
// against in r: the path and the query as the client sent them, before
|
|
||||||
// any decoding or re-encoding, split at the first ?, and a header's values
|
|
||||||
// joined by ", ", as HTTP joins those of a header sent more than once.
|
|
||||||
func value(target string, r *http.Request) string {
|
|
||||||
switch target {
|
|
||||||
case "path":
|
|
||||||
path, _, _ := strings.Cut(pathAndQuery(r), "?")
|
|
||||||
|
|
||||||
return path
|
|
||||||
case "query":
|
|
||||||
_, query, _ := strings.Cut(pathAndQuery(r), "?")
|
|
||||||
|
|
||||||
return query
|
|
||||||
case "method":
|
|
||||||
return r.Method
|
|
||||||
case "host":
|
|
||||||
return r.Host
|
|
||||||
case "user_agent":
|
|
||||||
return header(r, "User-Agent")
|
|
||||||
case "referer":
|
|
||||||
return header(r, "Referer")
|
|
||||||
default:
|
|
||||||
return header(r, strings.TrimPrefix(target, headerTarget))
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// pathAndQuery returns the target of r's request line, r.RequestURI, as
|
|
||||||
// the client sent it, less any scheme and host: a target with a scheme
|
|
||||||
// gives what follows the scheme and its :, and the host when // follows.
|
|
||||||
// So http://host/path, as a client sends it to a proxy, gives /path, and
|
|
||||||
// so does http:/path, which Go reads as a target with a scheme and no
|
|
||||||
// host. r.URL is not used: when the path holds a character it escapes,
|
|
||||||
// such as \ or a non-ASCII byte, it decodes the whole path and escapes it
|
|
||||||
// again, so that \ becomes %5C and %2e a dot.
|
|
||||||
func pathAndQuery(r *http.Request) string {
|
|
||||||
if !r.URL.IsAbs() {
|
|
||||||
return r.RequestURI
|
|
||||||
}
|
|
||||||
|
|
||||||
_, afterScheme, _ := strings.Cut(r.RequestURI, ":")
|
|
||||||
|
|
||||||
hostAndRest, hasHost := strings.CutPrefix(afterScheme, "//")
|
|
||||||
if !hasHost {
|
|
||||||
return afterScheme
|
|
||||||
}
|
|
||||||
|
|
||||||
start := strings.IndexAny(hostAndRest, "/?")
|
|
||||||
if start < 0 {
|
|
||||||
return ""
|
|
||||||
}
|
|
||||||
|
|
||||||
return hostAndRest[start:]
|
|
||||||
}
|
|
||||||
|
|
||||||
// header returns the values of r's header name joined by ", ", or "" if
|
|
||||||
// r has no such header.
|
|
||||||
func header(r *http.Request, name string) string {
|
|
||||||
return strings.Join(r.Header.Values(name), ", ")
|
|
||||||
}
|
|
||||||
|
|
||||||
// decodeOnce returns s with each percent escape, such as %2e, replaced by
|
|
||||||
// the byte it stands for. A % that is not followed by two hex digits is
|
|
||||||
// left as it is, so that a malformed escape cannot keep the rest of s
|
|
||||||
// from being decoded.
|
|
||||||
func decodeOnce(s string) string {
|
|
||||||
var decoded strings.Builder
|
|
||||||
|
|
||||||
for i := 0; i < len(s); i++ {
|
|
||||||
if s[i] == '%' && i+escapeLength <= len(s) {
|
|
||||||
b, err := hex.DecodeString(s[i+1 : i+escapeLength])
|
|
||||||
if err == nil {
|
|
||||||
decoded.Write(b)
|
|
||||||
|
|
||||||
i += escapeLength - 1
|
|
||||||
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
decoded.WriteByte(s[i])
|
|
||||||
}
|
|
||||||
|
|
||||||
return decoded.String()
|
|
||||||
}
|
|
||||||
@@ -1,687 +0,0 @@
|
|||||||
package rules_test
|
|
||||||
|
|
||||||
import (
|
|
||||||
"context"
|
|
||||||
"encoding/json"
|
|
||||||
"log/slog"
|
|
||||||
"maps"
|
|
||||||
"net/http"
|
|
||||||
"net/http/httptest"
|
|
||||||
"net/url"
|
|
||||||
"os"
|
|
||||||
"path/filepath"
|
|
||||||
"slices"
|
|
||||||
"strconv"
|
|
||||||
"testing"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
"sneak.berlin/go/smallwebwaf/internal/alerts"
|
|
||||||
"sneak.berlin/go/smallwebwaf/internal/rules"
|
|
||||||
)
|
|
||||||
|
|
||||||
const (
|
|
||||||
// What the process log says once Watch watches the directory, after
|
|
||||||
// each reading of the rule files, and for one that has an error.
|
|
||||||
watching = "watching the rule files for edits"
|
|
||||||
read = "read the rule files"
|
|
||||||
hasError = "a rule file has an error, and the rules stay as they were"
|
|
||||||
// maxLogLines is how many lines of the process log wait for a test to
|
|
||||||
// read them.
|
|
||||||
maxLogLines = 64
|
|
||||||
// browser is the user agent of an ordinary visitor.
|
|
||||||
browser = "Mozilla/5.0 (X11; Linux x86_64; rv:140.0) Gecko/20100101 Firefox/140.0"
|
|
||||||
// testFile is the rule file of a test that needs only one, and
|
|
||||||
// firstFile the first of a test's rule files.
|
|
||||||
testFile = "test.rules"
|
|
||||||
firstFile = "00-a.rules"
|
|
||||||
// userAgent is the header that carries the user agent.
|
|
||||||
userAgent = "User-Agent"
|
|
||||||
)
|
|
||||||
|
|
||||||
func TestEachTargetMatchesWhatItNames(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
for _, tc := range []struct {
|
|
||||||
name string
|
|
||||||
rule string // its target, action and regex
|
|
||||||
uri string // the request's path and query
|
|
||||||
header http.Header
|
|
||||||
want bool
|
|
||||||
}{
|
|
||||||
{"path as received", `path log ^/%2eenv$`, "/%2eenv", nil, true},
|
|
||||||
{"path not decoded", `path log ^/\.env$`, "/%2eenv", nil, false},
|
|
||||||
{"path without the query", `path log ^/a$`, "/a?b=c", nil, true},
|
|
||||||
{"query as received", `query log ^b=%2e$`, "/a?b=%2e", nil, true},
|
|
||||||
{"uri as received", `uri log ^/a\?b=%2e$`, "/a?b=%2e", nil, true},
|
|
||||||
{"uri decoded", `uri log (\.\./){2}`, "/a?f=%2e%2e%2f%2e%2e%2f", nil, true},
|
|
||||||
{
|
|
||||||
"uri decoded past malformed escapes", `uri log (\.\./){2}&h=%$`,
|
|
||||||
"/a?g=%zz&f=%2e%2e%2f%2e%2e%2f&h=%", nil, true,
|
|
||||||
},
|
|
||||||
{"uri decoded only once", `uri log ^/a\.b$`, "/a%252eb", nil, false},
|
|
||||||
{"method", `method log ^PUT$`, "/", nil, true},
|
|
||||||
{"host", `host log ^app\.example$`, "/", nil, true},
|
|
||||||
{
|
|
||||||
"user_agent", `user_agent log ^sqlmap/`, "/",
|
|
||||||
http.Header{userAgent: {"sqlmap/1.8"}}, true,
|
|
||||||
},
|
|
||||||
{
|
|
||||||
"user_agent sent twice", `user_agent log ^curl/8, sqlmap/`, "/",
|
|
||||||
http.Header{userAgent: {"curl/8", "sqlmap/1.8"}}, true,
|
|
||||||
},
|
|
||||||
{"user_agent missing", `user_agent log ^$`, "/", nil, true},
|
|
||||||
{
|
|
||||||
"referer", `referer log ^https://spam\.example/`, "/",
|
|
||||||
http.Header{"Referer": {"https://spam.example/buy"}}, true,
|
|
||||||
},
|
|
||||||
{
|
|
||||||
"a header sent twice", `header:x-api-version log ^2, 3$`, "/",
|
|
||||||
http.Header{"X-Api-Version": {"2", "3"}}, true,
|
|
||||||
},
|
|
||||||
{"a header missing", `header:X-Api-Version log ^$`, "/", nil, true},
|
|
||||||
} {
|
|
||||||
t.Run(tc.name, func(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
files := load(t, ruleFiles{testFile: "a-rule " + tc.rule + "\n"})
|
|
||||||
|
|
||||||
// Every request is a PUT, which the method rule looks for.
|
|
||||||
r := httptest.NewRequestWithContext(t.Context(), http.MethodPut,
|
|
||||||
"http://app.example"+tc.uri, nil)
|
|
||||||
maps.Copy(r.Header, tc.header)
|
|
||||||
|
|
||||||
got := len(files.Match(r)) == 1
|
|
||||||
if got != tc.want {
|
|
||||||
t.Errorf("%s matches %s: %t, want %t", tc.rule, tc.uri, got, tc.want)
|
|
||||||
}
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestPathMatchedAsTheClientSentIt(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
// Each path holds a character Go's URL type would escape again, \ or
|
|
||||||
// a non-ASCII byte, and each rule is written for the path as sent.
|
|
||||||
for _, tc := range []struct {
|
|
||||||
rule string // its target, action and regex
|
|
||||||
sent string // the path and query the client sent
|
|
||||||
}{
|
|
||||||
{`path log ^/\.\.\\\.\.\\windows\\win\.ini$`, `/..\..\windows\win.ini`},
|
|
||||||
{`path log ^/%2e%2e\\%2e%2e\\windows\\win\.ini$`, `/%2e%2e\%2e%2e\windows\win.ini`},
|
|
||||||
{`path log ^/café$`, "/café?x=1"},
|
|
||||||
{`uri log ^/%2e%2e\\%2e%2e\\boot\.ini\?x=1$`, `/%2e%2e\%2e%2e\boot.ini?x=1`},
|
|
||||||
} {
|
|
||||||
files := load(t, ruleFiles{testFile: "as-sent " + tc.rule + "\n"})
|
|
||||||
|
|
||||||
// The target in origin form, as traefik sends it, in absolute form,
|
|
||||||
// as a client sends it to a proxy, and with a scheme but no host,
|
|
||||||
// which Go reads as absolute form with no host, sending the app
|
|
||||||
// the path.
|
|
||||||
for _, target := range []string{
|
|
||||||
tc.sent, "http://app.example" + tc.sent, "http:" + tc.sent, "foo:" + tc.sent,
|
|
||||||
} {
|
|
||||||
r := httptest.NewRequestWithContext(t.Context(), http.MethodGet, target, nil)
|
|
||||||
wantMatched(t, files, r, "as-sent")
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestMatchingStopsAtTheFirstRuleThatRefuses(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
files := load(t, ruleFiles{testFile: `
|
|
||||||
every-path path log ^/
|
|
||||||
no-path path log ^$
|
|
||||||
first-refusal path block ^/probe
|
|
||||||
later-ban path ban ^/probe
|
|
||||||
after path log ^/
|
|
||||||
`})
|
|
||||||
|
|
||||||
// Every log rule that matches is noted, and the block rule ends the
|
|
||||||
// matching.
|
|
||||||
wantMatched(t, files, get(t, "/probe"), "every-path", "first-refusal")
|
|
||||||
wantMatched(t, files, get(t, "/page"), "every-path", "after")
|
|
||||||
|
|
||||||
// A ban rule ends it too.
|
|
||||||
files = load(t, ruleFiles{testFile: "ban path ban ^/\nlater path block ^/\n"})
|
|
||||||
wantMatched(t, files, get(t, "/"), "ban")
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestSpacesAndTabsEndingALineAreNotPartOfItsRegex(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
files := load(t, ruleFiles{testFile: "env-file path block ^/\\.env$ \t \n"})
|
|
||||||
|
|
||||||
wantMatched(t, files, get(t, "/.env"), "env-file")
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestFilesReadInNameOrderThenLineOrder(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
files := load(t, ruleFiles{
|
|
||||||
"50-b.rules": "b1 path log ^/\n\n# a comment\n # an indented one\nb2 path log ^/\n",
|
|
||||||
firstFile: "a1 path log ^/\r\n",
|
|
||||||
// None is a rule file.
|
|
||||||
"notes.txt": "notes, not rules\n",
|
|
||||||
"10-c.rules.bak": "an old copy\n",
|
|
||||||
"20-d.rules/keep": "a file in a directory\n",
|
|
||||||
})
|
|
||||||
|
|
||||||
wantMatched(t, files, get(t, "/"), "a1", "b1", "b2")
|
|
||||||
|
|
||||||
if files.Len() != 3 {
|
|
||||||
t.Errorf("%d rules loaded, want 3", files.Len())
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestFileWhoseNameStartsWithADotIsNotARuleFile(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
dir := writeFiles(t, ruleFiles{firstFile: "probe path block ^/probe\n"})
|
|
||||||
|
|
||||||
// The lock file Emacs makes beside a file while it is edited: a link to
|
|
||||||
// nothing, which cannot be read.
|
|
||||||
err := os.Symlink("user@host.1234:1700000000", filepath.Join(dir, ".#"+firstFile))
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("symlink: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
params, _ := newParams(dir)
|
|
||||||
|
|
||||||
files, err := rules.Load(params)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("load: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
wantMatched(t, files, get(t, "/probe"), "probe")
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestFaultStopsTheStartNamingTheFileAndLine(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
for _, tc := range []struct {
|
|
||||||
name string
|
|
||||||
content string
|
|
||||||
line int
|
|
||||||
want string
|
|
||||||
}{
|
|
||||||
{
|
|
||||||
"too few fields", "env-file path ban\n", 1,
|
|
||||||
"is not a rule: an id, a target, an action and a regex, " +
|
|
||||||
"separated by spaces or tabs",
|
|
||||||
},
|
|
||||||
{
|
|
||||||
// Else its regex would be a space, found in nearly every user agent.
|
|
||||||
"a regex of only spaces and tabs", "scanner user_agent ban\t \n", 1,
|
|
||||||
"is not a rule: an id, a target, an action and a regex, " +
|
|
||||||
"separated by spaces or tabs",
|
|
||||||
},
|
|
||||||
{
|
|
||||||
"an id of other characters", "# ids\n\nenv.file path ban ^/\n", 3,
|
|
||||||
`the id "env.file" is not an id of letters, digits, - and _`,
|
|
||||||
},
|
|
||||||
{
|
|
||||||
"an unknown target", "env-file paths ban ^/\n", 1,
|
|
||||||
`the target "paths" is not path, query, uri, method, host, ` +
|
|
||||||
"user_agent, referer or header:<Name>",
|
|
||||||
},
|
|
||||||
{
|
|
||||||
"a header without a name", "env-file header: ban ^/\n", 1,
|
|
||||||
`the target "header:" is not path, query, uri, method, host, ` +
|
|
||||||
"user_agent, referer or header:<Name>",
|
|
||||||
},
|
|
||||||
{
|
|
||||||
"a header name written with its colon", "sqlmap header:User-Agent: ban sqlmap\n", 1,
|
|
||||||
`the target "header:User-Agent:" has a character after header: ` +
|
|
||||||
"that no header name can have",
|
|
||||||
},
|
|
||||||
{
|
|
||||||
"a header name with a semicolon", "accept header:Accept;q log ^$\n", 1,
|
|
||||||
`the target "header:Accept;q" has a character after header: ` +
|
|
||||||
"that no header name can have",
|
|
||||||
},
|
|
||||||
{
|
|
||||||
"a header name with brackets", "x-header header:X(y) log ^$\n", 1,
|
|
||||||
`the target "header:X(y)" has a character after header: ` +
|
|
||||||
"that no header name can have",
|
|
||||||
},
|
|
||||||
{
|
|
||||||
"the Host header", "host-header header:host block ^$\n", 1,
|
|
||||||
`the target "header:host" names a header that Go's HTTP server ` +
|
|
||||||
"takes out of every request, so a rule never sees it; " +
|
|
||||||
"the request's host is the target host",
|
|
||||||
},
|
|
||||||
{
|
|
||||||
"the Transfer-Encoding header",
|
|
||||||
"# bodies sent in chunks\nchunked header:Transfer-Encoding block ^chunked$\n", 2,
|
|
||||||
`the target "header:Transfer-Encoding" names a header that Go's ` +
|
|
||||||
"HTTP server takes out of every request, so a rule never sees it",
|
|
||||||
},
|
|
||||||
{
|
|
||||||
"an unknown action", "env-file path deny ^/\n", 1,
|
|
||||||
`the action "deny" is not log, block or ban`,
|
|
||||||
},
|
|
||||||
{
|
|
||||||
"a regex that does not compile", "env-file path ban ^/(\n", 1,
|
|
||||||
"the regex does not compile: error parsing regexp: " +
|
|
||||||
"missing closing ): `^/(`",
|
|
||||||
},
|
|
||||||
} {
|
|
||||||
t.Run(tc.name, func(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
dir := writeFiles(t, ruleFiles{"00-default.rules": tc.content})
|
|
||||||
path := filepath.Join(dir, "00-default.rules")
|
|
||||||
|
|
||||||
wantRefused(t, dir, path+", line "+strconv.Itoa(tc.line)+": "+tc.want)
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestIDUsedTwiceStopsTheStartNamingBothPlaces(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
dir := writeFiles(t, ruleFiles{
|
|
||||||
"00-a.rules": "probe path log ^/a\n",
|
|
||||||
"50-b.rules": "other path log ^/b\nprobe path ban ^/c\n",
|
|
||||||
})
|
|
||||||
|
|
||||||
wantRefused(t, dir, filepath.Join(dir, "50-b.rules")+`, line 2: the id "probe" `+
|
|
||||||
"is already the id of the rule at "+filepath.Join(dir, "00-a.rules")+", line 1")
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestDirectoryThatDoesNotExistStopsTheStart(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
dir := filepath.Join(t.TempDir(), "rules.d")
|
|
||||||
|
|
||||||
wantRefused(t, dir, "SWWAF_RULES_DIR cannot be read: open "+dir+
|
|
||||||
": no such file or directory")
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestEmptyDirectoryLoadsNoRulesAndSaysSo(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
params, lines := newParams(writeFiles(t, ruleFiles{"00-default.rules": "# none\n"}))
|
|
||||||
|
|
||||||
files, err := rules.Load(params)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("load: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
line := lines.waitFor(t, read)
|
|
||||||
if files.Len() != 0 || line["rules"] != 0.0 {
|
|
||||||
t.Errorf("%d rules loaded, and the log says %v, want none", files.Len(), line)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestRuleFilesOffReadNothing(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
// SWWAF_RULES_DIR does not exist, which would stop the start.
|
|
||||||
params, _ := newParams(filepath.Join(t.TempDir(), "rules.d"))
|
|
||||||
params.Enabled = false
|
|
||||||
|
|
||||||
files, err := rules.Load(params)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("load: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
if files.Len() != 0 || files.Match(get(t, "/")) != nil {
|
|
||||||
t.Errorf("%d rules loaded with the rule files off", files.Len())
|
|
||||||
}
|
|
||||||
|
|
||||||
// It would watch until the test ends.
|
|
||||||
files.Watch(t.Context())
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestEditsTakenInWhileRunning(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
dir := writeFiles(t, ruleFiles{firstFile: "first path block ^/first\n"})
|
|
||||||
files, lines, _ := watch(t, dir)
|
|
||||||
|
|
||||||
// matches reports whether path matches a rule.
|
|
||||||
matches := func(path string) bool { return len(files.Match(get(t, path))) == 1 }
|
|
||||||
|
|
||||||
// A file added.
|
|
||||||
save(t, dir, "50-b.rules", "second path block ^/second\n")
|
|
||||||
lines.waitUntil(t, func() bool { return matches("/second") })
|
|
||||||
wantMatched(t, files, get(t, "/first"), "first")
|
|
||||||
|
|
||||||
// A file edited.
|
|
||||||
save(t, dir, firstFile, "first path block ^/edited\n")
|
|
||||||
lines.waitUntil(t, func() bool { return !matches("/first") })
|
|
||||||
wantMatched(t, files, get(t, "/edited"), "first")
|
|
||||||
|
|
||||||
// A file removed.
|
|
||||||
err := os.Remove(filepath.Join(dir, "50-b.rules"))
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("remove: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
lines.waitUntil(t, func() bool { return !matches("/second") })
|
|
||||||
wantMatched(t, files, get(t, "/edited"), "first")
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestBrokenEditKeepsTheRulesAsTheyWere(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
dir := writeFiles(t, ruleFiles{firstFile: "first path block ^/first\n"})
|
|
||||||
files, lines, queue := watch(t, dir)
|
|
||||||
|
|
||||||
// The edit's second line has an unknown action, so the rules stay as
|
|
||||||
// they were, the first line's earlier version included.
|
|
||||||
save(t, dir, firstFile, "first path block ^/edited\nsecond path bann ^/second\n")
|
|
||||||
|
|
||||||
line := lines.waitFor(t, hasError)
|
|
||||||
want := filepath.Join(dir, firstFile) +
|
|
||||||
`, line 2: the action "bann" is not log, block or ban`
|
|
||||||
|
|
||||||
if line["error"] != want || line["level"] != "ERROR" {
|
|
||||||
t.Errorf("logged %v, want an error %q", line, want)
|
|
||||||
}
|
|
||||||
|
|
||||||
// The error is raised as a file_error alert too, for the file.
|
|
||||||
wantFileError := func() {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
waiting := queue.Snapshot().Waiting[alerts.DestinationWebhook]
|
|
||||||
if len(waiting) != 1 || waiting[0].Event != alerts.EventFileError ||
|
|
||||||
waiting[0].Reason != hasError || waiting[0].Detail["error"] != want ||
|
|
||||||
waiting[0].Detail["file"] != filepath.Join(dir, firstFile) {
|
|
||||||
t.Errorf("alerts waiting %+v, want one file_error alert for %q", waiting, want)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
wantFileError()
|
|
||||||
|
|
||||||
wantMatched(t, files, get(t, "/first"), "first")
|
|
||||||
wantMatched(t, files, get(t, "/second"))
|
|
||||||
|
|
||||||
// Once mended, the file is read again, and raises no alert.
|
|
||||||
save(t, dir, firstFile, "first path block ^/edited\nsecond path ban ^/second\n")
|
|
||||||
lines.waitUntil(t, func() bool { return len(files.Match(get(t, "/second"))) == 1 })
|
|
||||||
wantMatched(t, files, get(t, "/edited"), "first")
|
|
||||||
wantFileError()
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestDefaultFileBansProbesAtTheSiteRootAlone(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
params, _ := newParams(filepath.Join("..", "..", "share", "rules.d"))
|
|
||||||
|
|
||||||
files, err := rules.Load(params)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("load the default file: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
// Probes sent by a browser, by the rule that refuses them.
|
|
||||||
for rule, targets := range map[string][]string{
|
|
||||||
"env-file": {"/.env", "/.env.production", "/.ENV"},
|
|
||||||
"vcs-dir": {"/.git/config", "/.git", "/.svn/entries"},
|
|
||||||
"secrets-dir": {"/.aws/credentials", "/.ssh/id_rsa"},
|
|
||||||
"secret-file": {"/.htpasswd", "/.DS_Store", "/.git-credentials"},
|
|
||||||
"editor-dir": {"/.vscode/sftp.json"},
|
|
||||||
"backup-file": {
|
|
||||||
"/wp-config.php.bak", "/index.php~", "/dump.sql", "/backup.sql.gz",
|
|
||||||
},
|
|
||||||
"log-file": {"/debug.log"},
|
|
||||||
"compose-file": {"/docker-compose.yml", "/compose.yaml"},
|
|
||||||
"php-shell": {"/shell.php"},
|
|
||||||
"path-traversal": {
|
|
||||||
"/static/../../etc/passwd", "/f?f=%2e%2e%2f%2e%2e%2fetc%2fpasswd",
|
|
||||||
},
|
|
||||||
} {
|
|
||||||
for _, target := range targets {
|
|
||||||
wantRefusedBy(t, files, target, browser, rule)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// Scanners, by their user agents.
|
|
||||||
for _, scanner := range []string{
|
|
||||||
"sqlmap/1.8.4#stable (https://sqlmap.org)",
|
|
||||||
"Mozilla/5.0 (compatible; Nuclei - Open-source project)",
|
|
||||||
} {
|
|
||||||
wantRefusedBy(t, files, "/", scanner, "scanner-agent")
|
|
||||||
}
|
|
||||||
|
|
||||||
// Ordinary requests to a code forge for files of those names deeper
|
|
||||||
// in its paths, and for other files at its root.
|
|
||||||
for _, target := range []string{
|
|
||||||
"/owner/repo/src/branch/main/.env.example",
|
|
||||||
"/owner/repo/src/branch/main/.env",
|
|
||||||
"/owner/repo/src/branch/main/.github/workflows/ci.yml",
|
|
||||||
"/owner/repo/src/branch/main/.vscode/settings.json",
|
|
||||||
"/owner/repo/src/branch/main/.htaccess",
|
|
||||||
"/owner/repo/src/branch/main/docker-compose.yml",
|
|
||||||
"/owner/repo/src/branch/main/db/schema.sql",
|
|
||||||
"/owner/repo/raw/branch/main/debug.log",
|
|
||||||
"/owner/repo.git/info/refs?service=git-upload-pack",
|
|
||||||
"/owner/repo/src/branch/main/docs/../README.md",
|
|
||||||
"/user/login?redirect_to=%2fowner%2frepo",
|
|
||||||
"/index.php",
|
|
||||||
"/.well-known/security.txt",
|
|
||||||
} {
|
|
||||||
r := get(t, target)
|
|
||||||
r.Header.Set(userAgent, browser)
|
|
||||||
|
|
||||||
matched := files.Match(r)
|
|
||||||
if len(matched) != 0 {
|
|
||||||
t.Errorf("%s matched %v, want no rule", target, ids(matched))
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// A request without a user agent is only noted.
|
|
||||||
wantMatched(t, files, get(t, "/"), "empty-agent")
|
|
||||||
}
|
|
||||||
|
|
||||||
// ruleFiles are files to write into a directory of rule files, by name.
|
|
||||||
type ruleFiles map[string]string
|
|
||||||
|
|
||||||
// writeFiles writes files into a new directory, and returns it.
|
|
||||||
func writeFiles(t *testing.T, files ruleFiles) string {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
dir := t.TempDir()
|
|
||||||
|
|
||||||
for name, content := range files {
|
|
||||||
path := filepath.Join(dir, name)
|
|
||||||
|
|
||||||
err := os.MkdirAll(filepath.Dir(path), 0o700)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("mkdir: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
err = os.WriteFile(path, []byte(content), 0o600)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("write %s: %v", name, err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
return dir
|
|
||||||
}
|
|
||||||
|
|
||||||
// save writes content to the rule file name in dir as an editor that
|
|
||||||
// saves by renaming does, so that the file is never seen half written.
|
|
||||||
func save(t *testing.T, dir, name, content string) {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
path := filepath.Join(dir, name)
|
|
||||||
|
|
||||||
err := os.WriteFile(path+".tmp", []byte(content), 0o600)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("write %s: %v", name, err)
|
|
||||||
}
|
|
||||||
|
|
||||||
err = os.Rename(path+".tmp", path)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("rename: %v", err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// newParams returns Params for the rule files in dir, switched on, with
|
|
||||||
// the process log in the processLog returned, and the alerts waiting in a
|
|
||||||
// queue for a webhook that is never sent them.
|
|
||||||
func newParams(dir string) (rules.Params, processLog) {
|
|
||||||
lines := make(processLog, maxLogLines)
|
|
||||||
|
|
||||||
return rules.Params{
|
|
||||||
Dir: dir,
|
|
||||||
Enabled: true,
|
|
||||||
ProcessLog: slog.New(slog.NewJSONHandler(lines, nil)),
|
|
||||||
Alerts: alerts.New(alerts.Params{
|
|
||||||
WebhookURL: &url.URL{Scheme: "https", Host: "alerts.example"},
|
|
||||||
Events: alerts.Events(),
|
|
||||||
Cooldown: 15 * time.Minute,
|
|
||||||
Now: time.Now,
|
|
||||||
}),
|
|
||||||
}, lines
|
|
||||||
}
|
|
||||||
|
|
||||||
// load writes files into a new directory and loads the rules in it.
|
|
||||||
func load(t *testing.T, files ruleFiles) *rules.Files {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
params, _ := newParams(writeFiles(t, files))
|
|
||||||
params.ProcessLog = slog.New(slog.DiscardHandler)
|
|
||||||
|
|
||||||
loaded, err := rules.Load(params)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("load: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
return loaded
|
|
||||||
}
|
|
||||||
|
|
||||||
// watch loads the rules in dir, runs their Watch until the test ends, and
|
|
||||||
// waits until it watches the directory. It returns the alerts' queue as
|
|
||||||
// well.
|
|
||||||
func watch(t *testing.T, dir string) (*rules.Files, processLog, *alerts.Queue) {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
params, lines := newParams(dir)
|
|
||||||
|
|
||||||
files, err := rules.Load(params)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("load: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
ctx, stop := context.WithCancel(t.Context())
|
|
||||||
stopped := make(chan struct{})
|
|
||||||
|
|
||||||
go func() {
|
|
||||||
files.Watch(ctx)
|
|
||||||
close(stopped)
|
|
||||||
}()
|
|
||||||
|
|
||||||
t.Cleanup(func() {
|
|
||||||
stop()
|
|
||||||
<-stopped
|
|
||||||
})
|
|
||||||
|
|
||||||
lines.waitFor(t, watching)
|
|
||||||
|
|
||||||
return files, lines, params.Alerts
|
|
||||||
}
|
|
||||||
|
|
||||||
// wantRefused checks that loading the rule files in dir fails with the
|
|
||||||
// error want.
|
|
||||||
func wantRefused(t *testing.T, dir, want string) {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
params, _ := newParams(dir)
|
|
||||||
|
|
||||||
_, err := rules.Load(params)
|
|
||||||
if err == nil || err.Error() != want {
|
|
||||||
t.Errorf("error %v, want %s", err, want)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// get returns a GET request for target, a path and an optional query, as
|
|
||||||
// smallwebwaf's server reads it, without a user agent.
|
|
||||||
func get(t *testing.T, target string) *http.Request {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
return httptest.NewRequestWithContext(t.Context(), http.MethodGet,
|
|
||||||
"http://app.example"+target, nil)
|
|
||||||
}
|
|
||||||
|
|
||||||
// wantRefusedBy checks that a GET request for target with the user agent
|
|
||||||
// sent matches rule alone, and that rule refuses it.
|
|
||||||
func wantRefusedBy(t *testing.T, files *rules.Files, target, sent, rule string) {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
r := get(t, target)
|
|
||||||
r.Header.Set(userAgent, sent)
|
|
||||||
|
|
||||||
matched := files.Match(r)
|
|
||||||
if len(matched) != 1 || matched[0].ID != rule || matched[0].Action == rules.ActionLog {
|
|
||||||
t.Errorf("%s from %q matched %v, want %s alone, refusing it", target,
|
|
||||||
sent, ids(matched), rule)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// wantMatched checks the ids of the rules r matches, in order.
|
|
||||||
func wantMatched(t *testing.T, files *rules.Files, r *http.Request, want ...string) {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
got := ids(files.Match(r))
|
|
||||||
if !slices.Equal(got, want) {
|
|
||||||
t.Errorf("%s matched %v, want %v", r.URL, got, want)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// ids returns the ids of matched.
|
|
||||||
func ids(matched []rules.Rule) []string {
|
|
||||||
got := make([]string, 0, len(matched))
|
|
||||||
for _, rule := range matched {
|
|
||||||
got = append(got, rule.ID)
|
|
||||||
}
|
|
||||||
|
|
||||||
return got
|
|
||||||
}
|
|
||||||
|
|
||||||
// processLog receives the lines of a process log, each a JSON object, for
|
|
||||||
// a test to wait for.
|
|
||||||
type processLog chan string
|
|
||||||
|
|
||||||
// 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. 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
|
|
||||||
}
|
|
||||||
|
|
||||||
// waitUntil waits for the rule files to be read until done reports true,
|
|
||||||
// as it does once they have been read after the test's last change. They
|
|
||||||
// can be read before then too, as they are once Watch starts watching.
|
|
||||||
func (l processLog) waitUntil(t *testing.T, done func() bool) {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
for !done() {
|
|
||||||
l.waitFor(t, read)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@@ -1,165 +0,0 @@
|
|||||||
package rules
|
|
||||||
|
|
||||||
import (
|
|
||||||
"context"
|
|
||||||
"log/slog"
|
|
||||||
"net/http"
|
|
||||||
"net/http/httptest"
|
|
||||||
"os"
|
|
||||||
"path/filepath"
|
|
||||||
"slices"
|
|
||||||
"testing"
|
|
||||||
"testing/synctest"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
"github.com/fsnotify/fsnotify"
|
|
||||||
|
|
||||||
"sneak.berlin/go/smallwebwaf/internal/alerts"
|
|
||||||
)
|
|
||||||
|
|
||||||
// The tests below run readAfterChanges 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 readAfterChanges waits again, so that every
|
|
||||||
// reading due by then is done. The test sends the changes itself, as the
|
|
||||||
// watch of a directory cannot run in a bubble.
|
|
||||||
|
|
||||||
func TestFileWrittenInTwoPartsTakenInOnlyWhole(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
synctest.Test(t, func(t *testing.T) {
|
|
||||||
dir := t.TempDir()
|
|
||||||
path := filepath.Join(dir, "50-app.rules")
|
|
||||||
writeFile(t, path, "first path block ^/first\n")
|
|
||||||
files := load(t, dir)
|
|
||||||
changes := run(t, files)
|
|
||||||
|
|
||||||
file, err := os.Create(path) //nolint:gosec // a file the test wrote
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("create: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
defer func() {
|
|
||||||
_ = file.Close()
|
|
||||||
}()
|
|
||||||
|
|
||||||
// The first part ends in the middle of a ban rule's regex, which,
|
|
||||||
// read then, would ban every request.
|
|
||||||
write(t, file, "first path block ^/first\nprobe path ban ^/")
|
|
||||||
|
|
||||||
changes <- fsnotify.Event{Name: path, Op: fsnotify.Write}
|
|
||||||
|
|
||||||
time.Sleep(quietTime - time.Nanosecond)
|
|
||||||
synctest.Wait()
|
|
||||||
wantMatched(t, files, "/anything")
|
|
||||||
|
|
||||||
// The second part starts the wait again.
|
|
||||||
write(t, file, `\.env$`+"\n")
|
|
||||||
|
|
||||||
changes <- fsnotify.Event{Name: path, Op: fsnotify.Write}
|
|
||||||
|
|
||||||
time.Sleep(quietTime - time.Nanosecond)
|
|
||||||
synctest.Wait()
|
|
||||||
wantMatched(t, files, "/.env")
|
|
||||||
|
|
||||||
time.Sleep(time.Nanosecond)
|
|
||||||
synctest.Wait()
|
|
||||||
wantMatched(t, files, "/.env", "probe")
|
|
||||||
wantMatched(t, files, "/anything")
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestEditSavedBeforeTheWatchStartsTakenIn(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
synctest.Test(t, func(t *testing.T) {
|
|
||||||
dir := t.TempDir()
|
|
||||||
path := filepath.Join(dir, "50-app.rules")
|
|
||||||
writeFile(t, path, "first path block ^/first\n")
|
|
||||||
files := load(t, dir)
|
|
||||||
|
|
||||||
// Saved after Load read the files, and before the directory was
|
|
||||||
// watched, so that no change is seen for it.
|
|
||||||
writeFile(t, path, "first path block ^/edited\n")
|
|
||||||
run(t, files)
|
|
||||||
time.Sleep(quietTime)
|
|
||||||
synctest.Wait()
|
|
||||||
wantMatched(t, files, "/edited", "first")
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
// load loads the rules in dir.
|
|
||||||
func load(t *testing.T, dir string) *Files {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
files, err := Load(Params{
|
|
||||||
Dir: dir, Enabled: true, ProcessLog: slog.New(slog.DiscardHandler),
|
|
||||||
Alerts: alerts.New(alerts.Params{}),
|
|
||||||
})
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("load: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
return files
|
|
||||||
}
|
|
||||||
|
|
||||||
// run runs files' readAfterChanges until the test ends, and returns the
|
|
||||||
// channel that sends it changes.
|
|
||||||
func run(t *testing.T, files *Files) chan<- fsnotify.Event {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
changes := make(chan fsnotify.Event)
|
|
||||||
ctx, stop := context.WithCancel(t.Context())
|
|
||||||
stopped := make(chan struct{})
|
|
||||||
|
|
||||||
go func() {
|
|
||||||
files.readAfterChanges(ctx, changes, nil)
|
|
||||||
close(stopped)
|
|
||||||
}()
|
|
||||||
|
|
||||||
t.Cleanup(func() {
|
|
||||||
stop()
|
|
||||||
<-stopped
|
|
||||||
})
|
|
||||||
|
|
||||||
return changes
|
|
||||||
}
|
|
||||||
|
|
||||||
// writeFile writes content to the file at path.
|
|
||||||
func writeFile(t *testing.T, path, content string) {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
err := os.WriteFile(path, []byte(content), 0o600)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("write %s: %v", path, err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// write writes text to the end of file.
|
|
||||||
func write(t *testing.T, file *os.File, text string) {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
_, err := file.WriteString(text)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("write: %v", err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// wantMatched checks the ids of the rules that a GET request for path
|
|
||||||
// matches, in order.
|
|
||||||
func wantMatched(t *testing.T, files *Files, path string, want ...string) {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
r := httptest.NewRequestWithContext(t.Context(), http.MethodGet,
|
|
||||||
"http://app.example"+path, nil)
|
|
||||||
|
|
||||||
matched := files.Match(r)
|
|
||||||
|
|
||||||
got := make([]string, 0, len(matched))
|
|
||||||
for _, rule := range matched {
|
|
||||||
got = append(got, rule.ID)
|
|
||||||
}
|
|
||||||
|
|
||||||
if !slices.Equal(got, want) {
|
|
||||||
t.Errorf("%s matched %v, want %v", path, got, want)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@@ -1,102 +0,0 @@
|
|||||||
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.
|
|
||||||
// It reads no other setting, nor a file that another names, so neither
|
|
||||||
// can fail it.
|
|
||||||
// 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()
|
|
||||||
|
|
||||||
listenAddr, upstreamURL, err := config.ListenAddrAndUpstreamURL(lookupEnv)
|
|
||||||
if err != nil {
|
|
||||||
return fmt.Errorf("invalid setting: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
// The settings have checked that the address has a port.
|
|
||||||
_, port, _ := net.SplitHostPort(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(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)
|
|
||||||
}
|
|
||||||
@@ -1,27 +0,0 @@
|
|||||||
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)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@@ -1,132 +0,0 @@
|
|||||||
package smallwebwaf_test
|
|
||||||
|
|
||||||
import (
|
|
||||||
"bytes"
|
|
||||||
"context"
|
|
||||||
"net"
|
|
||||||
"net/http"
|
|
||||||
"net/http/httptest"
|
|
||||||
"os"
|
|
||||||
"path/filepath"
|
|
||||||
"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(),
|
|
||||||
rulesDir: 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, "")
|
|
||||||
|
|
||||||
// The health check reads those two settings alone, here given as
|
|
||||||
// files: a removed or invalid token file, or an invalid value of
|
|
||||||
// another setting, does not fail it.
|
|
||||||
for _, other := range []struct{ name, value string }{
|
|
||||||
{"SWWAF_METRICS_TOKEN_FILE", filepath.Join(t.TempDir(), "removed")},
|
|
||||||
{"SWWAF_METRICS_TOKEN_FILE", writeFile(t, "too short\n")},
|
|
||||||
{"SWWAF_MODE", "neither"},
|
|
||||||
} {
|
|
||||||
wantHealthCheck(t, map[string]string{
|
|
||||||
listenAddr + "_FILE": writeFile(t, ":"+port+"\n"),
|
|
||||||
upstreamURL + "_FILE": writeFile(t, app.URL+"\n"),
|
|
||||||
other.name: other.value,
|
|
||||||
}, 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)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// writeFile writes contents to a file in a directory of its own, removed
|
|
||||||
// when the test ends, and returns the file's path.
|
|
||||||
func writeFile(t *testing.T, contents string) string {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
path := filepath.Join(t.TempDir(), "setting")
|
|
||||||
|
|
||||||
err := os.WriteFile(path, []byte(contents), 0o600)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("write %s: %v", path, err)
|
|
||||||
}
|
|
||||||
|
|
||||||
return path
|
|
||||||
}
|
|
||||||
@@ -1,7 +1,6 @@
|
|||||||
// Package smallwebwaf runs the smallwebwaf process: it reads the settings,
|
// Package smallwebwaf runs the smallwebwaf process: it reads the settings,
|
||||||
// the rule files, the lookup database and the state files, serves requests
|
// serves requests until it is told to stop, and then stops in an orderly
|
||||||
// until it is told to stop, and then stops in an orderly way, writing the
|
// way.
|
||||||
// state files.
|
|
||||||
package smallwebwaf
|
package smallwebwaf
|
||||||
|
|
||||||
import (
|
import (
|
||||||
@@ -16,14 +15,9 @@ import (
|
|||||||
"syscall"
|
"syscall"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"sneak.berlin/go/smallwebwaf/internal/alerts"
|
|
||||||
"sneak.berlin/go/smallwebwaf/internal/config"
|
"sneak.berlin/go/smallwebwaf/internal/config"
|
||||||
"sneak.berlin/go/smallwebwaf/internal/lookup"
|
|
||||||
"sneak.berlin/go/smallwebwaf/internal/proxy"
|
"sneak.berlin/go/smallwebwaf/internal/proxy"
|
||||||
"sneak.berlin/go/smallwebwaf/internal/remotelog"
|
|
||||||
"sneak.berlin/go/smallwebwaf/internal/requestlog"
|
"sneak.berlin/go/smallwebwaf/internal/requestlog"
|
||||||
"sneak.berlin/go/smallwebwaf/internal/rules"
|
|
||||||
"sneak.berlin/go/smallwebwaf/internal/state"
|
|
||||||
)
|
)
|
||||||
|
|
||||||
// shutdownTimeout is how long requests in progress may take to finish
|
// shutdownTimeout is how long requests in progress may take to finish
|
||||||
@@ -31,11 +25,6 @@ import (
|
|||||||
// runit and docker wait a little longer before they kill the process.
|
// runit and docker wait a little longer before they kill the process.
|
||||||
const shutdownTimeout = 5 * time.Second
|
const shutdownTimeout = 5 * time.Second
|
||||||
|
|
||||||
// remoteLogStopTimeout is how long, as smallwebwaf stops, the log lines
|
|
||||||
// still waiting are sent to SWWAF_LOG_REMOTE_URL before they are given
|
|
||||||
// up. stdout has carried them.
|
|
||||||
const remoteLogStopTimeout = 2 * time.Second
|
|
||||||
|
|
||||||
// Params are what Run needs from the process.
|
// Params are what Run needs from the process.
|
||||||
type Params struct {
|
type Params struct {
|
||||||
// Version is the version of the binary, set when it is built.
|
// Version is the version of the binary, set when it is built.
|
||||||
@@ -47,13 +36,8 @@ type Params struct {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// Main runs smallwebwaf until SIGTERM or SIGINT, and returns the
|
// Main runs smallwebwaf until SIGTERM or SIGINT, and returns the
|
||||||
// process's exit status. Run as `smallwebwaf healthcheck`, it is the
|
// process's exit status.
|
||||||
// container's health check instead.
|
|
||||||
func Main(version string) int {
|
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(),
|
ctx, stop := signal.NotifyContext(context.Background(),
|
||||||
syscall.SIGTERM, os.Interrupt)
|
syscall.SIGTERM, os.Interrupt)
|
||||||
defer stop()
|
defer stop()
|
||||||
@@ -65,12 +49,10 @@ func Main(version string) int {
|
|||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
// Run reads the settings, the rule files, the lookup database and the
|
// Run reads the settings, then serves requests until ctx is done. It
|
||||||
// state files, then serves requests until ctx is done. It returns the
|
// returns the process's exit status, 1 when smallwebwaf cannot start.
|
||||||
// process's exit status, 1 when smallwebwaf cannot start.
|
|
||||||
func Run(ctx context.Context, params Params) int {
|
func Run(ctx context.Context, params Params) int {
|
||||||
processLog := requestlog.NewProcessLogger(params.Stdout,
|
processLog := requestlog.NewProcessLogger(params.Stdout)
|
||||||
config.InstanceName(params.LookupEnv))
|
|
||||||
|
|
||||||
cfg, err := config.FromEnvironment(params.LookupEnv)
|
cfg, err := config.FromEnvironment(params.LookupEnv)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -79,56 +61,6 @@ func Run(ctx context.Context, params Params) int {
|
|||||||
return 1
|
return 1
|
||||||
}
|
}
|
||||||
|
|
||||||
// While SWWAF_LOG_REMOTE_URL is set, every line on stdout from here on
|
|
||||||
// is sent there too.
|
|
||||||
stdout := params.Stdout
|
|
||||||
|
|
||||||
var remote *remotelog.Sender
|
|
||||||
|
|
||||||
if cfg.LogRemoteURL != nil {
|
|
||||||
remote = newRemoteLogSender(cfg)
|
|
||||||
stdout = io.MultiWriter(params.Stdout, remote)
|
|
||||||
processLog = requestlog.NewProcessLogger(stdout, cfg.InstanceName)
|
|
||||||
|
|
||||||
stopSending := startSending(ctx, remote, processLog)
|
|
||||||
defer stopSending()
|
|
||||||
}
|
|
||||||
|
|
||||||
// The state files and the alerts give times in UTC.
|
|
||||||
now := func() time.Time { return time.Now().UTC() }
|
|
||||||
|
|
||||||
alertQueue := newAlertQueue(cfg, now, processLog)
|
|
||||||
|
|
||||||
ruleFiles, err := rules.Load(rules.Params{
|
|
||||||
Dir: cfg.RulesDir,
|
|
||||||
Enabled: cfg.RulesEnabled,
|
|
||||||
ProcessLog: processLog,
|
|
||||||
Alerts: alertQueue,
|
|
||||||
})
|
|
||||||
if err != nil {
|
|
||||||
processLog.Error("cannot use the rule files", "error", err.Error())
|
|
||||||
|
|
||||||
return 1
|
|
||||||
}
|
|
||||||
|
|
||||||
server, err := newServer(cfg, stdout, processLog, now, ruleFiles, alertQueue)
|
|
||||||
if err != nil {
|
|
||||||
processLog.Error("cannot use the lookup database", "error", err.Error())
|
|
||||||
|
|
||||||
return 1
|
|
||||||
}
|
|
||||||
|
|
||||||
if remote != nil {
|
|
||||||
server.Metrics.AddRemoteLog(remote)
|
|
||||||
}
|
|
||||||
|
|
||||||
files, err := loadStateFiles(cfg, server, alertQueue, now, processLog)
|
|
||||||
if err != nil {
|
|
||||||
processLog.Error("cannot use the state files", "error", err.Error())
|
|
||||||
|
|
||||||
return 1
|
|
||||||
}
|
|
||||||
|
|
||||||
listener, err := (&net.ListenConfig{}).Listen(ctx, "tcp", cfg.ListenAddr)
|
listener, err := (&net.ListenConfig{}).Listen(ctx, "tcp", cfg.ListenAddr)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
processLog.Error("cannot listen on SWWAF_LISTEN_ADDR",
|
processLog.Error("cannot listen on SWWAF_LISTEN_ADDR",
|
||||||
@@ -137,142 +69,24 @@ func Run(ctx context.Context, params Params) int {
|
|||||||
return 1
|
return 1
|
||||||
}
|
}
|
||||||
|
|
||||||
|
server := proxy.New(proxy.Params{
|
||||||
|
Config: cfg,
|
||||||
|
RequestLog: params.Stdout,
|
||||||
|
ProcessLog: processLog,
|
||||||
|
})
|
||||||
|
|
||||||
processLog.Info("starting",
|
processLog.Info("starting",
|
||||||
"version", params.Version,
|
"version", params.Version,
|
||||||
"address", listener.Addr().String(),
|
"address", listener.Addr().String(),
|
||||||
"settings", cfg)
|
"settings", cfg)
|
||||||
|
|
||||||
return serve(ctx, server, listener, files, ruleFiles, alertQueue, processLog)
|
return serve(ctx, server, listener, processLog)
|
||||||
}
|
}
|
||||||
|
|
||||||
// newServer returns the server smallwebwaf runs, with the metrics of the
|
// serve serves requests on listener until ctx is done, then gives the
|
||||||
// alerts, after reading the lookup database while SWWAF_LOOKUP_SOURCE is
|
// requests in progress shutdownTimeout to finish.
|
||||||
// file. A lookup database that cannot be read is an error.
|
|
||||||
func newServer(
|
|
||||||
cfg *config.Config, stdout io.Writer, processLog *slog.Logger,
|
|
||||||
now func() time.Time, ruleFiles *rules.Files, alertQueue *alerts.Queue,
|
|
||||||
) (*proxy.Server, error) {
|
|
||||||
var lookupFile *lookup.File
|
|
||||||
|
|
||||||
if cfg.LookupSource == "file" {
|
|
||||||
var err error
|
|
||||||
|
|
||||||
lookupFile, err = lookup.OpenFile(lookup.FileParams{
|
|
||||||
Path: cfg.LookupDBPath,
|
|
||||||
Now: now,
|
|
||||||
ProcessLog: processLog,
|
|
||||||
Alerts: alertQueue,
|
|
||||||
})
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
server := proxy.New(proxy.Params{
|
|
||||||
Config: cfg,
|
|
||||||
RequestLog: stdout,
|
|
||||||
ProcessLog: processLog,
|
|
||||||
GeoJSURL: lookup.URL,
|
|
||||||
LookupFile: lookupFile,
|
|
||||||
Now: now,
|
|
||||||
Rules: ruleFiles,
|
|
||||||
Alerts: alertQueue,
|
|
||||||
})
|
|
||||||
server.Metrics.AddAlerts(alertQueue)
|
|
||||||
|
|
||||||
if lookupFile != nil {
|
|
||||||
server.Metrics.AddLookupFile(lookupFile.LastRead, lookupFile.ReadFailures)
|
|
||||||
}
|
|
||||||
|
|
||||||
return server, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// loadStateFiles reads the state files into the parts of server and into
|
|
||||||
// alertQueue, as state.Load does.
|
|
||||||
func loadStateFiles(
|
|
||||||
cfg *config.Config, server *proxy.Server, alertQueue *alerts.Queue,
|
|
||||||
now func() time.Time, processLog *slog.Logger,
|
|
||||||
) (*state.Files, error) {
|
|
||||||
return state.Load(state.Params{
|
|
||||||
Dir: cfg.StateDir,
|
|
||||||
WriteDelay: cfg.StateWriteDelay,
|
|
||||||
CounterInterval: cfg.StateCounterInterval,
|
|
||||||
Ledger: server.Ledger,
|
|
||||||
Limiter: server.Limiter,
|
|
||||||
GeoJS: server.GeoJS,
|
|
||||||
Alerts: alertQueue,
|
|
||||||
Now: now,
|
|
||||||
ProcessLog: processLog,
|
|
||||||
Metrics: server.Metrics,
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
// newAlertQueue returns the queue of the alerts to the webhook, Slack and
|
|
||||||
// ntfy, with the settings for them.
|
|
||||||
func newAlertQueue(
|
|
||||||
cfg *config.Config, now func() time.Time, processLog *slog.Logger,
|
|
||||||
) *alerts.Queue {
|
|
||||||
return alerts.New(alerts.Params{
|
|
||||||
WebhookURL: cfg.AlertWebhookURL,
|
|
||||||
WebhookHeaders: cfg.AlertWebhookHeaders,
|
|
||||||
SlackURL: cfg.AlertSlackWebhookURL,
|
|
||||||
NtfyURL: cfg.AlertNtfyURL,
|
|
||||||
NtfyToken: cfg.AlertNtfyToken,
|
|
||||||
Events: cfg.AlertEvents,
|
|
||||||
Cooldown: cfg.AlertCooldown,
|
|
||||||
MaxPerHour: cfg.AlertMaxPerHour,
|
|
||||||
Instance: cfg.InstanceName,
|
|
||||||
Now: now,
|
|
||||||
ProcessLog: processLog,
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
// newRemoteLogSender returns a sender of the log lines to
|
|
||||||
// SWWAF_LOG_REMOTE_URL, with the settings for it.
|
|
||||||
func newRemoteLogSender(cfg *config.Config) *remotelog.Sender {
|
|
||||||
return remotelog.New(remotelog.Params{
|
|
||||||
URL: cfg.LogRemoteURL,
|
|
||||||
RootCAs: cfg.LogRemoteTLSCAs,
|
|
||||||
Buffer: cfg.LogRemoteBuffer,
|
|
||||||
Facility: cfg.LogRemoteFacility,
|
|
||||||
AppName: cfg.LogRemoteAppName,
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
// startSending runs remote until the function it returns is called, which
|
|
||||||
// then waits at most remoteLogStopTimeout for the lines still waiting to
|
|
||||||
// be sent. Sending goes on after ctx is done, so that the lines written
|
|
||||||
// while smallwebwaf stops are sent too.
|
|
||||||
func startSending(
|
|
||||||
ctx context.Context, remote *remotelog.Sender, processLog *slog.Logger,
|
|
||||||
) func() {
|
|
||||||
sending, stop := context.WithCancel(context.WithoutCancel(ctx))
|
|
||||||
sent := make(chan struct{})
|
|
||||||
|
|
||||||
go func() {
|
|
||||||
remote.Run(sending, processLog)
|
|
||||||
close(sent)
|
|
||||||
}()
|
|
||||||
|
|
||||||
return func() {
|
|
||||||
stop()
|
|
||||||
|
|
||||||
select {
|
|
||||||
case <-sent:
|
|
||||||
case <-time.After(remoteLogStopTimeout):
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// serve serves requests on listener, writes the state files as they are
|
|
||||||
// due, takes in an admin's edits of them, reads the rule files again as
|
|
||||||
// they change, and the lookup database when it is replaced, and sends the
|
|
||||||
// alerts, until ctx is done. Then it gives the requests in progress
|
|
||||||
// shutdownTimeout to finish, and writes every state file, alerts.json with
|
|
||||||
// the alerts still waiting.
|
|
||||||
func serve(
|
func serve(
|
||||||
ctx context.Context, server *proxy.Server, listener net.Listener,
|
ctx context.Context, server *http.Server, listener net.Listener,
|
||||||
files *state.Files, ruleFiles *rules.Files, alertQueue *alerts.Queue,
|
|
||||||
processLog *slog.Logger,
|
processLog *slog.Logger,
|
||||||
) int {
|
) int {
|
||||||
served := make(chan error, 1)
|
served := make(chan error, 1)
|
||||||
@@ -281,19 +95,6 @@ func serve(
|
|||||||
served <- server.Serve(listener)
|
served <- server.Serve(listener)
|
||||||
}()
|
}()
|
||||||
|
|
||||||
writing, stopWriting := context.WithCancel(ctx)
|
|
||||||
defer stopWriting()
|
|
||||||
|
|
||||||
written := inBackground(func() { files.Run(writing) })
|
|
||||||
watched := inBackground(func() { files.Watch(writing) })
|
|
||||||
rulesWatched := inBackground(func() { ruleFiles.Watch(writing) })
|
|
||||||
lookupFileWatched := inBackground(func() {
|
|
||||||
if server.LookupFile != nil {
|
|
||||||
server.LookupFile.Watch(writing)
|
|
||||||
}
|
|
||||||
})
|
|
||||||
alertsSent := inBackground(func() { alertQueue.Run(writing) })
|
|
||||||
|
|
||||||
select {
|
select {
|
||||||
case err := <-served:
|
case err := <-served:
|
||||||
processLog.Error("serving failed", "error", err.Error())
|
processLog.Error("serving failed", "error", err.Error())
|
||||||
@@ -323,41 +124,7 @@ func serve(
|
|||||||
return 1
|
return 1
|
||||||
}
|
}
|
||||||
|
|
||||||
// Run and Watch have ended, so nothing else reads or writes the
|
|
||||||
// files, and no alert is being sent, so that alerts.json keeps every
|
|
||||||
// alert not yet sent. 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
|
|
||||||
<-rulesWatched
|
|
||||||
<-lookupFileWatched
|
|
||||||
<-alertsSent
|
|
||||||
|
|
||||||
err = files.WriteAll()
|
|
||||||
if err != nil {
|
|
||||||
processLog.Error("writing the state files failed", "error", err.Error())
|
|
||||||
|
|
||||||
return 1
|
|
||||||
}
|
|
||||||
|
|
||||||
processLog.Info("stopped")
|
processLog.Info("stopped")
|
||||||
|
|
||||||
return 0
|
return 0
|
||||||
}
|
}
|
||||||
|
|
||||||
// inBackground runs task on a goroutine of its own, and returns a channel
|
|
||||||
// that is closed once task has returned.
|
|
||||||
func inBackground(task func()) <-chan struct{} {
|
|
||||||
done := make(chan struct{})
|
|
||||||
|
|
||||||
go func() {
|
|
||||||
task()
|
|
||||||
close(done)
|
|
||||||
}()
|
|
||||||
|
|
||||||
return done
|
|
||||||
}
|
|
||||||
|
|||||||
@@ -1,82 +0,0 @@
|
|||||||
package smallwebwaf
|
|
||||||
|
|
||||||
import (
|
|
||||||
"log/slog"
|
|
||||||
"net/url"
|
|
||||||
"testing"
|
|
||||||
"testing/synctest"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
"sneak.berlin/go/smallwebwaf/internal/remotelog"
|
|
||||||
)
|
|
||||||
|
|
||||||
// The stop's tests run in a synctest bubble, where the time package runs
|
|
||||||
// on a clock of the test's own, so that how long the stop takes can be
|
|
||||||
// told exactly. The sender is held up by its process log, not by the
|
|
||||||
// network: a goroutine of the bubble that waits on the network keeps that
|
|
||||||
// clock from moving on.
|
|
||||||
|
|
||||||
func TestStopWaitsForTheSenderToFinish(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
synctest.Test(t, func(t *testing.T) {
|
|
||||||
took := stopHeldSender(t, time.Second)
|
|
||||||
if took != time.Second {
|
|
||||||
t.Errorf("the stop took %s, want the second the sender took", took)
|
|
||||||
}
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestStopWaitsForTheSenderAtMostTwoSeconds(t *testing.T) {
|
|
||||||
t.Parallel()
|
|
||||||
|
|
||||||
synctest.Test(t, func(t *testing.T) {
|
|
||||||
took := stopHeldSender(t, time.Minute)
|
|
||||||
if took != 2*time.Second {
|
|
||||||
t.Errorf("the stop took %s, want 2s", took)
|
|
||||||
}
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
// heldLog holds each line written to it until it is closed.
|
|
||||||
type heldLog chan struct{}
|
|
||||||
|
|
||||||
// Write waits until the log is closed.
|
|
||||||
func (l heldLog) Write(p []byte) (int, error) {
|
|
||||||
<-l
|
|
||||||
|
|
||||||
return len(p), nil
|
|
||||||
}
|
|
||||||
|
|
||||||
// stopHeldSender starts sending to an endpoint the sender cannot connect
|
|
||||||
// to, holds the sender as it logs that failure until release has passed,
|
|
||||||
// stops the sending, and returns how long the stop took. It returns once
|
|
||||||
// the sender has ended, as a bubble must.
|
|
||||||
func stopHeldSender(t *testing.T, release time.Duration) time.Duration {
|
|
||||||
t.Helper()
|
|
||||||
|
|
||||||
log := make(heldLog)
|
|
||||||
sender := remotelog.New(remotelog.Params{
|
|
||||||
// No port is 65536, so each attempt to connect fails at once,
|
|
||||||
// before it reaches the network.
|
|
||||||
URL: &url.URL{Scheme: remotelog.SchemeTCP, Host: "127.0.0.1:65536"},
|
|
||||||
Buffer: 1,
|
|
||||||
})
|
|
||||||
|
|
||||||
stopSending := startSending(t.Context(), sender,
|
|
||||||
slog.New(slog.NewJSONHandler(log, nil)))
|
|
||||||
|
|
||||||
synctest.Wait()
|
|
||||||
time.AfterFunc(release, func() { close(log) })
|
|
||||||
|
|
||||||
stopped := time.Now()
|
|
||||||
|
|
||||||
stopSending()
|
|
||||||
|
|
||||||
took := time.Since(stopped)
|
|
||||||
|
|
||||||
time.Sleep(release)
|
|
||||||
synctest.Wait()
|
|
||||||
|
|
||||||
return took
|
|
||||||
}
|
|
||||||
File diff suppressed because it is too large
Load Diff
@@ -1,851 +0,0 @@
|
|||||||
// 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, lookups.json GeoJS's answers, and alerts.json the cooldowns,
|
|
||||||
// the hour under way and the alerts waiting for each destination. 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"
|
|
||||||
"maps"
|
|
||||||
"net/netip"
|
|
||||||
"os"
|
|
||||||
"path/filepath"
|
|
||||||
"slices"
|
|
||||||
"sync"
|
|
||||||
"time"
|
|
||||||
|
|
||||||
"github.com/fsnotify/fsnotify"
|
|
||||||
"sneak.berlin/go/smallwebwaf/internal/alerts"
|
|
||||||
"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"
|
|
||||||
alertsJSON = "alerts.json"
|
|
||||||
)
|
|
||||||
|
|
||||||
var (
|
|
||||||
errVersion = errors.New("unknown version")
|
|
||||||
// errMissing is for an entry without a field it needs.
|
|
||||||
errMissing = errors.New("has no")
|
|
||||||
errCause = errors.New("is not limit, attack or admin")
|
|
||||||
errDestination = errors.New("is not webhook, slack or ntfy")
|
|
||||||
errWaitingList = errors.New(`waiting is a list, but now lists the alerts by ` +
|
|
||||||
`destination: put the list under "webhook", as "waiting": {"webhook": [...]}, ` +
|
|
||||||
`or remove the file`)
|
|
||||||
)
|
|
||||||
|
|
||||||
// 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, GeoJS and Alerts hold the state. Alerts also
|
|
||||||
// receive a file_error alert for an edit set aside, and for a write
|
|
||||||
// that fails while smallwebwaf runs.
|
|
||||||
Ledger *bans.Ledger
|
|
||||||
Limiter *ratelimit.Limiter
|
|
||||||
GeoJS *lookup.GeoJS
|
|
||||||
Alerts *alerts.Queue
|
|
||||||
// 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, a ban an admin added may have no cause, which makes it an
|
|
||||||
// admin's, and lifted is left out until an admin lifts the ban. The ban
|
|
||||||
// endpoints answer with bans in this form too.
|
|
||||||
type BanEntry struct {
|
|
||||||
Netblock netip.Prefix `json:"netblock"`
|
|
||||||
Start time.Time `json:"start"`
|
|
||||||
Expires *time.Time `json:"expires"`
|
|
||||||
Cause string `json:"cause"`
|
|
||||||
Reason string `json:"reason,omitempty"`
|
|
||||||
Lifted *time.Time `json:"lifted,omitempty"`
|
|
||||||
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"`
|
|
||||||
}
|
|
||||||
|
|
||||||
// alertsFile is alerts.json, indented for an admin to read and edit.
|
|
||||||
type alertsFile struct {
|
|
||||||
Version int `json:"version"`
|
|
||||||
Cooldowns []alerts.Cooldown `json:"cooldowns"`
|
|
||||||
Hour alerts.Hour `json:"hour"`
|
|
||||||
Waiting map[string][]alerts.Alert `json:"waiting"`
|
|
||||||
}
|
|
||||||
|
|
||||||
// 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)
|
|
||||||
alertsRead, alertsErr := f.read(alertsJSON)
|
|
||||||
|
|
||||||
err = errors.Join(bansErr, clientsErr, lookupsErr, alertsErr)
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
params.ProcessLog.Info("read the state files", "directory", params.Dir,
|
|
||||||
"bans", bansRead, "clients", clientsRead, "lookups", lookupsRead,
|
|
||||||
"alerts_waiting", alertsRead)
|
|
||||||
|
|
||||||
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, raised as a file_error alert, 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(bansJSON, f.writeFile(bansJSON))
|
|
||||||
case <-interval.C:
|
|
||||||
for _, name := range []string{bansJSON, clientsJSON, lookupsJSON, alertsJSON} {
|
|
||||||
f.logFailure(name, f.writeFile(name))
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// 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), f.writeFile(alertsJSON))
|
|
||||||
}
|
|
||||||
|
|
||||||
// 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, alertsJSON:
|
|
||||||
f.fileChanged(name)
|
|
||||||
}
|
|
||||||
case err = <-watcher.Errors:
|
|
||||||
f.params.ProcessLog.Warn("watching the state files failed",
|
|
||||||
"error", err.Error())
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// logFailure logs a write of the state file name that failed, and raises
|
|
||||||
// a file_error alert for it.
|
|
||||||
func (f *Files) logFailure(name string, err error) {
|
|
||||||
if err != nil {
|
|
||||||
const failed = "writing the state files failed"
|
|
||||||
|
|
||||||
// Raised before it is logged, so that the alert is there once the
|
|
||||||
// log line is.
|
|
||||||
f.params.Alerts.Raise(alerts.Alert{
|
|
||||||
Event: alerts.EventFileError,
|
|
||||||
Reason: failed,
|
|
||||||
Detail: map[string]any{
|
|
||||||
"file": filepath.Join(f.params.Dir, name), "error": err.Error(),
|
|
||||||
},
|
|
||||||
})
|
|
||||||
f.params.ProcessLog.Error(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, true)
|
|
||||||
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, false)
|
|
||||||
}
|
|
||||||
|
|
||||||
// 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. edit is whether data is an admin's
|
|
||||||
// edit taken in while smallwebwaf runs, rather than the file read at the
|
|
||||||
// start. 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, edit bool) (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())
|
|
||||||
}
|
|
||||||
|
|
||||||
if edit {
|
|
||||||
f.params.Ledger.LoadEdit(held)
|
|
||||||
} else {
|
|
||||||
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)
|
|
||||||
case alertsJSON:
|
|
||||||
// waiting was a list, of the alerts waiting for the webhook, before
|
|
||||||
// alerts went to Slack and ntfy too.
|
|
||||||
var written struct {
|
|
||||||
Waiting json.RawMessage `json:"waiting"`
|
|
||||||
}
|
|
||||||
|
|
||||||
if json.Unmarshal(data, &written) == nil &&
|
|
||||||
bytes.HasPrefix(written.Waiting, []byte("[")) {
|
|
||||||
return 0, fmt.Errorf("%s: %w", path, errWaitingList)
|
|
||||||
}
|
|
||||||
|
|
||||||
var file alertsFile
|
|
||||||
|
|
||||||
err := parse(path, data, &file)
|
|
||||||
if err != nil {
|
|
||||||
return 0, err
|
|
||||||
}
|
|
||||||
|
|
||||||
f.params.Alerts.Load(alerts.State{
|
|
||||||
Cooldowns: file.Cooldowns, Hour: file.Hour, Waiting: file.Waiting,
|
|
||||||
})
|
|
||||||
|
|
||||||
for _, waiting := range file.Waiting {
|
|
||||||
entries += len(waiting)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
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, and raises a file_error alert for it. 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)
|
|
||||||
}
|
|
||||||
|
|
||||||
const setAside = "set aside an edit of a state file that does not parse"
|
|
||||||
|
|
||||||
// Raised before it is logged, so that the alert is there once the log
|
|
||||||
// line is.
|
|
||||||
f.params.Alerts.Raise(alerts.Alert{
|
|
||||||
Event: alerts.EventFileError,
|
|
||||||
Reason: setAside,
|
|
||||||
Detail: map[string]any{"file": path + ".bad", "error": parseErr.Error()},
|
|
||||||
})
|
|
||||||
f.params.ProcessLog.Error(setAside, "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:
|
|
||||||
file := bansFile{Version: version, Bans: BanEntries(f.params.Ledger.Snapshot())}
|
|
||||||
|
|
||||||
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())
|
|
||||||
case lookupsJSON:
|
|
||||||
return encodeOnePerLine("lookups", f.params.GeoJS.Snapshot())
|
|
||||||
default: // alerts.json
|
|
||||||
held := f.params.Alerts.Snapshot()
|
|
||||||
file := alertsFile{
|
|
||||||
Version: version, Cooldowns: held.Cooldowns, Hour: held.Hour,
|
|
||||||
Waiting: held.Waiting,
|
|
||||||
}
|
|
||||||
|
|
||||||
data, err := json.MarshalIndent(file, "", " ")
|
|
||||||
if err != nil {
|
|
||||||
return nil, err
|
|
||||||
}
|
|
||||||
|
|
||||||
return append(data, '\n'), nil
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// BanEntries returns held as bans.json lists them, an empty list for
|
|
||||||
// none.
|
|
||||||
func BanEntries(held []bans.Ban) []BanEntry {
|
|
||||||
entries := make([]BanEntry, 0, len(held))
|
|
||||||
for _, ban := range held {
|
|
||||||
entries = append(entries, newBanEntry(ban))
|
|
||||||
}
|
|
||||||
|
|
||||||
return entries
|
|
||||||
}
|
|
||||||
|
|
||||||
// newBanEntry returns ban as bans.json holds it.
|
|
||||||
func newBanEntry(ban bans.Ban) BanEntry {
|
|
||||||
entry := BanEntry{
|
|
||||||
Netblock: ban.Netblock, Start: ban.Start, Cause: ban.Cause, Reason: ban.Reason,
|
|
||||||
Notes: ban.Notes,
|
|
||||||
}
|
|
||||||
if !ban.Permanent() {
|
|
||||||
entry.Expires = &ban.Expires
|
|
||||||
}
|
|
||||||
|
|
||||||
if !ban.Lifted.IsZero() {
|
|
||||||
entry.Lifted = &ban.Lifted
|
|
||||||
}
|
|
||||||
|
|
||||||
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, Cause: e.Cause, Reason: e.Reason,
|
|
||||||
Notes: e.Notes,
|
|
||||||
}
|
|
||||||
if e.Expires != nil {
|
|
||||||
ban.Expires = *e.Expires
|
|
||||||
}
|
|
||||||
|
|
||||||
if e.Lifted != nil {
|
|
||||||
ban.Lifted = *e.Lifted
|
|
||||||
}
|
|
||||||
|
|
||||||
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. A cause other than limit,
|
|
||||||
// attack or admin, most likely misspelt, is refused too.
|
|
||||||
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")
|
|
||||||
case entry.Cause != "" && entry.Cause != bans.CauseLimit &&
|
|
||||||
entry.Cause != bans.CauseAttack && entry.Cause != bans.CauseAdmin:
|
|
||||||
return fmt.Errorf("entry %d's cause %q %w", i+1, entry.Cause, errCause)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
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
|
|
||||||
}
|
|
||||||
|
|
||||||
// check refuses a cooldown without its event or when its alert was sent,
|
|
||||||
// which would hold back no repeat, alerts waiting for a destination with
|
|
||||||
// another name than webhook, slack or ntfy, most likely misspelt, and an
|
|
||||||
// alert waiting without its event or its time.
|
|
||||||
func (f *alertsFile) check([]byte) error {
|
|
||||||
for i, cooldown := range f.Cooldowns {
|
|
||||||
switch {
|
|
||||||
case cooldown.Event == "":
|
|
||||||
return fmt.Errorf("cooldowns %w", missing(i, "event"))
|
|
||||||
case cooldown.Sent.IsZero():
|
|
||||||
return fmt.Errorf("cooldowns %w", missing(i, "sent"))
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
for _, destination := range slices.Sorted(maps.Keys(f.Waiting)) {
|
|
||||||
if !slices.Contains(alerts.Destinations(), destination) {
|
|
||||||
return fmt.Errorf("waiting %q %w", destination, errDestination)
|
|
||||||
}
|
|
||||||
|
|
||||||
for i, alert := range f.Waiting[destination] {
|
|
||||||
switch {
|
|
||||||
case alert.Event == "":
|
|
||||||
return fmt.Errorf("waiting %s %w", destination, missing(i, "event"))
|
|
||||||
case alert.Time.IsZero():
|
|
||||||
return fmt.Errorf("waiting %s %w", destination, missing(i, "time"))
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
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())
|
|
||||||
}
|
|
||||||
File diff suppressed because it is too large
Load Diff
@@ -1,35 +0,0 @@
|
|||||||
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)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@@ -131,7 +131,6 @@ main() {
|
|||||||
|
|
||||||
if missing make; then pkg_install gnumake make make make; fi
|
if missing make; then pkg_install gnumake make make make; fi
|
||||||
if missing git; then pkg_install git git git git; 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
|
if missing gofmt; then pkg_install go golang go go; fi
|
||||||
|
|
||||||
ensure_node
|
ensure_node
|
||||||
|
|||||||
+3
-2
@@ -16,8 +16,9 @@ main() {
|
|||||||
"$SCRIPT_DIR/check"
|
"$SCRIPT_DIR/check"
|
||||||
# Own line: a failing command substitution inside an argument does
|
# Own line: a failing command substitution inside an argument does
|
||||||
# not trip `set -e`, so the inline form degrades silently to an
|
# not trip `set -e`, so the inline form degrades silently to an
|
||||||
# empty constant. The VERSION build argument takes precedence over
|
# empty constant. VERSION is computed here because .dockerignore
|
||||||
# the version a build stage derives from the .git in the context.
|
# excludes .git, so `git describe` in a build stage yields an empty
|
||||||
|
# version without failing.
|
||||||
version="$(git describe --tags --always --dirty 2>/dev/null || true)"
|
version="$(git describe --tags --always --dirty 2>/dev/null || true)"
|
||||||
[ -n "$version" ] || version="unknown"
|
[ -n "$version" ] || version="unknown"
|
||||||
docker build --no-cache \
|
docker build --no-cache \
|
||||||
|
|||||||
+3
-2
@@ -12,8 +12,9 @@ main() {
|
|||||||
cd "$ROOT"
|
cd "$ROOT"
|
||||||
# Own line: a failing command substitution inside an argument does
|
# Own line: a failing command substitution inside an argument does
|
||||||
# not trip `set -e`, so the inline form degrades silently to an
|
# not trip `set -e`, so the inline form degrades silently to an
|
||||||
# empty constant. The VERSION build argument takes precedence over
|
# empty constant. VERSION is computed here because .dockerignore
|
||||||
# the version a build stage derives from the .git in the context.
|
# excludes .git, so `git describe` in a build stage yields an empty
|
||||||
|
# version without failing.
|
||||||
version="$(git describe --tags --always --dirty 2>/dev/null || true)"
|
version="$(git describe --tags --always --dirty 2>/dev/null || true)"
|
||||||
[ -n "$version" ] || version="unknown"
|
[ -n "$version" ] || version="unknown"
|
||||||
docker build --no-cache \
|
docker build --no-cache \
|
||||||
|
|||||||
@@ -1,145 +0,0 @@
|
|||||||
#!/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 a probe for /.env bans another client, which its next
|
|
||||||
# request bans for good, 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 run with SWWAF_LOOKUP_SOURCE=off, so that no address is sent
|
|
||||||
# to GeoJS. 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 <what fails> <command>...: 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 <text>...: a line of the container's output holds every text,
|
|
||||||
# in any order.
|
|
||||||
logged() {
|
|
||||||
lines="$(docker logs "$CONTAINER" 2>&1)"
|
|
||||||
for text in "$@"; do
|
|
||||||
lines="$(printf '%s\n' "$lines" | grep -F "$text")" || return 1
|
|
||||||
done
|
|
||||||
}
|
|
||||||
|
|
||||||
# start_container: run the app's container, with the state files on the
|
|
||||||
# volume, a rate limit of one request a minute and no client looked up,
|
|
||||||
# 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 \
|
|
||||||
--env SWWAF_LOOKUP_SOURCE=off \
|
|
||||||
"$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 ]
|
|
||||||
}
|
|
||||||
|
|
||||||
# refused_from <client> <path>: a request for path from client, as
|
|
||||||
# X-Forwarded-For names it, gets 403. smallwebwaf believes the header
|
|
||||||
# from docker's gateway, a private address.
|
|
||||||
refused_from() {
|
|
||||||
code="$(curl --silent --output /dev/null --write-out '%{http_code}' \
|
|
||||||
--max-time 10 --header "X-Forwarded-For: $1" "http://$address$2")" || 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"
|
|
||||||
|
|
||||||
refused_from 203.0.113.9 /.env || fail "a probe for /.env was not refused"
|
|
||||||
wait_for "smallwebwaf logged no ban for the probe" \
|
|
||||||
logged '"action":"banned"' '"rule_ids":["env-file"]'
|
|
||||||
refused_from 203.0.113.9 / || fail "the client of the probe was let through"
|
|
||||||
wait_for "the client's next request did not make its ban permanent" \
|
|
||||||
logged '"ban_expires":"permanent"'
|
|
||||||
echo "example-app: a probe for /.env bans the client, its next request for good"
|
|
||||||
|
|
||||||
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 "$@"
|
|
||||||
+1
-13
@@ -1,9 +1,6 @@
|
|||||||
#!/bin/sh
|
#!/bin/sh
|
||||||
# script/run: build bin/smallwebwaf with script/build and run it, with
|
# script/run: build bin/smallwebwaf with script/build and run it, with
|
||||||
# the settings in the environment. Unless SWWAF_STATE_DIR is set, the
|
# the settings in the environment.
|
||||||
# state files go in bin/state, beside the binary, and unless
|
|
||||||
# SWWAF_RULES_DIR is set, the rule files are those of share/rules.d,
|
|
||||||
# which the image ships.
|
|
||||||
set -eu
|
set -eu
|
||||||
|
|
||||||
SCRIPT_DIR="$(cd "$(dirname "$0")" && pwd -P)"
|
SCRIPT_DIR="$(cd "$(dirname "$0")" && pwd -P)"
|
||||||
@@ -11,15 +8,6 @@ ROOT="$(cd "$SCRIPT_DIR/.." && pwd -P)"
|
|||||||
|
|
||||||
main() {
|
main() {
|
||||||
"$SCRIPT_DIR/build"
|
"$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
|
|
||||||
if [ -z "${SWWAF_RULES_DIR+set}" ]; then
|
|
||||||
SWWAF_RULES_DIR="$ROOT/share/rules.d"
|
|
||||||
export SWWAF_RULES_DIR
|
|
||||||
fi
|
|
||||||
exec "$ROOT/bin/smallwebwaf"
|
exec "$ROOT/bin/smallwebwaf"
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -1,15 +0,0 @@
|
|||||||
# 00-default.rules: probes no real visitor sends, anchored at the site root
|
|
||||||
|
|
||||||
# id target action regex
|
|
||||||
env-file path ban (?i)^/\.env(\.[a-z]+)?$
|
|
||||||
vcs-dir path ban (?i)^/\.(git|svn|hg|bzr)(/|$)
|
|
||||||
secrets-dir path ban (?i)^/\.(aws|ssh|docker|kube)/
|
|
||||||
secret-file path ban (?i)^/\.(htpasswd|htaccess|npmrc|netrc|pgpass|git-credentials|bash_history|DS_Store)$
|
|
||||||
editor-dir path ban (?i)^/\.(vscode|idea)/
|
|
||||||
backup-file path ban (?i)^/[^/]+\.(php(\.[a-z0-9]+|~)|sql(\.[a-z0-9]+)?)$
|
|
||||||
log-file path ban (?i)^/(debug|error|access)\.log$
|
|
||||||
compose-file path ban (?i)^/(docker-)?compose\.ya?ml$
|
|
||||||
php-shell path ban (?i)^/(shell|c99|r57|wso|alfa)\.php$
|
|
||||||
scanner-agent user_agent ban (?i)\b(sqlmap|nikto|nuclei|masscan|zgrab|wpscan)\b
|
|
||||||
path-traversal uri block (\.\./){2,}
|
|
||||||
empty-agent user_agent log ^$
|
|
||||||
@@ -1,16 +0,0 @@
|
|||||||
#!/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 "$@"
|
|
||||||
Reference in New Issue
Block a user